From 1c1d49d2913db6b091d309df43a39db7a26b4b99 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Wed, 30 Sep 2026 19:00:00 -0700 Subject: [PATCH 001/161] [None][perf] Cut host work between speculative decoding steps One-model speculative decoding spends host time between two decode-step graph replays on small kernel launches and host-to-device copies. - On SM 100, a CUDA graph decode step of an engine whose attention metadata is a plain TrtllmAttentionMetadata writes its per-step inputs (overlap gathers, positions, prompt and KV lengths, KV block offsets) with one CuTe DSL kernel launch (StepInputStage). Eager steps on SM 100 run the overlap gathers as one kernel. - On SM 100, the speculative sampler moves the forward's outputs into its slot stores with one CuTe DSL kernel (SlotScatter). - The sampler's host copies of its stores run on its D2H side stream, and one-model speculative sampling runs on the execution stream. - CUDA graphs of a speculative engine read input ids and positions from the engine's buffers instead of copying them in at each replay. - Copies whose values did not change are skipped (gather ids, slot tables, previous-batch indices), as are the Philox seed and offset uploads of all-greedy batches and the unused num_accepted_draft_tokens upload. Other devices, and calls the kernels do not cover, keep the torch path. Signed-off-by: Vasanth Sabavat --- .../_torch/attention/backends/trtllm.py | 60 +- .../spec_step_copies/__init__.py | 18 + .../cute_dsl_kernels/spec_step_copies/op.py | 433 +++++++++ .../spec_step_copies_kernel.py | 362 ++++++++ .../_torch/pyexecutor/cuda_graph_runner.py | 42 +- .../kv_cache/kv_cache_manager_v2.py | 32 + .../_torch/pyexecutor/model_engine.py | 219 +++-- tensorrt_llm/_torch/pyexecutor/py_executor.py | 13 +- tensorrt_llm/_torch/speculative/interface.py | 7 +- .../_torch/speculative/spec_sampler_base.py | 145 ++- tensorrt_llm/_utils.py | 28 + .../test_lists/test-db/l0_b200.yml | 4 + .../cute_dsl_kernels/test_spec_step_copies.py | 822 ++++++++++++++++++ .../test_copy_to_device_if_changed.py | 119 +++ .../test_cuda_graph_capture_replay.py | 79 ++ .../executor/test_pytorch_model_engine.py | 128 +++ tests/unittest/_torch/helpers.py | 13 +- .../hw_agnostic/test_rng_window_counter.py | 38 +- 18 files changed, 2400 insertions(+), 162 deletions(-) create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/spec_step_copies/__init__.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/spec_step_copies/op.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/spec_step_copies/spec_step_copies_kernel.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/test_spec_step_copies.py create mode 100644 tests/unittest/_torch/executor/test_copy_to_device_if_changed.py diff --git a/tensorrt_llm/_torch/attention/backends/trtllm.py b/tensorrt_llm/_torch/attention/backends/trtllm.py index 6f792d64dd5a..3b835cfc3ef0 100644 --- a/tensorrt_llm/_torch/attention/backends/trtllm.py +++ b/tensorrt_llm/_torch/attention/backends/trtllm.py @@ -25,6 +25,7 @@ if TYPE_CHECKING: from tensorrt_llm.mapping import Mapping + from ...cute_dsl_kernels.spec_step_copies.op import StepInputStage from ...model_config import ModelConfig from ...speculative.interface import SpecMetadata from ...speculative.spec_tree_manager import SpecTreeManager @@ -123,6 +124,10 @@ class TrtllmAttentionMetadata(AttentionMetadata): workspace: Optional[torch.Tensor] = None cuda_graph_workspace: Optional[torch.Tensor] = None workspace_reclaimable: bool = field(default=True, init=False) + # Set by the model engine around prepare() on a CUDA graph decode step: + # the per-step host-to-device copies are staged on it (a StepInputStage) + # and written by its commit, instead of being copied here. + h2d_stage: Optional["StepInputStage"] = field(default=None, init=False) # TrtllmAttention needs to know the beam width to access to the cache indirection buffer, # when beam search is enabled. @@ -873,6 +878,29 @@ def restore_after_draft_forward(self, saved_state: dict | None) -> None: """Restore backend state modified for draft-forward execution.""" return None + def _copy_block_offsets(self, manager, dst: torch.Tensor, + max_blocks: Optional[int]) -> None: + """``manager``'s block-offset copy for this batch, staged on + ``h2d_stage`` when the engine set one and the manager can stage it + (``stage_batch_block_offsets``).""" + stage_copy = getattr(manager, "stage_batch_block_offsets", None) + if self.h2d_stage is not None and stage_copy is not None and stage_copy( + self.h2d_stage, dst, self.request_ids, self.beam_width, + self.num_contexts, self.num_seqs): + return + manager.copy_batch_block_offsets(dst, + self.request_ids, + self.beam_width, + self.num_contexts, + self.num_seqs, + max_blocks=max_blocks) + + def _copy_step_input(self, dst: torch.Tensor, src: torch.Tensor) -> None: + """``dst.copy_(src)`` for a per-step input, staged on ``h2d_stage`` + when the engine set one.""" + if self.h2d_stage is None or not self.h2d_stage.copy(dst, src): + dst.copy_(maybe_pin_memory(src), non_blocking=True) + def prepare(self) -> None: super().prepare() # Recomputed on first use this iteration; see mla_prepare_scheduler_buffers. @@ -901,8 +929,8 @@ def prepare(self) -> None: device='cpu', ) self.prompt_lens_cpu[:self.num_seqs].copy_(prompt_lens) - self.prompt_lens_cuda[:self.num_seqs].copy_( - self.prompt_lens_cpu[:self.num_seqs], non_blocking=True) + self._copy_step_input(self.prompt_lens_cuda[:self.num_seqs], + self.prompt_lens_cpu[:self.num_seqs]) # number of tokens in the kv cache for each sequence in the batch cached_token_lens = torch.tensor( @@ -955,9 +983,8 @@ def prepare(self) -> None: # the sequence length including the cached tokens and the input tokens. self.kv_lens[:self.num_seqs].copy_( kv_lens + self.kv_cache_params.num_extra_kv_tokens) - self.kv_lens_cuda[:self.num_seqs].copy_(maybe_pin_memory( - kv_lens[:self.num_seqs]), - non_blocking=True) + self._copy_step_input(self.kv_lens_cuda[:self.num_seqs], + kv_lens[:self.num_seqs]) # total kv lens for context requests and generation requests, without extra tokens self.host_total_kv_lens[0] = kv_lens[:self.num_contexts].sum().item() self.host_total_kv_lens[1] = kv_lens[self.num_contexts:self. @@ -995,24 +1022,14 @@ def prepare(self) -> None: if not spec_active and self.kv_cache_manager.tokens_per_block: max_blocks = ceil_div(max_kv_len, self.kv_cache_manager.tokens_per_block) - self.kv_cache_manager.copy_batch_block_offsets( - self.kv_cache_block_offsets, - self.request_ids, - self.beam_width, - self.num_contexts, - self.num_seqs, - max_blocks=max_blocks) + self._copy_block_offsets(self.kv_cache_manager, + self.kv_cache_block_offsets, max_blocks) # Also prepare draft KV cache block offsets if draft_kv_cache_manager exists if self.draft_kv_cache_manager is not None: - # Use the wrapper method which works for both V1 and V2 - self.draft_kv_cache_manager.copy_batch_block_offsets( - self.draft_kv_cache_block_offsets, - self.request_ids, - self.beam_width, - self.num_contexts, - self.num_seqs, - max_blocks=max_blocks) + self._copy_block_offsets(self.draft_kv_cache_manager, + self.draft_kv_cache_block_offsets, + max_blocks) # Don't pass self.kv_lens as kv_lens here because it includes extra # tokens. Use the actual KV length (without extra tokens) for @@ -1319,6 +1336,9 @@ def prepare_context_mla_with_cached_kv(self, self.max_ctx_cached_token_len = 0 self.max_ctx_kv_len = 0 self.max_ctx_seq_len = 0 + # The indptrs below are read only by the context MLA kernels, and + # every batch with context requests rewrites them. + return torch.cumsum(cached_token_lens[:self.num_contexts], dim=0, dtype=torch.int64, diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/spec_step_copies/__init__.py b/tensorrt_llm/_torch/cute_dsl_kernels/spec_step_copies/__init__.py new file mode 100644 index 000000000000..459e4e88bd94 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/spec_step_copies/__init__.py @@ -0,0 +1,18 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""The one-model speculative decoding step's eager copy passes as one CuTe DSL kernel each. + +Import :mod:`.op` for the host side; nothing is imported eagerly here so that the CuTe DSL dependency stays optional. +""" diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/spec_step_copies/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/spec_step_copies/op.py new file mode 100644 index 000000000000..c0b2e56c3658 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/spec_step_copies/op.py @@ -0,0 +1,433 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Host side of the one-model speculative decoding step's copy kernels (``spec_step_copies_kernel``). + +``SlotScatter`` moves a step's per-row outputs into the sampler's slot stores (one launch for the sampler's four +``index_copy_``). ``StepInputStage`` writes a decode step's per-step inputs with one launch: the overlap scheduler's +gathers from those stores, host-to-device copies whose values the kernel reads from a pinned host record, and KV cache +block-offset copies. Each kernel takes every size and address as a launch argument, so it is compiled once per process, +on the first launch, which must happen outside CUDA-graph capture. The kernels are called through TVM-FFI: a launch +costs a few microseconds of host time. + +Callers use them only where ``is_supported()`` holds. A call whose arguments are outside what the kernel covers launches +and stages nothing and returns False; the caller then does that work itself. +""" + +from __future__ import annotations + +from typing import Any, ClassVar + +import torch + +from ...._utils import get_sm_version, prefer_pinned +from ...cute_dsl_utils import IS_CUTLASS_DSL_AVAILABLE + + +def is_supported() -> bool: + """Whether the kernels run on the current device: SM 100 with the CuTe DSL installed.""" + return IS_CUTLASS_DSL_AVAILABLE and torch.cuda.is_available() and get_sm_version() == 100 + + +def _is_i32(t: torch.Tensor) -> bool: + """A contiguous CUDA int32 tensor.""" + return t.is_cuda and t.dtype == torch.int32 and t.is_contiguous() + + +def _is_pinned_i32(t: torch.Tensor) -> bool: + """A contiguous pinned host int32 tensor.""" + return not t.is_cuda and t.dtype == torch.int32 and t.is_contiguous() and t.is_pinned() + + +def _compile(name: str, args: tuple): + """Compile kernel ``name`` for its scalar ``args`` with TVM-FFI; a call then takes ``args`` and the CUDA stream + handle to launch on.""" + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + f"spec_step_copies: the {name} kernel must run once outside CUDA-graph capture first " + "(it compiles on its first launch)." + ) + import cutlass.cute as cute + + from . import spec_step_copies_kernel as kernel + + return cute.compile( + getattr(kernel, name), *args, cute.runtime.make_fake_stream(), options="--enable-tvm-ffi" + ) + + +def _gather_shapes_ok( + store_next_new_tokens: torch.Tensor, + store_next_draft: torch.Tensor, + store_lens: torch.Tensor, + rows: int, + tokens_per_row: int, + draft_width: int, +) -> bool: + """The overlap gather's stores are [width, slots, 1], [slots, draft width] and [slots], wide enough for a row.""" + num_slots = store_lens.numel() + return ( + store_lens.dim() == 1 + and store_next_new_tokens.dim() == 3 + and store_next_new_tokens.shape[1:] == (num_slots, 1) + and store_next_draft.dim() == 2 + and store_next_draft.shape[0] == num_slots + and 0 < tokens_per_row <= store_next_new_tokens.shape[0] + and 0 <= draft_width < tokens_per_row + and draft_width <= store_next_draft.shape[1] + and rows >= 0 + ) + + +class SlotScatter: + """The speculative sampler's store update as one kernel. + + ``store[..., slots[r]] = outputs[row_begin + r]`` for every row ``r``, each row padded with zeros or cut to its + store's width. + """ + + _kernel: ClassVar[Any] = None + + def scatter( + self, + outputs: dict[str, torch.Tensor], + row_begin: int, + rows: int, + slots: torch.Tensor, + store_new_tokens: torch.Tensor, + store_next_new_tokens: torch.Tensor, + store_lens: torch.Tensor, + store_next_draft: torch.Tensor, + ) -> bool: + """Launch the store update on the current stream. + + Args: + outputs: The forward's ``new_tokens`` [rows_total, a], ``next_new_tokens`` [rows_total, b], + ``new_tokens_lens`` [rows_total] and ``next_draft_tokens`` [rows_total, c], int32. + row_begin: The first row of ``outputs`` to move. + rows: The number of rows to move. + slots: The distinct slot of each moved row, int32 [>= rows] on the device. + store_new_tokens: int32 [width, num_slots, 1]. + store_next_new_tokens: int32 [width, num_slots, 1]. + store_lens: int32 [num_slots]. + store_next_draft: int32 [num_slots, width]. + + Returns: + False, launching nothing, when an argument is outside what the kernel covers (a non-contiguous or + non-int32 tensor, a store of another layout, fewer output rows or slots than ``rows``); True otherwise. + """ + if rows == 0: + return True + new_tokens = outputs["new_tokens"] + next_new_tokens = outputs["next_new_tokens"] + lens = outputs["new_tokens_lens"] + next_draft = outputs["next_draft_tokens"] + num_slots = store_lens.numel() + tensors = ( + new_tokens, + next_new_tokens, + lens, + next_draft, + slots, + store_new_tokens, + store_next_new_tokens, + store_lens, + store_next_draft, + ) + if not ( + all(_is_i32(t) for t in tensors) + and new_tokens.dim() == next_new_tokens.dim() == next_draft.dim() == 2 + and lens.dim() == store_lens.dim() == 1 + and min(t.shape[0] for t in (new_tokens, next_new_tokens, lens, next_draft)) + >= row_begin + rows + and row_begin >= 0 + and slots.numel() >= rows + and store_new_tokens.dim() == store_next_new_tokens.dim() == 3 + and store_new_tokens.shape[1:] == (num_slots, 1) + and store_next_new_tokens.shape[1:] == (num_slots, 1) + and store_next_draft.dim() == 2 + and store_next_draft.shape[0] == num_slots + ): + return False + columns = max( + store_new_tokens.shape[0], store_next_new_tokens.shape[0], store_next_draft.shape[1], 1 + ) + args = ( + new_tokens.data_ptr(), + next_new_tokens.data_ptr(), + lens.data_ptr(), + next_draft.data_ptr(), + slots.data_ptr(), + store_new_tokens.data_ptr(), + store_next_new_tokens.data_ptr(), + store_lens.data_ptr(), + store_next_draft.data_ptr(), + rows, + row_begin, + columns, + num_slots, + store_new_tokens.shape[0], + new_tokens.shape[1], + store_next_new_tokens.shape[0], + next_new_tokens.shape[1], + store_next_draft.shape[1], + next_draft.shape[1], + ) + if SlotScatter._kernel is None: + SlotScatter._kernel = _compile("scatter", args) + SlotScatter._kernel(*args, torch.cuda.current_stream().cuda_stream) + return True + + +class StepInputStage: + """One decode step's per-step device inputs, written by one ``stage_kernel`` launch. + + A step calls ``begin()``, then stages the overlap gathers (``gather``), host-to-device copies whose values go into a + pinned host record that the kernel reads in place (``copy``, up to ``max_copies``) and KV cache block-offset copies + (``block_copy``, up to ``max_block_copies``), then ``commit()`` launches one kernel on the current stream that + performs all of them. The record is reused every step: the first ``copy`` after a commit that staged copies waits + for that commit's kernel (in the overlap loop it has run by then: the host is at most one step ahead). + + Args: + capacity: The record's size in int32 values, the most a step's copies stage in total. + """ + + _kernel: ClassVar[Any] = None + + def __init__(self, capacity: int) -> None: + from . import spec_step_copies_kernel as kernel + + self.max_copies = kernel.STAGED_COPIES + self.max_block_copies = kernel.BLOCK_COPIES + self.record = torch.empty((capacity,), dtype=torch.int32, pin_memory=prefer_pinned()) + self._copies: list[tuple[int, int, int]] = [] + self._block_copies: list[tuple[int, ...]] = [] + self._gather: tuple[int, ...] | None = None + self._used = 0 + self._done: torch.cuda.Event | None = None + self._record_free = True + + def begin(self) -> None: + """Start a step: drop anything staged by a step that did not commit (its preparation raised).""" + self._copies = [] + self._block_copies = [] + self._gather = None + self._used = 0 + + def gather( + self, + store_next_new_tokens: torch.Tensor, + store_next_draft: torch.Tensor, + store_lens: torch.Tensor, + slots: torch.Tensor, + pos_indices: torch.Tensor, + rows: int, + tokens_per_row: int, + draft_width: int, + input_ids: torch.Tensor, + input_begin: int, + draft_tokens: torch.Tensor, + draft_begin: int, + pos_offsets: torch.Tensor, + pos_begin: int, + kv_offsets: torch.Tensor, + kv_begin: int, + ) -> bool: + """Stage the overlap scheduler's gathers of the rows whose request ran in the previous step. + + For ``r < rows`` with ``s = slots[r]`` and ``j < tokens_per_row``: + ``input_ids[input_begin + r * tokens_per_row + j] = store_next_new_tokens[j, s, 0]``, + ``pos_offsets[pos_begin + r * tokens_per_row + j] = store_lens[pos_indices[r * tokens_per_row + j]]``, + ``draft_tokens[draft_begin + r * draft_width + j] = store_next_draft[s, j]`` for ``j < draft_width`` and + ``kv_offsets[kv_begin + r] = store_lens[s] - tokens_per_row``. Every tensor is int32 on the device; the + stores are [width, slots, 1], [slots, width] and [slots]. + + Returns: + False, staging nothing, when an argument is outside what the kernel covers or a gather is already + staged; True otherwise. + """ + tensors = ( + store_next_new_tokens, + store_next_draft, + store_lens, + slots, + pos_indices, + input_ids, + draft_tokens, + pos_offsets, + kv_offsets, + ) + if not ( + self._gather is None + and all(_is_i32(t) for t in tensors) + and _gather_shapes_ok( + store_next_new_tokens, + store_next_draft, + store_lens, + rows, + tokens_per_row, + draft_width, + ) + and slots.numel() >= rows + and pos_indices.numel() >= rows * tokens_per_row + and input_ids.numel() >= input_begin + rows * tokens_per_row + and draft_tokens.numel() >= draft_begin + rows * draft_width + and pos_offsets.numel() >= pos_begin + rows * tokens_per_row + and kv_offsets.numel() >= kv_begin + rows + and min(input_begin, draft_begin, pos_begin, kv_begin) >= 0 + ): + return False + if rows == 0: + return True + self._gather = ( + store_next_new_tokens.data_ptr(), + store_next_draft.data_ptr(), + store_lens.data_ptr(), + slots.data_ptr(), + pos_indices.data_ptr(), + input_ids.data_ptr(), + draft_tokens.data_ptr(), + pos_offsets.data_ptr(), + kv_offsets.data_ptr(), + rows, + tokens_per_row, + draft_width, + store_lens.numel(), + store_next_draft.shape[1], + input_begin, + draft_begin, + pos_begin, + kv_begin, + ) + return True + + def copy(self, dst: torch.Tensor, values: torch.Tensor) -> bool: + """Stage ``dst[:n] = values`` for the ``n`` int32 host ``values``; ``dst`` is a 1-D int32 device tensor. + + Returns: + False, staging nothing, when ``dst`` is not a contiguous 1-D CUDA int32 tensor of at least ``n`` + values, the step already staged ``max_copies`` copies or the record cannot hold the values; True + otherwise. + """ + host = values.reshape(-1) + n = host.numel() + if not ( + _is_i32(dst) + and dst.dim() == 1 + and n <= dst.numel() + and host.dtype == torch.int32 + and not host.is_cuda + and len(self._copies) < self.max_copies + and self._used + n <= self.record.numel() + and self.record.is_pinned() + ): + return False + if n == 0: + return True + if not self._record_free: + self._done.synchronize() + self._record_free = True + self.record[self._used : self._used + n].copy_(host) + self._copies.append((dst.data_ptr(), self._used, n)) + self._used += n + return True + + def block_copy( + self, + offsets: torch.Tensor, + table: torch.Tensor, + copy_index: torch.Tensor, + index_scales: torch.Tensor, + kv_offset: torch.Tensor, + ) -> bool: + """Stage a KV cache manager's block-offset copy (``copy_batch_block_offsets_to_device``). + + Args: + offsets: The device block offsets, int32 [pools, seqs_cap, 2, blocks]. + table: The pinned host page table, int32 [pools, table_seqs, 2, blocks]. + copy_index: The table row of each of the step's sequences, pinned host int32 [seqs]. + index_scales: The per-pool page index scale, pinned host int32 [pools]. + kv_offset: The per-pool V offset, pinned host int32 [pools]. + + For every pool and sequence the K row is ``index_scale * page`` and the V row ``index_scale * page + + kv_offset`` over the table row ``copy_index[sequence]``; a bad page (-1) gives 0. + + Returns: + False, staging nothing, when an argument is outside what the kernel covers or the step already staged + ``max_block_copies`` block copies; True otherwise. + """ + if not ( + _is_i32(offsets) + and offsets.dim() == 4 + and table.dim() == 4 + and all(_is_pinned_i32(t) for t in (table, copy_index, index_scales, kv_offset)) + and len(self._block_copies) < self.max_block_copies + ): + return False + pools, table_seqs, kv, blocks = table.shape + seqs = copy_index.numel() + if not ( + kv == 2 + and offsets.shape[0] >= pools + and offsets.shape[1] >= seqs + and offsets.shape[2] == 2 + and offsets.shape[3] == blocks + and index_scales.numel() >= pools + and kv_offset.numel() >= pools + ): + return False + if pools * seqs * blocks == 0: + return True + self._block_copies.append( + ( + table.data_ptr(), + offsets.data_ptr(), + copy_index.data_ptr(), + index_scales.data_ptr(), + kv_offset.data_ptr(), + pools, + table_seqs, + offsets.shape[1], + blocks, + seqs, + ) + ) + return True + + def commit(self) -> None: + """Launch the staged gathers and copies as one kernel on the current stream, then start the next staging.""" + if self._gather is None and not self._copies and not self._block_copies: + return + gather = self._gather or (0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 1, 1, 0, 0, 0, 0) + copies = self._copies + [(0, 0, 0)] * (self.max_copies - len(self._copies)) + blocks = self._block_copies + [(0, 0, 0, 0, 0, 0, 0, 0, 0, 0)] * ( + self.max_block_copies - len(self._block_copies) + ) + args = ( + *gather, + self.record.data_ptr(), + *(c[0] for c in copies), + *(c[1] for c in copies), + *(c[2] for c in copies), + *(v for b in blocks for v in b), + ) + if StepInputStage._kernel is None: + StepInputStage._kernel = _compile("stage", args) + StepInputStage._kernel(*args, torch.cuda.current_stream().cuda_stream) + if self._copies: + if self._done is None: + self._done = torch.cuda.Event() + self._done.record() + self._record_free = False + self.begin() diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/spec_step_copies/spec_step_copies_kernel.py b/tensorrt_llm/_torch/cute_dsl_kernels/spec_step_copies/spec_step_copies_kernel.py new file mode 100644 index 000000000000..9780847687b4 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/spec_step_copies/spec_step_copies_kernel.py @@ -0,0 +1,362 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""The one-model speculative decoding step's eager copy passes, one kernel each. + +* ``scatter_kernel``: the sampler moves the forward's per-row outputs into its slot-indexed stores (``index_copy_`` + of new tokens, next new tokens, accepted lengths and next draft tokens, each padded or cut to its store width). +* ``stage_kernel``: one launch writes a decode step's per-step inputs: the overlap scheduler's gathers from those + stores for the rows whose request ran in the previous step (``index_select`` into input ids and draft tokens by + slot, into the position offsets by the per-token index list, and into the KV-length offsets by slot, minus the + tokens per step), up to four host-to-device copies whose values it reads in place from a pinned host record, and up + to two KV cache block-offset copies. + +Every buffer is int32 and passed as a raw address, so a call passes integers only. Row and element offsets are +arguments. One thread per element. +""" + +from __future__ import annotations + +import cuda.bindings.driver as cuda_driver +import cutlass +import cutlass.cute as cute + +THREADS = 128 + + +def _i32_at(address, index): + """A one-element int32 tensor at ``address + 4 * index`` (trace-time helper; global or mapped host memory).""" + ptr = cute.make_ptr( + cutlass.Int32, + address + cutlass.Int64(index) * cutlass.Int64(4), + cute.AddressSpace.gmem, + assumed_align=4, + ) + return cute.make_tensor(ptr, cute.make_layout((1,))) + + +@cute.kernel +def scatter_kernel( + out_new_tokens: cutlass.Int64, # int32 [rows_total, out_new_width], the forward's accepted tokens + out_next_new_tokens: cutlass.Int64, # int32 [rows_total, out_next_width] + out_lens: cutlass.Int64, # int32 [rows_total] + out_next_draft: cutlass.Int64, # int32 [rows_total, out_draft_width] + slots: cutlass.Int64, # int32 [rows] + store_new_tokens: cutlass.Int64, # int32 [new_width, num_slots] (token-major) + store_next_new_tokens: cutlass.Int64, # int32 [next_width, num_slots] + store_lens: cutlass.Int64, # int32 [num_slots] + store_next_draft: cutlass.Int64, # int32 [num_slots, draft_width] + rows: cutlass.Int32, + row_begin: cutlass.Int32, + columns: cutlass.Int32, + num_slots: cutlass.Int32, + new_width: cutlass.Int32, + out_new_width: cutlass.Int32, + next_width: cutlass.Int32, + out_next_width: cutlass.Int32, + draft_width: cutlass.Int32, + out_draft_width: cutlass.Int32, +): + tidx, _, _ = cute.arch.thread_idx() + bidx, _, _ = cute.arch.block_idx() + t = bidx * cutlass.Int32(THREADS) + tidx + if t < rows * columns: + row = t // columns + col = t - row * columns + src_row = row_begin + row + slot = _i32_at(slots, row)[0] + zero = cutlass.Int32(0) + if col < new_width: + value = zero + if col < out_new_width: + value = _i32_at(out_new_tokens, src_row * out_new_width + col)[0] + _i32_at(store_new_tokens, col * num_slots + slot)[0] = value + if col < next_width: + value = zero + if col < out_next_width: + value = _i32_at(out_next_new_tokens, src_row * out_next_width + col)[0] + _i32_at(store_next_new_tokens, col * num_slots + slot)[0] = value + if col < draft_width: + value = zero + if col < out_draft_width: + value = _i32_at(out_next_draft, src_row * out_draft_width + col)[0] + _i32_at(store_next_draft, slot * draft_width + col)[0] = value + if col == zero: + _i32_at(store_lens, slot)[0] = _i32_at(out_lens, src_row)[0] + + +@cute.jit +def scatter( + out_new_tokens: cutlass.Int64, + out_next_new_tokens: cutlass.Int64, + out_lens: cutlass.Int64, + out_next_draft: cutlass.Int64, + slots: cutlass.Int64, + store_new_tokens: cutlass.Int64, + store_next_new_tokens: cutlass.Int64, + store_lens: cutlass.Int64, + store_next_draft: cutlass.Int64, + rows: cutlass.Int32, + row_begin: cutlass.Int32, + columns: cutlass.Int32, + num_slots: cutlass.Int32, + new_width: cutlass.Int32, + out_new_width: cutlass.Int32, + next_width: cutlass.Int32, + out_next_width: cutlass.Int32, + draft_width: cutlass.Int32, + out_draft_width: cutlass.Int32, + stream: cuda_driver.CUstream, +) -> None: + scatter_kernel( + out_new_tokens, out_next_new_tokens, out_lens, out_next_draft, slots, store_new_tokens, + store_next_new_tokens, store_lens, store_next_draft, rows, row_begin, columns, num_slots, new_width, + out_new_width, next_width, out_next_width, draft_width, out_draft_width, + ).launch( + grid=[(rows * columns + THREADS - 1) // THREADS, 1, 1], + block=[THREADS, 1, 1], + stream=stream, + ) # fmt: skip + + +STAGED_COPIES = 4 +BLOCK_COPIES = 2 + + +@cute.kernel +def stage_kernel( + store_next_new_tokens: cutlass.Int64, # int32 [next_width, num_slots] (token-major) + store_next_draft: cutlass.Int64, # int32 [num_slots, draft_stride] + store_lens: cutlass.Int64, # int32 [num_slots] + slots: cutlass.Int64, # int32 [rows] + pos_indices: cutlass.Int64, # int32 [rows * tokens_per_row] + input_ids: cutlass.Int64, + draft_tokens: cutlass.Int64, + pos_offsets: cutlass.Int64, + kv_offsets: cutlass.Int64, + rows: cutlass.Int32, + tokens_per_row: cutlass.Int32, + draft_width: cutlass.Int32, + num_slots: cutlass.Int32, + draft_stride: cutlass.Int32, + input_begin: cutlass.Int32, + draft_begin: cutlass.Int32, + pos_begin: cutlass.Int32, + kv_begin: cutlass.Int32, + record: cutlass.Int64, # int32 values in pinned host memory, read in place + dst0: cutlass.Int64, + dst1: cutlass.Int64, + dst2: cutlass.Int64, + dst3: cutlass.Int64, + off0: cutlass.Int32, + off1: cutlass.Int32, + off2: cutlass.Int32, + off3: cutlass.Int32, + n0: cutlass.Int32, + n1: cutlass.Int32, + n2: cutlass.Int32, + n3: cutlass.Int32, + table0: cutlass.Int64, # block copy 0: int32 [pools, table_seqs, 2, blocks] in pinned host memory + offsets0: cutlass.Int64, # int32 [pools, offsets_seqs, 2, blocks] (device) + copy_index0: cutlass.Int64, # int32 [seqs], pinned host + index_scales0: cutlass.Int64, # int32 [pools], pinned host + kv_offset0: cutlass.Int64, # int32 [pools], pinned host + pools0: cutlass.Int32, + table_seqs0: cutlass.Int32, + offsets_seqs0: cutlass.Int32, + blocks0: cutlass.Int32, + seqs0: cutlass.Int32, + table1: cutlass.Int64, + offsets1: cutlass.Int64, + copy_index1: cutlass.Int64, + index_scales1: cutlass.Int64, + kv_offset1: cutlass.Int64, + pools1: cutlass.Int32, + table_seqs1: cutlass.Int32, + offsets_seqs1: cutlass.Int32, + blocks1: cutlass.Int32, + seqs1: cutlass.Int32, +): + """The overlap gathers for threads [0, rows * tokens_per_row) (for row r, slot s = slots[r] and token j: + ``input_ids[input_begin + r * tokens_per_row + j] = store_next_new_tokens[j, s]``, ``pos_offsets[pos_begin + r * + tokens_per_row + j] = store_lens[pos_indices[r * tokens_per_row + j]]``, ``draft_tokens[draft_begin + r * + draft_width + j] = store_next_draft[s, j]`` for j < draft_width, ``kv_offsets[kv_begin + r] = store_lens[s] - + tokens_per_row``); then up to four staged copies ``dst_i[j] = record[off_i + j]`` for j < n_i; then up to two KV + block-offset copies, each the + ``copyBatchBlockOffsetsToDeviceKernel`` of kvCacheManagerV2Utils.cu (the K and V rows of every (pool, sequence) + from the pinned host table's row ``copy_index[sequence]``: ``index_scale * page``, and ``+ kv_offset`` for V; a + bad page (-1) gives 0). One thread per int32 (per K / V pair for the block copies).""" + tidx, _, _ = cute.arch.thread_idx() + bidx, _, _ = cute.arch.block_idx() + t = bidx * cutlass.Int32(THREADS) + tidx + gathered = rows * tokens_per_row + if t < gathered: + row = t // tokens_per_row + col = t - row * tokens_per_row + slot = _i32_at(slots, row)[0] + _i32_at(input_ids, input_begin + t)[0] = _i32_at( + store_next_new_tokens, col * num_slots + slot + )[0] + _i32_at(pos_offsets, pos_begin + t)[0] = _i32_at(store_lens, _i32_at(pos_indices, t)[0])[0] + if col < draft_width: + _i32_at(draft_tokens, draft_begin + row * draft_width + col)[0] = _i32_at( + store_next_draft, slot * draft_stride + col + )[0] + if col == cutlass.Int32(0): + _i32_at(kv_offsets, kv_begin + row)[0] = _i32_at(store_lens, slot)[0] - tokens_per_row + else: + u = t - gathered + dst = cutlass.Int64(0) + src = cutlass.Int32(0) + j = cutlass.Int32(0) + active = cutlass.Int32(0) + if u < n0: + dst = dst0 + src = off0 + u + j = u + active = cutlass.Int32(1) + elif u < n0 + n1: + dst = dst1 + src = off1 + u - n0 + j = u - n0 + active = cutlass.Int32(1) + elif u < n0 + n1 + n2: + dst = dst2 + src = off2 + u - n0 - n1 + j = u - n0 - n1 + active = cutlass.Int32(1) + elif u < n0 + n1 + n2 + n3: + dst = dst3 + src = off3 + u - n0 - n1 - n2 + j = u - n0 - n1 - n2 + active = cutlass.Int32(1) + if active == cutlass.Int32(1): + _i32_at(dst, j)[0] = _i32_at(record, src)[0] + v = u - n0 - n1 - n2 - n3 + per0 = pools0 * seqs0 * blocks0 + per1 = pools1 * seqs1 * blocks1 + if v >= cutlass.Int32(0): + if v < per0 + per1: + table = table0 + offsets = offsets0 + copy_index = copy_index0 + index_scales = index_scales0 + kv_offset = kv_offset0 + table_seqs = table_seqs0 + offsets_seqs = offsets_seqs0 + blocks = blocks0 + seqs = seqs0 + w = v + if v >= per0: + table = table1 + offsets = offsets1 + copy_index = copy_index1 + index_scales = index_scales1 + kv_offset = kv_offset1 + table_seqs = table_seqs1 + offsets_seqs = offsets_seqs1 + blocks = blocks1 + seqs = seqs1 + w = v - per0 + pool = w // (seqs * blocks) + rest = w - pool * seqs * blocks + seq = rest // blocks + block = rest - seq * blocks + row = pool * table_seqs + _i32_at(copy_index, seq)[0] + page = _i32_at(table, row * cutlass.Int32(2) * blocks + block)[0] + key = cutlass.Int32(0) + value = cutlass.Int32(0) + if page != cutlass.Int32(-1): + key = _i32_at(index_scales, pool)[0] * page + value = key + _i32_at(kv_offset, pool)[0] + out = (pool * offsets_seqs + seq) * cutlass.Int32(2) * blocks + block + _i32_at(offsets, out)[0] = key + _i32_at(offsets, out + blocks)[0] = value + + +@cute.jit +def stage( + store_next_new_tokens: cutlass.Int64, + store_next_draft: cutlass.Int64, + store_lens: cutlass.Int64, + slots: cutlass.Int64, + pos_indices: cutlass.Int64, + input_ids: cutlass.Int64, + draft_tokens: cutlass.Int64, + pos_offsets: cutlass.Int64, + kv_offsets: cutlass.Int64, + rows: cutlass.Int32, + tokens_per_row: cutlass.Int32, + draft_width: cutlass.Int32, + num_slots: cutlass.Int32, + draft_stride: cutlass.Int32, + input_begin: cutlass.Int32, + draft_begin: cutlass.Int32, + pos_begin: cutlass.Int32, + kv_begin: cutlass.Int32, + record: cutlass.Int64, + dst0: cutlass.Int64, + dst1: cutlass.Int64, + dst2: cutlass.Int64, + dst3: cutlass.Int64, + off0: cutlass.Int32, + off1: cutlass.Int32, + off2: cutlass.Int32, + off3: cutlass.Int32, + n0: cutlass.Int32, + n1: cutlass.Int32, + n2: cutlass.Int32, + n3: cutlass.Int32, + table0: cutlass.Int64, + offsets0: cutlass.Int64, + copy_index0: cutlass.Int64, + index_scales0: cutlass.Int64, + kv_offset0: cutlass.Int64, + pools0: cutlass.Int32, + table_seqs0: cutlass.Int32, + offsets_seqs0: cutlass.Int32, + blocks0: cutlass.Int32, + seqs0: cutlass.Int32, + table1: cutlass.Int64, + offsets1: cutlass.Int64, + copy_index1: cutlass.Int64, + index_scales1: cutlass.Int64, + kv_offset1: cutlass.Int64, + pools1: cutlass.Int32, + table_seqs1: cutlass.Int32, + offsets_seqs1: cutlass.Int32, + blocks1: cutlass.Int32, + seqs1: cutlass.Int32, + stream: cuda_driver.CUstream, +) -> None: + total = ( + rows * tokens_per_row + + n0 + + n1 + + n2 + + n3 + + pools0 * seqs0 * blocks0 + + pools1 * seqs1 * blocks1 + ) + stage_kernel( + store_next_new_tokens, store_next_draft, store_lens, slots, pos_indices, input_ids, draft_tokens, pos_offsets, + kv_offsets, rows, tokens_per_row, draft_width, num_slots, draft_stride, input_begin, draft_begin, pos_begin, + kv_begin, record, dst0, dst1, dst2, dst3, off0, off1, off2, off3, n0, n1, n2, n3, table0, offsets0, + copy_index0, index_scales0, kv_offset0, pools0, table_seqs0, offsets_seqs0, blocks0, seqs0, table1, offsets1, + copy_index1, index_scales1, kv_offset1, pools1, table_seqs1, offsets_seqs1, blocks1, seqs1, + ).launch( + grid=[(total + THREADS - 1) // THREADS, 1, 1], + block=[THREADS, 1, 1], + stream=stream, + ) # fmt: skip diff --git a/tensorrt_llm/_torch/pyexecutor/cuda_graph_runner.py b/tensorrt_llm/_torch/pyexecutor/cuda_graph_runner.py index 4614477976e0..dc78b8a9f967 100644 --- a/tensorrt_llm/_torch/pyexecutor/cuda_graph_runner.py +++ b/tensorrt_llm/_torch/pyexecutor/cuda_graph_runner.py @@ -136,6 +136,13 @@ class CUDAGraphRunnerConfig: sparse_attention_config: Optional[BaseSparseAttentionConfig] = None enable_encoder_decoder_mixed_cuda_graph: bool = False enable_in_graph_sampling: bool = False + static_input_ids: Optional[torch.Tensor] = None + """The engine's own int32 input-id buffer. When given (with + ``static_position_ids``), the graphs' static inputs are views of the + engine's buffers: every write the engine makes is already the graphs' + input, and replay copies nothing. For engines that rewrite every token's + input each step.""" + static_position_ids: Optional[torch.Tensor] = None class CUDAGraphRunner: @@ -203,6 +210,25 @@ def _create_shared_static_tensors(self): max_total_tokens = self.config.max_num_tokens max_total_tokens = min(max_total_tokens, self.config.max_num_tokens) + if self.config.static_input_ids is not None: + input_ids = self.config.static_input_ids + position_ids = self.config.static_position_ids + if (self.config.use_mrope or position_ids is None + or input_ids.dtype != torch.int32 + or position_ids.dtype != torch.int32 or input_ids.dim() != 1 + or position_ids.dim() != 1 or not input_ids.is_contiguous() + or not position_ids.is_contiguous() + or min(input_ids.numel(), + position_ids.numel()) < max_total_tokens): + raise ValueError( + "CUDA graph static inputs from the engine's buffers need " + "two contiguous 1-D int32 buffers of at least " + f"{max_total_tokens} tokens and no MRoPE.") + self.shared_static_tensors = { + "input_ids": input_ids[:max_total_tokens], + "position_ids": position_ids[:max_total_tokens].unsqueeze(0), + } + return self.shared_static_tensors = { "input_ids": torch.ones((max_total_tokens, ), device="cuda", dtype=torch.int32), @@ -732,10 +758,12 @@ def capture(self, # Do not keep the eager result live from this runner across graph # setup/capture; release its reference before entering. output = None + # A captured forward does not run, so there is no in-place input + # change to undo (postprocess_fn is for the warmup forwards above). + # Undoing one here would shift the static inputs, which are the + # engine's own buffers when static_input_ids is set. with torch.cuda.graph(graph, pool=self.memory_pool): output = forward_fn(capture_inputs) - if postprocess_fn is not None: - postprocess_fn(capture_inputs) _restore_spec_decode_capture_state(attn_metadata, saved_kv_lens_cuda) @@ -765,7 +793,11 @@ def replay(self, key: KeyType, f"replay() got {seqlen} tokens for key {key}, but the graph " f"was captured for {expected_num_tokens} tokens. A shorter " "input_ids leaves the tail of the static input buffer stale.") - static_tensors["input_ids"][:seqlen].copy_(input_ids) + # With the engine's buffers as the static inputs (static_input_ids) + # the source is the static view itself, and there is nothing to copy. + static_input_ids = static_tensors["input_ids"][:seqlen] + if input_ids.data_ptr() != static_input_ids.data_ptr(): + static_input_ids.copy_(input_ids) position_ids = current_inputs["position_ids"] if self.config.use_mrope: @@ -812,7 +844,9 @@ def replay(self, key: KeyType, f"for key {key}, but expected {expected_position_ids_shape}. " "torch.Tensor.copy_() silently broadcasts mismatched shapes, " "which would corrupt the static input buffer.") - static_tensors["position_ids"][:, :seqlen].copy_(position_ids) + static_position_ids = static_tensors["position_ids"][:, :seqlen] + if position_ids.data_ptr() != static_position_ids.data_ptr(): + static_position_ids.copy_(position_ids) num_encoder_tokens = key.num_encoder_tokens if num_encoder_tokens: diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache/kv_cache_manager_v2.py b/tensorrt_llm/_torch/pyexecutor/kv_cache/kv_cache_manager_v2.py index ed0cc2a89afe..cb880fe0f69b 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache/kv_cache_manager_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache/kv_cache_manager_v2.py @@ -5762,6 +5762,38 @@ def update_resources( ) self._allocated_draft_lens.pop(req.py_request_id, None) + def stage_batch_block_offsets( + self, + stage, + dst_tensor: torch.Tensor, + request_ids: List[int], + beam_width: int, + num_contexts: int, + num_seqs: int, + ) -> bool: + """Stage ``copy_batch_block_offsets`` on a StepInputStage instead of launching its kernel. + + Returns False, staging nothing, where that copy is not this class's single-table copy on the current + stream: the per-layer page-table layout, a subclass that overrides ``copy_batch_block_offsets``, a manager + whose stream is not the current one, or a copy the stage does not cover. The caller then copies as usual. + """ + if ( + self._use_per_layer_page_tables + or type(self).copy_batch_block_offsets is not KVCacheManagerV2.copy_batch_block_offsets + or self._stream != torch.cuda.current_stream() + ): + return False + assert beam_width == 1, "beam_width must be 1 for KVCacheManagerV2" + copy_idx = self.index_mapper.get_copy_index(request_ids, num_contexts, beam_width) + assert copy_idx.shape[0] == num_seqs + return stage.block_copy( + dst_tensor, + self.host_kv_cache_block_offsets, + copy_idx, + self.index_scales, + self.kv_offset, + ) + def copy_batch_block_offsets( self, dst_tensor: torch.Tensor, diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 2e93692a3947..b67f11f8789a 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -19,8 +19,9 @@ from tensorrt_llm._torch.peft.lora.config import LoraConfig from tensorrt_llm._torch.peft.lora.manager import LoraModelConfig from tensorrt_llm._torch.pyexecutor.warmup_timer import _WarmupTimer -from tensorrt_llm._utils import (global_mpi_rank, maybe_pin_memory, nvtx_range, - prefer_pinned, release_gc) +from tensorrt_llm._utils import (copy_to_device_if_changed, global_mpi_rank, + maybe_pin_memory, nvtx_range, prefer_pinned, + release_gc) from tensorrt_llm.bindings.internal import \ batch_manager as batch_manager_bindings from tensorrt_llm.inputs.multimodal import (MultimodalParams, @@ -48,6 +49,7 @@ from ..autotuner import AutoTuner, autotune from ..compilation.backend import Backend from ..compilation.utils import capture_piecewise_cuda_graph +from ..cute_dsl_kernels.spec_step_copies import op as spec_step_copies from ..distributed import Distributed from ..distributed.communicator import init_pp_comm from ..memory_buffer_utils import clear_memory_buffers, with_shared_pool @@ -756,8 +758,6 @@ def __init__( self.gather_ids_cuda = torch.empty((self.max_num_tokens, ), dtype=torch.int, device='cuda') - self.num_accepted_draft_tokens_cuda = torch.empty( - (self.batch_size, ), dtype=torch.int, device='cuda') self.previous_pos_indices_cuda = torch.empty( (self.max_num_tokens, ), dtype=torch.int, device='cuda') self.previous_pos_id_offsets_cuda = torch.zeros( @@ -811,6 +811,11 @@ def __init__( dtype=torch.int, device='cuda') self._encoder_decoder_staged_request_ids: Optional[List[int]] = None + # Host copies of what previous_batch_indices_cuda and + # previous_pos_indices_cuda last received: a decode step whose batch is + # unchanged skips both copies. + self._staged_previous_batch_indices: Optional[List[int]] = None + self._staged_previous_pos_indices: Optional[List[int]] = None if self._fallback_to_engine: self.input_ids_cuda = torch.empty((self.max_num_tokens, ), dtype=torch.int, @@ -876,6 +881,12 @@ def __init__( self.kv_cache_dtype_byte_size = self.get_kv_cache_dtype_byte_size() self._prepare_inputs_event: Optional[torch.cuda.Event] = None + # Where the step-copy kernels run: writes a speculative decode step's + # per-step inputs with one launch. Positions (up to max_num_tokens) and + # two per-sequence lengths fit its record. + self._step_input_stage: Optional[spec_step_copies.StepInputStage] = ( + spec_step_copies.StepInputStage(3 * self.max_num_tokens) if + self.is_spec_decode and spec_step_copies.is_supported() else None) # Cache for enc-dec cross-attention stable generation steps. # Populated on the first CUDA-graph generation step; cleared whenever @@ -918,6 +929,13 @@ def _initialize_cuda_graph_runner(self) -> Optional[CUDAGraphRunner]: enable_encoder_decoder_mixed_cuda_graph=( enable_encoder_decoder_mixed_cuda_graph), enable_in_graph_sampling=self.enable_in_graph_sampling, + # A speculative engine rewrites every token's input id and position + # each step (its prepare never takes the steady fast path), so the + # graphs read those straight from the engine's buffers. + static_input_ids=(self.input_ids_cuda if self.is_spec_decode + and not self.use_mrope else None), + static_position_ids=(self.position_ids_cuda if self.is_spec_decode + and not self.use_mrope else None), ) return CUDAGraphRunner(config) @@ -4364,6 +4382,7 @@ def _prepare_encoder_decoder_inputs_fast( [:num_previous_batch_requests], non_blocking=True) self._encoder_decoder_staged_request_ids = staged_request_ids + self._staged_previous_batch_indices = None generation_begin = num_context_tokens generation_end = generation_begin + num_previous_batch_requests torch.index_select( @@ -4736,7 +4755,6 @@ def _prepare_tp_inputs( # permanently reads back a zero delta. mrope_dummy_seq_slot = get_mrope_dummy_seq_slot(self.max_num_tokens, self.mapping.pp_size) - num_accepted_draft_tokens = [] # per request is_enc_dec = self._is_encoder_decoder_model() cross_encoder_hidden_states: List[torch.Tensor] = [] cross_encoder_seq_lens: List[int] = [ @@ -4806,7 +4824,6 @@ def append_cross_attention_state(request: LlmRequest, gather_ids.append(len(input_ids) - 1) sequence_lengths.append(len(prompt_tokens)) - num_accepted_draft_tokens.append(len(prompt_tokens) - 1) prompt_lengths.append(len(prompt_tokens)) past_seen_token_num = begin_compute num_cached_tokens_per_seq.append(past_seen_token_num - @@ -5064,7 +5081,6 @@ def _helix_pack_extend(request, group: int) -> int: prompt_lengths.append(request.py_prompt_len) sequence_lengths.append(1 + num_draft_tokens) - num_accepted_draft_tokens.append(num_draft_tokens) gather_ids.extend( list( range(len(position_ids), @@ -5099,8 +5115,6 @@ def _helix_pack_extend(request, group: int) -> int: request.py_batch_idx = request.py_seq_slot sequence_lengths.append(runtime_tokens_per_gen_step) - num_accepted_draft_tokens.append( - request.py_num_accepted_draft_tokens) past_seen_token_num = request.max_beam_num_tokens - 1 draft_lens.append(runtime_draft_token_buffer_width) @@ -5160,8 +5174,6 @@ def _helix_pack_extend(request, group: int) -> int: gather_ids.append( len(input_ids) - 1 - (self.original_max_draft_len - request.py_num_accepted_draft_tokens)) - num_accepted_draft_tokens.append( - request.py_num_accepted_draft_tokens) sequence_lengths.append(1 + self.original_max_draft_len) prompt_lengths.append(request.py_prompt_len) @@ -5231,7 +5243,6 @@ def _helix_pack_extend(request, group: int) -> int: # overhead (saves ~3 append calls per request). draft_lens.extend([0] * (_n_gen * beam_width)) sequence_lengths.extend([1] * (_n_gen * beam_width)) - num_accepted_draft_tokens.extend([0] * (_n_gen * beam_width)) for request in generation_requests: request_ids.append(request.py_request_id) @@ -5403,18 +5414,36 @@ def _helix_pack_extend(request, group: int) -> int: previous_batch_len = len(previous_batch_indices) def previous_seq_slots_device(): - previous_batch_indices_host = torch.tensor( - previous_batch_indices, - dtype=torch.int, - pin_memory=prefer_pinned()) previous_slots = self.previous_batch_indices_cuda[: previous_batch_len] - previous_slots.copy_(previous_batch_indices_host, non_blocking=True) + if previous_batch_indices != self._staged_previous_batch_indices: + previous_batch_indices_host = torch.tensor( + previous_batch_indices, + dtype=torch.int, + pin_memory=prefer_pinned()) + previous_slots.copy_(previous_batch_indices_host, + non_blocking=True) + self._staged_previous_batch_indices = list( + previous_batch_indices) return previous_slots num_tokens = len(input_ids) num_draft_tokens = len(draft_tokens) total_num_tokens = len(position_ids) + # Where the step-copy kernels run, the overlap gathers below are one + # kernel. On a CUDA graph step whose attention metadata is a plain + # TrtllmAttentionMetadata (its prepare() reads none of these inputs on + # the device), the gathers, the positions and the metadata's prompt + # lengths, KV lengths and block offsets are staged instead and written + # by one launch right after attn_metadata.prepare(). + step_copies = (self._step_input_stage + if self.enable_spec_decode else None) + if step_copies is not None: + step_copies.begin() + stage = (step_copies if step_copies is not None + and attn_metadata.is_cuda_graph and not self.use_mrope + and type(attn_metadata) is TrtllmAttentionMetadata + and attn_metadata.fp4_mla_state is None else None) assert total_num_tokens <= self.max_num_tokens, ( f"total_num_tokens ({total_num_tokens}) should be less than or equal to max_num_tokens ({self.max_num_tokens})" ) @@ -5431,53 +5460,38 @@ def previous_seq_slots_device(): pin_memory=prefer_pinned()) self.draft_tokens_cuda[:len(draft_tokens)].copy_(draft_tokens, non_blocking=True) - if self.is_spec_decode and len(num_accepted_draft_tokens) > 0: - num_accepted_draft_tokens = torch.tensor(num_accepted_draft_tokens, - dtype=torch.int, - pin_memory=prefer_pinned()) - self.num_accepted_draft_tokens_cuda[:len( - num_accepted_draft_tokens)].copy_(num_accepted_draft_tokens, - non_blocking=True) if next_draft_tokens_device is not None: - # Initialize these two values to zeros - self.previous_pos_id_offsets_cuda *= 0 - self.previous_kv_lens_offsets_cuda *= 0 runtime_tokens_per_gen_step = self.get_runtime_tokens_per_gen_step( self.runtime_draft_len) runtime_draft_token_buffer_width = runtime_tokens_per_gen_step - 1 + num_extend_reqeust_wo_dummy = len(extend_requests) - len( + extend_dummy_requests) + # The offsets of requests without a previous batch and of dummy + # requests must read 0. On a CUDA graph step (no token padding) + # where every row is a request with a previous batch, the gathers + # below overwrite all rows _preprocess_inputs reads, so the buffers + # need no zeroing first. + if not (attn_metadata.is_cuda_graph and previous_batch_len + == num_extend_reqeust_wo_dummy and not extend_dummy_requests + and not generation_requests and not first_draft_requests + and scheduled_requests.num_context_requests == 0): + self.previous_pos_id_offsets_cuda.zero_() + self.previous_kv_lens_offsets_cuda.zero_() if previous_batch_len > 0: previous_slots = previous_seq_slots_device() - # previous input ids previous_batch_tokens = (previous_batch_len * runtime_tokens_per_gen_step) - new_tokens = new_tokens_device.transpose( - 0, - 1)[previous_slots, :runtime_tokens_per_gen_step].flatten() - self.input_ids_cuda[num_tokens:num_tokens + - previous_batch_tokens].copy_( - new_tokens, non_blocking=True) - - # previous draft tokens - previous_batch_draft_tokens = (previous_batch_len * - runtime_draft_token_buffer_width) - if runtime_draft_token_buffer_width > 0: - self.draft_tokens_cuda[ - num_draft_tokens:num_draft_tokens + - previous_batch_draft_tokens].copy_( - next_draft_tokens_device[ - previous_slots, : - runtime_draft_token_buffer_width].flatten(), - non_blocking=True) - # prepare data for the preprocess inputs - kv_len_offsets_device = (new_tokens_lens_device - - runtime_tokens_per_gen_step) - previous_pos_indices_host = torch.tensor( - previous_pos_indices, - dtype=torch.int, - pin_memory=prefer_pinned()) - self.previous_pos_indices_cuda[0:previous_batch_tokens].copy_( - previous_pos_indices_host, non_blocking=True) + if previous_pos_indices != self._staged_previous_pos_indices: + previous_pos_indices_host = torch.tensor( + previous_pos_indices, + dtype=torch.int, + pin_memory=prefer_pinned()) + self.previous_pos_indices_cuda[ + 0:previous_batch_tokens].copy_( + previous_pos_indices_host, non_blocking=True) + self._staged_previous_pos_indices = list( + previous_pos_indices) # The order of requests in a batch: [context requests, generation requests] # generation requests: ['requests that do not have previous batch', 'requests that already have previous batch', 'dummy requests'] @@ -5487,27 +5501,60 @@ def previous_seq_slots_device(): # Therefore, both of self.previous_pos_id_offsets_cuda and self.previous_kv_lens_offsets_cuda are also 3 segments. # For 1) 'requests that do not have previous batch': disable overlap scheduler or the first step in the generation server of disaggregated serving. # Set these requests' previous_pos_id_offsets and previous_kv_lens_offsets to '0' to skip the value changes in _preprocess_inputs. - # Already set to '0' during initialization. + # Zeroed above. # For 2) 'requests that already have previous batch': enable overlap scheduler. - # Set their previous_pos_id_offsets and previous_kv_lens_offsets according to new_tokens_lens_device and kv_len_offsets_device. + # Set their previous_pos_id_offsets and previous_kv_lens_offsets according to new_tokens_lens_device. # For 3) 'dummy requests': pad dummy requests for CUDA graph or attention dp. - # Already set to '0' during initialization. - - num_extend_reqeust_wo_dummy = len(extend_requests) - len( - extend_dummy_requests) - self.previous_pos_id_offsets_cuda[ - (num_extend_reqeust_wo_dummy - previous_batch_len) * - runtime_tokens_per_gen_step:num_extend_reqeust_wo_dummy * - runtime_tokens_per_gen_step].copy_( - new_tokens_lens_device[self.previous_pos_indices_cuda[ - 0:previous_batch_tokens]], - non_blocking=True) - - self.previous_kv_lens_offsets_cuda[ - num_extend_reqeust_wo_dummy - - previous_batch_len:num_extend_reqeust_wo_dummy].copy_( - kv_len_offsets_device[previous_slots], - non_blocking=True) + # Zeroed above. + previous_begin = (num_extend_reqeust_wo_dummy - + previous_batch_len) + if step_copies is not None and step_copies.gather( + new_tokens_device, next_draft_tokens_device, + new_tokens_lens_device, + self.previous_batch_indices_cuda, + self.previous_pos_indices_cuda, previous_batch_len, + runtime_tokens_per_gen_step, + runtime_draft_token_buffer_width, self.input_ids_cuda, + num_tokens, self.draft_tokens_cuda, num_draft_tokens, + self.previous_pos_id_offsets_cuda, + previous_begin * runtime_tokens_per_gen_step, + self.previous_kv_lens_offsets_cuda, previous_begin): + if stage is None: + step_copies.commit() + else: + # previous input ids + new_tokens = new_tokens_device.transpose(0, 1)[ + previous_slots, :runtime_tokens_per_gen_step].flatten() + self.input_ids_cuda[num_tokens:num_tokens + + previous_batch_tokens].copy_( + new_tokens, non_blocking=True) + + # previous draft tokens + previous_batch_draft_tokens = ( + previous_batch_len * runtime_draft_token_buffer_width) + if runtime_draft_token_buffer_width > 0: + self.draft_tokens_cuda[ + num_draft_tokens:num_draft_tokens + + previous_batch_draft_tokens].copy_( + next_draft_tokens_device[ + previous_slots, : + runtime_draft_token_buffer_width].flatten(), + non_blocking=True) + # prepare data for the preprocess inputs + kv_len_offsets_device = (new_tokens_lens_device - + runtime_tokens_per_gen_step) + self.previous_pos_id_offsets_cuda[ + previous_begin * runtime_tokens_per_gen_step: + num_extend_reqeust_wo_dummy * + runtime_tokens_per_gen_step].copy_( + new_tokens_lens_device[ + self.previous_pos_indices_cuda[ + 0:previous_batch_tokens]], + non_blocking=True) + self.previous_kv_lens_offsets_cuda[ + previous_begin:num_extend_reqeust_wo_dummy].copy_( + kv_len_offsets_device[previous_slots], + non_blocking=True) elif new_tokens_device is not None: seq_slots_device = previous_seq_slots_device() @@ -5572,16 +5619,18 @@ def previous_seq_slots_device(): final_position_ids = self.mrope_position_ids_cuda[:, :, : total_num_tokens] else: - self.position_ids_cuda[:total_num_tokens].copy_(host_position_ids, - non_blocking=True) + if stage is None or not stage.copy( + self.position_ids_cuda[:total_num_tokens], + host_position_ids): + self.position_ids_cuda[:total_num_tokens].copy_( + host_position_ids, non_blocking=True) final_position_ids = self.position_ids_cuda[: total_num_tokens].unsqueeze( 0) if self.enable_spec_decode: - self.gather_ids_cuda[:len(gather_ids)].copy_(torch.tensor( - gather_ids, dtype=torch.int, pin_memory=prefer_pinned()), - non_blocking=True) + copy_to_device_if_changed(self.gather_ids_cuda, + torch.tensor(gather_ids, dtype=torch.int)) if self.mapping.has_cp_helix(): # A non-None owned-count list is what arms @@ -5657,7 +5706,15 @@ def previous_seq_slots_device(): # pre-prepare counts so the steady-gen recording below stores values # that the per-step prepare() can re-clamp from scratch. num_cached_tokens_snapshot = list(num_cached_tokens_per_seq) - attn_metadata.prepare() + if stage is None: + attn_metadata.prepare() + else: + attn_metadata.h2d_stage = stage + try: + attn_metadata.prepare() + finally: + attn_metadata.h2d_stage = None + stage.commit() cross_attention_inputs = (self._prepare_enc_dec_cross_attn_inputs( cross_encoder_hidden_states, cross_encoder_seq_lens, @@ -5775,8 +5832,6 @@ def previous_seq_slots_device(): # num_generations / num_tokens / seq_lens are set above, before the # attention-DP allgather that must agree with prepare(). spec_metadata.host_position_ids = host_position_ids - spec_metadata.num_accepted_draft_tokens = self.num_accepted_draft_tokens_cuda[:len( - num_accepted_draft_tokens)] if context_prompt_lookahead is not None: spec_metadata.populate_context_prompt_lookahead( context_prompt_lookahead) diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index acd9da97a83d..5e4770d63b8b 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -72,7 +72,7 @@ from ..modules.decoder_layer import DecoderLayer from ..moe.expert_statistic import ExpertStatistic from ..speculative.drafter import Drafter -from ..speculative.spec_sampler_base import SampleStateTensorsSpec +from ..speculative.spec_sampler_base import SampleStateTensorsSpec, SpecSampler from ..speculative.speculation_gate import SpeculationGate from ..speculative.utils import update_draft_len from .adp_iter_stats import ADPIterStatsBuffer @@ -5420,7 +5420,16 @@ def _executor_loop_overlap(self): if can_queue: guided_decoder_failed_requests = None - with self.perf_manager.record_perf_events( + # The one-model speculative sampler only moves the forward's + # outputs into its stores, and the next forward reads those + # on the execution stream. Sampling on that stream too keeps + # the chain between two forwards on one stream (None: the + # current stream). + sample_stream = (self.execution_stream if isinstance( + self.sampler, SpecSampler) else None) + with torch.cuda.stream( + sample_stream + ), self.perf_manager.record_perf_events( None, gpu_sample_end) as sample_timing: with self._step_scope(scheduled_batch, phase="sample"): if self.guided_decoder is not None: diff --git a/tensorrt_llm/_torch/speculative/interface.py b/tensorrt_llm/_torch/speculative/interface.py index 37992e2f8835..d6d6864f48b4 100644 --- a/tensorrt_llm/_torch/speculative/interface.py +++ b/tensorrt_llm/_torch/speculative/interface.py @@ -509,8 +509,6 @@ class SpecMetadata: host_position_ids: Optional[torch.Tensor] = None # The gather ids for logits. gather_ids: Optional[torch.Tensor] = None - # The number of accepted draft tokens for each request. - num_accepted_draft_tokens: Optional[torch.Tensor] = None # The number of tokens for speculative model/layer num_tokens: int = 0 # The number of tokens for speculative model/layer of different rank @@ -863,6 +861,11 @@ def _upload(dst: torch.Tensor, values: list[int]) -> None: pin_memory=prefer_pinned()), non_blocking=True) + # An all-greedy batch runs the argmax graph, which reads none of these + # buffers. The offset windows above still advance, so skipping the + # copies changes no later batch's streams. + if self.is_all_greedy_sample: + return _upload(self.request_seeds, request_seeds) _upload(self.request_offsets, request_offsets) _upload(self.seeds, flat_seeds) diff --git a/tensorrt_llm/_torch/speculative/spec_sampler_base.py b/tensorrt_llm/_torch/speculative/spec_sampler_base.py index 030fb18f4fd9..8cd5ad530b76 100644 --- a/tensorrt_llm/_torch/speculative/spec_sampler_base.py +++ b/tensorrt_llm/_torch/speculative/spec_sampler_base.py @@ -27,7 +27,9 @@ import torch +from ..._utils import copy_to_device_if_changed from ...sampling_params import SamplingParams +from ..cute_dsl_kernels.spec_step_copies import op as spec_step_copies from ..pyexecutor.llm_request import LlmRequest, LlmRequestState, get_draft_token_length from ..pyexecutor.resource_manager import BaseResourceManager from ..pyexecutor.sampler import ( @@ -301,6 +303,16 @@ def __init__( # tree (K=6, T=60), MTP dynamic tree, PARD (T=2K-1) and the linear # modes -- none exceed it. self.max_accepted_path_len = args.max_draft_len + 1 + # Where the step-copy kernels run, one kernel moves a step's outputs + # into the stores below, reading each row's slot from _slot_table. + self._store_scatter: Optional[spec_step_copies.SlotScatter] = None + self._slot_table: Optional[torch.Tensor] = None + if spec_step_copies.is_supported(): + self._store_scatter = spec_step_copies.SlotScatter() + self._slot_table = torch.zeros((seq_slots,), dtype=torch.int32, device="cuda") + # Recorded after the last step's host copies of the stores, which read + # them on the D2H side stream. + self._store_copies_done: Optional[torch.cuda.Event] = None self.store = self.Store( new_tokens=int_tensor((self.max_accepted_path_len, seq_slots, self.max_beam_width)), next_new_tokens=int_tensor( @@ -371,6 +383,52 @@ def update_requests( req.py_rewind_len = runtime_draft_len - req.py_num_accepted_draft_tokens self._request_common_handling(req, next_draft_tokens_list, runtime_draft_len) + def _scatter_to_stores( + self, outputs: dict[str, torch.Tensor], num_skip: int, slots: list[int] + ) -> None: + """Move rows ``num_skip`` onward of the forward's outputs into the stores at ``slots``.""" + num_sampling_requests = len(slots) + slots_device = torch.as_tensor(slots, dtype=torch.long) + slots_device = slots_device.to(device="cuda", non_blocking=True) + + o_new_tokens = outputs["new_tokens"][num_skip : num_skip + num_sampling_requests] + o_new_tokens_lens = outputs["new_tokens_lens"][num_skip : num_skip + num_sampling_requests] + o_next_draft_tokens = outputs["next_draft_tokens"][ + num_skip : num_skip + num_sampling_requests + ] + o_next_new_tokens = outputs["next_new_tokens"][num_skip : num_skip + num_sampling_requests] + + # Pad or truncate to match fixed-size store buffers for index_copy_. + # The worker output width tracks runtime_draft_len, which dynamic draft + # length shrinks below the statically allocated store width. + new_tokens_width = self.store.new_tokens.shape[0] + next_new_tokens_width = self.store.next_new_tokens.shape[0] + draft_tokens_width = self.store.next_draft_tokens.shape[1] + if o_new_tokens.shape[1] < new_tokens_width: + o_new_tokens = torch.nn.functional.pad( + o_new_tokens, (0, new_tokens_width - o_new_tokens.shape[1]) + ) + elif o_new_tokens.shape[1] > new_tokens_width: + o_new_tokens = o_new_tokens[:, :new_tokens_width] + if o_next_draft_tokens.shape[1] < draft_tokens_width: + o_next_draft_tokens = torch.nn.functional.pad( + o_next_draft_tokens, (0, draft_tokens_width - o_next_draft_tokens.shape[1]) + ) + elif o_next_draft_tokens.shape[1] > draft_tokens_width: + o_next_draft_tokens = o_next_draft_tokens[:, :draft_tokens_width] + if o_next_new_tokens.shape[1] < next_new_tokens_width: + o_next_new_tokens = torch.nn.functional.pad( + o_next_new_tokens, (0, next_new_tokens_width - o_next_new_tokens.shape[1]) + ) + elif o_next_new_tokens.shape[1] > next_new_tokens_width: + o_next_new_tokens = o_next_new_tokens[:, :next_new_tokens_width] + + # Use index_copy_ for efficient copying (slots are unique) + self.store.new_tokens.squeeze(-1).T.index_copy_(0, slots_device, o_new_tokens) + self.store.next_new_tokens.squeeze(-1).T.index_copy_(0, slots_device, o_next_new_tokens) + self.store.new_tokens_lens.index_copy_(0, slots_device, o_new_tokens_lens) + self.store.next_draft_tokens.index_copy_(0, slots_device, o_next_draft_tokens) + def sample_async( self, scheduled_requests: ScheduledRequests, @@ -396,7 +454,6 @@ def sample_async( num_skip = len(scheduled_requests.context_requests_chunking) finished_context_requests = scheduled_requests.context_requests_last_chunk sampling_requests = finished_context_requests + scheduled_requests.generation_requests - num_sampling_requests = len(sampling_requests) # Snapshot each request's draft count for THIS step before # _add_dummy_draft_tokens below installs placeholder drafts on @@ -413,47 +470,27 @@ def sample_async( for r in sampling_requests ] - slots = torch.as_tensor([r.py_seq_slot for r in sampling_requests], dtype=torch.long) - slots = slots.to(device="cuda", non_blocking=True) - - o_new_tokens = outputs["new_tokens"][num_skip : num_skip + num_sampling_requests] - o_new_tokens_lens = outputs["new_tokens_lens"][num_skip : num_skip + num_sampling_requests] - o_next_draft_tokens = outputs["next_draft_tokens"][ - num_skip : num_skip + num_sampling_requests - ] - o_next_new_tokens = outputs["next_new_tokens"][num_skip : num_skip + num_sampling_requests] - runtime_draft_len = o_next_draft_tokens.shape[1] - - # Pad or truncate to match fixed-size store buffers for index_copy_. - # The worker output width tracks runtime_draft_len, which dynamic draft - # length shrinks below the statically allocated store width. - new_tokens_width = self.store.new_tokens.shape[0] - next_new_tokens_width = self.store.next_new_tokens.shape[0] - draft_tokens_width = self.store.next_draft_tokens.shape[1] - if o_new_tokens.shape[1] < new_tokens_width: - o_new_tokens = torch.nn.functional.pad( - o_new_tokens, (0, new_tokens_width - o_new_tokens.shape[1]) + runtime_draft_len = outputs["next_draft_tokens"].shape[1] + if self._store_copies_done is not None: + # The previous step's host copies may still be reading the stores. + torch.cuda.current_stream().wait_event(self._store_copies_done) + slots = [r.py_seq_slot for r in sampling_requests] + scattered = False + if self._store_scatter is not None: + # A step whose batch is unchanged leaves the slot table as it is. + copy_to_device_if_changed(self._slot_table, torch.tensor(slots, dtype=torch.int32)) + scattered = self._store_scatter.scatter( + outputs, + num_skip, + len(slots), + self._slot_table, + self.store.new_tokens, + self.store.next_new_tokens, + self.store.new_tokens_lens, + self.store.next_draft_tokens, ) - elif o_new_tokens.shape[1] > new_tokens_width: - o_new_tokens = o_new_tokens[:, :new_tokens_width] - if o_next_draft_tokens.shape[1] < draft_tokens_width: - o_next_draft_tokens = torch.nn.functional.pad( - o_next_draft_tokens, (0, draft_tokens_width - o_next_draft_tokens.shape[1]) - ) - elif o_next_draft_tokens.shape[1] > draft_tokens_width: - o_next_draft_tokens = o_next_draft_tokens[:, :draft_tokens_width] - if o_next_new_tokens.shape[1] < next_new_tokens_width: - o_next_new_tokens = torch.nn.functional.pad( - o_next_new_tokens, (0, next_new_tokens_width - o_next_new_tokens.shape[1]) - ) - elif o_next_new_tokens.shape[1] > next_new_tokens_width: - o_next_new_tokens = o_next_new_tokens[:, :next_new_tokens_width] - - # Use index_copy_ for efficient copying (slots are unique) - self.store.new_tokens.squeeze(-1).T.index_copy_(0, slots, o_new_tokens) - self.store.next_new_tokens.squeeze(-1).T.index_copy_(0, slots, o_next_new_tokens) - self.store.new_tokens_lens.index_copy_(0, slots, o_new_tokens_lens) - self.store.next_draft_tokens.index_copy_(0, slots, o_next_draft_tokens) + if not scattered: + self._scatter_to_stores(outputs, num_skip, slots) # Create sample state with async D2H copy device_tensors = SampleStateTensorsSpec( @@ -462,12 +499,26 @@ def sample_async( next_draft_tokens=self.store.next_draft_tokens, ) - host_tensors = SampleStateTensorsSpec( - new_tokens=self._copy_to_host(self.store.new_tokens), - new_tokens_lens=self._copy_to_host(self.store.new_tokens_lens), - next_draft_tokens=self._copy_to_host(self.store.next_draft_tokens), - ) - sampler_event = self._record_sampler_event() + if self._async_worker_active(): + host_tensors = SampleStateTensorsSpec( + new_tokens=self._copy_to_host(self.store.new_tokens), + new_tokens_lens=self._copy_to_host(self.store.new_tokens_lens), + next_draft_tokens=self._copy_to_host(self.store.next_draft_tokens), + ) + sampler_event = self._record_sampler_event() + else: + # The next step's input preparation waits on this stream and only + # needs the stores, so the host copies run on the side stream. + # update_requests syncs their event before reading them, and the + # next step's store update waits for it on this stream. + with self._make_side_stream_copier() as copier: + host_tensors = SampleStateTensorsSpec( + new_tokens=copier.stage_copy_to_host(self.store.new_tokens), + new_tokens_lens=copier.stage_copy_to_host(self.store.new_tokens_lens), + next_draft_tokens=copier.stage_copy_to_host(self.store.next_draft_tokens), + ) + self._store_copies_done = copier.event + sampler_event = self._record_sampler_event(side_stream_event=copier.event) # Add dummy draft tokens to context requests for KV cache preparation for request in finished_context_requests: diff --git a/tensorrt_llm/_utils.py b/tensorrt_llm/_utils.py index c875715c4501..b45bb486b740 100644 --- a/tensorrt_llm/_utils.py +++ b/tensorrt_llm/_utils.py @@ -1336,6 +1336,34 @@ def maybe_pin_memory(tensor: torch.Tensor) -> torch.Tensor: return tensor +def copy_to_device_if_changed(dst: torch.Tensor, host: torch.Tensor) -> None: + """Copy ``host`` into the leading elements of the device buffer ``dst`` + unless they already hold it. + + ``dst`` is the whole buffer (not a slice), and the values it last received + are kept on it, so a reallocated buffer starts over. Only for per-step + metadata buffers that nothing else writes: a decode step whose batch is + unchanged then enqueues no copy. "Nothing else" includes other tensor + objects over the same memory, such as the views a shared buffer pool hands + to each metadata object. The copy itself goes through fresh pinned staging, + so ``host`` may be reused right away. + """ + flat = host.reshape(-1) + n = flat.numel() + last = getattr(dst, "_last_host_values", None) + if last is not None and last.numel() >= n and torch.equal(last[:n], flat): + return + staging = torch.empty_like(flat, device="cpu", pin_memory=prefer_pinned()) + staging.copy_(flat) + dst.view(-1)[:n].copy_(staging, non_blocking=True) + if last is None or last.numel() <= n: + dst._last_host_values = flat.clone() + else: + last = last.clone() + last[:n] = flat + dst._last_host_values = last + + def async_tensor_h2d(data, dtype: torch.dtype, device: Union[str, torch.device]) -> torch.Tensor: """Build a CPU tensor from `data` and ship it to `device` with a diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index 2ca039299f9c..3a7ca13438e7 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -238,6 +238,10 @@ l0_b200: # accuracy deviation found so far is Blackwell-only), so it is not # hw-agnostic and needs its own Blackwell run. - unittest/_torch/speculative/test_fused_sampling_op.py + # One-model speculative decoding step-copy kernels and the engine's staged + # decode-step inputs (SM 100 only; they skip on every other list). + - unittest/_torch/cute_dsl_kernels/test_spec_step_copies.py + - unittest/_torch/executor/test_pytorch_model_engine.py -k "test_staged_spec_decode_graph_step" - unittest/_torch/thop/parallel TIMEOUT (90) - unittest/_torch/visual_gen/kernels/parallel - unittest/_torch/thop/serial diff --git a/tests/unittest/_torch/cute_dsl_kernels/test_spec_step_copies.py b/tests/unittest/_torch/cute_dsl_kernels/test_spec_step_copies.py new file mode 100644 index 000000000000..10408b60106d --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/test_spec_step_copies.py @@ -0,0 +1,822 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``spec_step_copies`` (the one-model speculative decoding step's copy kernels) against the torch ops each replaces, +bit for bit (int32 copies), at every split of R = 1 .. 8 requests x T = 1, 2, 4, 8 tokens, on buffers sized like an +engine's for max batch 8 and max draft length 7 (sampler stores [8, slots, 1], [8, slots, 1], [slots], [slots, 7] over +8 or 16 slots; draft-token buffer 56, KV-length offsets 8, per-token buffers max_num_tokens). + +* ``SlotScatter.scatter`` vs SpecSampler's torch store update (each output padded with zeros or cut to its store's + width, then four ``index_copy_`` by slot): outputs narrower than (dynamic draft length), as wide as and wider than the + stores, and mixed per field; row_begin 0 and > 0; shuffled distinct slot tables. +* ``StepInputStage.gather`` committed on its own vs PyTorchModelEngine._prepare_tp_inputs's torch overlap gathers (the + stores by slot into the input ids and draft tokens, the lengths by the per-token index list into the position + offsets, ``new_tokens_lens - T`` by slot into the KV-length offsets): the engine's offsets (the requests without a + previous batch first), zero and arbitrary offsets, a draft width below T - 1, the engine's and arbitrary index lists. +* ``StepInputStage`` vs those gathers plus what it stages instead of copying: ``dst.copy_(pinned, non_blocking=True)`` + per staged copy (positions, prompt and KV lengths, a fourth) and the KV cache manager's block-offset copy + (``copy_batch_block_offsets_to_device``, target and draft managers, bad pages included): the engine's step, all four + copy slots at offset destinations, the gathers alone, the copies alone. +* The calls the kernels do not cover launch nothing and return False, so the caller keeps its torch path. + +Every buffer starts random and is compared whole (what an op must not touch included), every case runs twice (reruns +bit-identical), and one CUDA graph per split family (R = 1 .. 8 at one T) captures each step's scatter, gather and +staged commit; it is replayed with every input (and the staged values in the stages' pinned records) rewritten in +place, bit-identical to the eager ops and to the torch ops on every replay. + +Table: ``python3 test_spec_step_copies.py report``. Timing: ``python3 test_spec_step_copies.py time`` (CUDA graphs of +back-to-back calls, median us per call over 15 replays, each kernel vs the torch ops it replaces, at every split). +""" + +import statistics +import sys + +import pytest +import torch + +MAX_BATCH = 8 # max_batch_size +MAX_DRAFT = 7 # max_draft_len (linear speculation: max_total_draft_tokens == max_draft_len) +STORE_WIDTH = MAX_DRAFT + 1 # the sampler's new_tokens / next_new_tokens stores +MAX_NUM_TOKENS = 8192 # the engine's per-token buffers (input ids, positions, previous_*) +DRAFT_BUFFER = MAX_DRAFT * MAX_BATCH # draft_tokens_cuda: max_draft_loop_tokens * batch_size +STAGE_CAPACITY = 3 * MAX_NUM_TOKENS # PyTorchModelEngine._get_step_input_stage +BLOCKS = 40 # KV blocks per sequence (copy_batch_block_offsets_to_device needs a multiple of 4) +TABLE_SEQS = 24 # rows of a KV cache manager's pinned host block-offset table +POOLS = (2, 1) # pools of the target and the draft KV cache managers +COPY_NAMES = ("position_ids", "prompt_lens", "kv_lens", "extra") # staged copies, in order + +ROWS = list(range(1, MAX_BATCH + 1)) +TOKENS = [1, 2, 4, 8] +SPLITS = [(r, t) for t in TOKENS for r in ROWS] + + +def _op(): + from tensorrt_llm._torch.cute_dsl_kernels.spec_step_copies import op + + return op + + +def _supported() -> bool: + return torch.cuda.is_available() and _op().is_supported() + + +pytestmark = pytest.mark.skipif( + not _supported(), reason="the step-copy kernels run on SM 100 with the CuTe DSL" +) + + +def _pinned_ok() -> bool: + """StepInputStage reads a pinned host record in place; it refuses to run where pinned memory is not preferred.""" + from tensorrt_llm._utils import prefer_pinned + + return prefer_pinned() + + +_eager = {} + + +def eager_stage(): + """The StepInputStage of the eager runs (never captured: a captured commit's stage cannot stage again).""" + if "stage" not in _eager: + _eager["stage"] = _op().StepInputStage(STAGE_CAPACITY) + return _eager["stage"] + + +def kernel_scatter(args) -> None: + """``SlotScatter.scatter`` on ``args``, which it must cover.""" + assert _op().SlotScatter().scatter(*args), "scatter declined a covered call" + + +def kernel_gather(stage, args) -> None: + """The overlap gathers on ``args`` (which they must cover) committed on their own through ``stage``.""" + stage.begin() + assert stage.gather(*args), "gather declined a covered call" + stage.commit() + + +# ---------------------------------------------------------------------------------------------------------------- +# The torch ops each kernel replaces +# ---------------------------------------------------------------------------------------------------------------- + + +def torch_scatter( + outputs, + num_skip, + num_sampling_requests, + slot_table, + store_new_tokens, + store_next_new_tokens, + store_new_tokens_lens, + store_next_draft_tokens, +): + """SpecSampler's torch store update (SlotScatter.scatter's arguments; the slots as the long tensor it builds from + the requests' seq slots).""" + slots = slot_table[:num_sampling_requests].long() + end = num_skip + num_sampling_requests + o_new_tokens = outputs["new_tokens"][num_skip:end] + o_new_tokens_lens = outputs["new_tokens_lens"][num_skip:end] + o_next_draft_tokens = outputs["next_draft_tokens"][num_skip:end] + o_next_new_tokens = outputs["next_new_tokens"][num_skip:end] + + def fit(t, width): # pad with zeros or truncate to the store width + if t.shape[1] < width: + return torch.nn.functional.pad(t, (0, width - t.shape[1])) + return t[:, :width] + + o_new_tokens = fit(o_new_tokens, store_new_tokens.shape[0]) + o_next_draft_tokens = fit(o_next_draft_tokens, store_next_draft_tokens.shape[1]) + o_next_new_tokens = fit(o_next_new_tokens, store_next_new_tokens.shape[0]) + store_new_tokens.squeeze(-1).T.index_copy_(0, slots, o_new_tokens) + store_next_new_tokens.squeeze(-1).T.index_copy_(0, slots, o_next_new_tokens) + store_new_tokens_lens.index_copy_(0, slots, o_new_tokens_lens) + store_next_draft_tokens.index_copy_(0, slots, o_next_draft_tokens) + + +def torch_gather( + new_tokens_device, + next_draft_tokens_device, + new_tokens_lens_device, + previous_batch_indices_cuda, + previous_pos_indices_cuda, + previous_batch_len, + runtime_tokens_per_gen_step, + runtime_draft_token_buffer_width, + input_ids_cuda, + num_tokens, + draft_tokens_cuda, + num_draft_tokens, + previous_pos_id_offsets_cuda, + pos_begin, + previous_kv_lens_offsets_cuda, + kv_begin, +): + """PyTorchModelEngine._prepare_tp_inputs's torch overlap gathers (StepInputStage.gather's arguments).""" + tokens = runtime_tokens_per_gen_step + width = runtime_draft_token_buffer_width + previous_slots = previous_batch_indices_cuda[:previous_batch_len] + previous_batch_tokens = previous_batch_len * tokens + new_tokens = new_tokens_device.transpose(0, 1)[previous_slots, :tokens].flatten() + input_ids_cuda[num_tokens : num_tokens + previous_batch_tokens].copy_( + new_tokens, non_blocking=True + ) + previous_batch_draft_tokens = previous_batch_len * width + if width > 0: + draft_tokens_cuda[num_draft_tokens : num_draft_tokens + previous_batch_draft_tokens].copy_( + next_draft_tokens_device[previous_slots, :width].flatten(), non_blocking=True + ) + kv_len_offsets_device = new_tokens_lens_device - tokens + previous_pos_id_offsets_cuda[pos_begin : pos_begin + previous_batch_tokens].copy_( + new_tokens_lens_device[previous_pos_indices_cuda[0:previous_batch_tokens]], + non_blocking=True, + ) + previous_kv_lens_offsets_cuda[kv_begin : kv_begin + previous_batch_len].copy_( + kv_len_offsets_device[previous_slots], non_blocking=True + ) + + +def torch_block_copy(offsets, table, copy_index, index_scales, kv_offset): + """KVCacheManagerV2.copy_batch_block_offsets's copy (StepInputStage.block_copy's arguments), current stream.""" + from tensorrt_llm.bindings.internal.batch_manager.kv_cache_manager_v2_utils import ( + copy_batch_block_offsets_to_device, + ) + + copy_batch_block_offsets_to_device( + table, offsets, copy_index, index_scales, kv_offset, torch.cuda.current_stream().cuda_stream + ) + + +# ---------------------------------------------------------------------------------------------------------------- +# One decode step's buffers +# ---------------------------------------------------------------------------------------------------------------- + + +def rand_i32(shape, g, low=-(1 << 30), high=1 << 30): + """Random host int32 (``g``: a CPU generator).""" + return torch.randint(low, high, tuple(shape), generator=g, dtype=torch.int32) + + +def shuffled_distinct(count, total, g): + """``count`` distinct indices below ``total`` (>= 2) in a shuffled order: never ascending and never index r at + position r, so a kernel that takes a row for its slot (or ignores the order) fails.""" + while True: + picked = torch.randperm(total, generator=g)[:count].tolist() + if all(p != r for r, p in enumerate(picked)) and (count < 2 or picked != sorted(picked)): + return picked + + +def engine_begins(rows, tokens): + """The engine's gather offsets (input ids, draft tokens, position offsets, KV-length offsets) when the batch's + other MAX_BATCH - rows generation requests (no previous batch; T tokens, T - 1 draft tokens each) come first.""" + first = MAX_BATCH - rows + return first * tokens, first * (tokens - 1), first * tokens, first + + +class Step: + """One decode step of ``rows`` requests x ``tokens`` tokens on engine-sized int32 buffers, all random so that a + stray or a missing write shows (index buffers hold valid indices past their used part): + + * the forward's outputs: new_tokens [N, a], next_new_tokens [N, b], new_tokens_lens [N], next_draft_tokens [N, c] + (default a = b = T, c = T - 1); + * the sampler's slot table [slots] and stores new_tokens / next_new_tokens [8, slots, 1], new_tokens_lens [slots], + next_draft_tokens [slots, 7]; + * the engine's previous_batch_indices / previous_pos_indices, and two output sets ("a": the gathers committed on + their own, "b": the gathers in a full staged step) of input_ids, draft_tokens [56], previous_pos_id_offsets and + previous_kv_lens_offsets [8] (grown only when an offset needs it); + * the staged copies' destinations (positions, prompt and KV lengths [slots], a fourth buffer) and host values; + * the target and draft KV cache managers' device block offsets [pools, slots, 2, 40] and pinned host tables + [pools, 24, 2, 40] (a fifth of the pages bad, -1), copy indices, index scales and KV offsets. + + ``chain``: the gather reads the scatter's slots in another order (the next step's inputs from this step's stores). + """ + + def __init__( + self, + rows, + tokens, + seed, + *, + num_slots=MAX_BATCH, + widths=None, + row_begin=0, + out_rows=None, + begins=(0, 0, 0, 0), + draft_width=None, + pos_list="engine", + chain=False, + copy_at=(0, 0, 0, 0), + ): + self.rows, self.tokens, self.num_slots = rows, tokens, num_slots + self.widths = widths or (tokens, tokens, tokens - 1) + self.row_begin, self.begins = row_begin, begins + self.draft_width = tokens - 1 if draft_width is None else draft_width + self.pos_list, self.chain, self.copy_at = pos_list, chain, copy_at + self.g = torch.Generator().manual_seed(seed) + n_out = max(MAX_BATCH, row_begin + rows) if out_rows is None else out_rows + assert n_out >= row_begin + rows + n_draft = max(DRAFT_BUFFER, begins[1] + rows * self.draft_width) + n_kv = max(MAX_BATCH, begins[3] + rows) + a, b, c = self.widths + device = { + "out_new": (n_out, a), + "out_next": (n_out, b), + "out_lens": (n_out,), + "out_draft": (n_out, c), + "slot_table": (num_slots,), + "st_new": (STORE_WIDTH, num_slots, 1), + "st_next": (STORE_WIDTH, num_slots, 1), + "st_lens": (num_slots,), + "st_draft": (num_slots, MAX_DRAFT), + "prev_slots": (MAX_NUM_TOKENS,), + "prev_pos": (MAX_NUM_TOKENS,), + "position_ids": (MAX_NUM_TOKENS,), + "prompt_lens": (num_slots,), + "kv_lens": (num_slots,), + "extra": (MAX_NUM_TOKENS,), + } + host = {} + for s in "ab": + device[f"input_ids_{s}"] = (MAX_NUM_TOKENS,) + device[f"draft_{s}"] = (n_draft,) + device[f"pos_off_{s}"] = (MAX_NUM_TOKENS,) + device[f"kv_off_{s}"] = (n_kv,) + for m, pools in zip("td", POOLS): + device[f"blk_{m}"] = (pools, num_slots, 2, BLOCKS) + host[f"tab_{m}"] = (pools, TABLE_SEQS, 2, BLOCKS) + host[f"cidx_{m}"] = (rows,) + host[f"scale_{m}"] = (pools,) + host[f"kvoff_{m}"] = (pools,) + self.bufs = {k: torch.empty(v, dtype=torch.int32, device="cuda") for k, v in device.items()} + for k, v in host.items(): + self.bufs[k] = torch.empty(v, dtype=torch.int32, pin_memory=True) + self.rewrite() + + def rewrite(self) -> None: + """New random contents for every buffer and host value, written in place (a captured graph keeps the + addresses).""" + g, rows, tokens, num_slots = self.g, self.rows, self.tokens, self.num_slots + b = self.bufs + for t in b.values(): + if t.is_cuda: + t.copy_(rand_i32(t.shape, g)) + slots = shuffled_distinct(rows, num_slots, g) + table = rand_i32((num_slots,), g, 0, num_slots) + table[:rows] = torch.tensor(slots, dtype=torch.int32) + b["slot_table"].copy_(table) + previous = slots[1:] + slots[:1] if self.chain else shuffled_distinct(rows, num_slots, g) + previous = torch.tensor(previous, dtype=torch.int32) + index = rand_i32((MAX_NUM_TOKENS,), g, 0, num_slots) + index[:rows] = previous + b["prev_slots"].copy_(index) + per_token = rand_i32((MAX_NUM_TOKENS,), g, 0, num_slots) + if self.pos_list == "engine": # each row's slot, once per token + per_token[: rows * tokens] = previous.repeat_interleave(tokens) + b["prev_pos"].copy_(per_token) + self.values = { + "position_ids": rand_i32((rows * tokens,), g, 0, 1 << 20), + "prompt_lens": rand_i32((rows,), g, 1, 1 << 20), + "kv_lens": rand_i32((rows,), g, 1, 1 << 20), + "extra": rand_i32((rows,), g), + } + self.pinned = {k: v.pin_memory() for k, v in self.values.items()} + for m in "td": + pages = rand_i32(b[f"tab_{m}"].shape, g, 0, 1 << 20) + pages[torch.rand(pages.shape, generator=g) < 0.2] = -1 + b[f"tab_{m}"].copy_(pages) + b[f"cidx_{m}"].copy_(torch.tensor(shuffled_distinct(rows, TABLE_SEQS, g))) + b[f"scale_{m}"].copy_(rand_i32(b[f"scale_{m}"].shape, g, 1, 8)) + b[f"kvoff_{m}"].copy_(rand_i32(b[f"kvoff_{m}"].shape, g, 0, 1 << 16)) + + def outputs(self, b): + return { + "new_tokens": b["out_new"], + "next_new_tokens": b["out_next"], + "new_tokens_lens": b["out_lens"], + "next_draft_tokens": b["out_draft"], + } + + def scatter_args(self, b): + """SlotScatter.scatter's arguments on the buffers ``b``.""" + return (self.outputs(b), self.row_begin, self.rows, b["slot_table"], b["st_new"], b["st_next"], + b["st_lens"], b["st_draft"]) # fmt: skip + + def gather_args(self, b, out): + """StepInputStage.gather's arguments on the buffers ``b``, into output set ``out``.""" + ib, db, pb, kb = self.begins + return (b["st_next"], b["st_draft"], b["st_lens"], b["prev_slots"], b["prev_pos"], self.rows, self.tokens, + self.draft_width, b[f"input_ids_{out}"], ib, b[f"draft_{out}"], db, b[f"pos_off_{out}"], pb, + b[f"kv_off_{out}"], kb) # fmt: skip + + def copies(self, b, count): + """The first ``count`` staged copies: (destination view, host value name).""" + return [ + (b[name][at : at + self.values[name].numel()], name) + for name, at in zip(COPY_NAMES[:count], self.copy_at) + ] + + def block_copies(self, b): + """StepInputStage.block_copy's arguments for the target and the draft KV cache managers.""" + return [ + (b[f"blk_{m}"], b[f"tab_{m}"], b[f"cidx_{m}"], b[f"scale_{m}"], b[f"kvoff_{m}"]) + for m in "td" + ] + + def run_stage(self, b, stage, gather=True, copies=3, blocks=True): + """The step's inputs through ``stage`` (a StepInputStage): one commit.""" + stage.begin() + if gather: + assert stage.gather(*self.gather_args(b, "b")), "gather declined" + for dst, name in self.copies(b, copies): + assert stage.copy(dst, self.values[name]), "copy declined" + if blocks: + for args in self.block_copies(b): + assert stage.block_copy(*args), "block copy declined" + stage.commit() + + def torch_stage(self, b, gather=True, copies=3, blocks=True): + """What the stage replaces: the gathers, a non-blocking copy from pinned memory per staged copy, and the KV + cache managers' block-offset copies.""" + if gather: + torch_gather(*self.gather_args(b, "b")) + for dst, name in self.copies(b, copies): + dst.copy_(self.pinned[name], non_blocking=True) + if blocks: + for args in self.block_copies(b): + torch_block_copy(*args) + + +def clone_bufs(bufs): + """A copy of every buffer (pinned host buffers stay pinned).""" + out = {} + for k, t in bufs.items(): + if t.is_cuda: + out[k] = t.clone() + else: + out[k] = torch.empty(t.shape, dtype=t.dtype, pin_memory=t.is_pinned()).copy_(t) + return out + + +def mismatches(x, y): + """Names of the buffers whose contents differ.""" + return [k for k in x if not torch.equal(x[k], y[k])] + + +def compare(bufs, run_torch, run_kernel): + """The buffers where the kernel differs from the torch ops and where a rerun differs from the kernel, each run on + its own copy of ``bufs``.""" + want, got, again = clone_bufs(bufs), clone_bufs(bufs), clone_bufs(bufs) + run_torch(want) + run_kernel(got) + run_kernel(again) + torch.cuda.synchronize() + return dict(torch=mismatches(want, got), rerun=mismatches(got, again)) + + +def result(case, **bad): + """A table row; ``bad`` maps each comparison (torch, eager, rerun) to the buffers where it differed.""" + failed = {k: v for k, v in bad.items() if v} + if not failed: + return dict(case=case, identical="yes (= " + ", = ".join(bad) + ")", ok=True) + detail = "; ".join(f"!= {k}: {', '.join(v)}" for k, v in failed.items()) + return dict(case=case, identical=f"no: {detail}", ok=False) + + +def _seed(op_index, rows, tokens, case): + return 100000 * op_index + 1000 * rows + 10 * tokens + case + + +# ---------------------------------------------------------------------------------------------------------------- +# The cases of one split +# ---------------------------------------------------------------------------------------------------------------- + + +def scatter_cases(rows, tokens): + """(label, Step arguments): the forward's outputs narrower than or as wide as the stores (dynamic draft length; + the engine's batch of 8 rows, the first 8 - R skipped), wider than every store, and mixed per field.""" + t = tokens + fit = "pad" if t < STORE_WIDTH else "exact" + return [ + (f"{fit}: widths ({t}, {t}, {t - 1}) -> (8, 8, 7), 8 slots, row_begin {MAX_BATCH - rows} of 8 rows", + dict(num_slots=MAX_BATCH, row_begin=MAX_BATCH - rows, out_rows=MAX_BATCH)), + (f"cut: widths ({t + 8}, {t + 8}, {t + 7}) -> (8, 8, 7), 16 slots, row_begin 0", + dict(num_slots=2 * MAX_BATCH, widths=(t + 8, t + 8, t + 7), out_rows=rows + 1)), + (f"mixed: widths ({t + 8}, {t}, {t + 7}) -> (8, 8, 7), 16 slots, row_begin 3", + dict(num_slots=2 * MAX_BATCH, widths=(t + 8, t, t + 7), row_begin=3, out_rows=rows + 4)), + ] # fmt: skip + + +def gather_cases(rows, tokens): + """(label, Step arguments): the engine's layout, zero offsets, and arbitrary offsets with a narrower draft width + and an arbitrary per-token index list.""" + begins = engine_begins(rows, tokens) + return [ + (f"engine: offsets {begins}, draft width {tokens - 1}, per-token list = slots x T, 8 slots", + dict(num_slots=MAX_BATCH, begins=begins)), + (f"offsets 0, draft width {tokens - 1}, per-token list = slots x T, 16 slots", + dict(num_slots=2 * MAX_BATCH)), + (f"offsets (37, 5, 11, 3), draft width {(tokens - 1) // 2}, arbitrary per-token list, 16 slots", + dict(num_slots=2 * MAX_BATCH, begins=(37, 5, 11, 3), draft_width=(tokens - 1) // 2, pos_list="any")), + ] # fmt: skip + + +def stage_cases(rows, tokens): + """(label, Step arguments, run_stage arguments).""" + begins = engine_begins(rows, tokens) + return [ + (f"engine step: gathers at {begins}, positions / prompt / KV lengths, target + draft block offsets", + dict(num_slots=MAX_BATCH, begins=begins), dict()), + ("4 copies to offset destinations, gathers at (37, 5, 11, 3), block offsets, 16 slots", + dict(num_slots=2 * MAX_BATCH, begins=(37, 5, 11, 3), pos_list="any", copy_at=(7, 1, 2, 3)), + dict(copies=4)), + ("gathers only, 16 slots", dict(num_slots=2 * MAX_BATCH), dict(copies=0, blocks=False)), + ("copies + block offsets only (no previous batch), 16 slots", + dict(num_slots=2 * MAX_BATCH, copy_at=(5, 3, 0, 0)), dict(gather=False)), + ] # fmt: skip + + +def measure_scatter(rows, tokens): + out = [] + for i, (case, kw) in enumerate(scatter_cases(rows, tokens)): + st = Step(rows, tokens, _seed(1, rows, tokens, i), **kw) + bad = compare( + st.bufs, + lambda b: torch_scatter(*st.scatter_args(b)), + lambda b: kernel_scatter(st.scatter_args(b)), + ) + out.append(result(case, **bad)) + return out + + +def measure_gather(rows, tokens): + stage = eager_stage() + out = [] + for i, (case, kw) in enumerate(gather_cases(rows, tokens)): + st = Step(rows, tokens, _seed(2, rows, tokens, i), **kw) + bad = compare( + st.bufs, + lambda b: torch_gather(*st.gather_args(b, "a")), + lambda b: kernel_gather(stage, st.gather_args(b, "a")), + ) + out.append(result(case, **bad)) + return out + + +def measure_stage(rows, tokens): + stage = eager_stage() + out = [] + for i, (case, kw, run_kw) in enumerate(stage_cases(rows, tokens)): + st = Step(rows, tokens, _seed(3, rows, tokens, i), **kw) + bad = compare( + st.bufs, + lambda b: st.torch_stage(b, **run_kw), + lambda b: st.run_stage(b, stage, **run_kw), + ) + out.append(result(case, **bad)) + return out + + +def measure_graph(tokens, replays=6): + """One CUDA graph for the split family R = 1 .. 8 at ``tokens``: each step's scatter, its gather (from the stores + the scatter wrote, as the next step reads them) and its staged commit, captured once and replayed with every + input rewritten in place; each replay against the eager ops and the torch ops on the same inputs, then a replay + of the last inputs against the first.""" + op = _op() + with_stage = _pinned_ok() + gather_stage = op.StepInputStage(STAGE_CAPACITY) # gathers only: its record is never read + steps = [ + Step(r, tokens, _seed(4, r, tokens, 0), num_slots=2 * MAX_BATCH, row_begin=MAX_BATCH - r, + out_rows=MAX_BATCH, begins=engine_begins(r, tokens), chain=True) + for r in ROWS + ] # fmt: skip + stage = eager_stage() if with_stage else None + # A captured commit reads its stage's record at every replay, so each step keeps its own stage. + captured = [op.StepInputStage(STAGE_CAPACITY) if with_stage else None for _ in steps] + + def run_kernels(st, b, step_stage): + kernel_scatter(st.scatter_args(b)) + kernel_gather(gather_stage, st.gather_args(b, "a")) + if step_stage is not None: + st.run_stage(b, step_stage) + + def run_torch(st, b): + torch_scatter(*st.scatter_args(b)) + torch_gather(*st.gather_args(b, "a")) + if with_stage: + st.torch_stage(b) + + stream = torch.cuda.Stream() + with torch.cuda.stream(stream): + for st in steps: # the kernels compile on their first call, outside capture + run_kernels(st, clone_bufs(st.bufs), stage) + torch.cuda.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + for st, step_stage in zip(steps, captured): + run_kernels(st, st.bufs, step_stage) + torch.cuda.synchronize() + # The staged values sit in each captured stage's pinned record in staging order; a replay reads them in place. + spans = [] + for st, step_stage in zip(steps, captured): + at, step_spans = 0, [] + for _, name in st.copies(st.bufs, 3) if with_stage else []: + n = st.values[name].numel() + assert torch.equal(step_stage.record[at : at + n], st.values[name]), "record layout" + step_spans.append((at, n, name)) + at += n + spans.append(step_spans) + what = "scatter, gather" + (", staged commit" if with_stage else "") + out = [] + before = None + for rep in range(replays): + for st, step_stage, step_spans in zip(steps, captured, spans): + st.rewrite() + for at, n, name in step_spans: + step_stage.record[at : at + n].copy_(st.values[name]) + before = [clone_bufs(st.bufs) for st in steps] + graph.replay() + torch.cuda.synchronize() + bad = dict(torch=[], eager=[]) + for st, start in zip(steps, before): + want, got = clone_bufs(start), clone_bufs(start) + run_torch(st, want) + run_kernels(st, got, stage) + torch.cuda.synchronize() + bad["torch"] += [f"{st.rows}x{tokens} {k}" for k in mismatches(want, st.bufs)] + bad["eager"] += [f"{st.rows}x{tokens} {k}" for k in mismatches(got, st.bufs)] + out.append(result(f"replay {rep}: {what}, every input rewritten", **bad)) + after = [clone_bufs(st.bufs) for st in steps] + for st, start in zip(steps, before): + for k, t in st.bufs.items(): + t.copy_(start[k]) + graph.replay() + torch.cuda.synchronize() + rerun = [ + f"{st.rows}x{tokens} {k}" for st, a in zip(steps, after) for k in mismatches(a, st.bufs) + ] + out.append(result("the last inputs replayed again", rerun=rerun)) + return [dict(family=f"R x {tokens}", **r) for r in out] + + +# ---------------------------------------------------------------------------------------------------------------- +# Tests +# ---------------------------------------------------------------------------------------------------------------- + + +@pytest.mark.parametrize("rows,tokens", SPLITS, ids=[f"{r}x{t}" for r, t in SPLITS]) +def test_scatter(rows, tokens): + bad = [r for r in measure_scatter(rows, tokens) if not r["ok"]] + assert not bad, bad + + +@pytest.mark.parametrize("rows,tokens", SPLITS, ids=[f"{r}x{t}" for r, t in SPLITS]) +def test_gather(rows, tokens): + bad = [r for r in measure_gather(rows, tokens) if not r["ok"]] + assert not bad, bad + + +@pytest.mark.parametrize("rows,tokens", SPLITS, ids=[f"{r}x{t}" for r, t in SPLITS]) +def test_stage(rows, tokens): + if not _pinned_ok(): + pytest.skip("StepInputStage needs pinned host memory (not preferred here)") + bad = [r for r in measure_stage(rows, tokens) if not r["ok"]] + assert not bad, bad + + +@pytest.mark.parametrize("tokens", TOKENS, ids=[f"Rx{t}" for t in TOKENS]) +def test_graph_replay(tokens): + bad = [r for r in measure_graph(tokens) if not r["ok"]] + assert not bad, bad + + +def test_contract(): + """rows == 0 launches nothing; every call the kernels do not cover returns False and launches nothing, so the + caller keeps its torch path; StepInputStage's per-step limits (max_copies = 4 copies, max_block_copies = 2 block + copies, one gather) decline the call past them.""" + op = _op() + assert op.is_supported() + st = Step(MAX_BATCH, STORE_WIDTH, _seed(6, 0, 0, 0)) + b = clone_bufs(st.bufs) + before = clone_bufs(b) + scatter = list(st.scatter_args(b)) + gather = list(st.gather_args(b, "a")) + stage = op.StepInputStage(STAGE_CAPACITY) + assert op.SlotScatter().scatter(*scatter[:2], 0, *scatter[3:]) + stage.begin() + assert stage.gather(*gather[:5], 0, *gather[6:]) + stage.commit() + torch.cuda.synchronize() + assert not mismatches(before, b), "rows == 0 wrote" + # Argument 7 (draft width) == tokens, argument 6 (tokens) wider than the 8-wide store, argument 15 (KV begin) + # ending past the engine's 8 KV-length offsets, and an int64 slot table (argument 3). + for index, value in ( + (7, STORE_WIDTH), + (6, STORE_WIDTH + 1), + (15, 1), + (3, b["prev_slots"].long()), + ): + args = list(gather) + args[index] = value + stage.begin() + assert not stage.gather(*args), f"gather argument {index} = {value!r} was staged" + # Outputs of 8 rows with rows 1 .. 8 asked, an int64 slot table, a non-contiguous output. + assert not op.SlotScatter().scatter(scatter[0], 1, *scatter[2:]) + assert not op.SlotScatter().scatter(*scatter[:3], b["slot_table"].long(), *scatter[4:]) + strided = dict( + scatter[0], next_draft_tokens=torch.empty_like(b["out_draft"]).t().contiguous().t() + ) + assert not op.SlotScatter().scatter(strided, *scatter[1:]) + if _pinned_ok(): + stage.begin() + for _ in range(stage.max_copies): + assert stage.copy(b["extra"][:2], torch.tensor([1, 2], dtype=torch.int32)) + assert not stage.copy(b["extra"][:2], torch.tensor([1, 2], dtype=torch.int32)) + stage.begin() + assert not stage.copy(b["extra"][:2], torch.tensor([1, 2])) # int64 values + assert not stage.copy( + b["out_new"], torch.tensor([1, 2], dtype=torch.int32) + ) # 2-D destination + assert stage.gather(*gather) + assert not stage.gather(*gather) + blocks = st.block_copies(b) + for args in blocks: + assert stage.block_copy(*args) + assert not stage.block_copy(*blocks[0]) + stage.begin() # drops the staged step: nothing was launched + unpinned = blocks[0][1].clone() # a pageable host table + assert not unpinned.is_pinned() + assert not stage.block_copy(blocks[0][0], unpinned, *blocks[0][2:]) + stage.begin() + torch.cuda.synchronize() + assert not mismatches(before, b), "a declined call wrote" + + +# ---------------------------------------------------------------------------------------------------------------- +# Error table (python3 test_spec_step_copies.py report) and timing (python3 test_spec_step_copies.py time) +# ---------------------------------------------------------------------------------------------------------------- + + +def report() -> int: + print(f"{torch.cuda.get_device_name()}") + ok_all = True + sections = [("SlotScatter.scatter", measure_scatter), ("StepInputStage.gather", measure_gather)] + if _pinned_ok(): + sections.append(("StepInputStage", measure_stage)) + else: + print("StepInputStage: skipped (pinned host memory not preferred here)") + for name, measure in sections: + print(f"\n## {name}\n") + print("| split | case | identical | result |") + print("| :-- | :-- | :-- | :-- |") + for rows, tokens in SPLITS: + for r in measure(rows, tokens): + ok_all &= r["ok"] + print(f"| {rows}x{tokens} | {r['case']} | {r['identical']} | {'PASS' if r['ok'] else 'FAIL'} |", + flush=True) # fmt: skip + print("\n## CUDA graph replays (one capture per family R = 1 .. 8)\n") + print("| family | case | identical | result |") + print("| :-- | :-- | :-- | :-- |") + for tokens in TOKENS: + for r in measure_graph(tokens): + ok_all &= r["ok"] + print(f"| {r['family']} | {r['case']} | {r['identical']} | {'PASS' if r['ok'] else 'FAIL'} |", + flush=True) # fmt: skip + print("\nALL PASS" if ok_all else "\nFAIL") + return 0 if ok_all else 1 + + +def time_graph(body, calls, replays=15): + """Per-call us of a CUDA graph of ``calls`` back-to-back ``body(i)``: median (min, max) over ``replays`` replays. + ``body(-1)`` runs first, outside capture.""" + stream = torch.cuda.Stream() + with torch.cuda.stream(stream): + body(-1) + torch.cuda.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + for i in range(calls): + body(i) + torch.cuda.synchronize() + for _ in range(3): + graph.replay() + torch.cuda.synchronize() + per_call = [] + for _ in range(replays): + start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + start.record() + graph.replay() + end.record() + torch.cuda.synchronize() + per_call.append(start.elapsed_time(end) * 1e3 / calls) + return statistics.median(per_call), min(per_call), max(per_call) + + +class StageArm: + """Back-to-back staged commits: fresh stages for every capture (a captured commit's stage cannot stage again), + and one of its own for the warm-up call.""" + + def __init__(self, step, bufs, calls): + self.step, self.bufs, self.calls = step, bufs, calls + self.warm = _op().StepInputStage(STAGE_CAPACITY) + self.fresh = [] + + def __call__(self, i): + if i < 0: + self.fresh = [_op().StepInputStage(STAGE_CAPACITY) for _ in range(self.calls)] + self.step.run_stage(self.bufs, self.warm) + else: + self.step.run_stage(self.bufs, self.fresh[i]) + + +def timing() -> None: + op = _op() + calls = 32 + with_stage = _pinned_ok() + print(f"{torch.cuda.get_device_name()}; graphs of {calls} back-to-back calls on one step's buffers (the engine's " + "layout), 15 replays: median (min-max) us per call") # fmt: skip + print("| split | scatter: torch | scatter: kernel | gather: torch | gather: kernel | stage: torch " + "| stage: kernel |") # fmt: skip + print("| :-- | --: | --: | --: | --: | --: | --: |") + for rows, tokens in SPLITS: + st = Step(rows, tokens, _seed(5, rows, tokens, 0), row_begin=MAX_BATCH - rows, out_rows=MAX_BATCH, + begins=engine_begins(rows, tokens)) # fmt: skip + b = st.bufs + gather_stage = op.StepInputStage(STAGE_CAPACITY) + arms = [ + lambda i: torch_scatter(*st.scatter_args(b)), + lambda i: kernel_scatter(st.scatter_args(b)), + lambda i: torch_gather(*st.gather_args(b, "a")), + lambda i: kernel_gather(gather_stage, st.gather_args(b, "a")), + ] + if with_stage: + arms += [lambda i: st.torch_stage(b), StageArm(st, b, calls)] + res = [[] for _ in arms] + for rep in range(3): # alternating order + order = range(len(arms)) if rep % 2 == 0 else reversed(range(len(arms))) + for a in order: + res[a].append(time_graph(arms[a], calls)) + cells = [] + for timings in res: + meds = sorted(x[0] for x in timings) + lo, hi = min(x[1] for x in timings), max(x[2] for x in timings) + cells.append(f"{meds[1]:.2f} ({lo:.2f}-{hi:.2f})") + cells += ["n/a"] * (6 - len(cells)) + print(f"| {rows}x{tokens} | " + " | ".join(cells) + " |", flush=True) + + +if __name__ == "__main__": + if len(sys.argv) > 1 and sys.argv[1] == "time": + timing() + elif len(sys.argv) > 1 and sys.argv[1] == "report": + sys.exit(report()) + else: + sys.exit(pytest.main([__file__, "-q", "-p", "no:cacheprovider", *sys.argv[1:]])) diff --git a/tests/unittest/_torch/executor/test_copy_to_device_if_changed.py b/tests/unittest/_torch/executor/test_copy_to_device_if_changed.py new file mode 100644 index 000000000000..c4305439b5a4 --- /dev/null +++ b/tests/unittest/_torch/executor/test_copy_to_device_if_changed.py @@ -0,0 +1,119 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``copy_to_device_if_changed``: which calls copy and what the buffer then holds. + +The skip decision does not depend on the device, so the buffer here is a host tensor: a skipped copy shows as a +sentinel written behind the function's back that survives the call. +""" + +import pytest +import torch + +from tensorrt_llm._utils import copy_to_device_if_changed + +pytestmark = pytest.mark.cpu_only + +SENTINEL = -7 + + +def _host(values: list[int]) -> torch.Tensor: + return torch.tensor(values, dtype=torch.int32) + + +def _tamper(dst: torch.Tensor) -> None: + """Overwrite the buffer without the function knowing, so only a real copy restores it.""" + dst.fill_(SENTINEL) + + +def test_first_call_copies() -> None: + dst = torch.zeros(8, dtype=torch.int32) + copy_to_device_if_changed(dst, _host([3, 1, 2])) + assert dst.tolist() == [3, 1, 2, 0, 0, 0, 0, 0] + + +def test_unchanged_values_skip_the_copy() -> None: + dst = torch.zeros(8, dtype=torch.int32) + copy_to_device_if_changed(dst, _host([3, 1, 2])) + _tamper(dst) + copy_to_device_if_changed(dst, _host([3, 1, 2])) + assert dst.eq(SENTINEL).all() + + +def test_changed_values_copy() -> None: + dst = torch.zeros(8, dtype=torch.int32) + copy_to_device_if_changed(dst, _host([3, 1, 2])) + _tamper(dst) + copy_to_device_if_changed(dst, _host([3, 1, 4])) + assert dst[:3].tolist() == [3, 1, 4] + assert dst[3:].eq(SENTINEL).all(), "only the leading elements are written" + + +def test_prefix_of_the_last_values_skips_the_copy() -> None: + dst = torch.zeros(8, dtype=torch.int32) + copy_to_device_if_changed(dst, _host([5, 6, 7, 8])) + _tamper(dst) + copy_to_device_if_changed(dst, _host([5, 6])) + assert dst.eq(SENTINEL).all() + + +def test_longer_values_copy_and_extend_what_is_kept() -> None: + dst = torch.zeros(8, dtype=torch.int32) + copy_to_device_if_changed(dst, _host([5, 6])) + copy_to_device_if_changed(dst, _host([5, 6, 7, 8])) + assert dst[:4].tolist() == [5, 6, 7, 8] + _tamper(dst) + copy_to_device_if_changed(dst, _host([5, 6, 7, 8])) + assert dst.eq(SENTINEL).all() + + +def test_shorter_changed_values_keep_the_tail() -> None: + # A shorter copy rewrites only its leading elements; the buffer's tail still holds the earlier values, and what + # the function keeps says so. + dst = torch.zeros(8, dtype=torch.int32) + copy_to_device_if_changed(dst, _host([5, 6, 7, 8])) + copy_to_device_if_changed(dst, _host([9, 9])) + assert dst[:4].tolist() == [9, 9, 7, 8] + _tamper(dst) + copy_to_device_if_changed(dst, _host([9, 9, 7, 8])) + assert dst.eq(SENTINEL).all() + copy_to_device_if_changed(dst, _host([9, 9, 7, 1])) + assert dst[:4].tolist() == [9, 9, 7, 1] + + +def test_host_values_may_be_reused_right_away() -> None: + dst = torch.zeros(4, dtype=torch.int32) + host = _host([1, 2, 3]) + copy_to_device_if_changed(dst, host) + host.fill_(0) + assert dst[:3].tolist() == [1, 2, 3] + copy_to_device_if_changed(dst, _host([1, 2, 3])) + assert dst[:3].tolist() == [1, 2, 3] + + +def test_a_new_buffer_starts_over() -> None: + first = torch.zeros(4, dtype=torch.int32) + copy_to_device_if_changed(first, _host([1, 2])) + second = torch.full((4,), SENTINEL, dtype=torch.int32) + copy_to_device_if_changed(second, _host([1, 2])) + assert second[:2].tolist() == [1, 2] + + +def test_two_dimensional_values_are_flattened() -> None: + dst = torch.zeros(6, dtype=torch.int32) + copy_to_device_if_changed(dst, torch.tensor([[1, 2], [3, 4]], dtype=torch.int32)) + assert dst[:4].tolist() == [1, 2, 3, 4] + _tamper(dst) + copy_to_device_if_changed(dst, _host([1, 2, 3, 4])) + assert dst.eq(SENTINEL).all() diff --git a/tests/unittest/_torch/executor/test_cuda_graph_capture_replay.py b/tests/unittest/_torch/executor/test_cuda_graph_capture_replay.py index 6c41f7821351..6b0a1b1d75d4 100644 --- a/tests/unittest/_torch/executor/test_cuda_graph_capture_replay.py +++ b/tests/unittest/_torch/executor/test_cuda_graph_capture_replay.py @@ -330,3 +330,82 @@ def forward_fn(inputs): logits_eager = forward_fn(self._make_inputs(attn_metadata, num_tokens, batch_size, value=2)) torch.testing.assert_close(logits_cuda_graph, logits_eager) + + +class TestEngineBuffersAsStaticInputs: + """With static_input_ids / static_position_ids, the graphs' static inputs are views of the engine's own buffers: + the engine's writes are the graphs' inputs, and replay copies an input only when it is another tensor.""" + + def test_graph_reads_the_engine_buffers(self): + batch_size = 1 + engine_input_ids = torch.zeros((8,), device="cuda", dtype=torch.int32) + engine_position_ids = torch.zeros((8,), device="cuda", dtype=torch.int32) + runner = create_mock_cuda_graph_runner( + batch_size, + max_num_tokens=8, + static_input_ids=engine_input_ids, + static_position_ids=engine_position_ids, + ) + assert runner.shared_static_tensors["input_ids"].data_ptr() == engine_input_ids.data_ptr() + assert ( + runner.shared_static_tensors["position_ids"].data_ptr() + == engine_position_ids.data_ptr() + ) + key = KeyType(batch_size=batch_size, draft_len=0, is_first_draft=False) + num_tokens = runner._get_num_tokens_for_key(key) + attn_metadata = object() + + def forward_fn(inputs): + return inputs["input_ids"] * 1000 + inputs["position_ids"][0] + + def engine_inputs(): + return { + "attn_metadata": attn_metadata, + "input_ids": engine_input_ids[:num_tokens], + "position_ids": engine_position_ids[:num_tokens].unsqueeze(0), + } + + engine_input_ids.fill_(3) + engine_position_ids.fill_(4) + runner.capture(key, forward_fn, engine_inputs()) + # The captured forward did not run, so the engine's buffers are as they were. + assert engine_input_ids.tolist() == [3] * 8 + assert engine_position_ids.tolist() == [4] * 8 + + engine_input_ids[:num_tokens] = 7 + engine_position_ids[:num_tokens] = 9 + output = runner.replay(key, engine_inputs()) + assert output.tolist() == [7009] * num_tokens + + # Another tensor as the input is copied into the engine's buffers, which the graph reads. + other = { + "attn_metadata": attn_metadata, + "input_ids": torch.full((num_tokens,), 5, device="cuda", dtype=torch.int32), + "position_ids": torch.full((1, num_tokens), 6, device="cuda", dtype=torch.int32), + } + output = runner.replay(key, other) + assert output.tolist() == [5006] * num_tokens + assert engine_input_ids[:num_tokens].tolist() == [5] * num_tokens + + @pytest.mark.parametrize( + "input_ids, position_ids, use_mrope", + [ + pytest.param(torch.int64, torch.int32, False, id="int64-input-ids"), + pytest.param(torch.int32, None, False, id="no-position-ids"), + pytest.param(torch.int32, torch.int32, True, id="mrope"), + pytest.param("short", torch.int32, False, id="short"), + ], + ) + def test_unusable_engine_buffers_are_rejected(self, input_ids, position_ids, use_mrope): + size = 0 if input_ids == "short" else 8 # the graphs need one token here + dtype = torch.int32 if input_ids == "short" else input_ids + with pytest.raises(ValueError): + create_mock_cuda_graph_runner( + 1, + use_mrope=use_mrope, + max_num_tokens=8, + static_input_ids=torch.zeros((size,), device="cuda", dtype=dtype), + static_position_ids=None + if position_ids is None + else torch.zeros((8,), device="cuda", dtype=position_ids), + ) diff --git a/tests/unittest/_torch/executor/test_pytorch_model_engine.py b/tests/unittest/_torch/executor/test_pytorch_model_engine.py index d86a355f6d21..e9b99e12dd82 100644 --- a/tests/unittest/_torch/executor/test_pytorch_model_engine.py +++ b/tests/unittest/_torch/executor/test_pytorch_model_engine.py @@ -1575,6 +1575,134 @@ def test_promoted_context_precedes_speculative_overlap_generation( [generation.py_seq_slot], 0) kv_cache_manager.shutdown() + def _check_staged_spec_decode_graph_step( + self, use_kv_cache_manager_v2: bool) -> None: + """A CUDA graph decode step of a speculative engine whose step inputs + are staged on the StepInputStage (one launch after the attention + metadata's prepare) writes what the torch path writes.""" + from tensorrt_llm._torch.attention.backends.trtllm import \ + TrtllmAttentionMetadata + from tensorrt_llm._torch.cute_dsl_kernels.spec_step_copies import \ + op as spec_step_copies + from tensorrt_llm._torch.pyexecutor.kv_cache.kv_cache_manager_v2 import \ + KVCacheManagerV2 + from tensorrt_llm.llmapi.llm_args import \ + KvCacheConfig as LlmKvCacheConfig + if not spec_step_copies.is_supported(): + self.skipTest("the step-copy kernels run on SM 100") + max_draft_len = 3 + tokens_per_step = max_draft_len + 1 + model_engine, kv_cache_manager = create_model_engine_and_kvcache( + spec_config=SADecodingConfig(max_draft_len=max_draft_len)) + if use_kv_cache_manager_v2: + kv_cache_manager.shutdown() + kv_cache_manager = KVCacheManagerV2( + LlmKvCacheConfig(max_tokens=512, enable_block_reuse=False), + tensorrt_llm.bindings.internal.batch_manager.CacheType.SELF, + num_layers=1, + num_kv_heads=model_engine.model.config.num_key_value_heads, + head_dim=model_engine.model.config.head_dim, + tokens_per_block=4, + max_seq_len=256, + max_batch_size=8, + mapping=Mapping(world_size=1, tp_size=1, rank=0), + dtype=tensorrt_llm.bindings.DataType.HALF) + model_engine.runtime_draft_len = max_draft_len + stage = model_engine._step_input_stage + self.assertIsNotNone(stage) + resource_manager = ResourceManager( + {ResourceManagerType.KV_CACHE_MANAGER: kv_cache_manager}) + + # Three requests that ran in the previous step, in slots out of order. + slots = [5, 2, 7] + requests = kv_cache_manager.add_dummy_requests( + [1, 2, 3], + token_nums=[10, 17, 33], + is_gen=True, + max_num_draft_tokens=max_draft_len) + for request, slot in zip(requests, slots): + request.is_dummy_request = False + request.py_seq_slot = slot + batch = ScheduledRequests() + batch.generation_requests = requests + attn_metadata = model_engine._set_up_attn_metadata(kv_cache_manager) + self.assertIs(type(attn_metadata), TrtllmAttentionMetadata) + graph_metadata = attn_metadata.create_cuda_graph_metadata( + len(requests), False, max_draft_len) + spec_metadata = Mock( + _force_non_greedy_for_capture=False, + context_prompt_lookahead_tokens=None, + ) + + # The sampler's slot stores as the previous step left them. + num_slots = 8 + generator = torch.Generator(device="cuda").manual_seed(0) + previous = SampleStateTensorsSpec( + new_tokens=torch.randint(0, + 1000, (tokens_per_step, num_slots, 1), + generator=generator, + dtype=torch.int32, + device="cuda"), + new_tokens_lens=torch.randint(1, + tokens_per_step + 1, (num_slots, ), + generator=generator, + dtype=torch.int32, + device="cuda"), + next_draft_tokens=torch.randint(0, + 1000, (num_slots, max_draft_len), + generator=generator, + dtype=torch.int32, + device="cuda"), + ) + + def step_inputs(step_input_stage): + """The buffers one step writes, from a filler they all start at.""" + model_engine._step_input_stage = step_input_stage + for request, slot in zip(requests, slots): + request.py_batch_idx = slot + written = (model_engine.input_ids_cuda, + model_engine.position_ids_cuda, + model_engine.draft_tokens_cuda, + model_engine.previous_pos_id_offsets_cuda, + model_engine.previous_kv_lens_offsets_cuda, + graph_metadata.prompt_lens_cuda, + graph_metadata.kv_lens_cuda, + graph_metadata.kv_cache_block_offsets) + for buffer in written: + buffer.fill_(-3) + model_engine._prepare_tp_inputs(scheduled_requests=batch, + kv_cache_manager=kv_cache_manager, + attn_metadata=graph_metadata, + spec_metadata=spec_metadata, + new_tensors_device=previous, + resource_manager=resource_manager) + torch.cuda.synchronize() + return [buffer.clone() for buffer in written] + + torch_path = step_inputs(None) + with patch.object(stage, "commit", wraps=stage.commit) as commit, \ + patch.object(stage, "block_copy", + wraps=stage.block_copy) as block_copy: + staged = step_inputs(stage) + commit.assert_called_once() + self.assertEqual(block_copy.call_count, + 1 if use_kv_cache_manager_v2 else 0) + for expected, actual in zip(torch_path, staged): + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + # The previous rows' input ids came from the stores, by slot. + num_tokens = len(requests) * tokens_per_step + expected_input_ids = previous.new_tokens[:, slots, 0].t().reshape(-1) + self.assertEqual(staged[0][:num_tokens].tolist(), + expected_input_ids.tolist()) + kv_cache_manager.shutdown() + + def test_staged_spec_decode_graph_step_matches_torch_path(self) -> None: + self._check_staged_spec_decode_graph_step(use_kv_cache_manager_v2=False) + + def test_staged_spec_decode_graph_step_matches_torch_path_kv_v2( + self) -> None: + self._check_staged_spec_decode_graph_step(use_kv_cache_manager_v2=True) + def test_multimodal_encoder_max_seq_len(self) -> None: class CapturingEncoder(torch.nn.Module, MultimodalEncoderMixin): diff --git a/tests/unittest/_torch/helpers.py b/tests/unittest/_torch/helpers.py index c2adb2b9b68d..7350a10e0dda 100644 --- a/tests/unittest/_torch/helpers.py +++ b/tests/unittest/_torch/helpers.py @@ -236,9 +236,12 @@ def block_scale_gemm(mat_a: torch.Tensor, mat_scale_a: torch.Tensor, return results.view_as(x) -def create_mock_cuda_graph_runner(batch_size: int, - use_mrope: bool = False, - max_num_tokens: int = 1): +def create_mock_cuda_graph_runner( + batch_size: int, + use_mrope: bool = False, + max_num_tokens: int = 1, + static_input_ids: Optional[torch.Tensor] = None, + static_position_ids: Optional[torch.Tensor] = None): config = CUDAGraphRunnerConfig( use_cuda_graph=True, cuda_graph_padding_enabled=False, @@ -256,7 +259,9 @@ def create_mock_cuda_graph_runner(batch_size: int, is_encoder_decoder=False, mapping=Mapping(), dist=None, - kv_cache_manager_key=ResourceManagerType.KV_CACHE_MANAGER) + kv_cache_manager_key=ResourceManagerType.KV_CACHE_MANAGER, + static_input_ids=static_input_ids, + static_position_ids=static_position_ids) return CUDAGraphRunner(config) diff --git a/tests/unittest/_torch/speculative/hw_agnostic/test_rng_window_counter.py b/tests/unittest/_torch/speculative/hw_agnostic/test_rng_window_counter.py index ae735562317f..23bab3ac0aec 100644 --- a/tests/unittest/_torch/speculative/hw_agnostic/test_rng_window_counter.py +++ b/tests/unittest/_torch/speculative/hw_agnostic/test_rng_window_counter.py @@ -20,7 +20,9 @@ import types from typing import Optional -from tensorrt_llm._torch.speculative.interface import SpecMetadata +import torch + +from tensorrt_llm._torch.speculative.interface import DEFAULT_SAMPLING_SEED, SpecMetadata MAX_DRAFT_LEN = 3 WINDOW = MAX_DRAFT_LEN + 1 @@ -180,3 +182,37 @@ def test_graph_copy_shares_the_counters() -> None: assert _offsets(meta, [_request(1)]) == [0] assert _offsets(graph_meta, [_request(1)]) == [WINDOW] assert _offsets(meta, [_request(1)]) == [2 * WINDOW] + + +# --- all-greedy batches ------------------------------------------------------ + + +def _populated(meta: SpecMetadata, requests: list[types.SimpleNamespace]) -> list[int]: + """Run _populate_request_rng_state on CPU buffers and return request_offsets.""" + for request in requests: + request.sampling_config = types.SimpleNamespace(seed=request.seed) + normalized = [(0.0, 0, 1.0, 0.0, WINDOW) for _ in requests] + meta._populate_request_rng_state(requests, normalized) + return meta.request_offsets[: len(requests)].tolist() + + +def test_all_greedy_batch_skips_the_copies_but_advances_the_windows() -> None: + # The argmax graph reads none of the Philox buffers, so an all-greedy batch + # leaves them as they are; its windows are still taken, so the next sampled + # batch gets the same offsets it would have had. + meta = _meta() + sentinel = -5 + meta.temperatures = torch.ones(8 * WINDOW) + meta.request_seeds = torch.full((8,), sentinel, dtype=torch.int64) + meta.request_offsets = torch.full((8,), sentinel, dtype=torch.int64) + meta.seeds = torch.full((8 * WINDOW,), sentinel, dtype=torch.int64) + meta.offsets = torch.full((8 * WINDOW,), sentinel, dtype=torch.int64) + + meta.is_all_greedy_sample = True + assert _populated(meta, [_request(0, seed=7), _request(1)]) == [sentinel, sentinel] + assert meta.seeds.eq(sentinel).all() and meta.offsets.eq(sentinel).all() + + meta.is_all_greedy_sample = False + assert _populated(meta, [_request(0, seed=7), _request(1)]) == [WINDOW, WINDOW] + assert meta.request_seeds[:2].tolist() == [7, DEFAULT_SAMPLING_SEED] + assert meta.offsets[: 2 * WINDOW].tolist() == [WINDOW] * (2 * WINDOW) From f6cb23b6fec9022728bc2277cddfe781a8bf0579 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Thu, 1 Oct 2026 00:08:13 -0700 Subject: [PATCH 002/161] [None][feat] KDA verify: optional per-token states in the V2 hybrid cache manager MambaHybridCacheManagerV2 takes kda_token_states (default off). With the KDA replay caches it then also keeps the fp32 state after every draft of the last verify round per slot (kda_state_tok, [layers, slots, num_spec, H, V, K]), passes the layer's view in the speculative state and moves it with the slot. A verify kernel can start the next round from the state after the accepted drafts instead of replaying them. The cache factory enables it when the model's backbone sets kda_token_states. Signed-off-by: Vasanth Sabavat --- tensorrt_llm/_torch/pyexecutor/_util.py | 8 ++ .../kv_cache/mamba_cache_manager.py | 55 +++++++- .../kv_cache/test_mamba_cache_manager.py | 127 +++++++++++++++++- 3 files changed, 185 insertions(+), 5 deletions(-) diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index f830d3c23191..2028050aa034 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -3048,6 +3048,14 @@ def _create_kv_cache_manager( if is_kda_mtp_verify_available(): kda_extra_kwargs["kda_replay_num_spec"] = ( spec_config.tokens_per_gen_step - 1) + # A model whose KDA verify starts from the state after the + # accepted drafts (instead of replaying them) asks for the + # state after every draft with `kda_token_states`. + backbone = getattr(getattr(model_engine, "model", None), + "model", None) + if (issubclass(kv_cache_manager_cls, MambaHybridCacheManagerV2) + and getattr(backbone, "kda_token_states", False)): + kda_extra_kwargs["kda_token_states"] = True if is_glm5_next: # The manager places an indexer buffer on every attention layer. from ..attention.backends.sparse.glm_kpool import \ diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache/mamba_cache_manager.py b/tensorrt_llm/_torch/pyexecutor/kv_cache/mamba_cache_manager.py index ad417618950b..ed5cef76dd35 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache/mamba_cache_manager.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache/mamba_cache_manager.py @@ -330,6 +330,14 @@ def use_kda_replay_update(self) -> bool: """Whether KDA fused verification owns per-slot replay caches.""" return getattr(self, "_use_kda_replay_update", False) + @property + def keeps_kda_token_states(self) -> bool: + """Whether the KDA replay caches come with the state after every + draft of the last verify round (``kda_state_tok``), so that a verify + kernel can start the next round from the state after the accepted + drafts instead of replaying them.""" + return False + @abstractmethod def get_conv_states(self, layer_idx: int) -> torch.Tensor: """Return conv states for specific layer. @@ -438,6 +446,14 @@ class SpeculativeState(State): kda_qkg_cache: torch.Tensor | None = None kda_v_cache: torch.Tensor | None = None kda_beta_cache: torch.Tensor | None = None + # Optional: the state after every draft of the last verify round, + # [slots, num_spec, H, V, K] fp32 (the SSM pool holds the state + # after the round's first, non-draft token). A verify kernel starts + # the next round from entry n - 1 when the round accepted n > 0 + # drafts (prev_num_accepted_tokens) instead of replaying them from + # the caches above; a slot reset for a new request has n = 0, so a + # previous owner's entries are never read. + kda_state_tok: torch.Tensor | None = None # Replay path: compact double-buffered cache # prev_num_accepted_tokens: # accepted tokens (always >= 1 if drafting). @@ -3024,6 +3040,7 @@ def __init__( conv_state_layout: Literal["x_b_c", "q_k_v"] = "x_b_c", kda_replay_num_spec: Optional[int] = None, qwen4_exp_ple_cache_params: Optional["Qwen4ExpPLECacheParams"] = None, + kda_token_states: bool = False, **kwargs, ) -> None: if conv_state_layout not in ("x_b_c", "q_k_v"): @@ -3040,6 +3057,9 @@ def __init__( self._kda_replay_num_spec = kda_replay_num_spec self._use_kda_replay_update = kda_replay_num_spec is not None + # Only with the KDA replay caches: also keep the state after every + # verify token (kda_state_tok). + self._kda_token_states = kda_token_states and self._use_kda_replay_update if self._use_kda_replay_update: if use_replay_state_update: raise ValueError( @@ -3420,6 +3440,10 @@ def get_disagg_role_layouts(self) -> Dict[DataRole, RoleLayout]: def use_gdn_cached_replay_all_layer_commit(self) -> bool: return getattr(self, "_use_gdn_cached_replay_all_layer_commit", False) + @property + def keeps_kda_token_states(self) -> bool: + return getattr(self, "kda_state_tok", None) is not None + def _setup_mtp_intermediate_states(self, spec_config, max_batch_size: int) -> None: if not self.use_kda_replay_update: @@ -3450,6 +3474,7 @@ def _allocate_pool_replay_buffers( self.kda_qkg_cache = None self.kda_v_cache = None self.kda_beta_cache = None + self.kda_state_tok = None if (not self.use_kda_replay_update or self.local_num_mamba_layers == 0): return allocated @@ -3504,8 +3529,17 @@ def allocate_dim_contiguous_conv_cache() -> torch.Tensor: self.ssm_state_shape[0]), device, ) + if self._kda_token_states: + self.kda_state_tok = torch.zeros( + (self.local_num_mamba_layers, cache_size, num_spec, + *self.ssm_state_shape), + dtype=torch.float32, + device=device, + ) + per_token = (", with per-token states" + if self.kda_state_tok is not None else "") logger.info("Mamba Cache (kda-replay) is allocated for " - f"{cache_size} state slots") + f"{cache_size} state slots{per_token}") return True def _commit_gdn_cached_replay_history_layers( @@ -3801,6 +3835,12 @@ def _max_resident_sequences(self) -> int: def _mamba_state_bytes_per_slot(self) -> int: base_bytes = self.local_num_mamba_layers * (self.ssm_bytes + self.conv_bytes) + if getattr(self, "_kda_token_states", False): + # kda_state_tok: an fp32 SSM state per draft of every slot. + base_bytes += (self.local_num_mamba_layers * + self._kda_replay_num_spec * + math.prod(self.ssm_state_shape) * + torch.float32.itemsize) local_ple_layers = sum(layer_id in self.pp_layers for layer_id in self._ple_layer_ids) if local_ple_layers == 0: @@ -4285,16 +4325,20 @@ def _relocate_kda_replay_slots(self, old_slots: List[int], destination_slots = torch.tensor([new for _, new in moves], dtype=torch.long, device=device) - replay_buffers = ( + replay_buffers = [ self.kda_conv_q, self.kda_conv_k, self.kda_conv_v, self.kda_qkg_cache, self.kda_v_cache, self.kda_beta_cache, - ) + ] + assert all(replay_buffer is not None + for replay_buffer in replay_buffers) + # The per-token verify states belong to the slot as well. + if self.kda_state_tok is not None: + replay_buffers.append(self.kda_state_tok) for replay_buffer in replay_buffers: - assert replay_buffer is not None replay_buffer.index_copy_( 1, destination_slots, @@ -4498,6 +4542,8 @@ def mamba_layer_cache( "kda_v_cache": self.kda_v_cache[layer_offset], "kda_beta_cache": self.kda_beta_cache[layer_offset], } + if self.kda_state_tok is not None: + spec_kwargs["kda_state_tok"] = self.kda_state_tok[layer_offset] if self.mamba_ssm_rand_seed is not None: spec_kwargs["mamba_ssm_rand_seed"] = self.mamba_ssm_rand_seed return PythonMambaCacheManager.SpeculativeState( @@ -4781,4 +4827,5 @@ def shutdown(self): self.kda_qkg_cache = None self.kda_v_cache = None self.kda_beta_cache = None + self.kda_state_tok = None super().shutdown() diff --git a/tests/unittest/_torch/executor/kv_cache/test_mamba_cache_manager.py b/tests/unittest/_torch/executor/kv_cache/test_mamba_cache_manager.py index 1f920c7b718c..0079ea8632a2 100644 --- a/tests/unittest/_torch/executor/kv_cache/test_mamba_cache_manager.py +++ b/tests/unittest/_torch/executor/kv_cache/test_mamba_cache_manager.py @@ -2,6 +2,7 @@ # SPDX-License-Identifier: Apache-2.0 """Regression tests for Python, Cpp, and V2 Mamba cache managers.""" +import math import os from types import SimpleNamespace from unittest.mock import MagicMock @@ -300,12 +301,16 @@ def _kimi_model_config() -> SimpleNamespace: quant_config=None, sparse_attention_config=None, get_num_mamba_layers=lambda: 2, + # get_kv_cache_manager_cls reads the mapping and spec config (helix check). + mapping=None, + spec_config=None, ) def _capture_kimi_v2_manager_ctor( monkeypatch: pytest.MonkeyPatch, spec_config=None, + model_engine=None, ) -> tuple[tuple, dict]: """Route a Kimi config through _create_kv_cache_manager with an explicit V2 manager and capture the constructor arguments.""" @@ -324,7 +329,7 @@ def __init__(self, *args: object, **kwargs: object) -> None: assert get_kv_cache_manager_cls(model_config, kv_cache_config) is MambaHybridCacheManagerV2 _create_kv_cache_manager( - model_engine=None, + model_engine=model_engine, kv_cache_manager_cls=RecordingV2Manager, mapping=Mapping(world_size=1, tp_size=1, pp_size=1), kv_cache_config=kv_cache_config, @@ -380,6 +385,33 @@ def test_kimi_explicit_v2_manager_enables_kda_replay( assert kwargs["kda_replay_num_spec"] == spec_config.tokens_per_gen_step - 1 assert kwargs["conv_state_layout"] == "q_k_v" + assert "kda_token_states" not in kwargs + + +@pytest.mark.parametrize("backbone_asks", [True, False]) +def test_kimi_explicit_v2_manager_kda_token_states_follow_the_model( + monkeypatch: pytest.MonkeyPatch, + backbone_asks: bool, +) -> None: + """The factory asks the V2 manager for per-token KDA verify states only when + the model's backbone sets ``kda_token_states``.""" + monkeypatch.setattr( + "tensorrt_llm._torch.modules.kimi_kda._kda_kernels.is_kda_mtp_verify_available", + lambda: True, + ) + backbone = SimpleNamespace(kda_token_states=True) if backbone_asks else SimpleNamespace() + model_engine = SimpleNamespace( + model=SimpleNamespace(model=backbone), + _max_cuda_graph_batch_size=4, + is_multimodal=True, + ) + + _, kwargs = _capture_kimi_v2_manager_ctor( + monkeypatch, MTPDecodingConfig(max_draft_len=3), model_engine + ) + + assert kwargs["kda_replay_num_spec"] == 3 + assert kwargs.get("kda_token_states", False) is backbone_asks def test_kimi_explicit_v2_manager_uses_qkv_convolution_layout( @@ -2353,6 +2385,7 @@ def _build_v2_hybrid_with_mamba_layer( mamba_n_groups=1, mamba_ssm_cache_dtype=torch.float16, kda_replay_num_spec=None, + kda_token_states=False, ): """Construct a real MambaHybridCacheManagerV2.""" mamba_mask = [True] * num_mamba_layers + [False] * num_attention_layers @@ -2407,6 +2440,7 @@ def _build_v2_hybrid_with_mamba_layer( dtype=dtype, conv_state_layout=conv_state_layout, kda_replay_num_spec=kda_replay_num_spec, + kda_token_states=kda_token_states, ) @@ -3629,6 +3663,43 @@ def test_v2_kda_replay_allocates_logical_slot_caches(): ) assert layer_cache.intermediate_ssm is None assert layer_cache.intermediate_conv_window is None + assert layer_cache.kda_state_tok is None + assert not mgr.keeps_kda_token_states + finally: + mgr.shutdown() + + +@skip_no_cuda +@pytest.mark.parametrize("kda_replay_num_spec", [2, None]) +def test_v2_kda_token_states_allocated_with_the_replay_caches(kda_replay_num_spec): + """With the replay caches, ``kda_token_states`` adds the fp32 state after every + draft of each slot, per layer; without them it allocates nothing.""" + mgr = _build_v2_hybrid_with_mamba_layer( + max_batch_size=4, + num_mamba_layers=2, + spec_config=MTPDecodingConfig(max_draft_len=2), + conv_state_layout="q_k_v", + mamba_d_conv=5, + mamba_num_heads=6, + mamba_n_groups=6, + mamba_ssm_cache_dtype=torch.float32, + kda_replay_num_spec=kda_replay_num_spec, + kda_token_states=True, + ) + try: + if kda_replay_num_spec is None: + assert getattr(mgr, "kda_state_tok", None) is None + assert not mgr.keeps_kda_token_states + return + assert mgr.keeps_kda_token_states + for layer_idx in range(2): + layer_cache = mgr.mamba_layer_cache(layer_idx) + cache_size = layer_cache.temporal.shape[0] + states = layer_cache.kda_state_tok + assert states.shape == (cache_size, 2, *layer_cache.temporal.shape[1:]) + assert states.dtype is torch.float32 + assert states.data_ptr() == mgr.kda_state_tok[layer_idx].data_ptr() + assert not states.any() finally: mgr.shutdown() @@ -3785,6 +3856,60 @@ def test_v2_kda_replay_relocates_live_slot_history(): mgr.shutdown() +@skip_no_cuda +def test_v2_kda_token_states_relocate_with_their_slot(): + mgr = _build_v2_hybrid_with_mamba_layer( + spec_config=MTPDecodingConfig(max_draft_len=2), + conv_state_layout="q_k_v", + mamba_n_groups=4, + mamba_ssm_cache_dtype=torch.float32, + kda_replay_num_spec=2, + kda_token_states=True, + ) + try: + states = mgr.kda_state_tok + assert states is not None + states.zero_() + states[:, 0].fill_(1.0) + states[:, 1].fill_(2.0) + source_zero, source_one = states[:, 0].clone(), states[:, 1].clone() + + mgr._relocate_kda_replay_slots([0, 1], [1, 2]) + + torch.testing.assert_close(states[:, 1], source_zero, rtol=0, atol=0) + torch.testing.assert_close(states[:, 2], source_one, rtol=0, atol=0) + finally: + mgr.shutdown() + # Released with the other replay buffers. + assert mgr.kda_state_tok is None + + +@skip_no_cuda +def test_v2_kda_token_states_count_in_the_per_slot_budget(): + """The capacity math sees the per-token states: a slot costs one more fp32 SSM state per draft and layer.""" + kwargs = dict( + num_mamba_layers=2, + spec_config=MTPDecodingConfig(max_draft_len=2), + conv_state_layout="q_k_v", + mamba_d_conv=5, + mamba_num_heads=6, + mamba_n_groups=6, + mamba_ssm_cache_dtype=torch.float32, + kda_replay_num_spec=2, + ) + plain = _build_v2_hybrid_with_mamba_layer(**kwargs) + try: + plain_bytes = plain._mamba_state_bytes_per_slot() + finally: + plain.shutdown() + with_states = _build_v2_hybrid_with_mamba_layer(kda_token_states=True, **kwargs) + try: + per_token = 2 * 2 * math.prod(with_states.ssm_state_shape) * 4 + assert with_states._mamba_state_bytes_per_slot() == plain_bytes + per_token + finally: + with_states.shutdown() + + @skip_no_cuda def test_v2_kda_state_index_setup_relocates_generation_history(): class StateCache: From 70ae50450a9c7f15e044bff51dc8e41a6135ffb2 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Thu, 1 Oct 2026 00:52:42 -0700 Subject: [PATCH 003/161] [None][feat] One-model speculative decoding: the worker produces the target logits SpecDecOneEngineForCausalLM asks its spec worker for the target logits (SpecWorkerBase.target_logits) instead of calling the logits processor itself. The default is the logits processor's fp32 [rows, vocab], so no worker changes behaviour. A worker whose acceptance reads another layout can override it; when its forward then returns this TP rank's vocabulary shard as "logits", it sets "logits_vocab_shard", and the engine all-gathers the logits before running logits post-processors (only when a scheduled request has one). Signed-off-by: Vasanth Sabavat --- .../_torch/models/modeling_speculative.py | 8 +- .../_torch/pyexecutor/model_engine.py | 24 +++ tensorrt_llm/_torch/speculative/interface.py | 19 ++ .../test_spec_worker_target_logits.py | 179 ++++++++++++++++++ 4 files changed, 227 insertions(+), 3 deletions(-) create mode 100644 tests/unittest/_torch/speculative/hw_agnostic/test_spec_worker_target_logits.py diff --git a/tensorrt_llm/_torch/models/modeling_speculative.py b/tensorrt_llm/_torch/models/modeling_speculative.py index 38788d8f06b4..ba54c65d95d1 100644 --- a/tensorrt_llm/_torch/models/modeling_speculative.py +++ b/tensorrt_llm/_torch/models/modeling_speculative.py @@ -1840,12 +1840,14 @@ def forward( hidden_states = hidden_states[:attn_metadata.num_tokens] if self.spec_worker is not None: - # get logits - logits = self.logits_processor.forward( + # The target logits, in the layout the worker's acceptance reads. + logits = self.spec_worker.target_logits( hidden_states[spec_metadata.gather_ids], self.lm_head, + self.logits_processor, attn_metadata, - True, + spec_metadata, + self.draft_model, ) # VLM wrappers (e.g. Qwen3VLModelBase) replace input_ids with diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 2e93692a3947..9a8ab34b02d4 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -6379,6 +6379,30 @@ def _execute_logit_post_processors(self, return logits_tensor = outputs["logits"] + if outputs.get("logits_vocab_shard", False): + # The spec worker kept the logits vocabulary-sharded (this TP + # rank's columns). A post-processor sees the whole vocabulary, so + # gather them first, only when a request has one. Every TP rank + # schedules the same requests, so all of them take the gather; + # under attention DP they would not. + if self.mapping.enable_attention_dp: + raise RuntimeError( + "vocabulary-sharded spec-worker logits need every TP rank " + "to schedule the same requests, which attention DP does " + "not") + if not any( + getattr(request, "py_logits_post_processors", None) + for request in scheduled_requests.all_requests()): + return + from ..distributed import allgather + logits_tensor = allgather(logits_tensor, self.mapping, + dim=-1).float() + # Drop the vocabulary's TP padding, as the LM head's own gather + # does. + vocab_size = getattr(getattr(self.model, "lm_head", None), + "num_embeddings", None) + if vocab_size is not None: + logits_tensor = logits_tensor[..., :vocab_size] logits_row_offset = 0 request_groups = ( diff --git a/tensorrt_llm/_torch/speculative/interface.py b/tensorrt_llm/_torch/speculative/interface.py index 37992e2f8835..53c0913592be 100644 --- a/tensorrt_llm/_torch/speculative/interface.py +++ b/tensorrt_llm/_torch/speculative/interface.py @@ -1648,6 +1648,25 @@ def forward(self, *args, **kwargs): def _forward_impl(self, *args, **kwargs): """Worker-specific forward logic, called by SpecWorkerBase.forward.""" + def target_logits(self, hidden_states: torch.Tensor, lm_head, + logits_processor, attn_metadata, spec_metadata, + draft_model) -> torch.Tensor: + """The target logits of the gathered rows ``hidden_states``, which + the model hands to ``forward`` as ``logits``: fp32 [rows, vocab]. + + A worker whose acceptance reads another layout overrides this. If its + ``forward`` then returns this TP rank's vocabulary shard as + ``"logits"``, it also returns ``"logits_vocab_shard": True``, and the + engine gathers the logits (on every TP rank together, so not under + attention DP) before any logits post-processor sees them; the + post-processors get a read-only view, since the acceptance already + ran inside ``forward``. Such a worker applies a guided decoder's + bitmask to its shard at the shard's column offset, or rejects guided + decoding. + """ + return logits_processor.forward(hidden_states, lm_head, attn_metadata, + True) + def _ensure_spec_dec_state_restored(self, attn_metadata, spec_metadata): """Restore attn-metadata spec-dec state if a failure skipped it. diff --git a/tests/unittest/_torch/speculative/hw_agnostic/test_spec_worker_target_logits.py b/tests/unittest/_torch/speculative/hw_agnostic/test_spec_worker_target_logits.py new file mode 100644 index 000000000000..c3a541a91654 --- /dev/null +++ b/tests/unittest/_torch/speculative/hw_agnostic/test_spec_worker_target_logits.py @@ -0,0 +1,179 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""The one-model spec worker's target-logits hook and the engine gather of vocabulary-sharded logits.""" + +from types import SimpleNamespace + +import pytest +import torch + +import tensorrt_llm._torch.distributed as distributed +from tensorrt_llm._torch.models.modeling_speculative import SpecDecOneEngineForCausalLM +from tensorrt_llm._torch.pyexecutor.model_engine import PyTorchModelEngine +from tensorrt_llm._torch.speculative.interface import SpecWorkerBase + +pytestmark = pytest.mark.cpu_only + + +class _RecordingLogitsProcessor: + def __init__(self, logits): + self.logits = logits + self.calls = [] + + def forward(self, *args): + self.calls.append(args) + return self.logits + + +def test_default_target_logits_are_the_logits_processors(): + logits = torch.randn(3, 16) + processor = _RecordingLogitsProcessor(logits) + hidden, lm_head, attn_metadata = torch.randn(3, 8), object(), object() + + out = SpecWorkerBase.target_logits( + object(), hidden, lm_head, processor, attn_metadata, object(), object() + ) + + assert out is logits + assert len(processor.calls) == 1 + args = processor.calls[0] + assert args[0] is hidden and args[1] is lm_head and args[2] is attn_metadata + assert args[3] is True + + +def _generation_request(post_processors, tokens=(1, 2, 3)): + return SimpleNamespace( + py_request_id=7, + py_logits_post_processors=post_processors, + py_beam_width=1, + get_beam_width_by_iter=lambda for_next_iteration: 1, + get_tokens=lambda beam: list(tokens), + ) + + +def _engine(vocab_size=None, attention_dp=False): + engine = object.__new__(PyTorchModelEngine) + engine.mapping = SimpleNamespace(is_last_pp_rank=lambda: True, enable_attention_dp=attention_dp) + engine.model = SimpleNamespace(lm_head=SimpleNamespace(num_embeddings=vocab_size)) + return engine + + +def _scheduled(*requests): + return SimpleNamespace( + context_requests=[], generation_requests=list(requests), all_requests=lambda: list(requests) + ) + + +def _recording_allgather(monkeypatch, vocab): + calls = [] + + def allgather(tensor, mapping, dim=-1): + calls.append((tensor, mapping, dim)) + return torch.zeros(tensor.shape[0], vocab, dtype=torch.bfloat16) + + monkeypatch.setattr(distributed, "allgather", allgather) + return calls + + +def test_sharded_logits_are_gathered_for_post_processors(monkeypatch): + calls = _recording_allgather(monkeypatch, vocab=32) + seen = [] + request = _generation_request( + [lambda req_id, rows, tokens, s, c: seen.append(tuple(rows.shape))] + ) + shard = torch.ones(1, 8, dtype=torch.bfloat16) + engine = _engine() + + engine._execute_logit_post_processors( + _scheduled(request), {"logits": shard, "logits_vocab_shard": True} + ) + + assert len(calls) == 1 + assert calls[0][0] is shard and calls[0][1] is engine.mapping and calls[0][2] == -1 + assert seen == [(1, 1, 32)] # the post-processor sees the whole vocabulary + + +def test_sharded_logits_without_post_processors_are_not_gathered(monkeypatch): + calls = _recording_allgather(monkeypatch, vocab=32) + + _engine()._execute_logit_post_processors( + _scheduled(_generation_request(None)), + {"logits": torch.ones(1, 8, dtype=torch.bfloat16), "logits_vocab_shard": True}, + ) + + assert calls == [] + + +def test_full_logits_are_never_gathered(monkeypatch): + calls = _recording_allgather(monkeypatch, vocab=32) + seen = [] + request = _generation_request( + [lambda req_id, rows, tokens, s, c: seen.append(tuple(rows.shape))] + ) + + _engine()._execute_logit_post_processors(_scheduled(request), {"logits": torch.ones(1, 32)}) + + assert calls == [] + assert seen == [(1, 1, 32)] + + +def test_sharded_logits_drop_the_vocab_padding(monkeypatch): + # 4 ranks x 8 columns hold a 30-token vocabulary: the gather's last 2 columns are TP padding. + _recording_allgather(monkeypatch, vocab=32) + seen = [] + request = _generation_request( + [lambda req_id, rows, tokens, s, c: seen.append(tuple(rows.shape))] + ) + + _engine(vocab_size=30)._execute_logit_post_processors( + _scheduled(request), + {"logits": torch.ones(1, 8, dtype=torch.bfloat16), "logits_vocab_shard": True}, + ) + + assert seen == [(1, 1, 30)] + + +def test_sharded_logits_under_attention_dp_raise(monkeypatch): + calls = _recording_allgather(monkeypatch, vocab=32) + request = _generation_request([lambda *args: None]) + + with pytest.raises(RuntimeError, match="attention DP"): + _engine(attention_dp=True)._execute_logit_post_processors( + _scheduled(request), + {"logits": torch.ones(1, 8, dtype=torch.bfloat16), "logits_vocab_shard": True}, + ) + assert calls == [] + + +def test_the_spec_model_hands_the_workers_target_logits_to_its_forward(): + hidden = torch.randn(4, 8) + gather_ids = torch.tensor([1, 3]) + calls = {} + + class Worker: + def target_logits(self, *args): + calls["target_logits"] = args + return "worker logits" + + def __call__(self, **kwargs): + calls["forward"] = kwargs + return "worker outputs" + + model = object.__new__(SpecDecOneEngineForCausalLM) + torch.nn.Module.__init__(model) + model.model = lambda **kwargs: hidden + model.layer_idx = -1 + model.spec_worker = Worker() + model.lm_head, model.logits_processor, model.draft_model = "lm_head", "processor", "drafter" + spec_metadata = SimpleNamespace(is_layer_capture=lambda layer: False, gather_ids=gather_ids) + attn_metadata = SimpleNamespace(padded_num_tokens=None, num_tokens=4) + + out = SpecDecOneEngineForCausalLM.forward( + model, attn_metadata, input_ids=torch.arange(4), spec_metadata=spec_metadata + ) + + assert out == "worker outputs" + rows, *rest = calls["target_logits"] + assert torch.equal(rows, hidden[gather_ids]) + assert rest == ["lm_head", "processor", attn_metadata, spec_metadata, "drafter"] + assert calls["forward"]["logits"] == "worker logits" From b44c537918a606c01f5187dfae94aeb5d1fbd2c6 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 15:41:56 -0700 Subject: [PATCH 004/161] [None][fix] spec_step_copies: store zeros past a row's accepted tokens The one-model worker's acceptance returns new_tokens as torch.empty [N, K + 1]. A context row writes only column 0 and accepts one token; a generation row writes every column. The scatter kernel copied every column of every row into the sampler's new_tokens store, so for context rows it read columns 1..K, which were never written; compute-sanitizer initcheck reports those reads. update_requests reads the store only up to new_tokens_lens, so no output changed. The scatter kernel now reads a row's new tokens only below its new_tokens_lens and stores zeros past them. A thread loads the row's length once; the column 0 thread's store_lens write reuses that load. The sampler's torch store update, used where the kernel does not run, is unchanged: nothing reads past new_tokens_lens. test_spec_step_copies: the reference zeroes the same columns, and the Step cases draw accepted lengths from 0 to past the output width. test_scatter_zeros_past_accepted_length covers a mixed batch laid out as the worker writes it, with a skipped chunked-context row. "poisoned" fills the never-written columns with a sentinel and fails on the old kernel; "never_written" leaves them unwritten for compute-sanitizer initcheck. Signed-off-by: Vasanth Sabavat --- .../cute_dsl_kernels/spec_step_copies/op.py | 3 +- .../spec_step_copies_kernel.py | 10 ++- .../cute_dsl_kernels/test_spec_step_copies.py | 64 ++++++++++++++++++- 3 files changed, 71 insertions(+), 6 deletions(-) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/spec_step_copies/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/spec_step_copies/op.py index c0b2e56c3658..9bc81fef63bb 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/spec_step_copies/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/spec_step_copies/op.py @@ -94,7 +94,8 @@ class SlotScatter: """The speculative sampler's store update as one kernel. ``store[..., slots[r]] = outputs[row_begin + r]`` for every row ``r``, each row padded with zeros or cut to its - store's width. + store's width. A row's ``new_tokens`` at or past its ``new_tokens_lens`` are stored as zeros: the forward writes + only column 0 of a context row, and readers of the store stop at that length. """ _kernel: ClassVar[Any] = None diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/spec_step_copies/spec_step_copies_kernel.py b/tensorrt_llm/_torch/cute_dsl_kernels/spec_step_copies/spec_step_copies_kernel.py index 9780847687b4..269da57f8d01 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/spec_step_copies/spec_step_copies_kernel.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/spec_step_copies/spec_step_copies_kernel.py @@ -15,7 +15,8 @@ """The one-model speculative decoding step's eager copy passes, one kernel each. * ``scatter_kernel``: the sampler moves the forward's per-row outputs into its slot-indexed stores (``index_copy_`` - of new tokens, next new tokens, accepted lengths and next draft tokens, each padded or cut to its store width). + of new tokens, next new tokens, accepted lengths and next draft tokens, each padded or cut to its store width; + a row's new tokens past its accepted length are written as zeros). * ``stage_kernel``: one launch writes a decode step's per-step inputs: the overlap scheduler's gathers from those stores for the rows whose request ran in the previous step (``index_select`` into input ids and draft tokens by slot, into the position offsets by the per-token index list, and into the KV-length offsets by slot, minus the @@ -76,11 +77,14 @@ def scatter_kernel( col = t - row * columns src_row = row_begin + row slot = _i32_at(slots, row)[0] + accepted = _i32_at(out_lens, src_row)[0] zero = cutlass.Int32(0) if col < new_width: + # Only the row's accepted tokens are read: past them a context row's forward output is never written. value = zero if col < out_new_width: - value = _i32_at(out_new_tokens, src_row * out_new_width + col)[0] + if col < accepted: + value = _i32_at(out_new_tokens, src_row * out_new_width + col)[0] _i32_at(store_new_tokens, col * num_slots + slot)[0] = value if col < next_width: value = zero @@ -93,7 +97,7 @@ def scatter_kernel( value = _i32_at(out_next_draft, src_row * out_draft_width + col)[0] _i32_at(store_next_draft, slot * draft_width + col)[0] = value if col == zero: - _i32_at(store_lens, slot)[0] = _i32_at(out_lens, src_row)[0] + _i32_at(store_lens, slot)[0] = accepted @cute.jit diff --git a/tests/unittest/_torch/cute_dsl_kernels/test_spec_step_copies.py b/tests/unittest/_torch/cute_dsl_kernels/test_spec_step_copies.py index 10408b60106d..a28ad5b659b1 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/test_spec_step_copies.py +++ b/tests/unittest/_torch/cute_dsl_kernels/test_spec_step_copies.py @@ -18,8 +18,11 @@ 8 or 16 slots; draft-token buffer 56, KV-length offsets 8, per-token buffers max_num_tokens). * ``SlotScatter.scatter`` vs SpecSampler's torch store update (each output padded with zeros or cut to its store's - width, then four ``index_copy_`` by slot): outputs narrower than (dynamic draft length), as wide as and wider than the - stores, and mixed per field; row_begin 0 and > 0; shuffled distinct slot tables. + width, the new tokens at or past a row's accepted length zeroed, then four ``index_copy_`` by slot): outputs + narrower than (dynamic draft length), as wide as and wider than the stores, and mixed per field; row_begin 0 and + > 0; shuffled distinct slot tables; accepted lengths from 0 to past the output width. A mixed context / generation + batch laid out as the one-model worker writes it (a context row's tokens past column 0 never written) stores zeros + there. * ``StepInputStage.gather`` committed on its own vs PyTorchModelEngine._prepare_tp_inputs's torch overlap gathers (the stores by slot into the input ids and draft tokens, the lengths by the per-token index list into the position offsets, ``new_tokens_lens - T`` by slot into the KV-length offsets): the engine's offsets (the requests without a @@ -135,6 +138,8 @@ def fit(t, width): # pad with zeros or truncate to the store width return t[:, :width] o_new_tokens = fit(o_new_tokens, store_new_tokens.shape[0]) + columns = torch.arange(o_new_tokens.shape[1], device=o_new_tokens.device) + o_new_tokens = torch.where(columns < o_new_tokens_lens[:, None], o_new_tokens, 0) o_next_draft_tokens = fit(o_next_draft_tokens, store_next_draft_tokens.shape[1]) o_next_new_tokens = fit(o_next_new_tokens, store_next_new_tokens.shape[0]) store_new_tokens.squeeze(-1).T.index_copy_(0, slots, o_new_tokens) @@ -309,6 +314,8 @@ def rewrite(self) -> None: for t in b.values(): if t.is_cuda: t.copy_(rand_i32(t.shape, g)) + # Accepted lengths from none to past the new-token width, so every row cuts its new tokens somewhere. + b["out_lens"].copy_(rand_i32(b["out_lens"].shape, g, 0, self.widths[0] + 2)) slots = shuffled_distinct(rows, num_slots, g) table = rand_i32((num_slots,), g, 0, num_slots) table[:rows] = torch.tensor(slots, dtype=torch.int32) @@ -638,6 +645,59 @@ def test_graph_replay(tokens): assert not bad, bad +POISON = 0x5EEDF00D # what the never-written columns hold in the poisoned case + + +@pytest.mark.parametrize("poison", [True, False], ids=["poisoned", "never_written"]) +def test_scatter_zeros_past_accepted_length(poison): + """A mixed batch as the one-model worker's acceptance leaves its outputs (``new_tokens`` is ``torch.empty`` + [N, K + 1]; a context row writes its first token only and accepts 1, a generation row writes every column): the + store holds each row's accepted tokens, then zeros, never a never-written column. ``poisoned``: those columns hold + POISON. ``never_written``: they keep the allocation's contents, so compute-sanitizer initcheck reports any read of + them. Row 0 is a context row whose chunk is not its last (row_begin 1 skips it).""" + op = _op() + g = torch.Generator().manual_seed(49) + skipped, num_contexts, accepted = 1, 3, [1, 3, MAX_DRAFT + 1, 5] + n = skipped + num_contexts + len(accepted) + rows, row_begin, num_slots = n - skipped, skipped, 2 * MAX_BATCH + width = MAX_DRAFT + 1 + contexts = skipped + num_contexts + tokens = rand_i32((n, width), g, 0, 1 << 20) + lens = torch.tensor([1] * contexts + accepted, dtype=torch.int32) + new_tokens = torch.empty((n, width), dtype=torch.int32, device="cuda") + if poison: + new_tokens.fill_(POISON) + new_tokens[:contexts, 0] = tokens[:contexts, 0].cuda() + new_tokens[contexts:] = tokens[contexts:].cuda() + outputs = { + "new_tokens": new_tokens, + "new_tokens_lens": lens.cuda(), + "next_new_tokens": rand_i32((n, width), g).cuda(), + "next_draft_tokens": rand_i32((n, MAX_DRAFT), g).cuda(), + } + slots = shuffled_distinct(rows, num_slots, g) + slot_table = torch.tensor(slots + [0] * (num_slots - rows), dtype=torch.int32, device="cuda") + stores = [ + torch.full(shape, -1, dtype=torch.int32, device="cuda") + for shape in ( + (STORE_WIDTH, num_slots, 1), + (STORE_WIDTH, num_slots, 1), + (num_slots,), + (num_slots, MAX_DRAFT), + ) + ] + assert op.SlotScatter().scatter(outputs, row_begin, rows, slot_table, *stores) + store_new, store_lens = stores[0][:, :, 0].T.cpu(), stores[2].cpu() + want = torch.full((num_slots, STORE_WIDTH), -1, dtype=torch.int32) + for r, s in enumerate(slots): + src = row_begin + r + want[s] = 0 + want[s, : lens[src]] = tokens[src, : lens[src]] + assert store_lens[s] == lens[src], (r, s) + assert not (store_new == POISON).any(), "a never-written column reached the store" + assert torch.equal(store_new, want), (store_new, want) + + def test_contract(): """rows == 0 launches nothing; every call the kernels do not cover returns False and launches nothing, so the caller keeps its torch path; StepInputStage's per-step limits (max_copies = 4 copies, max_block_copies = 2 block From 92d67f38d78dde4f756a661e9cc919ad2ed1bfed Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 15:42:23 -0700 Subject: [PATCH 005/161] [None][test] spec_step_copies tests: side stream waits for the current stream measure_graph (test_graph_replay) and time_graph build their inputs on the current stream, then make their first eager call on a fresh side stream without ordering it after that work. Under the stream-ordered cudaMallocAsync allocator the call can write memory before its allocation in the side stream's order; with any allocator it can read inputs whose initialization has not finished. The side stream now waits for the current stream first. Signed-off-by: Vasanth Sabavat --- tests/unittest/_torch/cute_dsl_kernels/test_spec_step_copies.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/unittest/_torch/cute_dsl_kernels/test_spec_step_copies.py b/tests/unittest/_torch/cute_dsl_kernels/test_spec_step_copies.py index a28ad5b659b1..a2a44221d1da 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/test_spec_step_copies.py +++ b/tests/unittest/_torch/cute_dsl_kernels/test_spec_step_copies.py @@ -562,6 +562,7 @@ def run_torch(st, b): st.torch_stage(b) stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) with torch.cuda.stream(stream): for st in steps: # the kernels compile on their first call, outside capture run_kernels(st, clone_bufs(st.bufs), stage) @@ -798,6 +799,7 @@ def time_graph(body, calls, replays=15): """Per-call us of a CUDA graph of ``calls`` back-to-back ``body(i)``: median (min, max) over ``replays`` replays. ``body(-1)`` runs first, outside capture.""" stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) with torch.cuda.stream(stream): body(-1) torch.cuda.synchronize() From f3cf23f5fd63b007e3a1b0ffabdebec8dfa2719a Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 15:44:12 -0700 Subject: [PATCH 006/161] [None][fix] KV cache manager V2: keep the SSM pool at its live floor with per-token KDA states With per-token KDA states, kda_state_tok holds num_spec recurrent states per SSM state slot, outside the cache quota, addressed by the request's SSM slot (_allocate_pool_replay_buffers). V2 sizes the SSM pool group by the typical step's ratio, so its slot count, and kda_state_tok with it, grows with the quota. Without block reuse a slot past the min-slots floor (resident sequences plus reserved dummies) never holds a request, so those slots only add memory. On the Kimi K3 stack (7 DSpark drafts, TP16, max_seq_len 4096, free_gpu_memory_fraction 0.25) the pool got 183 slots: 32.4 GiB of kda_state_tok beside a 13.2 GiB quota. It ran out of device memory above fraction 0.28 at max_seq_len 4096, and above 0.16 at 512. With per-token KDA states and block reuse off, the typical step's request capacity is now raised until its attention pages alone take the quota. The ratio then gives the SSM pool group fewer slots than its floor, so the min-slots constraint sets its size and attention gets the rest. Without per-token states (replay caches only), or with block reuse, the ratio sizes the pool as before. test_v2_kda_token_states_keep_ssm_pool_at_live_floor builds the V2 storage for a 128 MiB quota with 2 MiB state slots: with per-token states the SSM pool keeps its floor (plus one for max_util_for_resume) and attention takes the rest of the quota; without them, or with block reuse, the pool keeps its ratio share. Signed-off-by: Vasanth Sabavat --- .../kv_cache/mamba_cache_manager.py | 30 ++++++++ .../kv_cache/test_mamba_cache_manager.py | 72 +++++++++++++++++++ 2 files changed, 102 insertions(+) diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache/mamba_cache_manager.py b/tensorrt_llm/_torch/pyexecutor/kv_cache/mamba_cache_manager.py index ed5cef76dd35..6b60514f8dfa 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache/mamba_cache_manager.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache/mamba_cache_manager.py @@ -3985,6 +3985,32 @@ def _minimum_live_gpu_quota(self) -> int: state_quota + attention_block_quota, ) + def _ssm_pool_at_live_floor(self, kv_cache_config: KvCacheConfig) -> bool: + """Whether the SSM pool group keeps only its min-slots floor (one slot + per resident sequence plus the reserved dummies) instead of its share + of the typical step's ratio. + + With per-token KDA states, ``kda_state_tok`` takes num_spec states per + SSM slot outside the cache quota (``_allocate_pool_replay_buffers``), + addressed by the request's SSM slot. Without block reuse no SSM slot + holds anything but a live request or a reserved dummy, so slots past + the floor are never used and would only add that memory. + """ + return (getattr(self, "_kda_token_states", False) + and self.local_num_mamba_layers > 0 + and not kv_cache_config.enable_block_reuse + and self._attention_cache_bytes_per_token() > 0) + + def _quota_filling_request_capacity(self, gpu_quota: int) -> int: + """A typical request capacity whose resident requests' attention pages + alone take ``gpu_quota``. The typical step's ratio then gives the SSM + pool group fewer slots than its floor, so the min-slots constraint + sets its size and attention gets the rest of the quota.""" + bytes_per_token = (self._max_resident_sequences() * + self._attention_cache_bytes_per_token()) + # KVCacheDesc.capacity is a 32-bit int. + return min(-(-gpu_quota // bytes_per_token), 1 << 30) + def _build_cache_config( self, config: KVCacheManagerConfigPy) -> KVCacheManagerConfigPy: kv_cache_config = self.kv_cache_config @@ -4030,6 +4056,10 @@ def _build_cache_config( if config.initial_pool_ratio is None: typical_capacity = self._get_typical_request_capacity( kv_cache_config) + if self._ssm_pool_at_live_floor(kv_cache_config): + typical_capacity = max( + typical_capacity, + self._quota_filling_request_capacity(gpu_quota)) request_descs = self._typical_request_descs(typical_capacity, kv_cache_config) typical_step = BatchDesc(request_descs * diff --git a/tests/unittest/_torch/executor/kv_cache/test_mamba_cache_manager.py b/tests/unittest/_torch/executor/kv_cache/test_mamba_cache_manager.py index 0079ea8632a2..df05a69a015e 100644 --- a/tests/unittest/_torch/executor/kv_cache/test_mamba_cache_manager.py +++ b/tests/unittest/_torch/executor/kv_cache/test_mamba_cache_manager.py @@ -3747,6 +3747,78 @@ def test_v2_kda_replay_validates_configuration( ) +@pytest.mark.parametrize( + ("kda_token_states", "enable_block_reuse", "at_floor"), + [(True, False, True), (False, False, False), (True, True, False)], +) +def test_v2_kda_token_states_keep_ssm_pool_at_live_floor( + kda_token_states, enable_block_reuse, at_floor +): + """The per-token KDA states take memory per SSM slot outside the cache quota. Without block reuse the SSM pool + keeps only its live floor and attention gets the rest of the quota; otherwise (replay caches only, or block reuse) + the typical step's ratio sizes it.""" + mgr = object.__new__(MambaHybridCacheManagerV2) + mgr._generation_kv_capacity_headroom = 1 + mgr._has_cp_helix = False + mgr.kv_cache_type = CacheTypeCpp.SELF + mgr.head_dim_per_layer = [64, 64] + mgr.pp_layers = [0, 1] + mgr._mamba_layer_mask = [True, False] + # One 2 MiB grain per state slot, so the pool's slot count follows its grains (a Kimi K3 slot holds 27 MB). + mgr.ssm_bytes = 2 << 20 + mgr.conv_bytes = 64 << 10 + mgr.max_attention_window_vec = [128, 128] + mgr.max_batch_size = 2 + mgr.mapping = Mapping(world_size=1, rank=0, tp_size=1, pp_size=1) + mgr.max_seq_len = 128 + mgr.max_num_tokens = 128 + mgr.tokens_per_block = 32 + mgr.num_local_layers = 2 + mgr.local_num_mamba_layers = 1 + mgr._num_reserved_dummy_slots = 1 + mgr.dtype = DataType.HALF + mgr.enable_swa_scratch_reuse = False + mgr.enable_stats = False + mgr.num_extra_kv_tokens = 0 + mgr.get_layer_bytes_per_token = lambda **kwargs: 8 + # Layer 1's 256-byte pages of 32 tokens (_base_attention_layer_configs; layer 0 becomes the SSM layer). + mgr._attention_cache_bytes_per_token = lambda: 8 + mgr._use_kda_replay_update = True + mgr._kda_token_states = kda_token_states + # The per-token states count in each slot's state bytes (_mamba_state_bytes_per_slot): num_spec fp32 states of + # the SSM state shape, 2 MiB each here. + mgr._kda_replay_num_spec = 2 + mgr.ssm_state_shape = [8, 256, 256] + mgr.kv_cache_config = KvCacheConfig( + avg_seq_len=64, + enable_block_reuse=enable_block_reuse, + enable_partial_reuse=False, + ) + base_config = KVCacheManagerConfig( + tokens_per_block=32, + cache_tiers=[GpuCacheTierConfig(quota=128 << 20)], + layers=_base_attention_layer_configs(2), + ) + runtime_manager = RuntimeKVCacheManager(mgr._build_cache_config(base_config)) + try: + slots = {} + for stats in runtime_manager.get_storage_statistics(): + sizes = stats.slot_sizes if hasattr(stats, "slot_sizes") else stats.slot_size + role = "ssm" if mgr.ssm_bytes in [int(s) for s in sizes] else "attention" + slots[role] = int(stats.total) + finally: + runtime_manager.shutdown() + + floor = mgr._max_resident_sequences() + mgr._num_reserved_dummy_slots + if at_floor: + # The floor's min-slots constraint is scaled by 1 / max_util_for_resume (0.97). + assert floor <= slots["ssm"] <= floor + 1 + # Every other 2 MiB grain goes to attention's 256-byte pages. + assert slots["attention"] >= ((128 << 20) // (2 << 20) - (floor + 2)) * ((2 << 20) // 256) + else: + assert slots["ssm"] > 10 * floor + + def test_mamba_cache_manager_delegates_kda_replay_capability() -> None: mgr = object.__new__(MambaCacheManager) mgr._impl = SimpleNamespace(use_kda_replay_update=True) From 5e16d82b3171d04bc37c76dc17d372d25d82e025 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 15:47:02 -0700 Subject: [PATCH 007/161] [None][fix] MNNVL all-reduce: keep a grown workspace's predecessor alive for captured graphs Growing the MNNVL all-reduce workspace replaced it and freed the previous buffers, pointer tables and flags while CUDA graphs captured before the growth still launched on them. The replaced workspaces now stay alive with the process, and creating or growing a workspace during graph capture raises instead of failing inside the driver. test_mnnvl_workspace_growth_keeps_captured_graphs captures a one-token all-reduce, grows the workspace with an eager two-shot call, replays the graph and checks growth under capture is refused. Signed-off-by: Vasanth Sabavat --- tensorrt_llm/_torch/distributed/ops.py | 14 +++ .../_torch/multi_gpu/test_mnnvl_allreduce.py | 107 ++++++++++++++++++ 2 files changed, 121 insertions(+) diff --git a/tensorrt_llm/_torch/distributed/ops.py b/tensorrt_llm/_torch/distributed/ops.py index 6d20a9aa6394..31d06468876d 100644 --- a/tensorrt_llm/_torch/distributed/ops.py +++ b/tensorrt_llm/_torch/distributed/ops.py @@ -287,6 +287,10 @@ def _get_or_scale_allreduce_mnnvl_workspace( if mapping not in allreduce_mnnvl_workspaces or allreduce_mnnvl_workspaces[ mapping]["buffer_size_bytes"] < (buffer_size_bytes or 0): + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "MNNVL all-reduce workspace creation or growth during CUDA graph capture: run each shape once " + "outside capture first") # Initial buffer to be large enough to support 1024 tokens * 8192 hidden_dim init_buffer_size_bytes = max(1024 * 8192 * elem_size, buffer_size_bytes or 0) @@ -363,6 +367,11 @@ def _get_or_scale_allreduce_mnnvl_workspace( _initialize_allreduce_mnnvl_protocol(candidate_workspace) # Hand ownership of the communicator to the workspace. pending_comms.pop(mapping, None) + previous_workspace = allreduce_mnnvl_workspaces.get(mapping) + if previous_workspace is not None: + # CUDA graphs captured before this growth keep launching on the previous buffers and flags. + MNNVLAllReduce.allreduce_mnnvl_retired_workspaces.setdefault( + mapping, []).append(previous_workspace) allreduce_mnnvl_workspaces[mapping] = candidate_workspace return allreduce_mnnvl_workspaces[mapping] @@ -754,6 +763,11 @@ class MNNVLAllReduce(nn.Module): allreduce_mnnvl_workspaces: typing.ClassVar[dict[Mapping, _MnnvlWorkspace]] = {} + # Workspaces a larger one replaced. CUDA graphs captured before the growth still launch on their buffers and + # flags, so they stay alive with the process. The checkpoint hooks cover only the current workspace. + allreduce_mnnvl_retired_workspaces: typing.ClassVar[dict[ + Mapping, list[_MnnvlWorkspace]]] = {} + # Communicators split for a mapping whose workspace construction has not # succeeded yet. Ownership moves to the workspace once it is published, so # an entry here is never reachable from allreduce_mnnvl_workspaces. diff --git a/tests/unittest/_torch/multi_gpu/test_mnnvl_allreduce.py b/tests/unittest/_torch/multi_gpu/test_mnnvl_allreduce.py index f0abca2492bf..1a19977e4ad7 100644 --- a/tests/unittest/_torch/multi_gpu/test_mnnvl_allreduce.py +++ b/tests/unittest/_torch/multi_gpu/test_mnnvl_allreduce.py @@ -334,6 +334,96 @@ def mnnvl_checkpoint_worker_env(_: int): } +def mnnvl_growth_graph_forward(tensor_parallel_size: int, + tensor_parallel_rank: int) -> bool: + """A graph captured on the MNNVL workspace replays correctly after an eager call grew the workspace, and growth + during capture is refused.""" + env_names = ("TLLM_TEST_MNNVL", "TRTLLM_FORCE_MNNVL_AR") + previous_env = { + name: (name in os.environ, os.environ.get(name)) + for name in env_names + } + tensor_parallel_rank = tensorrt_llm.mpi_rank() + torch.cuda.set_device(tensor_parallel_rank) + os.environ["TLLM_TEST_MNNVL"] = "1" + os.environ["TRTLLM_FORCE_MNNVL_AR"] = "1" + mapping = None + try: + MPI.COMM_WORLD.barrier() + mapping = Mapping( + world_size=tensor_parallel_size, + tp_size=tensor_parallel_size, + rank=tensor_parallel_rank, + ) + MNNVLAllReduce.allreduce_mnnvl_workspaces.pop(mapping, None) + MNNVLAllReduce.allreduce_mnnvl_retired_workspaces.pop(mapping, None) + gc.collect() + MPI.COMM_WORLD.barrier() + expected = tensor_parallel_size * (tensor_parallel_size + 1) // 2 + with torch.inference_mode(): + allreduce = AllReduce( + mapping=mapping, + strategy=AllReduceStrategy.MNNVL, + dtype=torch.bfloat16, + ) + input_ = torch.full((1, 7168), + tensor_parallel_rank + 1, + dtype=torch.bfloat16, + device="cuda") + allreduce(input_) + torch.cuda.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + output = allreduce(input_) + graph.replay() + torch.cuda.synchronize() + torch.testing.assert_close(output, + torch.full_like(output, expected)) + + before = MNNVLAllReduce.allreduce_mnnvl_workspaces[mapping] + # Two-shot footprint 2 * 4096 * 7168 * 2 B: well past the initial 16 MiB per Lamport buffer. + big = torch.ones((4096, 7168), dtype=torch.bfloat16, device="cuda") + torch.testing.assert_close( + allreduce(big), + torch.full_like(big, tensor_parallel_size)) + after = MNNVLAllReduce.allreduce_mnnvl_workspaces[mapping] + assert after["buffer_size_bytes"] > before["buffer_size_bytes"] + assert MNNVLAllReduce.allreduce_mnnvl_retired_workspaces[ + mapping] == [before] + assert before["handle"].is_mapped() + del big + + for step in range(3): + input_.fill_(tensor_parallel_rank + 2 + step) + graph.replay() + torch.cuda.synchronize() + torch.testing.assert_close( + output, + torch.full_like(output, + expected + tensor_parallel_size * + (1 + step))) + + huge = torch.ones((8192, 7168), + dtype=torch.bfloat16, + device="cuda") + with pytest.raises(RuntimeError, + match="during CUDA graph capture"): + with torch.cuda.graph(torch.cuda.CUDAGraph()): + allreduce(huge) + return True + finally: + if mapping is not None: + MNNVLAllReduce.allreduce_mnnvl_workspaces.pop(mapping, None) + MNNVLAllReduce.allreduce_mnnvl_retired_workspaces.pop( + mapping, None) + gc.collect() + for name, (was_present, value) in previous_env.items(): + if was_present: + os.environ[name] = value + else: + os.environ.pop(name, None) + + def mnnvl_checkpoint_rejects_wrong_membership(world_size: int, world_rank: int) -> bool: env_names = ("TLLM_TEST_MNNVL", "TRTLLM_FORCE_MNNVL_AR") @@ -705,6 +795,23 @@ def test_mnnvl_checkpoint_rejects_wrong_group_membership( assert all(results) +@pytest.mark.skipif( + platform.machine().lower() != "aarch64" or torch.cuda.device_count() < 2 + or not MnnvlMemory.supports_mnnvl(), + reason="requires at least two GB200 GPUs with fabric-backed MNNVL", +) +@pytest.mark.parametrize("mpi_pool_executor", [2], indirect=True) +def test_mnnvl_workspace_growth_keeps_captured_graphs( + mpi_pool_executor) -> None: + tensor_parallel_size = mpi_pool_executor.num_workers + results = mpi_pool_executor.map( + mnnvl_growth_graph_forward, + [tensor_parallel_size] * tensor_parallel_size, + range(tensor_parallel_size), + ) + assert all(results) + + def _make_quant_scale(reference_norm: torch.Tensor, fusion_op: AllReduceFusionOp) -> torch.Tensor: amax = reference_norm.abs().max().float().clamp_min(1e-6) From f0ff9a83ff2ef63e700b669ea7f0b40a9117cec6 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Thu, 1 Oct 2026 22:11:00 -0700 Subject: [PATCH 008/161] MNNVL all-reduce: run the workspace-growth graph test in the GB200 multi-GPU list test_mnnvl_workspace_growth_keeps_captured_graphs needs fabric-backed MNNVL on two GB200 GPUs, like the checkpoint graph tests next to it. Signed-off-by: Vasanth Sabavat --- tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml b/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml index 25a83b2d12f5..5102abe33232 100644 --- a/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml +++ b/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml @@ -50,6 +50,7 @@ l0_gb200_multi_gpus: - unittest/_torch/moe/test_moe_comm.py::TestMoEComm::test_moe_comm_postquant - unittest/_torch/multi_gpu/test_mnnvl_allreduce.py::test_mnnvl_checkpoint_preserves_cuda_graph_addresses - unittest/_torch/multi_gpu/test_mnnvl_allreduce.py::test_mnnvl_checkpoint_rejects_wrong_group_membership + - unittest/_torch/multi_gpu/test_mnnvl_allreduce.py::test_mnnvl_workspace_growth_keeps_captured_graphs - unittest/_torch/moe/test_moe_comm.py::TestMoEComm::test_mnnvl_checkpoint_preserves_moe_graph_addresses - unittest/_torch/moe/test_moe_comm.py::TestMoEComm::test_mnnvl_engine_checkpoint_coordination - unittest/_torch/moe/test_moe_comm.py::TestMoEComm::test_mnnvl_checkpoint_failure_is_collective_and_bounded From 551ca164479d6294dc1bbf8dbb27f1370f254225 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 15:47:47 -0700 Subject: [PATCH 009/161] [None][chore] MNNVL all-reduce test: yapf the workspace-growth test The workspace-growth test came from a branch whose yapf settings differ from main's; format it as main's pre-commit does. No behaviour change. Signed-off-by: Vasanth Sabavat --- .../_torch/multi_gpu/test_mnnvl_allreduce.py | 20 +++++++------------ 1 file changed, 7 insertions(+), 13 deletions(-) diff --git a/tests/unittest/_torch/multi_gpu/test_mnnvl_allreduce.py b/tests/unittest/_torch/multi_gpu/test_mnnvl_allreduce.py index 1a19977e4ad7..4efd124b4820 100644 --- a/tests/unittest/_torch/multi_gpu/test_mnnvl_allreduce.py +++ b/tests/unittest/_torch/multi_gpu/test_mnnvl_allreduce.py @@ -384,8 +384,7 @@ def mnnvl_growth_graph_forward(tensor_parallel_size: int, # Two-shot footprint 2 * 4096 * 7168 * 2 B: well past the initial 16 MiB per Lamport buffer. big = torch.ones((4096, 7168), dtype=torch.bfloat16, device="cuda") torch.testing.assert_close( - allreduce(big), - torch.full_like(big, tensor_parallel_size)) + allreduce(big), torch.full_like(big, tensor_parallel_size)) after = MNNVLAllReduce.allreduce_mnnvl_workspaces[mapping] assert after["buffer_size_bytes"] > before["buffer_size_bytes"] assert MNNVLAllReduce.allreduce_mnnvl_retired_workspaces[ @@ -399,23 +398,18 @@ def mnnvl_growth_graph_forward(tensor_parallel_size: int, torch.cuda.synchronize() torch.testing.assert_close( output, - torch.full_like(output, - expected + tensor_parallel_size * - (1 + step))) - - huge = torch.ones((8192, 7168), - dtype=torch.bfloat16, - device="cuda") - with pytest.raises(RuntimeError, - match="during CUDA graph capture"): + torch.full_like( + output, expected + tensor_parallel_size * (1 + step))) + + huge = torch.ones((8192, 7168), dtype=torch.bfloat16, device="cuda") + with pytest.raises(RuntimeError, match="during CUDA graph capture"): with torch.cuda.graph(torch.cuda.CUDAGraph()): allreduce(huge) return True finally: if mapping is not None: MNNVLAllReduce.allreduce_mnnvl_workspaces.pop(mapping, None) - MNNVLAllReduce.allreduce_mnnvl_retired_workspaces.pop( - mapping, None) + MNNVLAllReduce.allreduce_mnnvl_retired_workspaces.pop(mapping, None) gc.collect() for name, (was_present, value) in previous_env.items(): if was_present: From df9d9bfc92aaed597c53d5c0a93299f6b1014b68 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 15:48:31 -0700 Subject: [PATCH 010/161] [None][test] MNNVL all-reduce: one-rank test of the workspace's capture guard Creating or growing the MNNVL all-reduce workspace under CUDA-graph capture raises before the collective allocation. Its only test was the last block of the 2-rank growth test, which runs on GB200 multi-GPU stages only. test_mnnvl_workspace_creation_refuses_graph_capture creates the workspace for a one-rank group inside a capture and expects the guard's RuntimeError. The guard runs before any MNNVL call, so one GPU suffices; without it the call enters the workspace construction under capture and fails inside it with another error. The test joins l0_b200's single-GPU PyTorch block. Signed-off-by: Vasanth Sabavat --- .../integration/test_lists/test-db/l0_b200.yml | 3 +++ .../_torch/multi_gpu/test_mnnvl_allreduce.py | 17 ++++++++++++++++- 2 files changed, 19 insertions(+), 1 deletion(-) diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index 2ca039299f9c..aba43618bb35 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -148,6 +148,9 @@ l0_b200: - unittest/_torch/modeling/test_kimi_linear_checkpoint.py # CPU self-test of the Kimi K3 disagg parity harness comparison logic. - test_kimi_k3_specdec.py::test_kimi_k3_disagg_parity_selftest + # ------------- distributed --------------- + # One process: the MNNVL workspace's capture guard raises before any MNNVL call. + - unittest/_torch/multi_gpu/test_mnnvl_allreduce.py::test_mnnvl_workspace_creation_refuses_graph_capture # ------------- modules (non-MoE) --------------- - unittest/_torch/modules/test_fused_add_rms_norm_quant.py - unittest/_torch/modules/test_fused_activation_quant.py diff --git a/tests/unittest/_torch/multi_gpu/test_mnnvl_allreduce.py b/tests/unittest/_torch/multi_gpu/test_mnnvl_allreduce.py index 4efd124b4820..ce053dba16e4 100644 --- a/tests/unittest/_torch/multi_gpu/test_mnnvl_allreduce.py +++ b/tests/unittest/_torch/multi_gpu/test_mnnvl_allreduce.py @@ -29,7 +29,8 @@ from tensorrt_llm._mnnvl_utils import MnnvlMemory from tensorrt_llm._torch.distributed import (AllReduce, AllReduceFusionOp, AllReduceParams) -from tensorrt_llm._torch.distributed.ops import MNNVLAllReduce +from tensorrt_llm._torch.distributed.ops import ( + MNNVLAllReduce, get_or_scale_allreduce_mnnvl_workspace) from tensorrt_llm.functional import AllReduceStrategy from tensorrt_llm.mapping import Mapping @@ -806,6 +807,20 @@ def test_mnnvl_workspace_growth_keeps_captured_graphs( assert all(results) +@pytest.mark.skipif(not torch.cuda.is_available(), reason="needs a GPU") +def test_mnnvl_workspace_creation_refuses_graph_capture() -> None: + """Creating or growing the MNNVL workspace allocates collectively, so under + CUDA-graph capture it raises before entering the allocation. One process, + a group of one: the guard runs before any MNNVL call.""" + mapping = Mapping(world_size=1, rank=0, tp_size=1) + MNNVLAllReduce.allreduce_mnnvl_workspaces.pop(mapping, None) + graph, stream = torch.cuda.CUDAGraph(), torch.cuda.Stream() + with torch.cuda.graph(graph, stream=stream): + with pytest.raises(RuntimeError, match="during CUDA graph capture"): + get_or_scale_allreduce_mnnvl_workspace(mapping, torch.bfloat16) + assert mapping not in MNNVLAllReduce.allreduce_mnnvl_workspaces + + def _make_quant_scale(reference_norm: torch.Tensor, fusion_op: AllReduceFusionOp) -> torch.Tensor: amax = reference_norm.abs().max().float().clamp_min(1e-6) From fd1bddcaf67cf4bc90a526df5eb82fd1dcd7656a Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Thu, 1 Oct 2026 21:41:57 -0700 Subject: [PATCH 011/161] MNNVL all-reduce: count Lamport arrivals per CTA in the RMSNorm-fused kernels oneshotAllreduceFusionKernel and rmsNormLamport used LamportFlags' cluster arrival, in which every thread arrives on the cluster barrier but only cluster rank 0's first warp waits, and then arrived on the cluster barrier again in the RMSNorm epilogue's cluster reduction. Each thread must arrive and wait once per phase: a second arrival counted toward the first phase lets the reduction read a peer CTA's partial sum before it is written, and counts the cluster in before a late CTA has read the buffer flags. Both kernels now count arrivals per CTA, as the attn_res kernel already does; their cluster barrier is left to the reduction. Signed-off-by: Vasanth Sabavat --- cpp/tensorrt_llm/common/lamportUtils.cuh | 2 ++ .../communicationKernels/mnnvlAllreduceKernels.cu | 9 ++++++--- 2 files changed, 8 insertions(+), 3 deletions(-) diff --git a/cpp/tensorrt_llm/common/lamportUtils.cuh b/cpp/tensorrt_llm/common/lamportUtils.cuh index 60639e2f4870..28a06d1d062a 100644 --- a/cpp/tensorrt_llm/common/lamportUtils.cuh +++ b/cpp/tensorrt_llm/common/lamportUtils.cuh @@ -261,6 +261,8 @@ public: if constexpr (UseCGA) { #if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) + // Every thread arrives on the cluster barrier but only rank 0's first warp waits, so a kernel using this + // must not arrive on the cluster barrier again (each thread arrives and waits once per phase). cg::cluster_group cluster = cg::this_cluster(); __cluster_barrier_arrive(); if (cluster.block_rank() == 0 && threadIdx.x < kWARP_SIZE) diff --git a/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllreduceKernels.cu b/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllreduceKernels.cu index a8ff3dd431a8..876836c81f7e 100644 --- a/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllreduceKernels.cu +++ b/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllreduceKernels.cu @@ -597,8 +597,9 @@ __global__ void __launch_bounds__(1024) oneshotAllreduceFusionKernel(MnnvlAllRed int threadOffset = token * params.tokenDim + packedIdx * kELTS_PER_THREAD; #endif - // We only use 1 stage for the oneshot allreduce - LamportFlags flag(params.bufferFlags, 1); + // We only use 1 stage for the oneshot allreduce. Arrivals are counted per CTA: the RMSNorm epilogue's cluster + // reduction below must be each thread's only phase of the cluster barrier. + LamportFlags flag(params.bufferFlags, 1); T* stagePtrMcast = reinterpret_cast(flag.getCurLamportBuf(params.mcastPtr, 0)); T* stagePtrLocal = reinterpret_cast(flag.getCurLamportBuf(params.inputPtrs[params.rank], 0)); bool const inBounds = packedIdx * kELTS_PER_THREAD < params.tokenDim; @@ -1016,7 +1017,9 @@ __global__ __launch_bounds__(1024) void rmsNormLamport(MnnvlAllReduceKernelParam T* smemResidual = reinterpret_cast(&smem[smemBufferSize]); T* smemGamma = reinterpret_cast(&smem[2 * smemBufferSize]); - LamportFlags flag(params.bufferFlags, MNNVLTwoShotStage::NUM_STAGES); + // Arrivals are counted per CTA: with UseCGA the RMSNorm cluster reduction below must be each thread's only phase + // of the cluster barrier. + LamportFlags flag(params.bufferFlags, MNNVLTwoShotStage::NUM_STAGES); T* input = reinterpret_cast( flag.getCurLamportBuf(reinterpret_cast(params.bufferInputPtr), MNNVLTwoShotStage::BROADCAST)); From 7a7ccf360fe011b7e9089099771365b51e7aed7a Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 00:03:12 -0700 Subject: [PATCH 012/161] MNNVL all-reduce: start every CTA of the cluster before the RMSNorm cluster reduction writes The RESIDUAL_RMS_NORM epilogues of oneshotAllreduceFusionKernel and rmsNormLamport write each CTA's partial sum into the other CTAs' shared memory through the cluster mapping before any cluster barrier wait, so a target CTA may not have started yet; racecheck reports every such write once the kernels' arrival counting no longer stalls it. A cluster.sync() before the writes makes every CTA of the cluster start first, as the attn_res kernel's early arrive / wait does. Signed-off-by: Vasanth Sabavat --- .../kernels/communicationKernels/mnnvlAllreduceKernels.cu | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllreduceKernels.cu b/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllreduceKernels.cu index 876836c81f7e..e60a2f233c15 100644 --- a/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllreduceKernels.cu +++ b/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllreduceKernels.cu @@ -711,6 +711,8 @@ __global__ void __launch_bounds__(1024) oneshotAllreduceFusionKernel(MnnvlAllRed fullSum = 0.F; // Need to reduce over the entire cluster int const blockRank = cluster.block_rank(); + // Every CTA of the cluster must have started before another CTA writes its shared memory. + cluster.sync(); if (threadIdx.x < numBlocks) { cluster.map_shared_rank(&sharedVal[0], threadIdx.x)[blockRank] = blockSum; @@ -1140,6 +1142,8 @@ __global__ __launch_bounds__(1024) void rmsNormLamport(MnnvlAllReduceKernelParam fullSum = 0.F; // Need to reduce over the entire cluster int const blockRank = cluster.block_rank(); + // Every CTA of the cluster must have started before another CTA writes its shared memory. + cluster.sync(); if (threadIdx.x < numBlocks) { cluster.map_shared_rank(&sharedVal[0], threadIdx.x)[blockRank] = blockSum; From f4672116460e12b094bb9548e0b34d628c74e77d Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Thu, 1 Oct 2026 21:10:23 -0700 Subject: [PATCH 013/161] [None][fix] Kimi K3 attn_res_fwd online kernel: order the ws_stats reads before the next chunk's writes attn_res_fwd_online_v2_kernel reuses one ws_stats buffer every chunk of candidates: lane 0 of each consumer warp writes its row, a named barrier orders the writes before every consumer's cross-warp reads, but nothing ordered those reads before the next chunk's writes. With more than one chunk (N > 4 at H 7168) a warp that finished reading could overwrite its row while a slower warp still read it, mixing the next chunk's statistics into this chunk's logits. Under compute-sanitizer racecheck, which serializes warps, every such call returned a wrong mixture (relative error 0.6-2.1 of max |ref|); natively it was not observed. Add the second barrier the persistent fork already has. Also order every lane's reads of a chunk's slots before lane 0 releases them to the producer (__syncwarp before the bar_consumed arrive), in online_v2 and in the persistent fork. Signed-off-by: Vasanth Sabavat --- .../kernels/kimiK3AttnRes/attnResFwd.cu | 15 +++++++++++++-- .../kimiK3AttnRes/attnResFwdPersistentFused.cu | 2 ++ 2 files changed, 15 insertions(+), 2 deletions(-) diff --git a/cpp/tensorrt_llm/kernels/kimiK3AttnRes/attnResFwd.cu b/cpp/tensorrt_llm/kernels/kimiK3AttnRes/attnResFwd.cu index e714466dbb3a..952b54da16ab 100644 --- a/cpp/tensorrt_llm/kernels/kimiK3AttnRes/attnResFwd.cu +++ b/cpp/tensorrt_llm/kernels/kimiK3AttnRes/attnResFwd.cu @@ -411,9 +411,14 @@ __global__ void __launch_bounds__(BLK, 1) attn_res_fwd_online_v2_kernel(bf16_t c float2 f[2] = {__bfloat1622float2(v2[0]), __bfloat1622float2(v2[1])}; if constexpr (FULL_N12) { - if (n == AN - 1 && lane == 0) + if (n == AN - 1) { - cute::arrive_barrier(plan.bar_consumed[chunk_slot]); + // Every lane's reads of the chunk's slots before lane 0 releases them. + __syncwarp(); + if (lane == 0) + { + cute::arrive_barrier(plan.bar_consumed[chunk_slot]); + } } } tmem_st_32dp32bNx<4>( @@ -492,6 +497,8 @@ __global__ void __launch_bounds__(BLK, 1) attn_res_fwd_online_v2_kernel(bf16_t c } if constexpr (!FULL_N12) { + // Every lane's reads of the chunk's slots before lane 0 releases them to the producer. + __syncwarp(); if (lane == 0) { cute::arrive_barrier(plan.bar_consumed[chunk_slot]); @@ -551,6 +558,10 @@ __global__ void __launch_bounds__(BLK, 1) attn_res_fwd_online_v2_kernel(bf16_t c { cross_warp_tail(lane); } + // ws_stats is one buffer reused every chunk. The barrier above orders this chunk's writes before the + // reads; this one orders the reads before the next chunk's writes, so a warp that has finished + // reading cannot overwrite its row while a slower warp still reads it. + cutlass::arch::NamedBarrier::sync(CONSUMER_THREADS, 0); float logit_n[N_CHUNK]; #pragma unroll for (int n = 0; n < N_CHUNK; n++) diff --git a/cpp/tensorrt_llm/kernels/kimiK3AttnRes/attnResFwdPersistentFused.cu b/cpp/tensorrt_llm/kernels/kimiK3AttnRes/attnResFwdPersistentFused.cu index 7245f724134a..32bfae257368 100644 --- a/cpp/tensorrt_llm/kernels/kimiK3AttnRes/attnResFwdPersistentFused.cu +++ b/cpp/tensorrt_llm/kernels/kimiK3AttnRes/attnResFwdPersistentFused.cu @@ -650,6 +650,8 @@ __global__ void __launch_bounds__(BLK, 1) default: __builtin_unreachable(); } } + // Every lane's reads of the chunk's slots before lane 0 releases them to the producer. + __syncwarp(); if (lane == 0) { mbarrier_arrive(plan.bar_consumed[chunk_slot]); From 26cb3c37cd70812e6d74040138709fc3b51cbe52 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 02:28:16 -0700 Subject: [PATCH 014/161] [None][fix] Kimi K3 attn_res: cross-proxy fence before a consumer releases a TMA-filled slot The online_v2, N = 1 tile and persistent fused kernels read their V (and delta) slots with generic-proxy shared loads and release them with an mbarrier arrive; the producer then refills them with cp.async.bulk (async proxy). Accesses to one location through two proxies need a cross-proxy fence: fence.proxy.async.shared::cta in every reading thread before the release. In the current SASS each release already issues after an instruction that consumes the last load, so this makes that ordering independent of instruction scheduling. Outputs are bit-identical. Signed-off-by: Vasanth Sabavat --- cpp/tensorrt_llm/kernels/kimiK3AttnRes/attnResFwd.cu | 9 +++++++++ .../kernels/kimiK3AttnRes/attnResFwdPersistentFused.cu | 3 +++ 2 files changed, 12 insertions(+) diff --git a/cpp/tensorrt_llm/kernels/kimiK3AttnRes/attnResFwd.cu b/cpp/tensorrt_llm/kernels/kimiK3AttnRes/attnResFwd.cu index 952b54da16ab..39a14df806b6 100644 --- a/cpp/tensorrt_llm/kernels/kimiK3AttnRes/attnResFwd.cu +++ b/cpp/tensorrt_llm/kernels/kimiK3AttnRes/attnResFwd.cu @@ -413,6 +413,9 @@ __global__ void __launch_bounds__(BLK, 1) attn_res_fwd_online_v2_kernel(bf16_t c { if (n == AN - 1) { + // The producer refills the slots through the async proxy (TMA): a + // cross-proxy fence orders this lane's generic-proxy reads before it. + asm volatile("fence.proxy.async.shared::cta;" ::: "memory"); // Every lane's reads of the chunk's slots before lane 0 releases them. __syncwarp(); if (lane == 0) @@ -497,6 +500,9 @@ __global__ void __launch_bounds__(BLK, 1) attn_res_fwd_online_v2_kernel(bf16_t c } if constexpr (!FULL_N12) { + // The producer refills the slots through the async proxy (TMA): a cross-proxy fence orders this + // lane's generic-proxy reads before it. + asm volatile("fence.proxy.async.shared::cta;" ::: "memory"); // Every lane's reads of the chunk's slots before lane 0 releases them to the producer. __syncwarp(); if (lane == 0) @@ -983,6 +989,9 @@ __global__ void __launch_bounds__(BLK, 1) } cutlass::arch::NamedBarrier::sync(CONSUMER_THREADS, 1); } + // The producer refills the tile through the async proxy (TMA): a cross-proxy fence orders this thread's + // generic-proxy reads before it. + asm volatile("fence.proxy.async.shared::cta;" ::: "memory"); cute::arrive_barrier(plan.bar_consumed[slot]); } } diff --git a/cpp/tensorrt_llm/kernels/kimiK3AttnRes/attnResFwdPersistentFused.cu b/cpp/tensorrt_llm/kernels/kimiK3AttnRes/attnResFwdPersistentFused.cu index 32bfae257368..ee734add0661 100644 --- a/cpp/tensorrt_llm/kernels/kimiK3AttnRes/attnResFwdPersistentFused.cu +++ b/cpp/tensorrt_llm/kernels/kimiK3AttnRes/attnResFwdPersistentFused.cu @@ -650,6 +650,9 @@ __global__ void __launch_bounds__(BLK, 1) default: __builtin_unreachable(); } } + // The producer refills the slots through the async proxy (TMA): a cross-proxy fence orders this lane's + // generic-proxy reads before it. + asm volatile("fence.proxy.async.shared::cta;" ::: "memory"); // Every lane's reads of the chunk's slots before lane 0 releases them to the producer. __syncwarp(); if (lane == 0) From 05e9e4f1f25555d9a91e89f95f344a07fa4d17d0 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 11:49:29 -0700 Subject: [PATCH 015/161] [None][test] Kimi K3 attn_res: check the cross-proxy fence before each slot release test_attn_res_proxy_fences.py reads attnResFwd.cu and attnResFwdPersistentFused.cu and checks every consumer release of a cp.async.bulk-filled slot (an arrive on bar_consumed). Each must follow fence.proxy.async.shared::cta, with only the warp sync and the lane-0 guard between the two. The fence orders the consumers' generic-proxy reads of the slot before the producer's async-proxy refill. Without it the order rests on instruction scheduling, which no test of values can see. The check fails on the sources from before the fence was added (four releases) and passes on the current ones. CPU only. Signed-off-by: Vasanth Sabavat --- .../test_attn_res_proxy_fences.py | 87 +++++++++++++++++++ 1 file changed, 87 insertions(+) create mode 100644 tests/unittest/_torch/modules/kimi_k3_attn_res/test_attn_res_proxy_fences.py diff --git a/tests/unittest/_torch/modules/kimi_k3_attn_res/test_attn_res_proxy_fences.py b/tests/unittest/_torch/modules/kimi_k3_attn_res/test_attn_res_proxy_fences.py new file mode 100644 index 000000000000..b1a83d52af7f --- /dev/null +++ b/tests/unittest/_torch/modules/kimi_k3_attn_res/test_attn_res_proxy_fences.py @@ -0,0 +1,87 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Source check of the cross-proxy fences in the Kimi K3 attention-residual kernels (PTX ISA, memory consistency +model, proxies). + +The consumers read the V (and delta) slots with generic-proxy shared loads and release each slot with an mbarrier +arrive on ``bar_consumed``. The producer then refills the slot with ``cp.async.bulk``, an async-proxy write. Accesses +to one location through two proxies need a cross-proxy fence: every reading thread issues +``fence.proxy.async.shared::cta`` after its last read and before the release. Without it the order rests on how the +compiler schedules the instructions. No test of values can see that, so the kernels' sources are read. + + pytest test_attn_res_proxy_fences.py +""" + +import pathlib +import re + +import pytest + +pytestmark = pytest.mark.cpu_only + +_SRC = pathlib.Path(__file__).resolve().parents[5] / "cpp/tensorrt_llm/kernels/kimiK3AttnRes" + +# The kernels whose consumers release cp.async.bulk-filled slots, with the number of releases each file has. +RELEASES = {"attnResFwd.cu": 3, "attnResFwdPersistentFused.cu": 1} + +RELEASE = re.compile(r"\b(cute::arrive_barrier|mbarrier_arrive)\(\s*plan\.bar_consumed\b") +FENCE = 'asm volatile("fence.proxy.async.shared::cta;" ::: "memory");' +# What may sit between the fence and the release: the warp sync that gathers every lane's reads before lane 0 releases, +# and the lane-0 guard. Any other statement could hold a read after the fence. +BETWEEN = re.compile(r"^(__syncwarp\(\);|if \(lane == 0\)|\{)$") + + +def _code(line): + return line.split("//", 1)[0].strip() + + +def violations(lines): + """(line number, release) of every ``bar_consumed`` release whose closest preceding statement, past the warp sync + and the lane-0 guard, is not the cross-proxy fence.""" + out = [] + for i, line in enumerate(lines): + if not RELEASE.search(_code(line)): + continue + for j in range(i - 1, -1, -1): + code = _code(lines[j]) + if not code or BETWEEN.match(code): + continue + if code != FENCE: + out.append((i + 1, line.strip())) + break + else: + out.append((i + 1, line.strip())) + return out + + +@pytest.mark.parametrize("name", sorted(RELEASES)) +def test_slot_release_follows_cross_proxy_fence(name): + lines = (_SRC / name).read_text().splitlines() + releases = [i for i, line in enumerate(lines) if RELEASE.search(_code(line))] + assert len(releases) == RELEASES[name], ( + f"{name}: {len(releases)} bar_consumed releases, expected {RELEASES[name]} " + "(update RELEASES if the kernels changed)" + ) + assert any("cp.async.bulk.shared::cta.global" in line for line in lines), ( + f"{name}: no cp.async.bulk refill; the fence may no longer be needed" + ) + assert violations(lines) == [], ( + f"{name}: slot releases without fence.proxy.async.shared::cta right before them: {violations(lines)}" + ) + + +def test_check_flags_a_release_without_the_fence(): + """The check itself: the release pattern without the fence, as before the fix, is reported.""" + lines = [ + "// Every lane's reads of the chunk's slots before lane 0 releases them.", + "__syncwarp();", + "if (lane == 0)", + "{", + " cute::arrive_barrier(plan.bar_consumed[chunk_slot]);", + "}", + ] + assert violations(lines) == [(5, "cute::arrive_barrier(plan.bar_consumed[chunk_slot]);")] + assert violations([FENCE] + lines) == [] + assert violations([FENCE, "float x = buf[0];"] + lines) == [ + (7, "cute::arrive_barrier(plan.bar_consumed[chunk_slot]);") + ] From af3e3841757b919537e1430da13d40a8ae8f8863 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Thu, 1 Oct 2026 23:54:42 -0700 Subject: [PATCH 016/161] [None][fix] kdaDecode legacy: __syncwarp before the lane-0 block-reduce stores In block_reduce_sum2_for and block_reduce_sum2_active_for, lanes of warp 0 read the per-warp partials scratch[lane] / scratch[kReduceWarps + lane], reduce them with shuffles, and lane 0 then stores the block totals into scratch[0] and scratch[1], which lane 1 read. Shuffles do not order shared memory, so the store and lane 1's read were unordered (compute-sanitizer racecheck: 32 warnings per run of test_k3_kda_decode_attn, all in kda_decode_fusion_compact_heads_kernel). Lane 0's totals depend on lane 1's loaded value, so the results were not affected; __syncwarp() makes the order explicit. Signed-off-by: Vasanth Sabavat --- cpp/tensorrt_llm/kernels/kdaDecode/kdaDecodeLegacy.cu | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/cpp/tensorrt_llm/kernels/kdaDecode/kdaDecodeLegacy.cu b/cpp/tensorrt_llm/kernels/kdaDecode/kdaDecodeLegacy.cu index ca77fc702e29..6d34478a8507 100644 --- a/cpp/tensorrt_llm/kernels/kdaDecode/kdaDecodeLegacy.cu +++ b/cpp/tensorrt_llm/kernels/kdaDecode/kdaDecodeLegacy.cu @@ -275,6 +275,8 @@ __device__ __forceinline__ Sum2 block_reduce_sum2_for(float x, float y, float* s block_y = lane < kReduceWarps ? scratch[kReduceWarps + lane] : 0.0f; block_x = warp_reduce_sum(block_x); block_y = warp_reduce_sum(block_y); + // Lane 0 overwrites partials that other lanes of this warp read above. + __syncwarp(); if (lane == 0) { scratch[0] = block_x; @@ -349,6 +351,8 @@ __device__ __forceinline__ Sum2 block_reduce_sum2_active_for(float x, float y, f block_y = lane < kReduceWarps ? scratch[kReduceWarps + lane] : 0.0f; block_x = warp_reduce_sum(block_x); block_y = warp_reduce_sum(block_y); + // Lane 0 overwrites partials that other lanes of this warp read above. + __syncwarp(); if (lane == 0) { scratch[0] = block_x; From 1e0dceb74e5e18d92294c6a62965c7043d87d3dd Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Thu, 1 Oct 2026 21:28:59 -0700 Subject: [PATCH 017/161] DFlash: start the dummy slot at context length 0 every step CUDA-graph padding rows and warmup dummies share the drafter's dummy slot. Their accepted tokens grow its context length like any request's, but nothing reset it, while their page-table rows are the padding request's pages with every other entry mapped to page 0, another request's page. Once the dummy length passed the padding request's own pages, the padding rows' context K / V (k3_ctx_kv, or the torch path) landed on that page: a live request's drafter context overwritten on every padded step, costing acceptance. prepare() now writes 0 to the dummy slot every step, with the evicted slots. test_dflash_dummy_slot.py: padded steps keep the dummy slot at 0 and the real slots untouched; an evicted request and the dummy reset in one step. It fails without the reset. Signed-off-by: Vasanth Sabavat --- tensorrt_llm/_torch/speculative/dflash.py | 4 + .../hw_agnostic/test_dflash_dummy_slot.py | 83 +++++++++++++++++++ 2 files changed, 87 insertions(+) create mode 100644 tests/unittest/_torch/speculative/hw_agnostic/test_dflash_dummy_slot.py diff --git a/tensorrt_llm/_torch/speculative/dflash.py b/tensorrt_llm/_torch/speculative/dflash.py index d8097dffbd58..d760e15fbc78 100644 --- a/tensorrt_llm/_torch/speculative/dflash.py +++ b/tensorrt_llm/_torch/speculative/dflash.py @@ -551,6 +551,10 @@ def prepare(self): evicted[slot] = 0 worker._req_ctx_pos.pop(rid, None) worker._free_slots.append(slot) + # The dummy slot starts every step empty. Its rows (CUDA-graph padding, warmup dummies) add their accepted + # tokens to its length like any request, and their table rows are the padding request's pages, the rest + # mapped to page 0 (another request's). Left growing, a padding row's context K / V would land there. + evicted[worker._dummy_slot] = 0 worker._write_ctx_len(evicted) # A disagg generation worker receives prompt KV instead of diff --git a/tests/unittest/_torch/speculative/hw_agnostic/test_dflash_dummy_slot.py b/tests/unittest/_torch/speculative/hw_agnostic/test_dflash_dummy_slot.py new file mode 100644 index 000000000000..37d837aa7c28 --- /dev/null +++ b/tests/unittest/_torch/speculative/hw_agnostic/test_dflash_dummy_slot.py @@ -0,0 +1,83 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""DFlash's dummy slot (CUDA-graph padding and warmup dummies) starts every step at context length 0. + +Padding rows add their accepted tokens to the dummy slot's length like any request, while their page-table rows are +the padding request's pages with the rest mapped to page 0, which another request owns. A dummy length that kept +growing would put the padding rows' context K / V on that page. ``DFlashSpecMetadata.prepare`` runs before every step +(eager or replayed graph); after it the dummy slot's length is 0 and every real slot's is untouched. +""" + +from functools import partial +from types import SimpleNamespace + +import pytest +import torch + +from tensorrt_llm._torch.speculative.dflash import DFlashSpecMetadata, DFlashWorker + +pytestmark = pytest.mark.skipif( + not torch.cuda.is_available(), reason="the slot lengths live on the GPU" +) + +FLOOR = 1 << 40 # request ids from here on are CUDA-graph padding dummies + + +def _worker(lengths): + w = SimpleNamespace( + _ctx_buf_inited=True, + _req_to_slot={11: 0, 12: 1, 13: 2}, + _req_ctx_pos={}, + _free_slots=[3], + _graph_dummy_id_floor=FLOOR, + _dummy_slot=len(lengths) - 1, + _ctx_len=torch.tensor(lengths, dtype=torch.long, device="cuda"), + _ctx_len_host=list(lengths), + _batch_to_slot=torch.zeros(8, dtype=torch.long, device="cuda"), + ) + w._write_ctx_len = partial(DFlashWorker._write_ctx_len, w) + w._assign_slot = lambda rid, *a, **k: None + return w + + +def _prepare(worker, request_ids, num_generations): + meta = SimpleNamespace( + request_ids=request_ids, + num_generations=num_generations, + batch_indices_cuda=torch.zeros(8, dtype=torch.int, device="cuda"), + _dflash_worker=worker, + ) + DFlashSpecMetadata.prepare(meta) + torch.cuda.synchronize() + + +def test_padding_steps_keep_the_dummy_slot_empty(): + worker = _worker([40, 300, 7, 0, 0]) + padded = [11, 12, 13, FLOOR + 1, FLOOR + 1] # three requests in a graph of five + for step in range(4): + # What the step's acceptance does to the slots of its rows (k3_ctx_kv / the torch path): + accepted. + _prepare(worker, padded, num_generations=5) + assert worker._ctx_len.tolist() == [40 + 8 * step, 300 + 8 * step, 7 + 8 * step, 0, 0], step + assert worker._batch_to_slot[:5].tolist() == [0, 1, 2, 4, 4] + assert worker._ctx_len_host[4] == 0 + worker._ctx_len[[0, 1, 2]] += 8 + worker._ctx_len[4] += 2 * 8 # both padding rows land on the dummy slot + + +def test_evicted_request_and_dummy_reset_together(): + worker = _worker([40, 300, 7, 0, 96]) + _prepare(worker, [11, 13, FLOOR + 1], num_generations=3) # request 12 left; one padding row + assert worker._ctx_len.tolist() == [40, 0, 7, 0, 0] + assert 1 in worker._free_slots and 12 not in worker._req_to_slot From 5db6a14b6296d63def0bcc0b398abd6401c204a6 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 17:48:35 -0700 Subject: [PATCH 018/161] [None][feat] Kimi K3 decode: KDA and MLA CuTe DSL kernels with their op tests The Kimi K3 decode kernels for the KDA and MLA layers, as of the K3 stack 82a110a92a, verbatim: - k3_kda_attn/: trtllm::k3_kda_qkvg (the KDA projection as one CTM kernel), trtllm::k3_kda_attn (projection and speculative verify of one request's 8 tokens in one launch) and trtllm::k3_kda_decode_attn (projection and T = 1 decode of up to 8 requests); - k3_kda_verify/: trtllm::k3_kda_verify (the KDA verify of a step's drafts over the per-token states of the V2 hybrid cache manager); - k3_mla/: trtllm::k3_mla_q / k3_mla_qkv / k3_mla_qkv_out (the MLA query path and the latent KV append) and trtllm::k3_mla_attn / _attn_out / _attn_vb_out (decode attention over the paged latent cache, optionally with v_b and the output gate); - k3_mla_decode_view in the CuTe DSL MLA backend: the attention metadata's view of R generation requests of T tokens that k3_mla_attn takes. The op tests come along; they carry the regression tests of the fixes made to these kernels: Lamport buffer index in 0..2 and launch counters across the int32 wrap, state pools addressed at 64-bit slot offsets (pools past 2 GiB), slot indices at any element offset, one shared Lamport set for the decode and verify launches, k3_mla_attn reading only the workspace words its own call wrote, and the no-cluster counters across the wrap. Signed-off-by: Vasanth Sabavat --- .../attention/backends/fmha/cute_dsl_mla.py | 86 + .../cute_dsl_kernels/k3_kda_attn/__init__.py | 19 + .../k3_kda_attn/k3_kda_attn_kernel.py | 1552 +++++++++++++++++ .../k3_kda_attn/k3_kda_decode_kernel.py | 802 +++++++++ .../_torch/cute_dsl_kernels/k3_kda_attn/op.py | 356 ++++ .../k3_kda_verify/__init__.py | 19 + .../k3_kda_verify/k3_kda_verify_kernel.py | 827 +++++++++ .../cute_dsl_kernels/k3_kda_verify/op.py | 205 +++ .../cute_dsl_kernels/k3_mla/__init__.py | 19 + .../k3_mla/k3_mla_attn_kernel.py | 1187 +++++++++++++ .../k3_mla/k3_mla_q_kernel.py | 1099 ++++++++++++ .../_torch/cute_dsl_kernels/k3_mla/op.py | 575 ++++++ .../kimi_k3/test_k3_kda_attn.py | 322 ++++ .../kimi_k3/test_k3_kda_decode_attn.py | 523 ++++++ .../kimi_k3/test_k3_kda_pools_past_2g.py | 212 +++ .../kimi_k3/test_k3_kda_verify.py | 588 +++++++ .../kimi_k3/test_k3_mla_attn.py | 471 +++++ .../kimi_k3/test_k3_mla_decode_view.py | 125 ++ .../cute_dsl_kernels/kimi_k3/test_k3_mla_q.py | 396 +++++ 19 files changed, 9383 insertions(+) create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_attn/__init__.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_attn/k3_kda_attn_kernel.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_attn/k3_kda_decode_kernel.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_attn/op.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_verify/__init__.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_verify/k3_kda_verify_kernel.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_verify/op.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/__init__.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/k3_mla_attn_kernel.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/k3_mla_q_kernel.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/op.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_kda_attn.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_kda_decode_attn.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_kda_pools_past_2g.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_kda_verify.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_attn.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_decode_view.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_q.py diff --git a/tensorrt_llm/_torch/attention/backends/fmha/cute_dsl_mla.py b/tensorrt_llm/_torch/attention/backends/fmha/cute_dsl_mla.py index ac97de1d550c..4f977580943d 100644 --- a/tensorrt_llm/_torch/attention/backends/fmha/cute_dsl_mla.py +++ b/tensorrt_llm/_torch/attention/backends/fmha/cute_dsl_mla.py @@ -40,6 +40,92 @@ _LOG2_E = math.log2(math.e) +def k3_mla_decode_view(attn, meta, num_tokens: int): + """The per-layer inputs of trtllm::k3_mla_qkv and trtllm::k3_mla_attn(_vb)_out for a bf16 decode step of R + generation requests of T = num_tokens / R tokens each (no context requests, R <= 8, T <= 8), or the reason it + does not apply (a string; never raises). The dict: ``pool`` (flat bf16 pool), ``row_stride``, ``page_table`` + (int32 [R, W], request i's pages in row i: a strided view of ``kv_cache_block_offsets``, row stride + ``page_table.stride(0)``), ``page_offset``, ``seq_len`` (int32 [R], the kv lengths including the step's tokens), + ``softmax_scale``, ``num_requests`` (R) and ``tokens_per_request`` (T). The page table and lengths are the + host-filled metadata buffers (``kv_cache_block_offsets``, ``kv_lens_cuda_runtime``), which the kernels may read + before their grid wait. T is taken as uniform, as the engine pads every generation request to the same number of + tokens.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_mla.op import MAX_REQUEST_TOKENS, MAX_REQUESTS + + num_requests = meta.num_generations + if ( + meta.num_contexts != 0 + or not 0 < num_requests <= MAX_REQUESTS + or num_tokens % num_requests + or not 0 < num_tokens // num_requests <= MAX_REQUEST_TOKENS + ): + return ( + f"{meta.num_contexts} context / {num_requests} generation requests, {num_tokens} tokens" + ) + if ( + meta.beam_width != 1 + or getattr(meta, "is_spec_dec_tree", False) + or getattr(meta, "is_spec_dec_dynamic_tree", False) + ): + return f"beam width {meta.beam_width} or a tree speculative mask" + if attn.num_heads % 6 or attn.kv_lora_rank != 512 or attn.qk_rope_head_dim != 64: + return f"heads {attn.num_heads}, latent {attn.kv_lora_rank}, rope {attn.qk_rope_head_dim}" + if meta.tokens_per_block != 64 or meta.helix_position_offsets is not None: + return f"page {meta.tokens_per_block}, helix {meta.helix_position_offsets is not None}" + if meta.kv_cache_manager is None or meta.kv_cache_block_offsets is None: + return "no KV cache manager / block offsets" + kv_pool = meta.kv_cache_manager.get_buffers(attn.layer_idx) + if kv_pool.dtype != torch.bfloat16: + return f"kv cache {kv_pool.dtype}" + packed_block = 1 + for size in kv_pool.shape[1:]: + packed_block *= size + block_stride = kv_pool.stride(0) + layers_in_pool = block_stride // packed_block if packed_block else 1 + layer_in_pool = 0 + if layers_in_pool > 1 and block_stride == layers_in_pool * packed_block: + layer_in_pool = kv_pool.storage_offset() // packed_block + kv_pool = kv_pool.as_strided( + (kv_pool.shape[0] * layers_in_pool, *kv_pool.shape[1:]), + (packed_block, *kv_pool.stride()[1:]), + 0, + ) + if ( + kv_pool.dim() != 5 + or kv_pool.shape[1] != 1 + or kv_pool.shape[3] != 1 + or kv_pool.shape[2] != 64 + ): + return f"pool layout {tuple(kv_pool.shape)}" + if not (kv_pool.is_contiguous() and kv_pool.stride(2) >= 576 and kv_pool.stride(2) % 8 == 0): + return f"pool strides {kv_pool.stride()}" + pool_idx = int(meta.host_kv_cache_pool_mapping[attn.get_local_layer_idx(meta), 0]) + gen = slice(meta.num_contexts, meta.num_contexts + num_requests) + page_table = meta.kv_cache_block_offsets[pool_idx, gen, 0, :] + seq_len = meta.kv_lens_cuda_runtime[gen] + if page_table.dtype != torch.int32 or seq_len.dtype != torch.int32: + return f"page table {page_table.dtype} / length {seq_len.dtype} not int32" + if ( + page_table.shape[0] != num_requests + or seq_len.shape[0] != num_requests + or page_table.stride(1) != 1 + ): + return f"page table {tuple(page_table.shape)} strides {page_table.stride()} / lengths {tuple(seq_len.shape)}" + softmax_scale = float( + 1.0 / (math.sqrt(attn.qk_nope_head_dim + attn.qk_rope_head_dim) * attn.q_scaling) + ) + return dict( + pool=kv_pool.view(-1), + row_stride=kv_pool.stride(2), + page_table=page_table, + page_offset=int(layer_in_pool), + seq_len=seq_len, + softmax_scale=softmax_scale, + num_requests=num_requests, + tokens_per_request=num_tokens // num_requests, + ) + + class CuteDslMlaFmha(PhasedFmha): """Blackwell CuTe DSL FMHA library for decode-only MLA.""" diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_attn/__init__.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_attn/__init__.py new file mode 100644 index 000000000000..cc433eba6f41 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_attn/__init__.py @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Kimi K3 KDA attention projection as one CTM kernel (``trtllm::k3_kda_qkvg``). + +Importing :mod:`.op` registers the torch op; nothing is imported eagerly here so that +the CuTe DSL dependency stays optional for every other model. +""" diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_attn/k3_kda_attn_kernel.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_attn/k3_kda_attn_kernel.py new file mode 100644 index 000000000000..c61c40dd9760 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_attn/k3_kda_attn_kernel.py @@ -0,0 +1,1552 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Kimi K3 KDA fused projection y = x W^T (TP16 rank slice: W bf16 [3208, 7168], x bf16 [T <= 8, 7168]) streamed in +three priority phases, so that a consumer of the projection can start on the rows it needs first. + +Rows of W (= columns of y): q [0, 768) | k [768, 1536) | v [1536, 2304) | og [2304, 3072) | f_a [3072, 3200) | +b [3200, 3206) | pad [3206, 3208). + +Grid: 26 clusters of 4 CTAs, 256 threads. Cluster k, rank r: + phase 1 (q | k | f_a, 26 tiles of 64 rows): tile k, K chunks [28 r, 28 r + 28) (split-K 4 over the cluster); + phase 2 (v, 12 tiles, then b): cluster k < 24: v tile k // 2, K half k % 2, chunks 56 (k % 2) + [14 r, 14 r + 14) + (split-K 8 = two clusters per tile); clusters 24 and 25: the b tile (8 rows), K half k - 24, likewise; + phase 3 (og, 12 tiles): clusters k < 24 as phase 2; clusters 24 and 25 have no phase 3. +Warps: 0 weight TMA (issued from launch, before the grid dependency, EVICT_FIRST), 1 activation TMA (after it), +2 TMEM allocation and the M = 64, N = 8 MMAs (weight = A, activation = B; each phase accumulates into its own 8 +TMEM columns so a phase's epilogue runs beside the next phase's MMAs), 3 idle, 4-7 epilogue: warp 4 + w loads TMEM +rows [16 w, 16 w + 16) of the tile (16x256b), which rank w of the cluster owns. The other ranks st.async them into +the owner's mailbox; the owner adds the four partials in rank order and publishes them Lamport-style (the data words +are the flags: no fence, no counter, so nothing waits for the stores to be acknowledged under the weight stream): + phase 1: bf16 bits into ``p1`` [3 buffers][8 tokens][1664 rows (q | k | f_a)]; + phases 2 and 3: the cluster's fp32 partial bits into ``part`` [3 buffers][region (v, og, b)][half][8][768]; the + projection is bf16(half 0 + half 1), the consumer's job. +Every word of a buffer holds the sentinel (all ones: a NaN no finite GEMV produces; a computed all-ones word is +stored as the canonical NaN instead) until its producer writes it, so a consumer polls the words it needs until none +is the sentinel. A launch reads e = ``epoch[cta]`` (the CTA's buffer index, its launch count mod 3) after the grid +dependency, writes buffer e and stores the sentinel into the same words of buffer (e + 1) % 3, which the next launch +writes and the launch before last read, then writes (e + 1) % 3 back at its end. The index stays in 0..2: a raw +launch count would turn negative after 2^31 launches and its signed remainder would index before the buffers. +""" + +from __future__ import annotations + +import cuda.bindings.driver as cuda_driver +import cutlass +import cutlass.base_dsl.array +import cutlass.cute as cute +import cutlass.experimental.cuda as cuda +from cutlass._mlir.dialects import llvm as _llvm +from cutlass._mlir.extras import types as _T +from cutlass.cutlass_dsl import dsl_user_op +from cutlass.experimental import primitives as prims + +try: + from cutlass.memory.smem import SmemAllocator +except ImportError: # older DSL layout + from cutlass.utils import SmemAllocator + +from ..k3_kda_verify.k3_kda_verify_kernel import _bf16, _butterfly, _st_async_f32, _store8, _test_wait_cluster + +K_IN = 7168 +HK = 768 # 6 local heads x 128 +PROJ_ROWS = 4 * HK + 128 + 8 # 3208 +V_ROW = 2 * HK +OG_ROW = 3 * HK +FA_ROW = 4 * HK +B_ROW = FA_ROW + 128 +TILE = 64 # weight rows per tile = the MMA's M +MMA_N = 8 +MMA_K = 16 +BOX_K = 64 # bf16 per 128-byte swizzled row +BOX_CH = 2 # 64-column chunks per weight box +BOX_ELEMS = TILE * BOX_K * BOX_CH # 16 KB +X_ELEMS = MMA_N * BOX_K * BOX_CH # 2 KB +STREAM_STAGES = 10 # weight ring depth of the projection alone (B1) +FUSED_STAGES = 12 # B2's stream CTAs: the verify role's buffers alias the same bytes (_SmemCarver) +P3_WINDOW = 3 # phase-3 weight boxes in flight per stream CTA +CLUSTER = 4 +FUSED_RAW_BYTES = ( + FUSED_STAGES * (BOX_ELEMS + X_ELEMS) * 2 + 3 * CLUSTER * 32 * 4 * 4 +) # the stream role's carve; the verify role's fits inside +STREAM_CLUSTERS = 26 +P1_BOXES = K_IN // (BOX_K * BOX_CH) // CLUSTER # 14 +P23_BOXES = P1_BOXES // 2 # 7 +THREADS = 256 +TMEM_COLS = 32 +PART_ROWS = HK +P1_ROWS = 2 * HK + 128 # q | k | f_a +P1_BUF = 8 * P1_ROWS +PART_BUF = 3 * 2 * 8 * PART_ROWS +BUFFERS = 3 +SENTINEL = -1 # all-ones bits +CANON_NAN16 = 0x7FC0 +CANON_NAN32 = 0x7FC00000 +EVICT_FIRST = 0x12F0000000000000 # createpolicy.fractional.L2::evict_first, fraction 1.0 + +# KDA verify role (stage B2): clusters 26-31 = local heads 0-5, cluster rank = V quarter (32 V rows, 4 per warp). +HD = 128 # key and value head dim +H_LOCAL = HK // HD +CONV_W = 4 +NT = 8 # verify tokens: 1 golden + 7 drafts (one request) +NUM_SPEC = NT - 1 +S_COLS = CONV_W - 1 + NUM_SPEC # conv-cache columns +ROWS_U = CONV_W - 1 + NT # raw conv inputs by position: 3 before token 0, then the tokens +V_CTA = HD // CLUSTER # 32 +REC_ROWS = V_CTA // 8 # V rows per warp in the recurrence +VEC = HD // 32 # keys per lane +# The recurrence's registers (one rmem array, static indices): the state [REC_ROWS][VEC], then the previous token's +# per-lane q-dot partials [REC_ROWS], then the current token's operands: v rows [REC_ROWS], q, decay * k, beta * k and +# decay [VEC] each. +R_Q = REC_ROWS * VEC +R_OP = R_Q + REC_ROWS +REC_REGS = R_OP + REC_ROWS + 4 * VEC +# The drafts' records, in each slot's per-token state region (fp32 words from the slot's start) in place of their +# full states: the row innovations vn [NUM_SPEC][H][V], then beta * k and the decay [NUM_SPEC][H][K]. The next launch +# rebuilds the state after the accepted drafts from the pool (the golden token's state) with the update's own +# arithmetic, S = fma(decay, S, vn * (beta * k)) draft by draft, so it is bit-identical to the one the drafts reached. +# k3_kda_verify writes and reads the same records. +CT_VN = 0 +CT_WB = NUM_SPEC * H_LOCAL * HD +CT_WD = CT_WB + NUM_SPEC * H_LOCAL * HD +CT_SLOT = NUM_SPEC * H_LOCAL * HD * HD # a slot's per-token state region +REC_CTA = ( + V_CTA + 2 * HD +) # one draft's records a verify CTA replays: vn of its 32 rows, beta * k and decay of 128 keys +PEND_WORDS = ( + 4 # pending counts a warp loads beside the slot index: pools up to 32 * PEND_WORDS slots +) +CP_CG = cutlass.base_dsl.array.LoadCacheModifier.CG # cp.async 16 B through L2 only +HEAD_CLUSTERS = H_LOCAL +FB_TMEM_COLS = 32 + +# Shared-memory descriptor units (16 bytes). +W_STAGE_U = (BOX_ELEMS * 2) >> 4 +W_CHUNK_U = (TILE * BOX_K * 2) >> 4 +X_STAGE_U = (X_ELEMS * 2) >> 4 +X_CHUNK_U = (MMA_N * BOX_K * 2) >> 4 +K_STEP_U = (MMA_K * 2) >> 4 + + +@dsl_user_op +def _mapa_u32(smem_ptr, peer, *, loc=None, ip=None): + """The shared::cluster address of this CTA's shared-memory location in cluster CTA ``peer``.""" + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [smem_ptr.toint(loc=loc, ip=ip).ir_value(), cutlass.Int32(peer).ir_value(loc=loc, ip=ip)], + "mapa.shared::cluster.u32 $0, $1, $2;", "=r,r,r", has_side_effects=False, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _st_async_v4_f32(dst, a, b, c, d, mbar, *, loc=None, ip=None): + """st.async of four fp32 to a shared::cluster address, completing ``mbar`` (shared::cluster) by 16 bytes.""" + _llvm.inline_asm( + None, + [cutlass.Int32(dst).ir_value(loc=loc, ip=ip), cutlass.Float32(a).ir_value(loc=loc, ip=ip), + cutlass.Float32(b).ir_value(loc=loc, ip=ip), cutlass.Float32(c).ir_value(loc=loc, ip=ip), + cutlass.Float32(d).ir_value(loc=loc, ip=ip), cutlass.Int32(mbar).ir_value(loc=loc, ip=ip)], + "st.async.shared::cluster.mbarrier::complete_tx::bytes.v4.f32 [$0], {$1, $2, $3, $4}, [$5];", "r,f,f,f,f,r", + has_side_effects=True, is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _bf16_bits(x, *, loc=None, ip=None): + """The bf16 rounding (round to nearest even) of an fp32, as its 16 bits in the low half of an int32.""" + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [cutlass.Float32(x).ir_value(loc=loc, ip=ip)], + "{ .reg .b16 h; cvt.rn.bf16.f32 h, $1; cvt.u32.u16 $0, h; }", "=r,f", has_side_effects=False, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _f32_bits(x, *, loc=None, ip=None): + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [cutlass.Float32(x).ir_value(loc=loc, ip=ip)], "mov.b32 $0, $1;", "=r,f", + has_side_effects=False, is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _st_b16(ptr, bits, *, loc=None, ip=None): + """st.global.b16 of the low 16 bits of an int32 at a global address (int64).""" + _llvm.inline_asm( + None, [cutlass.Int64(ptr).ir_value(loc=loc, ip=ip), cutlass.Int32(bits).ir_value(loc=loc, ip=ip)], + "{ .reg .b16 h; cvt.u16.u32 h, $1; st.global.b16 [$0], h; }", "l,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _st_f32(ptr, value, *, loc=None, ip=None): + """st.global.f32 of an fp32 at a global address (int64).""" + _llvm.inline_asm( + None, [cutlass.Int64(ptr).ir_value(loc=loc, ip=ip), cutlass.Float32(value).ir_value(loc=loc, ip=ip)], + "st.global.f32 [$0], $1;", "l,f", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _st_f32_if(ptr, value, pred, *, loc=None, ip=None): + """st.global.f32 of an fp32 at a global address (int64) where pred (int32) is non-zero.""" + _llvm.inline_asm( + None, [cutlass.Int64(ptr).ir_value(loc=loc, ip=ip), cutlass.Float32(value).ir_value(loc=loc, ip=ip), + cutlass.Int32(pred).ir_value(loc=loc, ip=ip)], + "{ .reg .pred p; setp.ne.s32 p, $2, 0; @p st.global.f32 [$0], $1; }", "l,f,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +def _lamport16(bits): + """bf16 bits with the sentinel pattern replaced by the canonical NaN.""" + return cutlass.select_(bits == cutlass.Int32(0xFFFF), cutlass.Int32(CANON_NAN16), bits) + + +def _lamport32(bits): + return cutlass.select_(bits == cutlass.Int32(SENTINEL), cutlass.Int32(CANON_NAN32), bits) + + +@cute.jit +def _publish( + p: cutlass.Constexpr[int], + mbox: cutlass.Array, + mbox_bar: cutlass.Array, + p1: cutlass.Array, + part: cutlass.Array, + a0, + a1, + a2, + a3, # this rank's partials: (r0, t0), (r0, t0 + 1), (r0 + 8, t0), (r0 + 8, t0 + 1) + rank, + lane, + cid, + half, + has_p3, + buf, + nbuf, +): + """Owner warp of rows [16 rank, 16 rank + 16) of the tile in phase p: wait for the three peers' partials, add + the four in rank order, and store them Lamport-style into buffer ``buf`` (phase 1: bf16 rows; phases 2-3: the + cluster's fp32 partial), the sentinel into the same words of buffer ``nbuf``. Clusters without rows in the phase + (24-25 in phase 3) store nothing.""" + while not _test_wait_cluster(mbox_bar.subview(p).data_ptr(), 0): + pass + s0 = cutlass.Float32(0.0) + s1 = cutlass.Float32(0.0) + s2 = cutlass.Float32(0.0) + s3 = cutlass.Float32(0.0) + for src in cutlass.range_constexpr(CLUSTER): + peer4 = mbox.load(idx=((p * CLUSTER + src) * 32 + lane) * 4, vector_size=4, alignment=16) + own = rank == cutlass.Int32(src) + s0 = s0 + cutlass.select_(own, a0, cutlass.Float32(peer4[0])) + s1 = s1 + cutlass.select_(own, a1, cutlass.Float32(peer4[1])) + s2 = s2 + cutlass.select_(own, a2, cutlass.Float32(peer4[2])) + s3 = s3 + cutlass.select_(own, a3, cutlass.Float32(peer4[3])) + r0 = rank * 16 + lane // 4 + t0 = (lane % 4) * 2 + if cutlass.const_expr(p == 0): + # p1 [buffer][token][q | k | f_a]: the tile's rows map to 64 cid + r (cid < 24: q, k; 24-25: f_a). + row_p = cid * cutlass.Int32(TILE) + r0 + cur = buf * cutlass.Int32(P1_BUF) + row_p + nxt = nbuf * cutlass.Int32(P1_BUF) + row_p + o00 = t0 * P1_ROWS + o10 = (t0 + 1) * P1_ROWS + _st_b16(p1.subview(cur + o00).data_ptr().toint(), _lamport16(_bf16_bits(s0))) + _st_b16(p1.subview(cur + o10).data_ptr().toint(), _lamport16(_bf16_bits(s1))) + _st_b16(p1.subview(cur + o00 + 8).data_ptr().toint(), _lamport16(_bf16_bits(s2))) + _st_b16(p1.subview(cur + o10 + 8).data_ptr().toint(), _lamport16(_bf16_bits(s3))) + _st_b16(p1.subview(nxt + o00).data_ptr().toint(), cutlass.Int32(0xFFFF)) + _st_b16(p1.subview(nxt + o10).data_ptr().toint(), cutlass.Int32(0xFFFF)) + _st_b16(p1.subview(nxt + o00 + 8).data_ptr().toint(), cutlass.Int32(0xFFFF)) + _st_b16(p1.subview(nxt + o10 + 8).data_ptr().toint(), cutlass.Int32(0xFFFF)) + else: + # Region 0 = v (phase 2, clusters < 24), 1 = og (phase 3), 2 = b (phase 2, clusters 24-25; rows 0-7). + if cutlass.const_expr(p == 1): + region = cutlass.select_(has_p3, cutlass.Int32(0), cutlass.Int32(2)) + lo_ok = has_p3 | (r0 < cutlass.Int32(8)) + else: + region = cutlass.Int32(1) + lo_ok = has_p3 + row_r = cutlass.select_(has_p3, (cid // cutlass.Int32(2)) * cutlass.Int32(TILE) + r0, r0) + slot = (region * cutlass.Int32(2) + half) * cutlass.Int32(8 * PART_ROWS) + row_r + cur_p = buf * cutlass.Int32(PART_BUF) + slot + nxt_p = nbuf * cutlass.Int32(PART_BUF) + slot + if lo_ok: + part.store(_lamport32(_f32_bits(s0)), idx=cur_p + t0 * PART_ROWS) + part.store(_lamport32(_f32_bits(s1)), idx=cur_p + (t0 + 1) * PART_ROWS) + part.store(cutlass.Int32(SENTINEL), idx=nxt_p + t0 * PART_ROWS) + part.store(cutlass.Int32(SENTINEL), idx=nxt_p + (t0 + 1) * PART_ROWS) + if has_p3: + part.store(_lamport32(_f32_bits(s2)), idx=cur_p + t0 * PART_ROWS + 8) + part.store(_lamport32(_f32_bits(s3)), idx=cur_p + (t0 + 1) * PART_ROWS + 8) + part.store(cutlass.Int32(SENTINEL), idx=nxt_p + t0 * PART_ROWS + 8) + part.store(cutlass.Int32(SENTINEL), idx=nxt_p + (t0 + 1) * PART_ROWS + 8) + + +@cute.jit +def _stream_role( + tma_w, + tma_x, + p1: cutlass.Array, + part: cutlass.Array, + epoch: cutlass.Array, + ring_w: cutlass.Array, + ring_x: cutlass.Array, + full: cutlass.Array, + empty: cutlass.Array, + acc_done: cutlass.Array, + mbox_bar: cutlass.Array, + mbox: cutlass.Array, + tmem_holder: cutlass.Array, + USE_PDL: cutlass.Constexpr[bool], + STAGES: cutlass.Constexpr[int], + lbx=None, +): + """One CTA of a stream cluster (cluster ids 0-25): the three projection phases of its tile rows. ``lbx`` is the + CTA's logical index in the fused grid (the block index when None).""" + tidx, _, _ = cute.arch.thread_idx() + bx = cute.arch.block_idx()[0] if lbx is None else lbx + lane = tidx % 32 + warp = cute.arch.make_warp_uniform(cute.arch.warp_idx()) + rank = cute.arch.block_idx_in_cluster() + # The f_a / b clusters (logical 24-25) take the grid's first cluster slots, so they are resident with the + # first CTAs even when a predecessor still holds SMs: every verify CTA needs their phase 1 (f_a). + cid = (bx // cutlass.Int32(CLUSTER) + cutlass.Int32(STREAM_CLUSTERS - 2)) % cutlass.Int32( + STREAM_CLUSTERS + ) + half = cid % cutlass.Int32(2) + has_p3 = cid < cutlass.Int32(2 * (HK // TILE)) + p1_row = cutlass.select_( + cid < cutlass.Int32(2 * HK // TILE), + cid * cutlass.Int32(TILE), + cutlass.Int32(FA_ROW) + (cid - cutlass.Int32(2 * HK // TILE)) * cutlass.Int32(TILE), + ) + p2_row = cutlass.select_( + has_p3, + cutlass.Int32(V_ROW) + (cid // cutlass.Int32(2)) * cutlass.Int32(TILE), + cutlass.Int32(B_ROW), + ) + p3_row = cutlass.Int32(OG_ROW) + (cid // cutlass.Int32(2)) * cutlass.Int32(TILE) + p1_chunk = rank * cutlass.Int32(P1_BOXES * BOX_CH) + p23_chunk = half * cutlass.Int32(CLUSTER * P23_BOXES * BOX_CH) + rank * cutlass.Int32( + P23_BOXES * BOX_CH + ) + + tma_ptr_w = tma_w.get_ptr() + tma_ptr_x = tma_x.get_ptr() + if warp == 0: + prims.prefetch_tensormap(tma_ptr_w) + prims.prefetch_tensormap(tma_ptr_x) + if prims.elect_sync(): + for s in cutlass.range_constexpr(STAGES): + prims.mbarrier_init(full.subview(s), 2) # the weight and the activation TMA + prims.mbarrier_init(empty.subview(s), 1) + for p in cutlass.range_constexpr(3): + prims.mbarrier_init(acc_done.subview(p), 1) + elif warp == 2: + prims.tcgen05_alloc(tmem_holder, TMEM_COLS) + prims.tcgen05_relinquish_alloc_permit() + elif warp == 3: + if prims.elect_sync(): + for p in cutlass.range_constexpr(3): + prims.mbarrier_init(mbox_bar.subview(p), 1) + prims.mbarrier_arrive_expect_tx(mbox_bar.subview(p), (CLUSTER - 1) * 32 * 16) + prims.fence_mbarrier_init() + # Cluster formation: the peers' mailboxes and their barriers are initialized and addressable from here on. + prims.barrier_cluster_arrive_relaxed() + prims.barrier_cluster_wait() + prims.barrier_cta_sync(0) + + if warp == 0: + # Weight boxes: phase 1 (14), phase 2 (7), phase 3 (7, clusters 0-23). The first STAGES go out at launch. + if prims.elect_sync(): + for i in cutlass.range_constexpr(P1_BOXES + P23_BOXES): + stage_w = i % STAGES + if cutlass.const_expr(i == STAGES and STAGES < P1_BOXES): + # Phase 1's boxes beyond the ring go to L2 now, before the grid dependency, so that their loads + # after it (once the first stages are consumed) hit L2 instead of HBM. + for i_pf in cutlass.range_constexpr(STAGES, P1_BOXES): + prims.cp_async_bulk_tensor_prefetch( + tma_ptr_w, + [cutlass.Int32(0), p1_row, p1_chunk + cutlass.Int32(i_pf * BOX_CH), cutlass.Int32(0), + cutlass.Int32(0)], + [], + ) # fmt: skip + if cutlass.const_expr(i >= STAGES): + while not cute.arch.mbarrier_test_wait( + empty.subview(stage_w).data_ptr(), (i // STAGES + 1) % 2 + ): + pass + if cutlass.const_expr(i < P1_BOXES): + row_w = p1_row + chunk_w = p1_chunk + cutlass.Int32(i * BOX_CH) + else: + row_w = p2_row + chunk_w = p23_chunk + cutlass.Int32((i - P1_BOXES) * BOX_CH) + prims.mbarrier_arrive_expect_tx(full.subview(stage_w), BOX_ELEMS * 2) + prims.cp_async_bulk_tensor_shared_cta_global( + ring_w.subview(stage_w * BOX_ELEMS), tma_ptr_w, + (cutlass.Int32(0), row_w, chunk_w, cutlass.Int32(0), cutlass.Int32(0)), full.subview(stage_w), + l2_cache_hint=EVICT_FIRST, + ) # fmt: skip + if has_p3: + for j in cutlass.range_constexpr(P23_BOXES): + i3 = P1_BOXES + P23_BOXES + j + stage_w3 = i3 % STAGES + # Phase 3 (the output gate, needed only after the recurrence) keeps at most P3_WINDOW boxes in + # flight: box i3 waits for box i3 - P3_WINDOW to be consumed, which also frees its own stage. A + # shallower HBM queue keeps the verify CTAs' phase-2 polls short while phase 3 streams. + jw = i3 - P3_WINDOW + while not cute.arch.mbarrier_test_wait( + empty.subview(jw % STAGES).data_ptr(), (jw // STAGES) % 2 + ): + pass + prims.mbarrier_arrive_expect_tx(full.subview(stage_w3), BOX_ELEMS * 2) + prims.cp_async_bulk_tensor_shared_cta_global( + ring_w.subview(stage_w3 * BOX_ELEMS), tma_ptr_w, + (cutlass.Int32(0), p3_row, p23_chunk + cutlass.Int32(j * BOX_CH), cutlass.Int32(0), + cutlass.Int32(0)), + full.subview(stage_w3), l2_cache_hint=EVICT_FIRST, + ) # fmt: skip + if cutlass.const_expr(USE_PDL): + # Dependents launch once the last weight box is in flight: they cannot take HBM bandwidth from the + # boxes this grid still needs. + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + elif warp == 1: + if cutlass.const_expr(USE_PDL): + prims.griddepcontrol(prims.GridDepAction.WAIT) + if prims.elect_sync(): + for i in cutlass.range_constexpr(P1_BOXES + P23_BOXES): + stage_x = i % STAGES + if cutlass.const_expr(i >= STAGES): + while not cute.arch.mbarrier_test_wait( + empty.subview(stage_x).data_ptr(), (i // STAGES + 1) % 2 + ): + pass + if cutlass.const_expr(i < P1_BOXES): + chunk_x = p1_chunk + cutlass.Int32(i * BOX_CH) + else: + chunk_x = p23_chunk + cutlass.Int32((i - P1_BOXES) * BOX_CH) + prims.mbarrier_arrive_expect_tx(full.subview(stage_x), X_ELEMS * 2) + for c in cutlass.range_constexpr(BOX_CH): + prims.cp_async_bulk_tensor_shared_cta_global( + ring_x.subview(stage_x * X_ELEMS + c * MMA_N * BOX_K), tma_ptr_x, + ((chunk_x + cutlass.Int32(c)) * cutlass.Int32(BOX_K), cutlass.Int32(0)), full.subview(stage_x), + ) # fmt: skip + if has_p3: + for j in cutlass.range_constexpr(P23_BOXES): + i3x = P1_BOXES + P23_BOXES + j + stage_x3 = i3x % STAGES + while not cute.arch.mbarrier_test_wait( + empty.subview(stage_x3).data_ptr(), (i3x // STAGES + 1) % 2 + ): + pass + prims.mbarrier_arrive_expect_tx(full.subview(stage_x3), X_ELEMS * 2) + for c in cutlass.range_constexpr(BOX_CH): + prims.cp_async_bulk_tensor_shared_cta_global( + ring_x.subview(stage_x3 * X_ELEMS + c * MMA_N * BOX_K), tma_ptr_x, + ((p23_chunk + cutlass.Int32(j * BOX_CH + c)) * cutlass.Int32(BOX_K), cutlass.Int32(0)), + full.subview(stage_x3), + ) # fmt: skip + elif warp == 2: + idesc = prims.Tcgen05InstrDesc.build( + c_dtype=cutlass.Float32, + a_dtype=cutlass.BFloat16, + b_dtype=cutlass.BFloat16, + n_dim=MMA_N, + m_dim=TILE, + ) + desc_w = prims.Tcgen05SmemDesc.build( + start_address=ring_w, leading_byte_offset=16, stride_byte_offset=1024, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + desc_x = prims.Tcgen05SmemDesc.build( + start_address=ring_x, leading_byte_offset=16, stride_byte_offset=1024, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + tmem_raw = tmem_holder.load() + for i in cutlass.range_constexpr(P1_BOXES + P23_BOXES): + stage_m = i % STAGES + phase_m = 0 if i < P1_BOXES else 1 + first_m = i == 0 or i == P1_BOXES + while not cute.arch.mbarrier_test_wait( + full.subview(stage_m).data_ptr(), (i // STAGES) % 2 + ): + pass + for kb in cutlass.range_constexpr(BOX_CH * (BOX_K // MMA_K)): + c = kb // (BOX_K // MMA_K) + kk = kb % (BOX_K // MMA_K) + if prims.elect_sync(): + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, prims.CTAGroup.CTA_1, + cutlass.inttoptr(tmem_raw + cutlass.Int32(MMA_N * phase_m), 6, cutlass.Int32), + desc_w + (stage_m * W_STAGE_U + c * W_CHUNK_U + kk * K_STEP_U), + desc_x + (stage_m * X_STAGE_U + c * X_CHUNK_U + kk * K_STEP_U), + idesc, not (first_m and kb == 0), + ) # fmt: skip + if prims.elect_sync(): + prims.tcgen05_commit(empty.subview(stage_m)) + if cutlass.const_expr(i == P1_BOXES - 1 or i == P1_BOXES + P23_BOXES - 1): + prims.tcgen05_commit(acc_done.subview(phase_m)) + if has_p3: + for j in cutlass.range_constexpr(P23_BOXES): + i3m = P1_BOXES + P23_BOXES + j + stage_m3 = i3m % STAGES + while not cute.arch.mbarrier_test_wait( + full.subview(stage_m3).data_ptr(), (i3m // STAGES) % 2 + ): + pass + for kb3 in cutlass.range_constexpr(BOX_CH * (BOX_K // MMA_K)): + c3 = kb3 // (BOX_K // MMA_K) + kk3 = kb3 % (BOX_K // MMA_K) + if prims.elect_sync(): + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, prims.CTAGroup.CTA_1, + cutlass.inttoptr(tmem_raw + cutlass.Int32(2 * MMA_N), 6, cutlass.Int32), + desc_w + (stage_m3 * W_STAGE_U + c3 * W_CHUNK_U + kk3 * K_STEP_U), + desc_x + (stage_m3 * X_STAGE_U + c3 * X_CHUNK_U + kk3 * K_STEP_U), + idesc, not (j == 0 and kb3 == 0), + ) # fmt: skip + if prims.elect_sync(): + prims.tcgen05_commit(empty.subview(stage_m3)) + if cutlass.const_expr(j == P23_BOXES - 1): + prims.tcgen05_commit(acc_done.subview(2)) + else: + # No phase-3 rows: the phase completes empty (its epilogue reduces unused columns and stores nothing). + if prims.elect_sync(): + prims.tcgen05_commit(acc_done.subview(2)) + elif warp >= 4: + w = ( + warp - 4 + ) # TMEM lanes 32 w .. 32 w + 31 = tile rows 16 w .. 16 w + 15, owned by cluster rank w + if cutlass.const_expr(USE_PDL): + prims.griddepcontrol(prims.GridDepAction.WAIT) + e = epoch.load(idx=bx) + buf = e % cutlass.Int32(BUFFERS) + nbuf = (e + cutlass.Int32(1)) % cutlass.Int32(BUFFERS) + tmem_raw_e = tmem_holder.load() + lane_base = (tmem_raw_e >> 16) + w * 32 + col_base = tmem_raw_e & 0xFFFF + for p in cutlass.range_constexpr(3): + while not cute.arch.mbarrier_test_wait(acc_done.subview(p).data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc = prims.tcgen05_ld( + "16x256b", + cutlass.inttoptr((lane_base << 16) | (col_base + p * MMA_N), 6, cutlass.Float32), + num=1, + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + if w != rank: + _st_async_v4_f32( + _mapa_u32(mbox.subview(((p * CLUSTER + rank) * 32 + lane) * 4).data_ptr(), w), + cutlass.Float32(acc[0]), cutlass.Float32(acc[1]), cutlass.Float32(acc[2]), cutlass.Float32(acc[3]), + _mapa_u32(mbox_bar.subview(p).data_ptr(), w), + ) # fmt: skip + else: + _publish(p, mbox, mbox_bar, p1, part, cutlass.Float32(acc[0]), cutlass.Float32(acc[1]), + cutlass.Float32(acc[2]), cutlass.Float32(acc[3]), rank, lane, cid, half, has_p3, buf, + nbuf) # fmt: skip + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + prims.barrier_cta_sync(1, thread_count=128) + if warp == 4: + prims.tcgen05_dealloc(cutlass.inttoptr(tmem_raw_e, 6, cutlass.Int32), TMEM_COLS) + if lane == 0: + epoch.store(nbuf, idx=bx) + + +def _rows_out(shq, lane): + """The warp's 4 row outputs from each lane's q-dot partials: two reduce-scatter rounds (bits 4, 3) then + butterflies (bits 2, 1, 0), the butterfly's pairs in its order (bit-identical); lanes 8 m hold row + 2 (m >> 1) + (m & 1).""" + up4 = (lane & cutlass.Int32(16)) != cutlass.Int32(0) + w0 = cutlass.select_(up4, shq[2], shq[0]) + cute.arch.shuffle_sync_bfly( + cutlass.select_(up4, shq[0], shq[2]), offset=16, mask=-1, mask_and_clamp=31 + ) + w1 = cutlass.select_(up4, shq[3], shq[1]) + cute.arch.shuffle_sync_bfly( + cutlass.select_(up4, shq[1], shq[3]), offset=16, mask=-1, mask_and_clamp=31 + ) + up3 = (lane & cutlass.Int32(8)) != cutlass.Int32(0) + xq = cutlass.select_(up3, w1, w0) + cute.arch.shuffle_sync_bfly( + cutlass.select_(up3, w0, w1), offset=8, mask=-1, mask_and_clamp=31 + ) + for offset in [4, 2, 1]: + xq = xq + cute.arch.shuffle_sync_bfly(xq, offset=offset, mask=-1, mask_and_clamp=31) + return xq + + +def _load_rec_operands(r_st: cutlass.Array, t, warp, lane, s_v, s_q, s_kd, s_bk, s_dec): + """Token t's recurrence operands into ``r_st[R_OP:]`` (v of the warp's rows, then the lane's keys of q, + decay * k, beta * k and decay).""" + v4 = s_v.load(idx=t * V_CTA + warp * REC_ROWS, vector_size=4, alignment=16) + for j in range(REC_ROWS): + r_st.store(cutlass.Float32(v4[j]), idx=R_OP + j) + for i in range(VEC): + c = t * HD + i * 32 + lane + r_st.store(s_q.load(idx=c), idx=R_OP + REC_ROWS + i) + r_st.store(s_kd.load(idx=c), idx=R_OP + REC_ROWS + VEC + i) + r_st.store(s_bk.load(idx=c), idx=R_OP + REC_ROWS + 2 * VEC + i) + r_st.store(s_dec.load(idx=c), idx=R_OP + REC_ROWS + 3 * VEC + i) + + +def _sent16(word): + """Whether either bf16 in an int32 word is the Lamport sentinel (all ones).""" + return ((word & cutlass.Int32(0xFFFF)) == cutlass.Int32(0xFFFF)) | ( + ((word >> 16) & cutlass.Int32(0xFFFF)) == cutlass.Int32(0xFFFF) + ) + + +def _sent32(word): + return word == cutlass.Int32(SENTINEL) + + +@cute.jit +def _poll_partials(part: cutlass.Array, idx, dst: cutlass.Array, dst_idx): + """Spin (relaxed, gpu scope, no sleep) until the four fp32 partial words at part[idx] are published, then store + them into dst[dst_idx : dst_idx + 4].""" + _poll_partials_from(part, idx, dst, dst_idx, cutlass.Int32(SENTINEL), cutlass.Int32(SENTINEL), + cutlass.Int32(SENTINEL), cutlass.Int32(SENTINEL)) # fmt: skip + + +@cute.jit +def _poll_partials_from(part: cutlass.Array, idx, dst: cutlass.Array, dst_idx, w0, w1, w2, w3): + """``_poll_partials`` starting from the words of an earlier read of part[idx]: re-reads only while one of them + is still the sentinel.""" + while _sent32(w0) | _sent32(w1) | _sent32(w2) | _sent32(w3): + vals = prims.load_ext( + part.subview(idx), dtype=cutlass.Int32, count=4, order="relaxed", scope="gpu" + ) + w0 = cutlass.Int32(vals[0]) + w1 = cutlass.Int32(vals[1]) + w2 = cutlass.Int32(vals[2]) + w3 = cutlass.Int32(vals[3]) + dst.store((w0.bitcast(cutlass.Float32), w1.bitcast(cutlass.Float32), w2.bitcast(cutlass.Float32), + w3.bitcast(cutlass.Float32)), idx=dst_idx, alignment=16) # fmt: skip + + +@cute.jit +def _head_role( + tma_wfb, + p1w: cutlass.Array, + part: cutlass.Array, + w_q: cutlass.Array, + w_k: cutlass.Array, + w_v: cutlass.Array, + a_log: cutlass.Array, + dt_bias: cutlass.Array, + onorm_w: cutlass.Array, + cs_q: cutlass.Array, + cs_k: cutlass.Array, + cs_v: cutlass.Array, + ssm: cutlass.Array, + state_tok: cutlass.Array, + slots: cutlass.Array, + pending: cutlass.Array, + out: cutlass.Array, + epoch: cutlass.Array, + smem_a: cutlass.Array, + smem_b: cutlass.Array, + bars: cutlass.Array, + tmem_holder: cutlass.Array, + s_uq: cutlass.Array, + s_uk: cutlass.Array, + s_uv: cutlass.Array, + s_wq: cutlass.Array, + s_wk: cutlass.Array, + s_wv: cutlass.Array, + s_dtb: cutlass.Array, + s_onw: cutlass.Array, + s_gr: cutlass.Array, + s_braw: cutlass.Array, + s_og: cutlass.Array, + s_q: cutlass.Array, + s_k: cutlass.Array, + s_dec: cutlass.Array, + s_kd: cutlass.Array, + s_bk: cutlass.Array, + s_beta: cutlass.Array, + s_v: cutlass.Array, + s_o: cutlass.Array, + s_ss: cutlass.Array, + s_rs: cutlass.Array, + s_vp: cutlass.Array, + s_ogp: cutlass.Array, + s_bp: cutlass.Array, + s_rec: cutlass.Array, + r_st: cutlass.Array, + ssm_stride, + pool_n, + lower_bound: cutlass.Constexpr[float], + scale: cutlass.Constexpr[float], + eps: cutlass.Constexpr[float], + USE_PDL: cutlass.Constexpr[bool], + lbx, +): + """One CTA of a KDA verify cluster (cluster ids 26-31): stage A's verify for local head h = cid - 26 and V rows + [32 q, 32 q + 32) (q = cluster rank), fed by the stream clusters' Lamport buffers of this launch: q, k and f_a + (phase 1), v and b (phase 2), the output gate (phase 3). Arithmetic, rounding and state contract as + ``k3_kda_verify`` (V split 4 instead of 8; the output norm adds the same eight 16-row partial sums in the same + order).""" + tidx, _, _ = cute.arch.thread_idx() + bx = lbx + lane = tidx % 32 + warp = cute.arch.make_warp_uniform(cute.arch.warp_idx()) + vq = cute.arch.block_idx_in_cluster() + h = bx // cutlass.Int32(CLUSTER) - cutlass.Int32(STREAM_CLUSTERS) + v0 = vq * V_CTA + ch0 = h * HD + w_full = bars.subview(0) + acc_done = bars.subview(1) + ss_ready = bars.subview(2) + + tma_ptr_w = tma_wfb.get_ptr() + if warp == 0: + prims.prefetch_tensormap(tma_ptr_w) + if prims.elect_sync(): + prims.mbarrier_init(w_full, 1) + prims.mbarrier_init(acc_done, 1) + elif warp == 2: + prims.tcgen05_alloc(tmem_holder, FB_TMEM_COLS) + prims.tcgen05_relinquish_alloc_permit() + elif warp == 3: + if prims.elect_sync(): + prims.mbarrier_init(ss_ready, 1) + prims.mbarrier_arrive_expect_tx(ss_ready, CLUSTER * 2 * NT * 4) + prims.fence_mbarrier_init() + prims.barrier_cluster_arrive_relaxed() + prims.barrier_cluster_wait() + prims.barrier_cta_sync(0) + if warp == 0: + if prims.elect_sync(): + prims.mbarrier_arrive_expect_tx(w_full, HD * HD * 2) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a, tma_ptr_w, (cutlass.Int32(0), ch0, cutlass.Int32(0), cutlass.Int32(0), cutlass.Int32(0)), + w_full, + ) # fmt: skip + + # ---- Before the grid dependency: nothing here is written by this launch's stream clusters or predecessors. + # Two round trips: first the slot index, every slot's accepted-draft count and what needs neither (conv weights, + # a_log, the norm and gate constants); then, as asynchronous copies into shared memory, everything the slot and + # its count select (the state, the conv caches' last positions and, below, the drafts' records). + slot = slots.load(idx=0) + pend_w = [] + for i_pw in cutlass.range_constexpr(PEND_WORDS): + pw_i = lane + cutlass.Int32(32 * i_pw) + pend_w.append( + cutlass.Int32(pending.load(idx=cutlass.select_(pw_i < pool_n, pw_i, cutlass.Int32(0)))) + ) + a_raw = a_log.load(idx=h) + QH = 3 * HD // 4 + VH = 3 * V_CTA // 4 + if tidx >= 2 * QH + VH + V_CTA: + c4_n = (tidx - (2 * QH + VH + V_CTA)) * 4 + prims.cp_async_shared_global(s_onw.subview(c4_n), onorm_w.subview(v0 + c4_n), 16, CP_CG) + if tidx < HD // 4: + prims.cp_async_shared_global( + s_dtb.subview(tidx * 4), dt_bias.subview(ch0 + tidx * 4), 16, CP_CG + ) + if tidx < HD: + wq4 = w_q.load(idx=(ch0 + tidx) * CONV_W, vector_size=CONV_W, alignment=16) + for w in cutlass.range_constexpr(CONV_W): + s_wq.store(cutlass.Float32(wq4[w]), idx=w * HD + tidx) + else: + ck_w = tidx - HD + wk4 = w_k.load(idx=(ch0 + ck_w) * CONV_W, vector_size=CONV_W, alignment=16) + for w in cutlass.range_constexpr(CONV_W): + s_wk.store(cutlass.Float32(wk4[w]), idx=w * HD + ck_w) + if tidx >= 2 * QH + VH: + if tidx < 2 * QH + VH + V_CTA: + cv_w = tidx - (2 * QH + VH) + wv4 = w_v.load(idx=(ch0 + v0 + cv_w) * CONV_W, vector_size=CONV_W, alignment=16) + for w in cutlass.range_constexpr(CONV_W): + s_wv.store(cutlass.Float32(wv4[w]), idx=w * V_CTA + cv_w) + pend_word = pend_w[PEND_WORDS - 1] + for i_pw in cutlass.range_constexpr(PEND_WORDS - 2, -1, -1): + pend_word = cutlass.Int32( + cutlass.select_(slot < cutlass.Int32(32 * (i_pw + 1)), pend_w[i_pw], pend_word) + ) + pend = cutlass.Int32(cute.arch.shuffle_sync(pend_word, offset=slot % cutlass.Int32(32))) + if pool_n > cutlass.Int32(32 * PEND_WORDS): + pend = cutlass.Int32(pending.load(idx=slot)) + pend = cutlass.Int32( + cutlass.select_(pend > cutlass.Int32(NUM_SPEC), cutlass.Int32(NUM_SPEC), pend) + ) + exp_a = cute.math.exp(a_raw, fastmath=True) + row_w = v0 + warp * REC_ROWS + # The slot's pool state and per-token region from their first element: the slot offset in 64 bits, once. + pool = ssm.subview(cutlass.Int64(slot) * ssm_stride) + tok = state_tok.subview(cutlass.Int64(slot) * CT_SLOT) + st_base = (h * HD + row_w) * HD + for r in cutlass.range_constexpr(REC_ROWS): + for i in cutlass.range_constexpr(VEC): + r_st.store(pool.load(idx=st_base + r * HD + i * 32 + lane), idx=r * VEC + i) + if tidx < QH: + m_q = tidx // (HD // 4) + c4_q = (tidx % (HD // 4)) * 4 + prims.cp_async_shared_global( + s_uq.subview(m_q * HD + c4_q), + cs_q.subview((slot * S_COLS + pend + m_q) * HK + ch0 + c4_q), + 16, + CP_CG, + ) + elif tidx < 2 * QH: + m_k = (tidx - QH) // (HD // 4) + c4_k = ((tidx - QH) % (HD // 4)) * 4 + prims.cp_async_shared_global( + s_uk.subview(m_k * HD + c4_k), + cs_k.subview((slot * S_COLS + pend + m_k) * HK + ch0 + c4_k), + 16, + CP_CG, + ) + elif tidx < 2 * QH + VH: + m_v = (tidx - 2 * QH) // (V_CTA // 4) + c4_v = ((tidx - 2 * QH) % (V_CTA // 4)) * 4 + prims.cp_async_shared_global( + s_uv.subview(m_v * V_CTA + c4_v), + cs_v.subview((slot * S_COLS + pend + m_v) * HK + ch0 + v0 + c4_v), + 16, + CP_CG, + ) + # The drafts the sampler accepted, replayed from their records onto the golden token's state. Every record this + # CTA may need (all NUM_SPEC drafts) comes into shared memory by one round trip first; a draft per iteration from + # global memory was one dependent round trip per accepted draft. + for it_rc in cutlass.range_constexpr((NUM_SPEC * REC_CTA // 4 + THREADS - 1) // THREADS): + q_rc = tidx + it_rc * THREADS + if q_rc < NUM_SPEC * REC_CTA // 4: + t_rc = q_rc // (REC_CTA // 4) + u_rc = (q_rc % (REC_CTA // 4)) * 4 + rec_rc = (t_rc * H_LOCAL + h) * HD + src_rc = cutlass.Int32( + cutlass.select_( + u_rc < V_CTA, + rec_rc + CT_VN + v0 + u_rc, + cutlass.select_( + u_rc < V_CTA + HD, + rec_rc + CT_WB + (u_rc - V_CTA), + rec_rc + CT_WD + (u_rc - V_CTA - HD), + ), + ) + ) + prims.cp_async_shared_global( + s_rec.subview(t_rc * REC_CTA + u_rc), tok.subview(src_rc), 16, CP_CG + ) + prims.cp_async_commit_group() + prims.cp_async_wait_group(0) + prims.barrier_cta_sync(0) + for t_acc in cutlass.range(pend, unroll=1): + rec_s = t_acc * REC_CTA + vn4_acc = s_rec.load(idx=rec_s + warp * REC_ROWS, vector_size=4, alignment=16) + vns_acc = [cutlass.Float32(vn4_acc[r]) for r in range(REC_ROWS)] + wbs_acc = [s_rec.load(idx=rec_s + V_CTA + i * 32 + lane) for i in range(VEC)] + wds_acc = [s_rec.load(idx=rec_s + V_CTA + HD + i * 32 + lane) for i in range(VEC)] + sts_acc = [r_st.load(idx=j) for j in range(REC_ROWS * VEC)] + for r in cutlass.range_constexpr(REC_ROWS): + for _pi in cutlass.range_constexpr(VEC // 2): + _p = _pi * 2 + vb0_acc, vb1_acc = cute.arch.mul_packed_f32x2( + (vns_acc[r], vns_acc[r]), (wbs_acc[_p], wbs_acc[_p + 1]) + ) + sts_acc[r * VEC + _p], sts_acc[r * VEC + _p + 1] = cute.arch.fma_packed_f32x2( + src_a=(wds_acc[_p], wds_acc[_p + 1]), src_b=(sts_acc[r * VEC + _p], sts_acc[r * VEC + _p + 1]), + src_c=(vb0_acc, vb1_acc), + ) # fmt: skip + for j in cutlass.range_constexpr(REC_ROWS * VEC): + r_st.store(sts_acc[j], idx=j) + # Every CTA of the head has read the head's q / k conv caches and per-key records once all four arrive here + # (release; the acquiring wait follows phase 1): from then on each CTA rewrites its quarter of both, in the idle + # and pre-recurrence slots instead of after the output norm. + prims.barrier_cluster_arrive() + + if cutlass.const_expr(USE_PDL): + prims.griddepcontrol(prims.GridDepAction.WAIT) + e = epoch.load(idx=bx) + buf = e % cutlass.Int32(BUFFERS) + if cutlass.const_expr(USE_PDL): + # The dependents may launch once the stream CTAs have also triggered (after their last weight box): they + # then take the SMs the stream leaves and run their pre-wait work beside the verify instead of after it. + if tidx == 0: + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + + # ---- Phase 1 (q, k, f_a of this launch): thread i polls 16-byte chunk i of the q | k rows, threads < 128 also + # chunk i of the f_a rows; q and k go to the conv inputs as fp32, f_a to the MMA's B operand (128-byte swizzle). + st_qk = tidx // HD + t_qk = (tidx % HD) // 16 + j_qk = tidx % 16 + a0 = cutlass.Int32(SENTINEL) + a1 = cutlass.Int32(SENTINEL) + a2 = cutlass.Int32(SENTINEL) + a3 = cutlass.Int32(SENTINEL) + qk_idx = ((buf * NT + t_qk) * P1_ROWS + st_qk * HK + ch0 + j_qk * 8) // 2 + while _sent16(a0) | _sent16(a1) | _sent16(a2) | _sent16(a3): + vqk = prims.load_ext( + p1w.subview(qk_idx), dtype=cutlass.Int32, count=4, order="relaxed", scope="gpu" + ) + a0 = cutlass.Int32(vqk[0]) + a1 = cutlass.Int32(vqk[1]) + a2 = cutlass.Int32(vqk[2]) + a3 = cutlass.Int32(vqk[3]) + if st_qk == 0: + _store8(s_uq, (a0, a1, a2, a3), (CONV_W - 1 + t_qk) * HD + j_qk * 8) + else: + _store8(s_uk, (a0, a1, a2, a3), (CONV_W - 1 + t_qk) * HD + j_qk * 8) + if tidx < HD: + t_fa = tidx // 16 + j_fa = tidx % 16 + f0 = cutlass.Int32(SENTINEL) + f1 = cutlass.Int32(SENTINEL) + f2 = cutlass.Int32(SENTINEL) + f3 = cutlass.Int32(SENTINEL) + fa_idx = ((buf * NT + t_fa) * P1_ROWS + 2 * HK + j_fa * 8) // 2 + while _sent16(f0) | _sent16(f1) | _sent16(f2) | _sent16(f3): + vfa = prims.load_ext( + p1w.subview(fa_idx), dtype=cutlass.Int32, count=4, order="relaxed", scope="gpu" + ) + f0 = cutlass.Int32(vfa[0]) + f1 = cutlass.Int32(vfa[1]) + f2 = cutlass.Int32(vfa[2]) + f3 = cutlass.Int32(vfa[3]) + # Row t_fa, 16-byte unit u of 64-column chunk c: byte 1024 c + 128 t_fa + 16 (u ^ t_fa). + sw_word = ((j_fa // 8) * 1024 + t_fa * 128 + ((j_fa % 8) ^ t_fa) * 16) // 4 + smem_b.store((f0, f1, f2, f3), idx=sw_word, alignment=16) + prims.fence_proxy("async_shared", space=prims.SharedSpace.shared_cta) + prims.barrier_cta_sync(0) + prims.barrier_cluster_wait() + + # ---- f_b on the tensor cores (warp 2 issues, warps 4-7 read TMEM), then the q/k pre-compute, one token per warp. + if warp == 2: + idesc = prims.Tcgen05InstrDesc.build( + c_dtype=cutlass.Float32, + a_dtype=cutlass.BFloat16, + b_dtype=cutlass.BFloat16, + n_dim=MMA_N, + m_dim=HD, + ) + desc_a_base = prims.Tcgen05SmemDesc.build( + start_address=smem_a, leading_byte_offset=16, stride_byte_offset=1024, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + desc_b_base = prims.Tcgen05SmemDesc.build( + start_address=smem_b, leading_byte_offset=16, stride_byte_offset=1024, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + tmem_acc = cutlass.inttoptr(tmem_holder.load(), 6, cutlass.Int32) + while not cute.arch.mbarrier_test_wait(w_full.data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) # f_a came from the other threads' smem stores + for kb in cutlass.range_constexpr(2 * (BOX_K // MMA_K)): + box = kb // (BOX_K // MMA_K) + within = kb % (BOX_K // MMA_K) + if prims.elect_sync(): + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, prims.CTAGroup.CTA_1, tmem_acc, + desc_a_base + (box * ((HD * BOX_K * 2) >> 4) + within * K_STEP_U), + desc_b_base + (box * ((MMA_N * BOX_K * 2) >> 4) + within * K_STEP_U), + idesc, kb != 0, + ) # fmt: skip + if prims.elect_sync(): + prims.tcgen05_commit(acc_done) + if warp >= 4: + while not cute.arch.mbarrier_test_wait(acc_done.data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc = prims.tcgen05_ld( + "32x32b", cutlass.inttoptr(tmem_holder.load(), 6, cutlass.Float32), num=MMA_N + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + c_fb = (warp - 4) * 32 + lane + for t_fb in cutlass.range_constexpr(NT): + s_gr.store(_bf16(cutlass.Float32(acc[t_fb])), idx=t_fb * HD + c_fb) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + prims.barrier_cta_sync(1, thread_count=128) + if warp == 4: + prims.tcgen05_dealloc( + cutlass.inttoptr(tmem_holder.load(), 6, cutlass.Int32), FB_TMEM_COLS + ) + # q and k of token w, one token per warp (the arithmetic of kda_mtp_decode's pre-compute warps). + tk = warp + pq = [cutlass.Float32(0.0)] * VEC + for i in cutlass.range_constexpr(VEC): + c = i * 32 + lane + conv = cutlass.Float32(0.0) + for w in cutlass.range_constexpr(CONV_W - 1): + conv += s_uq.load(idx=(tk + w) * HD + c) * s_wq.load(idx=w * HD + c) + conv += s_uq.load(idx=(tk + CONV_W - 1) * HD + c) * s_wq.load(idx=(CONV_W - 1) * HD + c) + ex = cute.math.exp(-conv, fastmath=True) + pq[i] = conv * cute.arch.rcp_approx(cutlass.Float32(1.0) + ex) + sum_q = cutlass.Float32(0.0) + for i in cutlass.range_constexpr(VEC): + sum_q += pq[i] * pq[i] + sum_q = _butterfly(sum_q) + rnorm_q = cute.math.rsqrt(sum_q + 1e-06, fastmath=True) * scale + for i in cutlass.range_constexpr(VEC): + s_q.store(pq[i] * rnorm_q, idx=tk * HD + i * 32 + lane) + pk = [cutlass.Float32(0.0)] * VEC + for i in cutlass.range_constexpr(VEC): + c = i * 32 + lane + conv = s_uk.load(idx=tk * HD + c) * s_wk.load(idx=c) + for w in cutlass.range_constexpr(1, CONV_W - 1): + conv += s_uk.load(idx=(tk + w) * HD + c) * s_wk.load(idx=w * HD + c) + conv += s_uk.load(idx=(tk + CONV_W - 1) * HD + c) * s_wk.load(idx=(CONV_W - 1) * HD + c) + pk[i] = conv * cute.arch.rcp_approx( + cutlass.Float32(1.0) + cute.math.exp(-conv, fastmath=True) + ) + sum_k = cutlass.Float32(0.0) + for i in cutlass.range_constexpr(VEC): + sum_k += pk[i] * pk[i] + sum_k = _butterfly(sum_k) + rnorm_k = cute.math.rsqrt(sum_k + 1e-06, fastmath=True) + for i in cutlass.range_constexpr(VEC): + s_k.store(pk[i] * rnorm_k, idx=tk * HD + i * 32 + lane) + if warp < 4: + # Conv caches for the next round, q and k, this CTA's 32 channels: column s = position s - 2 = row s + 1 of + # the position-indexed inputs. + t_cq = tidx + for it_qk in cutlass.range_constexpr((2 * S_COLS * V_CTA + 127) // 128): + item_qk = t_cq + it_qk * 128 + if item_qk < 2 * S_COLS * V_CTA: + sc_qk = (item_qk % (S_COLS * V_CTA)) // V_CTA + ch_qk = vq * V_CTA + item_qk % V_CTA + dst_qk = (slot * S_COLS + sc_qk) * HK + ch0 + ch_qk + if item_qk < S_COLS * V_CTA: + cs_q.store(s_uq.load(idx=(sc_qk + 1) * HD + ch_qk), idx=dst_qk) + else: + cs_k.store(s_uk.load(idx=(sc_qk + 1) * HD + ch_qk), idx=dst_qk) + prims.barrier_cta_sync(0) + + # ---- Phase 2 (v, b): the two K-half partials of this CTA's 32 v rows and of this head's b; the projection is + # their bf16-rounded sum. + if tidx < 2 * NT * (V_CTA // 4): + half_v = tidx // (NT * (V_CTA // 4)) + t_v = (tidx // (V_CTA // 4)) % NT + j_v = tidx % (V_CTA // 4) + _poll_partials( + part, + ((buf * 3 + 0) * 2 + half_v) * (8 * PART_ROWS) + t_v * PART_ROWS + ch0 + v0 + j_v * 4, + s_vp, + (half_v * NT + t_v) * V_CTA + j_v * 4, + ) + elif tidx < 2 * NT * (V_CTA // 4) + 2 * NT: + ib = tidx - 2 * NT * (V_CTA // 4) + half_b = ib // NT + t_b = ib % NT + wb = cutlass.Int32(SENTINEL) + while _sent32(wb): + wb = prims.load_ext( + part.subview(((buf * 3 + 2) * 2 + half_b) * (8 * PART_ROWS) + t_b * PART_ROWS + h), + dtype=cutlass.Int32, + order="relaxed", + scope="gpu", + ) + s_bp.store(wb.bitcast(cutlass.Float32), idx=half_b * NT + t_b) + prims.barrier_cta_sync(0) + t_cv = tidx // V_CTA + c_cv = tidx % V_CTA + s_uv.store(_bf16(s_vp.load(idx=t_cv * V_CTA + c_cv) + s_vp.load(idx=(NT + t_cv) * V_CTA + c_cv)), + idx=(CONV_W - 1 + t_cv) * V_CTA + c_cv) # fmt: skip + if tidx < NT: + s_braw.store(_bf16(s_bp.load(idx=tidx) + s_bp.load(idx=NT + tidx)), idx=tidx) + prims.barrier_cta_sync(0) + for it_cs in cutlass.range_constexpr((S_COLS * V_CTA + THREADS - 1) // THREADS): + item_cs = tidx + it_cs * THREADS + if item_cs < S_COLS * V_CTA: + s_cs = item_cs // V_CTA + v_cs = item_cs % V_CTA + cs_v.store( + s_uv.load(idx=(s_cs + 1) * V_CTA + v_cs), + idx=(slot * S_COLS + s_cs) * HK + ch0 + v0 + v_cs, + ) + # v (this CTA's 32 channels, one (token, channel) per thread) and beta. + conv_v = cutlass.Float32(0.0) + for w in cutlass.range_constexpr(CONV_W): + conv_v += s_uv.load(idx=(t_cv + w) * V_CTA + c_cv) * s_wv.load(idx=w * V_CTA + c_cv) + conv_v = conv_v * cute.arch.rcp_approx( + cutlass.Float32(1.0) + cute.math.exp(-conv_v, fastmath=True) + ) + s_v.store(conv_v, idx=t_cv * V_CTA + c_cv) + if tidx < NT: + b_pre = s_braw.load(idx=tidx) + s_beta.store( + cute.arch.rcp_approx(cutlass.Float32(1.0) + cute.math.exp(-b_pre, fastmath=True)), + idx=tidx, + ) + prims.barrier_cta_sync(0) + + # ---- The recurrence's per-token operands: warp t, token t: decay = exp(gate), decay * k, beta * k. + tg = warp + r_beta_g = s_beta.load(idx=tg) + for i_pair in cutlass.range_constexpr(VEC // 2): + dks = [] + for i in (i_pair * 2, i_pair * 2 + 1): + c = i * 32 + lane + g_raw = s_gr.load(idx=tg * HD + c) + s_dtb.load(idx=c) + xg = exp_a * g_raw + sig = cute.arch.rcp_approx(cutlass.Float32(1.0) + cute.math.exp(-xg, fastmath=True)) + dks.append( + (cute.math.exp(lower_bound * sig, fastmath=True), s_k.load(idx=tg * HD + c), c) + ) + (d0, k0v, c0), (d1, k1v, c1) = dks + bk0, bk1 = cute.arch.mul_packed_f32x2((r_beta_g, r_beta_g), (k0v, k1v)) + kd0, kd1 = cute.arch.mul_packed_f32x2((d0, d1), (k0v, k1v)) + s_dec.store(d0, idx=tg * HD + c0) + s_dec.store(d1, idx=tg * HD + c1) + s_kd.store(kd0, idx=tg * HD + c0) + s_kd.store(kd1, idx=tg * HD + c1) + s_bk.store(bk0, idx=tg * HD + c0) + s_bk.store(bk1, idx=tg * HD + c1) + prims.barrier_cta_sync(0) + # The drafts' per-key records (beta * k and the decay of tokens 1-7), this CTA's 32 keys. + if tidx < NUM_SPEC * V_CTA: + t_rq = tidx // V_CTA + key_rq = vq * V_CTA + tidx % V_CTA + rec_q = h * HD + t_rq * H_LOCAL * HD + key_rq + tok.store(s_bk.load(idx=(t_rq + 1) * HD + key_rq), idx=rec_q + CT_WB) + tok.store(s_dec.load(idx=(t_rq + 1) * HD + key_rq), idx=rec_q + CT_WD) + + # ---- Recurrence over the verify tokens: warp w owns V rows row_w .. row_w + 3, lane l keys 32 i + l. A rolled + # loop: in the model the MoE's stream has evicted this kernel's code from L2, and the 8 tokens fully unrolled + # (~1,100 instructions, run once per launch) stall on instruction fetch; one token's body is fetched once. The + # loop is software-pipelined so that consecutive tokens overlap: iteration tr reduces token tr - 1's q-dot + # partials over the warp beside token tr's k-dot butterflies (the chain that carries the state) and loads token + # tr + 1's operands. Every sum is formed from the same values in the same order as before (bit-identical). + # This lane's first state element in the pool (the golden token's state) and in the first draft's per-token + # state; element (r, i) sits (r HD + 32 i) floats past either. + pool_ptr = pool.subview(st_base + lane).data_ptr().toint() + # Lanes 0-3 store the warp's row innovations, one row each (a lane's address beyond them is never stored to). + vn_ptr = tok.subview(CT_VN + h * HD + row_w + lane).data_ptr().toint() + _load_rec_operands(r_st, cutlass.Int32(0), warp, lane, s_v, s_q, s_kd, s_bk, s_dec) + for r in cutlass.range_constexpr(REC_ROWS): + r_st.store(cutlass.Float32(0.0), idx=R_Q + r) + for tr in cutlass.range(NT, unroll=1): + vrows = [r_st.load(idx=R_OP + j) for j in range(REC_ROWS)] + wq = [r_st.load(idx=R_OP + REC_ROWS + i) for i in range(VEC)] + wk = [r_st.load(idx=R_OP + REC_ROWS + VEC + i) for i in range(VEC)] + wb = [r_st.load(idx=R_OP + REC_ROWS + 2 * VEC + i) for i in range(VEC)] + wd = [r_st.load(idx=R_OP + REC_ROWS + 3 * VEC + i) for i in range(VEC)] + shq_prev = [r_st.load(idx=R_Q + r) for r in range(REC_ROWS)] + # Token tr + 1's operands (the last iteration reloads token NT - 1's and does not use them). + t_next = cutlass.Int32(cutlass.select_(tr + 1 < NT, tr + 1, cutlass.Int32(NT - 1))) + _load_rec_operands(r_st, t_next, warp, lane, s_v, s_q, s_kd, s_bk, s_dec) + stw = [r_st.load(idx=j) for j in range(REC_ROWS * VEC)] + xq_prev = _rows_out(shq_prev, lane) + shk = [] + for r in cutlass.range_constexpr(REC_ROWS): + p1a = cutlass.Float32(0.0) + p2a = cutlass.Float32(0.0) + for _pi in cutlass.range_constexpr(VEC // 2): + _p = _pi * 2 + p1a, p2a = cute.arch.fma_packed_f32x2( + src_a=(stw[r * VEC + _p], stw[r * VEC + _p + 1]), + src_b=(wk[_p], wk[_p + 1]), + src_c=(p1a, p2a), + ) + shk.append(p1a + p2a) + for offset in [16, 8, 4, 2, 1]: + for r in cutlass.range_constexpr(REC_ROWS): + shk[r] += cute.arch.shuffle_sync_bfly( + shk[r], offset=offset, mask=-1, mask_and_clamp=31 + ) + shq = [] + vns = [] + for r in cutlass.range_constexpr(REC_ROWS): + vn = vrows[r] - shk[r] + vns.append(vn) + q1 = cutlass.Float32(0.0) + q2 = cutlass.Float32(0.0) + for _pi in cutlass.range_constexpr(VEC // 2): + _p = _pi * 2 + vb0, vb1 = cute.arch.mul_packed_f32x2((vn, vn), (wb[_p], wb[_p + 1])) + stw[r * VEC + _p], stw[r * VEC + _p + 1] = cute.arch.fma_packed_f32x2( + src_a=(wd[_p], wd[_p + 1]), src_b=(stw[r * VEC + _p], stw[r * VEC + _p + 1]), src_c=(vb0, vb1), + ) # fmt: skip + q1, q2 = cute.arch.fma_packed_f32x2( + src_a=(stw[r * VEC + _p], stw[r * VEC + _p + 1]), + src_b=(wq[_p], wq[_p + 1]), + src_c=(q1, q2), + ) + shq.append(q1 + q2) + # The golden token's state to the pool (tr = 0); a draft's record: this warp's row innovations (lanes 0-3). + # Predicated stores, no branch: they issue between the update's FMAs. + first = cutlass.Int32(cutlass.select_(tr == 0, cutlass.Int32(1), cutlass.Int32(0))) + for r in cutlass.range_constexpr(REC_ROWS): + for i in cutlass.range_constexpr(VEC): + _st_f32_if(pool_ptr + cutlass.Int64((r * HD + i * 32) * 4), stw[r * VEC + i], first) + vn_lane = cutlass.Float32( + cutlass.select_( + lane == 0, + vns[0], + cutlass.select_(lane == 1, vns[1], cutlass.select_(lane == 2, vns[2], vns[3])), + ) + ) + draft_rec = cutlass.Int32( + cutlass.select_((tr > 0) & (lane < 4), cutlass.Int32(1), cutlass.Int32(0)) + ) + _st_f32_if( + vn_ptr + cutlass.Int64(tr - 1) * cutlass.Int64(H_LOCAL * HD * 4), vn_lane, draft_rec + ) + for j in cutlass.range_constexpr(REC_ROWS * VEC): + r_st.store(stw[j], idx=j) + for r in cutlass.range_constexpr(REC_ROWS): + r_st.store(shq[r], idx=R_Q + r) + # Token tr - 1's outputs, from the first lane of each group of 8 (which all hold the same sum); at tr = 0 a + # placeholder in token 0's slot, which iteration 1 overwrites from the same lanes. + o_tok = cutlass.Int32(cutlass.select_(tr > 0, tr - 1, cutlass.Int32(0))) + if (lane & cutlass.Int32(7)) == cutlass.Int32(0): + s_o.store( + _bf16(xq_prev), + idx=o_tok * V_CTA + warp * REC_ROWS + (lane >> 4) * 2 + ((lane >> 3) & 1), + ) + xq_last = _rows_out([r_st.load(idx=R_Q + r) for r in range(REC_ROWS)], lane) + if (lane & cutlass.Int32(7)) == cutlass.Int32(0): + s_o.store( + _bf16(xq_last), idx=(NT - 1) * V_CTA + warp * REC_ROWS + (lane >> 4) * 2 + ((lane >> 3) & 1) + ) + + # ---- Phase 3 (the output gate), then the gated RMSNorm over V: every CTA sends its two 16-row sums of squares + # per token to all four CTAs (the eight partials of V split 8, added in the same order). Thread (token t, row v) + # reads its own gate's two K-half partials first, so their round trip overlaps the norm exchange. + t_out = tidx // V_CTA + v_y = tidx % V_CTA + og_at = (buf * 3 + 1) * 2 * (8 * PART_ROWS) + t_out * PART_ROWS + ch0 + v0 + v_y + og_h0 = cutlass.Int32( + prims.load_ext(part.subview(og_at), dtype=cutlass.Int32, order="relaxed", scope="gpu") + ) + og_h1 = cutlass.Int32(prims.load_ext(part.subview(og_at + 8 * PART_ROWS), dtype=cutlass.Int32, order="relaxed", + scope="gpu")) # fmt: skip + prims.barrier_cta_sync(0) + x_o = s_o.load(idx=tidx) + ss = x_o * x_o + for off_ss in [8, 4, 2, 1]: + ss = ss + cute.arch.shuffle_sync_bfly(ss, offset=off_ss, mask=-1, mask_and_clamp=31) + if lane % 16 == 0: + src_slot = vq * 2 + lane // 16 + for r in cutlass.range_constexpr(CLUSTER): + _st_async_f32( + _mapa_u32(s_ss.subview(src_slot * NT + t_out).data_ptr(), r), + ss, + _mapa_u32(ss_ready.data_ptr(), r), + ) + while not _test_wait_cluster(ss_ready.data_ptr(), 0): + pass + while _sent32(og_h0) | _sent32(og_h1): + og_h0 = cutlass.Int32( + prims.load_ext(part.subview(og_at), dtype=cutlass.Int32, order="relaxed", scope="gpu") + ) + og_h1 = cutlass.Int32(prims.load_ext(part.subview(og_at + 8 * PART_ROWS), dtype=cutlass.Int32, + order="relaxed", scope="gpu")) # fmt: skip + if tidx < NT: + total = cutlass.Float32(0.0) + for r in cutlass.range_constexpr(2 * CLUSTER): + total = total + s_ss.load(idx=r * NT + tidx) + s_rs.store(cute.math.rsqrt(total / HD + eps), idx=tidx) + prims.barrier_cta_sync(0) + z = _bf16(og_h0.bitcast(cutlass.Float32) + og_h1.bitcast(cutlass.Float32)) + gate = cutlass.Float32(1.0) / (cutlass.Float32(1.0) + cute.math.exp(-z, fastmath=True)) + y = s_o.load(idx=tidx) * s_rs.load(idx=t_out) * s_onw.load(idx=v_y) * gate + out.store(cutlass.BFloat16(y), idx=(t_out * H_LOCAL + h) * HD + v0 + v_y) + + if tidx == 0: + epoch.store((e + cutlass.Int32(1)) % cutlass.Int32(BUFFERS), idx=bx) + + +@cute.kernel +def k3_kda_attn_kernel( + tma_w: cutlass.GridConstant[ + cuda.TensorMap + ], # W [3208, 7168] bf16, 5-D, box 64 cols x 64 rows x 2 chunks + tma_x: cutlass.GridConstant[ + cuda.TensorMap + ], # x [T, 7168] bf16, box 64 cols x 8 rows (rows >= T read as 0) + p1: cutlass.Array, # int16 bits of bf16 [3][8][1664]: phase-1 rows (Lamport) + part: cutlass.Array, # int32 bits of fp32 [3][3][2][8][768]: v, og and b cluster partials (Lamport) + epoch: cutlass.Array, # int32 [CTAs]: each CTA's buffer index (launches completed mod 3) + USE_PDL: cutlass.Constexpr[bool], +): + ring_w = cutlass.Array( + cutlass.BFloat16, STREAM_STAGES * BOX_ELEMS, space=cutlass.AddressSpace.smem, alignment=1024 + ) + ring_x = cutlass.Array( + cutlass.BFloat16, STREAM_STAGES * X_ELEMS, space=cutlass.AddressSpace.smem, alignment=1024 + ) + full = cutlass.Array(cutlass.Int64, STREAM_STAGES, space=cutlass.AddressSpace.smem, alignment=8) + empty = cutlass.Array( + cutlass.Int64, STREAM_STAGES, space=cutlass.AddressSpace.smem, alignment=8 + ) + acc_done = cutlass.Array(cutlass.Int64, 3, space=cutlass.AddressSpace.smem, alignment=8) + mbox_bar = cutlass.Array(cutlass.Int64, 3, space=cutlass.AddressSpace.smem, alignment=8) + # [phase][source rank][lane][4]: the partials of the 16 rows this rank owns (its own slot stays unused). + mbox = cutlass.Array( + cutlass.Float32, 3 * CLUSTER * 32 * 4, space=cutlass.AddressSpace.smem, alignment=16 + ) + tmem_holder = cutlass.Array(cutlass.Int32, 1, space=cutlass.AddressSpace.smem) + + _stream_role( + tma_w, + tma_x, + p1, + part, + epoch, + ring_w, + ring_x, + full, + empty, + acc_done, + mbox_bar, + mbox, + tmem_holder, + USE_PDL, + STREAM_STAGES, + ) + + +@cute.jit +def k3_kda_qkvg( + w: cute.Tensor, # bf16 [3208, 7168] + x: cute.Tensor, # bf16 [T, 7168] + p1: cute.Tensor, # int16 [3 * 8 * 1664] + part: cute.Tensor, # int32 [3 * 3 * 2 * 8 * 768] + epoch: cute.Tensor, # int32 [104] + num_tokens: cutlass.Int32, + USE_PDL: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + """One launch of the three-phase projection (the stream role alone).""" + tma_w = cuda.create_tensor_map_tiled( + global_address=w.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[BOX_K, PROJ_ROWS, K_IN // BOX_K, 1, 1], + global_strides=[(K_IN * 2) // 16, (BOX_K * 2) // 16, (PROJ_ROWS * K_IN * 2) // 16, + (PROJ_ROWS * K_IN * 2) // 16], + box_dims=[BOX_K, TILE, BOX_CH, 1, 1], + swizzle=cuda.TensorMapSwizzle.s128b, + ) # fmt: skip + tma_x = cuda.create_tensor_map_tiled( + global_address=x.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[K_IN, num_tokens], + global_strides=[(K_IN * 2) // 16], + box_dims=[BOX_K, MMA_N], + swizzle=cuda.TensorMapSwizzle.s128b, + ) + k3_kda_attn_kernel(tma_w, tma_x, p1, part, epoch, USE_PDL).launch( + grid=(STREAM_CLUSTERS * CLUSTER, 1, 1), + block=(THREADS, 1, 1), + cluster=(CLUSTER, 1, 1), + stream=stream, + use_pdl=USE_PDL, + ) + + +class _SmemCarver: + """Typed views laid out one after another inside one raw shared-memory buffer: a CTA runs one role for its whole + life and every CTA of a cluster runs the same role, so the two roles share the bytes and a DSMEM address always + names the same view in the peer.""" + + def __init__(self, raw, raw_bytes: int): + self.raw = raw + self.cap = raw_bytes + self.off = 0 + + def view(self, dtype, n: int, align: int = 16) -> cutlass.Array: + self.off = -(-self.off // align) * align + base = self.raw if self.off == 0 else self.raw + self.off + ptr = cute.recast_ptr(base, dtype=dtype) + self.off += n * dtype.width // 8 + assert self.off <= self.cap, "a role's arrays exceed the shared raw buffer" + return cutlass.Array( + ptr, + shape=(n,), + dtype=dtype, + bounds_check=False, + addrspace=cutlass.AddressSpace.smem.value, + alignment=align, + ) + + +@cute.kernel +def k3_kda_attn_fused_kernel( + tma_w: cutlass.GridConstant[ + cuda.TensorMap + ], # W [3208, 7168] bf16, 5-D, box 64 cols x 64 rows x 2 chunks + tma_x: cutlass.GridConstant[cuda.TensorMap], # x [8, 7168] bf16, box 64 cols x 8 rows + tma_wfb: cutlass.GridConstant[ + cuda.TensorMap + ], # W_fb [768, 128] bf16, a head's 128 rows by one call + p1: cutlass.Array, # int16 bits of bf16 [3][8][1664] (Lamport, written by the stream role) + p1w: cutlass.Array, # the same memory as int32 words (polled by the verify role) + part: cutlass.Array, # int32 bits of fp32 [3][3][2][8][768] (Lamport) + epoch: cutlass.Array, # int32 [128]: each CTA's buffer index (launches completed mod 3) + w_q: cutlass.Array, # fp32 [768, 4] + w_k: cutlass.Array, + w_v: cutlass.Array, + a_log: cutlass.Array, # fp32 [6] + dt_bias: cutlass.Array, # fp32 [768] + onorm_w: cutlass.Array, # fp32 [128] + cs_q: cutlass.Array, # fp32 [pool][10][768]: raw conv inputs, column s = position s - 2 from the golden token + cs_k: cutlass.Array, + cs_v: cutlass.Array, + ssm: cutlass.Array, # fp32 [pool][6][128][128]: the state after the last golden token + state_tok: cutlass.Array, # fp32 [pool][7][6][128][128]: the state after each draft of the last round + slots: cutlass.Array, # int32 [1] + pending: cutlass.Array, # int32 [pool]: drafts the sampler accepted last round + out: cutlass.Array, # bf16 [8][6][128] + ssm_stride: cutlass.Int64, # fp32 elements between the pool's slots + pool_n: cutlass.Int32, # the pool's slots (entries of pending) + lower_bound: cutlass.Constexpr[float], + scale: cutlass.Constexpr[float], + eps: cutlass.Constexpr[float], + USE_PDL: cutlass.Constexpr[bool], +): + # The stream CTAs (clusters 0-25) and the verify CTAs (26-31) never share a cluster, so their large arrays alias + # one raw buffer and the stream's weight ring takes the bytes the verify buffers need elsewhere. + raw = SmemAllocator().allocate(FUSED_RAW_BYTES, byte_alignment=1024) + stream_carve = _SmemCarver(raw, FUSED_RAW_BYTES) + ring_w = stream_carve.view(cutlass.BFloat16, FUSED_STAGES * BOX_ELEMS, 1024) + ring_x = stream_carve.view(cutlass.BFloat16, FUSED_STAGES * X_ELEMS, 1024) + mbox = stream_carve.view(cutlass.Float32, 3 * CLUSTER * 32 * 4, 16) + verify_carve = _SmemCarver(raw, FUSED_RAW_BYTES) + smem_a = verify_carve.view(cutlass.BFloat16, HD * HD, 1024) + smem_b = verify_carve.view(cutlass.Int32, MMA_N * HD // 2, 1024) + s_uq = verify_carve.view(cutlass.Float32, ROWS_U * HD) + s_uk = verify_carve.view(cutlass.Float32, ROWS_U * HD) + s_uv = verify_carve.view(cutlass.Float32, ROWS_U * V_CTA) + s_wq = verify_carve.view(cutlass.Float32, CONV_W * HD) + s_wk = verify_carve.view(cutlass.Float32, CONV_W * HD) + s_wv = verify_carve.view(cutlass.Float32, CONV_W * V_CTA) + s_dtb = verify_carve.view(cutlass.Float32, HD) + s_onw = verify_carve.view(cutlass.Float32, V_CTA) + s_gr = verify_carve.view(cutlass.Float32, NT * HD) + s_braw = verify_carve.view(cutlass.Float32, NT) + s_og = verify_carve.view(cutlass.Float32, NT * V_CTA) + s_q = verify_carve.view(cutlass.Float32, NT * HD) + s_k = verify_carve.view(cutlass.Float32, NT * HD) + s_dec = verify_carve.view(cutlass.Float32, NT * HD) + s_kd = verify_carve.view(cutlass.Float32, NT * HD) + s_bk = verify_carve.view(cutlass.Float32, NT * HD) + s_beta = verify_carve.view(cutlass.Float32, NT) + s_v = verify_carve.view(cutlass.Float32, NT * V_CTA) + s_o = verify_carve.view(cutlass.Float32, NT * V_CTA) + s_ss = verify_carve.view(cutlass.Float32, 2 * CLUSTER * NT) + s_rs = verify_carve.view(cutlass.Float32, NT) + s_vp = verify_carve.view(cutlass.Float32, 2 * NT * V_CTA) + s_ogp = verify_carve.view(cutlass.Float32, 2 * NT * V_CTA) + s_bp = verify_carve.view(cutlass.Float32, 2 * NT) + s_rec = verify_carve.view(cutlass.Float32, NUM_SPEC * REC_CTA) + # Barriers and the TMEM holder stay static (both roles initialize their own). + full = cutlass.Array(cutlass.Int64, FUSED_STAGES, space=cutlass.AddressSpace.smem, alignment=8) + empty = cutlass.Array(cutlass.Int64, FUSED_STAGES, space=cutlass.AddressSpace.smem, alignment=8) + acc_done = cutlass.Array(cutlass.Int64, 3, space=cutlass.AddressSpace.smem, alignment=8) + mbox_bar = cutlass.Array(cutlass.Int64, 3, space=cutlass.AddressSpace.smem, alignment=8) + tmem_holder = cutlass.Array(cutlass.Int32, 1, space=cutlass.AddressSpace.smem) + bars = cutlass.Array(cutlass.Int64, 3, space=cutlass.AddressSpace.smem, alignment=8) + r_st = cutlass.Array(cutlass.Float32, REC_REGS, space=cutlass.AddressSpace.rmem) + + # Launch order by cluster: the f_a / b stream clusters, the verify clusters, then the q/k/v/og stream clusters. + # The CTAs of the first clusters become resident while the predecessor still holds SMs, so every verify CTA + # runs its pre-wait prologue (state, records, the drafts' replay) before the grid dependency instead of after + # it; the stream clusters that come last load their boxes after it. Roles use the logical index: + # stream CTAs 0-103 (the f_a / b clusters first), verify CTAs 104-127. + pbx, _, _ = cute.arch.block_idx() + lbx_head = pbx + cutlass.Int32((STREAM_CLUSTERS - 2) * CLUSTER) + lbx_qkv = pbx - cutlass.Int32(HEAD_CLUSTERS * CLUSTER) + lbx = cutlass.Int32( + cutlass.select_( + pbx < cutlass.Int32(2 * CLUSTER), + pbx, + cutlass.select_(pbx < cutlass.Int32((2 + HEAD_CLUSTERS) * CLUSTER), lbx_head, lbx_qkv), + ) + ) + if lbx < cutlass.Int32(STREAM_CLUSTERS * CLUSTER): + _stream_role(tma_w, tma_x, p1, part, epoch, ring_w, ring_x, full, empty, acc_done, mbox_bar, mbox, tmem_holder, + USE_PDL, FUSED_STAGES, lbx) # fmt: skip + else: + _head_role(tma_wfb, p1w, part, w_q, w_k, w_v, a_log, dt_bias, onorm_w, cs_q, cs_k, cs_v, ssm, state_tok, + slots, pending, out, epoch, smem_a, smem_b, bars, tmem_holder, s_uq, s_uk, s_uv, s_wq, s_wk, s_wv, + s_dtb, s_onw, s_gr, s_braw, s_og, s_q, s_k, s_dec, s_kd, s_bk, s_beta, s_v, s_o, s_ss, s_rs, s_vp, + s_ogp, s_bp, s_rec, r_st, ssm_stride, pool_n, lower_bound, scale, eps, USE_PDL, lbx) # fmt: skip + + +@cute.jit +def k3_kda_attn( + w: cute.Tensor, # bf16 [3208, 7168] + x: cute.Tensor, # bf16 [8, 7168] + w_fb: cute.Tensor, # bf16 [768, 128] + p1: cute.Tensor, # int16 [3 * 8 * 1664] + p1w: cute.Tensor, # the same memory as int32 + part: cute.Tensor, # int32 [3 * 3 * 2 * 8 * 768] + epoch: cute.Tensor, # int32 [128] + w_q: cute.Tensor, + w_k: cute.Tensor, + w_v: cute.Tensor, + a_log: cute.Tensor, + dt_bias: cute.Tensor, + onorm_w: cute.Tensor, + cs_q: cute.Tensor, + cs_k: cute.Tensor, + cs_v: cute.Tensor, + ssm: cute.Tensor, + state_tok: cute.Tensor, + slots: cute.Tensor, + pending: cute.Tensor, + out: cute.Tensor, + ssm_stride: cutlass.Int64, + lower_bound: cutlass.Constexpr[float], + scale: cutlass.Constexpr[float], + eps: cutlass.Constexpr[float], + USE_PDL: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + """One launch of the fused KDA projection + verify for one request of 8 verify tokens (stage B2).""" + tma_w = cuda.create_tensor_map_tiled( + global_address=w.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[BOX_K, PROJ_ROWS, K_IN // BOX_K, 1, 1], + global_strides=[(K_IN * 2) // 16, (BOX_K * 2) // 16, (PROJ_ROWS * K_IN * 2) // 16, + (PROJ_ROWS * K_IN * 2) // 16], + box_dims=[BOX_K, TILE, BOX_CH, 1, 1], + swizzle=cuda.TensorMapSwizzle.s128b, + ) # fmt: skip + tma_x = cuda.create_tensor_map_tiled( + global_address=x.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[K_IN, NT], + global_strides=[(K_IN * 2) // 16], + box_dims=[BOX_K, MMA_N], + swizzle=cuda.TensorMapSwizzle.s128b, + ) + tma_wfb = cuda.create_tensor_map_tiled( + global_address=w_fb.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[BOX_K, HK, HD // BOX_K, 1, 1], + global_strides=[(HD * 2) // 16, (BOX_K * 2) // 16, (HK * HD * 2) // 16, (HK * HD * 2) // 16], + box_dims=[BOX_K, HD, HD // BOX_K, 1, 1], + swizzle=cuda.TensorMapSwizzle.s128b, + ) # fmt: skip + k3_kda_attn_fused_kernel( + tma_w, tma_x, tma_wfb, p1, p1w, part, epoch, w_q, w_k, w_v, a_log, dt_bias, onorm_w, cs_q, cs_k, cs_v, ssm, + state_tok, slots, pending, out, ssm_stride, cutlass.Int32(cute.size(pending)), lower_bound, scale, eps, USE_PDL, + ).launch( + grid=((STREAM_CLUSTERS + HEAD_CLUSTERS) * CLUSTER, 1, 1), block=(THREADS, 1, 1), cluster=(CLUSTER, 1, 1), + stream=stream, use_pdl=USE_PDL, + # One CTA per SM (the shared-memory carve); without it ptxas may pick an occupancy-driven register target. + min_blocks_per_mp=1, + ) # fmt: skip diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_attn/k3_kda_decode_kernel.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_attn/k3_kda_decode_kernel.py new file mode 100644 index 000000000000..0955213cc57a --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_attn/k3_kda_decode_kernel.py @@ -0,0 +1,802 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Kimi K3 KDA plain decode in one launch: the fused projection of ``k3_kda_attn`` (TP16 rank slice, its 26 stream +clusters and three Lamport phases, unchanged) and the gated-delta decode of R <= 8 requests of one token each, on the +pools of ``trtllm::kda_decode``. + +The projection's 8 token rows are the R requests' tokens (rows >= R read as zero). Six head clusters (local heads +0-5) of 4 CTAs, cluster rank = V quarter; CTA (h, q) owns V rows [32 q, 32 q + 32) of local head h for every request. +Before ``griddepcontrol.wait`` (nothing there is written by this launch or its predecessors): the slots, the head's +W_fb block (TMA), the conv weights and gate constants, each request's conv windows (the head's q / k channels, this +CTA's 32 v channels) and each request's 32 state rows (16 KB) into shared memory. +After it, per request (warp r = request r where a step is per request): + phase 1 q, k and f_a rows; f_b on the tensor cores (M = 128 gate channels, N = 8 rows); + q, k = L2norm(SiLU(conv4(window, new))) (q also scaled); the q / k windows of this CTA's quarter rewritten + once all four CTAs of the head have read the old ones (cluster barrier); + phase 2 v and b (bf16 of the two K-half partials): v = SiLU(conv4) of this CTA's 32 channels, its windows + rewritten; beta = sigmoid(b); decay = exp(lower_bound * sigmoid(exp(A_log) * (g + dt_bias))); + update ``kda_decode``'s arithmetic and row / key layout: lane l keys 4 l .. 4 l + 3, warp w rows w, w + 8, + w + 16, w + 24; S *= decay; r = (v - S k) beta; S += k r; o = S q; S back to the pool; + phase 3 the output gate; the gated RMSNorm over V (each CTA's two 16-row sums of squares per row to the four CTAs + of the head); y = o * rms * w_norm * sigmoid(gate), bf16. +Pools: conv bf16 [slots][3 * 768][3] (q | k | v channels, the last three raw inputs, oldest first; slot stride in +elements, a multiple of 8); state fp32 [slots][6][128][128] (V rows, K contiguous; slot stride a multiple of 4). +""" + +from __future__ import annotations + +import cuda.bindings.driver as cuda_driver +import cutlass +import cutlass.cute as cute +import cutlass.experimental.cuda as cuda +from cutlass._mlir.dialects import llvm as _llvm +from cutlass._mlir.extras import types as _T +from cutlass.cutlass_dsl import dsl_user_op +from cutlass.experimental import primitives as prims + +try: + from cutlass.memory.smem import SmemAllocator +except ImportError: # older DSL layout + from cutlass.utils import SmemAllocator + +from ..k3_kda_verify.k3_kda_verify_kernel import _bf16, _butterfly, _st_async_f32, _store8 +from .k3_kda_attn_kernel import ( + BOX_CH, + BOX_ELEMS, + BOX_K, + BUFFERS, + CLUSTER, + CONV_W, + CP_CG, + FB_TMEM_COLS, + FUSED_RAW_BYTES, + FUSED_STAGES, + HD, + HEAD_CLUSTERS, + HK, + K_IN, + K_STEP_U, + MMA_K, + MMA_N, + P1_ROWS, + PART_ROWS, + PROJ_ROWS, + SENTINEL, + STREAM_CLUSTERS, + THREADS, + TILE, + V_CTA, + X_ELEMS, + _bf16_bits, + _mapa_u32, + _poll_partials, + _sent16, + _sent32, + _SmemCarver, + _st_b16, + _stream_role, +) + +NR = MMA_N # request rows of a launch (the projection's token rows) +WIN = CONV_W - 1 # raw inputs a conv window keeps +WIN_QK = HD * WIN # bf16 of one request's q (or k) window for a head: 128 channels x 3 +WIN_V = V_CTA * WIN # bf16 of one request's v window for a CTA's 32 channels +QK_CHUNKS = WIN_QK // 8 # 16-byte copies per window +V_CHUNKS = WIN_V // 8 +REQ_CHUNKS = 2 * QK_CHUNKS + V_CHUNKS +ST_ROWS = V_CTA * HD # fp32 of one request's state rows owned by a CTA (16 KB) +ST_CHUNKS = ST_ROWS // 4 +LANE_KEYS = HD // 32 # keys per lane in the update (kda_decode's float4) +CW_ITEMS = 2 * V_CTA * WIN # q and k window entries a CTA rewrites per request + + +@dsl_user_op +def _ld_bf16(smem_ptr, *, loc=None, ip=None): + """fp32 of one bf16 in shared memory by ld.shared.b16: a channel's three window entries are only 2-byte aligned, + so their loads must not be merged into one wider load.""" + return cutlass.Float32( + _llvm.inline_asm( + _T.f32(), [smem_ptr.toint(loc=loc, ip=ip).ir_value()], + "{ .reg .b16 h; ld.shared.b16 h, [$1]; cvt.f32.bf16 $0, h; }", "=f,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _test_wait_cluster(mbar, parity, *, loc=None, ip=None): + """mbarrier.test_wait.parity.acquire.cluster on a barrier of this CTA (``mbar``, its shared-memory pointer): whether + phase ``parity`` has completed, acquiring at cluster scope. The barrier is completed by other CTAs' st.async, whose + complete_tx releases at cluster scope.""" + done = cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [mbar.toint(loc=loc, ip=ip).ir_value(), cutlass.Int32(parity).ir_value(loc=loc, ip=ip)], + "{\n\t.reg .pred p;\n\tmbarrier.test_wait.parity.acquire.cluster.shared::cta.b64 p, [$1], $2;\n\t" + "selp.u32 $0, 1, 0, p;\n\t}", "=r,r,r,~{memory}", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + return done != cutlass.Int32(0) + + +@cute.jit +def _decode_head_role( + tma_wfb, + p1w: cutlass.Array, + part: cutlass.Array, + w_q: cutlass.Array, + w_k: cutlass.Array, + w_v: cutlass.Array, + a_log: cutlass.Array, + dt_bias: cutlass.Array, + onorm_w: cutlass.Array, + conv: cutlass.Array, + ssm: cutlass.Array, + slots: cutlass.Array, + out: cutlass.Array, + epoch: cutlass.Array, + smem_a: cutlass.Array, + smem_b: cutlass.Array, + bars: cutlass.Array, + tmem_holder: cutlass.Array, + s_slot: cutlass.Array, + s_winq: cutlass.Array, + s_wink: cutlass.Array, + s_winv: cutlass.Array, + s_wq: cutlass.Array, + s_wk: cutlass.Array, + s_wv: cutlass.Array, + s_dtb: cutlass.Array, + s_onw: cutlass.Array, + s_nq: cutlass.Array, + s_nk: cutlass.Array, + s_gr: cutlass.Array, + s_q: cutlass.Array, + s_k: cutlass.Array, + s_dec: cutlass.Array, + s_beta: cutlass.Array, + s_braw: cutlass.Array, + s_v: cutlass.Array, + s_o: cutlass.Array, + s_ss: cutlass.Array, + s_rs: cutlass.Array, + s_vp: cutlass.Array, + s_bp: cutlass.Array, + s_st: cutlass.Array, + n_req, + ssm_stride, + conv_stride, + lower_bound: cutlass.Constexpr[float], + scale: cutlass.Constexpr[float], + eps: cutlass.Constexpr[float], + USE_PDL: cutlass.Constexpr[bool], + lbx, +): + """One CTA of a head cluster (logical clusters 26-31): local head h = cluster - 26, V rows [32 q, 32 q + 32) + (q = cluster rank) of every request, fed by the stream clusters' Lamport buffers of this launch.""" + tidx, _, _ = cute.arch.thread_idx() + bx = lbx + lane = tidx % 32 + warp = cute.arch.make_warp_uniform(cute.arch.warp_idx()) + vq = cute.arch.block_idx_in_cluster() + h = bx // cutlass.Int32(CLUSTER) - cutlass.Int32(STREAM_CLUSTERS) + v0 = vq * V_CTA + ch0 = h * HD + w_full = bars.subview(0) + acc_done = bars.subview(1) + ss_ready = bars.subview(2) + + tma_ptr_w = tma_wfb.get_ptr() + if warp == 0: + prims.prefetch_tensormap(tma_ptr_w) + if prims.elect_sync(): + prims.mbarrier_init(w_full, 1) + prims.mbarrier_init(acc_done, 1) + elif warp == 2: + prims.tcgen05_alloc(tmem_holder, FB_TMEM_COLS) + prims.tcgen05_relinquish_alloc_permit() + elif warp == 3: + if prims.elect_sync(): + prims.mbarrier_init(ss_ready, 1) + prims.mbarrier_arrive_expect_tx(ss_ready, CLUSTER * 2 * NR * 4) + prims.fence_mbarrier_init() + prims.barrier_cluster_arrive_relaxed() + prims.barrier_cluster_wait() + prims.barrier_cta_sync(0) + if warp == 0: + if prims.elect_sync(): + prims.mbarrier_arrive_expect_tx(w_full, HD * HD * 2) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a, tma_ptr_w, (cutlass.Int32(0), ch0, cutlass.Int32(0), cutlass.Int32(0), cutlass.Int32(0)), + w_full, + ) # fmt: skip + + # ---- Before the grid dependency: the slots, the constants, then (asynchronous copies) every request's conv + # windows and its state rows. + if tidx < n_req: + s_slot.store(slots.load(idx=tidx), idx=tidx) + a_raw = a_log.load(idx=h) + if tidx < HD: + wq4 = w_q.load(idx=(ch0 + tidx) * CONV_W, vector_size=CONV_W, alignment=16) + wk4 = w_k.load(idx=(ch0 + tidx) * CONV_W, vector_size=CONV_W, alignment=16) + for w in cutlass.range_constexpr(CONV_W): + s_wq.store(cutlass.Float32(wq4[w]), idx=w * HD + tidx) + s_wk.store(cutlass.Float32(wk4[w]), idx=w * HD + tidx) + s_dtb.store(dt_bias.load(idx=ch0 + tidx), idx=tidx) + elif tidx < HD + V_CTA: + cv_w = tidx - HD + wv4 = w_v.load(idx=(ch0 + v0 + cv_w) * CONV_W, vector_size=CONV_W, alignment=16) + for w in cutlass.range_constexpr(CONV_W): + s_wv.store(cutlass.Float32(wv4[w]), idx=w * V_CTA + cv_w) + s_onw.store(onorm_w.load(idx=v0 + cv_w), idx=cv_w) + for i_o in cutlass.range_constexpr(NR * V_CTA // THREADS): + s_o.store(cutlass.Float32(0.0), idx=tidx + i_o * THREADS) + prims.barrier_cta_sync(0) + for it_w in cutlass.range_constexpr((NR * REQ_CHUNKS + THREADS - 1) // THREADS): + q_w = tidx + it_w * THREADS + r_w = q_w // REQ_CHUNKS + u_w = q_w % REQ_CHUNKS + if r_w < n_req: + conv_w = conv.subview(s_slot.load(idx=r_w) * conv_stride) + if u_w < QK_CHUNKS: + prims.cp_async_shared_global( + s_winq.subview(r_w * WIN_QK + u_w * 8), + conv_w.subview(ch0 * WIN + u_w * 8), + 16, + CP_CG, + ) + elif u_w < 2 * QK_CHUNKS: + u_k = u_w - QK_CHUNKS + prims.cp_async_shared_global( + s_wink.subview(r_w * WIN_QK + u_k * 8), + conv_w.subview((HK + ch0) * WIN + u_k * 8), + 16, + CP_CG, + ) + else: + u_v = u_w - 2 * QK_CHUNKS + prims.cp_async_shared_global( + s_winv.subview(r_w * WIN_V + u_v * 8), + conv_w.subview((2 * HK + ch0 + v0) * WIN + u_v * 8), + 16, + CP_CG, + ) + prims.cp_async_commit_group() + for r_s in cutlass.range_constexpr(NR): + if r_s < n_req: + st_src = ssm.subview(s_slot.load(idx=r_s) * ssm_stride + (ch0 + v0) * HD) + for j_s in cutlass.range_constexpr(ST_CHUNKS // THREADS): + c_s = (tidx + j_s * THREADS) * 4 + prims.cp_async_shared_global( + s_st.subview(r_s * ST_ROWS + c_s), st_src.subview(c_s), 16, CP_CG + ) + prims.cp_async_commit_group() + exp_a = cute.math.exp(a_raw, fastmath=True) + # The windows (the first group) have landed: every CTA of the head has read the head's q / k windows once all + # four arrive here (release; the acquiring wait follows phase 1), so each may then rewrite its quarter. + prims.cp_async_wait_group(1) + prims.barrier_cta_sync(0) + prims.barrier_cluster_arrive() + + if cutlass.const_expr(USE_PDL): + prims.griddepcontrol(prims.GridDepAction.WAIT) + e = epoch.load(idx=bx) + buf = e % cutlass.Int32(BUFFERS) + if cutlass.const_expr(USE_PDL): + if tidx == 0: + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + + # ---- Phase 1 (q, k, f_a): thread i polls 16-byte chunk i of the q | k rows, threads < 128 also chunk i of the + # f_a rows; q and k as fp32, f_a into the MMA's B operand (128-byte swizzle). + st_qk = tidx // HD + t_qk = (tidx % HD) // 16 + j_qk = tidx % 16 + a0 = cutlass.Int32(SENTINEL) + a1 = cutlass.Int32(SENTINEL) + a2 = cutlass.Int32(SENTINEL) + a3 = cutlass.Int32(SENTINEL) + qk_idx = ((buf * NR + t_qk) * P1_ROWS + st_qk * HK + ch0 + j_qk * 8) // 2 + while _sent16(a0) | _sent16(a1) | _sent16(a2) | _sent16(a3): + vqk = prims.load_ext( + p1w.subview(qk_idx), dtype=cutlass.Int32, count=4, order="relaxed", scope="gpu" + ) + a0 = cutlass.Int32(vqk[0]) + a1 = cutlass.Int32(vqk[1]) + a2 = cutlass.Int32(vqk[2]) + a3 = cutlass.Int32(vqk[3]) + if st_qk == 0: + _store8(s_nq, (a0, a1, a2, a3), t_qk * HD + j_qk * 8) + else: + _store8(s_nk, (a0, a1, a2, a3), t_qk * HD + j_qk * 8) + if tidx < HD: + t_fa = tidx // 16 + j_fa = tidx % 16 + f0 = cutlass.Int32(SENTINEL) + f1 = cutlass.Int32(SENTINEL) + f2 = cutlass.Int32(SENTINEL) + f3 = cutlass.Int32(SENTINEL) + fa_idx = ((buf * NR + t_fa) * P1_ROWS + 2 * HK + j_fa * 8) // 2 + while _sent16(f0) | _sent16(f1) | _sent16(f2) | _sent16(f3): + vfa = prims.load_ext( + p1w.subview(fa_idx), dtype=cutlass.Int32, count=4, order="relaxed", scope="gpu" + ) + f0 = cutlass.Int32(vfa[0]) + f1 = cutlass.Int32(vfa[1]) + f2 = cutlass.Int32(vfa[2]) + f3 = cutlass.Int32(vfa[3]) + # Row t_fa, 16-byte unit u of 64-column chunk c: byte 1024 c + 128 t_fa + 16 (u ^ t_fa). + sw_word = ((j_fa // 8) * 1024 + t_fa * 128 + ((j_fa % 8) ^ t_fa) * 16) // 4 + smem_b.store((f0, f1, f2, f3), idx=sw_word, alignment=16) + prims.fence_proxy("async_shared", space=prims.SharedSpace.shared_cta) + prims.barrier_cta_sync(0) + prims.barrier_cluster_wait() + + # ---- f_b on the tensor cores (warp 2 issues, warps 4-7 read TMEM), rounded to bf16 as the unfused GEMM's output. + if warp == 2: + idesc = prims.Tcgen05InstrDesc.build( + c_dtype=cutlass.Float32, + a_dtype=cutlass.BFloat16, + b_dtype=cutlass.BFloat16, + n_dim=MMA_N, + m_dim=HD, + ) + desc_a_base = prims.Tcgen05SmemDesc.build( + start_address=smem_a, leading_byte_offset=16, stride_byte_offset=1024, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + desc_b_base = prims.Tcgen05SmemDesc.build( + start_address=smem_b, leading_byte_offset=16, stride_byte_offset=1024, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + tmem_acc = cutlass.inttoptr(tmem_holder.load(), 6, cutlass.Int32) + while not cute.arch.mbarrier_test_wait(w_full.data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + for kb in cutlass.range_constexpr(2 * (BOX_K // MMA_K)): + box = kb // (BOX_K // MMA_K) + within = kb % (BOX_K // MMA_K) + if prims.elect_sync(): + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, prims.CTAGroup.CTA_1, tmem_acc, + desc_a_base + (box * ((HD * BOX_K * 2) >> 4) + within * K_STEP_U), + desc_b_base + (box * ((MMA_N * BOX_K * 2) >> 4) + within * K_STEP_U), + idesc, kb != 0, + ) # fmt: skip + if prims.elect_sync(): + prims.tcgen05_commit(acc_done) + if warp >= 4: + while not cute.arch.mbarrier_test_wait(acc_done.data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc = prims.tcgen05_ld( + "32x32b", cutlass.inttoptr(tmem_holder.load(), 6, cutlass.Float32), num=MMA_N + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + c_fb = (warp - 4) * 32 + lane + for t_fb in cutlass.range_constexpr(NR): + s_gr.store(_bf16(cutlass.Float32(acc[t_fb])), idx=t_fb * HD + c_fb) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + prims.barrier_cta_sync(1, thread_count=128) + if warp == 4: + prims.tcgen05_dealloc( + cutlass.inttoptr(tmem_holder.load(), 6, cutlass.Int32), FB_TMEM_COLS + ) + + # ---- q and k of request w (warp w): conv4 over the window and the new input, SiLU, then the L2 norm with + # kda_decode's sum order (per 32-channel group a butterfly, then (g0 + g2) + (g1 + g3)); q also scaled. + if warp < n_req: + tk = warp + pq = [cutlass.Float32(0.0)] * LANE_KEYS + pk = [cutlass.Float32(0.0)] * LANE_KEYS + for i in cutlass.range_constexpr(LANE_KEYS): + c = i * 32 + lane + cq = cutlass.Float32(0.0) + ck = cutlass.Float32(0.0) + for w in cutlass.range_constexpr(WIN): + cq += _ld_bf16(s_winq.subview(tk * WIN_QK + c * WIN + w).data_ptr()) * s_wq.load( + idx=w * HD + c + ) + ck += _ld_bf16(s_wink.subview(tk * WIN_QK + c * WIN + w).data_ptr()) * s_wk.load( + idx=w * HD + c + ) + cq += s_nq.load(idx=tk * HD + c) * s_wq.load(idx=WIN * HD + c) + ck += s_nk.load(idx=tk * HD + c) * s_wk.load(idx=WIN * HD + c) + pq[i] = cq * ( + cutlass.Float32(1.0) / (cutlass.Float32(1.0) + cute.math.exp(-cq, fastmath=True)) + ) + pk[i] = ck * ( + cutlass.Float32(1.0) / (cutlass.Float32(1.0) + cute.math.exp(-ck, fastmath=True)) + ) + gq = [_butterfly(pq[i] * pq[i]) for i in range(LANE_KEYS)] + gk = [_butterfly(pk[i] * pk[i]) for i in range(LANE_KEYS)] + rq = ( + cute.math.rsqrt( + (gq[0] + gq[2]) + (gq[1] + gq[3]) + cutlass.Float32(1e-6), fastmath=True + ) + * scale + ) + rk = cute.math.rsqrt( + (gk[0] + gk[2]) + (gk[1] + gk[3]) + cutlass.Float32(1e-6), fastmath=True + ) + for i in cutlass.range_constexpr(LANE_KEYS): + s_q.store(pq[i] * rq, idx=tk * HD + i * 32 + lane) + s_k.store(pk[i] * rk, idx=tk * HD + i * 32 + lane) + # The q / k windows of this CTA's quarter: the two newer raw inputs, then the token's (all four CTAs of the head + # read the old ones before the cluster wait above). + for it_c in cutlass.range_constexpr((NR * CW_ITEMS + THREADS - 1) // THREADS): + q_c = tidx + it_c * THREADS + r_c = q_c // CW_ITEMS + u_c = q_c % CW_ITEMS + if r_c < n_req: + sec_c = u_c // (V_CTA * WIN) + ch_c = vq * V_CTA + (u_c % (V_CTA * WIN)) // WIN + w_c = u_c % WIN + dst_c = ( + conv.subview( + s_slot.load(idx=r_c) * conv_stride + (sec_c * HK + ch0 + ch_c) * WIN + w_c + ) + .data_ptr() + .toint() + ) + # Single 16-bit stores: a channel's entries are only 2-byte aligned. + if w_c < WIN - 1: + if sec_c == 0: + _st_b16( + dst_c, + _bf16_bits( + _ld_bf16(s_winq.subview(r_c * WIN_QK + ch_c * WIN + w_c + 1).data_ptr()) + ), + ) + else: + _st_b16( + dst_c, + _bf16_bits( + _ld_bf16(s_wink.subview(r_c * WIN_QK + ch_c * WIN + w_c + 1).data_ptr()) + ), + ) + else: + if sec_c == 0: + _st_b16(dst_c, _bf16_bits(s_nq.load(idx=r_c * HD + ch_c))) + else: + _st_b16(dst_c, _bf16_bits(s_nk.load(idx=r_c * HD + ch_c))) + prims.barrier_cta_sync(0) + + # ---- Phase 2 (v, b): the two K-half partials of this CTA's 32 v rows and of this head's b; the projection is + # their bf16-rounded sum. + if tidx < 2 * NR * (V_CTA // 4): + half_v = tidx // (NR * (V_CTA // 4)) + t_v = (tidx // (V_CTA // 4)) % NR + j_v = tidx % (V_CTA // 4) + _poll_partials( + part, + ((buf * 3 + 0) * 2 + half_v) * (8 * PART_ROWS) + t_v * PART_ROWS + ch0 + v0 + j_v * 4, + s_vp, + (half_v * NR + t_v) * V_CTA + j_v * 4, + ) + elif tidx < 2 * NR * (V_CTA // 4) + 2 * NR: + ib = tidx - 2 * NR * (V_CTA // 4) + half_b = ib // NR + t_b = ib % NR + wb = cutlass.Int32(SENTINEL) + while _sent32(wb): + wb = prims.load_ext( + part.subview(((buf * 3 + 2) * 2 + half_b) * (8 * PART_ROWS) + t_b * PART_ROWS + h), + dtype=cutlass.Int32, + order="relaxed", + scope="gpu", + ) + s_bp.store(wb.bitcast(cutlass.Float32), idx=half_b * NR + t_b) + prims.barrier_cta_sync(0) + t_cv = tidx // V_CTA + c_cv = tidx % V_CTA + nv = _bf16(s_vp.load(idx=t_cv * V_CTA + c_cv) + s_vp.load(idx=(NR + t_cv) * V_CTA + c_cv)) + if tidx < NR: + s_braw.store(_bf16(s_bp.load(idx=tidx) + s_bp.load(idx=NR + tidx)), idx=tidx) + if t_cv < n_req: + # v (one (request, channel) per thread), then its window: the two newer raw inputs and the token's. + cv = cutlass.Float32(0.0) + wv_old = [ + _ld_bf16(s_winv.subview(t_cv * WIN_V + c_cv * WIN + w).data_ptr()) for w in range(WIN) + ] + for w in cutlass.range_constexpr(WIN): + cv += wv_old[w] * s_wv.load(idx=w * V_CTA + c_cv) + cv += nv * s_wv.load(idx=WIN * V_CTA + c_cv) + s_v.store( + cv + * (cutlass.Float32(1.0) / (cutlass.Float32(1.0) + cute.math.exp(-cv, fastmath=True))), + idx=t_cv * V_CTA + c_cv, + ) + dst_v = ( + conv.subview(s_slot.load(idx=t_cv) * conv_stride + (2 * HK + ch0 + v0 + c_cv) * WIN) + .data_ptr() + .toint() + ) + _st_b16(dst_v, _bf16_bits(wv_old[1])) + _st_b16(dst_v + cutlass.Int64(2), _bf16_bits(wv_old[2])) + _st_b16(dst_v + cutlass.Int64(4), _bf16_bits(nv)) + prims.barrier_cta_sync(0) + if tidx < n_req: + s_beta.store( + cutlass.Float32(1.0) + / (cutlass.Float32(1.0) + cute.math.exp(-s_braw.load(idx=tidx), fastmath=True)), + idx=tidx, + ) + # The decay of request w's keys (warp w; lane l keys 4 l .. 4 l + 3, the update's layout). + if warp < n_req: + for i in cutlass.range_constexpr(LANE_KEYS): + c = lane * LANE_KEYS + i + xg = exp_a * (s_gr.load(idx=warp * HD + c) + s_dtb.load(idx=c)) + sig = cutlass.Float32(1.0) / (cutlass.Float32(1.0) + cute.math.exp(-xg, fastmath=True)) + s_dec.store(cute.math.exp(lower_bound * sig, fastmath=True), idx=warp * HD + c) + prims.cp_async_wait_group(0) + prims.barrier_cta_sync(0) + + # ---- The update, request by request: warp w rows w, w + 8 and w + 16, w + 24 of this CTA's 32, lane l keys + # 4 l .. 4 l + 3 (kda_decode's layout and sum order); the rows back to the pool, the outputs to s_o. + for t in cutlass.range(n_req, unroll=1): + q4 = s_q.load(idx=t * HD + lane * LANE_KEYS, vector_size=4, alignment=16) + k4 = s_k.load(idx=t * HD + lane * LANE_KEYS, vector_size=4, alignment=16) + d4 = s_dec.load(idx=t * HD + lane * LANE_KEYS, vector_size=4, alignment=16) + beta_t = s_beta.load(idx=t) + st_dst = ssm.subview(s_slot.load(idx=t) * ssm_stride + (ch0 + v0) * HD) + for pr in cutlass.range_constexpr(2): + ra = warp + 16 * pr + rb = ra + 8 + sa4 = s_st.load( + idx=t * ST_ROWS + ra * HD + lane * LANE_KEYS, vector_size=4, alignment=16 + ) + sb4 = s_st.load( + idx=t * ST_ROWS + rb * HD + lane * LANE_KEYS, vector_size=4, alignment=16 + ) + sa = [cutlass.Float32(sa4[i]) * cutlass.Float32(d4[i]) for i in range(LANE_KEYS)] + sb = [cutlass.Float32(sb4[i]) * cutlass.Float32(d4[i]) for i in range(LANE_KEYS)] + ska = sa[0] * cutlass.Float32(k4[0]) + skb = sb[0] * cutlass.Float32(k4[0]) + for i in cutlass.range_constexpr(1, LANE_KEYS): + ska = ska + sa[i] * cutlass.Float32(k4[i]) + skb = skb + sb[i] * cutlass.Float32(k4[i]) + for offset in [16, 8, 4, 2, 1]: + ska = ska + cute.arch.shuffle_sync_bfly( + ska, offset=offset, mask=-1, mask_and_clamp=31 + ) + skb = skb + cute.arch.shuffle_sync_bfly( + skb, offset=offset, mask=-1, mask_and_clamp=31 + ) + res_a = (s_v.load(idx=t * V_CTA + ra) - ska) * beta_t + res_b = (s_v.load(idx=t * V_CTA + rb) - skb) * beta_t + for i in cutlass.range_constexpr(LANE_KEYS): + sa[i] = sa[i] + cutlass.Float32(k4[i]) * res_a + sb[i] = sb[i] + cutlass.Float32(k4[i]) * res_b + st_dst.store((sa[0], sa[1], sa[2], sa[3]), idx=ra * HD + lane * LANE_KEYS, alignment=16) + st_dst.store((sb[0], sb[1], sb[2], sb[3]), idx=rb * HD + lane * LANE_KEYS, alignment=16) + sqa = sa[0] * cutlass.Float32(q4[0]) + sqb = sb[0] * cutlass.Float32(q4[0]) + for i in cutlass.range_constexpr(1, LANE_KEYS): + sqa = sqa + sa[i] * cutlass.Float32(q4[i]) + sqb = sqb + sb[i] * cutlass.Float32(q4[i]) + for offset in [16, 8, 4, 2, 1]: + sqa = sqa + cute.arch.shuffle_sync_bfly( + sqa, offset=offset, mask=-1, mask_and_clamp=31 + ) + sqb = sqb + cute.arch.shuffle_sync_bfly( + sqb, offset=offset, mask=-1, mask_and_clamp=31 + ) + if lane == 0: + s_o.store(sqa, idx=t * V_CTA + ra) + s_o.store(sqb, idx=t * V_CTA + rb) + + # ---- Phase 3 (the output gate), then the gated RMSNorm over V: every CTA sends its two 16-row sums of squares + # per row to all four CTAs of the head. Thread (row t, V row v) reads its gate's two K-half partials first, so + # their round trip overlaps the norm exchange. + t_out = tidx // V_CTA + v_y = tidx % V_CTA + og_at = (buf * 3 + 1) * 2 * (8 * PART_ROWS) + t_out * PART_ROWS + ch0 + v0 + v_y + og_h0 = cutlass.Int32( + prims.load_ext(part.subview(og_at), dtype=cutlass.Int32, order="relaxed", scope="gpu") + ) + og_h1 = cutlass.Int32(prims.load_ext(part.subview(og_at + 8 * PART_ROWS), dtype=cutlass.Int32, order="relaxed", + scope="gpu")) # fmt: skip + prims.barrier_cta_sync(0) + x_o = s_o.load(idx=tidx) + ss = x_o * x_o + for off_ss in [8, 4, 2, 1]: + ss = ss + cute.arch.shuffle_sync_bfly(ss, offset=off_ss, mask=-1, mask_and_clamp=31) + if lane % 16 == 0: + src_slot = vq * 2 + lane // 16 + for r in cutlass.range_constexpr(CLUSTER): + _st_async_f32( + _mapa_u32(s_ss.subview(src_slot * NR + t_out).data_ptr(), r), + ss, + _mapa_u32(ss_ready.data_ptr(), r), + ) + # Completed by the head's st.async: acquire at cluster scope. + while not _test_wait_cluster(ss_ready.data_ptr(), 0): + pass + while _sent32(og_h0) | _sent32(og_h1): + og_h0 = cutlass.Int32( + prims.load_ext(part.subview(og_at), dtype=cutlass.Int32, order="relaxed", scope="gpu") + ) + og_h1 = cutlass.Int32(prims.load_ext(part.subview(og_at + 8 * PART_ROWS), dtype=cutlass.Int32, + order="relaxed", scope="gpu")) # fmt: skip + if tidx < NR: + total = cutlass.Float32(0.0) + for r in cutlass.range_constexpr(2 * CLUSTER): + total = total + s_ss.load(idx=r * NR + tidx) + s_rs.store(cute.math.rsqrt(total / HD + eps), idx=tidx) + prims.barrier_cta_sync(0) + if t_out < n_req: + z = _bf16(og_h0.bitcast(cutlass.Float32) + og_h1.bitcast(cutlass.Float32)) + gate = cutlass.Float32(1.0) / (cutlass.Float32(1.0) + cute.math.exp(-z, fastmath=True)) + y = s_o.load(idx=tidx) * s_rs.load(idx=t_out) * s_onw.load(idx=v_y) * gate + out.store(cutlass.BFloat16(y), idx=(t_out * (HK // HD) + h) * HD + v0 + v_y) + + if tidx == 0: + epoch.store((e + cutlass.Int32(1)) % cutlass.Int32(BUFFERS), idx=bx) + + +@cute.kernel +def k3_kda_decode_fused_kernel( + tma_w: cutlass.GridConstant[ + cuda.TensorMap + ], # W [3208, 7168] bf16, 5-D, box 64 cols x 64 rows x 2 chunks + tma_x: cutlass.GridConstant[ + cuda.TensorMap + ], # x [R, 7168] bf16, box 64 cols x 8 rows (rows >= R read as 0) + tma_wfb: cutlass.GridConstant[ + cuda.TensorMap + ], # W_fb [768, 128] bf16, a head's 128 rows by one call + p1: cutlass.Array, # int16 bits of bf16 [3][8][1664] (Lamport, written by the stream role) + p1w: cutlass.Array, # the same memory as int32 words (polled by the head role) + part: cutlass.Array, # int32 bits of fp32 [3][3][2][8][768] (Lamport) + epoch: cutlass.Array, # int32 [128]: each CTA's buffer index (launches completed mod 3) + w_q: cutlass.Array, # fp32 [768, 4] + w_k: cutlass.Array, + w_v: cutlass.Array, + a_log: cutlass.Array, # fp32 [6] + dt_bias: cutlass.Array, # fp32 [768] + onorm_w: cutlass.Array, # fp32 [128] + conv: cutlass.Array, # bf16 [slots][2304][3], slot stride conv_stride elements + ssm: cutlass.Array, # fp32 [slots][6][128][128], slot stride ssm_stride elements + slots: cutlass.Array, # int32 [R] + out: cutlass.Array, # bf16 [R][6][128] + n_req: cutlass.Int32, + ssm_stride: cutlass.Int64, + conv_stride: cutlass.Int64, + lower_bound: cutlass.Constexpr[float], + scale: cutlass.Constexpr[float], + eps: cutlass.Constexpr[float], + USE_PDL: cutlass.Constexpr[bool], +): + # The stream CTAs (clusters 0-25) and the head CTAs (26-31) never share a cluster, so their large arrays alias one + # raw buffer (k3_kda_attn's carve). + raw = SmemAllocator().allocate(FUSED_RAW_BYTES, byte_alignment=1024) + stream_carve = _SmemCarver(raw, FUSED_RAW_BYTES) + ring_w = stream_carve.view(cutlass.BFloat16, FUSED_STAGES * BOX_ELEMS, 1024) + ring_x = stream_carve.view(cutlass.BFloat16, FUSED_STAGES * X_ELEMS, 1024) + mbox = stream_carve.view(cutlass.Float32, 3 * CLUSTER * 32 * 4, 16) + head_carve = _SmemCarver(raw, FUSED_RAW_BYTES) + smem_a = head_carve.view(cutlass.BFloat16, HD * HD, 1024) + smem_b = head_carve.view(cutlass.Int32, MMA_N * HD // 2, 1024) + s_st = head_carve.view(cutlass.Float32, NR * ST_ROWS) + s_winq = head_carve.view(cutlass.BFloat16, NR * WIN_QK) + s_wink = head_carve.view(cutlass.BFloat16, NR * WIN_QK) + s_winv = head_carve.view(cutlass.BFloat16, NR * WIN_V) + s_slot = head_carve.view(cutlass.Int32, NR) + s_wq = head_carve.view(cutlass.Float32, CONV_W * HD) + s_wk = head_carve.view(cutlass.Float32, CONV_W * HD) + s_wv = head_carve.view(cutlass.Float32, CONV_W * V_CTA) + s_dtb = head_carve.view(cutlass.Float32, HD) + s_onw = head_carve.view(cutlass.Float32, V_CTA) + s_nq = head_carve.view(cutlass.Float32, NR * HD) + s_nk = head_carve.view(cutlass.Float32, NR * HD) + s_gr = head_carve.view(cutlass.Float32, NR * HD) + s_q = head_carve.view(cutlass.Float32, NR * HD) + s_k = head_carve.view(cutlass.Float32, NR * HD) + s_dec = head_carve.view(cutlass.Float32, NR * HD) + s_beta = head_carve.view(cutlass.Float32, NR) + s_braw = head_carve.view(cutlass.Float32, NR) + s_v = head_carve.view(cutlass.Float32, NR * V_CTA) + s_o = head_carve.view(cutlass.Float32, NR * V_CTA) + s_ss = head_carve.view(cutlass.Float32, 2 * CLUSTER * NR) + s_rs = head_carve.view(cutlass.Float32, NR) + s_vp = head_carve.view(cutlass.Float32, 2 * NR * V_CTA) + s_bp = head_carve.view(cutlass.Float32, 2 * NR) + # Barriers and the TMEM holder stay static (both roles initialize their own). + full = cutlass.Array(cutlass.Int64, FUSED_STAGES, space=cutlass.AddressSpace.smem, alignment=8) + empty = cutlass.Array(cutlass.Int64, FUSED_STAGES, space=cutlass.AddressSpace.smem, alignment=8) + acc_done = cutlass.Array(cutlass.Int64, 3, space=cutlass.AddressSpace.smem, alignment=8) + mbox_bar = cutlass.Array(cutlass.Int64, 3, space=cutlass.AddressSpace.smem, alignment=8) + tmem_holder = cutlass.Array(cutlass.Int32, 1, space=cutlass.AddressSpace.smem) + bars = cutlass.Array(cutlass.Int64, 3, space=cutlass.AddressSpace.smem, alignment=8) + + # Launch order = logical order: the f_a / b stream clusters, the q/k/v/og stream clusters, then the head clusters. + # Beside a predecessor that still holds SMs (the residual-update sandwich before this layer), the CTAs that launch + # first are the stream CTAs, which have the most bytes to move; the head CTAs' prologue is short. Stream CTAs 0-103, + # head CTAs 104-127. + lbx = cutlass.Int32(cute.arch.block_idx()[0]) + if lbx < cutlass.Int32(STREAM_CLUSTERS * CLUSTER): + _stream_role(tma_w, tma_x, p1, part, epoch, ring_w, ring_x, full, empty, acc_done, mbox_bar, mbox, tmem_holder, + USE_PDL, FUSED_STAGES, lbx) # fmt: skip + else: + _decode_head_role(tma_wfb, p1w, part, w_q, w_k, w_v, a_log, dt_bias, onorm_w, conv, ssm, slots, out, epoch, + smem_a, smem_b, bars, tmem_holder, s_slot, s_winq, s_wink, s_winv, s_wq, s_wk, s_wv, s_dtb, + s_onw, s_nq, s_nk, s_gr, s_q, s_k, s_dec, s_beta, s_braw, s_v, s_o, s_ss, s_rs, s_vp, + s_bp, s_st, n_req, ssm_stride, conv_stride, lower_bound, scale, eps, USE_PDL, + lbx) # fmt: skip + + +@cute.jit +def k3_kda_decode( + w: cute.Tensor, # bf16 [3208, 7168] + x: cute.Tensor, # bf16 [R, 7168] + w_fb: cute.Tensor, # bf16 [768, 128] + p1: cute.Tensor, # int16 [3 * 8 * 1664] + p1w: cute.Tensor, # the same memory as int32 + part: cute.Tensor, # int32 [3 * 3 * 2 * 8 * 768] + epoch: cute.Tensor, # int32 [128] + w_q: cute.Tensor, + w_k: cute.Tensor, + w_v: cute.Tensor, + a_log: cute.Tensor, + dt_bias: cute.Tensor, + onorm_w: cute.Tensor, + conv: cute.Tensor, + ssm: cute.Tensor, + slots: cute.Tensor, + out: cute.Tensor, + n_req: cutlass.Int32, + ssm_stride: cutlass.Int64, + conv_stride: cutlass.Int64, + lower_bound: cutlass.Constexpr[float], + scale: cutlass.Constexpr[float], + eps: cutlass.Constexpr[float], + USE_PDL: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + """One launch of the fused KDA projection + plain decode of ``n_req`` requests of one token.""" + tma_w = cuda.create_tensor_map_tiled( + global_address=w.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[BOX_K, PROJ_ROWS, K_IN // BOX_K, 1, 1], + global_strides=[(K_IN * 2) // 16, (BOX_K * 2) // 16, (PROJ_ROWS * K_IN * 2) // 16, + (PROJ_ROWS * K_IN * 2) // 16], + box_dims=[BOX_K, TILE, BOX_CH, 1, 1], + swizzle=cuda.TensorMapSwizzle.s128b, + ) # fmt: skip + tma_x = cuda.create_tensor_map_tiled( + global_address=x.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[K_IN, n_req], + global_strides=[(K_IN * 2) // 16], + box_dims=[BOX_K, MMA_N], + swizzle=cuda.TensorMapSwizzle.s128b, + ) + tma_wfb = cuda.create_tensor_map_tiled( + global_address=w_fb.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[BOX_K, HK, HD // BOX_K, 1, 1], + global_strides=[(HD * 2) // 16, (BOX_K * 2) // 16, (HK * HD * 2) // 16, (HK * HD * 2) // 16], + box_dims=[BOX_K, HD, HD // BOX_K, 1, 1], + swizzle=cuda.TensorMapSwizzle.s128b, + ) # fmt: skip + k3_kda_decode_fused_kernel( + tma_w, tma_x, tma_wfb, p1, p1w, part, epoch, w_q, w_k, w_v, a_log, dt_bias, onorm_w, conv, ssm, slots, out, + n_req, ssm_stride, conv_stride, lower_bound, scale, eps, USE_PDL, + ).launch( + grid=((STREAM_CLUSTERS + HEAD_CLUSTERS) * CLUSTER, 1, 1), block=(THREADS, 1, 1), cluster=(CLUSTER, 1, 1), + stream=stream, use_pdl=USE_PDL, + # One CTA per SM (the shared-memory carve); without it ptxas may pick an occupancy-driven register target. + min_blocks_per_mp=1, + ) # fmt: skip diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_attn/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_attn/op.py new file mode 100644 index 000000000000..a1f703c0df03 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_attn/op.py @@ -0,0 +1,356 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``trtllm::k3_kda_qkvg``: Kimi K3's fused KDA projection (TP16 slice, W [3208, 7168]) as a three-phase stream. + +Outputs (Lamport buffers, see the kernel module): ``p1`` int16 [3, 8, 1664] holds the bf16 bits of the q, k and f_a +columns; ``part`` int32 [3, 3, 2, 8, 768] the fp32 bits of the two K-half partial sums of v (region 0), og (1) and b +(2, first 8 rows), whose bf16-rounded sum is the projection. A launch writes buffer e = ``epoch[cta]`` and resets buffer +(e + 1) % 3 to the sentinel; ``epoch`` int32 [104] holds each CTA's launch count mod 3. The buffers must persist across +launches and start as all-ones (``p1``, ``part``) and zeros (``epoch``). The kernel compiles on the first call for each +configuration, which must happen outside CUDA-graph capture. +""" + +from __future__ import annotations + +import os +import threading +from typing import Dict + +import torch + +K_IN = 7168 +PROJ_ROWS = 3208 +P1_NUMEL = 3 * 8 * 1664 +PART_NUMEL = 3 * 3 * 2 * 8 * 768 +CTAS = 104 +FUSED_CTAS = 128 +NT = 8 # one request's verify tokens (golden + 7 drafts) + +_lock = threading.Lock() +_compiled: Dict[tuple, object] = {} + + +def _arg(t: torch.Tensor, align: int = 16): + from cutlass.cute.runtime import from_dlpack + + return from_dlpack(t.detach(), assumed_align=align).mark_layout_dynamic(leading_dim=t.dim() - 1) + + +def _index_arg(t: torch.Tensor): + """An int32 index tensor the kernels read element by element (``slots``, ``pending``), declared at its element's + alignment: the mixer passes slices such as ``state_indices[num_prefills:]``, which start on any 4-byte boundary.""" + return _arg(t, t.element_size()) + + +def use_pdl() -> bool: + return os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" + + +def make_buffers(device: torch.device, ctas: int = CTAS): + """Fresh (p1, part, epoch) for :func:`k3_kda_qkvg` (``ctas`` = 104) or :func:`k3_kda_attn` (``FUSED_CTAS``).""" + p1 = torch.full((P1_NUMEL,), -1, dtype=torch.int16, device=device) + part = torch.full((PART_NUMEL,), -1, dtype=torch.int32, device=device) + epoch = torch.zeros(ctas, dtype=torch.int32, device=device) + return p1, part, epoch + + +@torch.library.custom_op( + "trtllm::k3_kda_qkvg", mutates_args=("p1", "part", "epoch"), device_types="cuda" +) +def k3_kda_qkvg( + x: torch.Tensor, + w: torch.Tensor, + p1: torch.Tensor, + part: torch.Tensor, + epoch: torch.Tensor, +) -> None: + """x bf16 [T <= 8, 7168], w bf16 [3208, 7168] (both contiguous); writes p1 and part, advances epoch (see the + module doc).""" + import cuda.bindings.driver as cuda_driver + + from . import k3_kda_attn_kernel as kernel + + tokens = x.shape[0] + if ( + x.dtype != torch.bfloat16 + or w.dtype != torch.bfloat16 + or tuple(w.shape) != (PROJ_ROWS, K_IN) + or x.dim() != 2 + or x.shape[1] != K_IN + or not 1 <= tokens <= 8 + or not x.is_contiguous() + or not w.is_contiguous() + or p1.numel() != P1_NUMEL + or p1.dtype != torch.int16 + or part.numel() != PART_NUMEL + or part.dtype != torch.int32 + or epoch.numel() != CTAS + or epoch.dtype != torch.int32 + ): + raise ValueError( + f"k3_kda_qkvg: unsupported call x {tuple(x.shape)} {x.dtype}, w {tuple(w.shape)} {w.dtype}, " + f"p1 {p1.numel()} {p1.dtype}, part {part.numel()} {part.dtype}, epoch {epoch.numel()} {epoch.dtype}" + ) + args = (_arg(w), _arg(x), _arg(p1.view(-1)), _arg(part.view(-1)), _arg(epoch.view(-1))) + stream = cuda_driver.CUstream(torch.cuda.current_stream(x.device).cuda_stream) + pdl = use_pdl() + key = (pdl,) + fn = _compiled.get(key) + if fn is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "trtllm::k3_kda_qkvg must run once per configuration outside CUDA-graph capture first " + "(it compiles its kernel on the first call)." + ) + import cutlass.cute as cute + + with _lock: + fn = _compiled.get(key) + if fn is None: + fn = _compiled[key] = cute.compile(kernel.k3_kda_qkvg, *args, tokens, pdl, stream) + fn(*args, tokens, stream) + + +@k3_kda_qkvg.register_fake +def _(x, w, p1, part, epoch): + return None + + +@torch.library.custom_op( + "trtllm::k3_kda_attn", + mutates_args=("cs_q", "cs_k", "cs_v", "ssm", "state_tok", "p1", "part", "epoch"), + device_types="cuda", +) +def k3_kda_attn( + x: torch.Tensor, + w: torch.Tensor, + w_fb: torch.Tensor, + w_q: torch.Tensor, + w_k: torch.Tensor, + w_v: torch.Tensor, + a_log: torch.Tensor, + dt_bias: torch.Tensor, + onorm_w: torch.Tensor, + cs_q: torch.Tensor, + cs_k: torch.Tensor, + cs_v: torch.Tensor, + ssm: torch.Tensor, + state_tok: torch.Tensor, + slots: torch.Tensor, + pending: torch.Tensor, + p1: torch.Tensor, + part: torch.Tensor, + epoch: torch.Tensor, + num_spec: int, + lower_bound: float, + scale: float, + eps: float, +) -> torch.Tensor: + """Stage B2: the fused projection of ``x`` (bf16 [8, 7168], one request's golden token and 7 drafts) and the KDA + verify of ``trtllm::k3_kda_verify`` on it, in one launch; returns the gated-norm core output bf16 [8, 6, 128]. + Weights, pools and state contract as ``k3_kda_verify`` (``slots`` int32 [1]); ``p1``, ``part`` and ``epoch`` + from :func:`make_buffers` with ``FUSED_CTAS``, persistent across launches. The head CTAs read the slot's pools + before the grid-dependency wait, so a launch must not follow another launch on the same pools directly in the + stream: a kernel that waits (or a non-PDL one) in between, as the model's other layers are.""" + import cuda.bindings.driver as cuda_driver + + from ..k3_kda_verify.op import _flat, _slots_view + from . import k3_kda_attn_kernel as kernel + + if ( + tuple(x.shape) != (NT, K_IN) + or x.dtype != torch.bfloat16 + or not x.is_contiguous() + or tuple(w.shape) != (PROJ_ROWS, K_IN) + or w.dtype != torch.bfloat16 + or not w.is_contiguous() + or tuple(w_fb.shape) != (768, 128) + or w_fb.dtype != torch.bfloat16 + or not w_fb.is_contiguous() + or num_spec != NT - 1 + or tuple(ssm.shape[1:]) != (6, 128, 128) + or tuple(state_tok.shape[1:]) != (num_spec, 6, 128, 128) + or ssm.stride()[1:] != (128 * 128, 128, 1) + or not state_tok.is_contiguous() + or cs_q.shape[-1] != 3 + num_spec + or slots.numel() != 1 + or slots.dtype != torch.int32 + or pending.dtype != torch.int32 + or p1.numel() != P1_NUMEL + or part.numel() != PART_NUMEL + or epoch.numel() != FUSED_CTAS + or ssm.dtype != torch.float32 + or state_tok.dtype != torch.float32 + or any( + t.dtype != torch.float32 + for t in (w_q, w_k, w_v, a_log, dt_bias, onorm_w, cs_q, cs_k, cs_v) + ) + ): + raise ValueError( + f"k3_kda_attn: unsupported call x {tuple(x.shape)} {x.dtype}, w {tuple(w.shape)}, " + f"w_fb {tuple(w_fb.shape)}, ssm {tuple(ssm.shape)}, state_tok {tuple(state_tok.shape)}, " + f"cs_q {tuple(cs_q.shape)}, " + f"num_spec {num_spec}, slots {tuple(slots.shape)} {slots.dtype}, epoch {epoch.numel()}" + ) + out = torch.empty(NT, 6, 128, dtype=torch.bfloat16, device=x.device) + args = ( + _arg(w), _arg(x), _arg(w_fb), _arg(p1.view(-1)), _arg(p1.view(torch.int32)), _arg(part.view(-1)), + _arg(epoch.view(-1)), _arg(_flat(w_q)), _arg(_flat(w_k)), _arg(_flat(w_v)), _arg(_flat(a_log)), + _arg(_flat(dt_bias)), _arg(_flat(onorm_w)), _arg(_flat(cs_q)), _arg(_flat(cs_k)), _arg(_flat(cs_v)), + _arg(_slots_view(ssm)), _arg(_slots_view(state_tok)), _index_arg(slots), _index_arg(pending), + _arg(_flat(out)), + ) # fmt: skip + stream = cuda_driver.CUstream(torch.cuda.current_stream(x.device).cuda_stream) + pdl = use_pdl() + key = ("attn", float(lower_bound), float(scale), float(eps), pdl) + fn = _compiled.get(key) + if fn is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "trtllm::k3_kda_attn must run once per configuration outside CUDA-graph capture first " + "(it compiles its kernel on the first call)." + ) + import cutlass.cute as cute + + with _lock: + fn = _compiled.get(key) + if fn is None: + fn = _compiled[key] = cute.compile( + kernel.k3_kda_attn, *args, ssm.stride(0), float(lower_bound), float(scale), float(eps), pdl, + stream, + ) # fmt: skip + fn(*args, ssm.stride(0), stream) + return out + + +@k3_kda_attn.register_fake +def _(x, w, w_fb, w_q, w_k, w_v, a_log, dt_bias, onorm_w, cs_q, cs_k, cs_v, ssm, state_tok, slots, pending, p1, part, + epoch, num_spec, lower_bound, scale, eps): # fmt: skip + return x.new_empty((NT, 6, 128), dtype=torch.bfloat16) + + +@torch.library.custom_op( + "trtllm::k3_kda_decode_attn", + mutates_args=("conv", "ssm", "p1", "part", "epoch"), + device_types="cuda", +) +def k3_kda_decode_attn( + x: torch.Tensor, + w: torch.Tensor, + w_fb: torch.Tensor, + w_q: torch.Tensor, + w_k: torch.Tensor, + w_v: torch.Tensor, + a_log: torch.Tensor, + dt_bias: torch.Tensor, + onorm_w: torch.Tensor, + conv: torch.Tensor, + ssm: torch.Tensor, + slots: torch.Tensor, + p1: torch.Tensor, + part: torch.Tensor, + epoch: torch.Tensor, + lower_bound: float, + scale: float, + eps: float, +) -> torch.Tensor: + """The fused projection of ``x`` (bf16 [R <= 8, 7168], one token of each of R requests) and the KDA plain decode + of ``trtllm::kda_decode`` on it, in one launch; returns the gated-norm core output bf16 [R, 6, 128]. + + ``conv`` bf16 [slots, 2304, 3] (q | k | v channels, the last three raw inputs; each slot dense, the slot stride a + multiple of 8 elements) and ``ssm`` fp32 [slots, 6, 128, 128] (each slot dense, the slot stride a multiple of 4) + are updated in place at ``slots`` int32 [R]; ``w_q / w_k / w_v`` fp32 [768, 4]; ``a_log`` fp32 [6]; + ``dt_bias`` fp32 [768]; ``onorm_w`` fp32 [128]; ``p1``, ``part`` and ``epoch`` from :func:`make_buffers` with + ``FUSED_CTAS``, persistent across launches. The head CTAs read the slots' pools before the grid-dependency wait, + so a launch must not follow another launch on the same pools directly in the stream: a kernel that waits (or a + non-PDL one) in between, as the model's other layers are.""" + import cuda.bindings.driver as cuda_driver + + from ..k3_kda_verify.op import _flat, _slots_view + from . import k3_kda_decode_kernel as kernel + + n_req = x.shape[0] if x.dim() == 2 else 0 + if ( + x.dim() != 2 + or not 1 <= n_req <= NT + or x.shape[1] != K_IN + or x.dtype != torch.bfloat16 + or not x.is_contiguous() + or tuple(w.shape) != (PROJ_ROWS, K_IN) + or w.dtype != torch.bfloat16 + or not w.is_contiguous() + or tuple(w_fb.shape) != (768, 128) + or w_fb.dtype != torch.bfloat16 + or not w_fb.is_contiguous() + or conv.dtype != torch.bfloat16 + or conv.dim() != 3 + or tuple(conv.shape[1:]) != (3 * 768, 3) + or conv.stride()[1:] != (3, 1) + or conv.stride(0) % 8 != 0 + or conv.data_ptr() % 16 != 0 + or ssm.dtype != torch.float32 + or tuple(ssm.shape[1:]) != (6, 128, 128) + or ssm.stride()[1:] != (128 * 128, 128, 1) + or ssm.stride(0) % 4 != 0 + or ssm.data_ptr() % 16 != 0 + or slots.numel() != n_req + or slots.dtype != torch.int32 + or not slots.is_contiguous() + or p1.numel() != P1_NUMEL + or part.numel() != PART_NUMEL + or epoch.numel() != FUSED_CTAS + or any(t.dtype != torch.float32 for t in (w_q, w_k, w_v, a_log, dt_bias, onorm_w)) + ): + raise ValueError( + f"k3_kda_decode_attn: unsupported call x {tuple(x.shape)} {x.dtype}, w {tuple(w.shape)}, " + f"w_fb {tuple(w_fb.shape)}, conv {tuple(conv.shape)} {conv.dtype} strides {conv.stride()}, " + f"ssm {tuple(ssm.shape)} {ssm.dtype} strides {ssm.stride()}, slots {tuple(slots.shape)} {slots.dtype}, " + f"epoch {epoch.numel()}" + ) + out = torch.empty(n_req, 6, 128, dtype=torch.bfloat16, device=x.device) + args = ( + _arg(w), _arg(x), _arg(w_fb), _arg(p1.view(-1)), _arg(p1.view(torch.int32)), _arg(part.view(-1)), + _arg(epoch.view(-1)), _arg(_flat(w_q)), _arg(_flat(w_k)), _arg(_flat(w_v)), _arg(_flat(a_log)), + _arg(_flat(dt_bias)), _arg(_flat(onorm_w)), _arg(_slots_view(conv)), _arg(_slots_view(ssm)), + _index_arg(slots), _arg(_flat(out)), + ) # fmt: skip + runtime = (n_req, ssm.stride(0), conv.stride(0)) + stream = cuda_driver.CUstream(torch.cuda.current_stream(x.device).cuda_stream) + pdl = use_pdl() + key = ("decode", float(lower_bound), float(scale), float(eps), pdl) + fn = _compiled.get(key) + if fn is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "trtllm::k3_kda_decode_attn must run once per configuration outside CUDA-graph capture first " + "(it compiles its kernel on the first call)." + ) + import cutlass.cute as cute + + with _lock: + fn = _compiled.get(key) + if fn is None: + fn = _compiled[key] = cute.compile( + kernel.k3_kda_decode, *args, *runtime, float(lower_bound), float(scale), float(eps), pdl, + stream, + ) # fmt: skip + fn(*args, *runtime, stream) + return out + + +@k3_kda_decode_attn.register_fake +def _(x, w, w_fb, w_q, w_k, w_v, a_log, dt_bias, onorm_w, conv, ssm, slots, p1, part, epoch, lower_bound, scale, + eps): # fmt: skip + return x.new_empty((x.shape[0], 6, 128), dtype=torch.bfloat16) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_verify/__init__.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_verify/__init__.py new file mode 100644 index 000000000000..2892fd223c30 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_verify/__init__.py @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Kimi K3 KDA speculative verify as one CTM kernel (``trtllm::k3_kda_verify``). + +Importing :mod:`.op` registers the torch op; nothing is imported eagerly here so that +the CuTe DSL dependency stays optional for every other model. +""" diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_verify/k3_kda_verify_kernel.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_verify/k3_kda_verify_kernel.py new file mode 100644 index 000000000000..ad95b425e4ee --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_verify/k3_kda_verify_kernel.py @@ -0,0 +1,827 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Kimi K3 KDA speculative verify as one kernel: the forget-gate up-projection (f_b), the per-token pre-compute and +the delta-rule recurrence over the 1 + NUM_SPEC verify tokens, with a state committed after every verify token so +the next round starts from the accepted one instead of replaying the accepted drafts. + +Grid (H, N, 8), cluster (1, 1, 8), 256 threads. CTA (h, n, s) owns V rows [16 s, 16 s + 16) of local head h for +request n; warp w owns rows 16 s + 2 w and 16 s + 2 w + 1, and lane l keys k = 32 i + l (i < 4) of each. + +Before ``griddepcontrol.wait`` (written by the previous verify step, or constant): + the slot and P, the drafts the sampler accepted last round; the starting state into registers (P == 0: the pool + state, committed after the last golden token; else ``state_tok[slot, P - 1]``); the raw conv inputs at positions + -3..-1 before the first new token (conv-cache columns P..P+2) for q and k (128 channels) and v (this CTA's 16); + the conv weights, dt_bias, A_log and the output-norm weight; the head's W_fb^T block by TMA. +After the wait (this step's fused projection rows [q | k | v | onorm gate | f_a | b]): the new tokens' raw q, k, v, + f_a, b and output gate. +f_b: g[t, c] = bf16(sum_i f_a[t, i] W_fb[128 h + c, i]), rounded as the unfused bf16 GEMV output is. +Pre-compute: warp t = verify token t: q, k = l2norm(silu(conv4(u))) (q also scaled), the lower-bound gate, beta and + this CTA's 16 v channels. +Recurrence: the arithmetic of ``kda_mtp_decode``'s V-split path over the verify tokens, unrolled; the state after + token 0 (the golden token) is committed to the pool, the state after token t >= 1 to + ``state_tok[slot, t - 1]``. +Epilogue: the gated RMSNorm of the outputs over V (per-token sums of squares stored into every peer's shared + memory, one cluster barrier), bf16 output rows; CTA 0 rewrites the q/k conv cache (the window at the + golden token and the raw inputs of the drafts), every CTA the same for its v channels. + +With FOLD_FB False the gate comes from ``g_ext`` (the unfused f_b output); the kernel is then bit-exact against +``kda_mtp_decode`` fed the same gate and replaying the same accepted drafts. +""" + +from __future__ import annotations + +import cuda.bindings.driver as cuda_driver +import cutlass +import cutlass.cute as cute +import cutlass.experimental.cuda as cuda +from cutlass._mlir.dialects import llvm as _llvm +from cutlass._mlir.extras import types as _T +from cutlass.cutlass_dsl import dsl_user_op +from cutlass.experimental import primitives as prims + +K = 128 # key head dim; the value head dim V is the same +V = 128 +CONV_W = 4 +THREADS = 256 +WARPS = THREADS // 32 +V_SPLIT = 8 +V_CTA = V // V_SPLIT # V rows per CTA +ROWS = V_CTA // WARPS # V rows per warp +VEC = K // 32 # keys per lane +BF16_BYTES = 2 +CHUNK = 8 # bf16 per 16-byte load +SMEM = cutlass.AddressSpace.smem +RMEM = cutlass.AddressSpace.rmem + +assert ROWS == 2, "the recurrence below processes one pair of V rows per warp" + +# f_b on the tensor cores: D[c, t] = sum_i W_fb[128 h + c, i] f_a[t, i], M = 128 channels, N = 8 tokens, K = 128. +CTA_M = 128 +MMA_N = 8 +MMA_K = 16 +TMA_K_BOX = 64 +TMA_COPY_ITERS = K // TMA_K_BOX +K_BLOCKS_PER_HALF = TMA_K_BOX // MMA_K +LEADING = 16 +STRIDE = 8 * TMA_K_BOX * BF16_BYTES +A_HALF_ELEMS = CTA_M * TMA_K_BOX +B_HALF_ELEMS = MMA_N * TMA_K_BOX +STEP = (MMA_K * BF16_BYTES) >> 4 +A_BOX = A_HALF_ELEMS >> 3 +B_BOX = B_HALF_ELEMS >> 3 +TMEM_COLS = 32 + + +@dsl_user_op +def _mapa_u32(smem_ptr, peer, *, loc=None, ip=None): + """The shared::cluster address of this CTA's shared-memory location in cluster CTA ``peer``.""" + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [smem_ptr.toint(loc=loc, ip=ip).ir_value(), cutlass.Int32(peer).ir_value(loc=loc, ip=ip)], + "mapa.shared::cluster.u32 $0, $1, $2;", "=r,r,r", has_side_effects=False, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _test_wait_cluster(mbar, parity, *, loc=None, ip=None): + """mbarrier.test_wait.parity.acquire.cluster on a barrier of this CTA (``mbar``, its shared-memory pointer): whether + phase ``parity`` has completed, acquiring at cluster scope. For barriers that other CTAs' st.async complete (their + complete_tx releases at cluster scope; a CTA-scope acquire does not synchronize with it).""" + done = cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [mbar.toint(loc=loc, ip=ip).ir_value(), cutlass.Int32(parity).ir_value(loc=loc, ip=ip)], + "{\n\t.reg .pred p;\n\tmbarrier.test_wait.parity.acquire.cluster.shared::cta.b64 p, [$1], $2;\n\t" + "selp.u32 $0, 1, 0, p;\n\t}", "=r,r,r,~{memory}", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + return done != cutlass.Int32(0) + + +@dsl_user_op +def _try_wait_cluster(mbar, parity, *, loc=None, ip=None): + """mbarrier.try_wait.parity.acquire.cluster on a barrier of this CTA (``mbar``, its shared-memory pointer): as + ``_test_wait_cluster``, with try_wait's bounded suspend.""" + done = cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [mbar.toint(loc=loc, ip=ip).ir_value(), cutlass.Int32(parity).ir_value(loc=loc, ip=ip)], + "{\n\t.reg .pred p;\n\tmbarrier.try_wait.parity.acquire.cluster.shared::cta.b64 p, [$1], $2;\n\t" + "selp.u32 $0, 1, 0, p;\n\t}", "=r,r,r,~{memory}", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + return done != cutlass.Int32(0) + + +@dsl_user_op +def _st_async_f32(dst, value, mbar, *, loc=None, ip=None): + """st.async of one fp32 to a shared::cluster address, completing ``mbar`` (a shared::cluster address) by 4 bytes.""" + _llvm.inline_asm( + None, + [cutlass.Int32(dst).ir_value(loc=loc, ip=ip), cutlass.Float32(value).ir_value(loc=loc, ip=ip), + cutlass.Int32(mbar).ir_value(loc=loc, ip=ip)], + "st.async.shared::cluster.mbarrier::complete_tx::bytes.f32 [$0], $1, [$2];", "r,f,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +def _lo(word): + return (word << cutlass.Int32(16)).bitcast(cutlass.Float32) + + +def _hi(word): + return (word & cutlass.Int32(-65536)).bitcast(cutlass.Float32) + + +def _bf16(x): + return cutlass.Float32(cutlass.BFloat16(x)) + + +def _butterfly(value): + for offset in [16, 8, 4, 2, 1]: + value += cute.arch.shuffle_sync_bfly(value, offset=offset, mask=-1, mask_and_clamp=31) + return value + + +def _store8(dst, words, idx): + """The 8 bf16 of four int32 words, as fp32, into dst[idx : idx + 8].""" + dst.store( + (_lo(cutlass.Int32(words[0])), _hi(cutlass.Int32(words[0])), _lo(cutlass.Int32(words[1])), + _hi(cutlass.Int32(words[1]))), + idx=idx, alignment=16, + ) # fmt: skip + dst.store( + (_lo(cutlass.Int32(words[2])), _hi(cutlass.Int32(words[2])), _lo(cutlass.Int32(words[3])), + _hi(cutlass.Int32(words[3]))), + idx=idx + 4, alignment=16, + ) # fmt: skip + + +def _qk_pre(tk, lane, s_uq, s_wq, s_uk, s_wk, s_q, s_k, scale): + """q and k of verify token ``tk`` from the conv history in shared memory: conv4 + SiLU, then the L2 norm (q also + scaled), lane l keys 32 i + l (the arithmetic of kda_mtp_decode's pre-compute warps). Traced inline: its Python + loops unroll.""" + pq = [cutlass.Float32(0.0)] * VEC + for i in range(VEC): + c = i * 32 + lane + conv = cutlass.Float32(0.0) + for w in range(CONV_W - 1): + conv += s_uq.load(idx=(tk + w) * K + c) * s_wq.load(idx=w * K + c) + conv += s_uq.load(idx=(tk + CONV_W - 1) * K + c) * s_wq.load(idx=(CONV_W - 1) * K + c) + e = cute.math.exp(-conv, fastmath=True) + pq[i] = conv * cute.arch.rcp_approx(cutlass.Float32(1.0) + e) + sum_q = cutlass.Float32(0.0) + for i in range(VEC): + sum_q += pq[i] * pq[i] + sum_q = _butterfly(sum_q) + rnorm_q = cute.math.rsqrt(sum_q + 1e-06, fastmath=True) * scale + for i in range(VEC): + s_q.store(pq[i] * rnorm_q, idx=tk * K + i * 32 + lane) + pk = [cutlass.Float32(0.0)] * VEC + for i in range(VEC): + c = i * 32 + lane + conv = s_uk.load(idx=tk * K + c) * s_wk.load(idx=c) + for w in range(1, CONV_W - 1): + conv += s_uk.load(idx=(tk + w) * K + c) * s_wk.load(idx=w * K + c) + conv += s_uk.load(idx=(tk + CONV_W - 1) * K + c) * s_wk.load(idx=(CONV_W - 1) * K + c) + pk[i] = conv * cute.arch.rcp_approx( + cutlass.Float32(1.0) + cute.math.exp(-conv, fastmath=True) + ) + sum_k = cutlass.Float32(0.0) + for i in range(VEC): + sum_k += pk[i] * pk[i] + sum_k = _butterfly(sum_k) + rnorm_k = cute.math.rsqrt(sum_k + 1e-06, fastmath=True) + for i in range(VEC): + s_k.store(pk[i] * rnorm_k, idx=tk * K + i * 32 + lane) + + +@cute.kernel +def k3_kda_verify_kernel( + tma_wfb: cutlass.GridConstant[ + cuda.TensorMap + ], # W_fb [H * K, K] bf16, the head's 128 rows by one call + tma_fa: cutlass.GridConstant[ + cuda.TensorMap + ], # the projection rows' f_a columns, box 8 tokens x 64 + proj: cutlass.Array, # int32 words of the fused projection rows, bf16 [T, 2 * proj_words] + g_ext: cutlass.Array, # int32 words of bf16 [T, H * K], the unfused f_b output (FOLD_FB False only) + w_q: cutlass.Array, # fp32 [H * K, CONV_W] + w_k: cutlass.Array, + w_v: cutlass.Array, + a_log: cutlass.Array, # fp32 [H] + dt_bias: cutlass.Array, # fp32 [H * K] + onorm_w: cutlass.Array, # fp32 [V] + cs_q: cutlass.Array, # fp32 [pool][S][H * K]: raw conv inputs, column s = position s - 2 from the golden token + cs_k: cutlass.Array, + cs_v: cutlass.Array, + ssm: cutlass.Array, # fp32 [pool][H][V][K] at slot stride ssm_stride: the state after the last golden token + state_tok: cutlass.Array, # fp32 [pool][NUM_SPEC][H][V][K]: the state after draft t + 1 of the last round + slots: cutlass.Array, # int32 [N] + pending: cutlass.Array, # int32 [pool]: drafts the sampler accepted last round + out: cutlass.Array, # bf16 [T][H][V] + proj_words: cutlass.Int32, + ssm_stride: cutlass.Int64, # fp32 elements between the pool's slots (the manager interleaves conv states) + H: cutlass.Constexpr[int], + NUM_SPEC: cutlass.Constexpr[int], + lower_bound: cutlass.Constexpr[float], + scale: cutlass.Constexpr[float], + eps: cutlass.Constexpr[float], + FOLD_FB: cutlass.Constexpr[bool], + USE_PDL: cutlass.Constexpr[bool], +): + NT = NUM_SPEC + 1 # verify tokens per request + S = CONV_W - 1 + NUM_SPEC # conv-cache columns + HK = H * K + ROWS_U = ( + CONV_W - 1 + NT + ) # raw conv inputs by position: rows 0..2 before token 0, then the tokens + # The fused projection row [q | k | v | onorm gate | f_a | b | pad], in int32 words (bf16 pairs). + Q_W = 0 + K_W = HK // 2 + V_W = HK + OG_W = 3 * HK // 2 + FA_W = 2 * HK + B_COL = 4 * HK + K # bf16 column + assert NT <= WARPS and NT % 2 == 0 + + tidx, _, _ = cute.arch.thread_idx() + lane = tidx % 32 + warp = cute.arch.make_warp_uniform(cute.arch.warp_idx()) + h, n, vs = cute.arch.block_idx() + v0 = vs * V_CTA + ch0 = h * K + row0 = n * NT + + smem_a = cutlass.Array( + cutlass.BFloat16, CTA_M * K, space=SMEM, alignment=1024 + ) # W_fb rows, 128B-swizzled + smem_b = cutlass.Array( + cutlass.BFloat16, MMA_N * K, space=SMEM, alignment=1024 + ) # f_a, 128B-swizzled + w_full = cutlass.Array(cutlass.Int64, 1, space=SMEM, alignment=8) + x_full = cutlass.Array(cutlass.Int64, 1, space=SMEM, alignment=8) + acc_done = cutlass.Array(cutlass.Int64, 1, space=SMEM, alignment=8) + ss_ready = cutlass.Array( + cutlass.Int64, 1, space=SMEM, alignment=8 + ) # the peers' norm partials have landed + tmem_ptr_i32 = cutlass.Array(cutlass.Int32, 1, space=SMEM) + s_uq = cutlass.Array(cutlass.Float32, ROWS_U * K, space=SMEM, alignment=16) + s_uk = cutlass.Array(cutlass.Float32, ROWS_U * K, space=SMEM, alignment=16) + s_uv = cutlass.Array(cutlass.Float32, ROWS_U * V_CTA, space=SMEM, alignment=16) + s_wq = cutlass.Array(cutlass.Float32, CONV_W * K, space=SMEM, alignment=16) # [tap][channel] + s_wk = cutlass.Array(cutlass.Float32, CONV_W * K, space=SMEM, alignment=16) + s_wv = cutlass.Array(cutlass.Float32, CONV_W * V_CTA, space=SMEM, alignment=16) + s_dtb = cutlass.Array(cutlass.Float32, K, space=SMEM, alignment=16) + s_onw = cutlass.Array(cutlass.Float32, V_CTA, space=SMEM, alignment=16) + s_gr = cutlass.Array( + cutlass.Float32, NT * K, space=SMEM, alignment=16 + ) # f_b output, bf16-rounded + s_braw = cutlass.Array(cutlass.Float32, NT, space=SMEM, alignment=16) + s_og = cutlass.Array(cutlass.Float32, NT * V_CTA, space=SMEM, alignment=16) + s_q = cutlass.Array(cutlass.Float32, NT * K, space=SMEM, alignment=16) + s_k = cutlass.Array(cutlass.Float32, NT * K, space=SMEM, alignment=16) + s_dec = cutlass.Array(cutlass.Float32, NT * K, space=SMEM, alignment=16) # exp(gate) + s_kd = cutlass.Array(cutlass.Float32, NT * K, space=SMEM, alignment=16) # exp(gate) * k + s_bk = cutlass.Array(cutlass.Float32, NT * K, space=SMEM, alignment=16) # beta * k + s_beta = cutlass.Array(cutlass.Float32, NT, space=SMEM, alignment=16) + s_v = cutlass.Array(cutlass.Float32, NT * V_CTA, space=SMEM, alignment=16) + s_o = cutlass.Array( + cutlass.Float32, NT * V_CTA, space=SMEM, alignment=16 + ) # bf16-rounded raw outputs + s_ss = cutlass.Array( + cutlass.Float32, V_SPLIT * NT, space=SMEM, alignment=16 + ) # [source CTA][token] + s_rs = cutlass.Array(cutlass.Float32, NT, space=SMEM, alignment=16) + r_st = cutlass.Array(cutlass.Float32, ROWS * VEC, space=RMEM) + + # ---- Prologue: nothing here depends on the predecessor grid. + if cutlass.const_expr(USE_PDL): + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + tma_ptr_w = tma_wfb.get_ptr() + tma_ptr_x = tma_fa.get_ptr() + if cutlass.const_expr(FOLD_FB): + if warp == 0: + prims.prefetch_tensormap(tma_ptr_w) + prims.prefetch_tensormap(tma_ptr_x) + if prims.elect_sync(): + prims.mbarrier_init(w_full, 1) + prims.mbarrier_init(x_full, 1) + prims.mbarrier_init(acc_done, 1) + if warp == 2: + prims.tcgen05_alloc(tmem_ptr_i32, TMEM_COLS) + prims.tcgen05_relinquish_alloc_permit() + if warp == 3: + if prims.elect_sync(): + prims.mbarrier_init(ss_ready, 1) + prims.mbarrier_arrive_expect_tx(ss_ready, V_SPLIT * NT * 4) + prims.fence_mbarrier_init() + # Cluster formation: the peers' shared memory is addressable from here on. + prims.barrier_cluster_arrive_relaxed() + prims.barrier_cluster_wait() + prims.barrier_cta_sync(0) + if cutlass.const_expr(FOLD_FB): + if warp == 0: + if prims.elect_sync(): + prims.mbarrier_arrive_expect_tx(w_full, CTA_M * K * BF16_BYTES) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a, + tma_ptr_w, + (cutlass.Int32(0), ch0, cutlass.Int32(0), cutlass.Int32(0), cutlass.Int32(0)), + w_full, + ) + + slot = slots.load(idx=n) + pend = pending.load(idx=slot) + pend = cutlass.Int32( + cutlass.select_(pend > cutlass.Int32(NUM_SPEC), cutlass.Int32(NUM_SPEC), pend) + ) + exp_a = cute.math.exp(a_log.load(idx=h), fastmath=True) + row_a = v0 + warp * ROWS + # The warp owns both rows and lane l keys 32 i + l (the arithmetic of kda_mtp_decode). + # The drafts' records, in each slot's per-token state region in place of their full states: the row innovations + # vn [NUM_SPEC][H][V], then beta * k and the decay [NUM_SPEC][H][K]. The state after the accepted drafts is the + # pool's (the golden token's) with each accepted draft's update replayed, S = fma(decay, S, vn * (beta * k)): the + # recurrence's own arithmetic, so it is bit-identical to the one the drafts reached. k3_kda_attn uses the same + # records. + CT_VN = 0 + CT_WB = NUM_SPEC * H * V + CT_WD = CT_WB + NUM_SPEC * H * K + # The slot's pool state and per-token region from their first element: the slot offset in 64 bits, once. + pool = ssm.subview(cutlass.Int64(slot) * ssm_stride) + tok = state_tok.subview(cutlass.Int64(slot) * (NUM_SPEC * H * V * K)) + st_base = (h * V + row_a) * K + for r in cutlass.range_constexpr(ROWS): + for i in cutlass.range_constexpr(VEC): + r_st.store(pool.load(idx=st_base + r * K + i * 32 + lane), idx=r * VEC + i) + for t_acc in cutlass.range(pend, unroll=1): + rec_acc = (t_acc * H + h) * V + vns_acc = [tok.load(idx=rec_acc + CT_VN + row_a + r) for r in range(ROWS)] + wbs_acc = [tok.load(idx=CT_WB + (t_acc * H + h) * K + i * 32 + lane) for i in range(VEC)] + wds_acc = [tok.load(idx=CT_WD + (t_acc * H + h) * K + i * 32 + lane) for i in range(VEC)] + sts_acc = [r_st.load(idx=j) for j in range(ROWS * VEC)] + for r in cutlass.range_constexpr(ROWS): + for _pi in cutlass.range_constexpr(VEC // 2): + _p = _pi * 2 + vb0_acc, vb1_acc = cute.arch.mul_packed_f32x2( + (vns_acc[r], vns_acc[r]), (wbs_acc[_p], wbs_acc[_p + 1]) + ) + sts_acc[r * VEC + _p], sts_acc[r * VEC + _p + 1] = cute.arch.fma_packed_f32x2( + src_a=(wds_acc[_p], wds_acc[_p + 1]), src_b=(sts_acc[r * VEC + _p], sts_acc[r * VEC + _p + 1]), + src_c=(vb0_acc, vb1_acc), + ) # fmt: skip + for j in cutlass.range_constexpr(ROWS * VEC): + r_st.store(sts_acc[j], idx=j) + + # Conv weights, transposed to [tap][channel]: q by threads 0..127, k by 128..255. + if tidx < K: + wq4 = w_q.load(idx=(ch0 + tidx) * CONV_W, vector_size=CONV_W, alignment=16) + for w in cutlass.range_constexpr(CONV_W): + s_wq.store(cutlass.Float32(wq4[w]), idx=w * K + tidx) + else: + ck_w = tidx - K + wk4 = w_k.load(idx=(ch0 + ck_w) * CONV_W, vector_size=CONV_W, alignment=16) + for w in cutlass.range_constexpr(CONV_W): + s_wk.store(cutlass.Float32(wk4[w]), idx=w * K + ck_w) + # Conv history (positions -3..-1 = columns P..P+2), the v weights, the output-norm weight, dt_bias. + QH = 3 * K // 4 # float4 loads per q (or k) history + VH = 3 * V_CTA // 4 + if tidx < QH: + m_q = tidx // (K // 4) + c4_q = (tidx % (K // 4)) * 4 + s_uq.store(cs_q.load(idx=(slot * S + pend + m_q) * HK + ch0 + c4_q, vector_size=4, alignment=16), + idx=m_q * K + c4_q, alignment=16) # fmt: skip + elif tidx < 2 * QH: + m_k = (tidx - QH) // (K // 4) + c4_k = ((tidx - QH) % (K // 4)) * 4 + s_uk.store(cs_k.load(idx=(slot * S + pend + m_k) * HK + ch0 + c4_k, vector_size=4, alignment=16), + idx=m_k * K + c4_k, alignment=16) # fmt: skip + elif tidx < 2 * QH + VH: + m_v = (tidx - 2 * QH) // (V_CTA // 4) + c4_v = ((tidx - 2 * QH) % (V_CTA // 4)) * 4 + s_uv.store(cs_v.load(idx=(slot * S + pend + m_v) * HK + ch0 + v0 + c4_v, vector_size=4, alignment=16), + idx=m_v * V_CTA + c4_v, alignment=16) # fmt: skip + elif tidx < 2 * QH + VH + V_CTA: + cv_w = tidx - (2 * QH + VH) + wv4 = w_v.load(idx=(ch0 + v0 + cv_w) * CONV_W, vector_size=CONV_W, alignment=16) + for w in cutlass.range_constexpr(CONV_W): + s_wv.store(cutlass.Float32(wv4[w]), idx=w * V_CTA + cv_w) + elif tidx < 2 * QH + VH + 2 * V_CTA: + cv_n = tidx - (2 * QH + VH + V_CTA) + s_onw.store(onorm_w.load(idx=v0 + cv_n), idx=cv_n) + if tidx < K // 4: + s_dtb.store( + dt_bias.load(idx=ch0 + tidx * 4, vector_size=4, alignment=16), + idx=tidx * 4, + alignment=16, + ) + + # ---- This step's projection rows. + if cutlass.const_expr(USE_PDL): + prims.griddepcontrol(prims.GridDepAction.WAIT) + if cutlass.const_expr(FOLD_FB): + if warp == 1: + if prims.elect_sync(): + prims.mbarrier_arrive_expect_tx(x_full, MMA_N * K * BF16_BYTES) + for half in cutlass.range_constexpr(TMA_COPY_ITERS): + prims.cp_async_bulk_tensor_shared_cta_global( + smem_b.subview(half * B_HALF_ELEMS), tma_ptr_x, + (cutlass.Int32(FA_W * 2 + half * TMA_K_BOX), row0), x_full, + ) # fmt: skip + CH = K // CHUNK # 16-byte chunks per 128-channel row + for it in cutlass.range_constexpr((2 * NT * CH + THREADS - 1) // THREADS): + item = tidx + it * THREADS + if item < 2 * NT * CH: + stream_qk = item // (NT * CH) + t_qk = (item // CH) % NT + j_qk = item % CH + words_qk = proj.load( + idx=(row0 + t_qk) * proj_words + Q_W + ch0 // 2 + stream_qk * (K_W - Q_W) + j_qk * (CHUNK // 2), + vector_size=4, alignment=16, + ) # fmt: skip + if stream_qk == 0: + _store8(s_uq, words_qk, (CONV_W - 1 + t_qk) * K + j_qk * CHUNK) + else: + _store8(s_uk, words_qk, (CONV_W - 1 + t_qk) * K + j_qk * CHUNK) + VCH = V_CTA // CHUNK # 16-byte chunks per 16-channel v (or gate) slice + n_v = NT * VCH + n_g = 0 if FOLD_FB else NT * CH + for it in cutlass.range_constexpr((2 * n_v + NT + n_g + THREADS - 1) // THREADS): + item = tidx + it * THREADS + if item < n_v: + t_v = item // VCH + j_v = item % VCH + words_v = proj.load(idx=(row0 + t_v) * proj_words + V_W + (ch0 + v0) // 2 + j_v * (CHUNK // 2), + vector_size=4, alignment=16) # fmt: skip + _store8(s_uv, words_v, (CONV_W - 1 + t_v) * V_CTA + j_v * CHUNK) + elif item < 2 * n_v: + t_og = (item - n_v) // VCH + j_og = (item - n_v) % VCH + words_og = proj.load(idx=(row0 + t_og) * proj_words + OG_W + (ch0 + v0) // 2 + j_og * (CHUNK // 2), + vector_size=4, alignment=16) # fmt: skip + _store8(s_og, words_og, t_og * V_CTA + j_og * CHUNK) + elif item < 2 * n_v + NT: + t_b = item - 2 * n_v + word_b = proj.load(idx=(row0 + t_b) * proj_words + (B_COL + h) // 2) + b_in = _lo(word_b) + if (B_COL + h) % 2 == 1: + b_in = _hi(word_b) + s_braw.store(b_in, idx=t_b) + elif item < 2 * n_v + NT + n_g: + gi = item - (2 * n_v + NT) + t_g = gi // CH + j_g = gi % CH + words_g = g_ext.load(idx=(row0 + t_g) * (HK // 2) + ch0 // 2 + j_g * (CHUNK // 2), vector_size=4, + alignment=16) # fmt: skip + _store8(s_gr, words_g, t_g * K + j_g * CHUNK) + prims.barrier_cta_sync(0) + + # ---- f_b on the tensor cores (warp 2 issues, warps 4-7 read TMEM) beside the q/k pre-compute (warps 0-3). + if cutlass.const_expr(FOLD_FB): + if warp == 2: + idesc = prims.Tcgen05InstrDesc.build( + c_dtype=cutlass.Float32, + a_dtype=cutlass.BFloat16, + b_dtype=cutlass.BFloat16, + n_dim=MMA_N, + m_dim=CTA_M, + ) + desc_a_base = prims.Tcgen05SmemDesc.build( + start_address=smem_a, leading_byte_offset=LEADING, stride_byte_offset=STRIDE, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + desc_b_base = prims.Tcgen05SmemDesc.build( + start_address=smem_b, leading_byte_offset=LEADING, stride_byte_offset=STRIDE, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + tmem_acc = cutlass.inttoptr(tmem_ptr_i32.load(), 6, cutlass.Int32) + while not cute.arch.mbarrier_try_wait(w_full.data_ptr(), 0): + pass + while not cute.arch.mbarrier_try_wait(x_full.data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + for kb in cutlass.range_constexpr(TMA_COPY_ITERS * K_BLOCKS_PER_HALF): + box = kb // K_BLOCKS_PER_HALF + within = kb % K_BLOCKS_PER_HALF + if prims.elect_sync(): + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, prims.CTAGroup.CTA_1, tmem_acc, + desc_a_base + (box * A_BOX + within * STEP), desc_b_base + (box * B_BOX + within * STEP), + idesc, kb != 0, + ) # fmt: skip + if prims.elect_sync(): + prims.tcgen05_commit(acc_done) + if warp < 4: + # q and k of tokens 2w and 2w + 1 (the arithmetic of kda_mtp_decode's pre-compute warps); NT not a multiple of 4 + # rounds the tokens per warp up and the warps past NT idle. + qk_per_warp = (NT + 3) // 4 + for sub in cutlass.range_constexpr(qk_per_warp): + tk = warp * qk_per_warp + sub + if cutlass.const_expr(NT % 4 == 0): + _qk_pre(tk, lane, s_uq, s_wq, s_uk, s_wk, s_q, s_k, scale) + else: + if tk < NT: + _qk_pre(tk, lane, s_uq, s_wq, s_uk, s_wk, s_q, s_k, scale) + else: + if cutlass.const_expr(FOLD_FB): + while not cute.arch.mbarrier_try_wait(acc_done.data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc = prims.tcgen05_ld( + "32x32b", cutlass.inttoptr(tmem_ptr_i32.load(), 6, cutlass.Float32), num=MMA_N + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + c_fb = (warp - 4) * 32 + lane + for t_fb in cutlass.range_constexpr(NT): + s_gr.store(_bf16(cutlass.Float32(acc[t_fb])), idx=t_fb * K + c_fb) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + prims.barrier_cta_sync(1, thread_count=128) + if warp == 4: + prims.tcgen05_dealloc( + cutlass.inttoptr(tmem_ptr_i32.load(), 6, cutlass.Int32), TMEM_COLS + ) + # v (this CTA's 16 channels) and beta, one (token, channel) per thread of warps 4-7. + item_v = tidx - 128 + if item_v < NT * V_CTA: + t_cv = item_v // V_CTA + c_cv = item_v % V_CTA + conv_v = cutlass.Float32(0.0) + for w in cutlass.range_constexpr(CONV_W): + conv_v += s_uv.load(idx=(t_cv + w) * V_CTA + c_cv) * s_wv.load(idx=w * V_CTA + c_cv) + conv_v = conv_v * cute.arch.rcp_approx( + cutlass.Float32(1.0) + cute.math.exp(-conv_v, fastmath=True) + ) + s_v.store(conv_v, idx=t_cv * V_CTA + c_cv) + if item_v < NT: + b_pre = s_braw.load(idx=item_v) + s_beta.store( + cute.arch.rcp_approx(cutlass.Float32(1.0) + cute.math.exp(-b_pre, fastmath=True)), + idx=item_v, + ) + prims.barrier_cta_sync(0) + + # ---- The recurrence's per-token operands: warp t, token t: decay = exp(gate), decay * k, beta * k. + if warp < NT: + tg = warp + r_beta_g = s_beta.load(idx=tg) + for i_pair in cutlass.range_constexpr(VEC // 2): + dks = [] + for i in (i_pair * 2, i_pair * 2 + 1): + c = i * 32 + lane + g_raw = s_gr.load(idx=tg * K + c) + s_dtb.load(idx=c) + x = exp_a * g_raw + sig = cute.arch.rcp_approx(cutlass.Float32(1.0) + cute.math.exp(-x, fastmath=True)) + dks.append( + (cute.math.exp(lower_bound * sig, fastmath=True), s_k.load(idx=tg * K + c), c) + ) + (d0, k0v, c0), (d1, k1v, c1) = dks + bk0, bk1 = cute.arch.mul_packed_f32x2((r_beta_g, r_beta_g), (k0v, k1v)) + kd0, kd1 = cute.arch.mul_packed_f32x2((d0, d1), (k0v, k1v)) + s_dec.store(d0, idx=tg * K + c0) + s_dec.store(d1, idx=tg * K + c1) + s_kd.store(kd0, idx=tg * K + c0) + s_kd.store(kd1, idx=tg * K + c1) + s_bk.store(bk0, idx=tg * K + c0) + s_bk.store(bk1, idx=tg * K + c1) + prims.barrier_cta_sync(0) + + # ---- Recurrence over the verify tokens (register-resident state). + st = [r_st.load(idx=j) for j in range(ROWS * VEC)] + outs = [] + for tr in cutlass.range_constexpr(NT): + r_v_val = s_v.load(idx=tr * V_CTA + warp * ROWS + lane % ROWS) + r_q = [cutlass.Float32(0.0)] * VEC + r_k = [cutlass.Float32(0.0)] * VEC + r_decay = [cutlass.Float32(0.0)] * VEC + r_bk = [cutlass.Float32(0.0)] * VEC + for i in cutlass.range_constexpr(VEC): + ki = i * 32 + lane + r_q[i] = s_q.load(idx=tr * K + ki) + r_k[i] = s_kd.load(idx=tr * K + ki) + r_bk[i] = s_bk.load(idx=tr * K + ki) + r_decay[i] = s_dec.load(idx=tr * K + ki) + ra = 0 + rb = 1 + r_va = cute.arch.shuffle_sync(r_v_val, ra) + r_vb = cute.arch.shuffle_sync(r_v_val, rb) + shk_a1 = cutlass.Float32(0.0) + shk_a2 = cutlass.Float32(0.0) + shk_b1 = cutlass.Float32(0.0) + shk_b2 = cutlass.Float32(0.0) + for _pi in cutlass.range_constexpr(VEC // 2): + _p = _pi * 2 + shk_a1, shk_a2 = cute.arch.fma_packed_f32x2( + src_a=(st[ra * VEC + _p], st[ra * VEC + _p + 1]), + src_b=(r_k[_p], r_k[_p + 1]), + src_c=(shk_a1, shk_a2), + ) + shk_b1, shk_b2 = cute.arch.fma_packed_f32x2( + src_a=(st[rb * VEC + _p], st[rb * VEC + _p + 1]), + src_b=(r_k[_p], r_k[_p + 1]), + src_c=(shk_b1, shk_b2), + ) + shk_a = shk_a1 + shk_a2 + shk_b = shk_b1 + shk_b2 + for offset in [16, 8, 4, 2, 1]: + shk_a += cute.arch.shuffle_sync_bfly(shk_a, offset=offset, mask=-1, mask_and_clamp=31) + shk_b += cute.arch.shuffle_sync_bfly(shk_b, offset=offset, mask=-1, mask_and_clamp=31) + vn_a = r_va - shk_a + vn_b = r_vb - shk_b + shq_a1 = cutlass.Float32(0.0) + shq_a2 = cutlass.Float32(0.0) + shq_b1 = cutlass.Float32(0.0) + shq_b2 = cutlass.Float32(0.0) + for _pi in cutlass.range_constexpr(VEC // 2): + _p = _pi * 2 + vnbk_a0, vnbk_a1 = cute.arch.mul_packed_f32x2((vn_a, vn_a), (r_bk[_p], r_bk[_p + 1])) + vnbk_b0, vnbk_b1 = cute.arch.mul_packed_f32x2((vn_b, vn_b), (r_bk[_p], r_bk[_p + 1])) + st[ra * VEC + _p], st[ra * VEC + _p + 1] = cute.arch.fma_packed_f32x2( + src_a=(r_decay[_p], r_decay[_p + 1]), + src_b=(st[ra * VEC + _p], st[ra * VEC + _p + 1]), + src_c=(vnbk_a0, vnbk_a1), + ) + st[rb * VEC + _p], st[rb * VEC + _p + 1] = cute.arch.fma_packed_f32x2( + src_a=(r_decay[_p], r_decay[_p + 1]), + src_b=(st[rb * VEC + _p], st[rb * VEC + _p + 1]), + src_c=(vnbk_b0, vnbk_b1), + ) + shq_a1, shq_a2 = cute.arch.fma_packed_f32x2( + src_a=(st[ra * VEC + _p], st[ra * VEC + _p + 1]), + src_b=(r_q[_p], r_q[_p + 1]), + src_c=(shq_a1, shq_a2), + ) + shq_b1, shq_b2 = cute.arch.fma_packed_f32x2( + src_a=(st[rb * VEC + _p], st[rb * VEC + _p + 1]), + src_b=(r_q[_p], r_q[_p + 1]), + src_c=(shq_b1, shq_b2), + ) + # The q-dot feeds only the output, not the next token: its reduction waits until after the loop. + outs.append([shq_a1 + shq_a2, shq_b1 + shq_b2]) + # The golden token's state to the pool; a draft's record: the warp's row innovations (lane r, row r). + if cutlass.const_expr(tr == 0): + for r in cutlass.range_constexpr(ROWS): + for i in cutlass.range_constexpr(VEC): + pool.store(st[r * VEC + i], idx=st_base + r * K + i * 32 + lane) + else: + if lane < ROWS: + tok.store( + cutlass.Float32(cutlass.select_(lane == 0, vn_a, vn_b)), + idx=CT_VN + ((tr - 1) * H + h) * V + row_a + lane, + ) + + for offset in [16, 8, 4, 2, 1]: + for tr in cutlass.range_constexpr(NT): + for r in cutlass.range_constexpr(ROWS): + outs[tr][r] += cute.arch.shuffle_sync_bfly( + outs[tr][r], offset=offset, mask=-1, mask_and_clamp=31 + ) + if lane == 0: + for tr in cutlass.range_constexpr(NT): + s_o.store(_bf16(outs[tr][0]), idx=tr * V_CTA + warp * ROWS + 0) + s_o.store(_bf16(outs[tr][1]), idx=tr * V_CTA + warp * ROWS + 1) + + # ---- Output gated RMSNorm over V: per-token partial sums of squares to every CTA of the cluster. + prims.barrier_cta_sync(0) + if tidx < NT * V_CTA: + t_out = tidx // V_CTA + x_o = s_o.load(idx=tidx) + ss = x_o * x_o + for off_ss in [V_CTA >> d for d in range(1, V_CTA.bit_length()) if V_CTA >> d > 0]: + ss = ss + cute.arch.shuffle_sync_bfly(ss, offset=off_ss, mask=-1, mask_and_clamp=31) + if tidx % V_CTA == 0: + # st.async into every CTA of the cluster (this one too), completing on its mailbox barrier: no cluster + # barrier, so nothing waits for this CTA's state stores to drain. + for r in cutlass.range_constexpr(V_SPLIT): + _st_async_f32( + _mapa_u32(s_ss.subview(vs * NT + t_out).data_ptr(), r), + ss, + _mapa_u32(ss_ready.data_ptr(), r), + ) + # Every peer's partials are in this CTA's shared memory once its mailbox completes. Each peer sends only after + # its recurrence, i.e. after it consumed its reads of the q/k conv history, so CTA 0's rewrite below cannot race + # them; and no CTA exits before its own mailbox has received every peer's partials. + while not _try_wait_cluster(ss_ready.data_ptr(), 0): + pass + if tidx < NT: + total = cutlass.Float32(0.0) + for r in cutlass.range_constexpr(V_SPLIT): + total = total + s_ss.load(idx=r * NT + tidx) + s_rs.store(cute.math.rsqrt(total / V + eps), idx=tidx) + prims.barrier_cta_sync(0) + if tidx < NT * V_CTA: + t_y = tidx // V_CTA + v_y = tidx % V_CTA + z = s_og.load(idx=tidx) + gate = cutlass.Float32(1.0) / (cutlass.Float32(1.0) + cute.math.exp(-z, fastmath=True)) + y = s_o.load(idx=tidx) * s_rs.load(idx=t_y) * s_onw.load(idx=v_y) * gate + out.store(cutlass.BFloat16(y), idx=((row0 + t_y) * H + h) * V + v0 + v_y) + + # ---- The drafts' per-key records (beta * k and the decay), one CTA per head, after every peer's norm partials + # have arrived: each peer sends them after its recurrence, i.e. after its prologue read the last round's records. + if vs == 0: + for it_r in cutlass.range_constexpr((NUM_SPEC * K + THREADS - 1) // THREADS): + item_r = tidx + it_r * THREADS + if item_r < NUM_SPEC * K: + t_r = item_r // K + key_r = item_r % K + tok.store( + s_bk.load(idx=(t_r + 1) * K + key_r), + idx=CT_WB + (t_r * H + h) * K + key_r, + ) + tok.store( + s_dec.load(idx=(t_r + 1) * K + key_r), + idx=CT_WD + (t_r * H + h) * K + key_r, + ) + + # ---- Conv caches for the next round: column s = position s - 2 = row s + 1 of the position-indexed inputs. + if vs == 0: + if tidx < K: + for s in cutlass.range_constexpr(S): + cs_q.store(s_uq.load(idx=(s + 1) * K + tidx), idx=(slot * S + s) * HK + ch0 + tidx) + else: + ck_c = tidx - K + for s in cutlass.range_constexpr(S): + cs_k.store(s_uk.load(idx=(s + 1) * K + ck_c), idx=(slot * S + s) * HK + ch0 + ck_c) + for it in cutlass.range_constexpr((S * V_CTA + THREADS - 1) // THREADS): + item = tidx + it * THREADS + if item < S * V_CTA: + s_col = item // V_CTA + v_col = item % V_CTA + cs_v.store( + s_uv.load(idx=(s_col + 1) * V_CTA + v_col), + idx=(slot * S + s_col) * HK + ch0 + v0 + v_col, + ) + + +@cute.jit +def k3_kda_verify( + w_fb: cute.Tensor, # bf16 [H * K, K]: the f_b weight (out, in) + proj: cute.Tensor, # int32 words of the fused projection [T, proj_words] + g_ext: cute.Tensor, # int32 words of the unfused f_b output [T, H * K / 2] (FOLD_FB False) + w_q: cute.Tensor, + w_k: cute.Tensor, + w_v: cute.Tensor, + a_log: cute.Tensor, + dt_bias: cute.Tensor, + onorm_w: cute.Tensor, + cs_q: cute.Tensor, + cs_k: cute.Tensor, + cs_v: cute.Tensor, + ssm: cute.Tensor, + state_tok: cute.Tensor, + slots: cute.Tensor, + pending: cute.Tensor, + out: cute.Tensor, + proj_words: cutlass.Int32, + n_req: cutlass.Int32, + ssm_stride: cutlass.Int64, + H: cutlass.Constexpr[int], + NUM_SPEC: cutlass.Constexpr[int], + lower_bound: cutlass.Constexpr[float], + scale: cutlass.Constexpr[float], + eps: cutlass.Constexpr[float], + FOLD_FB: cutlass.Constexpr[bool], + USE_PDL: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + """One launch of the verify kernel for ``n_req`` requests of 1 + NUM_SPEC tokens each.""" + # W_fb [H K, K] as five TMA dimensions (64-element column chunk, row, chunk index, 1, 1): one call lands the + # head's 128 rows, both 128-byte-swizzled halves; strides in 16-byte units. + tma_wfb = cuda.create_tensor_map_tiled( + global_address=w_fb.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[TMA_K_BOX, H * K, K // TMA_K_BOX, 1, 1], + global_strides=[(K * BF16_BYTES) // 16, (TMA_K_BOX * BF16_BYTES) // 16, (H * K * K * BF16_BYTES) // 16, + (H * K * K * BF16_BYTES) // 16], + box_dims=[TMA_K_BOX, CTA_M, TMA_COPY_ITERS, 1, 1], + swizzle=cuda.TensorMapSwizzle.s128b, + ) # fmt: skip + # The projection rows as bf16 [T, 2 proj_words]: an 8-token x 64-column box of f_a per call. + tma_fa = cuda.create_tensor_map_tiled( + global_address=proj.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[proj_words * 2, n_req * (NUM_SPEC + 1)], + global_strides=[(proj_words * 4) // 16], + box_dims=[TMA_K_BOX, MMA_N], + swizzle=cuda.TensorMapSwizzle.s128b, + ) + k3_kda_verify_kernel( + tma_wfb, tma_fa, proj, g_ext, w_q, w_k, w_v, a_log, dt_bias, onorm_w, cs_q, cs_k, cs_v, ssm, state_tok, slots, + pending, out, proj_words, ssm_stride, H, NUM_SPEC, lower_bound, scale, eps, FOLD_FB, USE_PDL, + ).launch( + grid=(H, n_req, V_SPLIT), block=(THREADS, 1, 1), cluster=(1, 1, V_SPLIT), stream=stream, use_pdl=USE_PDL, + ) # fmt: skip diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_verify/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_verify/op.py new file mode 100644 index 000000000000..7e181d958fec --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_verify/op.py @@ -0,0 +1,205 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``trtllm::k3_kda_verify``: Kimi K3's KDA speculative verify for 1 + num_spec tokens per request, from the fused +projection rows to the gated-norm core output, with a committed state per verify token. + +State contract (differs from ``trtllm::kda_mtp_decode``'s replay caches): after a call the pool ``ssm`` holds the +state after each request's golden token and ``state_tok[slot, t - 1]`` the state after its draft t; the conv caches +hold the raw inputs at positions -2..num_spec around the golden token. The next call reads the state after the +drafts the sampler accepted (``pending[slot]``) directly, so nothing is replayed. The kernel compiles on the first +call for each configuration, which must happen outside CUDA-graph capture. +""" + +from __future__ import annotations + +import os +import threading +from typing import Dict, Optional + +import torch + +K = 128 +CONV_W = 4 + +_lock = threading.Lock() +_compiled: Dict[tuple, object] = {} + + +def _arg(t: torch.Tensor, align: int = 16): + from cutlass.cute.runtime import from_dlpack + + # detach(): DLPack refuses tensors that require grad (weights are parameters). + return from_dlpack(t.detach(), assumed_align=align).mark_layout_dynamic(leading_dim=t.dim() - 1) + + +def _index_arg(t: torch.Tensor): + """An int32 index tensor the kernel reads element by element (``slots``, ``pending``), declared at its element's + alignment: the mixer passes slices such as ``state_indices[num_prefills:]``, which start on any 4-byte boundary.""" + return _arg(t, t.element_size()) + + +def _flat(t: torch.Tensor) -> torch.Tensor: + """A 1-D view of a dense tensor's storage (a view with its last two dims transposed is fine), never a copy.""" + if t.is_contiguous(): + return t.view(-1) + swapped = t.transpose(-1, -2) + if swapped.is_contiguous(): + return swapped.reshape(-1) + raise ValueError(f"expected a dense tensor, got shape {tuple(t.shape)} strides {t.stride()}") + + +def _slots_view(t: torch.Tensor) -> torch.Tensor: + """A [slots, slot elements] view of a pool whose slots are dense but may be strided (the Mamba cache manager + coalesces each slot's per-layer states), never a copy. The kernels address it flat from the first slot at 64-bit + offsets; the view keeps every extent within the DSL's 32-bit sizes however far the last slot lies.""" + return t.view(t.shape[0], -1) + + +def _words(t: torch.Tensor) -> torch.Tensor: + return _flat(t).view(torch.int32) + + +def use_pdl() -> bool: + return os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" + + +@torch.library.custom_op( + "trtllm::k3_kda_verify", + mutates_args=("cs_q", "cs_k", "cs_v", "ssm", "state_tok"), + device_types="cuda", +) +def k3_kda_verify( + proj: torch.Tensor, + w_fb: torch.Tensor, + w_q: torch.Tensor, + w_k: torch.Tensor, + w_v: torch.Tensor, + a_log: torch.Tensor, + dt_bias: torch.Tensor, + onorm_w: torch.Tensor, + cs_q: torch.Tensor, + cs_k: torch.Tensor, + cs_v: torch.Tensor, + ssm: torch.Tensor, + state_tok: torch.Tensor, + slots: torch.Tensor, + pending: torch.Tensor, + num_spec: int, + lower_bound: float, + scale: float, + eps: float, + g_ext: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Gated-norm KDA core output [T, H, V] bf16 for T = N * (1 + num_spec) verify tokens. + + ``proj`` bf16 [T, 4 H K + K + H + pad]: the fused projection rows [q | k | v | onorm gate | f_a | b | pad]; + ``w_fb`` bf16 [H K, K], the f_b weight (out, in); ``w_q/w_k/w_v`` fp32 [H K, 4]; ``a_log`` fp32 [H]; + ``dt_bias`` fp32 [H K]; ``onorm_w`` fp32 [V]; ``cs_*`` fp32 [pool, H K, 3 + num_spec] (dim-contiguous); + ``ssm`` fp32 [pool, H, V, K] (each slot dense, slots at any stride); ``state_tok`` fp32 + [pool, num_spec, H, V, K]; ``slots`` int32 [N]; ``pending`` int32 [pool]. With ``g_ext`` (bf16 [T, H K], the + unfused f_b output) the gate is read instead of computed from f_a.""" + import cuda.bindings.driver as cuda_driver + + from . import k3_kda_verify_kernel as kernel + + n_req = slots.shape[0] + num_heads = w_fb.shape[0] // K + tokens = n_req * (1 + num_spec) + v_dim = onorm_w.shape[0] + if ( + proj.dtype != torch.bfloat16 + or proj.dim() != 2 + or proj.shape[0] != tokens + or proj.shape[1] % 8 != 0 + or proj.shape[1] < 4 * num_heads * K + K + num_heads + or not proj.is_contiguous() + or tuple(w_fb.shape) != (num_heads * K, K) + or w_fb.dtype != torch.bfloat16 + or not w_fb.is_contiguous() + or v_dim != K + or tuple(ssm.shape[1:]) != (num_heads, K, K) + or tuple(state_tok.shape[1:]) != (num_spec, num_heads, K, K) + or ssm.stride()[1:] != (K * K, K, 1) + or ssm.dtype != torch.float32 + or not state_tok.is_contiguous() + or state_tok.dtype != torch.float32 + or cs_q.shape[-1] != CONV_W - 1 + num_spec + or slots.dtype != torch.int32 + or pending.dtype != torch.int32 + or any( + t.dtype != torch.float32 + for t in (w_q, w_k, w_v, a_log, dt_bias, onorm_w, cs_q, cs_k, cs_v) + ) + ): + raise ValueError( + f"k3_kda_verify: unsupported call (w_q/w_k/w_v/a_log/dt_bias/onorm_w/cs_* must be fp32: " + f"{[str(t.dtype) for t in (w_q, w_k, w_v, a_log, dt_bias, onorm_w, cs_q, cs_k, cs_v)]}) " + f"proj {tuple(proj.shape)} {proj.dtype} contiguous {proj.is_contiguous()}, " + f"w_fb {tuple(w_fb.shape)} {w_fb.dtype} contiguous {w_fb.is_contiguous()}, onorm_w {tuple(onorm_w.shape)}, " + f"ssm {tuple(ssm.shape)} {ssm.dtype} strides {ssm.stride()}, state_tok {tuple(state_tok.shape)} " + f"{state_tok.dtype} contiguous {state_tok.is_contiguous()}, cs_q {tuple(cs_q.shape)}, num_spec {num_spec}, " + f"slots {tuple(slots.shape)} {slots.dtype}, pending {pending.dtype} (tokens expected {tokens})" + ) + fold_fb = g_ext is None + out = torch.empty(tokens, num_heads, v_dim, dtype=torch.bfloat16, device=proj.device) + args = ( + _arg(w_fb), + _arg(_words(proj)), + _arg(_words(proj if fold_fb else g_ext)), + _arg(_flat(w_q)), + _arg(_flat(w_k)), + _arg(_flat(w_v)), + _arg(_flat(a_log)), + _arg(_flat(dt_bias)), + _arg(_flat(onorm_w)), + _arg(_flat(cs_q)), + _arg(_flat(cs_k)), + _arg(_flat(cs_v)), + _arg(_slots_view(ssm)), + _arg(_slots_view(state_tok)), + _index_arg(slots), + _index_arg(pending), + _arg(_flat(out)), + ) + runtime = (proj.shape[1] // 2, n_req, ssm.stride(0)) + stream = cuda_driver.CUstream(torch.cuda.current_stream(proj.device).cuda_stream) + pdl = use_pdl() + key = (num_heads, num_spec, float(lower_bound), float(scale), float(eps), fold_fb, pdl) + fn = _compiled.get(key) + if fn is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "trtllm::k3_kda_verify must run once per configuration outside CUDA-graph capture first " + "(it compiles its kernel on the first call)." + ) + import cutlass.cute as cute + + with _lock: + fn = _compiled.get(key) + if fn is None: + fn = _compiled[key] = cute.compile( + kernel.k3_kda_verify, *args, *runtime, num_heads, num_spec, float(lower_bound), float(scale), + float(eps), fold_fb, pdl, stream, + ) # fmt: skip + fn(*args, *runtime, stream) + return out + + +@k3_kda_verify.register_fake +def _(proj, w_fb, w_q, w_k, w_v, a_log, dt_bias, onorm_w, cs_q, cs_k, cs_v, ssm, state_tok, slots, pending, + num_spec, lower_bound, scale, eps, g_ext=None): # fmt: skip + return proj.new_empty( + (proj.shape[0], w_fb.shape[0] // K, onorm_w.shape[0]), dtype=torch.bfloat16 + ) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/__init__.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/__init__.py new file mode 100644 index 000000000000..76bebda7a90b --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/__init__.py @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Kimi K3 MLA decode CTM kernels (``trtllm::k3_mla_q``). + +Importing :mod:`.op` registers the torch ops; nothing is imported eagerly here so that +the CuTe DSL dependency stays optional for every other model. +""" diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/k3_mla_attn_kernel.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/k3_mla_attn_kernel.py new file mode 100644 index 000000000000..9d179d489649 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/k3_mla_attn_kernel.py @@ -0,0 +1,1187 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +# ============================================================================= +# Kimi K3 MLA decode attention -- CTM (prims/cute) kernel: R requests of T <= 8 query tokens, 6 heads per rank +# ============================================================================= +# +# o[i, t, h] = softmax_k( scale * q[i, t, h] . kv_i[k] ) @ kv_i[k, :512] k <= L_i - T + t (bottom-right causal) +# q = fused_q [M = R T, 6, 576], request-major (request i's tokens are rows i T .. i T + T - 1); per cluster 48 +# query rows r = 8 h + t, head-major (a head's tokens are adjacent rows) +# kv_i = request i's pages of the paged latent cache, 64 rows x 576 (512 latent | 64 rope), L_i rows +# o = [M, 6, 512] bf16 +# +# One cluster of 16 CTAs per (request, group of 6 heads): grid (16 x groups, R), one group at TP16 (a rank holding more +# heads gets one cluster per group, each reading the request's whole cache). A cluster's work does not depend on R; the +# request only selects its page-table row, length, query / output rows and workspace slots. CTA c takes the 128-row KV +# tiles c, c + 16, ... (two pages each). +# Per tile: +# S = Q K^T A = Q [64 rows (48 real), 576] from shared memory, B = the tile's K [128 rows, 576] (chunk-major: +# page p's rows of chunk j at [j][64 p..]), M 64, N 128 (Q read once per chunk), into TMEM (an M = 64 +# accumulator: row 16 w + l in lane 32 w + l) +# softmax each row's 128 values over the 4 lanes of a quad (tcgen05.ld 16x256b): masked max, exp2, sum, +# P = bf16(p) into a 128B-swizzled K-major tile, fenced to the async proxy +# O = P V A = P [64, 128], B = V = the tile's latent columns read MN-major (N 256 per MMA), M 64 +# partial (m, l) and O / l (a convex combination of V rows: fp16 keeps 2^-11) of the tile's 48 rows into the +# CTA's global workspace slot; a CTA's later tiles (L > 2048) fold into its slot: +# O / l' = (O / l) (l a / l') + O_k (b / l'), l' = l a + l_k b, a, b = exp2(m - m'), exp2(m_k - m') +# Right after its last softmax a CTA st.async's its (m, l) per row into every CTA's mailbox; warps 1 and 3 turn the +# mailbox into the merge weights exp2(m_s - M) / sum_s exp2(m_s - M) l_s while O is computed and drained. Then a +# cluster barrier (release / acquire), and CTA c merges latent columns [32 c, 32 c + 32) of all rows over the slots +# (flash-decoding combine in fp32, one round of 16-byte loads) and stores o in bf16. The KV pages below L - T are +# TMA'd before griddepcontrol.wait (earlier steps wrote them); Q (3 chunk groups, so S starts on the first) and the +# last T rows' pages after it. O = P V runs in two N = 256 halves, the first drained while the second computes. +# Q is TMA'd as 8 token rows from the request's first token: rows t >= T (the next request's tokens, or zeros past M) +# are computed and written to the workspace slot like the live rows (so the merge's 16-byte loads read 32-byte sectors +# written whole; it weights these rows 0), and never stored. +# +# fuse_vb: the merge is split by (head, 4 tokens) over CTAs 0-11 instead of by columns, and each of those CTAs applies +# v_b to its 4 merged rows: W_vb[h] (128 x 512) TMA'd into the dead KV buffer after the CTA's last P V, o rounded to +# bf16 (the unfused path's attention output) into the dead P buffer as the B operand, y^T = W_vb[h] o^T (M 128, N 8, +# K 512) in TMEM, y [M, heads * 128] bf16 stored (the o_proj input). apply_gate: y is multiplied by the output gate's +# sigmoid (bf16, from the fused projection's gate columns) with the unfused path's rounding, bf16(bf16(y) * s). +# +# no_cluster (R x groups > CLUSTER_WAVE: more 16-CTA clusters than co-reside, which would run a second wave): the grid +# launches without a cluster, CTA c = block x mod 16, and the cluster's mailbox and barriers go through the workspace's +# tail instead. At its last tile a CTA stores its live rows' (m, l) into its slot of an fp32 exchange and arrives on its +# (request, head group)'s first counter; warps 1 and 3 wait for the launch's 16 arrivals and read the slots' (m, l) +# from there. After the drain every CTA arrives on the second counter, which CTAs 0-11 wait for instead of the cluster +# barrier. Counters only grow: a CTA reads its counter after the grid wait and before it arrives, so the launch's 16 +# arrivals end at (count & ~15) + 16. The merge sums the same values in the same order, so the outputs do not depend on +# the mode. All 16 R groups CTAs must co-reside (1 per SM), as their waits are spins. +# +# Warps: 0 TMA, 2 TMEM allocation (512 columns) + MMA, 4-7 softmax / partial writer, 1 and 3 merge weights, all 8 in +# the merge. +# ============================================================================= +"""CTM kernel: Kimi K3 MLA decode attention over the paged latent cache, split-KV over a 16-CTA cluster.""" + +from __future__ import annotations + +import math + +import cuda.bindings.driver as cuda_driver +import cutlass +import cutlass.cute as cute +import cutlass.experimental.cuda as cuda +from cutlass._mlir.dialects import llvm as _llvm +from cutlass._mlir.extras import types as _T +from cutlass.cutlass_dsl import dsl_user_op +from cutlass.experimental import primitives as prims + +HEADS = 6 +MAX_TOKENS = 8 # query tokens per request (T) +MAX_REQUESTS = 8 # requests per launch (R): the workspace holds R x groups x 16 slots +ROWS = HEADS * MAX_TOKENS # 48 query rows +# A softmax thread's rows r and r + 8 lie in one 16-row block, so they are both real or both past ROWS. +assert ROWS % 16 == 0 +QK = 576 +LATENT = 512 +PAGE = 64 +TILE = 128 # KV rows per tile = 2 pages +PAGES_PER_TILE = TILE // PAGE +CHUNK = 64 # 128-byte swizzle width in bf16 +QK_CHUNKS = QK // CHUNK # 9 +LAT_CHUNKS = LATENT // CHUNK # 8 +CLUSTER = 16 +MMA_M = 64 +MMA_K = 16 +THREADS = 256 +EPI_THREADS = 128 +TMEM_COLS = 512 +ELEM_BYTES = 2 +LAT_PER_CTA = LATENT // CLUSTER # 32 merged columns per CTA +LOG2E = 1.4426950408889634 +NEG_INF = float("-inf") + +# Shared-memory tiles, bf16 elements. +Q_CHUNK_ELEMS = ROWS * CHUNK # 3072 (6 KB): Q chunk j = rows 0..47 of columns [64 j, 64 j + 64) +Q_ELEMS = QK_CHUNKS * Q_CHUNK_ELEMS +KV_CHUNK_ELEMS = ( + TILE * CHUNK +) # 8192 (16 KB): both pages' rows of one 64-column chunk (chunk-major tile) +KV_PAGE_OFF = PAGE * CHUNK # page p's 64 rows start 8 KB into each chunk block +KV_ELEMS = QK_CHUNKS * KV_CHUNK_ELEMS +P_HALF_ELEMS = MMA_M * CHUNK # 4096: 64 rows x 64 KV columns +P_ELEMS = PAGES_PER_TILE * P_HALF_ELEMS +Q_OVERREAD_ELEMS = (MMA_M - ROWS) * CHUNK # the M = 64 MMA reads 16 rows past the last Q chunk + +# Descriptor offsets in 16-byte units. +LEADING = 16 +SBO = 8 * CHUNK * ELEM_BYTES # 1024 B between 8-row swizzle atoms +STEP_K = (MMA_K * ELEM_BYTES) >> 4 # K-major: 32 B per K16 step +STEP_MN = (2 * SBO) >> 4 # MN-major: 2 x SBO per K16 step (16 KV rows) +Q_CHUNK_U = (Q_CHUNK_ELEMS * ELEM_BYTES) >> 4 +KV_CHUNK_U = (KV_CHUNK_ELEMS * ELEM_BYTES) >> 4 +P_HALF_U = (P_HALF_ELEMS * ELEM_BYTES) >> 4 + +Q_GROUPS = 3 # Q TMA'd as 3 groups of 3 chunks, each on its own barrier +Q_GROUP_CHUNKS = QK_CHUNKS // Q_GROUPS +O_HALVES = LATENT // 256 +MERGE_HALF = CLUSTER // 2 # slots per thread of a merge pair +V_DIM = 128 # v_head_dim +VB_CTAS = ( + HEADS * 2 +) # fuse_vb: CTA c < 12 merges head c // 2, tokens 4 (c % 2) .. + 4, and applies v_b to them +VB_TOKENS = MAX_TOKENS // 2 +VB_ELEMS = V_DIM * LATENT + +io_dtype = cutlass.BFloat16 +ws_dtype = cutlass.Float16 # the partials O / l (convex combinations of V rows, |x| <= max |V|) +# Workspace: slot (request x groups + head group) x 16 + CTA. Slot layout, fragment-major: [64 groups of 8 columns][64 +# rows (the M = 64 accumulator)][8], so one warp's 16x256b fragment of a group (8 rows x 4 lanes x 2 values) is 128 +# contiguous bytes, and a row's 8 columns are 16. +WS_GROUP_ELEMS = MMA_M * 8 +WS_SLOT_ELEMS = (LATENT // 8) * WS_GROUP_ELEMS +# 16-CTA clusters of this kernel that co-reside on GB200 (cuOccupancyMaxActiveClusters at its shared memory: the eighth +# GPC holds 12 free SMs); a launch of more clusters takes the no_cluster mode. +CLUSTER_WAVE = 7 +# no_cluster's tail of the workspace, after the R x groups x 16 partial slots: an fp32 [slot][row][m, l] exchange, then +# two int32 counters per (request, head group), CTR_STRIDE words apart (a 128-byte line each). +CTR_STRIDE = 32 + + +def ws_sync_elems(groups: int) -> int: + """fp16 elements of the workspace's no_cluster tail (zeroed once: the counters start at 0).""" + exchange = MAX_REQUESTS * groups * CLUSTER * ROWS * 2 * 4 + counters = MAX_REQUESTS * groups * 2 * CTR_STRIDE * 4 + return (exchange + counters) // 2 + + +@dsl_user_op +def _mapa_u32(smem_ptr, peer, *, loc=None, ip=None): + """The shared::cluster address of this CTA's shared-memory location in cluster CTA ``peer``.""" + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [smem_ptr.toint(loc=loc, ip=ip).ir_value(), cutlass.Int32(peer).ir_value(loc=loc, ip=ip)], + "mapa.shared::cluster.u32 $0, $1, $2;", "=r,r,r", has_side_effects=False, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _st_async_v2(dst, a, b, mbar, *, loc=None, ip=None): + """st.async of two fp32 to a shared::cluster address, completing ``mbar`` (shared::cluster) by 8 bytes.""" + _llvm.inline_asm( + None, + [cutlass.Int32(dst).ir_value(loc=loc, ip=ip), cutlass.Float32(a).ir_value(loc=loc, ip=ip), + cutlass.Float32(b).ir_value(loc=loc, ip=ip), cutlass.Int32(mbar).ir_value(loc=loc, ip=ip)], + "st.async.shared::cluster.mbarrier::complete_tx::bytes.v2.f32 [$0], {$1, $2}, [$3];", "r,f,f,r", + has_side_effects=True, is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _red_add_release(addr, value, *, loc=None, ip=None): + """red.release.gpu.global.add.s32 on a global address (an arrival: no value back).""" + _llvm.inline_asm( + None, [cutlass.Int64(addr).ir_value(loc=loc, ip=ip), cutlass.Int32(value).ir_value(loc=loc, ip=ip)], + "red.release.gpu.global.add.s32 [$0], $1;", "l,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _ld_acquire(addr, *, loc=None, ip=None): + """ld.acquire.gpu.global.s32 of a global address.""" + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [cutlass.Int64(addr).ir_value(loc=loc, ip=ip)], + "ld.acquire.gpu.global.s32 $0, [$1];", "=r,l", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _try_wait_cluster(mbar, parity, *, loc=None, ip=None): + """mbarrier.try_wait.parity.acquire.cluster on a barrier of this CTA (``mbar``, its shared-memory pointer): whether + phase ``parity`` has completed, acquiring at cluster scope. The barrier is completed by other CTAs' st.async, whose + complete_tx releases at cluster scope.""" + done = cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [mbar.toint(loc=loc, ip=ip).ir_value(), cutlass.Int32(parity).ir_value(loc=loc, ip=ip)], + "{\n\t.reg .pred p;\n\tmbarrier.try_wait.parity.acquire.cluster.shared::cta.b64 p, [$1], $2;\n\t" + "selp.u32 $0, 1, 0, p;\n\t}", "=r,r,r,~{memory}", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + return done != cutlass.Int32(0) + + +@cute.kernel +def k3_mla_attn_kernel( + tma_q: cutlass.GridConstant[ + cuda.TensorMap + ], # fused_q [M, heads, 576]: box 64 cols x 8 tokens x 6 heads x 3 chunks + tma_kv: cutlass.GridConstant[ + cuda.TensorMap + ], # pool as [pages * 64 rows, 576]: box 64 cols x 64 rows x 9 chunks + tma_vb: cutlass.GridConstant[ + cuda.TensorMap + ], # fuse_vb: v_b_proj as [heads * 128 rows, 512]: box 64 x 128 x 8 + page_table: cutlass.Array, # int32: request i's pages (>= ceil(L_i / 64)) at [i * pt_stride, ...) + seq_len: cutlass.Array, # int32 [R]: L_i, request i's KV rows including its T new ones + ws_o: cutlass.Array, # fp16 [R, head groups, 16 CTAs, 64 column groups, 64 rows, 8] per-CTA partial O / l + out: cutlass.Array, # bf16 [M * heads * 512] (fuse_vb: [M * heads * 128]) + gate: cutlass.Array, # apply_gate: bf16 [M, gate_ld], sigmoid(g) of head h at columns gate_col0 + 128 h + tokens: cutlass.Int32, # T, the query tokens of each request (M = R T) + pt_stride: cutlass.Int32, # page-table elements between consecutive requests' rows + scale_log2: cutlass.Float32, # softmax scale * log2(e) + page_offset: cutlass.Int32, # added to every page-table entry (the layer's slot in an interleaved pool) + gate_col0: cutlass.Int32, + gate_ld: cutlass.Int32, + total_heads: cutlass.Constexpr[int], # heads of the rank: one cluster per 6 + fuse_vb: cutlass.Constexpr[bool], + apply_gate: cutlass.Constexpr[bool], + no_cluster: cutlass.Constexpr[bool], +): + tx, _, _ = cute.arch.thread_idx() + warp_id = cute.arch.warp_idx() + if cutlass.const_expr(not no_cluster): + cta = cute.arch.block_idx_in_cluster() + bid, req, _ = cute.arch.block_idx() + if cutlass.const_expr(no_cluster): + cta = bid % cutlass.Int32(CLUSTER) + hg = bid // cutlass.Int32(CLUSTER) # head group: heads 6 hg .. 6 hg + 5 + ws_slot0 = (req * cutlass.Int32(total_heads // HEADS) + hg) * cutlass.Int32(CLUSTER) + my_slot = ws_slot0 + cta + tok0 = req * tokens # the request's first row of q, out and gate + pt0 = req * pt_stride # the request's page-table row + ptr_q = tma_q.get_ptr() + ptr_kv = tma_kv.get_ptr() + ptr_vb = tma_vb.get_ptr() + + # Read before the grid dependency: the host writes the lengths and page table before the step. + kv_len = seq_len.load(idx=req) + n_tiles = (kv_len + cutlass.Int32(TILE - 1)) // cutlass.Int32(TILE) + n_pages = (kv_len + cutlass.Int32(PAGE - 1)) // cutlass.Int32(PAGE) + my_count = cutlass.Int32(0) + if cta < n_tiles: + my_count = (n_tiles - cta + cutlass.Int32(CLUSTER - 1)) // cutlass.Int32(CLUSTER) + q_rows = tokens * cutlass.Int32(HEADS) # live rows: 8 h + t with t < T + # Pages holding a row >= L - T are written this step (the cache append): they load after the wait. + fresh_page = (kv_len - tokens) // cutlass.Int32(PAGE) + # Slots (CTAs) that had a tile: min(n_tiles, 16). + n_valid = cutlass.Int32( + cutlass.select_(n_tiles < cutlass.Int32(CLUSTER), n_tiles, cutlass.Int32(CLUSTER)) + ) + + smem_kv = cutlass.Array(io_dtype, KV_ELEMS, space=cutlass.AddressSpace.smem, alignment=1024) + smem_q = cutlass.Array(io_dtype, Q_ELEMS, space=cutlass.AddressSpace.smem, alignment=1024) + # Right after Q so the M = 64 MMA's read past the last Q chunk stays in shared memory. + smem_p = cutlass.Array(io_dtype, P_ELEMS, space=cutlass.AddressSpace.smem, alignment=1024) + kv_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + q_full = cutlass.Array(cutlass.Int64, Q_GROUPS, space=cutlass.AddressSpace.smem, alignment=8) + s_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + p_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + o_full = cutlass.Array(cutlass.Int64, O_HALVES, space=cutlass.AddressSpace.smem, alignment=8) + o_drained = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + ml_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + vb_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + vb_done = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + tmem_ptr_i32 = cutlass.Array(cutlass.Int32, 1, space=cutlass.AddressSpace.smem) + # Mailbox: [source CTA][row] (m, l) of every CTA with a tile (st.async, completed by bytes); warps 1 and 3 then + # overwrite m with the merge weight exp2(m_s - M) / L. + ml_mail = cutlass.Array( + cutlass.Float32, CLUSTER * ROWS * 2, space=cutlass.AddressSpace.smem, alignment=16 + ) + if cutlass.const_expr(no_cluster): + # The workspace's tail: the (m, l) exchange [slot][row][2] and the counters; this CTA's targets in sync_target. + groups = total_heads // HEADS + tail = MAX_REQUESTS * groups * CLUSTER * WS_SLOT_ELEMS + ml_glob = cutlass.Array( + ws_o.data_ptr(tail), shape=(MAX_REQUESTS * groups * CLUSTER * ROWS * 2,), dtype=cutlass.Float32, + bounds_check=False, addrspace=cutlass.AddressSpace.gmem.value, alignment=16, + ) # fmt: skip + ctrs = cutlass.Array( + ws_o.data_ptr(tail + MAX_REQUESTS * groups * CLUSTER * ROWS * 4), + shape=(MAX_REQUESTS * groups * 2 * CTR_STRIDE,), dtype=cutlass.Int32, bounds_check=False, + addrspace=cutlass.AddressSpace.gmem.value, alignment=16, + ) # fmt: skip + ctr_ml = (req * cutlass.Int32(groups) + hg) * cutlass.Int32(2 * CTR_STRIDE) + ctr_o = ctr_ml + cutlass.Int32(CTR_STRIDE) + sync_target = cutlass.Array(cutlass.Int32, 2, space=cutlass.AddressSpace.smem, alignment=8) + + if warp_id == 0: + prims.prefetch_tensormap(ptr_q) + prims.prefetch_tensormap(ptr_kv) + if cutlass.const_expr(fuse_vb): + prims.prefetch_tensormap(ptr_vb) + if prims.elect_sync(): + prims.mbarrier_init(kv_full, 1) + for i in cutlass.range_constexpr(Q_GROUPS): + prims.mbarrier_init(q_full.subview(i), 1) + prims.mbarrier_init(s_full, 1) + prims.mbarrier_init(p_full, EPI_THREADS) + for i in cutlass.range_constexpr(O_HALVES): + prims.mbarrier_init(o_full.subview(i), 1) + prims.mbarrier_init(o_drained, EPI_THREADS) + prims.mbarrier_init(ml_full, 1) + prims.mbarrier_init(vb_full, 1) + prims.mbarrier_init(vb_done, 1) + if cutlass.const_expr(not no_cluster): + # Every CTA with a tile sends (m, l) of the M * 6 live rows (8 bytes each). + prims.mbarrier_arrive_expect_tx(ml_full, n_valid * q_rows * cutlass.Int32(8)) + if warp_id == 2: + prims.tcgen05_alloc(tmem_ptr_i32, TMEM_COLS) + prims.tcgen05_relinquish_alloc_permit() + prims.fence_mbarrier_init() + if cutlass.const_expr(not no_cluster): + prims.barrier_cluster_arrive_relaxed() + prims.barrier_cluster_wait() + prims.barrier_cta_sync(0) + tmem_base = tmem_ptr_i32.load() + + if warp_id == 0: + # ===================================================================== + # TMA: each tile's two pages (the pages written earlier before the grid + # wait), Q after it; the next tile once the previous one's O is done. + # ===================================================================== + if prims.elect_sync(): + for k in range(my_count): + tile = cta + k * cutlass.Int32(CLUSTER) + if k > 0: + # The KV buffer is free once the previous tile's P V is done (its last half's commit). + while not cute.arch.mbarrier_try_wait( + o_full.subview(O_HALVES - 1).data_ptr(), + (k - cutlass.Int32(1)) & cutlass.Int32(1), + ): + pass + prims.mbarrier_arrive_expect_tx(kv_full, KV_ELEMS * ELEM_BYTES) + for p in cutlass.range_constexpr(PAGES_PER_TILE): + page_idx = tile * cutlass.Int32(PAGES_PER_TILE) + cutlass.Int32(p) + page_c = cutlass.Int32( + cutlass.select_(page_idx < n_pages, page_idx, n_pages - cutlass.Int32(1)) + ) + row0 = (page_table.load(idx=pt0 + page_c) + page_offset) * cutlass.Int32(PAGE) + if k == 0: + if page_c >= fresh_page: + prims.griddepcontrol(prims.GridDepAction.WAIT) + # Chunk-major tile: page p's rows of chunk j at [j][64 p .. 64 p + 64), so one MMA covers both + # pages. + for j in cutlass.range_constexpr(QK_CHUNKS): + prims.cp_async_bulk_tensor_shared_cta_global( + smem_kv.subview(j * KV_CHUNK_ELEMS + p * KV_PAGE_OFF), + ptr_kv, + (cutlass.Int32(0), row0, cutlass.Int32(j)), + kv_full, + ) + if k == 0: + prims.griddepcontrol(prims.GridDepAction.WAIT) + for g in cutlass.range_constexpr(Q_GROUPS): + prims.mbarrier_arrive_expect_tx( + q_full.subview(g), Q_GROUP_CHUNKS * Q_CHUNK_ELEMS * ELEM_BYTES + ) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_q.subview(g * Q_GROUP_CHUNKS * Q_CHUNK_ELEMS), + ptr_q, + ( + cutlass.Int32(0), + tok0, + hg * cutlass.Int32(HEADS), + cutlass.Int32(g * Q_GROUP_CHUNKS), + ), + q_full.subview(g), + ) + # Dependents (v_b / o_proj) wait for this whole grid before reading its output. + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + if cutlass.const_expr(no_cluster): + if tx == cutlass.Int32(0): + prims.griddepcontrol(prims.GridDepAction.WAIT) + sync_target.store(ctrs.load(idx=ctr_o, is_volatile=True), idx=1) + if cutlass.const_expr(fuse_vb): + if cta < cutlass.Int32(VB_CTAS): + # W_vb[h] into the KV buffer once the CTA's last P V has read it. + if prims.elect_sync(): + if my_count > cutlass.Int32(0): + while not cute.arch.mbarrier_try_wait( + o_full.subview(O_HALVES - 1).data_ptr(), + (my_count - cutlass.Int32(1)) & cutlass.Int32(1), + ): + pass + prims.mbarrier_arrive_expect_tx(vb_full, VB_ELEMS * ELEM_BYTES) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_kv, + ptr_vb, + ( + cutlass.Int32(0), + (hg * cutlass.Int32(HEADS) + cta // cutlass.Int32(2)) + * cutlass.Int32(V_DIM), + cutlass.Int32(0), + ), + vb_full, + ) + elif warp_id == 2: + # ===================================================================== + # MMA: S = Q K^T (per page N 64), then O = P V (N 256 x 2), per tile. + # ===================================================================== + idesc_s = prims.Tcgen05InstrDesc.build( + c_dtype=cutlass.Float32, a_dtype=io_dtype, b_dtype=io_dtype, n_dim=TILE, m_dim=MMA_M + ) + idesc_o = prims.Tcgen05InstrDesc.build( + c_dtype=cutlass.Float32, + a_dtype=io_dtype, + b_dtype=io_dtype, + n_dim=256, + m_dim=MMA_M, + b_major=1, + ) + swz = prims.Tcgen05SmemSwizzle.SWIZZLE_128B + desc_q = prims.Tcgen05SmemDesc.build( + start_address=smem_q, leading_byte_offset=LEADING, stride_byte_offset=SBO, layout=swz + ) + desc_k = prims.Tcgen05SmemDesc.build( + start_address=smem_kv, leading_byte_offset=LEADING, stride_byte_offset=SBO, layout=swz + ) + desc_p = prims.Tcgen05SmemDesc.build( + start_address=smem_p, leading_byte_offset=LEADING, stride_byte_offset=SBO, layout=swz + ) + # V read MN-major: N = latent across 64-column chunks KV_CHUNK_ELEMS apart (leading byte offset). + desc_v = prims.Tcgen05SmemDesc.build( + start_address=smem_kv, + leading_byte_offset=KV_CHUNK_ELEMS * ELEM_BYTES, + stride_byte_offset=SBO, + layout=swz, + ) + tmem_s = cutlass.inttoptr(tmem_base, 6, cutlass.Int32) + for k in range(my_count): + phase = k & cutlass.Int32(1) + if k > 0: + # S overwrites O's first columns: the previous tile's O has been read out. + while not cute.arch.mbarrier_try_wait( + o_drained.data_ptr(), phase ^ cutlass.Int32(1) + ): + pass + while not cute.arch.mbarrier_try_wait(kv_full.data_ptr(), phase): + pass + # S per Q chunk group (each group's barrier completes once; later tiles pass it at once). + for g in cutlass.range_constexpr(Q_GROUPS): + while not cute.arch.mbarrier_try_wait(q_full.subview(g).data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + for jj in cutlass.range_constexpr(Q_GROUP_CHUNKS): + j = g * Q_GROUP_CHUNKS + jj + for kk in cutlass.range_constexpr(CHUNK // MMA_K): + if prims.elect_sync(): + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, prims.CTAGroup.CTA_1, tmem_s, + desc_q + (j * Q_CHUNK_U + kk * STEP_K), + desc_k + (j * KV_CHUNK_U + kk * STEP_K), + idesc_s, not (j == 0 and kk == 0), + ) # fmt: skip + if prims.elect_sync(): + prims.tcgen05_commit(s_full) + while not cute.arch.mbarrier_try_wait(p_full.data_ptr(), phase): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + # O in two N = 256 halves, each committed on its own barrier (the first is drained during the second). + for nh in cutlass.range_constexpr(O_HALVES): + for kk in cutlass.range_constexpr(TILE // MMA_K): + if prims.elect_sync(): + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, prims.CTAGroup.CTA_1, + cutlass.inttoptr(tmem_base + cutlass.Int32(nh * 256), 6, cutlass.Int32), + desc_p + ((kk // (CHUNK // MMA_K)) * P_HALF_U + (kk % (CHUNK // MMA_K)) * STEP_K), + desc_v + (nh * 4 * KV_CHUNK_U + kk * STEP_MN), + idesc_o, kk != 0, + ) # fmt: skip + if prims.elect_sync(): + prims.tcgen05_commit(o_full.subview(nh)) + elif warp_id >= 4: + # ===================================================================== + # Softmax and the partial (m, l, O) -> this CTA's workspace slot. + # 16x256b: lane l of warp w holds rows 16 w + l / 4 and + 8, columns + # 8 g + 2 (l % 4) + {0, 1} of every 8-column group g. + # ===================================================================== + lane = tx % 32 + w = warp_id - 4 + quad = lane % cutlass.Int32(4) + r0 = w * cutlass.Int32(16) + lane // cutlass.Int32(4) + r1 = r0 + cutlass.Int32(8) + # Last visible KV row of each query row (token t = r % 8): L - T + t. + lim0 = kv_len - tokens + r0 % cutlass.Int32(MAX_TOKENS) + lim1 = kv_len - tokens + r1 % cutlass.Int32(MAX_TOKENS) + # The slot takes every real row's partial (rows t >= T too), the mailboxes the live rows' (m, l). + real0 = r0 < cutlass.Int32(ROWS) + real1 = r1 < cutlass.Int32(ROWS) + live0 = real0 & (r0 % cutlass.Int32(MAX_TOKENS) < tokens) + live1 = real1 & (r1 % cutlass.Int32(MAX_TOKENS) < tokens) + # The CTA's running (m, l) per row over its tiles (the slot's O is relative to m). + m_run0 = cutlass.Float32(NEG_INF) + m_run1 = cutlass.Float32(NEG_INF) + l_run0 = cutlass.Float32(0.0) + l_run1 = cutlass.Float32(0.0) + for k in range(my_count): + phase = k & cutlass.Int32(1) + tile = cta + k * cutlass.Int32(CLUSTER) + kv0 = tile * cutlass.Int32(TILE) + while not cute.arch.mbarrier_try_wait(s_full.data_ptr(), phase): + pass + if cutlass.const_expr(no_cluster): + count_ml = ctrs.load(idx=ctr_ml, is_volatile=True) + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + s = prims.tcgen05_ld( + "16x256b", cutlass.inttoptr(tmem_base, 6, cutlass.Float32), num=TILE // 8 + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + # Masked, scaled scores and their row maxima (the quad's 4 lanes share the two rows). + m0 = cutlass.Float32(NEG_INF) + m1 = cutlass.Float32(NEG_INF) + v0 = [] + v1 = [] + for g in cutlass.range_constexpr(TILE // 8): + for e in cutlass.range_constexpr(2): + col = kv0 + cutlass.Int32(8 * g + e) + cutlass.Int32(2) * quad + a = cutlass.Float32(s[4 * g + e]) * scale_log2 + b = cutlass.Float32(s[4 * g + 2 + e]) * scale_log2 + a = cutlass.Float32(cutlass.select_(col <= lim0, a, cutlass.Float32(NEG_INF))) + b = cutlass.Float32(cutlass.select_(col <= lim1, b, cutlass.Float32(NEG_INF))) + v0.append(a) + v1.append(b) + m0 = cute.arch.fmax(m0, a) + m1 = cute.arch.fmax(m1, b) + for offset in (1, 2): + m0 = cute.arch.fmax(m0, cute.arch.shuffle_sync_bfly(m0, offset=offset)) + m1 = cute.arch.fmax(m1, cute.arch.shuffle_sync_bfly(m1, offset=offset)) + # A row with nothing visible in this tile: m = -inf, p = 0 (the merge weights it by 0). + base0 = cutlass.Float32( + cutlass.select_(m0 == cutlass.Float32(NEG_INF), cutlass.Float32(0.0), m0) + ) + base1 = cutlass.Float32( + cutlass.select_(m1 == cutlass.Float32(NEG_INF), cutlass.Float32(0.0), m1) + ) + l0 = cutlass.Float32(0.0) + l1 = cutlass.Float32(0.0) + for g in cutlass.range_constexpr(TILE // 8): + pv0 = [] + pv1 = [] + for e in cutlass.range_constexpr(2): + pa = cute.math.exp2(v0[2 * g + e] - base0, fastmath=True) + pb = cute.math.exp2(v1[2 * g + e] - base1, fastmath=True) + l0 = l0 + pa + l1 = l1 + pb + pv0.append(pa) + pv1.append(pb) + # P [row][kv] bf16, K-major 128B swizzle: half g // 8 (= page), chunk g % 8, elements 2 quad, +1. + half = g // 8 + chunk = g % 8 + pair0 = cutlass.Vector.from_elements( + (pv0[0].to(io_dtype), pv0[1].to(io_dtype)), io_dtype + ) + pair1 = cutlass.Vector.from_elements( + (pv1[0].to(io_dtype), pv1[1].to(io_dtype)), io_dtype + ) + off0 = ( + cutlass.Int32(half * P_HALF_ELEMS) + + r0 * cutlass.Int32(CHUNK) + + ((cutlass.Int32(chunk) ^ (r0 % cutlass.Int32(8))) * cutlass.Int32(8)) + + cutlass.Int32(2) * quad + ) + off1 = ( + cutlass.Int32(half * P_HALF_ELEMS) + + r1 * cutlass.Int32(CHUNK) + + ((cutlass.Int32(chunk) ^ (r1 % cutlass.Int32(8))) * cutlass.Int32(8)) + + cutlass.Int32(2) * quad + ) + smem_p.store(pair0, idx=off0, vector_size=2, alignment=4) + smem_p.store(pair1, idx=off1, vector_size=2, alignment=4) + for offset in (1, 2): + l0 = l0 + cute.arch.shuffle_sync_bfly(l0, offset=offset) + l1 = l1 + cute.arch.shuffle_sync_bfly(l1, offset=offset) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + prims.fence_proxy(prims.Proxy.ASYNC_SHARED, space=prims.SharedSpace.shared_cta) + prims.mbarrier_arrive(p_full) + # Fold scales (k > 0; the slot's m is finite then: an earlier tile of a CTA is fully visible). + m_new0 = cute.arch.fmax(m_run0, m0) + m_new1 = cute.arch.fmax(m_run1, m1) + fa0 = cute.math.exp2(m_run0 - m_new0, fastmath=True) + fa1 = cute.math.exp2(m_run1 - m_new1, fastmath=True) + fb0 = cute.math.exp2(m0 - m_new0, fastmath=True) + fb1 = cute.math.exp2(m1 - m_new1, fastmath=True) + first = k == cutlass.Int32(0) + l_new0 = cutlass.Float32(cutlass.select_(first, l0, l_run0 * fa0 + l0 * fb0)) + l_new1 = cutlass.Float32(cutlass.select_(first, l1, l_run1 * fa1 + l1 * fb1)) + # The slot holds O / l (rows with nothing visible: 0): scales of the old slot (ca) and of this tile (cb). + inv0 = cutlass.Float32( + cutlass.select_( + l_new0 > cutlass.Float32(0.0), + cutlass.Float32(1.0) / l_new0, + cutlass.Float32(0.0), + ) + ) + inv1 = cutlass.Float32( + cutlass.select_( + l_new1 > cutlass.Float32(0.0), + cutlass.Float32(1.0) / l_new1, + cutlass.Float32(0.0), + ) + ) + ca0 = cutlass.Float32(cutlass.select_(first, cutlass.Float32(0.0), l_run0 * fa0 * inv0)) + ca1 = cutlass.Float32(cutlass.select_(first, cutlass.Float32(0.0), l_run1 * fa1 * inv1)) + cb0 = cutlass.Float32(cutlass.select_(first, inv0, fb0 * inv0)) + cb1 = cutlass.Float32(cutlass.select_(first, inv1, fb1 * inv1)) + l_run0 = l_new0 + l_run1 = l_new1 + m_run0 = cutlass.Float32(cutlass.select_(first, m0, m_new0)) + m_run1 = cutlass.Float32(cutlass.select_(first, m1, m_new1)) + if k == my_count - cutlass.Int32(1): + if cutlass.const_expr(no_cluster): + # The CTA's final (m, l) of its live rows into its slot of the exchange, then its arrival. + if quad == cutlass.Int32(0): + if live0: + ml_glob.store( + cutlass.Vector.from_elements((m_run0, l_run0), cutlass.Float32), + idx=(my_slot * cutlass.Int32(ROWS) + r0) * cutlass.Int32(2), vector_size=2, + alignment=8, + ) # fmt: skip + if live1: + ml_glob.store( + cutlass.Vector.from_elements((m_run1, l_run1), cutlass.Float32), + idx=(my_slot * cutlass.Int32(ROWS) + r1) * cutlass.Int32(2), vector_size=2, + alignment=8, + ) # fmt: skip + prims.barrier_cta_sync(1, thread_count=EPI_THREADS) + if tx == cutlass.Int32(4 * 32): + _red_add_release(ctrs.data_ptr(ctr_ml).toint(), 1) + sync_target.store( + (count_ml & cutlass.Int32(-16)) + cutlass.Int32(16), idx=0 + ) + prims.mbarrier_arrive(ml_full) + else: + # The CTA's final (m, l) of its live rows into slot `cta` of every CTA's mailbox (this one too). + if quad == cutlass.Int32(0): + for dst in cutlass.range_constexpr(CLUSTER): + mb = _mapa_u32(ml_full.data_ptr(), dst) + if live0: + _st_async_v2( + _mapa_u32( + ml_mail.data_ptr((cta * cutlass.Int32(ROWS) + r0) * cutlass.Int32(2)), dst + ), + m_run0, l_run0, mb, + ) # fmt: skip + if live1: + _st_async_v2( + _mapa_u32( + ml_mail.data_ptr((cta * cutlass.Int32(ROWS) + r1) * cutlass.Int32(2)), dst + ), + m_run1, l_run1, mb, + ) # fmt: skip + base_o0 = ( + my_slot * cutlass.Int32(WS_SLOT_ELEMS) + + r0 * cutlass.Int32(8) + + cutlass.Int32(2) * quad + ) + base_o1 = ( + my_slot * cutlass.Int32(WS_SLOT_ELEMS) + + r1 * cutlass.Int32(8) + + cutlass.Int32(2) * quad + ) + # O / l -> ws_o[slot, row, :] in fp16 (folded into it for k > 0), 128 columns per load, each N = 256 half + # as soon as its MMAs are done. + for part in cutlass.range_constexpr(LATENT // 128): + if cutlass.const_expr(part % 2 == 0): + while not cute.arch.mbarrier_try_wait( + o_full.subview(part // 2).data_ptr(), phase + ): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + o = prims.tcgen05_ld( + "16x256b", + cutlass.inttoptr(tmem_base + cutlass.Int32(part * 128), 6, cutlass.Float32), + num=16, + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + if first: + for g in cutlass.range_constexpr(16): + col = (part * 16 + g) * WS_GROUP_ELEMS + if real0: + ws_o.store( + cutlass.Vector.from_elements( + ((cutlass.Float32(o[4 * g]) * cb0).to(ws_dtype), + (cutlass.Float32(o[4 * g + 1]) * cb0).to(ws_dtype)), + ws_dtype, + ), + idx=base_o0 + cutlass.Int32(col), vector_size=2, alignment=4, + ) # fmt: skip + if real1: + ws_o.store( + cutlass.Vector.from_elements( + ((cutlass.Float32(o[4 * g + 2]) * cb1).to(ws_dtype), + (cutlass.Float32(o[4 * g + 3]) * cb1).to(ws_dtype)), + ws_dtype, + ), + idx=base_o1 + cutlass.Int32(col), vector_size=2, alignment=4, + ) # fmt: skip + elif real0: + # This thread wrote these slot entries itself (program order): read, rescale, add, write back. + # Rows past ROWS are never written, so they are not read either; a thread's two rows are both + # real or both past ROWS (ROWS % 16 == 0), so warp 7 skips this whole branch. + old0 = [] + old1 = [] + for g in cutlass.range_constexpr(16): + col = (part * 16 + g) * WS_GROUP_ELEMS + old0.append( + ws_o.load(idx=base_o0 + cutlass.Int32(col), vector_size=2, alignment=4) + ) + old1.append( + ws_o.load(idx=base_o1 + cutlass.Int32(col), vector_size=2, alignment=4) + ) + for g in cutlass.range_constexpr(16): + col = (part * 16 + g) * WS_GROUP_ELEMS + ws_o.store( + cutlass.Vector.from_elements( + ( + (cutlass.Float32(old0[g][0]) * ca0 + + cutlass.Float32(o[4 * g]) * cb0).to(ws_dtype), + (cutlass.Float32(old0[g][1]) * ca0 + + cutlass.Float32(o[4 * g + 1]) * cb0).to(ws_dtype), + ), + ws_dtype, + ), + idx=base_o0 + cutlass.Int32(col), vector_size=2, alignment=4, + ) # fmt: skip + ws_o.store( + cutlass.Vector.from_elements( + ( + (cutlass.Float32(old1[g][0]) * ca1 + + cutlass.Float32(o[4 * g + 2]) * cb1).to(ws_dtype), + (cutlass.Float32(old1[g][1]) * ca1 + + cutlass.Float32(o[4 * g + 3]) * cb1).to(ws_dtype), + ), + ws_dtype, + ), + idx=base_o1 + cutlass.Int32(col), vector_size=2, alignment=4, + ) # fmt: skip + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + prims.mbarrier_arrive(o_drained) + if cutlass.const_expr(no_cluster): + if my_count == cutlass.Int32(0): + if tx == cutlass.Int32(4 * 32): + prims.griddepcontrol(prims.GridDepAction.WAIT) + count_ml0 = ctrs.load(idx=ctr_ml, is_volatile=True) + _red_add_release(ctrs.data_ptr(ctr_ml).toint(), 1) + sync_target.store((count_ml0 & cutlass.Int32(-16)) + cutlass.Int32(16), idx=0) + prims.mbarrier_arrive(ml_full) + else: + # ===================================================================== + # Warps 1 and 3: merge weights from the (m, l) mailbox while O runs. + # ===================================================================== + t1 = (warp_id // cutlass.Int32(2)) * cutlass.Int32(32) + tx % cutlass.Int32(32) + if cutlass.const_expr(no_cluster): + # This CTA's own arrival; the other CTAs' (m, l) are acquired through the counter. + while not cute.arch.mbarrier_try_wait(ml_full.data_ptr(), 0): + pass + target = sync_target.load(idx=0) + while _ld_acquire(ctrs.data_ptr(ctr_ml).toint()) - target < cutlass.Int32(0): + pass + else: + # Completed by the cluster's st.async: acquire at cluster scope. + while not _try_wait_cluster(ml_full.data_ptr(), 0): + pass + # A quad of lanes per row (16 rows per pass, 3 passes), lane j of the quad over slots j, j + 4, j + 8, j + 12: + # M = max_s m_s, weight_s = exp2(m_s - M) l_s / sum_s exp2(m_s - M) l_s (the slots hold O / l), over m_s. + sub = t1 % cutlass.Int32(4) + for ps in cutlass.range_constexpr(ROWS // 16): + row = cutlass.Int32(ps * 16) + t1 // cutlass.Int32(4) + live = (row < cutlass.Int32(ROWS)) & (row % cutlass.Int32(MAX_TOKENS) < tokens) + row_c = cutlass.Int32(cutlass.select_(live, row, cutlass.Int32(0))) + ms = [] + ls = [] + for i in cutlass.range_constexpr(CLUSTER // 4): + sl = sub + cutlass.Int32(4 * i) + ok = sl < n_valid + e = ( + cutlass.Int32(cutlass.select_(ok, sl, cutlass.Int32(0))) * cutlass.Int32(ROWS) + + row_c + ) * cutlass.Int32(2) + # Only a live lane's own slot entries are read as m: the m words are overwritten with the + # weights below by the lane that owns them. A lane past n_valid, or of a dead row, reads (and + # drops) an l word instead, which nothing writes here. + e_m = cutlass.Int32(cutlass.select_(ok & live, e, e + cutlass.Int32(1))) + if cutlass.const_expr(no_cluster): + eg = ws_slot0 * cutlass.Int32(ROWS * 2) + e + ms.append( + cutlass.Float32( + cutlass.select_( + ok, + ml_glob.load(idx=ws_slot0 * cutlass.Int32(ROWS * 2) + e_m), + cutlass.Float32(NEG_INF), + ) + ) + ) + ls.append( + cutlass.Float32( + cutlass.select_( + ok, ml_glob.load(idx=eg + cutlass.Int32(1)), cutlass.Float32(0.0) + ) + ) + ) + else: + ms.append( + cutlass.Float32( + cutlass.select_(ok, ml_mail.load(idx=e_m), cutlass.Float32(NEG_INF)) + ) + ) + ls.append( + cutlass.Float32( + cutlass.select_( + ok, ml_mail.load(idx=e + cutlass.Int32(1)), cutlass.Float32(0.0) + ) + ) + ) + mx = ms[0] + for i in cutlass.range_constexpr(1, CLUSTER // 4): + mx = cute.arch.fmax(mx, ms[i]) + for offset in (1, 2): + mx = cute.arch.fmax(mx, cute.arch.shuffle_sync_bfly(mx, offset=offset)) + bs = [cute.math.exp2(ms[i] - mx, fastmath=True) for i in range(CLUSTER // 4)] + den = bs[0] * ls[0] + for i in cutlass.range_constexpr(1, CLUSTER // 4): + den = den + bs[i] * ls[i] + for offset in (1, 2): + den = den + cute.arch.shuffle_sync_bfly(den, offset=offset) + inv = cutlass.Float32(1.0) / den + if live: + for i in cutlass.range_constexpr(CLUSTER // 4): + sl = sub + cutlass.Int32(4 * i) + if sl < n_valid: + ml_mail.store( + bs[i] * ls[i] * inv, + idx=(sl * cutlass.Int32(ROWS) + row) * cutlass.Int32(2), + ) + + # ========================================================================= + # Every tile's partial is in the workspace (release / acquire over the + # cluster, all threads; no_cluster: the second counter); CTA c merges + # latent columns [32 c, 32 c + 32). + # ========================================================================= + if cutlass.const_expr(no_cluster): + if cutlass.const_expr(fuse_vb): + merging = VB_CTAS + else: + merging = CLUSTER + prims.barrier_cta_sync(0) + if tx == cutlass.Int32(0): + _red_add_release(ctrs.data_ptr(ctr_o).toint(), 1) + if cta < cutlass.Int32(merging): + target_o = (sync_target.load(idx=1) & cutlass.Int32(-16)) + cutlass.Int32(16) + while _ld_acquire(ctrs.data_ptr(ctr_o).toint()) - target_o < cutlass.Int32(0): + pass + prims.barrier_cta_sync(0) + else: + prims.barrier_cluster_arrive() + prims.barrier_cluster_wait() + if cutlass.const_expr(fuse_vb): + if cta < cutlass.Int32(VB_CTAS): + # Rows 8 h + t of tokens t = 4 th + j (j < 4; adjacent rows), all 512 columns: thread -> row j = tx % 4, + # columns 8 (tx / 4) .. + 8, so a warp reads 8 column groups x 4 adjacent 16-byte rows; the slots in two + # batches of 8 (16 loads of 16 bytes in flight per batch). + h_loc = cta // cutlass.Int32(2) + j_row = tx % cutlass.Int32(VB_TOKENS) + t_tok = (cta % cutlass.Int32(2)) * cutlass.Int32(VB_TOKENS) + j_row + live = t_tok < tokens + # A token t >= T reads its own row too (weight 0); its column of the v_b product is never stored. + row = h_loc * cutlass.Int32(MAX_TOKENS) + t_tok + c8 = (tx // cutlass.Int32(VB_TOKENS)) * cutlass.Int32(8) + acc = [cutlass.Float32(0.0)] * 8 + for half in cutlass.range_constexpr(2): + vals = [] + wts = [] + for j in cutlass.range_constexpr(MERGE_HALF): + sl = cutlass.Int32(half * MERGE_HALF + j) + ok = sl < n_valid + sl_c = cutlass.Int32(cutlass.select_(ok, sl, cutlass.Int32(0))) + base = ( + (ws_slot0 + sl_c) * cutlass.Int32(WS_SLOT_ELEMS) + + (c8 // cutlass.Int32(8)) * cutlass.Int32(WS_GROUP_ELEMS) + + row * cutlass.Int32(8) + ) + vals.append(ws_o.load(idx=base, vector_size=8, alignment=16)) + wts.append( + cutlass.Float32( + cutlass.select_( + ok & live, + ml_mail.load( + idx=(sl_c * cutlass.Int32(ROWS) + row) * cutlass.Int32(2) + ), + cutlass.Float32(0.0), + ) + ) + ) + for j in cutlass.range_constexpr(MERGE_HALF): + for e in cutlass.range_constexpr(8): + acc[e] = acc[e] + wts[j] * cutlass.Float32(vals[j][e]) + # bf16 o (the unfused path's attention output) -> the v_b B tile in the dead P buffer (token j, columns + # c8 .. c8 + 8), K-major 128B swizzle: chunk c8 / 64, 16-byte vector ((c8 % 64) / 8) ^ j. + smem_p.store( + cutlass.Vector.from_elements(tuple(v.to(io_dtype) for v in acc), io_dtype), + idx=(c8 // cutlass.Int32(CHUNK)) * cutlass.Int32(8 * CHUNK) + j_row * cutlass.Int32(CHUNK) + + (((c8 % cutlass.Int32(CHUNK)) // cutlass.Int32(8)) ^ j_row) * cutlass.Int32(8), + vector_size=8, + alignment=16, + ) # fmt: skip + prims.fence_proxy(prims.Proxy.ASYNC_SHARED, space=prims.SharedSpace.shared_cta) + prims.barrier_cta_sync(0) + if cta < cutlass.Int32(VB_CTAS): + if warp_id == 2: + # y^T [128 v, 8 tokens] = W_vb[h] (M 128, K 512 from the KV buffer) x o^T (the B tile), fp32 in TMEM. + idesc_vb = prims.Tcgen05InstrDesc.build( + c_dtype=cutlass.Float32, a_dtype=io_dtype, b_dtype=io_dtype, n_dim=8, m_dim=128 + ) + swz = prims.Tcgen05SmemSwizzle.SWIZZLE_128B + desc_w = prims.Tcgen05SmemDesc.build( + start_address=smem_kv, + leading_byte_offset=LEADING, + stride_byte_offset=SBO, + layout=swz, + ) + desc_o = prims.Tcgen05SmemDesc.build( + start_address=smem_p, + leading_byte_offset=LEADING, + stride_byte_offset=SBO, + layout=swz, + ) + while not cute.arch.mbarrier_try_wait(vb_full.data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + for jc in cutlass.range_constexpr(LAT_CHUNKS): + for kk in cutlass.range_constexpr(CHUNK // MMA_K): + if prims.elect_sync(): + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, prims.CTAGroup.CTA_1, + cutlass.inttoptr(tmem_base, 6, cutlass.Int32), + desc_w + (jc * ((V_DIM * CHUNK * ELEM_BYTES) >> 4) + kk * STEP_K), + desc_o + (jc * ((8 * CHUNK * ELEM_BYTES) >> 4) + kk * STEP_K), + idesc_vb, not (jc == 0 and kk == 0), + ) # fmt: skip + if prims.elect_sync(): + prims.tcgen05_commit(vb_done) + elif warp_id >= 4: + # y[t, h, v]: lane v = 32 (warp - 4) + lane of the accumulator, tokens in its 8 columns. + v_idx = (warp_id - cutlass.Int32(4)) * cutlass.Int32(32) + tx % cutlass.Int32(32) + h_glob = hg * cutlass.Int32(HEADS) + cta // cutlass.Int32(2) + # The gate of this thread's tokens, loaded before the v_b wait so its latency hides under the MMA (a + # token past T reads the request's first row and is not used). + gate_pre = [] + if cutlass.const_expr(apply_gate): + for j in cutlass.range_constexpr(VB_TOKENS): + t_pre = (cta % cutlass.Int32(2)) * cutlass.Int32(VB_TOKENS) + cutlass.Int32( + j + ) + t_row = tok0 + cutlass.Int32( + cutlass.select_(t_pre < tokens, t_pre, cutlass.Int32(0)) + ) + gate_pre.append( + cutlass.Float32( + gate.load( + idx=t_row * gate_ld + + gate_col0 + + h_glob * cutlass.Int32(V_DIM) + + v_idx + ) + ) + ) + while not cute.arch.mbarrier_try_wait(vb_done.data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + yv = prims.tcgen05_ld( + "32x32b", cutlass.inttoptr(tmem_base, 6, cutlass.Float32), num=8 + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + for j in cutlass.range_constexpr(VB_TOKENS): + t_out = (cta % cutlass.Int32(2)) * cutlass.Int32(VB_TOKENS) + cutlass.Int32(j) + if t_out < tokens: + y_out = cutlass.Float32(yv[j]).to(io_dtype) + if cutlass.const_expr(apply_gate): + y_out = (cutlass.Float32(y_out) * gate_pre[j]).to(io_dtype) + out.store( + y_out, + idx=((tok0 + t_out) * cutlass.Int32(total_heads) + h_glob) + * cutlass.Int32(V_DIM) + + v_idx, + ) + # Every TMEM reader of this CTA has waited for its loads. + prims.barrier_cta_sync(0) + if warp_id == 4: + prims.tcgen05_dealloc(cutlass.inttoptr(tmem_base, 6, cutlass.Int32), TMEM_COLS) + else: + if warp_id == 4: + prims.tcgen05_dealloc(cutlass.inttoptr(tmem_base, 6, cutlass.Int32), TMEM_COLS) + # 4 groups of 8 columns x 48 rows = 192 items (item = 48 group + row, so consecutive pairs read consecutive + # 16-byte rows of a group); thread pair (2 p, 2 p + 1) takes items p and p + 128 (< 192), the even thread over + # slots [0, 8), the odd one over [8, 16): all 16 of a thread's workspace loads (8 fp16 each) are in flight + # together, then the pair's halves are added by one shuffle. Slots >= n_valid load slot 0 with weight 0, rows + # of tokens t >= T their own row with weight 0 (never stored). + hs = tx % cutlass.Int32(2) + pair = tx // cutlass.Int32(2) + vals = [] + wts = [] + for i in cutlass.range_constexpr(2): + item = pair + cutlass.Int32(128 * i) + item_c = cutlass.Int32(cutlass.select_(item < cutlass.Int32(ROWS * 4), item, pair)) + row = item_c % cutlass.Int32(ROWS) + live = (row % cutlass.Int32(MAX_TOKENS) < tokens) & (item < cutlass.Int32(ROWS * 4)) + grp = cta * cutlass.Int32(LAT_PER_CTA // 8) + item_c // cutlass.Int32(ROWS) + for j in cutlass.range_constexpr(MERGE_HALF): + sl = hs * cutlass.Int32(MERGE_HALF) + cutlass.Int32(j) + ok = sl < n_valid + sl_c = cutlass.Int32(cutlass.select_(ok, sl, cutlass.Int32(0))) + vals.append( + ws_o.load( + idx=(ws_slot0 + sl_c) * cutlass.Int32(WS_SLOT_ELEMS) + + grp * cutlass.Int32(WS_GROUP_ELEMS) + + row * cutlass.Int32(8), + vector_size=8, + alignment=16, + ) + ) + wts.append( + cutlass.Float32( + cutlass.select_( + ok & live, + ml_mail.load(idx=(sl_c * cutlass.Int32(ROWS) + row) * cutlass.Int32(2)), + cutlass.Float32(0.0), + ) + ) + ) + for i in cutlass.range_constexpr(2): + item = pair + cutlass.Int32(128 * i) + item_c = cutlass.Int32(cutlass.select_(item < cutlass.Int32(ROWS * 4), item, pair)) + row = item_c % cutlass.Int32(ROWS) + grp = cta * cutlass.Int32(LAT_PER_CTA // 8) + item_c // cutlass.Int32(ROWS) + acc = [] + for e in cutlass.range_constexpr(8): + a = cutlass.Float32(0.0) + for j in cutlass.range_constexpr(MERGE_HALF): + a = a + wts[i * MERGE_HALF + j] * cutlass.Float32(vals[i * MERGE_HALF + j][e]) + acc.append(a + cute.arch.shuffle_sync_bfly(a, offset=1)) + out_row = ( + (tok0 + row % cutlass.Int32(MAX_TOKENS)) * cutlass.Int32(total_heads) + + hg * cutlass.Int32(HEADS) + + row // cutlass.Int32(MAX_TOKENS) + ) + if hs == cutlass.Int32(0): + if (row % cutlass.Int32(MAX_TOKENS) < tokens) & (item < cutlass.Int32(ROWS * 4)): + out.store( + cutlass.Vector.from_elements(tuple(v.to(io_dtype) for v in acc), io_dtype), + idx=out_row * cutlass.Int32(LATENT) + grp * cutlass.Int32(8), + vector_size=8, + alignment=16, + ) + + +def _q_map(q, num_tokens, total_heads): + """fused_q [M, heads, 576] as (64-column chunk, token, head, chunk index): one call at (token i T, head 6 g, chunk + 3 j) lands the group's [3][6 heads x 8 tokens][64] (row 8 h + t holds token i T + t; tokens >= M zero).""" + return cuda.create_tensor_map_tiled( + global_address=q.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[CHUNK, num_tokens, total_heads, QK_CHUNKS], + global_strides=[ + (total_heads * QK * ELEM_BYTES) // 16, + (QK * ELEM_BYTES) // 16, + (CHUNK * ELEM_BYTES) // 16, + ], + box_dims=[CHUNK, MAX_TOKENS, HEADS, Q_GROUP_CHUNKS], + swizzle=cuda.TensorMapSwizzle.s128b, + ) + + +def _vb_map(w_vb, rows): + """v_b_proj [heads * 128, 512] as (64-column chunk, row, chunk index): one call lands a head's [8][128][64].""" + return cuda.create_tensor_map_tiled( + global_address=w_vb.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[CHUNK, rows, LAT_CHUNKS], + global_strides=[(LATENT * ELEM_BYTES) // 16, (CHUNK * ELEM_BYTES) // 16], + box_dims=[CHUNK, V_DIM, LAT_CHUNKS], + swizzle=cuda.TensorMapSwizzle.s128b, + ) + + +def _kv_map(pool, total_rows, row_stride): + """The latent pool as [rows, 576] (row stride `row_stride` elements): one call lands a page's 64-column chunk.""" + return cuda.create_tensor_map_tiled( + global_address=pool.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[CHUNK, total_rows, QK_CHUNKS], + global_strides=[(row_stride * ELEM_BYTES) // 16, (CHUNK * ELEM_BYTES) // 16], + box_dims=[CHUNK, PAGE, 1], + swizzle=cuda.TensorMapSwizzle.s128b, + ) + + +@cute.jit +def k3_mla_attn( + q: cute.Tensor, # [M * heads * 576] bf16, M = R T + pool: cute.Tensor, # the latent pool, flat bf16: row i of page p at (p * 64 + i) * row_stride + page_table: cute.Tensor, # int32, request i's pages at [i * pt_stride, ...) + seq_len: cute.Tensor, # int32 [R] + ws_o: cute.Tensor, # fp16 [MAX_REQUESTS * heads / 6 * 16 * WS_SLOT_ELEMS] + out: cute.Tensor, # bf16 [M * heads * 512] (fuse_vb: [M * heads * 128]) + w_vb: cute.Tensor, # fuse_vb: v_b_proj [heads * 128 * 512] bf16 (else any bf16 tensor, unused) + gate: cute.Tensor, # apply_gate: bf16 [M * gate_ld] (else any bf16 tensor, unused) + tokens: cutlass.Int32, # T <= 8 + num_requests: cutlass.Int32, # R <= MAX_REQUESTS + pt_stride: cutlass.Int32, + scale_log2: cutlass.Float32, + total_rows: cutlass.Int32, + page_offset: cutlass.Int32, + gate_col0: cutlass.Int32, + gate_ld: cutlass.Int32, + row_stride: cutlass.Constexpr[int], + total_heads: cutlass.Constexpr[int], + fuse_vb: cutlass.Constexpr[bool], + apply_gate: cutlass.Constexpr[bool], + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, + no_cluster: cutlass.Constexpr[bool] = False, +) -> None: + tma_q = _q_map(q, tokens * num_requests, total_heads) + tma_kv = _kv_map(pool, total_rows, row_stride) + tma_vb = _vb_map(w_vb, total_heads * V_DIM) if cutlass.const_expr(fuse_vb) else tma_kv + k3_mla_attn_kernel( + tma_q, + tma_kv, + tma_vb, + page_table, + seq_len, + ws_o, + out, + gate, + tokens, + pt_stride, + scale_log2, + page_offset, + gate_col0, + gate_ld, + total_heads, + fuse_vb, + apply_gate, + no_cluster, + ).launch( # fmt: skip + grid=(CLUSTER * (total_heads // HEADS), num_requests, 1), + block=(THREADS, 1, 1), + cluster=None if no_cluster else (CLUSTER, 1, 1), + stream=stream, + use_pdl=use_pdl, + ) + + +def softmax_scale_log2( + qk_nope_head_dim: int, qk_rope_head_dim: int, q_scaling: float = 1.0 +) -> float: + return LOG2E / (math.sqrt(qk_nope_head_dim + qk_rope_head_dim) * q_scaling) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/k3_mla_q_kernel.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/k3_mla_q_kernel.py new file mode 100644 index 000000000000..e39ea633654d --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/k3_mla_q_kernel.py @@ -0,0 +1,1099 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +# ============================================================================= +# Kimi K3 MLA decode query path -- CTM (prims/cute) kernel, M <= 64 tokens, any number of heads per rank (6 at TP16) +# ============================================================================= +# +# q_n[t] = bf16(q_a[t] * rsqrt(mean(q_a[t]^2) + eps) * w_qa) (q_a_layernorm, 1536) +# q[t, h] = bf16(q_n[t] @ W_qb[192 h : 192 h + 192]^T) (q_b_proj: 128 nope | 64 pe) +# fused_q[t, h] = [ bf16(q[t, h, :128] @ W_kb[h]^T) (512) | q[t, h, 128:] (64) ] (k_b absorb; K3 is NoPE) +# +# The grid's y index is a chunk of N tokens (rows N c .. N c + N - 1, the last chunk possibly shorter; N, the MMA's +# N, is the compile-time `mma_n`: 8, or 32 for the wider steps, whose chunks of 8 would need more co-resident clusters +# than fit): each chunk is the same N-token problem below, on its own rows, the weights read again per chunk (from L2 +# after the first). Every token's sums keep the same order at either N, so a 32-token chunk's rows are bit-identical +# to the four 8-token chunks of the same rows. Everything below describes one chunk of 8 tokens; at N 32 every +# per-token step runs for four groups of 8 tokens (token t of group g is row 8 g + t of the chunk). +# +# One cluster of 6 CTAs per head. Rank r holds (TMA'd EVICT_FIRST before griddepcontrol.wait) the head's W_qb +# nope rows (128) and pe rows (64) over k-tiles {2 r, 2 r + 1}, and for r < 4 the k_b rows W_kb[h][128 r:128 r+128]. +# After the wait the epilogue warps normalize the q_a rows (the RMS over all 1536 columns) and write this rank's 256 +# columns, bf16, into a resident 128B-swizzled B tile (the layout TMA would produce), fence it to the async proxy and +# release the MMA warp (fuse_rmsnorm_qkv_rope's resident-B prolog). cluster_rms: each rank reads only its 256 columns +# and st.async's its per-token partial sums of squares to the 6 ranks, which add them in rank order. The MMA warp +# runs the nope (M 128) and pe (M 64) partial products. Split-K reduce over the cluster: nope rows [32 w, 32 w + 32) +# (TMEM lanes of epilogue warp w) belong to rank w, pe rows [32 j, 32 j + 32) (lanes 0-15 of warps 2 j, 2 j + 1: an +# M = 64 accumulator puts row 16 w + l in lane 32 w + l) to rank 4 + j; non-owners store their fp32 partials into +# slot [rank] of the owner's mailbox through DSMEM and arrive (release, cluster scope); owners add the 6 partials in +# rank order and round once. Every wait on a barrier other CTAs complete acquires at cluster scope. +# pe owners store fused_q[..., 512:]. Each nope owner writes its bf16 32 x 8 slice, as 16-byte chunks, into the k_b B +# tile of ranks 0-3 and arrives on their barrier; those ranks fence the DSMEM writes to the async proxy, run the k_b +# MMA (M 128, N 8, K 128) and store fused_q[t, h, 128 r : 128 r + 128]. +# single_hop (8-token chunks only): instead, every rank st.async's its whole nope partial (128 rows x 8 tokens fp32) +# into the mailbox of each k_b rank (completing that rank's barrier by bytes); a k_b rank adds the 6 partials in rank +# order itself (the same fp32 sums, so the same bits), writes its B tile and runs the k_b MMA: one DSMEM hop instead +# of two. +# +# kv_mode (the KV half, warp 1 of CTA j < the chunk's tokens: token t = N c + j): token t's latent row +# ag[t, 1536:2048] RMS-normalized (flashinfer's order, like q_a) and its rope columns ag[t, 2048:2112] (K3 is NoPE: +# copied) form the 576-column cache row, stored (1) into the paged latent pool: the M tokens are R requests of T each +# (request-major), token t is token u = t - i T of request i = t / T, at position p = L_i - T + u (row +# (page_table[i][p / 64] + page_offset) * 64 + p % 64, 64-bit element index; not stored when p < 0), or (2) into +# kv_out [M, 576] (check builds). Page table and lengths are host-filled: read before the grid wait. The attention +# after this kernel reads the new rows after its own grid wait. +# +# Warps: 0 weight TMA (+ early dependent trigger), 1 KV half (kv_mode), 2 TMEM allocation + MMA, 4-7 norm prolog / +# reduce / epilogue, 3 idle. +# ============================================================================= +"""CTM kernel for Kimi K3 MLA's decode query path: q_a RMSNorm, q_b projection, k_b absorption -> fused_q.""" + +from __future__ import annotations + +import cuda.bindings.driver as cuda_driver +import cutlass +import cutlass.cute as cute +import cutlass.experimental.cuda as cuda +from cutlass._mlir.dialects import llvm as _llvm +from cutlass._mlir.extras import types as _T +from cutlass.cutlass_dsl import dsl_user_op +from cutlass.experimental import primitives as prims + +HEADS = 6 # per rank at TP16 (the kernel takes the rank's head count as a compile-time argument) +NOPE = 128 +PE = 64 +QK = NOPE + PE # 192 q_b rows per head +LATENT = 512 # kv_lora_rank +Q_LORA = 1536 # q_lora_rank = q_b's K +FUSED = LATENT + PE # 576 +PAGE = 64 # latent cache rows per page +CLUSTER = 6 # CTAs per head +K_TILES = Q_LORA // 128 # 12 +MY_TILES = K_TILES // CLUSTER # 2 q_b k-tiles per rank +KB_RANKS = LATENT // 128 # 4 ranks hold a 128-row k_b tile +MMA_N = 8 # tokens per chunk (the MMA's N), and the group of tokens every per-token step works on +WIDE_CHUNKS = ( + 16, + 32, +) # the other chunk sizes (cluster_rms, two-hop): 64 tokens are 12 clusters in chunks of 32 +CTA_K = 128 +MMA_K = 16 +TMA_K_BOX = 64 +TMA_COPY_ITERS = CTA_K // TMA_K_BOX +K_BLOCKS_PER_HALF = TMA_K_BOX // MMA_K +THREADS = 256 +EPI_THREADS = 128 +TMEM_NOPE = 0 # the nope accumulator's TMEM column; pe and k_b follow at N and 2 N (tmem_cols) +ELEM_BYTES = 2 +VEC = 8 +EVICT_FIRST = 0x12F0000000000000 + +LEADING = 16 +STRIDE = 8 * TMA_K_BOX * ELEM_BYTES # 1024 B between 8-row swizzle atoms +STEP = (MMA_K * ELEM_BYTES) >> 4 +HALF_128 = ( + 128 * TMA_K_BOX * ELEM_BYTES +) >> 4 # one 64-column half of a 128-row tile, 16-byte units +HALF_64 = (64 * TMA_K_BOX * ELEM_BYTES) >> 4 +TILE_128 = 2 * HALF_128 +TILE_64 = 2 * HALF_64 +NOPE_ELEMS = 128 * CTA_K # per k-tile +PE_ELEMS = 64 * CTA_K + +KB_MAIL_BYTES = ( + CLUSTER * 128 * MMA_N * 4 +) # the k_b rank's single-hop mailbox: 6 partials of 128 rows x 8 tokens + +io_dtype = cutlass.BFloat16 + + +def tmem_cols(n: int) -> int: + """TMEM columns of an n-token chunk: the nope, pe and k_b accumulators (n columns each), a power of 2 >= 32.""" + cols = 32 + while cols < 3 * n: + cols *= 2 + return cols + + +@dsl_user_op +def _mapa_u32(smem_ptr, peer, *, loc=None, ip=None): + """The shared::cluster address of this CTA's shared-memory location in cluster CTA ``peer``.""" + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [smem_ptr.toint(loc=loc, ip=ip).ir_value(), cutlass.Int32(peer).ir_value(loc=loc, ip=ip)], + "mapa.shared::cluster.u32 $0, $1, $2;", "=r,r,r", has_side_effects=False, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _st_async_v4(dst, a, b, c, d, mbar, *, loc=None, ip=None): + """st.async of four fp32 to a shared::cluster address, completing ``mbar`` (shared::cluster) by 16 bytes.""" + _llvm.inline_asm( + None, + [cutlass.Int32(dst).ir_value(loc=loc, ip=ip), cutlass.Float32(a).ir_value(loc=loc, ip=ip), + cutlass.Float32(b).ir_value(loc=loc, ip=ip), cutlass.Float32(c).ir_value(loc=loc, ip=ip), + cutlass.Float32(d).ir_value(loc=loc, ip=ip), cutlass.Int32(mbar).ir_value(loc=loc, ip=ip)], + "st.async.shared::cluster.mbarrier::complete_tx::bytes.v4.f32 [$0], {$1, $2, $3, $4}, [$5];", + "r,f,f,f,f,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _st_async_f32(dst, value, mbar, *, loc=None, ip=None): + """st.async of one fp32 to a shared::cluster address, completing ``mbar`` (shared::cluster) by 4 bytes.""" + _llvm.inline_asm( + None, + [cutlass.Int32(dst).ir_value(loc=loc, ip=ip), cutlass.Float32(value).ir_value(loc=loc, ip=ip), + cutlass.Int32(mbar).ir_value(loc=loc, ip=ip)], + "st.async.shared::cluster.mbarrier::complete_tx::bytes.f32 [$0], $1, [$2];", "r,f,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _st_async_v4_b32(dst, w0, w1, w2, w3, mbar, *, loc=None, ip=None): + """st.async of 16 bytes (four 32-bit words) to a shared::cluster address, completing ``mbar`` by 16 bytes.""" + _llvm.inline_asm( + None, + [cutlass.Int32(dst).ir_value(loc=loc, ip=ip), cutlass.Int32(w0).ir_value(loc=loc, ip=ip), + cutlass.Int32(w1).ir_value(loc=loc, ip=ip), cutlass.Int32(w2).ir_value(loc=loc, ip=ip), + cutlass.Int32(w3).ir_value(loc=loc, ip=ip), cutlass.Int32(mbar).ir_value(loc=loc, ip=ip)], + "st.async.shared::cluster.mbarrier::complete_tx::bytes.v4.b32 [$0], {$1, $2, $3, $4}, [$5];", + "r,r,r,r,r,r", has_side_effects=True, is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _try_wait_cluster(mbar, parity, *, loc=None, ip=None): + """mbarrier.try_wait.parity.acquire.cluster on a barrier of this CTA (``mbar``, its shared-memory pointer): whether + phase ``parity`` has completed, acquiring at cluster scope. The barrier is completed by other CTAs' st.async, whose + complete_tx releases at cluster scope.""" + done = cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [mbar.toint(loc=loc, ip=ip).ir_value(), cutlass.Int32(parity).ir_value(loc=loc, ip=ip)], + "{\n\t.reg .pred p;\n\tmbarrier.try_wait.parity.acquire.cluster.shared::cta.b64 p, [$1], $2;\n\t" + "selp.u32 $0, 1, 0, p;\n\t}", "=r,r,r,~{memory}", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + return done != cutlass.Int32(0) + + +def _swz(t, k, n=MMA_N): + """Element index of the 8-column vector holding (row t < n, column k < 128) in an n x 128 bf16 K-major tile, + 128B swizzle (two 64-column halves of 8-row swizzle atoms); add k % 8 for the element itself.""" + half = k // cutlass.Int32(TMA_K_BOX) + chunk = (k % cutlass.Int32(TMA_K_BOX)) // cutlass.Int32(VEC) + phase = t if n == MMA_N else t % cutlass.Int32(8) + return ( + half * cutlass.Int32(n * TMA_K_BOX) + + t * cutlass.Int32(TMA_K_BOX) + + (chunk ^ phase) * cutlass.Int32(VEC) + ) + + +def _plus(base, offset: int): + """base + offset; a zero offset traces nothing, so 8-token chunks keep exactly their per-token instructions.""" + return base if offset == 0 else base + cutlass.Int32(offset) + + +@cute.kernel +def k3_mla_q_kernel( + tma_nope: cutlass.GridConstant[cuda.TensorMap], # W_qb [1152, 1536], box 128 rows + tma_pe: cutlass.GridConstant[cuda.TensorMap], # W_qb [1152, 1536], box 64 rows + tma_kb: cutlass.GridConstant[cuda.TensorMap], # W_kb [6 * 512, 128], box 128 rows + ag: cutlass.Array, # [M, ag_cols] bf16: q_a in columns [0, 1536) + w_qa: cutlass.Array, # [1536] bf16, q_a_layernorm weight + fused_q: cutlass.Array, # [M, heads, 576] bf16, out + w_kv: cutlass.Array, # [512] bf16, kv_a_layernorm weight (kv_mode) + kv_pool: cutlass.Array, # the latent pool, flat bf16: row i of page p at (p * 64 + i) * row_stride (kv_mode 1) + page_table: cutlass.Array, # int32, request i's pages at [i * pt_stride, ...) (kv_mode 1) + seq_len: cutlass.Array, # int32 [R], request i's KV length including its T tokens of this step (kv_mode 1) + kv_out: cutlass.Array, # [M, 576] bf16 (kv_mode 2) + total_tokens: cutlass.Int32, # M = R T + tokens: cutlass.Int32, # T (kv_mode 1) + pt_stride: cutlass.Int32, # page-table elements between consecutive requests' rows (kv_mode 1) + eps: cutlass.Float32, + kv_eps: cutlass.Float32, + page_offset: cutlass.Int32, # added to every page-table entry (the layer's slot in an interleaved pool) + ag_cols: cutlass.Constexpr[int], + total_heads: cutlass.Constexpr[int], + trigger_early: cutlass.Constexpr[bool], + single_hop: cutlass.Constexpr[bool], + kv_mode: cutlass.Constexpr[int], # 0 off, 1 into the pool, 2 into kv_out + row_stride: cutlass.Constexpr[int], + cluster_rms: cutlass.Constexpr[ + bool + ], # the q_a RMS from the cluster's partial sums (each rank reads its 256 columns) + mma_n: cutlass.Constexpr[ + int + ], # tokens per chunk: MMA_N, or one of WIDE_CHUNKS (cluster_rms, two-hop) +): + tx, _, _ = cute.arch.thread_idx() + bx, chunk, _ = cute.arch.block_idx() + warp_id = cute.arch.warp_idx() + rank = cute.arch.block_idx_in_cluster() + head = bx // cutlass.Int32(CLUSTER) + groups = mma_n // MMA_N # per-token steps run for each group of 8 of the chunk's tokens + half_n = ( + mma_n * TMA_K_BOX * ELEM_BYTES + ) >> 4 # one 64-column half of the B tile, 16-byte units + b_elems = mma_n * CTA_K + col_pe, col_kb = ( + mma_n, + 2 * mma_n, + ) # TMEM columns of the pe and k_b accumulators (nope at column 0) + # This chunk's tokens: rows tok0 .. tok0 + num_tokens - 1 of ag and fused_q. + tok0 = chunk * cutlass.Int32(mma_n) + num_tokens = cutlass.Int32( + cutlass.select_( + total_tokens - tok0 < cutlass.Int32(mma_n), total_tokens - tok0, cutlass.Int32(mma_n) + ) + ) + ptr_nope = tma_nope.get_ptr() + ptr_pe = tma_pe.get_ptr() + ptr_kb = tma_kb.get_ptr() + + # Same allocation order in every CTA: mapa() addresses the peers' mailbox, B2 tile and barriers. + smem_wn = cutlass.Array( + io_dtype, MY_TILES * NOPE_ELEMS, space=cutlass.AddressSpace.smem, alignment=1024 + ) + smem_wp = cutlass.Array( + io_dtype, MY_TILES * PE_ELEMS, space=cutlass.AddressSpace.smem, alignment=1024 + ) + smem_wk = cutlass.Array(io_dtype, NOPE_ELEMS, space=cutlass.AddressSpace.smem, alignment=1024) + smem_b = cutlass.Array( + io_dtype, MY_TILES * b_elems, space=cutlass.AddressSpace.smem, alignment=1024 + ) + smem_b2 = cutlass.Array(io_dtype, b_elems, space=cutlass.AddressSpace.smem, alignment=1024) + stage_n = cutlass.Array(io_dtype, mma_n * 32, space=cutlass.AddressSpace.smem, alignment=16) + stage_o = cutlass.Array(io_dtype, mma_n * 128, space=cutlass.AddressSpace.smem, alignment=16) + # [source rank][token][row within the owned 32] fp32 partials (the own rank's slot stays unused). + mailbox = cutlass.Array( + cutlass.Float32, CLUSTER * 32 * mma_n, space=cutlass.AddressSpace.smem, alignment=16 + ) + norm_part = cutlass.Array( + cutlass.Float32, 4 * mma_n, space=cutlass.AddressSpace.smem, alignment=16 + ) + rrms = cutlass.Array(cutlass.Float32, mma_n, space=cutlass.AddressSpace.smem, alignment=16) + # cluster_rms: [source rank][token] partial sums of squares over the source's 256 q_a columns. + norm_mail = cutlass.Array( + cutlass.Float32, CLUSTER * mma_n, space=cutlass.AddressSpace.smem, alignment=16 + ) + norm_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + w_full = cutlass.Array( + cutlass.Int64, 2 * MY_TILES + 1, space=cutlass.AddressSpace.smem, alignment=8 + ) + b_ready = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + acc_done = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + mail_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + qn_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + acc2_done = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + tmem_ptr_i32 = cutlass.Array(cutlass.Int32, 1, space=cutlass.AddressSpace.smem) + # single_hop: [source rank][nope row][token] fp32 partials, completed by bytes; the B tile's local readiness. + # (Not in 32-token builds, which are two-hop: 96 KB.) + kb_mail = None + if cutlass.const_expr(mma_n == MMA_N): + kb_mail = cutlass.Array( + cutlass.Float32, CLUSTER * 128 * MMA_N, space=cutlass.AddressSpace.smem, alignment=16 + ) + kb_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + b2_ready = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + + if warp_id == 0: + prims.prefetch_tensormap(ptr_nope) + prims.prefetch_tensormap(ptr_pe) + prims.prefetch_tensormap(ptr_kb) + if prims.elect_sync(): + for i in cutlass.range_constexpr(2 * MY_TILES + 1): + prims.mbarrier_init(w_full.subview(i), 1) + prims.mbarrier_init(b_ready, EPI_THREADS) + prims.mbarrier_init(acc_done, 1) + # Both two-hop mailboxes complete by bytes (st.async), armed here once. + prims.mbarrier_init(mail_full, 1) + prims.mbarrier_init(qn_full, 1) + prims.mbarrier_init(acc2_done, 1) + prims.mbarrier_init(kb_full, 1) + prims.mbarrier_init(b2_ready, EPI_THREADS) + if cutlass.const_expr(single_hop): + if rank < cutlass.Int32(KB_RANKS): + prims.mbarrier_arrive_expect_tx(kb_full, KB_MAIL_BYTES) + prims.mbarrier_arrive_expect_tx(mail_full, (CLUSTER - 1) * mma_n * 32 * 4) + if cutlass.const_expr(cluster_rms): + prims.mbarrier_init(norm_full, 1) + prims.mbarrier_arrive_expect_tx(norm_full, CLUSTER * mma_n * 4) + if rank < cutlass.Int32(KB_RANKS): + # 4 owners x the chunk's tokens x 32 columns, bf16. + prims.mbarrier_arrive_expect_tx(qn_full, KB_RANKS * mma_n * 32 * ELEM_BYTES) + if warp_id == 2: + prims.tcgen05_alloc(tmem_ptr_i32, tmem_cols(mma_n)) + prims.tcgen05_relinquish_alloc_permit() + prims.fence_mbarrier_init() + # Cluster formation: the peers' shared memory and barriers are addressable from here on. + prims.barrier_cluster_arrive_relaxed() + prims.barrier_cluster_wait() + prims.barrier_cta_sync(0) + tmem_base = tmem_ptr_i32.load() + + if warp_id == 0: + # ===================================================================== + # Weights, ahead of the grid dependency: q_b nope and pe rows over this + # rank's two k-tiles, and (ranks 0-3) the k_b tile. + # ===================================================================== + if prims.elect_sync(): + row_nope = head * cutlass.Int32(QK) + for i in cutlass.range_constexpr(MY_TILES): + k = rank * cutlass.Int32(MY_TILES) + cutlass.Int32(i) + prims.mbarrier_arrive_expect_tx(w_full.subview(i), NOPE_ELEMS * ELEM_BYTES) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_wn.subview(i * NOPE_ELEMS), + ptr_nope, + ( + cutlass.Int32(0), + row_nope, + k * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), + cutlass.Int32(0), + ), + w_full.subview(i), + l2_cache_hint=EVICT_FIRST, + ) + prims.mbarrier_arrive_expect_tx(w_full.subview(MY_TILES + i), PE_ELEMS * ELEM_BYTES) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_wp.subview(i * PE_ELEMS), + ptr_pe, + ( + cutlass.Int32(0), + row_nope + cutlass.Int32(NOPE), + k * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), + cutlass.Int32(0), + ), + w_full.subview(MY_TILES + i), + l2_cache_hint=EVICT_FIRST, + ) + if rank < cutlass.Int32(KB_RANKS): + prims.mbarrier_arrive_expect_tx( + w_full.subview(2 * MY_TILES), NOPE_ELEMS * ELEM_BYTES + ) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_wk, + ptr_kb, + ( + cutlass.Int32(0), + head * cutlass.Int32(LATENT) + rank * cutlass.Int32(128), + cutlass.Int32(0), + cutlass.Int32(0), + cutlass.Int32(0), + ), + w_full.subview(2 * MY_TILES), + l2_cache_hint=EVICT_FIRST, + ) + if cutlass.const_expr(trigger_early): + # Dependents may launch now; they wait for this whole grid before reading fused_q. + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + elif warp_id == 1: + if cutlass.const_expr(kv_mode != 0): + # ===================================================================== + # KV half: CTA j < the chunk's tokens, token tok0 + j's cache row (see + # kv_mode above). + # ===================================================================== + if bx < num_tokens: + tok = tok0 + bx + lane = tx % 32 + # Lane l: latent vectors l and l + 32 (8 columns each) and, for l < 8, rope vector l. + wk = [] + for j in cutlass.range_constexpr(2): + wk.append( + w_kv.load( + idx=(lane + cutlass.Int32(32 * j)) * cutlass.Int32(VEC), + vector_size=VEC, + alignment=16, + ) + ) + dst = cutlass.Int64(tok) * cutlass.Int64(FUSED) + keep = cutlass.Boolean(True) + if cutlass.const_expr(kv_mode == 1): + req = tok // tokens + pos = seq_len.load(idx=req) - tokens + (tok - req * tokens) + keep = pos >= cutlass.Int32(0) + pos_c = cutlass.Int32(cutlass.select_(keep, pos, cutlass.Int32(0))) + page = ( + page_table.load(idx=req * pt_stride + pos_c // cutlass.Int32(PAGE)) + + page_offset + ) + dst = ( + cutlass.Int64(page) * cutlass.Int64(PAGE) + + cutlass.Int64(pos_c % cutlass.Int32(PAGE)) + ) * cutlass.Int64(row_stride) + prims.griddepcontrol(prims.GridDepAction.WAIT) + src = tok * cutlass.Int32(ag_cols) + cutlass.Int32(Q_LORA) + xs = [] + for j in cutlass.range_constexpr(2): + xs.append( + ag.load( + idx=src + (lane + cutlass.Int32(32 * j)) * cutlass.Int32(VEC), + vector_size=VEC, + alignment=16, + ) + ) + pe_v = cutlass.Int32( + cutlass.select_(lane < cutlass.Int32(PE // VEC), lane, cutlass.Int32(0)) + ) + pe = ag.load( + idx=src + cutlass.Int32(LATENT) + pe_v * cutlass.Int32(VEC), + vector_size=VEC, + alignment=16, + ) + ss = cutlass.Float32(0.0) + for j in cutlass.range_constexpr(2): + for e in cutlass.range_constexpr(VEC): + xf = cutlass.Float32(xs[j][e]) + ss = ss + xf * xf + for offset in (16, 8, 4, 2, 1): + ss = ss + cute.arch.shuffle_sync_bfly(ss, offset=offset) + r = cute.math.rsqrt(ss * cutlass.Float32(1.0 / LATENT) + kv_eps, fastmath=True) + dst_arr = kv_pool if cutlass.const_expr(kv_mode == 1) else kv_out + if keep: + for j in cutlass.range_constexpr(2): + outs = [] + for e in cutlass.range_constexpr(VEC): + # flashinfer's RMSNorm order: (x * rrms) * w in fp32, one bf16 rounding. + outs.append( + (cutlass.Float32(xs[j][e]) * r * cutlass.Float32(wk[j][e])).to( + io_dtype + ) + ) + dst_arr.store( + cutlass.Vector.from_elements(tuple(outs), io_dtype), + idx=dst + + cutlass.Int64((lane + cutlass.Int32(32 * j)) * cutlass.Int32(VEC)), + vector_size=VEC, + alignment=16, + ) + if lane < cutlass.Int32(PE // VEC): + dst_arr.store( + pe, + idx=dst + + cutlass.Int64(cutlass.Int32(LATENT) + lane * cutlass.Int32(VEC)), + vector_size=VEC, + alignment=16, + ) + elif warp_id == 2: + # ===================================================================== + # MMA: q_b partials (nope M 128, pe M 64) once B is normalized; then, + # on ranks 0-3, the k_b product once q_nope has arrived. + # ===================================================================== + idesc128 = prims.Tcgen05InstrDesc.build( + c_dtype=cutlass.Float32, a_dtype=io_dtype, b_dtype=io_dtype, n_dim=mma_n, m_dim=128 + ) + idesc64 = prims.Tcgen05InstrDesc.build( + c_dtype=cutlass.Float32, a_dtype=io_dtype, b_dtype=io_dtype, n_dim=mma_n, m_dim=64 + ) + swz = prims.Tcgen05SmemSwizzle.SWIZZLE_128B + desc_wn = prims.Tcgen05SmemDesc.build( + start_address=smem_wn, + leading_byte_offset=LEADING, + stride_byte_offset=STRIDE, + layout=swz, + ) + desc_wp = prims.Tcgen05SmemDesc.build( + start_address=smem_wp, + leading_byte_offset=LEADING, + stride_byte_offset=STRIDE, + layout=swz, + ) + desc_wk = prims.Tcgen05SmemDesc.build( + start_address=smem_wk, + leading_byte_offset=LEADING, + stride_byte_offset=STRIDE, + layout=swz, + ) + desc_b = prims.Tcgen05SmemDesc.build( + start_address=smem_b, leading_byte_offset=LEADING, stride_byte_offset=STRIDE, layout=swz + ) + desc_b2 = prims.Tcgen05SmemDesc.build( + start_address=smem_b2, + leading_byte_offset=LEADING, + stride_byte_offset=STRIDE, + layout=swz, + ) + tmem_nope = cutlass.inttoptr(tmem_base + cutlass.Int32(TMEM_NOPE), 6, cutlass.Int32) + tmem_pe = cutlass.inttoptr(tmem_base + cutlass.Int32(col_pe), 6, cutlass.Int32) + tmem_kb = cutlass.inttoptr(tmem_base + cutlass.Int32(col_kb), 6, cutlass.Int32) + while not cute.arch.mbarrier_try_wait(b_ready.data_ptr(), 0): + pass + for i in cutlass.range_constexpr(MY_TILES): + while not cute.arch.mbarrier_try_wait(w_full.subview(i).data_ptr(), 0): + pass + while not cute.arch.mbarrier_try_wait(w_full.subview(MY_TILES + i).data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + for kb in cutlass.range_constexpr(TMA_COPY_ITERS * K_BLOCKS_PER_HALF): + box = kb // K_BLOCKS_PER_HALF + within = kb % K_BLOCKS_PER_HALF + d_b = desc_b + (i * 2 * half_n + box * half_n + within * STEP) + acc = not (i == 0 and kb == 0) + if prims.elect_sync(): + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, prims.CTAGroup.CTA_1, tmem_nope, + desc_wn + (i * TILE_128 + box * HALF_128 + within * STEP), d_b, idesc128, acc, + ) # fmt: skip + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, prims.CTAGroup.CTA_1, tmem_pe, + desc_wp + (i * TILE_64 + box * HALF_64 + within * STEP), d_b, idesc64, acc, + ) # fmt: skip + if prims.elect_sync(): + prims.tcgen05_commit(acc_done) + if rank < cutlass.Int32(KB_RANKS): + while not cute.arch.mbarrier_try_wait(w_full.subview(2 * MY_TILES).data_ptr(), 0): + pass + if cutlass.const_expr(single_hop): + # This rank's epilogue wrote the B tile (generic stores, each writer fenced to the async proxy). + while not cute.arch.mbarrier_try_wait(b2_ready.data_ptr(), 0): + pass + else: + # q_nope arrives from the 4 owners' st.async stores, then feeds the tensor core (async proxy). + while not _try_wait_cluster(qn_full.data_ptr(), 0): + pass + prims.fence_proxy(prims.Proxy.ASYNC_SHARED, space=prims.SharedSpace.shared_cta) + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + for kb in cutlass.range_constexpr(TMA_COPY_ITERS * K_BLOCKS_PER_HALF): + box = kb // K_BLOCKS_PER_HALF + within = kb % K_BLOCKS_PER_HALF + if prims.elect_sync(): + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, prims.CTAGroup.CTA_1, tmem_kb, + desc_wk + (box * HALF_128 + within * STEP), desc_b2 + (box * half_n + within * STEP), + idesc128, kb != 0, + ) # fmt: skip + if prims.elect_sync(): + prims.tcgen05_commit(acc2_done) + elif warp_id >= 4: + tid = tx - cutlass.Int32(EPI_THREADS) + lane = tx % 32 + w = warp_id - 4 + # ===================================================================== + # Norm prolog: the RMS of every token's q_a row (1536 columns), then this + # rank's 256 columns normalized into the resident B tile. + # ===================================================================== + last_row = num_tokens - cutlass.Int32(1) + if cutlass.const_expr(cluster_rms): + # Thread tid: token t = tid / 16 of each group, this rank's vectors c and c + 16 (c = tid % 16; columns + # 256 rank + 8 c). Each rank reads only its own 256 columns; the per-token partial sums of squares go to + # every rank of the cluster (st.async into slot [rank][t], completing norm_full by bytes), which adds the + # 6 in rank order: the same total, hence the same rrms, on every rank. + cr_t_me = tid // cutlass.Int32(16) + cr_c_me = tid % cutlass.Int32(16) + cr_wk = [] + for cr_j in cutlass.range_constexpr(2): + cr_wk.append( + w_qa.load( + idx=(rank * cutlass.Int32(32) + cr_c_me + cutlass.Int32(16 * cr_j)) + * cutlass.Int32(VEC), + vector_size=VEC, + alignment=16, + ) + ) + prims.griddepcontrol(prims.GridDepAction.WAIT) + cr_t = [] + cr_live_t = [] + cr_row_t = [] + for cr_g in cutlass.range_constexpr(groups): + cr_t.append(_plus(cr_t_me, MMA_N * cr_g)) + cr_live_t.append(cr_t[cr_g] < num_tokens) + cr_row_t.append( + cutlass.Int32(cutlass.select_(cr_live_t[cr_g], cr_t[cr_g], last_row)) + ) + cr_xk = [] + cr_part = [] + for cr_g in cutlass.range_constexpr(groups): + cr_part_g = cutlass.Float32(0.0) + for cr_j in cutlass.range_constexpr(2): + cr_xv = ag.load( + idx=(tok0 + cr_row_t[cr_g]) * cutlass.Int32(ag_cols) + + (rank * cutlass.Int32(32) + cr_c_me + cutlass.Int32(16 * cr_j)) + * cutlass.Int32(VEC), + vector_size=VEC, + alignment=16, + ) + cr_xk.append(cr_xv) + for cr_e in cutlass.range_constexpr(VEC): + cr_xf = cutlass.Float32( + cutlass.select_( + cr_live_t[cr_g], cutlass.Float32(cr_xv[cr_e]), cutlass.Float32(0.0) + ) + ) + cr_part_g = cr_part_g + cr_xf * cr_xf + cr_part.append(cr_part_g) + for cr_g in cutlass.range_constexpr(groups): + for cr_off in (8, 4, 2, 1): + cr_part[cr_g] = cr_part[cr_g] + cute.arch.shuffle_sync_bfly( + cr_part[cr_g], offset=cr_off + ) + if cr_c_me == cutlass.Int32(0): + for cr_g in cutlass.range_constexpr(groups): + for cr_dst in cutlass.range_constexpr(CLUSTER): + _st_async_f32( + _mapa_u32(norm_mail.data_ptr(rank * cutlass.Int32(mma_n) + cr_t[cr_g]), cr_dst), + cr_part[cr_g], _mapa_u32(norm_full.data_ptr(), cr_dst), + ) # fmt: skip + while not _try_wait_cluster(norm_full.data_ptr(), 0): + pass + for cr_g in cutlass.range_constexpr(groups): + cr_total = norm_mail.load(idx=cr_t[cr_g]) + for cr_src in cutlass.range_constexpr(1, CLUSTER): + cr_total = cr_total + norm_mail.load( + idx=cutlass.Int32(cr_src * mma_n) + cr_t[cr_g] + ) + cr_r = cutlass.Float32( + cutlass.select_( + cr_live_t[cr_g], + cute.math.rsqrt( + cr_total * cutlass.Float32(1.0 / Q_LORA) + eps, fastmath=True + ), + cutlass.Float32(0.0), + ) + ) + for cr_j in cutlass.range_constexpr(2): + cr_outs = [] + for cr_e in cutlass.range_constexpr(VEC): + # flashinfer's RMSNorm order: (x * rrms) * w in fp32, one bf16 rounding. + cr_outs.append( + ( + cutlass.Float32(cr_xk[2 * cr_g + cr_j][cr_e]) + * cr_r + * cutlass.Float32(cr_wk[cr_j][cr_e]) + ).to(io_dtype) + ) + cr_k_local = (cr_c_me + cutlass.Int32(16 * cr_j)) * cutlass.Int32(VEC) + smem_b.store( + cutlass.Vector.from_elements(tuple(cr_outs), io_dtype), + idx=(cr_k_local // cutlass.Int32(CTA_K)) * cutlass.Int32(b_elems) + + _swz(cr_t[cr_g], cr_k_local % cutlass.Int32(CTA_K), mma_n), + vector_size=VEC, + alignment=16, + ) + else: + # RMS pass: thread tid loads 16-byte vector v = tid + 128 j (j = 1 only for tid < 64) of every token's + # 1536-column q_a row. This rank's own 32 vectors (v = 32 rank + lane) sit in warp rank % 4, pass rank // 4: + # that warp normalizes them from registers, so q_a is read once. + owner = w == rank % cutlass.Int32(4) + own_j = rank // cutlass.Int32(4) + # The norm weight of the owner lane's 8 columns is a weight: loaded ahead of the grid dependency. + wv = w_qa.load( + idx=(rank * cutlass.Int32(32) + lane) * cutlass.Int32(VEC), + vector_size=VEC, + alignment=16, + ) + prims.griddepcontrol(prims.GridDepAction.WAIT) + # Rows past M and vectors past the row are loaded clamped and count zero (no loads out of bounds, no + # values escaping runtime branches). + sums = [cutlass.Float32(0.0)] * MMA_N + xs = [] + for t in cutlass.range_constexpr(MMA_N): + row_t = cutlass.Int32( + cutlass.select_(cutlass.Int32(t) < num_tokens, cutlass.Int32(t), last_row) + ) + for j in cutlass.range_constexpr((Q_LORA // VEC + EPI_THREADS - 1) // EPI_THREADS): + v = tid + cutlass.Int32(j * EPI_THREADS) + live = (cutlass.Int32(t) < num_tokens) & (v < cutlass.Int32(Q_LORA // VEC)) + v_c = cutlass.Int32( + cutlass.select_(v < cutlass.Int32(Q_LORA // VEC), v, cutlass.Int32(0)) + ) + xv = ag.load( + idx=(tok0 + row_t) * cutlass.Int32(ag_cols) + v_c * cutlass.Int32(VEC), + vector_size=VEC, + alignment=16, + ) + xs.append(xv) + for e in cutlass.range_constexpr(VEC): + xf = cutlass.Float32( + cutlass.select_(live, cutlass.Float32(xv[e]), cutlass.Float32(0.0)) + ) + sums[t] = sums[t] + xf * xf + for t in cutlass.range_constexpr(MMA_N): + for offset in (16, 8, 4, 2, 1): + sums[t] = sums[t] + cute.arch.shuffle_sync_bfly(sums[t], offset=offset) + if lane == 0: + for t in cutlass.range_constexpr(MMA_N): + norm_part.store(sums[t], idx=w * cutlass.Int32(MMA_N) + cutlass.Int32(t)) + prims.barrier_cta_sync(1, thread_count=EPI_THREADS) + if tid < cutlass.Int32(MMA_N): + total = norm_part.load(idx=tid) + for ww in cutlass.range_constexpr(1, 4): + total = total + norm_part.load(idx=cutlass.Int32(ww * MMA_N) + tid) + rrms.store( + cute.math.rsqrt(total * cutlass.Float32(1.0 / Q_LORA) + eps, fastmath=True), + idx=tid, + ) + prims.barrier_cta_sync(1, thread_count=EPI_THREADS) + # The owner warp: 8 tokens x its lane's 8 columns (CTA column 8 lane) into the B tile (rows >= M are zero). + n_pass = (Q_LORA // VEC + EPI_THREADS - 1) // EPI_THREADS + if owner: + for jo in cutlass.range_constexpr(n_pass): + if own_j == cutlass.Int32(jo): + for t in cutlass.range_constexpr(MMA_N): + xv = xs[t * n_pass + jo] + r = cutlass.Float32( + cutlass.select_( + cutlass.Int32(t) < num_tokens, + rrms.load(idx=cutlass.Int32(t)), + cutlass.Float32(0.0), + ) + ) + outs = [] + for e in cutlass.range_constexpr(VEC): + # flashinfer's RMSNorm order: (x * rrms) * w in fp32, one bf16 rounding. + outs.append( + (cutlass.Float32(xv[e]) * r * cutlass.Float32(wv[e])).to( + io_dtype + ) + ) + k_local = lane * cutlass.Int32(VEC) + smem_b.store( + cutlass.Vector.from_elements(tuple(outs), io_dtype), + idx=(k_local // cutlass.Int32(CTA_K)) * cutlass.Int32(b_elems) + + _swz(cutlass.Int32(t), k_local % cutlass.Int32(CTA_K)), + vector_size=VEC, + alignment=16, + ) + prims.fence_proxy(prims.Proxy.ASYNC_SHARED, space=prims.SharedSpace.shared_cta) + prims.mbarrier_arrive(b_ready) + + # ===================================================================== + # q_b partials -> row owners (DSMEM), rank-order sum, bf16. + # ===================================================================== + while not cute.arch.mbarrier_try_wait(acc_done.data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc_n = prims.tcgen05_ld( + "32x32b", + cutlass.inttoptr(tmem_base + cutlass.Int32(TMEM_NOPE), 6, cutlass.Float32), + num=mma_n, + ) + acc_p = prims.tcgen05_ld( + "32x32b", + cutlass.inttoptr(tmem_base + cutlass.Int32(col_pe), 6, cutlass.Float32), + num=mma_n, + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + # nope rows 32 w + lane -> rank w; pe rows 16 w + lane (lane < 16) -> rank 4 + w // 2, local row + # 16 (w % 2) + lane. + pe_owner = cutlass.Int32(KB_RANKS) + w // cutlass.Int32(2) + pe_row = (w % cutlass.Int32(2)) * cutlass.Int32(16) + lane + if cutlass.const_expr(single_hop): + # Nope row 32 w + lane, 8 tokens -> slot [rank] of every k_b rank's mailbox (this rank's too). + base_kb = (rank * cutlass.Int32(128) + w * cutlass.Int32(32) + lane) * cutlass.Int32( + MMA_N + ) + for dst in cutlass.range_constexpr(KB_RANKS): + mbar = _mapa_u32(kb_full.data_ptr(), dst) + for h4 in cutlass.range_constexpr(MMA_N // 4): + _st_async_v4( + _mapa_u32(kb_mail.data_ptr(base_kb + cutlass.Int32(4 * h4)), dst), + acc_n[4 * h4], acc_n[4 * h4 + 1], acc_n[4 * h4 + 2], acc_n[4 * h4 + 3], mbar, + ) # fmt: skip + else: + if w != rank: + # st.async into the owner's mailbox [source][token][row]: a warp's 32 lanes fill 128 contiguous bytes + # per token; the owner's barrier completes when all 5 sources' bytes have landed. + mb_n = _mapa_u32(mail_full.data_ptr(), w) + for t in cutlass.range_constexpr(mma_n): + _st_async_f32( + _mapa_u32(mailbox.data_ptr((rank * cutlass.Int32(mma_n) + cutlass.Int32(t)) * cutlass.Int32(32) + + lane), w), + acc_n[t], mb_n, + ) # fmt: skip + if lane < cutlass.Int32(16): + if pe_owner != rank: + mb_p = _mapa_u32(mail_full.data_ptr(), pe_owner) + for t in cutlass.range_constexpr(mma_n): + _st_async_f32( + _mapa_u32(mailbox.data_ptr((rank * cutlass.Int32(mma_n) + cutlass.Int32(t)) * cutlass.Int32(32) + + pe_row), pe_owner), + acc_p[t], mb_p, + ) # fmt: skip + if cutlass.const_expr(single_hop): + if rank < cutlass.Int32(KB_RANKS): + # Nope row tid: the 6 partials in rank order (the owner's sum of the two-hop path), bf16, into the + # k_b B tile (token t, column tid), fenced to the async proxy for the MMA warp. + while not _try_wait_cluster(kb_full.data_ptr(), 0): + pass + parts = [] + for src in cutlass.range_constexpr(CLUSTER): + for h4 in cutlass.range_constexpr(MMA_N // 4): + parts.append( + kb_mail.load( + idx=(cutlass.Int32(src * 128) + tid) * cutlass.Int32(MMA_N) + + cutlass.Int32(4 * h4), + vector_size=4, + alignment=16, + ) + ) + for t in cutlass.range_constexpr(MMA_N): + total = cutlass.Float32(0.0) + for src in cutlass.range_constexpr(CLUSTER): + total = total + cutlass.Float32(parts[src * (MMA_N // 4) + t // 4][t % 4]) + smem_b2.store( + total.to(io_dtype), + idx=_swz(cutlass.Int32(t), tid) + tid % cutlass.Int32(VEC), + ) + prims.fence_proxy(prims.Proxy.ASYNC_SHARED, space=prims.SharedSpace.shared_cta) + prims.mbarrier_arrive(b2_ready) + elif cutlass.const_expr(mma_n == MMA_N): + if w == rank: + # Rank w < 4 owns nope rows 32 w + lane: sum the 6 partials, round, stage for the k_b broadcast. + while not _try_wait_cluster(mail_full.data_ptr(), 0): + pass + for t in cutlass.range_constexpr(MMA_N): + total = cutlass.Float32(0.0) + for src in cutlass.range_constexpr(CLUSTER): + part = mailbox.load(idx=cutlass.Int32((src * MMA_N + t) * 32) + lane) + total = total + cutlass.Float32( + cutlass.select_( + rank == cutlass.Int32(src), cutlass.Float32(acc_n[t]), part + ) + ) + stage_n.store(total.to(io_dtype), idx=cutlass.Int32(t * 32) + lane) + prims.bar_warp_sync(0xFFFFFFFF) + # Lane l: token l // 4, 8 of the owner's 32 columns (16 bytes) -> st.async into the k_b B tile of ranks + # 0-3 (this one too), completing their q_nope barrier. + ct = lane // cutlass.Int32(4) + c4 = lane % cutlass.Int32(4) + stage_w = cutlass.Array( + cutlass.inttoptr(stage_n.data_ptr().toint(), 3, cutlass.Int32), + shape=MMA_N * 32 // 2, + ) + words = stage_w.load( + idx=(ct * cutlass.Int32(32) + c4 * cutlass.Int32(VEC)) // cutlass.Int32(2), + vector_size=4, + alignment=16, + ) + dst_idx = _swz(ct, rank * cutlass.Int32(32) + c4 * cutlass.Int32(VEC)) + for dst in cutlass.range_constexpr(KB_RANKS): + _st_async_v4_b32( + _mapa_u32(smem_b2.data_ptr(dst_idx), dst), words[0], words[1], words[2], words[3], + _mapa_u32(qn_full.data_ptr(), dst), + ) # fmt: skip + else: + if rank < cutlass.Int32(KB_RANKS): + # Rank r < 4 owns nope rows 32 r + lane. Its warp r puts its own partials into the mailbox's slot [r]; + # then epilogue warp v sums the 6 partials of tokens 8 v .. 8 v + 7 in rank order, rounds, stages and + # sends them as above (each group's sums and transfers are an 8-token chunk's). + if w == rank: + for t in cutlass.range_constexpr(mma_n): + mailbox.store( + cutlass.Float32(acc_n[t]), + idx=(rank * cutlass.Int32(mma_n) + cutlass.Int32(t)) * cutlass.Int32(32) + + lane, + ) + prims.barrier_cta_sync(1, thread_count=EPI_THREADS) + while not _try_wait_cluster(mail_full.data_ptr(), 0): + pass + if w < cutlass.Int32(groups): + for t in cutlass.range_constexpr(MMA_N): + g_tok = w * cutlass.Int32(MMA_N) + cutlass.Int32(t) + g_total = cutlass.Float32(0.0) + for src in cutlass.range_constexpr(CLUSTER): + g_total = g_total + cutlass.Float32( + mailbox.load( + idx=(cutlass.Int32(src * mma_n) + g_tok) * cutlass.Int32(32) + + lane + ) + ) + stage_n.store(g_total.to(io_dtype), idx=g_tok * cutlass.Int32(32) + lane) + prims.bar_warp_sync(0xFFFFFFFF) + g_ct = w * cutlass.Int32(MMA_N) + lane // cutlass.Int32(4) + g_c4 = lane % cutlass.Int32(4) + g_stage = cutlass.Array( + cutlass.inttoptr(stage_n.data_ptr().toint(), 3, cutlass.Int32), + shape=mma_n * 32 // 2, + ) + g_words = g_stage.load( + idx=(g_ct * cutlass.Int32(32) + g_c4 * cutlass.Int32(VEC)) + // cutlass.Int32(2), + vector_size=4, + alignment=16, + ) + g_dst = _swz(g_ct, rank * cutlass.Int32(32) + g_c4 * cutlass.Int32(VEC), mma_n) + for dst in cutlass.range_constexpr(KB_RANKS): + _st_async_v4_b32( + _mapa_u32(smem_b2.data_ptr(g_dst), dst), g_words[0], g_words[1], g_words[2], g_words[3], + _mapa_u32(qn_full.data_ptr(), dst), + ) # fmt: skip + if rank >= cutlass.Int32(KB_RANKS): + if pe_owner == rank: + if lane < cutlass.Int32(16): + # Rank 4 + j owns pe rows 32 j + pe_row: sum, round, store fused_q[t, head, 512 + row]. + while not _try_wait_cluster(mail_full.data_ptr(), 0): + pass + p = (rank - cutlass.Int32(KB_RANKS)) * cutlass.Int32(32) + pe_row + for t in cutlass.range_constexpr(mma_n): + total = cutlass.Float32(0.0) + for src in cutlass.range_constexpr(CLUSTER): + part = mailbox.load(idx=cutlass.Int32((src * mma_n + t) * 32) + pe_row) + total = total + cutlass.Float32( + cutlass.select_( + rank == cutlass.Int32(src), cutlass.Float32(acc_p[t]), part + ) + ) + if cutlass.Int32(t) < num_tokens: + fused_q.store( + total.to(io_dtype), + idx=(tok0 + cutlass.Int32(t)) * cutlass.Int32(total_heads * FUSED) + + head * cutlass.Int32(FUSED) + + cutlass.Int32(LATENT) + + p, + ) + # ===================================================================== + # k_b product (ranks 0-3): fused_q[t, head, 128 rank + i]. + # ===================================================================== + if rank < cutlass.Int32(KB_RANKS): + while not cute.arch.mbarrier_try_wait(acc2_done.data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc_k = prims.tcgen05_ld( + "32x32b", + cutlass.inttoptr(tmem_base + cutlass.Int32(col_kb), 6, cutlass.Float32), + num=mma_n, + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + row = w * cutlass.Int32(32) + lane + for t in cutlass.range_constexpr(mma_n): + stage_o.store( + cutlass.Float32(acc_k[t]).to(io_dtype), idx=cutlass.Int32(t * 128) + row + ) + # Every TMEM reader has waited for its loads; the staging tile is complete. + prims.barrier_cta_sync(1, thread_count=EPI_THREADS) + if warp_id == 4: + prims.tcgen05_dealloc(cutlass.inttoptr(tmem_base, 6, cutlass.Int32), tmem_cols(mma_n)) + if rank < cutlass.Int32(KB_RANKS): + o_tok = tid // cutlass.Int32(16) + oc = tid % cutlass.Int32(16) + for g in cutlass.range_constexpr(groups): + o_tok_g = _plus(o_tok, MMA_N * g) + if o_tok_g < num_tokens: + out_vec = stage_o.load( + idx=o_tok_g * cutlass.Int32(128) + oc * cutlass.Int32(VEC), + vector_size=VEC, + alignment=16, + ) + fused_q.store( + out_vec, + idx=(tok0 + o_tok_g) * cutlass.Int32(total_heads * FUSED) + + head * cutlass.Int32(FUSED) + + rank * cutlass.Int32(128) + + oc * cutlass.Int32(VEC), + vector_size=VEC, + alignment=16, + ) + + +def _weight_map(w, n_rows, k_in, box_rows): + """W [n_rows, k_in] as five TMA dimensions (64-column chunk, row, chunk index, 1, 1): one call per 128-column + k-tile lands both 128-byte-swizzled halves of `box_rows` rows.""" + return cuda.create_tensor_map_tiled( + global_address=w.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[TMA_K_BOX, n_rows, k_in // TMA_K_BOX, 1, 1], + global_strides=[ + (k_in * ELEM_BYTES) // 16, + (TMA_K_BOX * ELEM_BYTES) // 16, + (n_rows * k_in * ELEM_BYTES) // 16, + (n_rows * k_in * ELEM_BYTES) // 16, + ], + box_dims=[TMA_K_BOX, box_rows, TMA_COPY_ITERS, 1, 1], + swizzle=cuda.TensorMapSwizzle.s128b, + ) + + +@cute.jit +def k3_mla_q( + w_qb: cute.Tensor, # [heads * 192, 1536] bf16 + w_kb: cute.Tensor, # [heads * 512, 128] bf16 (k_b_proj_trans rows) + ag: cute.Tensor, # [M * ag_cols] bf16: q_a in columns [0, 1536) of each row + w_qa: cute.Tensor, # [1536] bf16 + fused_q: cute.Tensor, # [M * heads * 576] bf16 + w_kv: cute.Tensor, # [512] bf16 (kv_mode) + kv_pool: cute.Tensor, # flat bf16 latent pool (kv_mode 1) + page_table: cute.Tensor, # int32, request i's pages at [i * pt_stride, ...) (kv_mode 1) + seq_len: cute.Tensor, # int32 [R] (kv_mode 1) + kv_out: cute.Tensor, # [M * 576] bf16 (kv_mode 2) + num_tokens: cutlass.Int32, # M + tokens: cutlass.Int32, # T = M / R (kv_mode 1) + pt_stride: cutlass.Int32, # (kv_mode 1) + eps: cutlass.Float32, + kv_eps: cutlass.Float32, + page_offset: cutlass.Int32, + ag_cols: cutlass.Constexpr[int], + total_heads: cutlass.Constexpr[int], + trigger_early: cutlass.Constexpr[bool], + single_hop: cutlass.Constexpr[bool], + kv_mode: cutlass.Constexpr[int], + row_stride: cutlass.Constexpr[int], + cluster_rms: cutlass.Constexpr[bool], + use_pdl: cutlass.Constexpr[bool], + mma_n: cutlass.Constexpr[ + int + ], # tokens per chunk: MMA_N, or one of WIDE_CHUNKS with cluster_rms, two-hop + stream: cuda_driver.CUstream, +) -> None: + if cutlass.const_expr( + mma_n != MMA_N and (mma_n not in WIDE_CHUNKS or single_hop or not cluster_rms) + ): + raise ValueError( + f"k3_mla_q: chunks of {mma_n} tokens need one of {WIDE_CHUNKS}, cluster_rms and two-hop" + ) + tma_nope = _weight_map(w_qb, total_heads * QK, Q_LORA, 128) + tma_pe = _weight_map(w_qb, total_heads * QK, Q_LORA, 64) + tma_kb = _weight_map(w_kb, total_heads * LATENT, NOPE, 128) + k3_mla_q_kernel( + tma_nope, + tma_pe, + tma_kb, + ag, + w_qa, + fused_q, + w_kv, + kv_pool, + page_table, + seq_len, + kv_out, + num_tokens, + tokens, + pt_stride, + eps, + kv_eps, + page_offset, + ag_cols, + total_heads, + trigger_early, + single_hop, + kv_mode, + row_stride, + cluster_rms, + mma_n, + ).launch( # fmt: skip + grid=(total_heads * CLUSTER, (num_tokens + mma_n - 1) // mma_n, 1), + block=(THREADS, 1, 1), + cluster=(CLUSTER, 1, 1), + stream=stream, + use_pdl=use_pdl, + ) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/op.py new file mode 100644 index 000000000000..6d953a76427e --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/op.py @@ -0,0 +1,575 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Torch ops of the Kimi K3 MLA decode CTM kernels. + +A decode step is M = R x T tokens: R <= 8 requests of T <= 8 tokens each, request-major (the rows of request i are +i T .. i T + T - 1). Request i's KV pages are row i of ``page_table`` (int32 [R, W] with unit column stride and any row +stride, e.g. a slice of the attention metadata's kv_cache_block_offsets; [W] for one request) and its KV length, +including its T new tokens, is ``seq_len[i]`` (int32 [R]). + +``trtllm::k3_mla_q``: the decode query path (M <= 64; 6 heads per rank at TP16, 24 at TP4): q_a RMSNorm, q_b projection +and k_b absorption in one launch, producing the attention's ``fused_q`` [M, heads * 576] = [q_nope @ W_kb^T | q_pe] per +head; ``trtllm::k3_mla_qkv`` also stores the KV half (kv_a RMSNorm, rope columns) into the paged latent cache. +``trtllm::k3_mla_attn`` and its ``_out`` / ``_vb_out`` forms: the attention of every request over its pages. Compiled +on the first call for its shapes, which must happen outside CUDA-graph capture. +""" + +from __future__ import annotations + +import os +import threading +from typing import Dict, Optional + +import torch + +MAX_TOKENS = 64 # tokens of a k3_mla_q / k3_mla_qkv call +MAX_REQUESTS = 8 +MAX_REQUEST_TOKENS = 8 # T +# k3_mla_q's tokens per chunk (the MMA's N) by call size, as (largest call, chunk): a chunk is a cluster of 6 CTAs per +# head, 23 such clusters are resident at once on GB200, and a wider chunk takes longer per CTA. 8 up to 24 tokens (18 +# clusters), 16 up to 48 (18), 32 up to 64 (12). +CHUNK_TOKENS = ((24, 8), (48, 16), (64, 32)) + +_lock = threading.Lock() +_compiled: Dict[tuple, object] = {} + + +def _arg(t: torch.Tensor, align: int = 16): + from cutlass.cute.runtime import from_dlpack + + return from_dlpack(t.detach(), assumed_align=align).mark_layout_dynamic(leading_dim=t.dim() - 1) + + +def _use_pdl() -> bool: + return os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" + + +_arg_dummies: Dict[tuple, torch.Tensor] = {} + + +def _dummy(device: torch.device, dtype: torch.dtype) -> torch.Tensor: + """A tensor argument the build does not touch (the kv arguments when the KV half is off).""" + key = (device.index, dtype) + t = _arg_dummies.get(key) + if t is None: + t = _arg_dummies[key] = torch.zeros(8, dtype=dtype, device=device) + return t + + +def _pool_base(pool: torch.Tensor, row_stride: int) -> torch.Tensor: + """The pool's first page as the kernels' pointer carrier: they index the pool with 64-bit offsets (or through a + tensor map sized separately), and a flat view of a multi-GB pool would not fit a 32-bit dynamic shape.""" + return pool.view(-1)[: 64 * row_stride] + + +def _requests(page_table: torch.Tensor, seq_len: torch.Tensor, num_tokens: int): + """``(rows, row_stride, R, T)`` of a step of ``num_tokens`` tokens over R = ``seq_len.numel()`` requests: the + page-table rows as one flat int32 view (request i's row at ``i * row_stride``), or None when the arguments do not + describe R requests of T = num_tokens / R tokens (see the module docstring).""" + num_requests = seq_len.numel() + if not ( + page_table.dtype == seq_len.dtype == torch.int32 + and page_table.device == seq_len.device + and seq_len.dim() == 1 + and seq_len.is_contiguous() + and 0 < num_requests <= num_tokens + and num_tokens % num_requests == 0 + and page_table.dim() in (1, 2) + and page_table.stride(-1) == 1 + and page_table.shape[-1] > 0 + ): + return None + if page_table.dim() == 1: + return (page_table, 0, 1, num_tokens) if num_requests == 1 else None + width = page_table.shape[1] + row_stride = page_table.stride(0) if num_requests > 1 else width + if page_table.shape[0] != num_requests or row_stride < width: + return None + rows = page_table.as_strided(((num_requests - 1) * row_stride + width,), (1,)) + return rows, row_stride, num_requests, num_tokens // num_requests + + +def supports_kv(ag: torch.Tensor, w_kv: torch.Tensor, pool: torch.Tensor, row_stride: int, page_table: torch.Tensor, + seq_len: torch.Tensor) -> bool: # fmt: skip + """Whether ``k3_mla_qkv`` can store the KV half of the ``ag.shape[0]`` tokens: the latent (512) and rope (64) + columns after q_a in ``ag``, a dense bf16 pool with rows of ``row_stride`` elements, int32 page-table rows and + lengths of R requests (see the module docstring).""" + from . import k3_mla_q_kernel as kernel + + return ( + ag.shape[1] >= kernel.Q_LORA + kernel.FUSED + and tuple(w_kv.shape) == (kernel.LATENT,) + and w_kv.dtype == torch.bfloat16 + and w_kv.is_contiguous() + and pool.dtype == torch.bfloat16 + and pool.is_contiguous() + and row_stride >= kernel.FUSED + and row_stride % 8 == 0 + and pool.numel() >= 64 * row_stride + and _requests(page_table, seq_len, ag.shape[0]) is not None + ) + + +def supports_q( + ag: torch.Tensor, w_qa: torch.Tensor, w_qb: torch.Tensor, w_kb: torch.Tensor +) -> bool: + """Whether ``k3_mla_q`` runs: q_a in the first 1536 columns of dense bf16 rows, the TP16 per-rank shapes.""" + from . import k3_mla_q_kernel as kernel + + return ( + ag.is_cuda + and ag.dtype == w_qa.dtype == w_qb.dtype == w_kb.dtype == torch.bfloat16 + and ag.dim() == 2 + and 0 < ag.shape[0] <= MAX_TOKENS + and ag.shape[1] >= kernel.Q_LORA + and ag.shape[1] % 8 == 0 + and ag.is_contiguous() + and tuple(w_qa.shape) == (kernel.Q_LORA,) + and w_kb.dim() == 3 + and tuple(w_kb.shape[1:]) == (kernel.LATENT, kernel.NOPE) + and 0 < w_kb.shape[0] * kernel.CLUSTER <= 148 + and tuple(w_qb.shape) == (w_kb.shape[0] * kernel.QK, kernel.Q_LORA) + and w_qa.is_contiguous() + and w_qb.is_contiguous() + and w_kb.is_contiguous() + ) + + +def _launch_q( + ag, w_qa, eps, w_qb, w_kb, trigger_early=True, single_hop=False, kv=None, cluster_rms=True +) -> torch.Tensor: + """``kv``: None (query path only) or dict(w, eps) plus either pool, row_stride, page_table, page_offset, seq_len + (the KV half into the paged pool) or out (into a dense [M, 576] tensor). ``cluster_rms``: each rank of a head's + cluster reads only its 256 q_a columns and the ranks exchange per-token partial sums of squares (st.async; False: + every CTA reads all 1536 columns). ``single_hop``: the single-hop k_b reduce (same bits, slower on GB200). + The call runs in chunks of CHUNK_TOKENS' size for its token count when cluster_rms and two-hop (and, with the KV + half, when the heads launch a CTA per token of a chunk), else in chunks of 8: every token's bits are the same.""" + import cuda.bindings.driver as cuda_driver + + if not supports_q(ag, w_qa, w_qb, w_kb): + raise ValueError( + f"k3_mla_q: unsupported call ag {tuple(ag.shape)} {ag.dtype}, w_qa {tuple(w_qa.shape)}, " + f"w_qb {tuple(w_qb.shape)}, w_kb {tuple(w_kb.shape)}" + ) + from . import k3_mla_q_kernel as kernel + + num_tokens, ag_cols = ag.shape + heads = w_kb.shape[0] + mma_n = next(n for limit, n in CHUNK_TOKENS if num_tokens <= limit) + if single_hop or not cluster_rms or (kv is not None and heads * kernel.CLUSTER < mma_n): + mma_n = kernel.MMA_N + out = torch.empty(num_tokens, heads * kernel.FUSED, dtype=torch.bfloat16, device=ag.device) + bf16_dummy, i32_dummy = _dummy(ag.device, torch.bfloat16), _dummy(ag.device, torch.int32) + kv_mode, row_stride, kv_eps, page_offset = 0, kernel.FUSED, 0.0, 0 + tokens, pt_stride = num_tokens, 0 + w_kv, kv_pool, page_table, seq_len, kv_out = ( + bf16_dummy, + bf16_dummy, + i32_dummy, + i32_dummy, + bf16_dummy, + ) + if kv is not None: + w_kv, kv_eps = kv["w"], float(kv["eps"]) + if heads * kernel.CLUSTER < kernel.MMA_N: + raise ValueError( + f"k3_mla_q: the KV half takes {kernel.MMA_N} CTAs per 8 tokens, {heads} heads launch " + f"{heads * kernel.CLUSTER}" + ) + if "out" in kv: + kv_mode, kv_out = 2, kv["out"] + if not ( + kv_out.is_contiguous() + and kv_out.dtype == torch.bfloat16 + and kv_out.numel() == num_tokens * kernel.FUSED + ): + raise ValueError( + f"k3_mla_q: kv out {tuple(kv_out.shape)} {kv_out.dtype} is not a dense [M, 576] bf16" + ) + else: + kv_mode, row_stride, page_offset = 1, int(kv["row_stride"]), int(kv["page_offset"]) + page_table, seq_len = kv["page_table"], kv["seq_len"] + if not supports_kv(ag, w_kv, kv["pool"], row_stride, page_table, seq_len): + raise ValueError( + f"k3_mla_q: unsupported KV half: ag {tuple(ag.shape)}, w_kv {tuple(w_kv.shape)} {w_kv.dtype}, pool " + f"{kv['pool'].dtype} rows of {row_stride}, page table {tuple(page_table.shape)} " + f"{page_table.dtype} strides {page_table.stride()}, lengths {tuple(seq_len.shape)} {seq_len.dtype}" + ) + page_table, pt_stride, _, tokens = _requests(page_table, seq_len, num_tokens) + kv_pool = _pool_base(kv["pool"], row_stride) + args = ( + _arg(w_qb), + _arg(w_kb.view(heads * kernel.LATENT, kernel.NOPE)), + _arg(ag.view(-1)), + _arg(w_qa), + _arg(out.view(-1)), + _arg(w_kv), + _arg(kv_pool), + # int32 page-table rows and lengths: views into the metadata buffers, read by scalar loads. + _arg(page_table, 4), + _arg(seq_len, 4), + _arg(kv_out.view(-1)), + ) + stream = cuda_driver.CUstream(torch.cuda.current_stream(ag.device).cuda_stream) + use_pdl = _use_pdl() + key = ( + "k3_mla_q", + ag_cols, + heads, + trigger_early, + single_hop, + kv_mode, + row_stride, + cluster_rms, + use_pdl, + mma_n, + ) + fn = _compiled.get(key) + if fn is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "trtllm::k3_mla_q must run once outside CUDA-graph capture first (it compiles its kernel)." + ) + import cutlass.cute as cute + + with _lock: + fn = _compiled.get(key) + if fn is None: + fn = _compiled[key] = cute.compile( + kernel.k3_mla_q, *args, num_tokens, tokens, pt_stride, float(eps), kv_eps, page_offset, ag_cols, + heads, trigger_early, single_hop, kv_mode, row_stride, cluster_rms, use_pdl, mma_n, stream, + ) # fmt: skip + fn(*args, num_tokens, tokens, pt_stride, float(eps), kv_eps, page_offset, stream) + return out + + +@torch.library.custom_op("trtllm::k3_mla_q", mutates_args=()) +def k3_mla_q( + ag: torch.Tensor, + w_qa: torch.Tensor, + eps: float, + w_qb: torch.Tensor, + w_kb: torch.Tensor, + trigger_early: bool = True, +) -> torch.Tensor: + """``fused_q`` [M <= 64, heads * 576] bf16 from ``ag`` [M, C >= 1536] (q_a = ag[:, :1536]): per head, + ``[bf16(q_nope @ w_kb[h]^T) | q_pe]`` with ``q = bf16(rmsnorm(q_a) @ w_qb^T)``; each bf16 rounding of the + unfused RMSNorm -> q_b -> bmm chain is kept. Every 8-token chunk of rows computes as an M <= 8 call would.""" + return _launch_q(ag, w_qa, eps, w_qb, w_kb, trigger_early) + + +@k3_mla_q.register_fake +def _(ag, w_qa, eps, w_qb, w_kb, trigger_early=True): + return ag.new_empty((ag.shape[0], w_kb.shape[0] * 576), dtype=torch.bfloat16) + + +@torch.library.custom_op("trtllm::k3_mla_qkv", mutates_args=("pool",)) +def k3_mla_qkv( + ag: torch.Tensor, + w_qa: torch.Tensor, + eps: float, + w_qb: torch.Tensor, + w_kb: torch.Tensor, + w_kv: torch.Tensor, + kv_eps: float, + pool: torch.Tensor, + row_stride: int, + page_table: torch.Tensor, + page_offset: int, + seq_len: torch.Tensor, + trigger_early: bool = True, +) -> torch.Tensor: + """``k3_mla_q`` plus the KV half in the same launch: the cache row ``[bf16(rmsnorm(ag[t, 1536:2048]) * w_kv) | + ag[t, 2048:2112]]`` (K3 is NoPE) of token t = i T + u (token u of request i, see the module docstring) stored into + the paged latent ``pool`` at position ``pos = seq_len[i] - T + u``, row ``(page_table[i][pos // 64] + + page_offset) * 64 + pos % 64`` of ``row_stride`` elements (``seq_len[i] >= T``; a token with ``pos < 0`` is not + stored). The page table and lengths are read before the grid dependency wait (written before the graph runs).""" + kv = dict(w=w_kv, eps=kv_eps, pool=pool, row_stride=row_stride, page_table=page_table, page_offset=page_offset, + seq_len=seq_len) # fmt: skip + return _launch_q(ag, w_qa, eps, w_qb, w_kb, trigger_early, kv=kv) + + +@k3_mla_qkv.register_fake +def _( + ag, + w_qa, + eps, + w_qb, + w_kb, + w_kv, + kv_eps, + pool, + row_stride, + page_table, + page_offset, + seq_len, + trigger_early=True, +): + return ag.new_empty((ag.shape[0], w_kb.shape[0] * 576), dtype=torch.bfloat16) + + +@torch.library.custom_op("trtllm::k3_mla_qkv_out", mutates_args=("kv_out",)) +def k3_mla_qkv_out( + ag: torch.Tensor, + w_qa: torch.Tensor, + eps: float, + w_qb: torch.Tensor, + w_kb: torch.Tensor, + w_kv: torch.Tensor, + kv_eps: float, + kv_out: torch.Tensor, + trigger_early: bool = True, +) -> torch.Tensor: + """``k3_mla_qkv`` with the cache rows stored densely into ``kv_out`` [M, 576] instead of the pool (checks).""" + return _launch_q( + ag, w_qa, eps, w_qb, w_kb, trigger_early, kv=dict(w=w_kv, eps=kv_eps, out=kv_out) + ) + + +@k3_mla_qkv_out.register_fake +def _(ag, w_qa, eps, w_qb, w_kb, w_kv, kv_eps, kv_out, trigger_early=True): + return ag.new_empty((ag.shape[0], w_kb.shape[0] * 576), dtype=torch.bfloat16) + + +# --------------------------------------------------------------------------------------------------------------- +# trtllm::k3_mla_attn: decode attention over the paged latent cache (R <= 8 requests of T <= 8 tokens, one cluster of +# 16 CTAs per request and 6 heads) +# --------------------------------------------------------------------------------------------------------------- +_workspaces: Dict[tuple, torch.Tensor] = {} + + +def _attn_workspace(device: torch.device, groups: int) -> torch.Tensor: + """The per-CTA partials of every (request, head group), then the no_cluster mode's (m, l) exchange and arrival + counters (zeroed): allocated once, for MAX_REQUESTS, on the first call.""" + from . import k3_mla_attn_kernel as kernel + + key = (device.index, groups) + ws = _workspaces.get(key) + if ws is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "trtllm::k3_mla_attn must run once outside CUDA-graph capture first (it allocates its workspace)." + ) + slots = kernel.MAX_REQUESTS * groups * kernel.CLUSTER + ws = _workspaces[key] = torch.empty( + slots * kernel.WS_SLOT_ELEMS + kernel.ws_sync_elems(groups), dtype=torch.float16, device=device + ) + ws[slots * kernel.WS_SLOT_ELEMS :].zero_() + return ws + + +def supports_attn( + q: torch.Tensor, + pool: torch.Tensor, + row_stride: int, + page_table: torch.Tensor, + seq_len: torch.Tensor, +) -> bool: + """Whether ``k3_mla_attn`` takes the call: ``q`` [M = R T, heads * 576] bf16 with heads a multiple of 6, R <= 8 + requests of T <= 8 tokens (page-table rows and lengths as in the module docstring), a dense bf16 pool with rows of + ``row_stride`` elements.""" + from . import k3_mla_attn_kernel as kernel + + if not ( + q.is_cuda + and q.dtype == pool.dtype == torch.bfloat16 + and q.dim() == 2 + and q.shape[1] % (kernel.HEADS * kernel.QK) == 0 + and q.is_contiguous() + and pool.is_contiguous() + and row_stride >= kernel.QK + and row_stride % 8 == 0 + ): + return False + requests = _requests(page_table, seq_len, q.shape[0]) + return ( + requests is not None + and requests[2] <= kernel.MAX_REQUESTS + and requests[3] <= kernel.MAX_TOKENS + ) + + +def _launch_attn( + q, + pool, + row_stride, + page_table, + seq_len, + softmax_scale, + out=None, + page_offset=0, + w_vb=None, + gate=None, + gate_col0=0, +): + import cuda.bindings.driver as cuda_driver + + from . import k3_mla_attn_kernel as kernel + + if not supports_attn(q, pool, row_stride, page_table, seq_len): + raise ValueError( + f"k3_mla_attn: unsupported call q {tuple(q.shape)} {q.dtype}, row_stride {row_stride}, page table " + f"{tuple(page_table.shape)} {page_table.dtype} strides {page_table.stride()}, lengths " + f"{tuple(seq_len.shape)} {seq_len.dtype}" + ) + num_tokens = q.shape[0] + page_rows, pt_stride, num_requests, tokens = _requests(page_table, seq_len, num_tokens) + total_heads = q.shape[1] // kernel.QK + fuse_vb = w_vb is not None + if fuse_vb and not ( + w_vb.dtype == torch.bfloat16 + and w_vb.is_contiguous() + and tuple(w_vb.shape) == (total_heads, kernel.V_DIM, kernel.LATENT) + ): + raise ValueError( + f"k3_mla_attn: v_b weight {tuple(w_vb.shape)} {w_vb.dtype} is not a dense [heads, 128, 512] bf16" + ) + width = kernel.V_DIM if fuse_vb else kernel.LATENT + apply_gate = gate is not None + if apply_gate and not ( + fuse_vb + and gate.dtype == torch.bfloat16 + and gate.dim() == 2 + and gate.shape[0] == num_tokens + and gate.stride(1) == 1 + and 0 <= gate_col0 + and gate_col0 + total_heads * kernel.V_DIM <= gate.shape[1] + ): + raise ValueError( + f"k3_mla_attn: gate {tuple(gate.shape)} {gate.dtype} col0 {gate_col0} does not fit the v_b output" + ) + groups = total_heads // kernel.HEADS + ws_o = _attn_workspace(q.device, groups) + # More 16-CTA clusters than co-reside would run a second wave: launch without a cluster instead when every CTA fits + # on the SMs at once (one per SM; its waits are spins). + clusters = num_requests * groups + no_cluster = ( + clusters > kernel.CLUSTER_WAVE + and clusters * kernel.CLUSTER <= torch.cuda.get_device_properties(q.device).multi_processor_count + ) + if out is None: + out = torch.empty(num_tokens, total_heads * width, dtype=torch.bfloat16, device=q.device) + elif not ( + out.is_contiguous() + and out.dtype == torch.bfloat16 + and out.numel() == num_tokens * total_heads * width + ): + raise ValueError( + f"k3_mla_attn: output {tuple(out.shape)} {out.dtype} is not a dense [M, heads * {width}] bf16" + ) + total_rows = pool.numel() // row_stride + gate_flat = gate.as_strided((gate.numel(),), (1,)) if apply_gate else q.view(-1) + gate_ld = gate.stride(0) if apply_gate else 0 + args = (_arg(q.view(-1)), _arg(_pool_base(pool, row_stride)), _arg(page_rows, 4), _arg(seq_len, 4), _arg(ws_o), + _arg(out.view(-1)), _arg((w_vb if fuse_vb else q).view(-1)), _arg(gate_flat)) # fmt: skip + stream = cuda_driver.CUstream(torch.cuda.current_stream(q.device).cuda_stream) + use_pdl = _use_pdl() + scale_log2 = float(softmax_scale) * kernel.LOG2E + key = ("k3_mla_attn", row_stride, total_heads, fuse_vb, apply_gate, use_pdl, no_cluster) + fn = _compiled.get(key) + if fn is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "trtllm::k3_mla_attn must run once outside CUDA-graph capture first (it compiles its kernel)." + ) + import cutlass.cute as cute + + with _lock: + fn = _compiled.get(key) + if fn is None: + fn = _compiled[key] = cute.compile( + kernel.k3_mla_attn, *args, tokens, num_requests, pt_stride, scale_log2, total_rows, + int(page_offset), int(gate_col0), int(gate_ld), row_stride, total_heads, fuse_vb, apply_gate, + use_pdl, stream, no_cluster, + ) # fmt: skip + fn( + *args, + tokens, + num_requests, + pt_stride, + scale_log2, + total_rows, + int(page_offset), + int(gate_col0), + int(gate_ld), + stream, + ) + return out + + +@torch.library.custom_op("trtllm::k3_mla_attn", mutates_args=()) +def k3_mla_attn( + q: torch.Tensor, + pool: torch.Tensor, + row_stride: int, + page_table: torch.Tensor, + seq_len: torch.Tensor, + softmax_scale: float, +) -> torch.Tensor: + """MLA decode attention of R <= 8 requests of T <= 8 tokens: ``q`` [M = R T, heads * 576] (``fused_q``, heads a + multiple of 6, request-major) against the paged latent cache ``pool`` (flat bf16; row i of page p at ``(p * 64 + + i) * row_stride``, 512 latent then 64 rope columns), request i's pages ``page_table[i]`` and length ``seq_len[i]`` + = L_i (rows including its T new ones; see the module docstring), causal bottom-right (token t of request i sees + rows <= L_i - T + t). Returns ``[M, heads * 512]`` bf16. The page table and lengths are read before the grid + dependency wait (they must be written before the CUDA graph runs); q and the pages holding rows >= L_i - T after + it. Request i's rows are computed as the R = 1 call on its own rows, pages and length would compute them.""" + return _launch_attn(q, pool, row_stride, page_table, seq_len, softmax_scale) + + +@torch.library.custom_op("trtllm::k3_mla_attn_out", mutates_args=("out",)) +def k3_mla_attn_out( + q: torch.Tensor, + pool: torch.Tensor, + row_stride: int, + page_table: torch.Tensor, + page_offset: int, + seq_len: torch.Tensor, + softmax_scale: float, + out: torch.Tensor, +) -> None: + """``k3_mla_attn`` into ``out`` (dense [M, heads * 512] bf16), with ``page_offset`` added to every page-table + entry (the layer's slot in a layer-interleaved pool).""" + _launch_attn( + q, pool, row_stride, page_table, seq_len, softmax_scale, out=out, page_offset=page_offset + ) + + +@torch.library.custom_op("trtllm::k3_mla_attn_vb_out", mutates_args=("out",)) +def k3_mla_attn_vb_out( + q: torch.Tensor, + pool: torch.Tensor, + row_stride: int, + page_table: torch.Tensor, + page_offset: int, + seq_len: torch.Tensor, + softmax_scale: float, + w_vb: torch.Tensor, + out: torch.Tensor, + gate: Optional[torch.Tensor] = None, + gate_col0: int = 0, +) -> None: + """``k3_mla_attn`` with v_b applied in the same launch: ``out`` [M, heads * 128] = per head + ``bf16(bf16(o) @ w_vb[h]^T)`` for the attention output o, ``w_vb`` = v_b_proj [heads, 128, 512] bf16. With ``gate`` + (bf16 [M, C], sigmoid of head h's gate at columns ``gate_col0 + 128 h``) the output is ``bf16(y * s)``, the + unfused output gate.""" + _launch_attn( + q, pool, row_stride, page_table, seq_len, softmax_scale, out=out, page_offset=page_offset, w_vb=w_vb, + gate=gate, gate_col0=gate_col0, + ) # fmt: skip + + +@k3_mla_attn.register_fake +def _(q, pool, row_stride, page_table, seq_len, softmax_scale): + return q.new_empty((q.shape[0], q.shape[1] // 576 * 512), dtype=torch.bfloat16) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_kda_attn.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_kda_attn.py new file mode 100644 index 000000000000..61f0debe682c --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_kda_attn.py @@ -0,0 +1,322 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``trtllm::k3_kda_attn`` (Kimi K3's fused KDA projection and speculative verify of one request's 8 tokens in one +launch, the DSpark batch-1 path) at the per-rank TP16 shape (W [3208, 7168], 6 heads, K = V = 128, conv width 4). + +The reference is the same x through ``trtllm::k3_kda_qkvg`` (the projection stream alone, the same stream clusters; +its Lamport buffers decoded into the ``[q | k | v | og | f_a | b]`` rows) and ``trtllm::k3_kda_verify`` on those rows, +on a copy of the same pools: the same arithmetic, so outputs, conv caches, pool state and the drafts' records are +compared bit for bit. A request keeps its slot for several rounds, so its pending count (the drafts the previous +round accepted) runs through 0..7, then the next request starts on another slot; slots outside the batch stay +untouched; pools dense and strided. Also: CUDA-graph replays with rewritten inputs, the launch counter near 2^31 +(buffers inside guard bands: nothing written outside them, the same bits as a run from zero), and the slot and the +pending counts given as slices of longer index tensors at any element offset. + + pytest test_k3_kda_attn.py +""" + +import pytest +import torch + +H = 6 +K = V = 128 +W = 4 +HK = H * K +PROJ = 4 * HK + K + H + 2 # 3208 +K_IN = 7168 +NUM_SPEC = 7 +NT = NUM_SPEC + 1 +LOWER_BOUND = -5.0 +EPS = 1e-5 +SCALE = K**-0.5 +POOL = 11 +SLOT_ROUNDS = 9 # pending 0, then 1, 6, 3, 0, 5, 2, 7, 4 (pending_schedule) + + +def _sm100() -> bool: + return torch.cuda.is_available() and torch.cuda.get_device_capability()[0] == 10 + + +pytestmark = pytest.mark.skipif(not _sm100(), reason="needs SM100 (tcgen05, TMA, clusters)") + + +def _ops(): + import tensorrt_llm # noqa: F401 (registers the trtllm ops) + from tensorrt_llm._torch.cute_dsl_kernels.k3_kda_attn import op # noqa: F401 + from tensorrt_llm._torch.cute_dsl_kernels.k3_kda_verify import op as _verify # noqa: F401 + + return op + + +def make_weights(seed: int) -> dict: + g = torch.Generator(device="cuda").manual_seed(seed) + + def rnd(*s, scale=1.0): + return torch.randn(*s, generator=g, device="cuda") * scale + + return { + "w": rnd(PROJ, K_IN, scale=0.02).bfloat16(), "w_fb": (rnd(HK, K) * 0.05).bfloat16(), + "w_q": rnd(HK, W, scale=0.3), "w_k": rnd(HK, W, scale=0.3), "w_v": rnd(HK, W, scale=0.3), + "a_log": rnd(H, scale=0.5), "dt_bias": rnd(HK, scale=0.5), "onorm_w": (1 + 0.1 * rnd(V)).float(), + } # fmt: skip + + +def make_pools(seed: int, layout: str) -> dict: + """Conv caches [pool, HK, W - 1 + NUM_SPEC] (dim-contiguous), the SSM state [pool, H, V, K] (dense, or each slot + followed by the cache manager's conv bytes), the drafts' records and the pending counts.""" + g = torch.Generator(device="cuda").manual_seed(seed) + s = W - 1 + NUM_SPEC + p = {name: (torch.randn(POOL, s, HK, generator=g, device="cuda") * 0.5).transpose(1, 2) + for name in ("cs_q", "cs_k", "cs_v")} # fmt: skip + state = torch.randn(POOL, H, V, K, generator=g, device="cuda") * 0.05 + if layout == "strided": + buf = torch.zeros(POOL, state[0].numel() + 3 * HK * (W - 1), device="cuda") + view = buf[:, : state[0].numel()].view(state.shape) + view.copy_(state) + state = view + p["state"] = state + p["state_tok"] = torch.zeros(POOL, NUM_SPEC, H, V, K, device="cuda") + p["pending"] = torch.zeros(POOL, dtype=torch.int32, device="cuda") + return p + + +def clone_pools(p: dict) -> dict: + out = {} + for name, t in p.items(): + if name.startswith("cs_"): + out[name] = t.transpose(1, 2).clone().transpose(1, 2) + elif not t.is_contiguous(): + buf = torch.zeros(t.shape[0], t.stride(0), dtype=t.dtype, device=t.device) + view = buf[:, : t[0].numel()].view(t.shape) + view.copy_(t) + out[name] = view + else: + out[name] = t.clone() + return out + + +POOL_NAMES = ("cs_q", "cs_k", "cs_v", "state", "state_tok") + + +def pending_schedule(rnd: int) -> int: + return (5 * rnd + 1) % (NUM_SPEC + 1) + + +class Fused: + """k3_kda_attn on its own pools and Lamport buffers.""" + + def __init__(self, wt, pools, bufs=None): + op = _ops() + self.wt, self.p = wt, pools + self.bufs = ( + bufs if bufs is not None else op.make_buffers(torch.device("cuda"), op.FUSED_CTAS) + ) + + def __call__(self, x, slot): + wt, p = self.wt, self.p + return torch.ops.trtllm.k3_kda_attn( + x, wt["w"], wt["w_fb"], wt["w_q"], wt["w_k"], wt["w_v"], wt["a_log"], wt["dt_bias"], wt["onorm_w"], + p["cs_q"], p["cs_k"], p["cs_v"], p["state"], p["state_tok"], slot, p["pending"], *self.bufs, NUM_SPEC, + LOWER_BOUND, SCALE, EPS, + ) # fmt: skip + + +class Unfused: + """k3_kda_qkvg's rows (its Lamport buffers decoded), then k3_kda_verify on them.""" + + def __init__(self, wt, pools): + op = _ops() + self.wt, self.p = wt, pools + self.bufs = op.make_buffers(torch.device("cuda"), op.CTAS) + + def rows(self, x): + p1, part, epoch = self.bufs + buf = int(epoch[0].item()) % 3 + torch.ops.trtllm.k3_kda_qkvg(x, self.wt["w"], p1, part, epoch) + qkfa = p1.view(3, 8, 2 * HK + K)[buf].view(torch.bfloat16) + parts = part.view(3, 3, 2, 8, HK)[buf].view(torch.float32) + v, og, b = ((parts[r, 0] + parts[r, 1]).bfloat16() for r in range(3)) + rows = torch.zeros(NT, PROJ, dtype=torch.bfloat16, device="cuda") + rows[:, : 2 * HK] = qkfa[:, : 2 * HK] + rows[:, 2 * HK : 3 * HK] = v + rows[:, 3 * HK : 4 * HK] = og + rows[:, 4 * HK : 4 * HK + K] = qkfa[:, 2 * HK :] + rows[:, 4 * HK + K : 4 * HK + K + H] = b[:, :H] + return rows + + def __call__(self, x, slot): + wt, p = self.wt, self.p + return torch.ops.trtllm.k3_kda_verify( + self.rows(x), wt["w_fb"], wt["w_q"], wt["w_k"], wt["w_v"], wt["a_log"], wt["dt_bias"], wt["onorm_w"], + p["cs_q"], p["cs_k"], p["cs_v"], p["state"], p["state_tok"], slot, p["pending"], NUM_SPEC, LOWER_BOUND, + SCALE, EPS, None, + ) # fmt: skip + + +def same_pools(a: dict, b: dict) -> bool: + return all(torch.equal(a[n], b[n]) for n in POOL_NAMES) + + +def others_untouched(p: dict, before: dict, slot: int) -> bool: + others = [s for s in range(POOL) if s != slot] + return all(torch.equal(p[n][others], before[n][others]) for n in POOL_NAMES) + + +def slot_order(seed: int): + return torch.randperm(POOL, generator=torch.Generator().manual_seed(seed)).tolist() + + +@pytest.mark.parametrize("layout", ["dense", "strided"]) +def test_rounds(layout): + """Two requests, one after the other, each on its slot for SLOT_ROUNDS rounds: bits of the unfused path.""" + with torch.inference_mode(): + wt = make_weights(100) + pools = make_pools(200, layout) + fused, ref = Fused(wt, clone_pools(pools)), Unfused(wt, clone_pools(pools)) + gen = torch.Generator(device="cuda").manual_seed(400) + order = slot_order(300) + for rnd in range(2 * SLOT_ROUNDS): + s = order[rnd // SLOT_ROUNDS] + slot = torch.tensor([s], dtype=torch.int32, device="cuda") + x = torch.randn(NT, K_IN, generator=gen, device="cuda").bfloat16() + before = {n: fused.p[n].clone() for n in POOL_NAMES} + got, want = fused(x, slot), ref(x, slot) + torch.cuda.synchronize() + pend = int(fused.p["pending"][s]) + assert torch.equal(got, want), (rnd, pend) + assert bool(torch.isfinite(got.float()).all()), (rnd, pend) + assert same_pools(fused.p, ref.p), (rnd, pend) + assert others_untouched(fused.p, before, s), (rnd, pend) + for p in (fused.p, ref.p): + p["pending"][s] = pending_schedule(rnd) + + +def test_graph_replay(): + """One captured call replayed with x, the slot and the pending counts rewritten in place: the bits of eager calls on + a copy of the pools.""" + op = _ops() + with torch.inference_mode(): + wt = make_weights(110) + pools = make_pools(210, "dense") + eager, graphed = Fused(wt, clone_pools(pools)), Fused(wt, clone_pools(pools)) + gen = torch.Generator(device="cuda").manual_seed(410) + order = slot_order(310) + x_in = torch.zeros(NT, K_IN, dtype=torch.bfloat16, device="cuda") + s_in = torch.zeros(1, dtype=torch.int32, device="cuda") + # Compiles outside capture, on scratch pools and buffers. + Fused(wt, clone_pools(pools), op.make_buffers(torch.device("cuda"), op.FUSED_CTAS))( + x_in, s_in + ) + torch.cuda.synchronize() + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + graph = torch.cuda.CUDAGraph() + with torch.cuda.stream(stream), torch.cuda.graph(graph, stream=stream): + y = graphed(x_in, s_in) + torch.cuda.current_stream().wait_stream(stream) + for rnd in range(2 * SLOT_ROUNDS): + s = order[rnd // SLOT_ROUNDS] + slot = torch.tensor([s], dtype=torch.int32, device="cuda") + x = torch.randn(NT, K_IN, generator=gen, device="cuda").bfloat16() + x_in.copy_(x) + s_in.copy_(slot) + graph.replay() + want = eager(x, slot) + torch.cuda.synchronize() + assert torch.equal(y, want), rnd + assert same_pools(graphed.p, eager.p), rnd + for p in (graphed.p, eager.p): + p["pending"][s] = pending_schedule(rnd) + + +def _alone(path, x, slot): + """One launch, complete before the next: the head CTAs read the slot's pools before their grid-dependency wait, + so a launch right behind another on the same slot could read them mid-update (the model never runs one layer's + KDA twice in a row; its next launch on these pools is a step later).""" + out = path(x, slot) + torch.cuda.synchronize() + return out + + +def _guarded(numel: int, dtype, fill: int): + """``numel`` all-ones words with a whole set of the op's buffers of ``fill`` words on each side.""" + guard = numel + 4096 + big = torch.full((guard + numel + guard,), fill, dtype=dtype, device="cuda") + big[guard : guard + numel] = -1 + return big, big[guard : guard + numel], guard + + +def test_epoch_wrap(): + """Every CTA's counter preset to 2^31 - 2, as after 2^31 launches on a device (the buffers are shared by every KDA + layer): four launches write nothing outside the buffers, give the bits of a run from zero, and keep the counter + in 0..2.""" + op = _ops() + with torch.inference_mode(): + wt = make_weights(120) + pools = make_pools(220, "dense") + gen = torch.Generator(device="cuda").manual_seed(420) + xs = [torch.randn(NT, K_IN, generator=gen, device="cuda").bfloat16() for _ in range(4)] + slot = torch.tensor([3], dtype=torch.int32, device="cuda") + fresh = Fused(wt, clone_pools(pools)) + want = [_alone(fresh, x, slot) for x in xs] + big1, p1, g1 = _guarded(op.P1_NUMEL, torch.int16, 0x1234) + big2, part, g2 = _guarded(op.PART_NUMEL, torch.int32, 0x12345678) + epoch = torch.full((op.FUSED_CTAS,), 2**31 - 2, dtype=torch.int32, device="cuda") + wrapped = Fused(wt, clone_pools(pools), (p1, part, epoch)) + got = [_alone(wrapped, x, slot) for x in xs] + for big, guard, numel, fill in ( + (big1, g1, op.P1_NUMEL, 0x1234), + (big2, g2, op.PART_NUMEL, 0x12345678), + ): + assert bool((big[:guard] == fill).all()), "stores before the buffers" + assert bool((big[guard + numel :] == fill).all()), "stores after the buffers" + assert all(torch.equal(a, b) for a, b in zip(got, want)) + assert same_pools(wrapped.p, fresh.p) + assert bool(((epoch >= 0) & (epoch < 3)).all()), epoch.unique().tolist() + + +def _at_offset(t: torch.Tensor, offset: int) -> torch.Tensor: + """``t``'s values in a longer tensor, starting at element ``offset`` (a slice such as + ``state_indices[num_prefills:]``: its data pointer is aligned to the element only).""" + buf = torch.zeros(offset + t.numel(), dtype=t.dtype, device=t.device) + buf[offset:] = t + view = buf[offset:] + assert view.data_ptr() % 16 == offset * t.element_size() % 16 + return view + + +@pytest.mark.parametrize("offset", [0, 1, 2, 3]) +@pytest.mark.parametrize("which", ["slot", "pending"]) +def test_index_offset(which, offset): + """The slot or the pending counts as a slice of a longer index tensor that starts at element ``offset``: over two + rounds, the second on a nonzero pending count, the outputs and the pools bit for bit those of the same indices in + tensors of their own.""" + with torch.inference_mode(): + wt = make_weights(130) + pools = make_pools(230, "dense") + gen = torch.Generator(device="cuda").manual_seed(430) + s = 5 + slot = torch.tensor([s], dtype=torch.int32, device="cuda") + ref, got = Fused(wt, clone_pools(pools)), Fused(wt, clone_pools(pools)) + got_slot = _at_offset(slot, offset) if which == "slot" else slot + if which == "pending": + got.p["pending"] = _at_offset(got.p["pending"], offset) + for rnd in range(2): + x = torch.randn(NT, K_IN, generator=gen, device="cuda").bfloat16() + want, out = _alone(ref, x, slot), _alone(got, x, got_slot) + assert torch.equal(out, want), rnd + assert same_pools(got.p, ref.p), rnd + for p in (ref.p, got.p): + p["pending"][s] = pending_schedule(rnd) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_kda_decode_attn.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_kda_decode_attn.py new file mode 100644 index 000000000000..5aad370eabd3 --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_kda_decode_attn.py @@ -0,0 +1,523 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``trtllm::k3_kda_decode_attn`` (Kimi K3's fused KDA projection + plain decode of R requests of one token) at the +per-rank TP16 shape (6 heads, K = V = 128, conv width 4), against ``trtllm::kda_decode``. + +The reference is the model's plain-decode path fed the very projection rows the fused kernel computes: the stream +alone (``trtllm::k3_kda_qkvg``, the same stream clusters) on the same x, its Lamport buffers decoded into the +``[q | k | v | og | f_a | b]`` rows, f_b as a bf16 ``F.linear``, then ``trtllm::kda_decode`` (native, indexed state +pool, packed conv pool updated in place). A float64 torch decode on the same rows bounds both. + +Checks over rounds of distinct slots from the same initial pools: every request's output against kda_decode and +float64 (fp32 tolerance), the state rows against both, the conv pool against kda_decode bit for bit (raw inputs), +slots outside the batch untouched; pools dense and with strided slots (the cache manager's interleaving); slots +given as a slice of a longer index tensor at any element offset; the +schedule twice bit-identical; CUDA-graph replays with rewritten inputs; the launch counter near 2^31 (nothing +written outside the buffers); launches interleaved with ``trtllm::k3_kda_attn`` on one shared buffer set, as the +model runs them. + + pytest test_k3_kda_decode_attn.py + python3 test_k3_kda_decode_attn.py report +""" + +import sys + +import pytest +import torch +import torch.nn.functional as F + +H = 6 +K = V = 128 +HK = H * K +W = 4 +PROJ = 3208 +K_IN = 7168 +LOWER_BOUND = -5.0 +EPS = 1e-5 +SCALE = K**-0.5 +POOL = 11 +ROUNDS = 4 +NUM_SPEC = 7 # trtllm::k3_kda_attn verifies one request's NUM_SPEC + 1 tokens per launch +TOL_OUT = 2e-2 +TOL_STATE = 1e-3 +CONV_PAD = 128 # extra bf16 per conv slot in the strided layout (a multiple of 8) +SSM_PAD = ( + 3 * HK * (W - 1) +) # extra fp32 per state slot in the strided layout (the manager's conv bytes) + + +def _sm100() -> bool: + return torch.cuda.is_available() and torch.cuda.get_device_capability()[0] == 10 + + +pytestmark = pytest.mark.skipif(not _sm100(), reason="needs SM100 (tcgen05, TMA, clusters)") + + +def _ops(): + import tensorrt_llm # noqa: F401 (registers the trtllm ops) + from tensorrt_llm._torch.cute_dsl_kernels.k3_kda_attn import op # noqa: F401 + + +def make_weights(seed: int) -> dict: + g = torch.Generator(device="cuda").manual_seed(seed) + + def rnd(*s, scale=1.0): + return torch.randn(*s, generator=g, device="cuda") * scale + + conv = ( + rnd(3, HK, W, scale=0.3).bfloat16().float() + ) # bf16 values: the native path reads them as bf16 + return { + "w": rnd(PROJ, K_IN, scale=0.02).bfloat16(), "w_fb": rnd(HK, K, scale=0.05).bfloat16(), + "w_q": conv[0].contiguous(), "w_k": conv[1].contiguous(), "w_v": conv[2].contiguous(), + "w_t": [conv[i].t().bfloat16().contiguous() for i in range(3)], + "a_log": rnd(H, scale=0.5), "dt_bias": rnd(HK, scale=0.5), "onorm_w": (1 + 0.1 * rnd(V)).float(), + } # fmt: skip + + +def make_pools(seed: int, layout: str) -> dict: + """conv bf16 [POOL, 3 HK, W - 1] and the state fp32 [POOL, H, V, K]; "strided": each slot padded.""" + g = torch.Generator(device="cuda").manual_seed(seed) + conv = (torch.randn(POOL, 3 * HK, W - 1, generator=g, device="cuda") * 0.5).bfloat16() + state = torch.randn(POOL, H, V, K, generator=g, device="cuda") * 0.05 + if layout == "strided": + conv_buf = torch.zeros( + POOL, 3 * HK * (W - 1) + CONV_PAD, dtype=torch.bfloat16, device="cuda" + ) + conv_view = conv_buf[:, : 3 * HK * (W - 1)].view(POOL, 3 * HK, W - 1) + conv_view.copy_(conv) + conv = conv_view + st_buf = torch.zeros(POOL, H * V * K + SSM_PAD, device="cuda") + st_view = st_buf[:, : H * V * K].view(POOL, H, V, K) + st_view.copy_(state) + state = st_view + return {"conv": conv, "state": state} + + +def clone_pools(p: dict) -> dict: + out = {} + for name, t in p.items(): + if t.is_contiguous(): + out[name] = t.clone() + else: + buf = torch.zeros(t.shape[0], t.stride(0), dtype=t.dtype, device=t.device) + view = buf[:, : t[0].numel()].view(t.shape) + view.copy_(t) + out[name] = view + return out + + +def make_slots(num_requests: int, seed: int) -> torch.Tensor: + perm = torch.randperm(POOL, generator=torch.Generator().manual_seed(seed)) + return perm[:num_requests].to(torch.int32).cuda() + + +def rel(a: torch.Tensor, b: torch.Tensor) -> float: + return ((a.float() - b.float()).abs().max() / b.float().abs().max().clamp_min(1e-6)).item() + + +class Rows: + """The projection rows the fused kernel computes for x: the stream alone, its Lamport buffers decoded.""" + + def __init__(self, wt): + from tensorrt_llm._torch.cute_dsl_kernels.k3_kda_attn import op + + self.wt = wt + self.bufs = op.make_buffers(torch.device("cuda"), op.CTAS) + + def __call__(self, x: torch.Tensor) -> torch.Tensor: + p1, part, epoch = self.bufs + buf = int(epoch[0].item()) % 3 + torch.ops.trtllm.k3_kda_qkvg(x, self.wt["w"], p1, part, epoch) + n = x.shape[0] + qkfa = p1.view(3, 8, 2 * HK + K)[buf, :n].view(torch.bfloat16) + parts = part.view(3, 3, 2, 8, HK)[buf].view(torch.float32) + v, og, b = ((parts[r, 0, :n] + parts[r, 1, :n]).bfloat16() for r in range(3)) + rows = torch.zeros(n, PROJ, dtype=torch.bfloat16, device="cuda") + rows[:, : 2 * HK] = qkfa[:, : 2 * HK] + rows[:, 2 * HK : 3 * HK] = v + rows[:, 3 * HK : 4 * HK] = og + rows[:, 4 * HK : 4 * HK + K] = qkfa[:, 2 * HK :] + rows[:, 4 * HK + K : 4 * HK + K + H] = b[:, :H] + return rows + + +class FusedPath: + def __init__(self, wt, pools): + from tensorrt_llm._torch.cute_dsl_kernels.k3_kda_attn import op + + self.wt, self.p = wt, pools + self.bufs = op.make_buffers(torch.device("cuda"), op.FUSED_CTAS) + + def __call__(self, x, slots): + wt, p = self.wt, self.p + return torch.ops.trtllm.k3_kda_decode_attn( + x, wt["w"], wt["w_fb"], wt["w_q"], wt["w_k"], wt["w_v"], wt["a_log"], wt["dt_bias"], wt["onorm_w"], + p["conv"], p["state"], slots, *self.bufs, LOWER_BOUND, SCALE, EPS, + ) # fmt: skip + + +class NativePath: + """The model's plain decode (``forward_decode``): f_b as a bf16 F.linear, then trtllm::kda_decode.""" + + def __init__(self, wt, pools): + self.wt, self.p = wt, pools + + def __call__(self, rows, slots): + from tensorrt_llm._torch.modules.kimi_kda._kda_decode import run_kda_decode_fusion_cuda + + wt, p = self.wt, self.p + n = rows.shape[0] + + def heads(cols): + return cols.unflatten(-1, (H, K)).unsqueeze(0) + + g = F.linear(rows[:, 4 * HK : 4 * HK + K], wt["w_fb"]) + out = torch.empty(n, 1, H, V, dtype=torch.bfloat16, device="cuda") + conv = p["conv"] + run_kda_decode_fusion_cuda( + x_q=heads(rows[:, :HK]), x_k=heads(rows[:, HK : 2 * HK]), x_v=heads(rows[:, 2 * HK : 3 * HK]), + w_q_t=wt["w_t"][0], w_k_t=wt["w_t"][1], w_v_t=wt["w_t"][2], bias_q=None, bias_k=None, bias_v=None, + cs_q=conv[:, :HK], cs_k=conv[:, HK : 2 * HK], cs_v=conv[:, 2 * HK :], A_log=wt["a_log"], + g=heads(g), dt_bias=wt["dt_bias"], beta=rows[:, 4 * HK + K : 4 * HK + K + H].unsqueeze(0), + state=p["state"], onorm_g=heads(rows[:, 3 * HK : 4 * HK]), onorm_weight=wt["onorm_w"], out=out, + ssm_state_indices=slots, scale=SCALE, onorm_eps=EPS, lower_bound=LOWER_BOUND, + use_beta_sigmoid_in_kernel=True, update_conv_cache=True, + ) # fmt: skip + return out.view(n, H, V) + + +class F64Path: + """The plain KDA decode in float64 torch on the same rows (f_b in float64, rounded to bf16 as the GEMM's + output): conv4 + SiLU, q / k L2 norm (q scaled), beta sigmoid, the lower-bound gate, + S <- S d + beta (v - (S d) k) k^T, o = S q, the gated RMSNorm.""" + + def __init__(self, wt, pools): + self.wt = wt + self.conv = pools["conv"].double().clone() + self.state = pools["state"].double().clone() + + def __call__(self, rows, slots): + wt = self.wt + r = rows.double() + n = rows.shape[0] + conv_w = torch.stack([wt["w_q"], wt["w_k"], wt["w_v"]]).double() # [3, HK, W] + g = (r[:, 4 * HK : 4 * HK + K] @ wt["w_fb"].double().t()).bfloat16().double() + outs = torch.empty(n, H, V, dtype=torch.float64, device="cuda") + for i, s in enumerate(slots.tolist()): + new = r[i, : 3 * HK].view(3, HK) + win = self.conv[s].view(3, HK, W - 1) + u = torch.cat([win, new.unsqueeze(-1)], dim=-1) # [3, HK, W], oldest first + act = (u * conv_w).sum(-1) + act = act * torch.sigmoid(act) + self.conv[s] = u[:, :, 1:].reshape(3 * HK, W - 1) + q, k, v = (act[j].view(H, K) for j in range(3)) + q = q / torch.sqrt((q * q).sum(-1, keepdim=True) + 1e-6) * SCALE + k = k / torch.sqrt((k * k).sum(-1, keepdim=True) + 1e-6) + beta = torch.sigmoid(r[i, 4 * HK + K : 4 * HK + K + H]) + xg = torch.exp(wt["a_log"].double()).unsqueeze(-1) * ( + g[i].view(H, K) + wt["dt_bias"].double().view(H, K) + ) + decay = torch.exp(LOWER_BOUND * torch.sigmoid(xg)) + st = self.state[s] * decay.unsqueeze(1) # [H, V, K] * decay per key + res = (v - (st * k.unsqueeze(1)).sum(-1)) * beta.unsqueeze(-1) # [H, V] + st = st + res.unsqueeze(-1) * k.unsqueeze(1) + self.state[s] = st + o = (st * q.unsqueeze(1)).sum(-1) # [H, V] + rms = torch.rsqrt((o * o).mean(-1, keepdim=True) + EPS) + gate = torch.sigmoid(r[i, 3 * HK : 4 * HK].view(H, V)) + outs[i] = o * rms * wt["onorm_w"].double() * gate + return outs + + +def run_schedule(num_requests, layout="dense", seed=0, rounds=ROUNDS): + _ops() + wt = make_weights(100 + seed) + pools = make_pools(200 + seed, layout) + fused = FusedPath(wt, clone_pools(pools)) + native = NativePath(wt, clone_pools(pools)) + ref = F64Path(wt, pools) + rows_of = Rows(wt) + gen = torch.Generator(device="cuda").manual_seed(400 + seed) + metrics, outs = [], [] + for rnd in range(rounds): + slots = make_slots(num_requests, 300 + seed + rnd) + x = torch.randn(num_requests, K_IN, generator=gen, device="cuda").bfloat16() + state_before = fused.p["state"].clone() + conv_before = fused.p["conv"].clone() + out_f = fused(x, slots) + rows = rows_of(x) + out_n = native(rows, slots) + out_r = ref(rows, slots) + torch.cuda.synchronize() + idx = slots.long() + others = torch.ones(POOL, dtype=torch.bool, device="cuda") + others[idx] = False + m = dict(round=rnd) + m["out_vs_native"] = rel(out_f, out_n) + m["out_vs_f64"] = rel(out_f, out_r) + m["native_vs_f64"] = rel(out_n, out_r) + m["state_vs_f64"] = max(rel(fused.p["state"][s], ref.state[s]) for s in idx.tolist()) + m["state_vs_native"] = max( + rel(fused.p["state"][s], native.p["state"][s]) for s in idx.tolist() + ) + m["conv_bits_vs_native"] = torch.equal(fused.p["conv"], native.p["conv"]) + m["others_untouched"] = torch.equal( + fused.p["state"][others], state_before[others] + ) and torch.equal(fused.p["conv"][others], conv_before[others]) + metrics.append(m) + outs.append(out_f.clone()) + return metrics, outs + + +def round_ok(m) -> bool: + return (m["out_vs_native"] <= TOL_OUT and m["out_vs_f64"] <= TOL_OUT and m["state_vs_f64"] <= TOL_STATE + and m["state_vs_native"] <= TOL_STATE and m["conv_bits_vs_native"] and m["others_untouched"]) # fmt: skip + + +@pytest.mark.parametrize("num_requests", range(1, 9)) +@pytest.mark.parametrize("layout", ["dense", "strided"]) +def test_split(num_requests, layout): + with torch.inference_mode(): + metrics, _ = run_schedule(num_requests, layout) + bad = [m for m in metrics if not round_ok(m)] + assert not bad, bad + + +@pytest.mark.parametrize("num_requests", [1, 8]) +def test_deterministic(num_requests): + with torch.inference_mode(): + _, a = run_schedule(num_requests, seed=5) + _, b = run_schedule(num_requests, seed=5) + assert all(torch.equal(x, y) for x, y in zip(a, b)) + + +@pytest.mark.parametrize("num_requests", [1, 3, 8]) +def test_graph_replay(num_requests): + """Rounds captured once in a CUDA graph (static x and slots rewritten in place before every replay) against the + same rounds run eagerly from the same pools.""" + _ops() + with torch.inference_mode(): + wt = make_weights(7) + pools = make_pools(8, "dense") + eager = FusedPath(wt, clone_pools(pools)) + graphed = FusedPath(wt, clone_pools(pools)) + gen = torch.Generator(device="cuda").manual_seed(9) + xs = [ + torch.randn(num_requests, K_IN, generator=gen, device="cuda").bfloat16() + for _ in range(ROUNDS) + ] + slots = [make_slots(num_requests, 10 + r) for r in range(ROUNDS)] + want = [eager(x, s).clone() for x, s in zip(xs, slots)] + x_in = xs[0].clone() + s_in = slots[0].clone() + graphed(x_in, s_in) # compiles; this launch's result is discarded with the pools below + graphed.p = clone_pools(pools) + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + graph = torch.cuda.CUDAGraph() + with torch.cuda.stream(stream), torch.cuda.graph(graph, stream=stream): + y = graphed(x_in, s_in) + torch.cuda.current_stream().wait_stream(stream) + got = [] + for x, s in zip(xs, slots): + x_in.copy_(x) + s_in.copy_(s) + graph.replay() + got.append(y.clone()) + torch.cuda.synchronize() + assert all(torch.equal(a, b) for a, b in zip(want, got)) + assert torch.equal(eager.p["state"], graphed.p["state"]) and torch.equal( + eager.p["conv"], graphed.p["conv"] + ) + + +def _alone(path, x, slots): + """One launch, complete before the next: the head CTAs read the slots' pools before their grid-dependency wait, + so a launch right behind another on the same slots could read them mid-update (the model never runs one layer's + KDA twice in a row; its next launch on these pools is a step later).""" + out = path(x, slots) + torch.cuda.synchronize() + return out + + +def test_epoch_wrap(): + """Every CTA's counter preset to 2^31 - 2, as after 2^31 launches on a device (the buffers are shared by every KDA + layer): four launches write nothing outside the buffers (each inside guard bands), give the bits of a run from + zero, and keep the counter in 0..2.""" + _ops() + from tensorrt_llm._torch.cute_dsl_kernels.k3_kda_attn import op + + with torch.inference_mode(): + wt = make_weights(13) + pools = make_pools(14, "dense") + gen = torch.Generator(device="cuda").manual_seed(15) + xs = [torch.randn(4, K_IN, generator=gen, device="cuda").bfloat16() for _ in range(4)] + slots = make_slots(4, 16) + fresh = FusedPath(wt, clone_pools(pools)) + want = [_alone(fresh, x, slots) for x in xs] + bands = [] + for numel, dtype, fill in ( + (op.P1_NUMEL, torch.int16, 0x1234), + (op.PART_NUMEL, torch.int32, 0x12345678), + ): + guard = numel + 4096 # a whole set of buffers on each side + big = torch.full((guard + numel + guard,), fill, dtype=dtype, device="cuda") + big[guard : guard + numel] = -1 + bands.append((big, guard, numel, fill)) + epoch = torch.full((op.FUSED_CTAS,), 2**31 - 2, dtype=torch.int32, device="cuda") + wrapped = FusedPath(wt, clone_pools(pools)) + wrapped.bufs = tuple(big[guard : guard + numel] for big, guard, numel, _ in bands) + ( + epoch, + ) + got = [_alone(wrapped, x, slots) for x in xs] + for big, guard, numel, fill in bands: + assert bool((big[:guard] == fill).all()), "stores before the buffers" + assert bool((big[guard + numel :] == fill).all()), "stores after the buffers" + assert all(torch.equal(a, b) for a, b in zip(got, want)) + assert torch.equal(wrapped.p["state"], fresh.p["state"]) and torch.equal( + wrapped.p["conv"], fresh.p["conv"] + ) + assert bool(((epoch >= 0) & (epoch < 3)).all()), epoch.unique().tolist() + + +@pytest.mark.parametrize("offset", [0, 1, 2, 3]) +def test_slots_offset(offset): + """``slots`` as a slice of a longer index tensor that starts at element ``offset`` (the mixer passes + ``state_indices[num_prefills:]`` in a step with prefills, so the data pointer is only 4-byte aligned): the output + and the pools bit for bit those of the same slots in a tensor of their own.""" + _ops() + with torch.inference_mode(): + wt = make_weights(17) + pools = make_pools(18, "dense") + gen = torch.Generator(device="cuda").manual_seed(19) + x = torch.randn(3, K_IN, generator=gen, device="cuda").bfloat16() + slots = make_slots(3, 20) + idx = torch.zeros(offset + slots.numel(), dtype=torch.int32, device="cuda") + idx[offset:] = slots + view = idx[offset:] + ref = FusedPath(wt, clone_pools(pools)) + got = FusedPath(wt, clone_pools(pools)) + want = _alone(ref, x, slots) + out = _alone(got, x, view) + assert view.data_ptr() % 16 == 4 * offset % 16 + assert torch.equal(out, want) + assert torch.equal(got.p["state"], ref.p["state"]) and torch.equal(got.p["conv"], ref.p["conv"]) + + +def make_verify_pools(seed: int) -> dict: + """``trtllm::k3_kda_attn``'s pools for a layer (test_k3_kda_attn.py's dense layout): the conv caches + [POOL, HK, W - 1 + NUM_SPEC] (dim-contiguous), the SSM state, the drafts' records and the pending counts.""" + g = torch.Generator(device="cuda").manual_seed(seed) + p = {name: (torch.randn(POOL, W - 1 + NUM_SPEC, HK, generator=g, device="cuda") * 0.5).transpose(1, 2) + for name in ("cs_q", "cs_k", "cs_v")} # fmt: skip + p["state"] = torch.randn(POOL, H, V, K, generator=g, device="cuda") * 0.05 + p["state_tok"] = torch.zeros(POOL, NUM_SPEC, H, V, K, device="cuda") + p["pending"] = torch.zeros(POOL, dtype=torch.int32, device="cuda") + return p + + +def clone_verify_pools(p: dict) -> dict: + return { + name: t.transpose(1, 2).clone().transpose(1, 2) if name.startswith("cs_") else t.clone() + for name, t in p.items() + } + + +class VerifyPath: + """``trtllm::k3_kda_attn``: the same projection and the speculative verify of one request's NUM_SPEC + 1 tokens.""" + + def __init__(self, wt, pools, bufs): + self.wt, self.p, self.bufs = wt, pools, bufs + + def __call__(self, x, slot): + wt, p = self.wt, self.p + return torch.ops.trtllm.k3_kda_attn( + x, wt["w"], wt["w_fb"], wt["w_q"], wt["w_k"], wt["w_v"], wt["a_log"], wt["dt_bias"], wt["onorm_w"], + p["cs_q"], p["cs_k"], p["cs_v"], p["state"], p["state_tok"], slot, p["pending"], *self.bufs, NUM_SPEC, + LOWER_BOUND, SCALE, EPS, + ) # fmt: skip + + +def test_shared_buffers_with_verify(): + """The model keeps one (p1, part, epoch) set per device for this op and ``trtllm::k3_kda_attn`` + (kimi_kda_mixer.py): plain-decode and verify launches interleaved in stream order on one set give the bits of the + same launches on a set of their own, one at a time. Both run the same stream role (all 8 token rows published, + the same words re-armed) and advance every CTA's index once per launch, so either may follow the other.""" + _ops() + from tensorrt_llm._torch.cute_dsl_kernels.k3_kda_attn import op + + dev = torch.device("cuda") + with torch.inference_mode(): + wt = make_weights(23) + pools = make_pools(24, "dense") + vpools = make_verify_pools(25) + gen = torch.Generator(device="cuda").manual_seed(26) + calls = [] + for i, r in enumerate((1, 4, 8, 3)): + x = torch.randn(r, K_IN, generator=gen, device="cuda").bfloat16() + calls.append(("decode", x, make_slots(r, 27 + i))) + x = torch.randn(NUM_SPEC + 1, K_IN, generator=gen, device="cuda").bfloat16() + calls.append(("verify", x, torch.tensor([i], dtype=torch.int32, device="cuda"))) + own = { + "decode": FusedPath(wt, clone_pools(pools)), + "verify": VerifyPath( + wt, clone_verify_pools(vpools), op.make_buffers(dev, op.FUSED_CTAS) + ), + } + want = [_alone(own[kind], x, s) for kind, x, s in calls] + shared_bufs = op.make_buffers(dev, op.FUSED_CTAS) + shared = { + "decode": FusedPath(wt, clone_pools(pools)), + "verify": VerifyPath(wt, clone_verify_pools(vpools), shared_bufs), + } + shared["decode"].bufs = shared_bufs + got = [shared[kind](x, s) for kind, x, s in calls] + torch.cuda.synchronize() + assert all(torch.equal(a, b) for a, b in zip(got, want)) + for kind, path in shared.items(): + assert all(torch.equal(t, own[kind].p[name]) for name, t in path.p.items()), kind + epoch = shared_bufs[2] + assert bool((epoch == epoch[0]).all()), epoch.unique().tolist() + + +def report() -> int: + """Per split and layout: the worst round's errors.""" + print(f"{torch.cuda.get_device_name()}; {ROUNDS} rounds per split; rel = max |a - b| / max |b|") + print("| R | layout | out vs kda_decode | out vs f64 | kda_decode vs f64 | state vs f64 | state vs kda_decode | " + "conv bits | others untouched | result |") # fmt: skip + print("| --: | :-- | --: | --: | --: | --: | --: | :-- | :-- | :-- |") + ok_all = True + with torch.inference_mode(): + for layout in ("dense", "strided"): + for r in range(1, 9): + metrics, _ = run_schedule(r, layout) + ok = all(round_ok(m) for m in metrics) + ok_all &= ok + worst = {k: max(m[k] for m in metrics) for k in ("out_vs_native", "out_vs_f64", "native_vs_f64", + "state_vs_f64", "state_vs_native")} # fmt: skip + conv_ok = all(m["conv_bits_vs_native"] for m in metrics) + others_ok = all(m["others_untouched"] for m in metrics) + print(f"| {r} | {layout} | {worst['out_vs_native']:.2e} | {worst['out_vs_f64']:.2e} | " + f"{worst['native_vs_f64']:.2e} | {worst['state_vs_f64']:.2e} | {worst['state_vs_native']:.2e} | " + f"{conv_ok} | {others_ok} | {'PASS' if ok else 'FAIL'} |", flush=True) # fmt: skip + print("ALL PASS" if ok_all else "FAIL") + return 0 if ok_all else 1 + + +if __name__ == "__main__": + if len(sys.argv) > 1 and sys.argv[1] == "report": + sys.exit(report()) + sys.exit(pytest.main([__file__, "-q", "-p", "no:cacheprovider", *sys.argv[1:]])) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_kda_pools_past_2g.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_kda_pools_past_2g.py new file mode 100644 index 000000000000..c9884b6c384b --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_kda_pools_past_2g.py @@ -0,0 +1,212 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""The fused KDA ops (``trtllm::k3_kda_decode_attn``, ``trtllm::k3_kda_attn``, ``trtllm::k3_kda_verify``) on pools +whose last slot starts past element 2^31, at the per-rank TP16 shape (6 heads, K = V = 128, conv width 4). + +The Mamba cache manager coalesces a rank's KDA layers inside each slot, so a layer's state pool is a view at slot +stride layers x 98,304 fp32 elements (its conv pool at layers x 6,912 bf16): with enough slots the pools reach past +2^31 elements. Here each pool is the last layer's view of such a pool ([slots, layers, ...] underneath), and the +per-token verify states are dense at 7 x 98,304 elements per slot; 3,122 slots put the last slot of each past 2^31. + +Each op runs once on requests in the pools' first, middle and last slots (``k3_kda_attn``: one request, the last slot) +and once on small pools holding copies of those slots; the outputs and every slot the call updates must agree bit for +bit. About 19 GB of device memory. + + pytest test_k3_kda_pools_past_2g.py +""" + +import pytest +import torch + +H = 6 +K = V = 128 +HK = H * K +W = 4 +PROJ = 4 * HK + K + H + 2 # the fused [q | k | v | onorm gate | f_a | b | pad] row (3208 columns) +K_IN = 7168 +NUM_SPEC = 7 +NT = NUM_SPEC + 1 +LOWER_BOUND = -5.0 +EPS = 1e-5 +SCALE = K**-0.5 +SLOTS = 3122 +SSM_LAYERS = 8 +CONV_LAYERS = 100 +TWO_G = 2**31 + + +def _sm100() -> bool: + return torch.cuda.is_available() and torch.cuda.get_device_capability()[0] == 10 + + +pytestmark = pytest.mark.skipif(not _sm100(), reason="needs SM100 (tcgen05, TMA, clusters)") + + +def _ops(): + import tensorrt_llm # noqa: F401 (registers the trtllm ops) + from tensorrt_llm._torch.cute_dsl_kernels.k3_kda_attn import op + from tensorrt_llm._torch.cute_dsl_kernels.k3_kda_verify import op as _verify # noqa: F401 + + return op + + +def make_weights(seed: int) -> dict: + g = torch.Generator(device="cuda").manual_seed(seed) + + def rnd(*s, scale=1.0): + return torch.randn(*s, generator=g, device="cuda") * scale + + return { + "w": rnd(PROJ, K_IN, scale=0.02).bfloat16(), "w_fb": rnd(HK, K, scale=0.05).bfloat16(), + "w_q": rnd(HK, W, scale=0.3), "w_k": rnd(HK, W, scale=0.3), "w_v": rnd(HK, W, scale=0.3), + "a_log": rnd(H, scale=0.5), "dt_bias": rnd(HK, scale=0.5), "onorm_w": (1 + 0.1 * rnd(V)).float(), + } # fmt: skip + + +def layer_view(layers: int, shape: tuple, dtype: torch.dtype) -> torch.Tensor: + """The last layer's view of a pool that keeps ``layers`` layers' states in each of its slots.""" + return torch.zeros(SLOTS, layers, *shape, dtype=dtype, device="cuda")[:, -1] + + +def last_slot_offset(t: torch.Tensor) -> int: + return (t.shape[0] - 1) * t.stride(0) + + +def verify_pools() -> dict: + """Conv caches [slots, HK, W - 1 + NUM_SPEC] (dim-contiguous), the SSM state, the drafts' records, the pending + counts.""" + p = { + name: torch.zeros(SLOTS, W - 1 + NUM_SPEC, HK, device="cuda").transpose(1, 2) + for name in ("cs_q", "cs_k", "cs_v") + } + p["ssm"] = layer_view(SSM_LAYERS, (H, V, K), torch.float32) + p["state_tok"] = torch.zeros(SLOTS, NUM_SPEC, H, V, K, device="cuda") + p["pending"] = torch.zeros(SLOTS, dtype=torch.int32, device="cuda") + assert last_slot_offset(p["ssm"]) >= TWO_G and last_slot_offset(p["state_tok"]) >= TWO_G + return p + + +def fill_verify_slots(p: dict, used: list, pending: list, seed: int) -> dict: + """Random contents in the used slots; returns small pools holding copies of them, in the same layouts.""" + g = torch.Generator(device="cuda").manual_seed(seed) + for i, s in enumerate(used): + for name in ("cs_q", "cs_k", "cs_v"): + p[name][s] = torch.randn(HK, W - 1 + NUM_SPEC, generator=g, device="cuda") * 0.5 + p["ssm"][s] = torch.randn(H, V, K, generator=g, device="cuda") * 0.05 + p["state_tok"][s] = torch.rand(NUM_SPEC, H, V, K, generator=g, device="cuda") * 0.5 + p["pending"][s] = pending[i] + small = { + name: torch.stack([p[name][s].t() for s in used]).transpose(1, 2) + for name in ("cs_q", "cs_k", "cs_v") + } + for name in ("ssm", "state_tok", "pending"): + small[name] = p[name][used] + return small + + +def same(a: torch.Tensor, b: torch.Tensor) -> bool: + """Bit-equal (NaN-safe).""" + bits = torch.int16 if a.element_size() == 2 else torch.int32 + return torch.equal(a.contiguous().view(bits), b.contiguous().view(bits)) + + +def check_slots(p: dict, small: dict, used: list, names: tuple) -> None: + for i, s in enumerate(used): + for name in names: + assert same(p[name][s], small[name][i]), f"slot {s} {name}" + + +def test_decode_attn(): + op = _ops() + with torch.inference_mode(): + wt = make_weights(1) + p = { + "ssm": layer_view(SSM_LAYERS, (H, V, K), torch.float32), + "conv": layer_view(CONV_LAYERS, (3 * HK, W - 1), torch.bfloat16), + } + assert last_slot_offset(p["ssm"]) >= TWO_G and last_slot_offset(p["conv"]) >= TWO_G + used = [0, SLOTS // 2, SLOTS - 1] + g = torch.Generator(device="cuda").manual_seed(2) + for s in used: + p["ssm"][s] = torch.randn(H, V, K, generator=g, device="cuda") * 0.05 + p["conv"][s] = (torch.randn(3 * HK, W - 1, generator=g, device="cuda") * 0.5).bfloat16() + small = {name: t[used] for name, t in p.items()} + x = torch.randn(len(used), K_IN, generator=g, device="cuda").bfloat16() + + def call(q: dict, slots: list) -> torch.Tensor: + bufs = op.make_buffers(torch.device("cuda"), op.FUSED_CTAS) + return torch.ops.trtllm.k3_kda_decode_attn( + x, wt["w"], wt["w_fb"], wt["w_q"], wt["w_k"], wt["w_v"], wt["a_log"], wt["dt_bias"], wt["onorm_w"], + q["conv"], q["ssm"], torch.tensor(slots, dtype=torch.int32, device="cuda"), *bufs, LOWER_BOUND, + SCALE, EPS, + ) # fmt: skip + + want = call(small, list(range(len(used)))) + got = call(p, used) + torch.cuda.synchronize() + assert same(got, want) + check_slots(p, small, used, ("ssm", "conv")) + + +def test_attn(): + op = _ops() + with torch.inference_mode(): + wt = make_weights(3) + p = verify_pools() + used = [SLOTS - 1] + small = fill_verify_slots(p, used, [5], 4) + x = torch.randn( + NT, K_IN, generator=torch.Generator(device="cuda").manual_seed(5), device="cuda" + ).bfloat16() + + def call(q: dict, slot: int) -> torch.Tensor: + bufs = op.make_buffers(torch.device("cuda"), op.FUSED_CTAS) + return torch.ops.trtllm.k3_kda_attn( + x, wt["w"], wt["w_fb"], wt["w_q"], wt["w_k"], wt["w_v"], wt["a_log"], wt["dt_bias"], wt["onorm_w"], + q["cs_q"], q["cs_k"], q["cs_v"], q["ssm"], q["state_tok"], + torch.tensor([slot], dtype=torch.int32, device="cuda"), q["pending"], *bufs, NUM_SPEC, LOWER_BOUND, + SCALE, EPS, + ) # fmt: skip + + want = call(small, 0) + got = call(p, used[0]) + torch.cuda.synchronize() + assert same(got, want) + check_slots(p, small, used, ("cs_q", "cs_k", "cs_v", "ssm", "state_tok")) + + +def test_verify(): + _ops() + with torch.inference_mode(): + wt = make_weights(6) + p = verify_pools() + used = [0, SLOTS // 2, SLOTS - 1] + small = fill_verify_slots(p, used, [0, 3, 7], 7) + g = torch.Generator(device="cuda").manual_seed(8) + proj = torch.randn(len(used) * NT, PROJ, generator=g, device="cuda").bfloat16() + + def call(q: dict, slots: list) -> torch.Tensor: + return torch.ops.trtllm.k3_kda_verify( + proj, wt["w_fb"], wt["w_q"], wt["w_k"], wt["w_v"], wt["a_log"], wt["dt_bias"], wt["onorm_w"], + q["cs_q"], q["cs_k"], q["cs_v"], q["ssm"], q["state_tok"], + torch.tensor(slots, dtype=torch.int32, device="cuda"), q["pending"], NUM_SPEC, LOWER_BOUND, SCALE, + EPS, + ) # fmt: skip + + want = call(small, list(range(len(used)))) + got = call(p, used) + torch.cuda.synchronize() + assert same(got, want) + check_slots(p, small, used, ("cs_q", "cs_k", "cs_v", "ssm", "state_tok")) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_kda_verify.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_kda_verify.py new file mode 100644 index 000000000000..d7136753192e --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_kda_verify.py @@ -0,0 +1,588 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``trtllm::k3_kda_verify`` (Kimi K3's KDA speculative verify with a committed state per verify token) for R requests +of T = 1 + num_spec tokens, at the per-rank TP16 shape (6 heads, K = V = 128, conv width 4). + +Two references, run over the same schedule of rounds from the same initial pools, with distinct slots and a pending +(accepted-draft) count per request that changes every round: + +* main's path, as ``KimiDeltaAttention.forward_verify_fused`` calls it: f_b as a bf16 ``F.linear``, then + ``trtllm::kda_mtp_decode`` replaying each request's accepted drafts from its replay caches + (``cu_seqlens[n] = n T - pending[n]``, ``num_accepted_tokens = pending``), then the gated RMSNorm + (``rms_norm_gated_token_major``, sigmoid gate). The report says which ``kda_mtp_decode`` was loaded. +* a float64 torch delta rule (the "fp32 reference"; float64 so that no TF32 enters) that keeps each request's + committed token history: conv4 + SiLU, q / k L2 norm, beta sigmoid, the lower-bound gate, + S <- S d + beta (v - (S d) k) k^T, o = S q, the gated RMSNorm. + +Checks per round: the outputs of every request against both references (fp32 tolerance), the committed pool state +(after each golden token) against the history, the conv caches against main's (raw inputs: bit-exact); with the gate +taken from the unfused f_b output (``g_ext``) the pool state against main's bit for bit. The drafts' records in +``state_tok`` are checked through the next round, which starts from the accepted ones. Also: requests isolated, the +schedule twice bit-identical, CUDA-graph replays with rewritten inputs, the slots and the pending counts given as +slices of longer index tensors at any element offset. + + pytest test_k3_kda_verify.py + python3 test_k3_kda_verify.py report | time +""" + +import inspect +import statistics +import sys + +import pytest +import torch +import torch.nn.functional as F + +H = 6 +K = V = 128 +W = 4 +HK = H * K +PROJ = 4 * HK + K + H + 2 # the fused [q | k | v | onorm gate | f_a | b | pad] row (3208 columns) +LOWER_BOUND = -5.0 +EPS = 1e-5 +SCALE = K**-0.5 +POOL = 11 +ROUNDS = 6 +TOL_OUT = 2e-2 # bf16 outputs; main rounds the core to bf16 before its norm +# fp32 states against the fp32 history: approximate exp / rcp in the kernels, and the bf16 f_b gate (one bf16 ulp +# moves a decay by up to ~1 %; main's kda_mtp_decode measures 2.05e-4 at 6x8) +TOL_STATE = 1e-3 + +SPLITS = [(r, t) for t in (8, 4, 2) for r in range(1, 9)] + + +def _sm100() -> bool: + return torch.cuda.is_available() and torch.cuda.get_device_capability()[0] == 10 + + +pytestmark = pytest.mark.skipif(not _sm100(), reason="needs SM100 (tcgen05, TMA, clusters)") + + +def _ops(): + import tensorrt_llm # noqa: F401 (registers the trtllm ops) + import tensorrt_llm._torch.custom_ops.cute_dsl_kimi_k3_kda_mtp_ops # noqa: F401 (kda_mtp_decode) + from tensorrt_llm._torch.cute_dsl_kernels.k3_kda_verify import op # noqa: F401 (k3_kda_verify) + + +def main_kda_variant() -> str: + """'main' when the loaded kda_mtp_decode has main's signature, else 'port' (a modified one).""" + _ops() + from tensorrt_llm._torch.custom_ops import cute_dsl_kimi_k3_kda_mtp_ops as mtp + + return ( + "port" + if "accepted_by_slot" in inspect.signature(mtp.kda_mtp_decode_impl).parameters + else "main" + ) + + +def make_weights(seed: int) -> dict: + g = torch.Generator(device="cuda").manual_seed(seed) + + def rnd(*s, scale=1.0): + return torch.randn(*s, generator=g, device="cuda") * scale + + return { + "w_q": rnd(HK, W, scale=0.3), "w_k": rnd(HK, W, scale=0.3), "w_v": rnd(HK, W, scale=0.3), + "a_log": rnd(H, scale=0.5), "dt_bias": rnd(HK, scale=0.5), "onorm_w": (1 + 0.1 * rnd(V)).float(), + "w_fb": (rnd(HK, K) * 0.05).bfloat16(), + } # fmt: skip + + +def strided(state: torch.Tensor) -> torch.Tensor: + """The same values in the Mamba cache manager's layout: each slot's SSM state followed by its conv state bytes.""" + pool = torch.zeros(state.shape[0], state[0].numel() + 3 * HK * (W - 1), device=state.device) + view = pool[:, : state[0].numel()].view(state.shape) + view.copy_(state) + return view + + +def make_pools(seed: int, num_spec: int, layout: str) -> dict: + """The initial pools: conv caches [pool, HK, W - 1 + num_spec] (dim-contiguous), the SSM state [pool, H, V, K].""" + g = torch.Generator(device="cuda").manual_seed(seed) + s = W - 1 + num_spec + p = {name: (torch.randn(POOL, s, HK, generator=g, device="cuda") * 0.5).transpose(1, 2) + for name in ("cs_q", "cs_k", "cs_v")} # fmt: skip + p["state"] = torch.randn(POOL, H, V, K, generator=g, device="cuda") * 0.05 + if layout == "strided": + p["state"] = strided(p["state"]) + return p + + +def clone_pools(p: dict) -> dict: + out = {} + for name, t in p.items(): + if name.startswith("cs_"): + out[name] = ( + t.transpose(1, 2).clone().transpose(1, 2) + ) # a copy in the same dim-contiguous layout + elif name == "state" and not t.is_contiguous(): + out[name] = strided(t) + else: + out[name] = t.clone() + return out + + +def split_proj(proj: torch.Tensor): + t = proj.shape[0] + x_q, x_k, x_v, og = (proj[:, i * HK : (i + 1) * HK] for i in range(4)) + f_a = proj[:, 4 * HK : 4 * HK + K] + beta = proj[:, 4 * HK + K : 4 * HK + K + H] + return t, x_q, x_k, x_v, og, f_a, beta + + +class MainPath: + """main's verify (forward_verify_fused): its replay caches, a per-request pending count.""" + + def __init__(self, wt, pools, slots, num_spec): + self.wt, self.p, self.slots, self.num_spec = wt, pools, slots, num_spec + self.p["qkg_cache"] = torch.zeros(POOL, num_spec, 3, HK, device="cuda") + self.p["v_cache"] = torch.zeros(POOL, num_spec, HK, device="cuda") + self.p["beta_cache"] = torch.zeros(POOL, num_spec, H, device="cuda") + + def __call__(self, proj, pending_req): + from tensorrt_llm._torch.modules.mamba.layernorm_gated import rms_norm_gated_token_major + + t, x_q, x_k, x_v, og, f_a, beta = split_proj(proj) + n, steps = self.slots.numel(), self.num_spec + 1 + g = F.linear(f_a, self.wt["w_fb"]) + cu = torch.arange(0, (n + 1) * steps, steps, dtype=torch.int32, device="cuda") + cu[:n].sub_(pending_req) + o = torch.ops.trtllm.kda_mtp_decode( + x_q=x_q.view(1, t, H, K), x_k=x_k.view(1, t, H, K), x_v=x_v.view(1, t, H, V), w_q=self.wt["w_q"], + w_k=self.wt["w_k"], w_v=self.wt["w_v"], cs_q=self.p["cs_q"], cs_k=self.p["cs_k"], cs_v=self.p["cs_v"], + g=g.view(1, t, H, K), beta=beta.contiguous().view(1, t, H), A_log=self.wt["a_log"], + dt_bias=self.wt["dt_bias"], recurrent_state=self.p["state"], qkg_cache=self.p["qkg_cache"], + v_cache=self.p["v_cache"], beta_cache=self.p["beta_cache"], ssm_state_indices=self.slots, cu_seqlens=cu, + num_spec=self.num_spec, num_accepted_tokens=pending_req, lower_bound=LOWER_BOUND, scale=SCALE, + ) # fmt: skip + core = rms_norm_gated_token_major(o.reshape(-1, V), og.reshape(t, H, V), self.wt["onorm_w"], EPS, + gate_activation="sigmoid") # fmt: skip + return core.view(t, H, V), g + + +class K3Path: + """k3_kda_verify: a committed state per verify token, a per-slot pending count.""" + + def __init__(self, wt, pools, slots, num_spec): + self.wt, self.p, self.slots, self.num_spec = wt, pools, slots, num_spec + self.p["state_tok"] = torch.zeros(POOL, num_spec, H, V, K, device="cuda") + self.p["pending"] = torch.zeros(POOL, dtype=torch.int32, device="cuda") + + def __call__(self, proj, g_ext=None): + p = self.p + return torch.ops.trtllm.k3_kda_verify( + proj, self.wt["w_fb"], self.wt["w_q"], self.wt["w_k"], self.wt["w_v"], self.wt["a_log"], + self.wt["dt_bias"], self.wt["onorm_w"], p["cs_q"], p["cs_k"], p["cs_v"], p["state"], p["state_tok"], + self.slots, p["pending"], self.num_spec, LOWER_BOUND, SCALE, EPS, g_ext, + ) # fmt: skip + + +def conv_silu(win, raw, c, w): + """Channel group ``c`` (q, k, v) of the causal conv over the window (oldest first) and the new raw input, SiLU.""" + x = win[0][c] * w[:, 0] + win[1][c] * w[:, 1] + win[2][c] * w[:, 2] + raw[c] * w[:, 3] + return x * torch.sigmoid(x) + + +class Fp32Path: + """The delta rule over each request's committed history (raw conv inputs and the state after the last committed + token), in float64: no TF32 even where cuBLAS is told to use it.""" + + def __init__(self, wt, pools, slots, num_spec): + self.wt, self.num_spec = wt, num_spec + self.slots = slots.tolist() + # Committed raw inputs, oldest first: the conv caches' window columns 0..2 (q, k, v). + self.seq = [[torch.stack([pools[c][s, :, i].double() for c in ("cs_q", "cs_k", "cs_v")]) for i in range(3)] + for s in self.slots] # fmt: skip + self.state = [pools["state"][s].double().clone() for s in self.slots] + self.last = None + + def __call__(self, proj): + """Outputs [T, H, V] and the per-token states of every request.""" + t_total, x_q, x_k, x_v, og, f_a, beta_raw = split_proj(proj) + steps = self.num_spec + 1 + # f_b's output is bf16 in the model. + g_all = F.linear(f_a.double(), self.wt["w_fb"].double()).bfloat16().double() + wq, wk, wv = (self.wt[n].double() for n in ("w_q", "w_k", "w_v")) + exp_a = self.wt["a_log"].double().exp() + dt_bias = self.wt["dt_bias"].double().view(H, K) + onorm_w = self.wt["onorm_w"].double() + out = torch.empty(t_total, H, V, dtype=torch.float64, device="cuda") + self.last = [] + for n in range(len(self.slots)): + seq = list(self.seq[n]) + s_cur = self.state[n] + states, raws = [], [] + for t in range(steps): + row = n * steps + t + raw = torch.stack([x_q[row].double(), x_k[row].double(), x_v[row].double()]) + win = seq[-3:] + q = conv_silu(win, raw, 0, wq).view(H, K) + k = conv_silu(win, raw, 1, wk).view(H, K) + v = conv_silu(win, raw, 2, wv).view(H, V) + q = q * torch.rsqrt((q * q).sum(-1, keepdim=True) + 1e-6) * SCALE + k = k * torch.rsqrt((k * k).sum(-1, keepdim=True) + 1e-6) + beta = torch.sigmoid(beta_raw[row].double()) + gk = LOWER_BOUND * torch.sigmoid(exp_a[:, None] * (g_all[row].view(H, K) + dt_bias)) + decay = gk.exp() + sd = s_cur * decay[:, None, :] + vn = v - torch.einsum("hvk,hk->hv", sd, k) + s_cur = sd + beta[:, None, None] * vn[:, :, None] * k[:, None, :] + o = torch.einsum("hvk,hk->hv", s_cur, q) + rms = torch.rsqrt((o * o).mean(-1, keepdim=True) + EPS) + out[row] = o * rms * onorm_w * torch.sigmoid(og[row].double().view(H, V)) + states.append(s_cur) + raws.append(raw) + seq.append(raw) + self.last.append((states, raws)) + return out + + def commit(self, pending): + """The sampler accepted ``pending[n]`` drafts of the last round: the golden token and those drafts commit.""" + for n, p in enumerate(pending): + states, raws = self.last[n] + self.seq[n] = (self.seq[n] + raws[: p + 1])[-3:] + self.state[n] = states[p] + + +def make_slots(num_requests: int, seed: int) -> torch.Tensor: + perm = torch.randperm(POOL, generator=torch.Generator().manual_seed(seed)) + return perm[:num_requests].to(torch.int32).cuda() + + +def pending_schedule(num_requests: int, num_spec: int, rnd: int): + """Accepted drafts per request after round ``rnd``: every count 0..num_spec, different per request.""" + return [(3 * n + 5 * rnd + 1 + (rnd * n) % 3) % (num_spec + 1) for n in range(num_requests)] + + +def rel(a: torch.Tensor, b: torch.Tensor) -> float: + return ((a.float() - b.float()).abs().max() / b.float().abs().max().clamp_min(1e-6)).item() + + +def run_schedule(num_requests, steps, fold=True, layout="dense", seed=0, rounds=ROUNDS): + """Main's path, k3_kda_verify and the fp32 history over ``rounds`` rounds; per-round metrics.""" + _ops() + num_spec = steps - 1 + wt = make_weights(100 + seed) + pools = make_pools(200 + seed, num_spec, layout) + slots = make_slots(num_requests, 300 + seed) + main = MainPath(wt, clone_pools(pools), slots, num_spec) + k3 = K3Path(wt, clone_pools(pools), slots, num_spec) + ref = Fp32Path(wt, pools, slots, num_spec) + pending_req = torch.zeros(num_requests, dtype=torch.int32, device="cuda") + gen = torch.Generator(device="cuda").manual_seed(400 + seed) + rows, outs = [], [] + for rnd in range(rounds): + proj = torch.randn(num_requests * steps, PROJ, generator=gen, device="cuda").bfloat16() + out_main, g = main(proj, pending_req) + out_k3 = k3(proj, None if fold else g.contiguous()) + out_ref = ref(proj) + torch.cuda.synchronize() + r = dict(round=rnd, pending=pending_req.tolist()) + r["out_vs_main"] = rel(out_k3, out_main) + r["out_vs_fp32"] = rel(out_k3, out_ref) + r["main_vs_fp32"] = rel(out_main, out_ref) + # The pool state after the golden token (state_tok holds the drafts' compact records, which the next round's + # outputs check against the history). + st_err, main_st_err, st_bits = 0.0, 0.0, True + for n, s in enumerate(slots.tolist()): + states, _ = ref.last[n] + st_err = max(st_err, rel(k3.p["state"][s], states[0])) + main_st_err = max(main_st_err, rel(main.p["state"][s], states[0])) + st_bits &= torch.equal(k3.p["state"][s], main.p["state"][s]) + r["state_vs_fp32"] = st_err + r["main_state_vs_fp32"] = main_st_err + r["state_bits_vs_main"] = st_bits + r["conv_bits_vs_main"] = all( + torch.equal(k3.p[c], main.p[c]) for c in ("cs_q", "cs_k", "cs_v") + ) + rows.append(r) + outs.append(out_k3.clone()) + nxt = pending_schedule(num_requests, num_spec, rnd) + pending_req.copy_(torch.tensor(nxt, dtype=torch.int32)) + k3.p["pending"][slots.long()] = pending_req + ref.commit(nxt) + return rows, outs + + +def round_ok(r, fold: bool) -> bool: + ok = (r["out_vs_main"] <= TOL_OUT and r["out_vs_fp32"] <= TOL_OUT and r["state_vs_fp32"] <= TOL_STATE + and r["conv_bits_vs_main"]) # fmt: skip + if not fold: # the same gate as main's: the recurrence's arithmetic is kda_mtp_decode's + ok &= r["state_bits_vs_main"] + return ok + + +@pytest.mark.parametrize("num_requests,steps", SPLITS, ids=[f"{r}x{t}" for r, t in SPLITS]) +def test_split(num_requests, steps): + with torch.inference_mode(): + rows, outs = run_schedule(num_requests, steps, fold=True) + again, outs2 = run_schedule(num_requests, steps, fold=True) + for r in rows: + assert round_ok(r, True), r + assert all(torch.equal(a, b) for a, b in zip(outs, outs2)), "the schedule twice differs" + + +@pytest.mark.parametrize("num_requests,steps", [(1, 8), (3, 4), (8, 8), (8, 2)]) +@pytest.mark.parametrize("layout", ["dense", "strided"]) +def test_g_ext(num_requests, steps, layout): + """The gate from the unfused f_b output (main's bf16 F.linear): the pool state bit-exact against main's.""" + with torch.inference_mode(): + rows, _ = run_schedule(num_requests, steps, fold=False, layout=layout, seed=1) + for r in rows: + assert round_ok(r, False), r + + +@pytest.mark.parametrize("num_requests,steps", [(4, 8), (8, 4)]) +def test_strided_fold(num_requests, steps): + with torch.inference_mode(): + rows, _ = run_schedule(num_requests, steps, fold=True, layout="strided", seed=2) + for r in rows: + assert round_ok(r, True), r + + +@pytest.mark.parametrize("num_requests,steps", [(4, 8), (8, 2)]) +def test_isolation(num_requests, steps): + """A change in request 0's rows changes only request 0's outputs, pool state and per-token states.""" + _ops() + with torch.inference_mode(): + num_spec = steps - 1 + wt = make_weights(7) + pools = make_pools(8, num_spec, "dense") + slots = make_slots(num_requests, 9) + a = K3Path(wt, clone_pools(pools), slots, num_spec) + b = K3Path(wt, clone_pools(pools), slots, num_spec) + pend = torch.tensor( + pending_schedule(num_requests, num_spec, 3), dtype=torch.int32, device="cuda" + ) + for path in (a, b): + path.p["pending"][slots.long()] = pend + proj = torch.randn(num_requests * steps, PROJ, generator=torch.Generator(device="cuda").manual_seed(10), + device="cuda").bfloat16() # fmt: skip + proj_c = proj.clone() + proj_c[1, HK + 5] += ( + 1.0 # request 0, token 1, k channel 5: its outputs and states from token 1 on + ) + out_a, out_b = a(proj), b(proj_c) + torch.cuda.synchronize() + assert not torch.equal(out_a[:steps], out_b[:steps]) + assert torch.equal(out_a[steps:], out_b[steps:]) + s0 = int(slots[0]) + others = [int(s) for s in slots[1:]] + assert not torch.equal(a.p["state_tok"][s0], b.p["state_tok"][s0]) + for s in others: + assert torch.equal(a.p["state"][s], b.p["state"][s]) + assert torch.equal(a.p["state_tok"][s], b.p["state_tok"][s]) + + +@pytest.mark.parametrize("num_requests,steps", [(8, 8), (4, 2), (1, 8)]) +def test_graph_replay(num_requests, steps): + """One captured call replayed over the schedule with the projection rows and pending counts rewritten in place, + bit-identical to eager calls on a copy of the pools.""" + _ops() + with torch.inference_mode(): + num_spec = steps - 1 + wt = make_weights(11) + pools = make_pools(12, num_spec, "dense") + slots = make_slots(num_requests, 13) + graphed = K3Path(wt, clone_pools(pools), slots, num_spec) + eager = K3Path(wt, clone_pools(pools), slots, num_spec) + proj = torch.zeros(num_requests * steps, PROJ, dtype=torch.bfloat16, device="cuda") + holder = {} + snap = clone_pools( + {k: v for k, v in graphed.p.items() if k in ("cs_q", "cs_k", "cs_v", "state")} + ) + snap_tok = graphed.p["state_tok"].clone() + graphed(proj) # compiles outside capture + torch.cuda.synchronize() + for name in ("cs_q", "cs_k", "cs_v", "state"): + graphed.p[name].copy_(snap[name]) + graphed.p["state_tok"].copy_(snap_tok) + stream = torch.cuda.Stream() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + holder["out"] = graphed(proj) + gen = torch.Generator(device="cuda").manual_seed(14) + for rnd in range(ROUNDS): + proj.copy_(torch.randn(proj.shape, generator=gen, device="cuda").bfloat16()) + graph.replay() + want = eager(proj.clone()) + torch.cuda.synchronize() + assert torch.equal(holder["out"], want), f"round {rnd}" + for name in ("cs_q", "cs_k", "cs_v", "state", "state_tok"): + assert torch.equal(graphed.p[name], eager.p[name]), f"round {rnd} {name}" + nxt = torch.tensor( + pending_schedule(num_requests, num_spec, rnd), dtype=torch.int32, device="cuda" + ) + graphed.p["pending"][slots.long()] = nxt + eager.p["pending"][slots.long()] = nxt + + +def _at_offset(t: torch.Tensor, offset: int) -> torch.Tensor: + """``t``'s values in a longer tensor, starting at element ``offset`` (a slice such as + ``state_indices[num_contexts:]``: its data pointer is aligned to the element only).""" + buf = torch.zeros(offset + t.numel(), dtype=t.dtype, device=t.device) + buf[offset:] = t + view = buf[offset:] + assert view.data_ptr() % 16 == offset * t.element_size() % 16 + return view + + +@pytest.mark.parametrize("offset", [0, 1, 2, 3]) +@pytest.mark.parametrize("which", ["slots", "pending"]) +def test_index_offset(which, offset): + """The slots or the pending counts as a slice of a longer index tensor that starts at element ``offset``: over two + rounds, the second on nonzero pending counts, the outputs and the pools bit for bit those of the same indices in + tensors of their own.""" + _ops() + with torch.inference_mode(): + num_requests, steps = 3, 8 + num_spec = steps - 1 + wt = make_weights(15) + pools = make_pools(16, num_spec, "dense") + slots = make_slots(num_requests, 17) + ref = K3Path(wt, clone_pools(pools), slots, num_spec) + got = K3Path( + wt, + clone_pools(pools), + _at_offset(slots, offset) if which == "slots" else slots, + num_spec, + ) + if which == "pending": + got.p["pending"] = _at_offset(got.p["pending"], offset) + gen = torch.Generator(device="cuda").manual_seed(18) + for rnd in range(2): + proj = torch.randn(num_requests * steps, PROJ, generator=gen, device="cuda").bfloat16() + want, out = ref(proj), got(proj) + torch.cuda.synchronize() + assert torch.equal(out, want), rnd + for name in ("cs_q", "cs_k", "cs_v", "state", "state_tok"): + assert torch.equal(got.p[name], ref.p[name]), (rnd, name) + nxt = torch.tensor( + pending_schedule(num_requests, num_spec, rnd), dtype=torch.int32, device="cuda" + ) + ref.p["pending"][slots.long()] = nxt + got.p["pending"][slots.long()] = nxt + + +# ---------------------------------------------------------------------------------------------------------------- +# Report and timing (python3 test_k3_kda_verify.py report | time) +# ---------------------------------------------------------------------------------------------------------------- + + +def report() -> int: + with torch.inference_mode(): + print(f"{torch.cuda.get_device_name()}; kda_mtp_decode: {main_kda_variant()}") + print("| split | mode | layout | rounds | k3 vs main | k3 vs fp32 | main vs fp32 | state vs fp32 " + "| main state vs fp32 | state bits = main | conv bits = main | deterministic | result |") # fmt: skip + print("| :-- | :-- | :-- | --: | --: | --: | --: | --: | --: | :-- | :-- | :-- | :-- |") + ok_all = True + cases = [(r, t, True, "dense") for r, t in SPLITS] + cases += [ + (r, t, False, layout) + for r, t in [(1, 8), (3, 4), (8, 8), (8, 2)] + for layout in ("dense", "strided") + ] + cases += [(4, 8, True, "strided"), (8, 4, True, "strided")] + for r, t, fold, layout in cases: + rows, outs = run_schedule( + r, t, fold=fold, layout=layout, seed=0 if fold and layout == "dense" else 1 + ) + _, outs2 = run_schedule( + r, t, fold=fold, layout=layout, seed=0 if fold and layout == "dense" else 1 + ) + det = all(torch.equal(a, b) for a, b in zip(outs, outs2)) + ok = all(round_ok(x, fold) for x in rows) and det + ok_all &= ok + print(f"| {r}x{t} | {'fold' if fold else 'g_ext'} | {layout} | {len(rows)} | " + f"{max(x['out_vs_main'] for x in rows):.2e} | {max(x['out_vs_fp32'] for x in rows):.2e} | " + f"{max(x['main_vs_fp32'] for x in rows):.2e} | {max(x['state_vs_fp32'] for x in rows):.2e} | " + f"{max(x['main_state_vs_fp32'] for x in rows):.2e} | " + f"{all(x['state_bits_vs_main'] for x in rows)} | {all(x['conv_bits_vs_main'] for x in rows)} | " + f"{det} | {'PASS' if ok else 'FAIL'} |", flush=True) # fmt: skip + print("ALL PASS" if ok_all else "FAIL") + return 0 if ok_all else 1 + + +def time_graph(body, calls, replays=15): + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(stream): + body(0) + torch.cuda.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + for i in range(calls): + body(i) + torch.cuda.synchronize() + for _ in range(3): + graph.replay() + torch.cuda.synchronize() + per_call = [] + for _ in range(replays): + start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + start.record() + graph.replay() + end.record() + torch.cuda.synchronize() + per_call.append(start.elapsed_time(end) * 1e3 / calls) + return statistics.median(per_call), min(per_call), max(per_call) + + +def timing(layers: int = 48) -> None: + """Graphs of back-to-back calls over ``layers`` copies of the weights and pools (HBM-cold), steady pending.""" + _ops() + print(f"{torch.cuda.get_device_name()}; kda_mtp_decode: {main_kda_variant()}; graphs over {layers} layer copies, " + "15 replays: median (min-max) us per call") # fmt: skip + print("| split | main: f_b + kda_mtp_decode + gated norm | k3_kda_verify |") + print("| :-- | --: | --: |") + with torch.inference_mode(): + for r, t in SPLITS: + num_spec = t - 1 + slots = make_slots(r, 1) + mains, k3s = [], [] + for i in range(layers): + wt = make_weights(1000 + i) + pools = make_pools(2000 + i, num_spec, "dense") + mains.append(MainPath(wt, clone_pools(pools), slots, num_spec)) + k3s.append(K3Path(wt, pools, slots, num_spec)) + pend = torch.tensor(pending_schedule(r, num_spec, 2), dtype=torch.int32, device="cuda") + for k3 in k3s: + k3.p["pending"][slots.long()] = pend + proj = torch.randn(r * t, PROJ, device="cuda").bfloat16() + arms = [lambda i: mains[i % layers](proj, pend), lambda i: k3s[i % layers](proj)] + res = [[], []] + for rep in range(3): + for a in (0, 1) if rep % 2 == 0 else (1, 0): + res[a].append(time_graph(arms[a], layers)) + cells = [] + for a in range(2): + meds = sorted(x[0] for x in res[a]) + cells.append( + f"{meds[1]:.2f} ({min(x[1] for x in res[a]):.2f}-{max(x[2] for x in res[a]):.2f})" + ) + print(f"| {r}x{t} | " + " | ".join(cells) + " |", flush=True) + mains.clear() + k3s.clear() + torch.cuda.empty_cache() + + +if __name__ == "__main__": + if len(sys.argv) > 1 and sys.argv[1] == "report": + sys.exit(report()) + elif len(sys.argv) > 1 and sys.argv[1] == "time": + timing() + else: + sys.exit(pytest.main([__file__, "-q", "-p", "no:cacheprovider", *sys.argv[1:]])) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_attn.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_attn.py new file mode 100644 index 000000000000..2b2afe5e5bc5 --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_attn.py @@ -0,0 +1,471 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""trtllm::k3_mla_attn / k3_mla_attn_out / k3_mla_attn_vb_out (Kimi K3 MLA decode attention) at the TP16 shape (6 +heads per rank, latent 512 + rope 64, bf16 pool of 64-row pages) for decode steps of R requests x T tokens: against a +torch float64 reference with per-request bottom-right causal masks and against the stock CuTe DSL MLA decode +(trtllm::cute_dsl_mla_decode_fp16_blackwell); request i's rows bit-identical to the one-request call on its own +rows, pages and length; reruns bit-identical. + +Batch-1 identity against the unmodified kernel: set ``K3_BASE_TRTLLM`` to an unmodified ``tensorrt_llm`` package +directory (one request's outputs must match its single-request kernel bit for bit).""" + +import importlib.util +import math +import os + +import pytest +import torch + + +def _is_sm100() -> bool: + if not torch.cuda.is_available(): + return False + major, minor = torch.cuda.get_device_capability() + return major * 10 + minor in (100, 103) + + +pytestmark = pytest.mark.skipif(not _is_sm100(), reason="needs an SM100 / SM103 GPU") + +H, LAT, ROPE, PAGE, V = 6, 512, 64, 64, 128 +DQK = LAT + ROPE +SCALE = 1.0 / math.sqrt(128 + ROPE) +GATE_COL0 = 2112 + +# Every split of up to 8 tokens and the DSpark verify steps of R requests x 8 tokens. +SPLITS = [(1, 1), (2, 1), (3, 1), (4, 1), (5, 1), (8, 1), (1, 8), (2, 4), (4, 2)] + [ + (r, 8) for r in range(2, 9) +] +SPLIT_IDS = [f"{r}x{t}" for r, t in SPLITS] +# The TP4 shape (24 heads: 4 clusters per request, the workspace slots of several head groups) on a few splits. +ATTN_CASES = [(r, t, H) for r, t in SPLITS] + [ + (r, t, 4 * H) for r, t in ((2, 4), (8, 1), (3, 8), (8, 8)) +] +# KV lengths, assigned to the requests of a step in turn: the step's rows crossing a page boundary (579, 1989), only +# the step's tokens (8, 1), one 128-row tile minus one, > 16 tiles (several per CTA of the cluster), short contexts. +LENGTHS = (1100, 64 * 9 + 3, 8, 2049, 127, 4100, 64 * 31 + 5, 300, 64, 1) + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def _make_case(seed, num_requests, tokens, heads=H, row_stride=DQK, layers=1, slot=0, lens=None): + """A step of `num_requests` x `tokens`: q [M, heads, 576] (request-major), a pool of `layers` interleaved layer + slots with rows of `row_stride` elements, the page table as rows of an int32 [R, 2, W] buffer (the layout of + kv_cache_block_offsets: row stride 2 W; entries past a request's pages name a page no request owns), lengths + (`lens`, or LENGTHS in turn).""" + gen = torch.Generator(device="cuda").manual_seed(seed) + cpu_gen = torch.Generator().manual_seed(seed) + if lens is None: + lens = [max(tokens, LENGTHS[(i + seed) % len(LENGTHS)]) for i in range(num_requests)] + pages = [(n + PAGE - 1) // PAGE for n in lens] + total_pages = sum(pages) + 3 + width = max(pages) + 2 + pool = torch.randn(total_pages * layers, PAGE, row_stride, generator=gen, device="cuda") * 0.5 + pool = pool.bfloat16() + perm = (torch.randperm(total_pages, generator=cpu_gen) * layers).to(torch.int32) + offsets = torch.empty(num_requests, 2, width, dtype=torch.int32).fill_(int(perm[-1])) + start = 0 + for i, n in enumerate(pages): + offsets[i, :, :n] = perm[start : start + n] + start += n + page_table = offsets.cuda()[:, 0, :] + seq_len = torch.tensor(lens, dtype=torch.int32, device="cuda") + q = torch.randn(num_requests * tokens, heads, DQK, generator=gen, device="cuda") * 0.5 + return q.bfloat16(), pool, page_table, slot, seq_len + + +def _reference(q, pool, page_table, page_offset, seq_len, tokens): + """float64 attention of each request's tokens over its pages; token t sees rows <= L - T + t. [M, heads, 512].""" + outs = [] + for i, length in enumerate(seq_len.tolist()): + pages = page_table[i, : (length + PAGE - 1) // PAGE].long() + page_offset + kv = pool[pages].reshape(-1, pool.shape[-1])[:length, :DQK].double() + qi = q[i * tokens : (i + 1) * tokens].double() + s = torch.einsum("thd,ld->thl", qi, kv) * SCALE + limit = length - tokens + torch.arange(tokens, device=q.device) + hidden = torch.arange(length, device=q.device)[None, :] > limit[:, None] + s = s.masked_fill(hidden[:, None, :], float("-inf")) + outs.append(torch.einsum("thl,ld->thd", torch.softmax(s, dim=-1), kv[:, :LAT])) + return torch.cat(outs) + + +def _max_rel(a, b, tokens, num_requests): + """max over requests of max |a - b| / max |b| (per request, so a short context is not hidden by a long one).""" + err = 0.0 + for i in range(num_requests): + ai, bi = ( + a[i * tokens : (i + 1) * tokens].double(), + b[i * tokens : (i + 1) * tokens].double(), + ) + err = max(err, (ai - bi).abs().max().item() / max(bi.abs().max().item(), 1e-6)) + return err + + +def _attn_out(q, pool, row_stride, page_table, page_offset, seq_len): + from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import op # noqa: F401 + + m, heads = q.shape[0], q.shape[1] + out = torch.empty(m, heads * LAT, dtype=torch.bfloat16, device="cuda") + torch.ops.trtllm.k3_mla_attn_out( + q.reshape(m, -1), pool.view(-1), row_stride, page_table, page_offset, seq_len, SCALE, out + ) + return out.view(m, heads, LAT) + + +def _attn_vb(q, pool, row_stride, page_table, page_offset, seq_len, w_vb, gate=None): + from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import op # noqa: F401 + + m, heads = q.shape[0], q.shape[1] + y = torch.empty(m, heads * V, dtype=torch.bfloat16, device="cuda") + torch.ops.trtllm.k3_mla_attn_vb_out( + q.reshape(m, -1), + pool.view(-1), + row_stride, + page_table, + page_offset, + seq_len, + SCALE, + w_vb, + y, + gate, + GATE_COL0, + ) + return y + + +@pytest.mark.parametrize( + "num_requests,tokens,heads", ATTN_CASES, ids=[f"{r}x{t}-h{h}" for r, t, h in ATTN_CASES] +) +def test_attn_out(num_requests, tokens, heads): + """k3_mla_attn_out against the float64 reference, request by request against the one-request call, and rerun. Odd + seeds use an interleaved pool (2 layers, the layer's slot as the page offset) with rows of 640 elements.""" + seed = 11 * num_requests + tokens + layers, slot, row_stride = (2, 1, 640) if seed % 2 else (1, 0, DQK) + q, pool, page_table, page_offset, seq_len = _make_case( + seed, num_requests, tokens, heads, row_stride, layers, slot + ) + out = _attn_out(q, pool, row_stride, page_table, page_offset, seq_len) + again = _attn_out(q, pool, row_stride, page_table, page_offset, seq_len) + rows = [slice(i * tokens, (i + 1) * tokens) for i in range(num_requests)] + alone = torch.cat( + [_attn_out(q[r], pool, row_stride, page_table[i : i + 1], page_offset, seq_len[i : i + 1]) + for i, r in enumerate(rows)] + ) # fmt: skip + ref = _reference(q, pool, page_table, page_offset, seq_len, tokens) + torch.cuda.synchronize() + assert _max_rel(out, ref, tokens, num_requests) <= 1e-2 + assert torch.equal(_bits(out), _bits(again)) + assert torch.equal(_bits(out), _bits(alone)) + + +@pytest.mark.parametrize("num_requests,tokens", SPLITS, ids=SPLIT_IDS) +def test_attn_returns(num_requests, tokens): + """k3_mla_attn (its own output; a flat page-table row when R = 1) gives k3_mla_attn_out's bits.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import op # noqa: F401 + + q, pool, page_table, _, seq_len = _make_case(5 + num_requests, num_requests, tokens) + m = q.shape[0] + table = page_table[0] if num_requests == 1 else page_table + out = torch.ops.trtllm.k3_mla_attn(q.view(m, -1), pool.view(-1), DQK, table, seq_len, SCALE) + ref = _attn_out(q, pool, DQK, page_table, 0, seq_len) + assert torch.equal(_bits(out), _bits(ref.view(m, -1))) + + +@pytest.mark.parametrize("gated", [False, True], ids=["plain", "gated"]) +@pytest.mark.parametrize("num_requests,tokens", SPLITS, ids=SPLIT_IDS) +def test_attn_vb(num_requests, tokens, gated): + """k3_mla_attn_vb_out against the reference (attention output rounded to bf16, v_b in float64), request by request + against the one-request call, rerun; with the gate, bit-identical to torch's bf16 y * s.""" + seed = 7 * num_requests + tokens + q, pool, page_table, page_offset, seq_len = _make_case(seed, num_requests, tokens) + gen = torch.Generator(device="cuda").manual_seed(seed + 1) + m, heads = q.shape[0], q.shape[1] + w_vb = (torch.randn(heads, V, LAT, generator=gen, device="cuda") * 0.05).bfloat16() + ag = ( + torch.rand(m, GATE_COL0 + heads * V, generator=gen, device="cuda").bfloat16() + if gated + else None + ) + y = _attn_vb(q, pool, DQK, page_table, page_offset, seq_len, w_vb, ag) + again = _attn_vb(q, pool, DQK, page_table, page_offset, seq_len, w_vb, ag) + rows = [slice(i * tokens, (i + 1) * tokens) for i in range(num_requests)] + alone = torch.cat( + [_attn_vb(q[r], pool, DQK, page_table[i : i + 1], page_offset, seq_len[i : i + 1], w_vb, + ag[r] if gated else None) for i, r in enumerate(rows)] + ) # fmt: skip + o_ref = _reference(q, pool, page_table, page_offset, seq_len, tokens).bfloat16().double() + y_ref = torch.einsum("thc,hvc->thv", o_ref, w_vb.double()).reshape(m, heads * V) + torch.cuda.synchronize() + if gated: + plain = _attn_vb(q, pool, DQK, page_table, page_offset, seq_len, w_vb) + assert torch.equal(_bits(y), _bits(plain * ag[:, GATE_COL0:])) + y_ref = y_ref.bfloat16().double() * ag[:, GATE_COL0:].double() + assert _max_rel(y, y_ref, tokens, num_requests) <= 1e-2 + assert torch.equal(_bits(y), _bits(again)) + assert torch.equal(_bits(y), _bits(alone)) + + +@pytest.mark.parametrize("num_requests,tokens", SPLITS, ids=SPLIT_IDS) +def test_attn_vs_stock(num_requests, tokens): + """k3_mla_attn_out against the stock CuTe DSL MLA decode on the same step (batch R, seq_len_q T); both within 1e-2 + of the float64 reference.""" + import cutlass + + from tensorrt_llm._torch.custom_ops.cute_dsl_custom_ops import CuteDSLNVMlaDecodeBlackwellRunner + + q, pool, page_table, page_offset, seq_len = _make_case( + 3 * num_requests + tokens, num_requests, tokens + ) + m, heads = q.shape[0], q.shape[1] + size = CuteDSLNVMlaDecodeBlackwellRunner.get_max_padded_workspace_size( + heads, tokens, LAT, num_requests, cutlass.Float32 + ) + workspace = torch.empty(max(size, 1), dtype=torch.int8, device="cuda") + kv = pool[:, :, :DQK] + stock = torch.empty(num_requests, tokens, heads, LAT, dtype=torch.bfloat16, device="cuda") + qv = q.view(num_requests, tokens, heads, DQK) + torch.ops.trtllm.cute_dsl_mla_decode_fp16_blackwell( + qv[..., :LAT].permute(2, 3, 1, 0), qv[..., LAT:].permute(2, 3, 1, 0), kv[..., :LAT].permute(1, 2, 0), + kv[..., LAT:].permute(1, 2, 0), (page_table + page_offset).transpose(0, 1), seq_len, + stock.permute(2, 3, 1, 0), workspace, heads, tokens, PAGE, SCALE, 1.0, num_requests, None, None, + ) # fmt: skip + out = _attn_out(q, pool, DQK, page_table, page_offset, seq_len) + ref = _reference(q, pool, page_table, page_offset, seq_len, tokens) + torch.cuda.synchronize() + stock = stock.view(m, heads, LAT) + assert _max_rel(stock, ref, tokens, num_requests) <= 1e-2 + assert _max_rel(out, stock, tokens, num_requests) <= 1e-2 + + +def test_attn_rejects(): + """Calls outside R <= 8 requests of T <= 8 tokens, or page-table rows / lengths that do not match, are refused.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import op + + q, pool, page_table, _, seq_len = _make_case(1, 2, 4) + q2, flat = q.view(8, -1), pool.view(-1) + assert op.supports_attn(q2, flat, DQK, page_table, seq_len) + assert not op.supports_attn(q2, flat, DQK, page_table[:1], seq_len) # 1 row, 2 lengths + assert not op.supports_attn(q2, flat, DQK, page_table[0], seq_len) # a flat row, 2 lengths + assert op.supports_attn(q2[:6], flat, DQK, page_table, seq_len) # 2 requests x 3 tokens + assert not op.supports_attn(q2[:7], flat, DQK, page_table, seq_len) # 7 tokens over 2 requests + q16, _, table16, _, len16 = _make_case(2, 1, 16) + assert not op.supports_attn(q16.view(16, -1), flat, DQK, table16, len16) # T = 16 + q9, _, table9, _, len9 = _make_case(3, 9, 1) + assert not op.supports_attn(q9.view(9, -1), flat, DQK, table9, len9) # R = 9 + with pytest.raises(ValueError): + _attn_out(q[:7], pool, DQK, page_table, 0, seq_len) + + +# Both launch modes (clusters; no_cluster past CLUSTER_WAVE clusters) at the TP16 and TP4 head counts, with folds +# (3x5: 4100 rows, 8x1, 8x8) and CTAs without a tile (1x8-h24: 5 tiles). +POISON_CASES = [(1, 1, H), (3, 5, H), (8, 1, H), (8, 8, H), (1, 8, 4 * H), (2, 4, 4 * H)] + + +def _fill_workspace(heads, value): + """Fill the data words of the attention workspace: the per-CTA partial slots (fp16) and the no_cluster (m, l) + exchange (fp32). The arrival counters after them keep their values.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import k3_mla_attn_kernel as kernel + from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import op + + groups = heads // kernel.HEADS + ws = op._attn_workspace(torch.device("cuda", torch.cuda.current_device()), groups) + slots = kernel.MAX_REQUESTS * groups * kernel.CLUSTER + partials = slots * kernel.WS_SLOT_ELEMS + exchange = slots * kernel.ROWS * 2 # fp32 words + ws[:partials].fill_(value) + ws[partials : partials + 2 * exchange].view(torch.float32).fill_(value) + + +@pytest.mark.parametrize( + "num_requests,tokens,heads", POISON_CASES, ids=[f"{r}x{t}-h{h}" for r, t, h in POISON_CASES] +) +def test_attn_workspace_poison(num_requests, tokens, heads): + """A call reads only workspace words it wrote itself: with the partial slots and the (m, l) exchange refilled + with NaN before each call, k3_mla_attn_out and the gated k3_mla_attn_vb_out give the bits of the same calls on a + zero-filled workspace.""" + seed = 13 * num_requests + tokens + q, pool, page_table, page_offset, seq_len = _make_case(seed, num_requests, tokens, heads) + gen = torch.Generator(device="cuda").manual_seed(seed + 1) + m = q.shape[0] + w_vb = (torch.randn(heads, V, LAT, generator=gen, device="cuda") * 0.05).bfloat16() + ag = torch.rand(m, GATE_COL0 + heads * V, generator=gen, device="cuda").bfloat16() + outs = {} + try: + for value in (0.0, float("nan")): + _fill_workspace(heads, value) + o = _attn_out(q, pool, DQK, page_table, page_offset, seq_len) + _fill_workspace(heads, value) + y = _attn_vb(q, pool, DQK, page_table, page_offset, seq_len, w_vb, ag) + outs[value == 0.0] = (o, y) + finally: + _fill_workspace(heads, 0.0) + torch.cuda.synchronize() + for got, want in zip(outs[False], outs[True]): + assert not torch.isnan(got).any() + assert torch.equal(_bits(got), _bits(want)) + + +# no_cluster launches (more clusters than co-reside): 8 requests at 6 heads, 2 requests at 24 heads. +WRAP_CASES = [(8, 1, H), (8, 8, H), (2, 4, 4 * H)] + + +def _set_counters(heads, value): + """Set every no_cluster arrival counter of the attention workspace (int32 words after the (m, l) exchange).""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import k3_mla_attn_kernel as kernel + from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import op + + groups = heads // kernel.HEADS + ws = op._attn_workspace(torch.device("cuda", torch.cuda.current_device()), groups) + slots = kernel.MAX_REQUESTS * groups * kernel.CLUSTER + start = ( + slots * kernel.WS_SLOT_ELEMS + 2 * slots * kernel.ROWS * 2 + ) # fp16 elements before the counters + ws[start:].view(torch.int32).fill_(value) + + +@pytest.mark.parametrize( + "num_requests,tokens,heads", WRAP_CASES, ids=[f"{r}x{t}-h{h}" for r, t, h in WRAP_CASES] +) +def test_attn_counter_wrap(num_requests, tokens, heads): + """The no_cluster arrival counters only grow (16 per launch) and are compared by signed difference: calls with the + counters just below the int32 wrap (2^31 - 16, and -16 just below 0) give the bits of calls on zeroed ones.""" + seed = 19 * num_requests + tokens + q, pool, page_table, page_offset, seq_len = _make_case(seed, num_requests, tokens, heads) + gen = torch.Generator(device="cuda").manual_seed(seed + 1) + w_vb = (torch.randn(heads, V, LAT, generator=gen, device="cuda") * 0.05).bfloat16() + outs = {} + try: + for start in (0, 2**31 - 16, -16): + _set_counters(heads, start) + # Three launches: each crosses the wrap point once the counters start 16 below it. + outs[start] = [ + _attn_out(q, pool, DQK, page_table, page_offset, seq_len), + _attn_vb(q, pool, DQK, page_table, page_offset, seq_len, w_vb), + _attn_out(q, pool, DQK, page_table, page_offset, seq_len), + ] + finally: + _set_counters(heads, 0) + torch.cuda.synchronize() + for start in (2**31 - 16, -16): + for got, want in zip(outs[start], outs[0]): + assert torch.equal(_bits(got), _bits(want)) + + +# ---------------------------------------------------------------------------------------------------------------- +# Batch-1 identity against the unmodified kernel (K3_BASE_TRTLLM: an unmodified tensorrt_llm package directory). +# ---------------------------------------------------------------------------------------------------------------- + +_base = {} + + +def _base_kernel(): + """The unmodified single-request kernel module, loaded from K3_BASE_TRTLLM.""" + if "mod" not in _base: + path = os.path.join( + os.environ["K3_BASE_TRTLLM"], + "_torch", + "cute_dsl_kernels", + "k3_mla", + "k3_mla_attn_kernel.py", + ) + spec = importlib.util.spec_from_file_location("k3_mla_attn_kernel_base", path) + mod = importlib.util.module_from_spec(spec) + spec.loader.exec_module(mod) + _base["mod"] = mod + return _base["mod"] + + +def _base_attn(q, pool, row_stride, page_row, page_offset, seq_len, out, w_vb=None, gate=None): + """The unmodified op's launch (one request: one 16-byte aligned page-table row, seq_len [1]); q [M, heads * 576].""" + import cuda.bindings.driver as cuda_driver + import cutlass.cute as cute + from cutlass.cute.runtime import from_dlpack + + kern = _base_kernel() + + def arg(t): + return from_dlpack(t.detach(), assumed_align=16).mark_layout_dynamic( + leading_dim=t.dim() - 1 + ) + + heads = q.shape[1] // kern.QK + groups = heads // kern.HEADS + if ("ws", groups) not in _base: + _base["ws", groups] = torch.empty( + groups * kern.CLUSTER * kern.WS_SLOT_ELEMS, dtype=torch.float16, device=q.device + ) + fuse_vb, apply_gate = w_vb is not None, gate is not None + gate_flat = gate.as_strided((gate.numel(),), (1,)) if apply_gate else q.view(-1) + args = (arg(q.view(-1)), arg(pool.view(-1)[: PAGE * row_stride]), arg(page_row.reshape(-1)), + arg(seq_len.reshape(-1)), arg(_base["ws", groups]), arg(out.view(-1)), + arg((w_vb if fuse_vb else q).view(-1)), arg(gate_flat)) # fmt: skip + scalars = (q.shape[0], SCALE * kern.LOG2E, pool.numel() // row_stride, page_offset, GATE_COL0, + gate.stride(0) if apply_gate else 0) # fmt: skip + use_pdl = os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" + stream = cuda_driver.CUstream(torch.cuda.current_stream().cuda_stream) + key = (row_stride, heads, fuse_vb, apply_gate, use_pdl) + fn = _base.get(key) + if fn is None: + fn = _base[key] = cute.compile(kern.k3_mla_attn, *args, *scalars, row_stride, heads, fuse_vb, apply_gate, + use_pdl, stream) # fmt: skip + fn(*args, *scalars, stream) + return out + + +@pytest.mark.skipif( + not os.environ.get("K3_BASE_TRTLLM"), reason="K3_BASE_TRTLLM (unmodified package) not set" +) +@pytest.mark.parametrize("tokens", [8, 1, 2, 4, 7]) +def test_batch1_identity(tokens): + """One request: k3_mla_attn_out and k3_mla_attn_vb_out (plain and gated) bit-identical to the unmodified kernel + for every length in LENGTHS, rows of 576 and interleaved rows of 640 with a page offset, the page-table row given + flat, as a [1, W] row, and at a 4-byte offset (a row of a wider table).""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import op # noqa: F401 + + for i, length in enumerate(LENGTHS): + layers, slot, row_stride = (2, 1, 640) if i % 2 else (1, 0, DQK) + q, pool, table, off, seq_len = _make_case( + 900 + 10 * tokens + i, + 1, + tokens, + H, + row_stride, + layers, + slot, + lens=[max(tokens, length)], + ) + gen = torch.Generator(device="cuda").manual_seed(1900 + 10 * tokens + i) + w_vb = (torch.randn(H, V, LAT, generator=gen, device="cuda") * 0.05).bfloat16() + ag = torch.rand(tokens, GATE_COL0 + H * V, generator=gen, device="cuda").bfloat16() + shifted = torch.zeros(table.shape[1] + 1, dtype=torch.int32, device="cuda") + shifted[1:] = table[0] + q2 = q.view(tokens, -1) + want = _base_attn(q2, pool, row_stride, table[0], off, seq_len, + torch.empty(tokens, H * LAT, dtype=torch.bfloat16, device="cuda")) # fmt: skip + for form, row in ( + ("flat", table[0]), + ("[1, W]", table[:1]), + ("4-byte offset", shifted[1:]), + ): + got = _attn_out(q, pool, row_stride, row, off, seq_len).view(tokens, -1) + assert torch.equal(_bits(got), _bits(want)), f"attn_out L {length} {form}" + for gate in (None, ag): + want = _base_attn(q2, pool, row_stride, table[0], off, seq_len, + torch.empty(tokens, H * V, dtype=torch.bfloat16, device="cuda"), w_vb, gate) # fmt: skip + got = _attn_vb(q, pool, row_stride, table[:1], off, seq_len, w_vb, gate) + assert torch.equal(_bits(got), _bits(want)), ( + f"attn_vb L {length} gate {gate is not None}" + ) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_decode_view.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_decode_view.py new file mode 100644 index 000000000000..8b5004947b2e --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_decode_view.py @@ -0,0 +1,125 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""k3_mla_decode_view: the Kimi K3 MLA decode kernels' per-layer inputs from the attention metadata for R generation +requests of T tokens (views of the metadata buffers, no copies), and a reason string for every step it cannot take.""" + +from types import SimpleNamespace + +import pytest +import torch + +from tensorrt_llm._torch.attention.backends.fmha.cute_dsl_mla import k3_mla_decode_view + +MAX_SEQS, MAX_PAGES, POOL_PAGES = 16, 32, 40 + + +def _attn(num_heads=6): + return SimpleNamespace( + num_heads=num_heads, kv_lora_rank=512, qk_rope_head_dim=64, qk_nope_head_dim=128, q_scaling=1.0, + layer_idx=3, get_local_layer_idx=lambda meta: 1, + ) # fmt: skip + + +def _meta(num_generations, num_contexts=0, pool=None, **overrides): + if pool is None: + pool = torch.zeros(POOL_PAGES, 1, 64, 1, 576, dtype=torch.bfloat16) + meta = SimpleNamespace( + num_contexts=num_contexts, + num_generations=num_generations, + beam_width=1, + tokens_per_block=64, + helix_position_offsets=None, + kv_cache_manager=SimpleNamespace(get_buffers=lambda layer_idx: pool), + kv_cache_block_offsets=torch.arange(2 * MAX_SEQS * 2 * MAX_PAGES, dtype=torch.int32).view( + 2, MAX_SEQS, 2, MAX_PAGES + ), + host_kv_cache_pool_mapping=torch.tensor([[0, 0], [1, 0]], dtype=torch.int32), + kv_lens_cuda_runtime=torch.arange(100, 100 + MAX_SEQS, dtype=torch.int32), + ) + for key, value in overrides.items(): + setattr(meta, key, value) + return meta + + +SPLITS = [(1, 1), (2, 1), (3, 1), (4, 1), (5, 1), (8, 1), (1, 8), (2, 4), (4, 2)] + [ + (r, 8) for r in range(2, 9) +] + + +@pytest.mark.cpu_only +@pytest.mark.parametrize("num_requests,tokens", SPLITS, ids=[f"{r}x{t}" for r, t in SPLITS]) +def test_view_requests(num_requests, tokens): + """R x T: the R page-table rows of the layer's pool as a strided view of kv_cache_block_offsets, the R lengths as + a view of kv_lens_cuda_runtime, R and T.""" + meta = _meta(num_requests) + view = k3_mla_decode_view(_attn(), meta, num_requests * tokens) + assert isinstance(view, dict), view + table = view["page_table"] + assert tuple(table.shape) == (num_requests, MAX_PAGES) and table.stride() == (2 * MAX_PAGES, 1) + assert table.data_ptr() == meta.kv_cache_block_offsets[1, 0, 0].data_ptr() + assert torch.equal(table, meta.kv_cache_block_offsets[1, :num_requests, 0]) + assert view["seq_len"].data_ptr() == meta.kv_lens_cuda_runtime.data_ptr() + assert view["seq_len"].shape == (num_requests,) + assert (view["num_requests"], view["tokens_per_request"]) == (num_requests, tokens) + assert view["row_stride"] == 576 and view["page_offset"] == 0 + assert view["pool"].numel() == POOL_PAGES * 64 * 576 + assert view["softmax_scale"] == pytest.approx(192**-0.5) + + +@pytest.mark.cpu_only +def test_view_interleaved_pool(): + """A layer's view of a layer-interleaved pool: the whole pool with the layer's slot as the page offset.""" + layers = 3 + pools = torch.zeros(POOL_PAGES, layers, 1, 64, 1, 640, dtype=torch.bfloat16) + view = k3_mla_decode_view(_attn(), _meta(2, pool=pools[:, 2]), 16) + assert isinstance(view, dict), view + assert view["page_offset"] == 2 and view["row_stride"] == 640 + assert view["pool"].numel() == POOL_PAGES * layers * 64 * 640 + + +@pytest.mark.cpu_only +@pytest.mark.parametrize( + "num_generations,num_tokens,overrides,attn_heads", + [ + (2, 16, dict(num_contexts=1), 6), # a context request + (0, 8, {}, 6), # no generation request + (9, 9, {}, 6), # more than 8 requests + (1, 16, {}, 6), # more than 8 tokens per request + (2, 7, {}, 6), # tokens not uniform over the requests + (2, 8, dict(beam_width=2), 6), + (2, 8, dict(is_spec_dec_tree=True), 6), + (2, 8, dict(tokens_per_block=32), 6), + (2, 8, dict(helix_position_offsets=torch.zeros(1)), 6), + (2, 8, dict(kv_cache_manager=None), 6), + (2, 8, {}, 8), # heads not a multiple of 6 + ], +) +def test_view_reasons(num_generations, num_tokens, overrides, attn_heads): + """Every step the kernels cannot take returns a reason string (and does not raise).""" + overrides = dict(overrides) + num_contexts = overrides.pop("num_contexts", 0) + reason = k3_mla_decode_view( + _attn(attn_heads), _meta(num_generations, num_contexts, **overrides), num_tokens + ) + assert isinstance(reason, str) and reason + + +@pytest.mark.cpu_only +def test_view_pool_reasons(): + """A quantized or differently laid out pool returns a reason string.""" + fp8 = torch.zeros(POOL_PAGES, 1, 64, 1, 576, dtype=torch.float8_e4m3fn) + assert isinstance(k3_mla_decode_view(_attn(), _meta(1, pool=fp8), 1), str) + narrow = torch.zeros(POOL_PAGES, 1, 64, 1, 512, dtype=torch.bfloat16) + assert isinstance(k3_mla_decode_view(_attn(), _meta(1, pool=narrow), 1), str) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_q.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_q.py new file mode 100644 index 000000000000..6dfad07ca19b --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_q.py @@ -0,0 +1,396 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""trtllm::k3_mla_q / k3_mla_qkv / k3_mla_qkv_out (Kimi K3 MLA decode query path: q_a RMSNorm, q_b, k_b absorb, and the +KV half into the paged latent cache) at the TP16 shape (6 heads per rank, q_lora 1536, latent 512 + rope 64), M <= 64 +tokens: fused_q against the unfused chain (the model's RMSNorm, the q_b GEMM, the k_b bmm) and a reference with +the same bf16 roundings, every 8-token chunk bit-identical to the 8-token call on its rows; the KV rows of R requests +x T tokens at each request's positions in a sentinel-filled pool, nothing else written. + +Batch-1 identity against the unmodified kernel: set ``K3_BASE_TRTLLM`` to an unmodified ``tensorrt_llm`` package +directory (one request's M <= 8 outputs, and the pool, must match its kernel bit for bit).""" + +import importlib.util +import os + +import pytest +import torch + + +def _is_sm100() -> bool: + if not torch.cuda.is_available(): + return False + major, minor = torch.cuda.get_device_capability() + return major * 10 + minor in (100, 103) + + +pytestmark = pytest.mark.skipif(not _is_sm100(), reason="needs an SM100 / SM103 GPU") + +H, NOPE, PE, QK, LAT, QL, V = 6, 128, 64, 192, 512, 1536, 128 +DQK, PAGE = LAT + PE, 64 +EPS = KV_EPS = 1e-6 +SENTINEL = 0x7F7F # bf16 bits of the largest finite value: no cache row of the tests holds it + +SPLITS = [(1, 1), (2, 1), (3, 1), (4, 1), (5, 1), (8, 1), (1, 8), (2, 4), (4, 2)] + [ + (r, 8) for r in range(2, 9) +] +SPLITS += [(3, 5), (7, 3)] # chunks of 8 tokens that cut through requests +SPLIT_IDS = [f"{r}x{t}" for r, t in SPLITS] +# KV lengths, assigned to the requests of a step in turn: the step's rows crossing a page boundary (579, 1989), only +# the step's tokens (8), short and long contexts. +LENGTHS = (1100, 64 * 9 + 3, 8, 2049, 127, 4100, 64 * 31 + 5, 300, 64) + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def _weights(seed, heads=H): + gen = torch.Generator(device="cuda").manual_seed(seed) + w_qa = (1.0 + 0.1 * torch.randn(QL, generator=gen, device="cuda")).bfloat16() + w_qb = (torch.randn(heads * QK, QL, generator=gen, device="cuda") * 0.03).bfloat16() + w_kb = (torch.randn(heads, LAT, NOPE, generator=gen, device="cuda") * 0.08).bfloat16() + w_kv = (1.0 + 0.1 * torch.randn(LAT, generator=gen, device="cuda")).bfloat16() + return w_qa, w_qb, w_kb, w_kv + + +def _ag(seed, m, heads=H): + """The fused projection's rows: [q_a 1536 | kv_a latent 512 | rope 64 | gate heads * 128].""" + gen = torch.Generator(device="cuda").manual_seed(seed) + return (torch.randn(m, QL + DQK + heads * V, generator=gen, device="cuda") * 0.7).bfloat16() + + +def _rms_norm(x, w, eps): + from tensorrt_llm._torch.modules.rms_norm import RMSNorm + + norm = RMSNorm(hidden_size=x.shape[1], eps=eps, dtype=torch.bfloat16).cuda() + norm.weight.data.copy_(w) + return norm(x.contiguous()) + + +def _unfused_q(ag, w_qa, w_qb, w_kb): + """The model's unfused chain: q_a_layernorm, q_b_proj (GEMM), bmm with k_b_proj_trans, q_pe copied.""" + m, heads = ag.shape[0], w_kb.shape[0] + q = torch.matmul(_rms_norm(ag[:, :QL], w_qa, EPS), w_qb.t()).view(m, heads, QK) + q_abs = torch.bmm(q[..., :NOPE].transpose(0, 1), w_kb.transpose(1, 2)).transpose(0, 1) + return torch.cat([q_abs, q[..., NOPE:]], dim=-1).reshape(m, heads * DQK) + + +def _reference_q(ag, w_qa, w_qb, w_kb): + """The norm in fp32, the GEMMs in float64 (cuBLAS may run fp32 GEMMs in TF32), bf16 rounding where the model + rounds (norm output, q_b output, q_abs).""" + m, heads = ag.shape[0], w_kb.shape[0] + x = ag[:, :QL].float() + qn = (x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + EPS) * w_qa.float()).bfloat16() + q = (qn.double() @ w_qb.double().t()).bfloat16().view(m, heads, QK) + q_abs = torch.einsum("thd,hcd->thc", q[..., :NOPE].double(), w_kb.double()).bfloat16() + return torch.cat([q_abs, q[..., NOPE:]], dim=-1).reshape(m, heads * DQK) + + +def _reference_kv(ag, w_kv): + """The cache rows in fp32, flashinfer's order ((x * rrms) * w, one bf16 rounding); the rope columns copied.""" + x = ag[:, QL : QL + LAT].float() + r = torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + KV_EPS) + return torch.cat([((x * r) * w_kv.float()).bfloat16(), ag[:, QL + LAT : QL + DQK]], dim=-1) + + +def _max_rel(a, b): + return (a.float() - b.float()).abs().max().item() / b.float().abs().max().item() + + +def _q(ag, w_qa, w_qb, w_kb): + from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import op # noqa: F401 + + return torch.ops.trtllm.k3_mla_q(ag, w_qa, EPS, w_qb, w_kb, True) + + +# Chunks of 8 tokens up to 24, of 16 up to 48 and of 32 up to 64 (op.CHUNK_TOKENS), with full and short last chunks. +# The TP4 shape (24 heads, 144 CTAs per chunk) at M 8, 40 and 64. +Q_CASES = [(m, H) for m in (1, 2, 3, 5, 8, 9, 15, 16, 24, 25, 31, 32, 33, 40, 48, 49, 57, 64)] + [ + (m, 4 * H) for m in (8, 40, 64) +] + + +@pytest.mark.parametrize("m,heads", Q_CASES, ids=[f"m{m}-h{h}" for m, h in Q_CASES]) +def test_q(m, heads): + """fused_q against the reference and the unfused chain (q_abs and q_pe), rerun; each 8-token chunk of rows + bit-identical to the call on those rows alone.""" + w_qa, w_qb, w_kb, _ = _weights(20261001, heads) + ag = _ag(m, m, heads) + y = _q(ag, w_qa, w_qb, w_kb) + again = _q(ag, w_qa, w_qb, w_kb) + chunks = torch.cat([_q(ag[c : c + 8].contiguous(), w_qa, w_qb, w_kb) for c in range(0, m, 8)]) + yu = _unfused_q(ag, w_qa, w_qb, w_kb) + yr = _reference_q(ag, w_qa, w_qb, w_kb) + torch.cuda.synchronize() + for part in (slice(0, LAT), slice(LAT, DQK)): + yk = y.view(m, heads, DQK)[..., part] + assert _max_rel(yk, yr.view(m, heads, DQK)[..., part]) <= 2e-2 + assert _max_rel(yk, yu.view(m, heads, DQK)[..., part]) <= 2e-2 + assert torch.equal(_bits(y), _bits(again)) + assert torch.equal(_bits(y), _bits(chunks)) + + +def _kv_case(seed, num_requests, tokens, row_stride=DQK, layers=1, lens=None): + """A sentinel-filled pool of `layers` interleaved layer slots (rows of `row_stride`), the page table as rows of an + int32 [R, 2, W] buffer (kv_cache_block_offsets' layout; entries past a request's pages name a page no request + owns), the lengths.""" + cpu_gen = torch.Generator().manual_seed(seed) + if lens is None: + lens = [max(tokens, LENGTHS[(i + seed) % len(LENGTHS)]) for i in range(num_requests)] + pages = [(n + PAGE - 1) // PAGE for n in lens] + total_pages = sum(pages) + 5 + width = max(pages) + 2 + pool = torch.full( + (total_pages * layers * PAGE * row_stride,), SENTINEL, dtype=torch.int16, device="cuda" + ) + perm = (torch.randperm(total_pages, generator=cpu_gen) * layers).to(torch.int32) + offsets = torch.empty(num_requests, 2, width, dtype=torch.int32).fill_(int(perm[-1])) + start = 0 + for i, n in enumerate(pages): + offsets[i, :, :n] = perm[start : start + n] + start += n + return ( + pool.view(torch.bfloat16), + offsets.cuda()[:, 0, :], + torch.tensor(lens, dtype=torch.int32, device="cuda"), + ) + + +def _kv_rows(page_table, page_offset, seq_len, tokens): + """(token, pool row) of every stored token: token u of request i at position L_i - T + u, if >= 0.""" + out = [] + for i, length in enumerate(seq_len.tolist()): + for u in range(tokens): + pos = length - tokens + u + if pos >= 0: + out.append( + ( + i * tokens + u, + (int(page_table[i, pos // PAGE]) + page_offset) * PAGE + pos % PAGE, + ) + ) + return out + + +def _qkv(ag, weights, pool, row_stride, page_table, page_offset, seq_len): + from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import op # noqa: F401 + + w_qa, w_qb, w_kb, w_kv = weights + return torch.ops.trtllm.k3_mla_qkv( + ag, + w_qa, + EPS, + w_qb, + w_kb, + w_kv, + KV_EPS, + pool, + row_stride, + page_table, + page_offset, + seq_len, + True, + ) + + +@pytest.mark.parametrize("num_requests,tokens", SPLITS, ids=SPLIT_IDS) +def test_qkv(num_requests, tokens): + """The KV rows of every request at its positions (latent against the fp32 reference and the model's RMSNorm, rope + columns bit-exact), nothing else in the pool written, fused_q bit-identical to k3_mla_q, the dense variant's rows + bit-identical, rerun. Odd seeds: an interleaved pool (2 layers, slot 1) with rows of 640 elements.""" + seed = 13 * num_requests + tokens + layers, slot, row_stride = (2, 1, 640) if seed % 2 else (1, 0, DQK) + m = num_requests * tokens + weights = _weights(seed) + ag = _ag(seed + 1, m) + pool, page_table, seq_len = _kv_case(seed, num_requests, tokens, row_stride, layers) + y = _qkv(ag, weights, pool, row_stride, page_table, slot, seq_len) + stored = _kv_rows(page_table, slot, seq_len, tokens) + toks = torch.tensor([t for t, _ in stored], device="cuda") + rows = torch.tensor([r for _, r in stored], device="cuda") + pool_rows = pool.view(-1, row_stride) + got = pool_rows[rows, :DQK].clone() + rest = pool_rows.clone() + rest[rows] = torch.full_like(rest[rows].view(torch.int16), SENTINEL).view(torch.bfloat16) + ref = _reference_kv(ag, weights[3])[toks] + model = _rms_norm(ag[:, QL : QL + LAT], weights[3], KV_EPS)[toks] + dense = torch.empty(m, DQK, dtype=torch.bfloat16, device="cuda") + y_dense = torch.ops.trtllm.k3_mla_qkv_out( + ag, weights[0], EPS, weights[1], weights[2], weights[3], KV_EPS, dense, True + ) + y_q = _q(ag, *weights[:3]) + _qkv(ag, weights, pool, row_stride, page_table, slot, seq_len) + torch.cuda.synchronize() + assert len(stored) == m + assert _max_rel(got[:, :LAT], ref[:, :LAT]) <= 1e-2 + assert int((_bits(got[:, :LAT]) != _bits(model)).sum()) <= max(4, model.numel() // 1000) + assert torch.equal(_bits(got[:, LAT:]), _bits(ag[toks, QL + LAT : QL + DQK])) + assert bool((rest.view(torch.int16) == SENTINEL).all()) + assert torch.equal(_bits(y), _bits(y_q)) and torch.equal(_bits(y_dense), _bits(y_q)) + assert torch.equal(_bits(dense[toks]), _bits(got)) + assert torch.equal(_bits(pool_rows[rows, :DQK]), _bits(got)) + + +def test_qkv_short_length(): + """A request whose length is below T (positions < 0) has only its tokens at positions >= 0 stored.""" + num_requests, tokens = 2, 4 + weights = _weights(5) + ag = _ag(6, num_requests * tokens) + pool, page_table, seq_len = _kv_case(7, num_requests, tokens, lens=[2, 100]) + _qkv(ag, weights, pool, DQK, page_table, 0, seq_len) + stored = _kv_rows(page_table, 0, seq_len, tokens) + rows = torch.tensor([r for _, r in stored], device="cuda") + pool_rows = pool.view(-1, DQK) + ref = _reference_kv(ag, weights[3])[[t for t, _ in stored]] + rest = pool_rows.clone() + rest[rows] = torch.full_like(rest[rows].view(torch.int16), SENTINEL).view(torch.bfloat16) + torch.cuda.synchronize() + assert len(stored) == 6 + assert _max_rel(pool_rows[rows, :LAT], ref[:, :LAT]) <= 1e-2 + assert bool((rest.view(torch.int16) == SENTINEL).all()) + + +def test_q_rejects(): + """More than 64 tokens, or page-table rows / lengths that do not describe R requests of the call's tokens.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import op + + w_qa, w_qb, w_kb, w_kv = _weights(1) + assert op.supports_q(_ag(1, 64), w_qa, w_qb, w_kb) + assert not op.supports_q(_ag(1, 65), w_qa, w_qb, w_kb) + pool, page_table, seq_len = _kv_case(2, 2, 4) + ag = _ag(2, 8) + assert op.supports_kv(ag, w_kv, pool, DQK, page_table, seq_len) + assert not op.supports_kv( + ag[:7], w_kv, pool, DQK, page_table, seq_len + ) # 7 tokens over 2 requests + assert not op.supports_kv(ag, w_kv, pool, DQK, page_table[:1], seq_len) # 1 row, 2 lengths + assert not op.supports_kv(ag, w_kv, pool, DQK, page_table[0], seq_len) # a flat row, 2 lengths + assert not op.supports_kv(ag, w_kv, pool, DQK, page_table.long(), seq_len) # int64 table + + +# ---------------------------------------------------------------------------------------------------------------- +# Batch-1 identity against the unmodified kernel (K3_BASE_TRTLLM: an unmodified tensorrt_llm package directory). +# ---------------------------------------------------------------------------------------------------------------- + +_base = {} + + +def _base_kernel(): + """The unmodified kernel module, loaded from K3_BASE_TRTLLM.""" + if "mod" not in _base: + path = os.path.join( + os.environ["K3_BASE_TRTLLM"], + "_torch", + "cute_dsl_kernels", + "k3_mla", + "k3_mla_q_kernel.py", + ) + spec = importlib.util.spec_from_file_location("k3_mla_q_kernel_base", path) + mod = importlib.util.module_from_spec(spec) + spec.loader.exec_module(mod) + _base["mod"] = mod + return _base["mod"] + + +def _base_q(ag, w_qa, w_qb, w_kb, w_kv=None, kv=None): + """The unmodified op's launch (M <= 8, one request): ``kv`` None (k3_mla_q), dict(pool, row_stride, page_row, + page_offset, seq_len) (k3_mla_qkv) or dict(out) (k3_mla_qkv_out). Returns fused_q.""" + import cuda.bindings.driver as cuda_driver + import cutlass.cute as cute + from cutlass.cute.runtime import from_dlpack + + kern = _base_kernel() + + def arg(t): + return from_dlpack(t.detach(), assumed_align=16).mark_layout_dynamic( + leading_dim=t.dim() - 1 + ) + + m, ag_cols = ag.shape + heads = w_kb.shape[0] + out = torch.empty(m, heads * kern.FUSED, dtype=torch.bfloat16, device=ag.device) + if "dummies" not in _base: + _base["dummies"] = (torch.zeros(8, dtype=torch.bfloat16, device="cuda"), + torch.zeros(8, dtype=torch.int32, device="cuda")) # fmt: skip + bf16_dummy, i32_dummy = _base["dummies"] + kv_mode, row_stride, kv_eps, page_offset = 0, kern.FUSED, 0.0, 0 + w_kv_t, pool, page_row, seq_len, kv_out = ( + bf16_dummy, + bf16_dummy, + i32_dummy, + i32_dummy, + bf16_dummy, + ) + if kv is not None: + w_kv_t, kv_eps = w_kv, KV_EPS + if "out" in kv: + kv_mode, kv_out = 2, kv["out"] + else: + kv_mode, row_stride, page_offset = 1, kv["row_stride"], kv["page_offset"] + pool = kv["pool"].view(-1)[: PAGE * row_stride] + page_row, seq_len = kv["page_row"], kv["seq_len"] + args = (arg(w_qb), arg(w_kb.view(heads * kern.LATENT, kern.NOPE)), arg(ag.view(-1)), arg(w_qa), + arg(out.view(-1)), arg(w_kv_t), arg(pool), arg(page_row.reshape(-1)), arg(seq_len.reshape(-1)), + arg(kv_out.view(-1))) # fmt: skip + scalars = (m, m, 0, EPS, kv_eps, page_offset) # M, T (one request), page-table row stride + use_pdl = os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" + stream = cuda_driver.CUstream(torch.cuda.current_stream().cuda_stream) + key = (ag_cols, heads, kv_mode, row_stride, use_pdl) + fn = _base.get(key) + if fn is None: + # trigger_early True, single_hop False, cluster_rms True: the unmodified op's defaults. + fn = _base[key] = cute.compile(kern.k3_mla_q, *args, *scalars, ag_cols, heads, True, False, kv_mode, + row_stride, True, use_pdl, stream) # fmt: skip + fn(*args, *scalars, stream) + return out + + +@pytest.mark.skipif( + not os.environ.get("K3_BASE_TRTLLM"), reason="K3_BASE_TRTLLM (unmodified package) not set" +) +@pytest.mark.parametrize("m", [1, 2, 5, 8]) +def test_batch1_identity(m): + """One request of M <= 8 tokens: k3_mla_q, k3_mla_qkv (fused_q and the whole pool: rows of 576, and interleaved + rows of 640 with a page offset; the page-table row flat and as [1, W]; the new rows crossing a page) and + k3_mla_qkv_out (fused_q and the dense rows) bit-identical to the unmodified kernel.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import op # noqa: F401 + + weights = _weights(700 + m) + ag = _ag(800 + m, m) + assert torch.equal(_bits(_q(ag, *weights[:3])), _bits(_base_q(ag, *weights[:3]))) + for layers, slot, row_stride in ((1, 0, DQK), (2, 1, 640)): + pool, table, seq_len = _kv_case(900 + m, 1, m, row_stride, layers, lens=[64 * 9 + 3]) + want_pool = pool.clone() + kv = dict( + pool=want_pool, + row_stride=row_stride, + page_row=table[0], + page_offset=slot, + seq_len=seq_len, + ) + want = _base_q(ag, *weights[:3], weights[3], kv) + for form, row in (("flat", table[0]), ("[1, W]", table[:1])): + got_pool = pool.clone() + got = _qkv(ag, weights, got_pool, row_stride, row, slot, seq_len) + assert torch.equal(_bits(got), _bits(want)), f"qkv fused_q rows of {row_stride} {form}" + assert torch.equal(_bits(got_pool), _bits(want_pool)), ( + f"qkv pool rows of {row_stride} {form}" + ) + want_rows = torch.empty(m, DQK, dtype=torch.bfloat16, device="cuda") + want = _base_q(ag, *weights[:3], weights[3], dict(out=want_rows)) + got_rows = torch.empty_like(want_rows) + got = torch.ops.trtllm.k3_mla_qkv_out(ag, weights[0], EPS, weights[1], weights[2], weights[3], KV_EPS, got_rows, + True) # fmt: skip + assert torch.equal(_bits(got), _bits(want)) and torch.equal(_bits(got_rows), _bits(want_rows)) From b8edc4adf8ab7bbf7eb8ed33cddf886748fbed32 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 18:09:03 -0700 Subject: [PATCH 019/161] [None][feat] modeling_v2: Kimi K3 MXFP4 target skeleton, sm_100, tp16 + moe tp4 x ep4 Route KimiK3ForConditionalGeneration to a new target, kimi_k3_mxfp4__sm_100__tp16_moetp4ep4: SM 10.0, the text_config shape (93, 7168, 896, 3584), no global quantization (the NVFP4 requant reads MIXED_PRECISION and does not route here), tp 16 with moe_tp 4 x moe_ep 4 and no attention DP. The target is text only. It checks its settings at construction (SM, topology, bf16, bf16 KV, AUTO or MNNVL all-reduce) and on the first forward (tokens_per_block 64, the V2 hybrid manager, block reuse off), and refuses multimodal input. Every step runs the built-in Kimi K3 text model as the generic path, listed in UNCERTIFIED_GENERIC_CALLS; the fused decode path's dispatch is in place and empty. The weight loader hands language_model.* to the built-in loader and fails on any key outside it and the vision tower's predicted non-load. Routing tests cover the match, the requant, three topologies, a wrong depth and SM 10.3. Signed-off-by: Vasanth Sabavat --- .../modeling_v2/_router_index.py | 1 + .../modeling_v2/models/kimi_k3_vl/__init__.py | 3 + .../__init__.py | 3 + .../modeling.py | 267 ++++++++++++++++++ .../weights.py | 42 +++ .../modeling_v2/models/kimi_k3_vl/routing.py | 123 ++++++++ .../modeling_v2/test_modeling_v2_routing.py | 99 +++++++ 7 files changed, 538 insertions(+) create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/__init__.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/__init__.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/weights.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/routing.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/_router_index.py b/tensorrt_llm/_torch/_experimental/modeling_v2/_router_index.py index d1b2a7840899..16c6b9769eb6 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/_router_index.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/_router_index.py @@ -59,6 +59,7 @@ MODELING_V2_ROUTERS = { "GptOssForCausalLM": "models.gpt_oss.routing", "DeepseekV3ForCausalLM": "models.deepseek_v3.routing", + "KimiK3ForConditionalGeneration": "models.kimi_k3_vl.routing", } #: Backends whose model construction reaches ``modeling_v2_resolve``. The diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/__init__.py new file mode 100644 index 000000000000..8d4f159338fc --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/__init__.py @@ -0,0 +1,3 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""KimiK3ForConditionalGeneration targets.""" diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/__init__.py new file mode 100644 index 000000000000..474e865afce4 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/__init__.py @@ -0,0 +1,3 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3 (MXFP4) / sm_100 / tp16 attention, routed experts moe_tp 4 x moe_ep 4.""" diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py new file mode 100644 index 000000000000..e02a387828b3 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -0,0 +1,267 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""ModelingV2 target: Kimi K3 (MXFP4) / sm_100 / tp16 attention, routed experts moe_tp 4 x moe_ep 4. + +Kimi K3's language model: 93 layers, hidden 7168. Every fourth layer (3, 7, ..., 91) is MLA attention with 96 +query heads; the others are Kimi Delta Attention (KDA), a gated linear-attention recurrence with per-request state. +Layer 0's MLP is dense. Every other layer's MLP is a latent MoE: 896 routed experts, top-16, expert width 3584. +Attention residuals mix each sublayer's input from a bank of earlier outputs. The routed experts are MXFP4 (the +checkpoint's compressed-tensors default, read as W4A8_MXFP4_MXFP8); everything else is bf16, and so is the KV pool. + +`tp16_moetp4ep4` is `tensor_parallel_size: 16` with `moe_tensor_parallel_size: 4` and +`moe_expert_parallel_size: 4`, no attention data parallelism, all 16 GPUs in one NVLink domain. Attention is +head-split (6 MLA query heads per rank), and every rank holds a quarter of the width of a quarter of the experts. + +**Each step takes one of two paths, chosen on the host from the step's shape** (`_step_path`): + +* **The fused decode path**: pure decode steps of at most 8 tokens without speculation, or at most 64 with DSpark + (8 requests x 1 + 7 drafts), on the K3 decode kernels' catalog entries. The state those kernels share (MNNVL + workspace, sandwich and MoE Lamport buffers, KDA / MLA scratch) lives in typed objects this target creates + collectively in `post_load_weights`, before any graph capture. Until those entries exist `_fused_decode` stays + None, and every step takes the generic path. +* **The generic path**: prefill, mixed steps, and decode steps above those bounds, on the built-in Kimi K3 text + model, whose modules and ops have no catalog entries yet. `UNCERTIFIED_GENERIC_CALLS` names them. + +**What this target asserts rather than adapts**: SM 10.0; the topology above; bf16 weights and a bf16 KV pool; +tokens_per_block 64 (the MLA generation kernels K3's 96 heads reach exist only at 64); the V2 hybrid KV / state +manager, which holds the KDA states, with block reuse off; an all-reduce strategy of AUTO or MNNVL. The +construction-time ones fail in `__init__`, the per-engine ones on the first forward, each naming the setting. + +**Text only.** The checkpoint is the vision-language wrapper. This target builds and loads no vision tower (its +weights are a predicted non-load, `weights.py`), and a step carrying multimodal input raises. + +**Speculative decoding** goes through the stock one-engine shell: DSpark or DFlash with an external drafter +checkpoint, and SA. The worker and its kernels stay upstream code; this target does not own a worker. +""" + +import copy +from typing import Any, Literal, Optional + +import torch + +from tensorrt_llm._torch.attention.backends.interface import AttentionMetadata +from tensorrt_llm._torch.model_config import ModelConfig +from tensorrt_llm._torch.models.modeling_kimi_linear import KimiLinearForCausalLM +from tensorrt_llm._torch.models.modeling_utils import register_auto_model +from tensorrt_llm._torch.pyexecutor.kv_cache.mamba_cache_manager import MambaHybridCacheManagerV2 +from tensorrt_llm.functional import AllReduceStrategy + +from . import weights as _weights + +# The GPU architecture this target IS. Routing will not send another one here, but a direct instantiation could, +# and the certification is per arch. +_SM = (10, 0) + +#: Every trtllm op this target reaches for, in its forward and in the weight load. Declared here, asserted in +#: tests/unittest/_torch/modeling_v2. Today these are the K3-specific ops of the generic path (attention residuals, +#: KDA, the router and fused-A GEMMs); the fused decode path adds its own. +REQUIRED_TRTLLM_OPS = ( + "attn_res_fwd", + "attn_res_rmsnorm_fwd", + "attn_res_add_rmsnorm_fwd", + "attn_res_add_rmsnorm_persistent_fwd", + "kda_prefill", + "kda_decode", + "kda_mtp_decode", + "dsv3_router_gemm_op", + "dsv3_fused_a_gemm_op", +) + +#: The engine surface the first forward checks before this target relies on it: per object, the attributes read. +#: A renamed field upstream fails here, loudly, instead of reading as a default. +REQUIRED_ENGINE_FIELDS = { + "attn_metadata": ( + "num_contexts", + "num_seqs", + "num_tokens", + "tokens_per_block", + "kv_cache_manager", + ), + "kv_cache_manager": ("enable_block_reuse",), +} + +#: Calls the generic path makes outside the catalog, declared so they are not consumed silently. A call leaves this +#: list when a catalog entry replaces it. +UNCERTIFIED_GENERIC_CALLS = ( + "tensorrt_llm._torch.models.modeling_kimi_linear.KimiLinearForCausalLM", +) + +# The fused decode path's token bounds per step: one token per request without speculation (the K3 decode kernels +# are built for up to 8 rows), and 1 + 7 drafts per request with DSpark at batch 8. +_FUSED_MAX_TOKENS = 8 +_FUSED_MAX_TOKENS_SPEC = 64 + +# The MLA generation kernels for K3's 96 query heads exist only at a 64-token page (the built-in model's own +# get_model_defaults sets it for the same reason). +_TOKENS_PER_BLOCK = 64 + +_LANG_PREFIX = "language_model." + + +def _text_model_config(model_config: ModelConfig) -> ModelConfig: + """The language model's ModelConfig: the checkpoint's text_config, with quant exclusions renamed to match. + + The checkpoint names its language-model modules `language_model.`; the text model's are ``, with + `layers.*` under `model.`. + """ + config = model_config.pretrained_config + text = copy.copy(model_config) + text._frozen = False + text.pretrained_config = config.text_config + excluded = text.quant_config.exclude_modules + if excluded: + text.quant_config = copy.copy(text.quant_config) + renamed = [] + for name in excluded: + if name.startswith(_LANG_PREFIX): + name = name[len(_LANG_PREFIX) :] + if name.startswith("layers."): + name = "model." + name + renamed.append(name) + text.quant_config.exclude_modules = renamed + text.skip_create_weights_in_init = True + text._frozen = True + return text + + +def _check_construction(model_config: ModelConfig) -> None: + """The settings this target is built for that are fixed before the first step.""" + capability = torch.cuda.get_device_capability() + assert capability == _SM, ( + f"this target is certified on sm_{_SM[0]}{_SM[1]}, running on sm_{capability[0]}{capability[1]}" + ) + mapping = model_config.mapping + topology = ( + mapping.world_size, + mapping.tp_size, + mapping.pp_size, + mapping.moe_tp_size, + mapping.moe_ep_size, + mapping.enable_attention_dp, + ) + assert topology == (16, 16, 1, 4, 4, False), ( + "the tp16_moetp4ep4 target needs world_size 16, tensor_parallel_size 16, pipeline_parallel_size 1, " + "moe_tensor_parallel_size 4, moe_expert_parallel_size 4 and enable_attention_dp false; the engine built " + f"(world, tp, pp, moe_tp, moe_ep, attention_dp) = {topology}" + ) + assert model_config.torch_dtype == torch.bfloat16, ( + f"this target computes in bf16; the engine resolved dtype {model_config.torch_dtype}" + ) + kv_algo = model_config.quant_config.kv_cache_quant_algo + assert kv_algo is None, ( + f"this target's MLA kernels read a bf16 KV pool; kv_cache_config.dtype resolved to {kv_algo}" + ) + strategy = model_config.allreduce_strategy + assert strategy in (AllReduceStrategy.AUTO, AllReduceStrategy.MNNVL), ( + f"this target runs its all-reduces over MNNVL; allreduce_strategy is {strategy.name}" + ) + + +@register_auto_model("ModelingV2KimiK3Mxfp4Sm100Tp16Moetp4ep4") +class ModelingV2KimiK3Mxfp4Sm100Tp16Moetp4ep4(KimiLinearForCausalLM): + """The registration shell: the built-in Kimi K3 text model as the generic path, behind this target's checks.""" + + @classmethod + def get_preferred_kv_cache_manager_version(cls, pretrained_config: Any = None) -> Literal["V2"]: + """The V2 hybrid manager holds the KDA states; the step contract requires it.""" + return "V2" + + def __init__(self, model_config: ModelConfig): + config = model_config.pretrained_config + assert getattr(config, "text_config", None) is not None, ( + "this target loads the KimiK3ForConditionalGeneration checkpoint, whose language model is its " + "text_config" + ) + _check_construction(model_config) + super().__init__(_text_model_config(model_config)) + self._step_checked = False + # The fused decode path and the state its kernels share, built in post_load_weights once the catalog entries + # it calls exist. None: every step takes the generic path. + self._fused_decode = None + # The executor reads generation settings (eos_token_id, ...) off the model config the engine holds, which + # must therefore be the text config, as the built-in wrapper leaves it. + model_config._frozen = False + model_config.pretrained_config = self.config + model_config._frozen = True + + def load_weights(self, weights, *args, **kwargs): + _weights.load(self, weights) + + def _check_step_contract(self, attn_metadata: AttentionMetadata) -> None: + """First-forward checks of the engine surface and the per-engine settings.""" + objects = { + "attn_metadata": attn_metadata, + "kv_cache_manager": attn_metadata.kv_cache_manager, + } + missing = [ + f"{owner}.{name}" + for owner, names in REQUIRED_ENGINE_FIELDS.items() + for name in names + if not hasattr(objects[owner], name) + ] + assert not missing, f"engine fields this target reads are missing: {missing}" + assert attn_metadata.tokens_per_block == _TOKENS_PER_BLOCK, ( + f"this target needs kv_cache_config.tokens_per_block {_TOKENS_PER_BLOCK}; the engine built " + f"{attn_metadata.tokens_per_block}" + ) + manager = attn_metadata.kv_cache_manager + assert isinstance(manager, MambaHybridCacheManagerV2), ( + "this target needs the V2 hybrid KV / state cache manager " + "(kv_cache_config.use_kv_cache_manager_v2); the engine built " + f"{type(manager).__name__}" + ) + assert not manager.enable_block_reuse, ( + "this target runs with kv_cache_config.enable_block_reuse false; the engine enabled it" + ) + self._step_checked = True + + def _step_path(self, attn_metadata: AttentionMetadata, spec_metadata) -> str: + """`"fused"` for a pure decode step within the fused path's bounds, else `"generic"`. + + Read on the host from per-step integers only. A CUDA graph is captured per decode batch shape, and every + input here is fixed by that shape, so a captured step and its replays take the same path. + """ + if attn_metadata.num_contexts: + return "generic" + bound = _FUSED_MAX_TOKENS if spec_metadata is None else _FUSED_MAX_TOKENS_SPEC + return "fused" if attn_metadata.num_tokens <= bound else "generic" + + def forward( + self, + attn_metadata: AttentionMetadata, + input_ids: Optional[torch.IntTensor] = None, + position_ids: Optional[torch.IntTensor] = None, + inputs_embeds: Optional[torch.FloatTensor] = None, + return_context_logits: bool = False, + spec_metadata=None, + resource_manager=None, + **kwargs, + ) -> torch.Tensor: + if kwargs.pop("multimodal_params", None): + raise ValueError( + "this Kimi K3 target is text only: it loads no vision tower, and a request carried image input" + ) + if not self._step_checked: + self._check_step_contract(attn_metadata) + if ( + self._fused_decode is not None + and self._step_path(attn_metadata, spec_metadata) == "fused" + ): + return self._fused_decode( + attn_metadata=attn_metadata, + input_ids=input_ids, + position_ids=position_ids, + spec_metadata=spec_metadata, + resource_manager=resource_manager, + **kwargs, + ) + return super().forward( + attn_metadata, + input_ids, + position_ids, + inputs_embeds, + return_context_logits, + spec_metadata, + resource_manager, + **kwargs, + ) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/weights.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/weights.py new file mode 100644 index 000000000000..b4177f698d8b --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/weights.py @@ -0,0 +1,42 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Weight loading: Kimi K3 (MXFP4) / sm_100 / tp16_moetp4ep4. + +The checkpoint is the vision-language wrapper's. `language_model.*` holds the language model; `vision_tower.*` and +`mm_projector.*` hold the vision tower and its projector. This target is text only, so the keys split three ways: + +* `language_model.*` goes to the built-in text model's loader with the prefix stripped. That loader streams the + routed experts one at a time, keeps this rank's slice (a quarter of the width of a quarter of the experts under + moe_tp 4 x moe_ep 4), and checks its own key coverage. +* The vision tower and the projector are a predicted non-load: listed here, never read. +* Any other key fails the load, naming it, rather than being dropped. +""" + +from tensorrt_llm._torch.models.checkpoints.base_weight_loader import ConsumableWeightsDict +from tensorrt_llm._torch.models.modeling_kimi_linear import KimiLinearForCausalLM +from tensorrt_llm._torch.models.modeling_utils import filter_weights + +_LANG_PREFIX = "language_model." + +# The vision tower's and the projector's key families: in the checkpoint, not loaded by this text-only target. +PREDICTED_NON_LOAD = ("vision_tower.", "mm_projector.") + + +def load(model, weights) -> None: + """Load the language model's weights into `model`, the target shell, and check every other key is predicted.""" + unknown = sorted( + k + for k in weights.keys() + if not k.startswith(_LANG_PREFIX) and not k.startswith(PREDICTED_NON_LOAD) + ) + assert not unknown, ( + f"{len(unknown)} checkpoint key(s) are neither language-model weights nor a predicted non-load, " + f"first {unknown[:5]}" + ) + lm_weights = ConsumableWeightsDict(filter_weights(_LANG_PREFIX[:-1], weights)) + assert len(lm_weights), f"the checkpoint has no {_LANG_PREFIX}* keys" + checkpoint_dir = getattr(weights, "checkpoint_dir", None) + if checkpoint_dir is not None: + lm_weights.checkpoint_dir = checkpoint_dir + lm_weights.checkpoint_prefix = _LANG_PREFIX + KimiLinearForCausalLM.load_weights(model, lm_weights) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/routing.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/routing.py new file mode 100644 index 000000000000..048c790a0b84 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/routing.py @@ -0,0 +1,123 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Where a KimiK3ForConditionalGeneration config lands. Read this file and you know. + +One forward-reading decision tree per architecture family: the criteria are +evaluated in the order a reader would ask them, and every branch that does not +end in a target returns None (in ``auto`` the engine then uses the built-in +Kimi K3 implementation; in ``require`` it raises, quoting the trace below). + +The checkpoint declares the vision-language wrapper, and the language model's +shape lives in its ``text_config``, so that is what this tree reads. The +targets are text-only: they load no vision tower and refuse image input. +""" + +from __future__ import annotations + +from typing import Any, Optional + +from ..._router_index import NULL_TRACE, ModelingV2Context, Trace + +# The one GPU architecture these targets are written for. sm is part of a +# target's identity, not a knob: a different SM is a different target. The K3 +# decode kernels (the tcgen05 GEMVs, the fused MoE) are built for GB200 only. +_SM = (10, 0) + +# Config-shape fingerprint -> checkpoint identity; see the note in the sibling +# gpt_oss routing module for what shape-sniffing does and does not pin. +# +# (num_hidden_layers, hidden_size, num_experts, routed_expert_hidden_size), all +# read from ``text_config``. The NVFP4 requant of the same checkpoint has this +# shape too; the quantization criterion below is what keeps it out. +_CHECKPOINTS = { + (93, 7168, 896, 3584): "kimi_k3_mxfp4", +} + +_TARGETS = { + ("kimi_k3_mxfp4", "tp16_moetp4ep4"): "ModelingV2KimiK3Mxfp4Sm100Tp16Moetp4ep4", +} + +# Synthetic architecture name -> the module whose import registers it. +TARGET_MODULES = { + "ModelingV2KimiK3Mxfp4Sm100Tp16Moetp4ep4": ( + "models.kimi_k3_vl.kimi_k3_mxfp4__sm_100__tp16_moetp4ep4.modeling" + ), +} + + +def _field(config: Any, name: str) -> Any: + """Read ``name`` off a sub-config that may be an object or a plain dict.""" + if isinstance(config, dict): + return config.get(name) + return getattr(config, name, None) + + +def _mxfp4(quant_config: Any) -> bool: + """Whether the checkpoint is the MXFP4 one these targets load. + + The MXFP4 checkpoint declares its quantization only inside + ``text_config.quantization_config`` (compressed-tensors), which the model + config does not surface, so it reads as no quantization at all. The NVFP4 + requant ships ``hf_quant_config.json`` and reads as MIXED_PRECISION; its + experts would land in a loader that reads packed MXFP4 tensors. + """ + return quant_config is None or quant_config.quant_algo is None + + +def _parallel(m) -> Optional[str]: + """Name the parallel topology, or None if no target implements it. + + Attention and the dense layers are sharded 16 ways; the routed experts are + split 4 ways by tensor and 4 ways by expert. The expert split decides which + expert weights each rank loads and into which shapes, so it selects a + target rather than a runtime branch. + """ + if ( + m.world_size == 16 + and m.tp_size == 16 + and m.pp_size == 1 + and m.moe_tp_size == 4 + and m.moe_ep_size == 4 + and not m.enable_attention_dp + ): + return "tp16_moetp4ep4" + return None + + +def route(ctx: ModelingV2Context, trace: Trace = NULL_TRACE) -> Optional[str]: + c, m = ctx.pretrained_config, ctx.mapping + + if not trace.check("sm", ctx.sm, ctx.sm == _SM): + return None + + t = getattr(c, "text_config", None) + if not trace.check("text_config", type(t).__name__, t is not None): + return None + + shape = tuple( + _field(t, name) + for name in ( + "num_hidden_layers", + "hidden_size", + "num_experts", + "routed_expert_hidden_size", + ) + ) + ckpt = trace.resolve("shape", shape, _CHECKPOINTS.get(shape)) + if ckpt is None: + return None + + quant_algo = getattr(ctx.quant_config, "quant_algo", None) + if not trace.check("quant", quant_algo, _mxfp4(ctx.quant_config)): + return None + + parallel = trace.resolve( + "parallel", + f"ws={m.world_size} tp={m.tp_size} pp={m.pp_size} moe_tp={m.moe_tp_size} " + f"moe_ep={m.moe_ep_size} attention_dp={m.enable_attention_dp}", + _parallel(m), + ) + if parallel is None: + return None + + return _TARGETS.get((ckpt, parallel)) diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py index f40cb6428713..9728ba55d985 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py @@ -235,3 +235,102 @@ def test_an_unknown_mode_raises_rather_than_falling_back(monkeypatch): monkeypatch.setenv(MODELING_V2_ENV, "yes") with pytest.raises(ValueError, match="not a modeling_v2 mode"): ModelingV2Mode.from_env() + + +_SM100 = (10, 0) +_TP16_MOETP4EP4 = dict(world_size=16, tp_size=16, moe_tp_size=4, moe_ep_size=4) +_K3_TARGET = "ModelingV2KimiK3Mxfp4Sm100Tp16Moetp4ep4" + + +def _k3_config(text_as_dict=False, **text_overrides): + """The shape Kimi K3's own config.json declares: a vision-language wrapper + whose language model is ``text_config``.""" + text = dict( + model_type="kimi_linear", + num_hidden_layers=93, + hidden_size=7168, + num_experts=896, + routed_expert_hidden_size=3584, + ) + text.update(text_overrides) + return PretrainedConfig( + architectures=["KimiK3ForConditionalGeneration"], + model_type="kimi_k3", + text_config=text if text_as_dict else PretrainedConfig(**text), + ) + + +@pytest.fixture +def _on_sm100(monkeypatch): + """Route as if this were a GB200.""" + monkeypatch.setattr(torch.cuda, "get_device_capability", lambda *a, **k: _SM100) + + +@pytest.mark.usefixtures("_on_sm100") +@pytest.mark.parametrize("mode", ["auto", "require"]) +@pytest.mark.parametrize("text_as_dict", [False, True], ids=["text-config", "text-dict"]) +def test_kimi_k3_tp16_moetp4ep4_matches(monkeypatch, mode, text_as_dict): + _set_mode(monkeypatch, mode) + config = _model_config(_k3_config(text_as_dict=text_as_dict), **_TP16_MOETP4EP4) + assert modeling_v2_resolve(config) == _K3_TARGET + + +@pytest.mark.usefixtures("_on_sm100") +def test_kimi_k3_target_registers_and_counts_as_external(): + config = _model_config(_k3_config(), **_TP16_MOETP4EP4) + cls = get_registered_model_class(modeling_v2_resolve(config)) + assert cls is not None and cls.__name__ == _K3_TARGET + assert not _is_builtin_model_class(cls) + + +@pytest.mark.usefixtures("_on_sm100") +def test_kimi_k3_nvfp4_requant_does_not_match(monkeypatch): + """The NVFP4 requant has the MXFP4 checkpoint's shape; its quantization + is what keeps it out of a target whose loader reads packed MXFP4 experts.""" + from tensorrt_llm.models.modeling_utils import QuantConfig + from tensorrt_llm.quantization.mode import QuantAlgo + + config = ModelConfig( + pretrained_config=_k3_config(), + mapping=Mapping(**_TP16_MOETP4EP4), + quant_config=QuantConfig(quant_algo=QuantAlgo.MIXED_PRECISION), + ) + assert modeling_v2_resolve(config) is None + + _set_mode(monkeypatch, "require") + with pytest.raises(ValueError, match="quant"): + modeling_v2_resolve(config) + + +@pytest.mark.usefixtures("_on_sm100") +@pytest.mark.parametrize( + "text_overrides, mapping_kwargs, missed", + [ + # a Kimi-family checkpoint of another depth + (dict(num_hidden_layers=61), _TP16_MOETP4EP4, "shape"), + # route B's expert split: experts tensor-parallel 16 ways, no target yet + (dict(), dict(world_size=16, tp_size=16, moe_tp_size=16, moe_ep_size=1), "parallel"), + # attention data parallelism splits the requests, not the heads + (dict(), dict(_TP16_MOETP4EP4, enable_attention_dp=True), "parallel"), + # one tray instead of four + (dict(), dict(world_size=4, tp_size=4, moe_tp_size=1, moe_ep_size=4), "parallel"), + ], +) +def test_kimi_k3_near_misses_do_not_match(monkeypatch, text_overrides, mapping_kwargs, missed): + config = _model_config(_k3_config(**text_overrides), **mapping_kwargs) + assert modeling_v2_resolve(config) is None + + _set_mode(monkeypatch, "require") + with pytest.raises(ValueError, match=missed): + modeling_v2_resolve(config) + + +def test_kimi_k3_on_another_sm_does_not_match(monkeypatch): + """The autouse fixture routes as a GB300 (sm 10.3); the K3 target is sm 10.0 only.""" + config = _model_config(_k3_config(), **_TP16_MOETP4EP4) + assert modeling_v2_resolve(config) is None + + _set_mode(monkeypatch, "require") + with pytest.raises(ValueError) as excinfo: + modeling_v2_resolve(config) + assert "sm" in str(excinfo.value) and "(10, 3)" in str(excinfo.value) From b45a9468b68ce5504f7ab3b89644684a9dc23325 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 18:15:12 -0700 Subject: [PATCH 020/161] [None][feat] Kimi K3 MLA attention: a caller-owned workspace trtllm::k3_mla_attn, k3_mla_attn_out and k3_mla_attn_vb_out kept their workspace (the per-CTA partials and the no-cluster mode's exchange and arrival counters) in a module-level dict keyed by device and head-group count, allocated on the first eager call. They now take it as an argument, named in mutates_args: make_attn_workspace(device, groups) allocates and arms one (the counters zeroed, as before) and refuses to run under CUDA-graph capture; a call checks that the workspace is the one for its device and head-group count and raises ValueError before launching otherwise. The kernel and the words it reads and writes are unchanged. The op test passes a workspace; the workspace poison and counter wrap tests now run on their own explicit workspaces, and a malformed workspace is refused without being touched. Signed-off-by: Vasanth Sabavat --- .../_torch/cute_dsl_kernels/k3_mla/op.py | 96 +++++++++++++------ .../kimi_k3/test_k3_mla_attn.py | 94 +++++++++++------- 2 files changed, 126 insertions(+), 64 deletions(-) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/op.py index 6d953a76427e..96cb022aa777 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/op.py @@ -22,8 +22,9 @@ ``trtllm::k3_mla_q``: the decode query path (M <= 64; 6 heads per rank at TP16, 24 at TP4): q_a RMSNorm, q_b projection and k_b absorption in one launch, producing the attention's ``fused_q`` [M, heads * 576] = [q_nope @ W_kb^T | q_pe] per head; ``trtllm::k3_mla_qkv`` also stores the KV half (kv_a RMSNorm, rope columns) into the paged latent cache. -``trtllm::k3_mla_attn`` and its ``_out`` / ``_vb_out`` forms: the attention of every request over its pages. Compiled -on the first call for its shapes, which must happen outside CUDA-graph capture. +``trtllm::k3_mla_attn`` and its ``_out`` / ``_vb_out`` forms: the attention of every request over its pages, over a +caller-owned workspace (:func:`make_attn_workspace`). Compiled on the first call for its shapes, which must happen +outside CUDA-graph capture. """ from __future__ import annotations @@ -347,29 +348,48 @@ def _(ag, w_qa, eps, w_qb, w_kb, w_kv, kv_eps, kv_out, trigger_early=True): # trtllm::k3_mla_attn: decode attention over the paged latent cache (R <= 8 requests of T <= 8 tokens, one cluster of # 16 CTAs per request and 6 heads) # --------------------------------------------------------------------------------------------------------------- -_workspaces: Dict[tuple, torch.Tensor] = {} +def attn_workspace_elems(groups: int) -> int: + """fp16 elements of a ``k3_mla_attn`` workspace for calls of ``groups`` head groups (heads / 6): the per-CTA + partials of MAX_REQUESTS x groups x 16 slots, then the no_cluster mode's (m, l) exchange and arrival counters.""" + from . import k3_mla_attn_kernel as kernel + + return ( + kernel.MAX_REQUESTS * groups * kernel.CLUSTER * kernel.WS_SLOT_ELEMS + + kernel.ws_sync_elems(groups) + ) -def _attn_workspace(device: torch.device, groups: int) -> torch.Tensor: - """The per-CTA partials of every (request, head group), then the no_cluster mode's (m, l) exchange and arrival - counters (zeroed): allocated once, for MAX_REQUESTS, on the first call.""" +def make_attn_workspace(device: torch.device, groups: int) -> torch.Tensor: + """A new workspace for the ``k3_mla_attn`` calls of ``groups`` head groups on ``device``: fp16 + [attn_workspace_elems(groups)], the partial slots uninitialized (a call reads only the words it wrote) and the + no_cluster tail zeroed (the arrival counters start at 0). It allocates, so it refuses to run under CUDA-graph + capture; the zeroing is ordered on the device's current stream.""" from . import k3_mla_attn_kernel as kernel - key = (device.index, groups) - ws = _workspaces.get(key) - if ws is None: - if torch.cuda.is_current_stream_capturing(): - raise RuntimeError( - "trtllm::k3_mla_attn must run once outside CUDA-graph capture first (it allocates its workspace)." - ) - slots = kernel.MAX_REQUESTS * groups * kernel.CLUSTER - ws = _workspaces[key] = torch.empty( - slots * kernel.WS_SLOT_ELEMS + kernel.ws_sync_elems(groups), dtype=torch.float16, device=device - ) - ws[slots * kernel.WS_SLOT_ELEMS :].zero_() + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError("k3_mla_attn workspaces allocate: make them outside CUDA-graph capture.") + if groups < 1: + raise ValueError(f"k3_mla_attn workspace: {groups} head groups") + ws = torch.empty(attn_workspace_elems(groups), dtype=torch.float16, device=device) + ws[kernel.MAX_REQUESTS * groups * kernel.CLUSTER * kernel.WS_SLOT_ELEMS :].zero_() return ws +def _check_attn_workspace(workspace: torch.Tensor, device: torch.device, groups: int) -> None: + elems = attn_workspace_elems(groups) + if not ( + workspace.dtype == torch.float16 + and workspace.dim() == 1 + and workspace.is_contiguous() + and workspace.numel() == elems + and workspace.device == device + ): + raise ValueError( + f"k3_mla_attn: workspace {tuple(workspace.shape)} {workspace.dtype} on {workspace.device} is not one for " + f"{groups} head group(s) on {device}: fp16 [{elems}] (make_attn_workspace)" + ) + + def supports_attn( q: torch.Tensor, pool: torch.Tensor, @@ -408,6 +428,7 @@ def _launch_attn( page_table, seq_len, softmax_scale, + workspace, out=None, page_offset=0, w_vb=None, @@ -451,7 +472,7 @@ def _launch_attn( f"k3_mla_attn: gate {tuple(gate.shape)} {gate.dtype} col0 {gate_col0} does not fit the v_b output" ) groups = total_heads // kernel.HEADS - ws_o = _attn_workspace(q.device, groups) + _check_attn_workspace(workspace, q.device, groups) # More 16-CTA clusters than co-reside would run a second wave: launch without a cluster instead when every CTA fits # on the SMs at once (one per SM; its waits are spins). clusters = num_requests * groups @@ -472,8 +493,8 @@ def _launch_attn( total_rows = pool.numel() // row_stride gate_flat = gate.as_strided((gate.numel(),), (1,)) if apply_gate else q.view(-1) gate_ld = gate.stride(0) if apply_gate else 0 - args = (_arg(q.view(-1)), _arg(_pool_base(pool, row_stride)), _arg(page_rows, 4), _arg(seq_len, 4), _arg(ws_o), - _arg(out.view(-1)), _arg((w_vb if fuse_vb else q).view(-1)), _arg(gate_flat)) # fmt: skip + args = (_arg(q.view(-1)), _arg(_pool_base(pool, row_stride)), _arg(page_rows, 4), _arg(seq_len, 4), + _arg(workspace), _arg(out.view(-1)), _arg((w_vb if fuse_vb else q).view(-1)), _arg(gate_flat)) # fmt: skip stream = cuda_driver.CUstream(torch.cuda.current_stream(q.device).cuda_stream) use_pdl = _use_pdl() scale_log2 = float(softmax_scale) * kernel.LOG2E @@ -509,7 +530,7 @@ def _launch_attn( return out -@torch.library.custom_op("trtllm::k3_mla_attn", mutates_args=()) +@torch.library.custom_op("trtllm::k3_mla_attn", mutates_args=("workspace",)) def k3_mla_attn( q: torch.Tensor, pool: torch.Tensor, @@ -517,6 +538,7 @@ def k3_mla_attn( page_table: torch.Tensor, seq_len: torch.Tensor, softmax_scale: float, + workspace: torch.Tensor, ) -> torch.Tensor: """MLA decode attention of R <= 8 requests of T <= 8 tokens: ``q`` [M = R T, heads * 576] (``fused_q``, heads a multiple of 6, request-major) against the paged latent cache ``pool`` (flat bf16; row i of page p at ``(p * 64 + @@ -524,11 +546,15 @@ def k3_mla_attn( = L_i (rows including its T new ones; see the module docstring), causal bottom-right (token t of request i sees rows <= L_i - T + t). Returns ``[M, heads * 512]`` bf16. The page table and lengths are read before the grid dependency wait (they must be written before the CUDA graph runs); q and the pages holding rows >= L_i - T after - it. Request i's rows are computed as the R = 1 call on its own rows, pages and length would compute them.""" - return _launch_attn(q, pool, row_stride, page_table, seq_len, softmax_scale) + it. Request i's rows are computed as the R = 1 call on its own rows, pages and length would compute them. + ``workspace``: from :func:`make_attn_workspace` for this device and heads / 6 head groups. A call writes and reads + its partials there, all after its grid dependency wait, and in the no_cluster mode (more than CLUSTER_WAVE + clusters, all of whose CTAs fit on the SMs) adds 16 to the arrival counters of its requests' head groups; calls + on one workspace must run one at a time.""" + return _launch_attn(q, pool, row_stride, page_table, seq_len, softmax_scale, workspace) -@torch.library.custom_op("trtllm::k3_mla_attn_out", mutates_args=("out",)) +@torch.library.custom_op("trtllm::k3_mla_attn_out", mutates_args=("out", "workspace")) def k3_mla_attn_out( q: torch.Tensor, pool: torch.Tensor, @@ -538,15 +564,24 @@ def k3_mla_attn_out( seq_len: torch.Tensor, softmax_scale: float, out: torch.Tensor, + workspace: torch.Tensor, ) -> None: """``k3_mla_attn`` into ``out`` (dense [M, heads * 512] bf16), with ``page_offset`` added to every page-table entry (the layer's slot in a layer-interleaved pool).""" _launch_attn( - q, pool, row_stride, page_table, seq_len, softmax_scale, out=out, page_offset=page_offset + q, + pool, + row_stride, + page_table, + seq_len, + softmax_scale, + workspace, + out=out, + page_offset=page_offset, ) -@torch.library.custom_op("trtllm::k3_mla_attn_vb_out", mutates_args=("out",)) +@torch.library.custom_op("trtllm::k3_mla_attn_vb_out", mutates_args=("out", "workspace")) def k3_mla_attn_vb_out( q: torch.Tensor, pool: torch.Tensor, @@ -557,6 +592,7 @@ def k3_mla_attn_vb_out( softmax_scale: float, w_vb: torch.Tensor, out: torch.Tensor, + workspace: torch.Tensor, gate: Optional[torch.Tensor] = None, gate_col0: int = 0, ) -> None: @@ -565,11 +601,11 @@ def k3_mla_attn_vb_out( (bf16 [M, C], sigmoid of head h's gate at columns ``gate_col0 + 128 h``) the output is ``bf16(y * s)``, the unfused output gate.""" _launch_attn( - q, pool, row_stride, page_table, seq_len, softmax_scale, out=out, page_offset=page_offset, w_vb=w_vb, - gate=gate, gate_col0=gate_col0, + q, pool, row_stride, page_table, seq_len, softmax_scale, workspace, out=out, page_offset=page_offset, + w_vb=w_vb, gate=gate, gate_col0=gate_col0, ) # fmt: skip @k3_mla_attn.register_fake -def _(q, pool, row_stride, page_table, seq_len, softmax_scale): +def _(q, pool, row_stride, page_table, seq_len, softmax_scale, workspace): return q.new_empty((q.shape[0], q.shape[1] // 576 * 512), dtype=torch.bfloat16) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_attn.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_attn.py index 2b2afe5e5bc5..843ca267e52b 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_attn.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_attn.py @@ -114,18 +114,37 @@ def _max_rel(a, b, tokens, num_requests): return err -def _attn_out(q, pool, row_stride, page_table, page_offset, seq_len): +_workspaces = {} + + +def _workspace(heads): + """This module's attention workspace for ``heads`` (heads / 6 head groups): one per group count on the current + device, made on first use and shared by every call of that shape, as a target shares one across its layers.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import op + + groups = heads // H + if groups not in _workspaces: + _workspaces[groups] = op.make_attn_workspace( + torch.device("cuda", torch.cuda.current_device()), groups + ) + return _workspaces[groups] + + +def _attn_out(q, pool, row_stride, page_table, page_offset, seq_len, workspace=None): from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import op # noqa: F401 m, heads = q.shape[0], q.shape[1] out = torch.empty(m, heads * LAT, dtype=torch.bfloat16, device="cuda") torch.ops.trtllm.k3_mla_attn_out( - q.reshape(m, -1), pool.view(-1), row_stride, page_table, page_offset, seq_len, SCALE, out - ) + q.reshape(m, -1), pool.view(-1), row_stride, page_table, page_offset, seq_len, SCALE, out, + _workspace(heads) if workspace is None else workspace, + ) # fmt: skip return out.view(m, heads, LAT) -def _attn_vb(q, pool, row_stride, page_table, page_offset, seq_len, w_vb, gate=None): +def _attn_vb( + q, pool, row_stride, page_table, page_offset, seq_len, w_vb, gate=None, workspace=None +): from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import op # noqa: F401 m, heads = q.shape[0], q.shape[1] @@ -140,6 +159,7 @@ def _attn_vb(q, pool, row_stride, page_table, page_offset, seq_len, w_vb, gate=N SCALE, w_vb, y, + _workspace(heads) if workspace is None else workspace, gate, GATE_COL0, ) @@ -179,7 +199,9 @@ def test_attn_returns(num_requests, tokens): q, pool, page_table, _, seq_len = _make_case(5 + num_requests, num_requests, tokens) m = q.shape[0] table = page_table[0] if num_requests == 1 else page_table - out = torch.ops.trtllm.k3_mla_attn(q.view(m, -1), pool.view(-1), DQK, table, seq_len, SCALE) + out = torch.ops.trtllm.k3_mla_attn( + q.view(m, -1), pool.view(-1), DQK, table, seq_len, SCALE, _workspace(H) + ) ref = _attn_out(q, pool, DQK, page_table, 0, seq_len) assert torch.equal(_bits(out), _bits(ref.view(m, -1))) @@ -251,7 +273,8 @@ def test_attn_vs_stock(num_requests, tokens): def test_attn_rejects(): - """Calls outside R <= 8 requests of T <= 8 tokens, or page-table rows / lengths that do not match, are refused.""" + """Calls outside R <= 8 requests of T <= 8 tokens, or page-table rows / lengths that do not match, are refused; + so is a workspace that is not one for the call's head groups and device (its words are left as they were).""" from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import op q, pool, page_table, _, seq_len = _make_case(1, 2, 4) @@ -267,6 +290,13 @@ def test_attn_rejects(): assert not op.supports_attn(q9.view(9, -1), flat, DQK, table9, len9) # R = 9 with pytest.raises(ValueError): _attn_out(q[:7], pool, DQK, page_table, 0, seq_len) + other = _workspace(4 * H) # four head groups' layout; the call has one + before = other.clone() + for workspace in (other, other.float(), _workspace(H)[:-8]): + with pytest.raises(ValueError, match="workspace"): + _attn_out(q, pool, DQK, page_table, 0, seq_len, workspace) + torch.cuda.synchronize() + assert torch.equal(other.view(torch.int16), before.view(torch.int16)) # Both launch modes (clusters; no_cluster past CLUSTER_WAVE clusters) at the TP16 and TP4 head counts, with folds @@ -274,14 +304,12 @@ def test_attn_rejects(): POISON_CASES = [(1, 1, H), (3, 5, H), (8, 1, H), (8, 8, H), (1, 8, 4 * H), (2, 4, 4 * H)] -def _fill_workspace(heads, value): - """Fill the data words of the attention workspace: the per-CTA partial slots (fp16) and the no_cluster (m, l) +def _fill_workspace(ws, heads, value): + """Fill the data words of an attention workspace: the per-CTA partial slots (fp16) and the no_cluster (m, l) exchange (fp32). The arrival counters after them keep their values.""" from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import k3_mla_attn_kernel as kernel - from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import op groups = heads // kernel.HEADS - ws = op._attn_workspace(torch.device("cuda", torch.cuda.current_device()), groups) slots = kernel.MAX_REQUESTS * groups * kernel.CLUSTER partials = slots * kernel.WS_SLOT_ELEMS exchange = slots * kernel.ROWS * 2 # fp32 words @@ -296,22 +324,22 @@ def test_attn_workspace_poison(num_requests, tokens, heads): """A call reads only workspace words it wrote itself: with the partial slots and the (m, l) exchange refilled with NaN before each call, k3_mla_attn_out and the gated k3_mla_attn_vb_out give the bits of the same calls on a zero-filled workspace.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import op + seed = 13 * num_requests + tokens q, pool, page_table, page_offset, seq_len = _make_case(seed, num_requests, tokens, heads) gen = torch.Generator(device="cuda").manual_seed(seed + 1) m = q.shape[0] w_vb = (torch.randn(heads, V, LAT, generator=gen, device="cuda") * 0.05).bfloat16() ag = torch.rand(m, GATE_COL0 + heads * V, generator=gen, device="cuda").bfloat16() + ws = op.make_attn_workspace(torch.device("cuda", torch.cuda.current_device()), heads // H) outs = {} - try: - for value in (0.0, float("nan")): - _fill_workspace(heads, value) - o = _attn_out(q, pool, DQK, page_table, page_offset, seq_len) - _fill_workspace(heads, value) - y = _attn_vb(q, pool, DQK, page_table, page_offset, seq_len, w_vb, ag) - outs[value == 0.0] = (o, y) - finally: - _fill_workspace(heads, 0.0) + for value in (0.0, float("nan")): + _fill_workspace(ws, heads, value) + o = _attn_out(q, pool, DQK, page_table, page_offset, seq_len, ws) + _fill_workspace(ws, heads, value) + y = _attn_vb(q, pool, DQK, page_table, page_offset, seq_len, w_vb, ag, ws) + outs[value == 0.0] = (o, y) torch.cuda.synchronize() for got, want in zip(outs[False], outs[True]): assert not torch.isnan(got).any() @@ -322,13 +350,11 @@ def test_attn_workspace_poison(num_requests, tokens, heads): WRAP_CASES = [(8, 1, H), (8, 8, H), (2, 4, 4 * H)] -def _set_counters(heads, value): - """Set every no_cluster arrival counter of the attention workspace (int32 words after the (m, l) exchange).""" +def _set_counters(ws, heads, value): + """Set every no_cluster arrival counter of an attention workspace (int32 words after the (m, l) exchange).""" from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import k3_mla_attn_kernel as kernel - from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import op groups = heads // kernel.HEADS - ws = op._attn_workspace(torch.device("cuda", torch.cuda.current_device()), groups) slots = kernel.MAX_REQUESTS * groups * kernel.CLUSTER start = ( slots * kernel.WS_SLOT_ELEMS + 2 * slots * kernel.ROWS * 2 @@ -342,22 +368,22 @@ def _set_counters(heads, value): def test_attn_counter_wrap(num_requests, tokens, heads): """The no_cluster arrival counters only grow (16 per launch) and are compared by signed difference: calls with the counters just below the int32 wrap (2^31 - 16, and -16 just below 0) give the bits of calls on zeroed ones.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import op + seed = 19 * num_requests + tokens q, pool, page_table, page_offset, seq_len = _make_case(seed, num_requests, tokens, heads) gen = torch.Generator(device="cuda").manual_seed(seed + 1) w_vb = (torch.randn(heads, V, LAT, generator=gen, device="cuda") * 0.05).bfloat16() + ws = op.make_attn_workspace(torch.device("cuda", torch.cuda.current_device()), heads // H) outs = {} - try: - for start in (0, 2**31 - 16, -16): - _set_counters(heads, start) - # Three launches: each crosses the wrap point once the counters start 16 below it. - outs[start] = [ - _attn_out(q, pool, DQK, page_table, page_offset, seq_len), - _attn_vb(q, pool, DQK, page_table, page_offset, seq_len, w_vb), - _attn_out(q, pool, DQK, page_table, page_offset, seq_len), - ] - finally: - _set_counters(heads, 0) + for start in (0, 2**31 - 16, -16): + _set_counters(ws, heads, start) + # Three launches: each crosses the wrap point once the counters start 16 below it. + outs[start] = [ + _attn_out(q, pool, DQK, page_table, page_offset, seq_len, ws), + _attn_vb(q, pool, DQK, page_table, page_offset, seq_len, w_vb, workspace=ws), + _attn_out(q, pool, DQK, page_table, page_offset, seq_len, ws), + ] torch.cuda.synchronize() for start in (2**31 - 16, -16): for got, want in zip(outs[start], outs[0]): From f51e780c36108789382c4fe121e96424a4a2c05d Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 18:15:13 -0700 Subject: [PATCH 021/161] [None][feat] modeling_v2 catalog: Kimi K3's stateful KDA and MLA decode entries New category ssm/ (KDA decode) and two attention/ entries, each a contract, a wrapper and a GPU test that drives the op on a real cache manager: - ssm/kda_decode: the one-token KDA decode (C++), on the V2 hybrid manager's pools; - ssm/k3_kda_attn (+ k3_kda_qkvg), ssm/k3_kda_decode_attn: Kimi K3's fused KDA projection with the verify or the plain decode in one launch, over a caller-owned K3KdaBuffers (the projection's Lamport set, one per device, shared by both ops); - ssm/k3_kda_verify: the KDA verify from the per-token states that the V2 hybrid manager keeps per slot; - attention/k3_mla_qkv (+ k3_mla_q, k3_mla_qkv_out): the MLA decode query path and the latent KV append into the paged cache; - attention/k3_mla_attn_vb_out (+ k3_mla_attn_out, k3_mla_attn): MLA decode attention with v_b and the output gate, over a caller-owned K3MlaAttnWorkspace. Every stateful contract has a ## State section (contents and size, creator, sharing, call order, what a later launch reads, re-arm), and every test runs layers x steps on one state object, a CUDA-graph capture replayed with rewritten inputs, two objects interleaved and a negative control. Signed-off-by: Vasanth Sabavat --- .../catalog/attention/k3_mla_attn_vb_out.md | 177 +++++++ .../catalog/attention/k3_mla_attn_vb_out.py | 79 ++++ .../attention/k3_mla_attn_workspace.py | 45 ++ .../catalog/attention/k3_mla_qkv.md | 154 ++++++ .../catalog/attention/k3_mla_qkv.py | 74 +++ .../modeling_v2/catalog/index.yaml | 29 ++ .../modeling_v2/catalog/ssm/__init__.py | 3 + .../modeling_v2/catalog/ssm/k3_kda_attn.md | 195 ++++++++ .../modeling_v2/catalog/ssm/k3_kda_attn.py | 75 +++ .../modeling_v2/catalog/ssm/k3_kda_buffers.py | 52 ++ .../catalog/ssm/k3_kda_decode_attn.md | 127 +++++ .../catalog/ssm/k3_kda_decode_attn.py | 58 +++ .../modeling_v2/catalog/ssm/k3_kda_verify.md | 144 ++++++ .../modeling_v2/catalog/ssm/k3_kda_verify.py | 60 +++ .../modeling_v2/catalog/ssm/kda_decode.md | 135 ++++++ .../modeling_v2/catalog/ssm/kda_decode.py | 74 +++ .../test_modeling_v2_k3_mla_attn_vb_out.py | 445 ++++++++++++++++++ .../attention/test_modeling_v2_k3_mla_qkv.py | 366 ++++++++++++++ .../_torch/modeling_v2/ssm/_kda_cells.py | 329 +++++++++++++ .../ssm/test_modeling_v2_k3_kda_attn.py | 292 ++++++++++++ .../test_modeling_v2_k3_kda_decode_attn.py | 358 ++++++++++++++ .../ssm/test_modeling_v2_k3_kda_verify.py | 202 ++++++++ .../ssm/test_modeling_v2_kda_decode.py | 194 ++++++++ 23 files changed, 3667 insertions(+) create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_attn_vb_out.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_attn_vb_out.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_attn_workspace.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_qkv.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_qkv.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/__init__.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_attn.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_attn.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_buffers.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_decode_attn.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_decode_attn.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_verify.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_verify.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/kda_decode.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/kda_decode.py create mode 100644 tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_mla_attn_vb_out.py create mode 100644 tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_mla_qkv.py create mode 100644 tests/unittest/_torch/modeling_v2/ssm/_kda_cells.py create mode 100644 tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_k3_kda_attn.py create mode 100644 tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_k3_kda_decode_attn.py create mode 100644 tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_k3_kda_verify.py create mode 100644 tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_kda_decode.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_attn_vb_out.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_attn_vb_out.md new file mode 100644 index 000000000000..d6cf18aa0050 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_attn_vb_out.md @@ -0,0 +1,177 @@ +--- +receipts: {} +--- + +# k3_mla_attn_vb_out + +**Wraps** `torch.ops.trtllm.k3_mla_attn_vb_out` (one call). The same contract covers two sibling wrappers, one call +each: `k3_mla_attn_out` (`torch.ops.trtllm.k3_mla_attn_out`, the attention output itself) and `k3_mla_attn` +(`torch.ops.trtllm.k3_mla_attn`, the same into a new tensor, page offset 0). All three take a caller-owned +`K3MlaAttnWorkspace` (*State*). + +## Semantics + +Kimi K3's MLA decode attention for `R <= 8` requests of `T <= 8` tokens over the paged latent cache, with v_b and +the output gate in the same launch. `q` is `fused_q` (from `attention/k3_mla_qkv`), request-major: rows +`i T .. i T + T - 1` are request `i`'s tokens. Request `i`'s cache rows `kv_i[k]`, `k < L_i = seq_len[i]`, are row +`(page_table[i][k // 64] + page_offset) * 64 + k % 64` of `pool` (512 latent columns, then 64 rope columns). Per +request `i`, token `t` and head `h`: + +``` +o[t, h] = softmax_k( softmax_scale * q[t, h] . kv_i[k] ) @ kv_i[k, :512] k <= L_i - T + t (bottom-right causal) +y[t, h] = bf16( bf16(o[t, h]) @ w_vb[h]^T ) # v_b, 128 per head +out[t, 128 h : 128 h + 128] = y[t, h] # without gate + = bf16(y[t, h] * s[t, h]) # with gate: s = gate[t, gate_col0 + 128 h .. + 128] +``` + +`gate` holds the output gate's sigmoid values (the op multiplies, it does not apply the sigmoid); the gated output is +bit-identical to `bf16(plain * s)` (certified). `k3_mla_attn_out` writes `o` itself (`[M, heads * 512]`, bf16) into +`out`; `k3_mla_attn` returns it in a new tensor and is bit-identical to `k3_mla_attn_out` at page offset 0 +(certified). + +How it computes: one cluster of 16 CTAs per (request, group of 6 heads), split-KV over 128-row tiles; fp32 softmax; +each CTA's partial `O / l` (a convex combination of V rows) is kept in fp16, and the 16 partials are merged in fp32 +in a fixed order. More than `CLUSTER_WAVE` = 7 clusters (`R x heads / 6`) would not co-reside on GB200; when all +their CTAs fit on the SMs at once (one per SM: at most 9 clusters on 148 SMs) the launch takes the **no-cluster +mode**, which exchanges the merge statistics through the workspace instead of cluster shared memory. The merge sums +the same values in the same order in both modes, so the outputs do not depend on the mode. Request `i`'s rows are +computed as the one-request call on its own rows, pages and length would compute them (the op test, +`test_k3_mla_attn.py`). + +Accuracy (certified bound): every request's output within `1e-2` (max relative, per request) of a float64 +reference: the attention in float64 with the causal mask above, rounded to bf16; v_b in float64; the gate applied +to the bf16-rounded product. + +Fusion boundary. Inside: the attention over the cache, v_b, the output gate. Outside: `fused_q` and the cache append +(`attention/k3_mla_qkv`, which must store the step's rows before this call reads them), and the output projection +that consumes `out`. + +## Signature + +```python +def k3_mla_attn_vb_out( + q: torch.Tensor, + pool: torch.Tensor, + row_stride: int, + page_table: torch.Tensor, + page_offset: int, + seq_len: torch.Tensor, + softmax_scale: float, + w_vb: torch.Tensor, + out: torch.Tensor, + workspace: K3MlaAttnWorkspace, + gate: Optional[torch.Tensor] = None, + gate_col0: int = 0, +) -> None + +def k3_mla_attn_out(q, pool, row_stride, page_table, page_offset, seq_len, softmax_scale, out, + workspace: K3MlaAttnWorkspace) -> None + +def k3_mla_attn(q, pool, row_stride, page_table, seq_len, softmax_scale, + workspace: K3MlaAttnWorkspace) -> torch.Tensor +``` + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `q` | `[M, heads * 576]`, `M = R T`, `R <= 8`, `T <= 8`, `heads` a multiple of 6 (6 at TP16, 24 at TP4) | bf16 | contiguous | CUDA | +| `pool` | the layer's paged latent pool, flat | bf16 | contiguous | CUDA | +| `row_stride` | elements per cache row, `>= 576`, a multiple of 8 | Python int | — | — | +| `page_table` | `[R, W]` with unit column stride and any row stride `>= W` (rows of `kv_cache_block_offsets`), or `[W]` when `R = 1` | int32 | see shape | CUDA | +| `page_offset` | the layer's slot in a layer-interleaved pool | Python int | — | — | +| `seq_len` | `[R]`, `L_i >= T`: each request's length including its `T` new rows | int32 | contiguous | CUDA | +| `softmax_scale` | `1 / (sqrt(qk_nope + qk_rope) * q_scaling)` (the view's) | Python float | — | — | +| `w_vb` | `[heads, 128, 512]` (`v_b_proj`) | bf16 | contiguous | CUDA | +| `out` | `[M, heads * 128]` (`k3_mla_attn_out`: `[M, heads * 512]`) | bf16 | contiguous | CUDA | +| `workspace` | a `K3MlaAttnWorkspace` of `q`'s device for `heads / 6` head groups | — | — | — | +| `gate` | `None`, or `[M, C]` with unit column stride and `gate_col0 + heads * 128 <= C` | bf16 | rows of any stride | CUDA | +| `gate_col0` | `>= 0` | Python int | — | — | +| returns | `None` (`k3_mla_attn`: `[M, heads * 512]`, newly allocated) | bf16 | contiguous | = `q.device` | + +Certified at `(R, T, heads)` = `(3, 8, 6)` (3 clusters), `(8, 1, 6)` and `(2, 8, 24)` (8 clusters: the no-cluster +mode), with 64 to 4100 cached rows per request before a step (pages crossed, several tiles per CTA), over a real +`KVCacheManager` (below). + +## State + +**Object.** `K3MlaAttnWorkspace` (`catalog/attention/k3_mla_attn_workspace.py`), one per device and head-group +count (`heads / 6`: 1 at TP16, 4 at TP4), owned by the caller. The pool is the KV cache manager's (as +`attention/k3_mla_qkv`); this call only reads it. + +**Contents and size.** One fp16 buffer, `attn_workspace_elems(groups)` elements: the per-CTA partials of 8 requests +x `groups` x 16 CTAs (64 KiB each: 8 MiB per head group), then the no-cluster mode's fp32 `(m, l)` exchange (48 KiB +per group) and two int32 arrival counters per (request, head group), each on its own 128-byte line (2 KiB per +group): 8,439,808 bytes per head group, 33,759,232 at TP4. The size does not depend on `R` or `T`: every call fits. + +**Who creates it, and when.** The target, in `post_load_weights`, with `K3MlaAttnWorkspace.create(device, groups)`: + +- eager: it allocates, so it refuses to run under CUDA-graph capture (`RuntimeError`, certified); +- it zeroes the counters and leaves the partials as allocated (a call reads only words it wrote, below), then + synchronizes the device, so the workspace is ready on any stream. + +**Which ops may share one object.** All three forms, for every MLA layer of the device with the same head-group +count: one workspace serves a model's layers. Two workspaces are independent (certified: the layers' calls +alternating between two workspaces over three no-cluster steps; every output bit-identical to the same call on a new +workspace, each workspace's counters counting only its own launches). + +**Call-order invariant.** Calls on one workspace run one at a time. Every access a call makes to the workspace +(partials, exchange, counters) follows its grid-dependency wait, so a call may directly follow another on the same +stream; the earlier call must have completed by the time the later one passes its wait. That holds for calls on one +stream when every kernel between them waits on its predecessor or is launched without PDL, as the kernels of a decode +step do. Calls that overlap on one workspace (two streams, or a graph replay beside an eager call) overwrite each +other's partials and counters; they are outside the contract (not measured). + +**What a later launch reads.** Only the counters, and only in the no-cluster mode: a launch reads each of its +(request, head group) counters after its grid-dependency wait and before it arrives, then waits for 16 arrivals past +`count & ~15`. The partials and the exchange are written and read within one call: the op test refills them with NaN +before every call and the outputs keep their bits (`test_k3_mla_attn.py::test_attn_workspace_poison`). Before its +grid-dependency wait a launch reads only the page table, the lengths and the cache pages before the one holding row +`L_i - T` (earlier steps wrote them); `q` and the pages from that one on after it. + +**How it is re-armed.** Never. The counters only grow: each no-cluster launch adds 16 to both counters of each of its +requests' head groups and leaves the others; cluster-mode launches leave them all (certified: counters equal to 16 x +the launches after a sequence of eager calls and graph replays). They are compared by signed difference, so they run +through the int32 wrap (op test: counters started at `2^31 - 16` and at `-16` give the bits of zeroed ones). A +launch must start with its counters at a multiple of 16, which every complete launch leaves; `create()` starts them +at 0. + +**Why the test drives call sequences.** One workspace serves every layer and step, eagerly and inside captured +graphs, and its counters carry over from launch to launch. The test runs 3 layers x 2 decode steps on one shared +workspace eagerly, then one step's 3 calls captured as a CUDA graph on the same workspace and replayed for 2 more +steps with rewritten `q` and the metadata prepared for each step. Every output is +compared bit for bit with the same call on a new workspace and within `1e-2` with the reference; in the no-cluster +mode the shared workspace's counters are checked at the end. + +**What a wrong workspace does.** Measured (the test's negative control): a one-head-group call given a workspace +laid out for four head groups, an fp32 tensor, or a buffer 8 elements short raises `ValueError` naming the expected +workspace, before any launch; neither `out` nor the workspace changes. + +## Metadata consumed + +Through `k3_mla_decode_view` (see `attention/k3_mla_qkv`): the page table and lengths are views of the prepared +metadata's `kv_cache_block_offsets` and `kv_lens_cuda_runtime`, `page_offset` the layer's slot, `row_stride` the +pool's token stride, `softmax_scale` from the attention module's `qk_nope_head_dim`, `qk_rope_head_dim` and +`q_scaling`. + +## Preconditions + +- `q`, the page table and the lengths describe `R <= 8` requests of the same `T <= 8` tokens, `heads` a multiple of 6, + `row_stride >= 576` and a multiple of 8; `w_vb`, `out` and `gate` as above. Anything else raises `ValueError` + before any launch (the op test, `test_attn_rejects`), as does a workspace of another layout, dtype or device + (certified, above). +- `seq_len[i] >= T`: the step's `T` new rows are in the cache (stored by `k3_mla_qkv` first). +- The first call of each configuration compiles the kernel and must run outside CUDA-graph capture (it raises + `RuntimeError` inside one). +- The no-cluster mode needs its 16 x `R x groups` CTAs co-resident (one per SM); the op takes it only when they fit, + and otherwise launches clusters in more than one wave. +- SM 100 / SM 103 only (CuTe DSL `tcgen05` kernel). + +## Notes + +- The op-level test is `tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_attn.py`: every split of up to 8 + tokens and the R x 8 verify steps against the float64 reference and the stock CuTe DSL MLA decode + (`trtllm::cute_dsl_mla_decode_fp16_blackwell`), the TP4 head count, interleaved pools of rows of 640, the + workspace poison and counter-wrap checks above, refusals, and (with `K3_BASE_TRTLLM`) batch-1 identity against + the unmodified single-request kernel. +- `k3_mla_attn` takes no page offset: it fits a pool with one layer per block, or layer 0 of an interleaved one. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_attn_vb_out.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_attn_vb_out.py new file mode 100644 index 000000000000..8f52772f1565 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_attn_vb_out.py @@ -0,0 +1,79 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3's MLA decode attention over the paged latent cache, with v_b and the output gate in the same launch, over +a caller-owned :class:`K3MlaAttnWorkspace`; also the plain forms (the attention output itself).""" + +from typing import Optional + +import torch + +from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import ( # noqa: F401 — registers the ops + op as _k3_mla_op, +) + +from .k3_mla_attn_workspace import K3MlaAttnWorkspace + + +def k3_mla_attn_vb_out( + q: torch.Tensor, + pool: torch.Tensor, + row_stride: int, + page_table: torch.Tensor, + page_offset: int, + seq_len: torch.Tensor, + softmax_scale: float, + w_vb: torch.Tensor, + out: torch.Tensor, + workspace: K3MlaAttnWorkspace, + gate: Optional[torch.Tensor] = None, + gate_col0: int = 0, +) -> None: + """Write ``out`` [M, heads * 128] = per head ``bf16(bf16(o) @ w_vb[h]^T)`` (times the gate's sigmoid with + ``gate``) for the decode attention o of R <= 8 requests of T <= 8 tokens over their pages of ``pool``. Uses + ``workspace``'s partial slots and, past 7 clusters, advances its counters.""" + torch.ops.trtllm.k3_mla_attn_vb_out( + q, + pool, + row_stride, + page_table, + page_offset, + seq_len, + softmax_scale, + w_vb, + out, + workspace.buffer, + gate, + gate_col0, + ) + + +def k3_mla_attn_out( + q: torch.Tensor, + pool: torch.Tensor, + row_stride: int, + page_table: torch.Tensor, + page_offset: int, + seq_len: torch.Tensor, + softmax_scale: float, + out: torch.Tensor, + workspace: K3MlaAttnWorkspace, +) -> None: + """Write the attention output itself, ``out`` [M, heads * 512] bf16 (no v_b), over ``workspace``.""" + torch.ops.trtllm.k3_mla_attn_out( + q, pool, row_stride, page_table, page_offset, seq_len, softmax_scale, out, workspace.buffer + ) + + +def k3_mla_attn( + q: torch.Tensor, + pool: torch.Tensor, + row_stride: int, + page_table: torch.Tensor, + seq_len: torch.Tensor, + softmax_scale: float, + workspace: K3MlaAttnWorkspace, +) -> torch.Tensor: + """Return the attention output [M, heads * 512] bf16 in a new tensor (page offset 0), over ``workspace``.""" + return torch.ops.trtllm.k3_mla_attn( + q, pool, row_stride, page_table, seq_len, softmax_scale, workspace.buffer + ) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_attn_workspace.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_attn_workspace.py new file mode 100644 index 000000000000..3514dcad6265 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_attn_workspace.py @@ -0,0 +1,45 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Caller-owned state of Kimi K3's MLA decode attention: the per-CTA partial slots and the no-cluster mode's arrival +counters of every ``attention/k3_mla_attn_vb_out`` call (and its ``k3_mla_attn`` / ``k3_mla_attn_out`` forms). + +A state type, not an entry: it launches nothing per call. Its constructor is eager; the target builds one per device +and head-group count in ``post_load_weights`` (before any CUDA-graph capture) and passes it to every MLA attention +call of that shape. The contract is the ``## State`` section of ``k3_mla_attn_vb_out.md``. +""" + +from __future__ import annotations + +from dataclasses import dataclass + +import torch + + +@dataclass(eq=False) +class K3MlaAttnWorkspace: + """One device's attention workspace for calls of ``groups`` head groups (the rank's heads / 6): fp16 partial + slots for 8 requests x ``groups`` x 16 CTAs, then the no-cluster mode's fp32 (m, l) exchange and its int32 + arrival counters, which only grow. Calls on one workspace must run one at a time (see the contract's + ``## State``).""" + + buffer: torch.Tensor + """fp16 [attn_workspace_elems(groups)]: the words the op's ``workspace`` argument names.""" + groups: int + + @classmethod + def create(cls, device: torch.device, groups: int) -> "K3MlaAttnWorkspace": + """Allocate and arm a workspace for ``groups`` head groups on ``device``: the counters zeroed, the partial + slots left as allocated (a call reads only the words it wrote). Eager: it allocates, so it refuses to run + under CUDA-graph capture; it returns once the arming is complete, so the workspace is ready on any stream.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import op + + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError("K3MlaAttnWorkspace.create allocates: call it before capture") + device = torch.device(device) + buffer = op.make_attn_workspace(device, groups) + torch.cuda.synchronize(device) + return cls(buffer=buffer, groups=groups) + + @property + def device(self) -> torch.device: + return self.buffer.device diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_qkv.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_qkv.md new file mode 100644 index 000000000000..114d1d4d8a86 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_qkv.md @@ -0,0 +1,154 @@ +--- +receipts: {} +--- + +# k3_mla_qkv + +**Wraps** `torch.ops.trtllm.k3_mla_qkv` (one call). The same contract covers two sibling wrappers, one call each: +`k3_mla_q` (`torch.ops.trtllm.k3_mla_q`, the query path alone) and `k3_mla_qkv_out` +(`torch.ops.trtllm.k3_mla_qkv_out`, the cache rows stored densely instead of into the pool). + +## Semantics + +Kimi K3's MLA decode query path for `M <= 64` tokens in one launch, plus the step's latent KV rows stored into the +paged latent cache. Per token `t` and head `h`, from the fused projection's rows `ag` (fp32 statistics, one bf16 +rounding wherever the unfused model rounds): + +``` +q_n[t] = bf16(rmsnorm(ag[t, :1536]) * w_qa) # q_a_layernorm +q[t, h] = bf16(q_n[t] @ w_qb[192 h : 192 h + 192]^T) # q_b_proj: 128 nope | 64 pe +fused_q[t, h] = [ bf16(q[t, h, :128] @ w_kb[h]^T) | q[t, h, 128:] ] # k_b absorption (512) | q_pe (64) +row[t] = [ bf16(rmsnorm(ag[t, 1536:2048]) * w_kv) | ag[t, 2048:2112] ] # kv_a_layernorm (512) | rope (64) +``` + +Kimi K3 is NoPE: `q_pe` and the rope columns are copied, not rotated. `fused_q` is returned in a new tensor +`[M, heads * 576]`; `row[t]` is stored into `pool`. The `M` tokens are `R` requests of `T = M / R` tokens each, +request-major: token `t = i T + u` is token `u` of request `i`, at position `p = seq_len[i] - T + u`, stored as row +`(page_table[i][p // 64] + page_offset) * 64 + p % 64` of `row_stride` elements (64-bit element index). A token +with `p < 0` is not stored. Nothing else in `pool` is written (certified: every layer's whole pool compared after +every step). + +- `k3_mla_q` returns the same `fused_q` bits and stores nothing (certified). +- `k3_mla_qkv_out` returns the same `fused_q` bits and stores `row[t]` densely into `kv_out[t]` (`[M, 576]`), the + same bits `k3_mla_qkv` stores into the pool (certified). +- `fused_q` rows are computed in chunks of 8, 16 or 32 tokens by call size; every 8-token chunk of rows is + bit-identical to the call on those 8 rows alone (the op test, `test_k3_mla_q.py::test_q`). + +Accuracy (certified bounds): `fused_q` within `2e-2` (max relative, per half) of a float64 reference that keeps the +model's bf16 roundings; the latent columns of `row` within `1e-2` of an fp32 reference (`(x * rrms) * w`, one bf16 +rounding); the rope columns bit-exact. + +Fusion boundary. Inside: both RMSNorms, the q_b GEMM, the k_b absorption, the cache append. Outside: the fused +projection that produced `ag` (`[q_a 1536 | kv_a latent 512 | rope 64 | gate heads * 128]`), the attention that reads +`fused_q` and the cache (`attention/k3_mla_attn_vb_out`), and the output projection. + +## Signature + +```python +def k3_mla_qkv( + ag: torch.Tensor, + w_qa: torch.Tensor, + eps: float, + w_qb: torch.Tensor, + w_kb: torch.Tensor, + w_kv: torch.Tensor, + kv_eps: float, + pool: torch.Tensor, + row_stride: int, + page_table: torch.Tensor, + page_offset: int, + seq_len: torch.Tensor, + trigger_early: bool = True, +) -> torch.Tensor + +def k3_mla_q(ag, w_qa, eps, w_qb, w_kb, trigger_early=True) -> torch.Tensor + +def k3_mla_qkv_out(ag, w_qa, eps, w_qb, w_kb, w_kv, kv_eps, kv_out, trigger_early=True) -> torch.Tensor +``` + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `ag` | `[M, C]`, `M` = `R T` up to 64; `C >= 2112`, a multiple of 8 (Kimi K3: `C = 2112 + heads * 128`) | bf16 | contiguous | CUDA | +| `w_qa`, `w_kv` | `[1536]`, `[512]` | bf16 | contiguous | CUDA | +| `eps`, `kv_eps` | scalar | Python float | — | — | +| `w_qb` | `[heads * 192, 1536]`, `heads` 6 (TP16) or 24 (TP4); the op takes 2 to 24 | bf16 | contiguous | CUDA | +| `w_kb` | `[heads, 512, 128]` (`k_b_proj_trans`) | bf16 | contiguous | CUDA | +| `pool` | the layer's paged latent pool, flat; at least 64 rows | bf16 | contiguous | CUDA | +| `row_stride` | elements per cache row, `>= 576`, a multiple of 8 (576 for the manager's MLA pool) | Python int | — | — | +| `page_table` | `[R, W]` int32 with unit column stride and any row stride `>= W` (rows of `kv_cache_block_offsets`), or `[W]` when `R = 1` | int32 | see shape | CUDA | +| `page_offset` | the layer's slot in a layer-interleaved pool, added to every page-table entry | Python int | — | — | +| `seq_len` | `[R]`, each request's length including its `T` new tokens | int32 | contiguous | CUDA | +| `trigger_early` | default `True`: dependents may launch early (they wait for the whole grid before reading) | Python bool | — | — | +| `kv_out` (`k3_mla_qkv_out`) | `[M, 576]` | bf16 | contiguous | CUDA | +| returns | `fused_q` `[M, heads * 576]`, newly allocated | bf16 | contiguous | = `ag.device` | + +Certified at `R` up to 8 and `T` 1 and 8 (no speculation; a DSpark verify step), 6 and 24 heads, over a real +`KVCacheManager` (below). + +## State + +**Object.** None of its own (kind P: it writes the caller's cache). `pool` is the KV cache manager's MLA latent +pool: `KVCacheManager.get_buffers(layer)` of a `SELFKONLY` manager (one latent head of 576, `kv_factor` 1, 64 tokens +per block), owned and sized by the manager. The addressing (`pool`, `row_stride`, `page_table`, `page_offset`, +`seq_len`) is the attention metadata's view of one generation step: `k3_mla_decode_view(attn, metadata, M)` +(`attention/backends/fmha/cute_dsl_mla.py`), computed per layer after `metadata.prepare()`. + +**What a call writes.** One 576-column row per token, at that token's position in its request (positions below 0 +skipped); nothing else (certified, above). The manager keeps every layer in one pool, interleaved by block: a layer's +rows sit in its slot of each block, and `page_offset` names the slot. + +**Which calls may share one object.** Every layer's call writes the same pool, kept apart by `page_offset` (each +layer's slot) and by position (each step's tokens). Two managers, such as two models' caches, are independent +(certified: calls alternating between two managers step by step and layer by layer; each pool holds exactly its own +rows). + +**Call order.** Calls of one step write disjoint rows, so their order does not matter; a step's rows must be stored +before the attention of the same layer reads them, which it does after its grid-dependency wait, so the attention +may directly follow this call. + +**What a launch reads before its grid-dependency wait.** The page-table rows and the lengths: host-filled metadata +buffers (`kv_cache_block_offsets`, `kv_lens_cuda_runtime`, written by `TrtllmAttentionMetadata.prepare()` before the +step runs or its graph replays). They must not be produced by the kernel directly before this one in the stream. +`ag` is read after the wait. + +**Re-arm.** Nothing to re-arm. + +**Why the test drives call sequences.** A misplaced row is silent: it lands in another layer's slot or another +position, and only a later read sees it. The test therefore keeps each layer's expected image from the manager's own +block ids (`get_batch_cache_indices`), not from the op's addressing, and compares the whole pool of every layer after +every step: 3 layers x 2 decode steps eagerly, then one step captured as a CUDA graph and replayed for 2 more steps +with rewritten inputs and the metadata prepared for each step. + +**What a wrong slot does.** Measured (the test's negative control): layer 1's call given layer 0's `page_offset`. +Nothing raises; layer 0's rows at the step's positions now hold layer 1's values, and layer 1's slot is not written. +The op trusts `page_offset`. + +## Metadata consumed + +Through `k3_mla_decode_view`: request `i`'s pages are `kv_cache_block_offsets[pool, num_contexts + i, 0, :]` (a +strided view, no copy), its length `kv_lens_cuda_runtime[num_contexts + i]`, `page_offset` the layer's index in the +pool, `row_stride` the pool's token stride. The view applies to generation-only steps (`num_contexts == 0`) of `R +<= 8` requests of the same `T <= 8` tokens, beam width 1, no tree speculation, 64-token pages, a bf16 pool, and +`heads` a multiple of 6 with `kv_lora_rank` 512 and `qk_rope_head_dim` 64; otherwise it returns a reason string and +the target takes its generic path. + +## Preconditions + +- `M <= 64`, and `page_table` / `seq_len` describe `R` requests of `T = M / R` tokens (`M % R == 0`; a `[W]` table + only when `R = 1`); `heads * 6 <= 148` (at most 24 heads) and, with the cache half, `heads >= 2`. Out of contract + calls raise `ValueError` before any launch, the pool unchanged (certified). +- `seq_len[i] >= T` for a request whose tokens should all be stored; a smaller length stores only the positions + `>= 0` (the op test, `test_qkv_short_length`). +- The first call of each configuration compiles the kernel and must run outside CUDA-graph capture (it raises + `RuntimeError` inside one). +- SM 100 / SM 103 only (CuTe DSL `tcgen05` kernel). + +## Notes + +- The op-level test is `tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_q.py`: `fused_q` at `M` 1 to 64 + against the reference and the unfused chain (the model's RMSNorm, the q_b GEMM, the k_b bmm), the cache rows of R + x T steps in a sentinel-filled pool, an interleaved pool of rows of 640, the short-length case, refusals, and (with + `K3_BASE_TRTLLM`) batch-1 identity against the unmodified single-request kernel. +- The weights are TMA'd before the grid-dependency wait (EVICT_FIRST); the q_a and kv_a rows after it. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_qkv.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_qkv.py new file mode 100644 index 000000000000..0e1345ff6f3a --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_qkv.py @@ -0,0 +1,74 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3's MLA decode query path (q_a RMSNorm, q_b projection, k_b absorption) in one launch, which also stores the +step's latent KV rows into the paged latent cache; also the query path alone and the form that stores the rows +densely.""" + +import torch + +from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import ( # noqa: F401 — registers the ops + op as _k3_mla_op, +) + + +def k3_mla_qkv( + ag: torch.Tensor, + w_qa: torch.Tensor, + eps: float, + w_qb: torch.Tensor, + w_kb: torch.Tensor, + w_kv: torch.Tensor, + kv_eps: float, + pool: torch.Tensor, + row_stride: int, + page_table: torch.Tensor, + page_offset: int, + seq_len: torch.Tensor, + trigger_early: bool = True, +) -> torch.Tensor: + """Return ``fused_q`` [M, heads * 576] bf16 of the M <= 64 rows of ``ag``, and store each token's cache row + ``[bf16(rmsnorm(ag[t, 1536:2048]) * w_kv) | ag[t, 2048:2112]]`` into ``pool`` at its request's position.""" + return torch.ops.trtllm.k3_mla_qkv( + ag, + w_qa, + eps, + w_qb, + w_kb, + w_kv, + kv_eps, + pool, + row_stride, + page_table, + page_offset, + seq_len, + trigger_early, + ) + + +def k3_mla_q( + ag: torch.Tensor, + w_qa: torch.Tensor, + eps: float, + w_qb: torch.Tensor, + w_kb: torch.Tensor, + trigger_early: bool = True, +) -> torch.Tensor: + """Return ``fused_q`` alone (no cache row stored): the same bits as ``k3_mla_qkv``'s.""" + return torch.ops.trtllm.k3_mla_q(ag, w_qa, eps, w_qb, w_kb, trigger_early) + + +def k3_mla_qkv_out( + ag: torch.Tensor, + w_qa: torch.Tensor, + eps: float, + w_qb: torch.Tensor, + w_kb: torch.Tensor, + w_kv: torch.Tensor, + kv_eps: float, + kv_out: torch.Tensor, + trigger_early: bool = True, +) -> torch.Tensor: + """Return ``fused_q`` and store the cache rows densely into ``kv_out`` [M, 576] instead of the pool.""" + return torch.ops.trtllm.k3_mla_qkv_out( + ag, w_qa, eps, w_qb, w_kb, w_kv, kv_eps, kv_out, trigger_early + ) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml index 7382f31d99a1..f3013bff9fc6 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml @@ -172,6 +172,35 @@ entries: impl: tensorrt_llm.bindings.internal.thop.attention summary: "Full attention core with fully explicit state (pybind binding, approved exception): paged KV-cache append + causal/padding-masked GQA FMHA over a caller-owned pool (bf16, or fp8-e4m3 with per-tensor kv scales) and explicit length/offset tensors, written into a caller buffer, on either context execution path (use_paged_context_fmha selects packed-QKV context FMHA, or paged-KV context FMHA so a context call may run over a cached prefix — KV-cache reuse and chunked prefill), with optional per-query-head attention sinks (one extra softmax-denominator logit, dropped from the output) and an optional per-call sliding window (attention_window_size keys ending at the query's absolute position; a pure mask — the append stays at absolute positions, so cyclic pool reuse is the caller's page mapping); MLA mode runs context prefill (in-kernel RoPE + latent append), no-append context over explicit K/V (latent_cache=None: cached-KV prefixes and chunked partial passes with softmax-stats output), and generation latent-MQA decode as separate calls over a paged latent pool (bf16 or fp8-e4m3)" + # Kimi K3 MLA decode (sm_100). The state type attention/k3_mla_attn_workspace.py is not an entry: it launches nothing + # per call. + - path: attention/k3_mla_qkv.py + impl: torch.ops.trtllm.k3_mla_qkv + summary: "Kimi K3's MLA decode query path for at most 64 tokens in one launch (q_a RMSNorm, q_b projection and k_b absorption into fused_q [M, heads x 576]) plus the step's latent KV rows (kv_a RMSNorm, rope columns) stored into the paged latent cache at each token's row; the page table is read before the grid-dependency wait, so under CUDA graphs it must be written before the graph runs; siblings: k3_mla_q (the query path alone) and k3_mla_qkv_out (the cache rows into a dense buffer)" + + - path: attention/k3_mla_attn_vb_out.py + impl: torch.ops.trtllm.k3_mla_attn_vb_out + summary: "Kimi K3's MLA decode attention of at most 8 requests of at most 8 tokens over the paged latent cache (heads a multiple of 6, causal bottom-right, a request's rows exactly those of its own R = 1 call) with v_b and the optional sigmoid output gate in the same launch, over a caller-owned K3MlaAttnWorkspace: per-CTA partials and the no-cluster mode's arrival counters, one workspace per device and head-group count shared by every layer, its calls one at a time in stream order; siblings: k3_mla_attn_out (the attention output itself) and k3_mla_attn" + + # ─── ssm ─────────────────────────────────────────────────────── + # KDA (Kimi Delta Attention) decode. The state type ssm/k3_kda_buffers.py is not an entry: it launches nothing per + # call. The pools every entry here updates belong to the cache manager (the V2 hybrid manager's per-layer views). + - path: ssm/kda_decode.py + impl: torch.ops.trtllm.kda_decode + summary: "One-token KDA decode of B requests: the causal depthwise conv over each request's conv window, the gated delta rule on its fp32 recurrent state and the gated RMS norm, the state (and the conv window) updated in place at each request's slot of the cache manager's pools; at Kimi K3's rank slice (6 heads, K = V = 128, conv width 4) B <= 5 runs the cluster kernel and 6 <= B <= 24 the legacy compact-heads kernel" + + - path: ssm/k3_kda_attn.py + impl: torch.ops.trtllm.k3_kda_attn + summary: "Kimi K3's KDA layer for one speculative-decoding request in one launch: the fused input projection of its golden token and 7 drafts and the KDA verify of ssm/k3_kda_verify on it, bit for bit the unfused pair, writing the slot's state, per-draft states and conv window in place, over a caller-owned K3KdaBuffers (three Lamport buffers and per-CTA indices, one set per device shared with ssm/k3_kda_decode_attn, its launches one at a time in stream order); the head CTAs read the slot's pools before the grid-dependency wait, so a launch must not directly follow another on the same pools; sibling: k3_kda_qkvg (the projection stream alone, a 104-CTA set of its own)" + + - path: ssm/k3_kda_decode_attn.py + impl: torch.ops.trtllm.k3_kda_decode_attn + summary: "Kimi K3's KDA layer for plain decode of R <= 8 requests in one launch: the fused input projection of one token per request and the one-token KDA decode of ssm/kda_decode on it, the state and conv window updated in place at each request's slot, over the same caller-owned K3KdaBuffers as ssm/k3_kda_attn (either may follow the other on one set); the head CTAs read the slots' pools before the grid-dependency wait" + + - path: ssm/k3_kda_verify.py + impl: torch.ops.trtllm.k3_kda_verify + summary: "Kimi K3's KDA speculative verify of N requests of 1 + num_spec tokens from the fused projection rows to the gated-norm core output, starting each request from the state after the drafts the sampler accepted last round (the per-token states of the V2 hybrid manager, selected by prev_num_accepted_tokens, read before the grid-dependency wait) and committing the state after every verify token, so the next round needs no replay" + # ─── moe ─────────────────────────────────────────────────────── - path: moe/noaux_tc_op.py impl: torch.ops.trtllm.noaux_tc_op diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/__init__.py new file mode 100644 index 000000000000..a0f916e7970b --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/__init__.py @@ -0,0 +1,3 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""State-space (KDA) entries.""" diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_attn.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_attn.md new file mode 100644 index 000000000000..505145f1ad39 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_attn.md @@ -0,0 +1,195 @@ +--- +receipts: {} +--- + +# k3_kda_attn + +**Wraps** `torch.ops.trtllm.k3_kda_attn` (one call), over a caller-owned `K3KdaBuffers` +(`catalog/ssm/k3_kda_buffers.py`). Sibling in this contract: `k3_kda_qkvg`, which wraps +`torch.ops.trtllm.k3_kda_qkvg` (the projection stream alone, one call). + +Kimi K3's KDA layer for one speculative-decoding request, at the TP16 rank slice: the fused input projection of the +request's golden token and its 7 drafts and the KDA verify of `ssm/k3_kda_verify` on it, in one launch. + +## Semantics + +`x` bf16 `[8, 7168]` is one request's golden token and 7 drafts. One launch computes: + +``` +# 1. projection, streamed by 26 clusters of 4 CTAs (w = [q | k | v | og | f_a | b | pad], 3208 rows) +y = x @ w^T +# q, k, f_a: bf16 of the four split-K partials of a cluster, summed in rank order +# v, og, b: bf16(half_0 + half_1), the fp32 sums of the two K halves +# 2. verify (exactly ssm/k3_kda_verify on y with the gate folded in, g_ext = None): +# starting state S = ssm[slot] if P == 0 else state_tok[slot, P - 1], P = pending[slot] +# per token t = 0..7: g = bf16(f_a @ w_fb^T); q, k, v = SiLU(conv4(window, new raw)); q, k L2-normalized, +# q *= scale; beta = sigmoid(b); decay = exp(lower_bound * sigmoid(exp(a_log) * (g + dt_bias))); +# S *= decay (per key); S += beta (v - S k) k^T; o = S q; +# out[t] = o * rsqrt(mean(o^2) + eps) * onorm_w * sigmoid(og) +``` + +and returns `out` bf16 `[8, 6, 128]`, the gated-norm core output. In place, at the request's slot: + +- `ssm[slot]` = the state after the golden token (t = 0); +- `state_tok[slot, t - 1]` = the state after draft t, t = 1..7; +- `cs_q` / `cs_k` / `cs_v[slot]` = the raw inputs at positions -2..7 around the golden token (the next round's + window starts at column P). + +The result is bit for bit that of `k3_kda_qkvg` followed by `ssm/k3_kda_verify` on the decoded rows, on copies of the +same pools (certified, every pool word). Repeated runs are bit-identical. + +**`k3_kda_qkvg`** runs phase 1 alone on `x` bf16 `[T <= 8, 7168]` and publishes the rows: buffer +`e = buffers.epoch[0]` (before the call) holds q, k and f_a as bf16 bits in `p1`, and v, og and b as the fp32 bits of +the two K-half partials in `part`. The projection is bf16 of their sum, the consumer's job. Measured against a float64 +`x @ w^T`: TBD(tray: test_modeling_v2_k3_kda_attn.py::test_qkvg_rows_match_the_projection). + +## Signature + +```python +def k3_kda_attn( + x: torch.Tensor, + w: torch.Tensor, + w_fb: torch.Tensor, + w_q: torch.Tensor, + w_k: torch.Tensor, + w_v: torch.Tensor, + a_log: torch.Tensor, + dt_bias: torch.Tensor, + onorm_w: torch.Tensor, + cs_q: torch.Tensor, + cs_k: torch.Tensor, + cs_v: torch.Tensor, + ssm: torch.Tensor, + state_tok: torch.Tensor, + slots: torch.Tensor, + pending: torch.Tensor, + buffers: K3KdaBuffers, + num_spec: int, + lower_bound: float, + scale: float, + eps: float, +) -> torch.Tensor + +def k3_kda_qkvg(x: torch.Tensor, w: torch.Tensor, buffers: K3KdaBuffers) -> None +``` + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `x` | `[8, 7168]` (`k3_kda_qkvg`: `[T <= 8, 7168]`) | bf16 | contiguous | CUDA | +| `w` | `[3208, 7168]`: q, k, v, og (768 rows each), f_a (128), b (6), pad (2) | bf16 | contiguous | CUDA | +| `w_fb` | `[768, 128]` (f_b, out x in) | bf16 | contiguous | CUDA | +| `w_q`, `w_k`, `w_v` | `[768, 4]` conv taps, oldest input first | fp32 | dense | CUDA | +| `a_log` | `[6]` | fp32 | dense | CUDA | +| `dt_bias` | `[768]` | fp32 | dense | CUDA | +| `onorm_w` | `[128]` | fp32 | dense | CUDA | +| `cs_q`, `cs_k`, `cs_v` | `[pool, 768, 3 + num_spec]` | fp32 | channel stride 1 (dim-contiguous) | CUDA | +| `ssm` | `[pool, 6, 128, 128]` (V rows, K contiguous) | fp32 | each slot dense, slots at any stride | CUDA | +| `state_tok` | `[pool, num_spec, 6, 128, 128]` | fp32 | contiguous | CUDA | +| `slots` | `[1]`: the request's slot | int32 | any element offset | CUDA | +| `pending` | `[pool]`: drafts the sampler accepted last round, per slot | int32 | any element offset | CUDA | +| `buffers` | `K3KdaBuffers` made with `ctas=FUSED_CTAS` (128) | — | — | CUDA | +| `num_spec` | 7 | Python int | — | — | +| `lower_bound`, `scale`, `eps` | scalars (Kimi K3: -5.0, 128^-0.5, 1e-5) | Python float | — | — | +| returns | `[8, 6, 128]` | bf16 | contiguous, fresh | CUDA | + +`mutates_args`: `cs_q`, `cs_k`, `cs_v`, `ssm`, `state_tok`, and the set's `p1`, `part`, `epoch`. `k3_kda_qkvg` takes a +set made with `ctas=CTAS` (104) and mutates its `p1`, `part`, `epoch`; it returns None. + +## State + +**Object.** `K3KdaBuffers` (`catalog/ssm/k3_kda_buffers.py`), one per device, owned by the caller. + +**Contents and size.** The projection's three Lamport buffers and each CTA's buffer index: + +- `p1` int16 `[3 * 8 * 1664]`: per buffer, 8 token rows of the q, k and f_a columns (79,872 bytes); +- `part` int32 `[3 * 3 * 2 * 8 * 768]`: per buffer, the fp32 bits of the two K-half partials of v, og and b + (442,368 bytes); +- `epoch` int32 `[128]` (`[104]` for `k3_kda_qkvg`): each CTA's launch count mod 3. + +Every buffer word holds the sentinel (all ones, a NaN no finite GEMV produces; a computed all-ones word is stored as +the canonical NaN instead) until the launch's producer CTAs write it. The consumer CTAs poll the words they need until +none is the sentinel: the data words are the flags, with no fence or counter. 8 token rows are published per launch, +whatever the batch. + +**Who creates it, and when.** The target, in `post_load_weights`, with `K3KdaBuffers.create(device)`: + +- eager: it allocates, so it refuses to run under CUDA-graph capture; +- it arms every buffer word to the sentinel and every index to 0; +- `ctas` is `FUSED_CTAS` (128) for this op and `ssm/k3_kda_decode_attn`, or `CTAS` (104) for `k3_kda_qkvg`; any other + count raises. No environment variable is read. + +**Which ops may share one object.** `k3_kda_attn` and `ssm/k3_kda_decode_attn`, of every KDA layer on the device: +one set per device serves them all. Both run the same stream role, publish all 8 token rows, re-arm the same words +and move every CTA's index once per launch, so either may follow the other in any order. Certified: plain-decode and +verify launches interleaved on one set give the bits of the same launches on sets of their own; layers x steps on one +set, a set per layer, and two sets alternating between launches agree bit for bit. `k3_kda_qkvg` needs a set of its +own (104 CTAs); a 104-CTA set passed here raises `ValueError` before anything is written. + +**Call-order invariant.** Launches on one set run one at a time, in stream order: never on two streams at once. Each +launch reads `e = epoch[cta]` after its grid-dependency wait, writes buffer e, stores the sentinel into the same words +of buffer (e + 1) % 3 (the next launch's, which the launch before last read) and leaves `epoch[cta] = (e + 1) % 3`. +After any launch every CTA's index is equal. + +**What a later launch reads.** `epoch`, as the previous launch left it, and its own buffer's words, which the previous +launch re-armed to the sentinel. + +**How it is re-armed.** By the previous launch (above); the first launch on a new set writes buffer 0, which `create` +armed. The index stays in 0..2: a raw launch count would turn negative after 2^31 launches and its signed remainder +would index before the buffers. Certified with every index preset to 2^31 - 2: four launches give the bits of a fresh +set's and leave every index in 0..2 (the op's own test also checks with guard bands that nothing is written outside the +buffers). + +**The pools.** Not part of the object: they belong to the cache manager. With the V2 hybrid manager built with the +KDA replay caches and `kda_token_states`, layer `l`'s are `mamba_layer_cache(l).kda_conv_q` / `kda_conv_k` / +`kda_conv_v` (`cs_*`), `get_ssm_states(l)` (`ssm`, its slots strided by the manager's per-slot coalescing) and +`mamba_layer_cache(l).kda_state_tok`; `pending` is `prev_num_accepted_tokens`, one record shared by every layer and +written by the sampler's acceptance between steps. A launch on layer `l` writes only layer `l`'s views, at its slot. + +**The PDL rule.** The head CTAs read the slot's pools (state, records, conv window) before their grid-dependency wait. +A launch must therefore not directly follow another launch on the same pools in the stream: a kernel that waits, or a +non-PDL kernel, must sit between them, as the model's other layers do. This op lets its dependents launch only once +its head CTAs have passed their own wait, i.e. once the launch before it has completed, so one launch of this op (or of +`ssm/k3_kda_decode_attn`) on other pools in between is enough. A kernel that lets its dependents launch before its +wait (`ssm/k3_kda_verify` does, in its prologue) is not. + +**Why the test drives call sequences.** The state is carried by the pools from round to round: `pending` selects the +next round's starting state and conv window, and a round's per-draft states are only read by the next one. The test +therefore runs two requests one after the other, each for 9 rounds with pending running through 0..7, on every layer +of a real manager, against the unfused path bit for bit and layer 0 also against a float64 verify over the request's +committed history; then layers x rounds on one shared set against a set per layer, and a round of every layer captured +once and replayed with rewritten inputs and records. + +**What a wrong order does.** Two rounds of one request in swapped order (the test's negative control): nothing raises, +and the second round's output and the slot's state are those of a different history. Measured on sm_100: +TBD(tray: test_modeling_v2_k3_kda_attn.py::test_swapped_rounds_are_silently_wrong, its printed rel diffs). + +## Metadata consumed + +`slots` and `pending`. Both are read before the grid-dependency wait, so under CUDA graphs they must be written before +the graph runs (by the step's preparation and the previous step's acceptance), never by a kernel of the same step +ahead of this launch. + +## Preconditions + +- sm_100: the kernel uses tcgen05, TMA and clusters. +- The op checks `x`, `w`, `w_fb` (shape, dtype, contiguity), `num_spec` = 7, `ssm` / `state_tok` (shape, fp32, + `ssm.stride()[1:] == (16384, 128, 1)`, `state_tok` contiguous), `cs_q`'s window width, the fp32 dtypes, `slots` + (one int32 element), `pending` (int32) and the set's sizes (`epoch` of 128), and raises `ValueError` before a launch. + A conv cache that is not dense raises `ValueError` too. The other shapes in the table (the fp32 weights, `cs_k` / + `cs_v`, `pending` covering the slot, the slot inside the pool) are the caller's obligation. +- Every tensor the kernel addresses is 16-byte aligned except `slots` / `pending` (see *Notes*); the DSL checks the + assumed alignment at the call. +- The first call of each configuration (`lower_bound`, `scale`, `eps` and whether PDL is on) compiles the kernel and + must run outside CUDA-graph capture; under capture it raises `RuntimeError`. PDL follows `TRTLLM_ENABLE_PDL` (default + on), read at every call. +- The slot is not in use by another launch on the same pools (one request per launch; a batch is a sequence of + launches). + +## Notes + +- `slots` and `pending` may be slices of longer index tensors starting at any element (the mixer passes + `state_indices[num_prefills:]`): the kernel reads them at their element's alignment. The op's own test certifies + the bits at offsets 0..3. +- The pools are addressed at 64-bit slot offsets, so pools past 2 GiB are fine (the op's own test, + `test_k3_kda_pools_past_2g.py`). +- Repeated runs, and graph replays against eager runs, are bit-identical. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_attn.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_attn.py new file mode 100644 index 000000000000..fce495f9baf5 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_attn.py @@ -0,0 +1,75 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3's KDA projection and speculative verify of one request's 8 tokens in one launch, over a caller-owned +:class:`K3KdaBuffers`; and ``k3_kda_qkvg``, the projection stream alone.""" + +import torch + +from tensorrt_llm._torch.cute_dsl_kernels.k3_kda_attn import ( # noqa: F401 — registers the ops + op as _k3_kda_attn_op, +) +from tensorrt_llm._torch.cute_dsl_kernels.k3_kda_verify import ( # noqa: F401 — k3_kda_attn's helpers + op as _k3_kda_verify_op, +) + +from .k3_kda_buffers import K3KdaBuffers + + +def k3_kda_attn( + x: torch.Tensor, + w: torch.Tensor, + w_fb: torch.Tensor, + w_q: torch.Tensor, + w_k: torch.Tensor, + w_v: torch.Tensor, + a_log: torch.Tensor, + dt_bias: torch.Tensor, + onorm_w: torch.Tensor, + cs_q: torch.Tensor, + cs_k: torch.Tensor, + cs_v: torch.Tensor, + ssm: torch.Tensor, + state_tok: torch.Tensor, + slots: torch.Tensor, + pending: torch.Tensor, + buffers: K3KdaBuffers, + num_spec: int, + lower_bound: float, + scale: float, + eps: float, +) -> torch.Tensor: + """Return the gated-norm core output bf16 [8, 6, 128] of one request's golden token and 7 drafts ``x`` bf16 + [8, 7168]: the projection ``x w^T`` and the verify of ``ssm/k3_kda_verify`` on it, in one launch. Updates the + slot's conv caches, state and per-draft states in place; advances ``buffers`` by one launch.""" + return torch.ops.trtllm.k3_kda_attn( + x, + w, + w_fb, + w_q, + w_k, + w_v, + a_log, + dt_bias, + onorm_w, + cs_q, + cs_k, + cs_v, + ssm, + state_tok, + slots, + pending, + buffers.p1, + buffers.part, + buffers.epoch, + num_spec, + lower_bound, + scale, + eps, + ) + + +def k3_kda_qkvg(x: torch.Tensor, w: torch.Tensor, buffers: K3KdaBuffers) -> None: + """The projection ``x w^T`` of T <= 8 tokens (``x`` bf16 [T, 7168], ``w`` bf16 [3208, 7168]) published into + ``buffers`` (made with ``ctas=CTAS``): buffer ``e = buffers.epoch[0]`` (before the call) holds the rows; advances + ``buffers`` by one launch. Returns None.""" + torch.ops.trtllm.k3_kda_qkvg(x, w, buffers.p1, buffers.part, buffers.epoch) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_buffers.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_buffers.py new file mode 100644 index 000000000000..d2869f3c5afa --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_buffers.py @@ -0,0 +1,52 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Caller-owned state of Kimi K3's fused KDA projection: one set of its three Lamport buffers and per-CTA indices. + +A state type, not an entry: it launches nothing per call. Its constructor is eager; the target builds one per +device in ``post_load_weights`` (before any CUDA-graph capture) and passes it to every ``ssm/k3_kda_attn`` and +``ssm/k3_kda_decode_attn`` call on that device, of every KDA layer. ``ssm/k3_kda_attn``'s sibling +``k3_kda_qkvg`` (the projection stream alone) takes a set of its own, made with ``ctas=CTAS``. The contract is the +``## State`` section of ``k3_kda_attn.md``. +""" + +from __future__ import annotations + +from dataclasses import dataclass + +import torch + +from tensorrt_llm._torch.cute_dsl_kernels.k3_kda_attn import op as _op + +CTAS = _op.CTAS # k3_kda_qkvg's grid: 26 clusters of 4 CTAs +FUSED_CTAS = _op.FUSED_CTAS # k3_kda_attn's and k3_kda_decode_attn's grid + + +@dataclass(eq=False) +class K3KdaBuffers: + """The projection's Lamport set: three buffers of every published word, each word the sentinel (all ones) until a + launch writes it, and each CTA's buffer index. A launch writes buffer ``e = epoch[cta]``, re-arms buffer + ``(e + 1) % 3`` (the next launch's) to the sentinel and leaves ``epoch[cta] = (e + 1) % 3``, so launches on one + set run one at a time, in stream order, and every launch moves every index once (see the contract's + ``## State``).""" + + p1: torch.Tensor + """int16 [3 * 8 * 1664]: per buffer, 8 token rows of the q, k and f_a columns (bf16 bits).""" + part: torch.Tensor + """int32 [3 * 3 * 2 * 8 * 768]: per buffer, the fp32 bits of the two K-half partials of v, og and b.""" + epoch: torch.Tensor + """int32 [ctas]: each CTA's buffer index (its launch count mod 3).""" + ctas: int + + @classmethod + def create(cls, device, ctas: int = FUSED_CTAS) -> "K3KdaBuffers": + """Allocate and arm a set for ``ctas`` CTAs (``FUSED_CTAS`` for k3_kda_attn / k3_kda_decode_attn, ``CTAS`` + for k3_kda_qkvg) on ``device``: every buffer word the sentinel, every index 0. Eager: it allocates, so it + refuses to run under CUDA-graph capture.""" + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError("K3KdaBuffers.create allocates: call it before CUDA-graph capture") + if ctas not in (CTAS, FUSED_CTAS): + raise ValueError( + f"K3KdaBuffers: ctas must be {CTAS} (k3_kda_qkvg) or {FUSED_CTAS}, got {ctas}" + ) + p1, part, epoch = _op.make_buffers(torch.device(device), ctas) + return cls(p1=p1, part=part, epoch=epoch, ctas=ctas) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_decode_attn.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_decode_attn.md new file mode 100644 index 000000000000..11cc7085fa88 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_decode_attn.md @@ -0,0 +1,127 @@ +--- +receipts: {} +--- + +# k3_kda_decode_attn + +**Wraps** `torch.ops.trtllm.k3_kda_decode_attn` (one call), over a caller-owned `K3KdaBuffers` +(`catalog/ssm/k3_kda_buffers.py`), the same object `ssm/k3_kda_attn` takes. + +Kimi K3's KDA layer for plain (non-speculative) decode, at the TP16 rank slice: the fused input projection of one +token of each of R <= 8 requests and the one-token KDA decode of `ssm/kda_decode` on it, in one launch. + +## Semantics + +`x` bf16 `[R, 7168]`, one token of each of R requests. One launch computes: + +``` +# 1. projection: k3_kda_attn's stream (26 clusters of 4 CTAs, three Lamport phases), rows >= R read as zero +y = x @ w^T # q, k, f_a, v, og, b rounded to bf16 as in ssm/k3_kda_attn +# 2. per request r, on slot s = slots[r] (six head clusters of 4 CTAs, one V quarter per CTA): +g = bf16(f_a @ w_fb^T) +q, k, v = SiLU(conv4(conv[s] window, new raw)) # q, k L2-normalized, q *= scale +conv[s] = the window shifted by one: its last two raw inputs, then the new one +beta = sigmoid(b); decay = exp(lower_bound * sigmoid(exp(a_log) * (g + dt_bias))) +S = ssm[s] * decay (per key); S += beta (v - S k) k^T; ssm[s] = S; o = S q +out[r] = o * rsqrt(mean(o^2) + eps) * onorm_w * sigmoid(og) +``` + +and returns `out` bf16 `[R, 6, 128]`, the gated-norm core output. The decode is `ssm/kda_decode`'s arithmetic, row and +key layout. Certified against the model's unfused path (the projection stream alone, f_b as a bf16 `F.linear`, then +`ssm/kda_decode` on a copy of the pools) and a float64 decode, both at fp32 tolerance for the output (2e-2 relative, +bf16 outputs) and the state rows (1e-3); the conv windows bit for bit (raw bf16 inputs). Repeated runs are +bit-identical. + +## Signature + +```python +def k3_kda_decode_attn( + x: torch.Tensor, + w: torch.Tensor, + w_fb: torch.Tensor, + w_q: torch.Tensor, + w_k: torch.Tensor, + w_v: torch.Tensor, + a_log: torch.Tensor, + dt_bias: torch.Tensor, + onorm_w: torch.Tensor, + conv: torch.Tensor, + ssm: torch.Tensor, + slots: torch.Tensor, + buffers: K3KdaBuffers, + lower_bound: float, + scale: float, + eps: float, +) -> torch.Tensor +``` + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `x` | `[R, 7168]`, 1 <= R <= 8 | bf16 | contiguous | CUDA | +| `w` | `[3208, 7168]` (as `ssm/k3_kda_attn`) | bf16 | contiguous | CUDA | +| `w_fb` | `[768, 128]` | bf16 | contiguous | CUDA | +| `w_q`, `w_k`, `w_v` | `[768, 4]` conv taps, oldest input first | fp32 | dense | CUDA | +| `a_log` | `[6]` | fp32 | dense | CUDA | +| `dt_bias` | `[768]` | fp32 | dense | CUDA | +| `onorm_w` | `[128]` | fp32 | dense | CUDA | +| `conv` | `[slots, 2304, 3]`: q, k, v channels, the last three raw inputs oldest first | bf16 | each slot dense; slot stride a multiple of 8 elements; 16-byte aligned base | CUDA | +| `ssm` | `[slots, 6, 128, 128]` (V rows, K contiguous) | fp32 | each slot dense; slot stride a multiple of 4; 16-byte aligned base | CUDA | +| `slots` | `[R]` | int32 | contiguous | CUDA | +| `buffers` | `K3KdaBuffers` made with `ctas=FUSED_CTAS` (128) | — | — | CUDA | +| `lower_bound`, `scale`, `eps` | scalars (Kimi K3: -5.0, 128^-0.5, 1e-5) | Python float | — | — | +| returns | `[R, 6, 128]` | bf16 | contiguous, fresh | CUDA | + +`mutates_args`: `conv`, `ssm`, and the set's `p1`, `part`, `epoch`. + +## State + +**Object.** `K3KdaBuffers`, shared with `ssm/k3_kda_attn`: one per device for both ops and every KDA layer. Its +contents, creation, re-arming, the 128-CTA requirement and the call-order invariant (launches on one set run one at a +time, in stream order; each launch moves every CTA's index once) are in `k3_kda_attn.md`'s `## State`. Certified here: +plain-decode launches and `ssm/k3_kda_attn` verify launches interleaved on one set give the bits of the same launches +on sets of their own; layers x steps on one set, a set per layer and two sets alternating agree bit for bit; the index +preset to 2^31 - 2 gives the bits of a fresh set and stays in 0..2; a 104-CTA set raises `ValueError` before anything +is written; `create` under capture raises `RuntimeError`. + +**The pools.** The cache manager's plain-decode states, not part of the object: with the V2 hybrid manager +(`conv_state_layout="q_k_v"`, bf16 conv states, fp32 SSM states), layer `l`'s `get_conv_states(l)` (`conv`) and +`get_ssm_states(l)` (`ssm`), their slots strided by the manager's per-slot coalescing; `slots` from +`get_state_indices`. A launch on layer `l` writes only layer `l`'s views, at its slots: certified on a real manager, +with the other slots of the layer and every slot of the other layers bit-unchanged. A batch names each slot once. + +**The PDL rule.** As `ssm/k3_kda_attn`: the head CTAs read the slots' pools (the conv windows, 32 state rows per CTA) +before their grid-dependency wait, so a launch must not directly follow another launch on the same pools in the +stream. A kernel that waits, or a non-PDL kernel, sits between them in the model; one launch of this op (or of +`ssm/k3_kda_attn`) on other pools is also enough, since its dependents launch only after its own wait. + +**Why the test drives call sequences.** Each step's state and conv window are the next step's input, so the test +runs layers x steps of a real manager, and captures one step of every layer once and replays it with rewritten +inputs: the replays give the bits of the same steps run eagerly on a copy of the pools. + +**What a wrong order does.** Two steps of one request in swapped order (the negative control): nothing raises, and the +second step's output and the final state are those of a different history. Measured on sm_100: +TBD(tray: test_modeling_v2_k3_kda_decode_attn.py::test_swapped_steps_are_silently_wrong, its printed rel diffs). + +## Metadata consumed + +`slots`, read before the grid-dependency wait: under CUDA graphs it must be written before the graph runs. + +## Preconditions + +- sm_100: the kernel uses tcgen05, TMA and clusters. +- The op checks `x` (rank 2, 1..8 rows, 7168 columns, bf16, contiguous), `w`, `w_fb`, `conv` (bf16, `[*, 2304, 3]`, + strides `(*, 3, 1)`, slot stride a multiple of 8, 16-byte base), `ssm` (fp32, `[*, 6, 128, 128]`, dense slots, slot + stride a multiple of 4, 16-byte base), `slots` (R int32 elements, contiguous), the set's sizes (`epoch` of 128) and + the fp32 dtypes, and raises `ValueError` before a launch. The fp32 weights' shapes and the slots' range are the + caller's obligation. The V2 hybrid manager's per-layer views meet the layout (certified on a real manager). +- The first call of each configuration (`lower_bound`, `scale`, `eps` and whether PDL is on) compiles the kernel and + must run outside CUDA-graph capture; under capture it raises `RuntimeError`. PDL follows `TRTLLM_ENABLE_PDL` + (default on), read at every call. + +## Notes + +- The pools are addressed at 64-bit slot offsets, so pools past 2 GiB are fine (the op's own test, + `test_k3_kda_pools_past_2g.py`). `slots` may be a slice of a longer index tensor starting at any element; the op's + own test certifies offsets 0..3. +- Within the fp32 tolerances above, not bitwise, against `ssm/kda_decode`: the two read the same rows but sum the + decode in their own orders. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_decode_attn.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_decode_attn.py new file mode 100644 index 000000000000..1d29d58f23d8 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_decode_attn.py @@ -0,0 +1,58 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3's KDA projection and plain one-token decode of up to 8 requests in one launch, over a caller-owned +:class:`K3KdaBuffers`.""" + +import torch + +from tensorrt_llm._torch.cute_dsl_kernels.k3_kda_attn import ( # noqa: F401 — registers the op + op as _k3_kda_attn_op, +) +from tensorrt_llm._torch.cute_dsl_kernels.k3_kda_verify import ( # noqa: F401 — the op's helpers + op as _k3_kda_verify_op, +) + +from .k3_kda_buffers import K3KdaBuffers + + +def k3_kda_decode_attn( + x: torch.Tensor, + w: torch.Tensor, + w_fb: torch.Tensor, + w_q: torch.Tensor, + w_k: torch.Tensor, + w_v: torch.Tensor, + a_log: torch.Tensor, + dt_bias: torch.Tensor, + onorm_w: torch.Tensor, + conv: torch.Tensor, + ssm: torch.Tensor, + slots: torch.Tensor, + buffers: K3KdaBuffers, + lower_bound: float, + scale: float, + eps: float, +) -> torch.Tensor: + """Return the gated-norm core output bf16 [R, 6, 128] of one token of each of R <= 8 requests ``x`` bf16 + [R, 7168]: the projection ``x w^T`` and the plain decode of ``ssm/kda_decode`` on it, in one launch. Updates the + slots' conv and state pools in place; advances ``buffers`` by one launch.""" + return torch.ops.trtllm.k3_kda_decode_attn( + x, + w, + w_fb, + w_q, + w_k, + w_v, + a_log, + dt_bias, + onorm_w, + conv, + ssm, + slots, + buffers.p1, + buffers.part, + buffers.epoch, + lower_bound, + scale, + eps, + ) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_verify.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_verify.md new file mode 100644 index 000000000000..7a819e96b565 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_verify.md @@ -0,0 +1,144 @@ +--- +receipts: {} +--- + +# k3_kda_verify + +**Wraps** `torch.ops.trtllm.k3_kda_verify` (one call). + +Kimi K3's KDA speculative verify of N requests of 1 + num_spec tokens (a golden token and its drafts), from the fused +projection rows to the gated-norm core output, committing the state after every verify token so that the next round +starts from the drafts the sampler accepted instead of replaying them. + +## Semantics + +`proj` bf16 `[N (1 + num_spec), cols]` holds each request's 1 + num_spec rows of the fused projection +`[q | k | v | og | f_a | b | pad]` (H K, H K, H K, H K, K, H columns, then padding). Grid (H, N, 8): a cluster of 8 +CTAs per (head, request), each owning 16 V rows. For request n on slot `s = slots[n]`, with `P = pending[s]`: + +``` +S = ssm[s] if P == 0 else state_tok[s, P - 1] # the state after the last accepted token +window = conv caches cs_*[s] columns P..P+2 # raw inputs at positions -3..-1 +per token t = 0..num_spec: + g = bf16(f_a[t] @ w_fb^T) (or g_ext[t]: the unfused f_b output) + q, k, v = SiLU(conv4(window, raw[t])); q, k L2-normalized (q *= scale) + beta = sigmoid(b[t]); decay = exp(lower_bound * sigmoid(exp(a_log) * (g + dt_bias))) + S *= decay (per key); S += beta (v - S k) k^T; o = S q + out[t] = o * rsqrt(mean(o^2) + eps) * onorm_w * sigmoid(og[t]) +ssm[s] = the state after t = 0 (the golden token); state_tok[s, t - 1] = the state after draft t +cs_*[s] = the raw inputs at positions -2..num_spec around the golden token +``` + +Returns `out` bf16 `[N (1 + num_spec), H, 128]`. The recurrence is `kda_mtp_decode`'s V-split arithmetic over the +verify tokens, unrolled: with `g_ext` the op is bit-exact against `kda_mtp_decode` fed the same gate and replaying the +same accepted drafts (the op's own test). Certified here against a float64 verify over each request's committed +history: every round's outputs within 2e-2 relative (bf16) and the committed states within 1e-3, the per-draft states +and conv caches through the next round, which starts from them. Repeated runs, and graph replays against eager runs, +are bit-identical. + +## Signature + +```python +def k3_kda_verify( + proj: torch.Tensor, + w_fb: torch.Tensor, + w_q: torch.Tensor, + w_k: torch.Tensor, + w_v: torch.Tensor, + a_log: torch.Tensor, + dt_bias: torch.Tensor, + onorm_w: torch.Tensor, + cs_q: torch.Tensor, + cs_k: torch.Tensor, + cs_v: torch.Tensor, + ssm: torch.Tensor, + state_tok: torch.Tensor, + slots: torch.Tensor, + pending: torch.Tensor, + num_spec: int, + lower_bound: float, + scale: float, + eps: float, + g_ext: Optional[torch.Tensor] = None, +) -> torch.Tensor +``` + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `proj` | `[N (1 + num_spec), cols]`, cols a multiple of 8 and >= 4 H K + K + H | bf16 | contiguous | CUDA | +| `w_fb` | `[H K, K]` (f_b, out x in) | bf16 | contiguous | CUDA | +| `w_q`, `w_k`, `w_v` | `[H K, 4]` conv taps, oldest input first | fp32 | dense | CUDA | +| `a_log` | `[H]` | fp32 | dense | CUDA | +| `dt_bias` | `[H K]` | fp32 | dense | CUDA | +| `onorm_w` | `[K]` | fp32 | dense | CUDA | +| `cs_q`, `cs_k`, `cs_v` | `[pool, H K, 3 + num_spec]` | fp32 | channel stride 1 (dim-contiguous) | CUDA | +| `ssm` | `[pool, H, K, K]` (V rows, K contiguous) | fp32 | each slot dense, slots at any stride | CUDA | +| `state_tok` | `[pool, num_spec, H, K, K]` | fp32 | contiguous | CUDA | +| `slots` | `[N]` | int32 | any element offset | CUDA | +| `pending` | `[pool]`: drafts the sampler accepted last round, per slot | int32 | any element offset | CUDA | +| `num_spec` | drafts per request | Python int | — | — | +| `lower_bound`, `scale`, `eps` | scalars (Kimi K3: -5.0, 128^-0.5, 1e-5) | Python float | — | — | +| `g_ext` | `[N (1 + num_spec), H K]` or None (fold f_b into the kernel) | bf16 | contiguous | CUDA | +| returns | `[N (1 + num_spec), H, K]` | bf16 | contiguous, fresh | CUDA | + +K = V = 128, conv width 4. Kimi K3's TP16 rank slice is H = 6, num_spec = 7. `mutates_args`: `cs_q`, `cs_k`, `cs_v`, +`ssm`, `state_tok`. + +## State + +**Object.** None of its own (stateful kinds P and R): the op updates the caller's pools in place and reads the +record the sampler wrote after the previous round. No buffer persists in the op. + +**The pools and the record.** The cache manager's: with the V2 hybrid manager built with the KDA replay caches and +`kda_token_states` (`kda_replay_num_spec` = num_spec), layer `l`'s `mamba_layer_cache(l).kda_conv_q` / `kda_conv_k` +/ `kda_conv_v` (`cs_*`), `get_ssm_states(l)` (`ssm`, its slots strided by the manager's per-slot coalescing) and +`mamba_layer_cache(l).kda_state_tok`. `pending` is the manager's `prev_num_accepted_tokens`: one record shared by +every layer, written between steps by the sampler's acceptance. A call on layer `l` writes only layer `l`'s views, at +its slots: certified on a real manager, the other slots bit-unchanged. + +**Call-order invariant.** Per layer, one call per verify round, in round order, with `pending` updated between rounds +(after the sampler accepts) and not during one. A call reads what the previous round's call on that layer left at the +slot: the state after the last accepted token (`ssm` if none was accepted, else `state_tok[s, P - 1]`) and the conv +window starting at column P. The per-draft states of a round are read only by the next round. + +**The PDL rule.** The kernel reads the slot, P, the starting state and the conv window before its grid-dependency +wait, and lets its dependents launch in its prologue, before that wait. So a launch must not follow another launch on +the same pools without a kernel that waits (or a non-PDL kernel) in between, and the launch in between must not be +one that lets its dependents launch before its own wait either: one `k3_kda_verify` launch on other pools is not +enough. In the model, a layer's other kernels sit between its launches. + +**Why the test drives call sequences.** The record and the per-draft states make each round depend on the previous +one's acceptance, which no single call can check. The test therefore runs 6 rounds of every layer with a pending count +per request that changes every round (every count 0..num_spec over the rounds), against the float64 history, and +captures a round of every layer once, replaying it with rewritten rows and records: the replays give the bits of the +same rounds run eagerly, and the schedule run twice gives the same bits. + +**What a wrong order does.** Two rounds of one request in swapped order (the negative control): nothing raises, and +the second round's outputs and the slot's state are those of a different history. Measured on sm_100: +TBD(tray: test_modeling_v2_k3_kda_verify.py::test_swapped_rounds_are_silently_wrong, its printed rel diffs). + +## Metadata consumed + +`slots` and `pending`, both read before the grid-dependency wait: under CUDA graphs they must be written before the +graph runs (the step's preparation and the previous step's acceptance). + +## Preconditions + +- sm_100: the kernel uses tcgen05, TMA and clusters. +- The op checks `proj` (bf16, rank 2, N (1 + num_spec) rows, columns), `w_fb`, V = K = 128, `ssm` / `state_tok` + (shapes, fp32, `ssm.stride()[1:] == (K K, K, 1)`, `state_tok` contiguous), `cs_q`'s window width, the int32 index + tensors and the fp32 dtypes, and raises `ValueError` before a launch. A conv cache that is not dense raises + `ValueError` too. The fp32 weights' shapes, `cs_k` / `cs_v`, `pending` covering the slots and the slots' range are + the caller's obligation. +- The first call of each configuration (H, num_spec, `lower_bound`, `scale`, `eps`, with or without `g_ext`, and + whether PDL is on) compiles the kernel and must run outside CUDA-graph capture; under capture it raises + `RuntimeError`. PDL follows `TRTLLM_ENABLE_PDL` (default on), read at every call. +- A call names each slot once. + +## Notes + +- `slots` and `pending` may be slices of longer index tensors starting at any element (the mixer passes + `state_indices[num_prefills:]`); the op's own test certifies the bits at offsets 0..3. +- The pools are addressed at 64-bit slot offsets, so pools past 2 GiB are fine (the op's own test, + `test_k3_kda_pools_past_2g.py`). +- `ssm/k3_kda_attn` computes this verify fused with the layer's projection, for one request of 8 tokens. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_verify.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_verify.py new file mode 100644 index 000000000000..822f0d29e23c --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_verify.py @@ -0,0 +1,60 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3's KDA speculative verify of N requests of 1 + num_spec tokens, from the fused projection rows to the +gated-norm core output, committing the state after every verify token.""" + +from typing import Optional + +import torch + +from tensorrt_llm._torch.cute_dsl_kernels.k3_kda_verify import ( # noqa: F401 — registers the op + op as _k3_kda_verify_op, +) + + +def k3_kda_verify( + proj: torch.Tensor, + w_fb: torch.Tensor, + w_q: torch.Tensor, + w_k: torch.Tensor, + w_v: torch.Tensor, + a_log: torch.Tensor, + dt_bias: torch.Tensor, + onorm_w: torch.Tensor, + cs_q: torch.Tensor, + cs_k: torch.Tensor, + cs_v: torch.Tensor, + ssm: torch.Tensor, + state_tok: torch.Tensor, + slots: torch.Tensor, + pending: torch.Tensor, + num_spec: int, + lower_bound: float, + scale: float, + eps: float, + g_ext: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Return the gated-norm core output bf16 [N (1 + num_spec), H, 128] of the verify tokens whose fused + projection rows are ``proj``; updates the slots' conv caches, state and per-draft states in place.""" + return torch.ops.trtllm.k3_kda_verify( + proj, + w_fb, + w_q, + w_k, + w_v, + a_log, + dt_bias, + onorm_w, + cs_q, + cs_k, + cs_v, + ssm, + state_tok, + slots, + pending, + num_spec, + lower_bound, + scale, + eps, + g_ext, + ) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/kda_decode.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/kda_decode.md new file mode 100644 index 000000000000..07bca5fd177f --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/kda_decode.md @@ -0,0 +1,135 @@ +--- +receipts: {} +--- + +# kda_decode + +**Wraps** `torch.ops.trtllm.kda_decode` (one call). + +One-token KDA (Kimi Delta Attention) decode of B requests: the causal depthwise conv over each request's conv window, +the gated delta rule on its recurrent state, and the gated RMS norm, with the state (and optionally the conv windows) +updated in place at each request's slot. Certified here at Kimi K3's TP16 rank slice (6 heads, K = V = 128, conv +width 4) on a real cache manager's pools. + +## Semantics + +For request b (slot `s = ssm_state_indices[b]`), head h, with `x_q`, `x_k`, `x_v` the new raw q, k, v inputs: + +``` +q, k, v = SiLU(conv4(window[s], new raw) + bias) # taps w_*_t[j], oldest input first; bias_* added before SiLU +q = q / sqrt(sum(q^2) + 1e-6) * scale; k = k / sqrt(sum(k^2) + 1e-6) +beta = sigmoid(beta) if apply_beta_sigmoid else beta +decay = exp(lower_bound * sigmoid(exp(a_log) * (g + dt_bias))) if use_lower_bound + exp(-exp(a_log) * softplus(g + dt_bias)) otherwise +S = state[s] * decay (per key); S += beta (v - S k) k^T; state[s] = S; o = S q +output = o * rsqrt(mean(o^2) + onorm_eps) * onorm_weight * sigmoid(onorm_g) if apply_onorm, else o +``` + +With `update_conv_cache` the conv windows at the slot shift by one (the last two raw inputs, then the new one). The +float64 reference of the test is this arithmetic; the op matches it within fp32 tolerance (outputs 2e-2 relative, +bf16; state rows 1e-3), and the conv windows bit for bit (raw bf16 inputs). + +**Kernels.** On sm_100 and sm_103 the dispatcher picks by the workload B x H: up to 32 the four-CTA cluster kernel, up +to 144 the legacy compact-heads kernel, above that per-architecture choices among the optimized single-CTA and bulk +kernels and the legacy many-heads kernel. At Kimi K3's 6 heads: B <= 5 runs the cluster kernel and 6 <= B <= 24 the +legacy compact-heads kernel (whose block reduction #19830 orders with `__syncwarp`). The test runs B = 1..8, so both. +Other architectures run the legacy kernels. + +## Signature + +```python +def kda_decode( + x_q: torch.Tensor, + x_k: torch.Tensor, + x_v: torch.Tensor, + w_q_t: torch.Tensor, + w_k_t: torch.Tensor, + w_v_t: torch.Tensor, + bias_q: torch.Tensor, + bias_k: torch.Tensor, + bias_v: torch.Tensor, + conv_state_q: torch.Tensor, + conv_state_k: torch.Tensor, + conv_state_v: torch.Tensor, + a_log: torch.Tensor, + g: torch.Tensor, + dt_bias: torch.Tensor, + beta: torch.Tensor, + onorm_g: torch.Tensor, + onorm_weight: torch.Tensor, + ssm_state_indices: Optional[torch.Tensor], + state: torch.Tensor, + apply_onorm: bool, + update_conv_cache: bool, + use_lower_bound: bool, + apply_beta_sigmoid: bool, + lower_bound: float, + scale: float, + onorm_eps: float, + output: torch.Tensor, +) -> None +``` + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `x_q`, `x_k`, `x_v` | `[1, B, H, 128]` | bf16 | head and channel axes packed; any row stride | CUDA | +| `w_q_t`, `w_k_t`, `w_v_t` | `[4, H * 128]` conv taps, oldest input first | bf16 | contiguous | CUDA | +| `bias_q`, `bias_k`, `bias_v` | `[H * 128]` (zeros: no bias) | bf16 | contiguous | CUDA | +| `conv_state_q/k/v` | with `update_conv_cache`: `[slots, H * 128, 3]` section views of one packed `[slots, 3 H 128, 3]` pool (`q | k | v`); else `[B, H * 128, 3]` | bf16 | packed: equal slot strides >= `3 H 128 * 3`, strides `(*, 3, 1)`; else contiguous | CUDA | +| `a_log` | `[H]` | fp32 | contiguous | CUDA | +| `g` | `[1, B, H, 128]` (the gate before `dt_bias`) | bf16 | head and channel axes packed | CUDA | +| `dt_bias` | `[H * 128]` | fp32 | contiguous | CUDA | +| `beta` | `[1, B, H]` | bf16 | head axis stride 1 | CUDA | +| `onorm_g` | `[1, B, H, 128]` (the output gate) | bf16 | head and channel axes packed | CUDA | +| `onorm_weight` | `[128]` | fp32 | contiguous | CUDA | +| `ssm_state_indices` | `[B]` slots, or None (state rows 0..B-1) | int32 | contiguous | CUDA | +| `state` | `[slots >= B, H, 128, 128]` (V rows, K contiguous) | fp32 | each slot dense; slot stride a multiple of 4; 16-byte aligned base | CUDA | +| `apply_onorm`, `update_conv_cache`, `use_lower_bound`, `apply_beta_sigmoid` | scalars (Kimi K3: all True) | Python bool | — | — | +| `lower_bound`, `scale`, `onorm_eps` | scalars (Kimi K3: -5.0, 128^-0.5, 1e-5) | Python float | — | — | +| `output` | `[B, 1, H, 128]` | bf16 | contiguous | CUDA | + +H is one of 1, 2, 3, 4, 6, 8, 12, 16, 24, 32, 48, 64, 96 and the same for q / k and v. Returns None; writes `output`. +Schema mutations (`Tensor(a!)`): `conv_state_q`, `conv_state_k`, `conv_state_v`, `state`, `output`. + +## State + +**Object.** None of its own (stateful kind P): the op updates the caller's pools in place, the recurrent `state` at +each request's slot and, with `update_conv_cache`, the conv windows there. Nothing persists in the op between calls. + +**The pools.** The cache manager's: with the V2 hybrid manager (`conv_state_layout="q_k_v"`, bf16 conv states, fp32 +SSM states), layer `l`'s `get_conv_states(l)` split into its q, k, v sections and `get_ssm_states(l)`, their slots +strided by the manager's per-slot coalescing; `ssm_state_indices` from `get_state_indices`. A call on layer `l` writes +only layer `l`'s views, at its slots: certified on a real manager, with the other slots of the layer and every slot of +the other layers bit-unchanged. + +**Call-order invariant.** One call per layer per decode step, in step order; a batch names each slot once. A call +reads the state and conv window that the previous step's call on that layer left at the slot. + +**Why the test drives call sequences.** Each step's state is the next step's input, so the test runs layers x steps +on a real manager and captures a step of every layer once, replaying it with rewritten inputs: the replays give the +bits of the same steps run eagerly on a copy of the pools. + +**What a wrong order does.** Two steps of one request in swapped order (the negative control): nothing raises, and the +second step's output and the final state are those of a different history. Measured at Kimi K3's shape on sm_100: +TBD(tray: test_modeling_v2_kda_decode.py::test_swapped_steps_are_silently_wrong, its printed rel diffs). + +## Metadata consumed + +`ssm_state_indices` (the requests' slots). None of the attention metadata. + +## Preconditions + +- Every check in the table is the op's (`TORCH_CHECK`): a violation raises `RuntimeError` before the launch, with the + pools unchanged. Among them: the state base must be 16-byte aligned and its slot stride a multiple of 4 floats (the + kernels move state with 16-byte accesses at `slot * stride(0)`); `ssm_state_indices` must be int32. +- On sm_100 and sm_103, the optimized kernels the dispatcher picks for small and large workloads require + `apply_onorm`, `use_lower_bound` and `apply_beta_sigmoid`; with any of them off such a call raises `RuntimeError` + ("Optimized KDA decode requires ...") before the launch. Certified with the lower bound off at B = 2, H = 6. +- The conv inputs, gates and outputs are bf16 and the state fp32; K = V = 128 and conv width 4 only. + +## Notes + +- Kimi K3's KDA layer calls this op through `run_kda_decode_fusion_cuda` + (`_torch/modules/kimi_kda/_kda_decode.py`), which fills the optional bias, gate and norm arguments with cached + dummies; this entry exposes the schema raw. +- `ssm/k3_kda_decode_attn` computes the same decode fused with the layer's projection, for up to 8 requests. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/kda_decode.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/kda_decode.py new file mode 100644 index 000000000000..09dd96c8d95f --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/kda_decode.py @@ -0,0 +1,74 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""One-token KDA decode of B requests: the causal conv, the gated delta rule on each request's state slot and the +gated RMS norm, the conv and state pools updated in place.""" + +from typing import Optional + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + + +def kda_decode( + x_q: torch.Tensor, + x_k: torch.Tensor, + x_v: torch.Tensor, + w_q_t: torch.Tensor, + w_k_t: torch.Tensor, + w_v_t: torch.Tensor, + bias_q: torch.Tensor, + bias_k: torch.Tensor, + bias_v: torch.Tensor, + conv_state_q: torch.Tensor, + conv_state_k: torch.Tensor, + conv_state_v: torch.Tensor, + a_log: torch.Tensor, + g: torch.Tensor, + dt_bias: torch.Tensor, + beta: torch.Tensor, + onorm_g: torch.Tensor, + onorm_weight: torch.Tensor, + ssm_state_indices: Optional[torch.Tensor], + state: torch.Tensor, + apply_onorm: bool, + update_conv_cache: bool, + use_lower_bound: bool, + apply_beta_sigmoid: bool, + lower_bound: float, + scale: float, + onorm_eps: float, + output: torch.Tensor, +) -> None: + """Writes the decode output into ``output`` [B, 1, HV, 128] bf16; updates ``state`` at ``ssm_state_indices`` + and, with ``update_conv_cache``, the conv pools there. Returns None.""" + torch.ops.trtllm.kda_decode( + x_q, + x_k, + x_v, + w_q_t, + w_k_t, + w_v_t, + bias_q, + bias_k, + bias_v, + conv_state_q, + conv_state_k, + conv_state_v, + a_log, + g, + dt_bias, + beta, + onorm_g, + onorm_weight, + ssm_state_indices, + state, + apply_onorm, + update_conv_cache, + use_lower_bound, + apply_beta_sigmoid, + lower_bound, + scale, + onorm_eps, + output=output, + ) diff --git a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_mla_attn_vb_out.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_mla_attn_vb_out.py new file mode 100644 index 000000000000..c966efaa08ae --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_mla_attn_vb_out.py @@ -0,0 +1,445 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the k3_mla_attn_vb_out catalog entry (and its k3_mla_attn / k3_mla_attn_out forms). + +The op reads the paged latent cache through addressing tensors the runtime derives from a KVCacheManager and a +prepared TrtllmAttentionMetadata, and keeps its partials and arrival counters in a caller-owned K3MlaAttnWorkspace. +The test builds that state for real: an MLA (SELFKONLY, kv_factor 1) KVCacheManager of three layers at Kimi K3's +latent width (576 = 512 + 64) and page size (64), every layer's rows written through the manager's own block ids, +a prepared TrtllmAttentionMetadata of R generation requests of T tokens, one TRTLLM MLA attention object per layer +(k3_mla_decode_view gives the op's pool, row stride, page table, layer slot and lengths) and workspaces made by +K3MlaAttnWorkspace.create. The reference reads each request's rows back through the manager's block ids, not +through the op's addressing. + +Covered, at Kimi K3's per-rank shapes (6 heads at TP16, 24 at TP4; latent 512, rope 64, v_head 128), in both +launch modes (up to CLUSTER_WAVE = 7 clusters of 16 CTAs; more take the no-cluster mode, which uses the workspace's +arrival counters): + +1. Cells: the gated output against a float64 reference (attention output rounded to bf16, v_b in float64, the gate + applied as bf16(y * s)); the plain v_b output, k3_mla_attn_out and k3_mla_attn on the same step. +2. Call sequences: layers x decode steps on one shared workspace eagerly, then the same steps captured once as a + CUDA graph and replayed with rewritten inputs and advanced lengths; two workspaces interleaved across layers. + Every output is bit-identical to the same call on a new workspace, and each workspace's counters account for + exactly its own no-cluster launches (16 per launch per request and head group). +3. Negative control: a workspace laid out for another head-group count, or of another dtype or size, is refused + with ValueError before any launch, and neither the output nor the workspace changes. +""" + +from typing import List + +import pytest +import torch + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.k3_mla_attn_vb_out import ( + k3_mla_attn, + k3_mla_attn_out, + k3_mla_attn_vb_out, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.k3_mla_attn_workspace import ( + K3MlaAttnWorkspace, +) +from tensorrt_llm._torch.attention.backends.fmha.cute_dsl_mla import k3_mla_decode_view +from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata +from tensorrt_llm._torch.attention.backends.utils import create_attention +from tensorrt_llm._torch.metadata import KVCacheParams +from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager +from tensorrt_llm.bindings import DataType +from tensorrt_llm.bindings.internal.batch_manager import CacheType +from tensorrt_llm.llmapi.llm_args import KvCacheConfig +from tensorrt_llm.mapping import Mapping + + +def _is_sm100() -> bool: + if not torch.cuda.is_available(): + return False + major, minor = torch.cuda.get_device_capability() + return major * 10 + minor in (100, 103) + + +pytestmark = pytest.mark.skipif(not _is_sm100(), reason="needs an SM100 / SM103 GPU") + +H, NOPE, LAT, PE, QL, V = 6, 128, 512, 64, 1536, 128 +DQK, PAGE = LAT + PE, 64 +SCALE = 1.0 / (NOPE + PE) ** 0.5 +GATE_COL0 = QL + DQK # the gate's columns in the fused projection's rows +LAYERS = 3 +MAX_REQUESTS = 8 +# Request lengths before the first step, assigned in turn: rows crossing a page boundary (579, 1989), > 2048 rows +# (several tiles per CTA), a short context, exactly one page. +LENGTHS = (1100, 64 * 9 + 3, 2049, 127, 4100, 64 * 31 + 5, 300, 64) + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def _no_cluster(num_requests: int, heads: int) -> bool: + """Whether a call of R requests takes the no-cluster mode: more 16-CTA clusters than co-reside, all of whose CTAs + fit on the SMs at once.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import k3_mla_attn_kernel as kernel + + clusters = num_requests * heads // H + sms = torch.cuda.get_device_properties(torch.cuda.current_device()).multi_processor_count + return clusters > kernel.CLUSTER_WAVE and clusters * kernel.CLUSTER <= sms + + +class _MlaCache: + """Real op state: an MLA (SELFKONLY) paged KV cache manager of LAYERS layers holding R requests that grow by T + tokens per decode step, every row they will reach written through the manager's block ids, a + TrtllmAttentionMetadata over it and one TRTLLM MLA attention object per layer for k3_mla_decode_view.""" + + def __init__(self, num_requests: int, tokens: int, steps: int, heads: int, seed: int): + self.tokens, self.heads = tokens, heads + self.request_ids = list(range(num_requests)) + self.cached = [LENGTHS[(i + seed) % len(LENGTHS)] for i in range(num_requests)] + final = [n + steps * tokens for n in self.cached] + pages = sum((n + PAGE - 1) // PAGE for n in final) + self.mgr = KVCacheManager( + KvCacheConfig(max_tokens=(pages + 16) * PAGE, enable_block_reuse=False), + CacheType.SELFKONLY, + num_layers=LAYERS, + num_kv_heads=1, + head_dim=DQK, + tokens_per_block=PAGE, + max_seq_len=((max(final) + PAGE - 1) // PAGE + 1) * PAGE, + max_batch_size=MAX_REQUESTS, + mapping=Mapping(world_size=1, tp_size=1, rank=0), + dtype=DataType.BF16, + ) + self.mgr.add_dummy_requests(self.request_ids, token_nums=final) + self.attn = [ + create_attention( + "TRTLLM", + layer_idx=layer, + num_heads=heads, + head_dim=DQK, + num_kv_heads=1, + is_mla_enable=True, + q_lora_rank=QL, + kv_lora_rank=LAT, + qk_nope_head_dim=NOPE, + qk_rope_head_dim=PE, + v_head_dim=V, + ) + for layer in range(LAYERS) + ] + gen = torch.Generator(device="cuda").manual_seed(seed) + for layer in range(LAYERS): + pool = self.mgr.get_buffers(layer) + assert pool.dtype == torch.bfloat16 and tuple(pool.shape[1:]) == (1, PAGE, 1, DQK) + for blocks, length in zip(self._blocks(layer), final): + for p in range((length + PAGE - 1) // PAGE): + rows = torch.randn(PAGE, DQK, generator=gen, device="cuda") * 0.5 + pool[blocks[p], 0, :, 0] = rows.bfloat16() + + def _blocks(self, layer: int) -> List[List[int]]: + return self.mgr.get_batch_cache_indices(self.request_ids, layer) + + def metadata(self) -> TrtllmAttentionMetadata: + return TrtllmAttentionMetadata( + max_num_requests=MAX_REQUESTS, max_num_tokens=8192, kv_cache_manager=self.mgr + ) + + def prepare(self, md: TrtllmAttentionMetadata) -> None: + """The metadata of the next decode step: every request's cached rows plus its T new tokens.""" + md.seq_lens = torch.tensor([self.tokens] * len(self.request_ids), dtype=torch.int) + md.num_contexts = 0 + md.request_ids = self.request_ids + md.prompt_lens = [c + self.tokens for c in self.cached] + md.kv_cache_params = KVCacheParams( + use_cache=True, num_cached_tokens_per_seq=list(self.cached) + ) + md.prepare() + + def view(self, md: TrtllmAttentionMetadata, layer: int) -> dict: + view = k3_mla_decode_view(self.attn[layer], md, len(self.request_ids) * self.tokens) + assert isinstance(view, dict), view + assert view["row_stride"] == DQK and view["page_offset"] == self.mgr.layer_offsets[layer] + assert view["softmax_scale"] == pytest.approx(SCALE) + return view + + def advance(self) -> None: + self.cached = [c + self.tokens for c in self.cached] + + def reference(self, layer: int, q: torch.Tensor) -> torch.Tensor: + """float64 attention of each request's T tokens over its rows [0, L) read through the manager's block ids; + token t sees rows <= L - T + t. [M, heads, 512].""" + pool, blocks, outs = self.mgr.get_buffers(layer), self._blocks(layer), [] + t = self.tokens + for i, cached in enumerate(self.cached): + length = cached + t + pages = [blocks[i][p] for p in range((length + PAGE - 1) // PAGE)] + kv = pool[pages, 0, :, 0].reshape(-1, DQK)[:length].double() + qi = q[i * t : (i + 1) * t].view(t, self.heads, DQK).double() + s = torch.einsum("thd,ld->thl", qi, kv) * SCALE + limit = length - t + torch.arange(t, device="cuda") + hidden = torch.arange(length, device="cuda")[None, :] > limit[:, None] + s = s.masked_fill(hidden[:, None, :], float("-inf")) + outs.append(torch.einsum("thl,ld->thd", torch.softmax(s, dim=-1), kv[:, :LAT])) + return torch.cat(outs) + + def shutdown(self) -> None: + self.mgr.shutdown() + + +def _inputs(seed: int, m: int, heads: int): + """q (fused_q rows), the v_b weight and the gate (sigmoid values at GATE_COL0 + 128 h).""" + gen = torch.Generator(device="cuda").manual_seed(seed) + q = (torch.randn(m, heads * DQK, generator=gen, device="cuda") * 0.5).bfloat16() + w_vb = (torch.randn(heads, V, LAT, generator=gen, device="cuda") * 0.05).bfloat16() + gate = torch.rand(m, GATE_COL0 + heads * V, generator=gen, device="cuda").bfloat16() + return q, w_vb, gate + + +def _vb(view: dict, q, w_vb, gate, workspace: K3MlaAttnWorkspace) -> torch.Tensor: + m, heads = q.shape[0], q.shape[1] // DQK + out = torch.empty(m, heads * V, dtype=torch.bfloat16, device="cuda") + k3_mla_attn_vb_out( + q, + view["pool"], + view["row_stride"], + view["page_table"], + view["page_offset"], + view["seq_len"], + view["softmax_scale"], + w_vb, + out, + workspace, + gate, + GATE_COL0, + ) + return out + + +def _vb_reference(cache: _MlaCache, layer: int, q, w_vb, gate) -> torch.Tensor: + m, heads = q.shape[0], q.shape[1] // DQK + o = cache.reference(layer, q).bfloat16().double() + y = torch.einsum("thc,hvc->thv", o, w_vb.double()).reshape(m, heads * V) + if gate is None: + return y + return y.bfloat16().double() * gate[:, GATE_COL0:].double() + + +def _max_rel(a, b, tokens: int) -> float: + """max over requests of max |a - b| / max |b| (per request, so a short context is not hidden by a long one).""" + err = 0.0 + for i in range(a.shape[0] // tokens): + ai, bi = ( + a[i * tokens : (i + 1) * tokens].double(), + b[i * tokens : (i + 1) * tokens].double(), + ) + err = max(err, (ai - bi).abs().max().item() / max(bi.abs().max().item(), 1e-6)) + return err + + +def _counters(workspace: K3MlaAttnWorkspace) -> torch.Tensor: + """The no-cluster arrival counters, [8 requests, groups, 2] (the (m, l) exchange's and the drain's).""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import k3_mla_attn_kernel as kernel + + g = workspace.groups + slots = kernel.MAX_REQUESTS * g * kernel.CLUSTER + # fp16 elements before the counters: the partial slots, then the fp32 (m, l) exchange. + start = slots * kernel.WS_SLOT_ELEMS + slots * kernel.ROWS * 4 + words = workspace.buffer[start:].view(torch.int32) + return words.view(kernel.MAX_REQUESTS, g, 2, kernel.CTR_STRIDE)[..., 0].clone() + + +def _check_counters(workspace: K3MlaAttnWorkspace, num_requests: int, launches: int) -> None: + """After `launches` no-cluster launches of R requests, every counter of those requests' head groups is + 16 x launches and every other counter is 0.""" + ctrs = _counters(workspace) + assert bool((ctrs[:num_requests] == 16 * launches).all()), ctrs[:num_requests] + assert bool((ctrs[num_requests:] == 0).all()) + + +# (R, T, heads): T 1 (decode without speculation) and 8 (a DSpark verify step); the cluster and no-cluster modes; +# the TP4 head count (4 head groups: several clusters per request). +CASES = [(3, 8, H), (8, 1, H), (2, 8, 4 * H)] +CASE_IDS = [f"{r}x{t}-h{h}" for r, t, h in CASES] +EAGER_STEPS, GRAPH_STEPS = 2, 2 + + +@pytest.mark.parametrize("num_requests,tokens,heads", CASES, ids=CASE_IDS) +def test_cells(num_requests, tokens, heads): + """One step of layer 0: the gated output within 1e-2 (per request) of the float64 reference and rerun-identical; + the plain v_b output within 1e-2 of its reference, and the gated one bit-identical to bf16(plain * s); + k3_mla_attn_out within 1e-2 of the attention reference and k3_mla_attn bit-identical to it.""" + seed = 17 * num_requests + tokens + heads + cache = _MlaCache(num_requests, tokens, 1, heads, seed) + m = num_requests * tokens + try: + md = cache.metadata() + cache.prepare(md) + view = cache.view(md, 0) + workspace = K3MlaAttnWorkspace.create(torch.device("cuda"), heads // H) + q, w_vb, gate = _inputs(seed, m, heads) + y = _vb(view, q, w_vb, gate, workspace) + again = _vb(view, q, w_vb, gate, workspace) + plain = _vb(view, q, w_vb, None, workspace) + o = torch.empty(m, heads * LAT, dtype=torch.bfloat16, device="cuda") + k3_mla_attn_out( + q, view["pool"], DQK, view["page_table"], view["page_offset"], view["seq_len"], + view["softmax_scale"], o, workspace, + ) # fmt: skip + assert view["page_offset"] == 0 # layer 0: k3_mla_attn takes no page offset + o_new = k3_mla_attn( + q, + view["pool"], + DQK, + view["page_table"], + view["seq_len"], + view["softmax_scale"], + workspace, + ) + torch.cuda.synchronize() + assert _max_rel(y, _vb_reference(cache, 0, q, w_vb, gate), tokens) <= 1e-2 + assert _max_rel(plain, _vb_reference(cache, 0, q, w_vb, None), tokens) <= 1e-2 + assert torch.equal(_bits(y), _bits(again)) + assert torch.equal(_bits(y), _bits(plain * gate[:, GATE_COL0:])) + ref_o = cache.reference(0, q).reshape(m, heads * LAT) + assert _max_rel(o, ref_o, tokens) <= 1e-2 + assert torch.equal(_bits(o_new), _bits(o)) + finally: + cache.shutdown() + + +@pytest.mark.parametrize("num_requests,tokens,heads", CASES, ids=CASE_IDS) +def test_layers_by_steps_eager_then_graph(num_requests, tokens, heads): + """LAYERS layers x EAGER_STEPS decode steps on one shared workspace, eagerly; then one step's LAYERS calls + captured as a CUDA graph on the same workspace and replayed for GRAPH_STEPS steps with rewritten q and the + metadata prepared for each step. Every output within 1e-2 of the reference and bit-identical to the same call + on a new workspace; in the no-cluster mode the shared workspace's counters end at 16 x its launches.""" + seed = 23 * num_requests + tokens + heads + groups = heads // H + cache = _MlaCache(num_requests, tokens, EAGER_STEPS + GRAPH_STEPS, heads, seed) + m = num_requests * tokens + shared = K3MlaAttnWorkspace.create(torch.device("cuda"), groups) + launches = 0 + try: + weights = [_inputs(seed + layer, m, heads)[1:] for layer in range(LAYERS)] + md = cache.metadata() + for step in range(EAGER_STEPS): + cache.prepare(md) + for layer in range(LAYERS): + q = _inputs(1000 * step + 10 * layer + seed, m, heads)[0] + view = cache.view(md, layer) + y = _vb(view, q, *weights[layer], shared) + launches += 1 + fresh = _vb(view, q, *weights[layer], K3MlaAttnWorkspace.create("cuda", groups)) + ref = _vb_reference(cache, layer, q, *weights[layer]) + torch.cuda.synchronize() + assert torch.equal(_bits(y), _bits(fresh)), (step, layer) + assert _max_rel(y, ref, tokens) <= 1e-2, (step, layer) + cache.advance() + + graph_md = cache.metadata().create_cuda_graph_metadata( + num_requests, max_draft_tokens=tokens - 1 + ) + cache.prepare(graph_md) + views = [cache.view(graph_md, layer) for layer in range(LAYERS)] + static_q = [ + torch.zeros(m, heads * DQK, dtype=torch.bfloat16, device="cuda") for _ in range(LAYERS) + ] + static_y = [ + torch.empty(m, heads * V, dtype=torch.bfloat16, device="cuda") for _ in range(LAYERS) + ] + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + for layer in range(LAYERS): + w_vb, gate = weights[layer] + k3_mla_attn_vb_out( + static_q[layer], views[layer]["pool"], DQK, views[layer]["page_table"], + views[layer]["page_offset"], views[layer]["seq_len"], SCALE, w_vb, static_y[layer], + shared, gate, GATE_COL0, + ) # fmt: skip + for step in range(EAGER_STEPS, EAGER_STEPS + GRAPH_STEPS): + if step > EAGER_STEPS: + cache.prepare(graph_md) + qs = [_inputs(1000 * step + 10 * layer + seed, m, heads)[0] for layer in range(LAYERS)] + for layer in range(LAYERS): + static_q[layer].copy_(qs[layer]) + graph.replay() + launches += LAYERS + for layer in range(LAYERS): + fresh = _vb(views[layer], qs[layer], *weights[layer], + K3MlaAttnWorkspace.create("cuda", groups)) # fmt: skip + ref = _vb_reference(cache, layer, qs[layer], *weights[layer]) + torch.cuda.synchronize() + assert torch.equal(_bits(static_y[layer]), _bits(fresh)), (step, layer) + assert _max_rel(static_y[layer], ref, tokens) <= 1e-2, (step, layer) + cache.advance() + torch.cuda.synchronize() + _check_counters(shared, num_requests, launches if _no_cluster(num_requests, heads) else 0) + finally: + cache.shutdown() + + +def test_two_workspaces_interleaved(): + """Two workspaces with the layers' calls alternating between them over 3 no-cluster steps (8 requests at TP16): + every output bit-identical to the same call on a new workspace, and each workspace's counters at 16 x its own + launches.""" + num_requests, tokens = 8, 1 + assert _no_cluster(num_requests, H) + cache = _MlaCache(num_requests, tokens, 3, H, 11) + pair = [K3MlaAttnWorkspace.create(torch.device("cuda"), 1) for _ in range(2)] + launches = [0, 0] + try: + weights = _inputs(12, num_requests, H)[1:] + md = cache.metadata() + for step in range(3): + cache.prepare(md) + for layer in range(LAYERS): + k = (step * LAYERS + layer) % 2 + q = _inputs(100 * step + layer, num_requests, H)[0] + view = cache.view(md, layer) + y = _vb(view, q, *weights, pair[k]) + launches[k] += 1 + fresh = _vb(view, q, *weights, K3MlaAttnWorkspace.create("cuda", 1)) + torch.cuda.synchronize() + assert torch.equal(_bits(y), _bits(fresh)), (step, layer) + cache.advance() + for workspace, n in zip(pair, launches): + _check_counters(workspace, num_requests, n) + finally: + cache.shutdown() + + +def test_workspace_of_another_shape_is_refused(): + """Negative control. A TP16 call (one head group) given a workspace laid out for four head groups, an fp32 + tensor of the same size, or a buffer 8 elements short raises ValueError before any launch: neither the output + nor the workspace changes.""" + cache = _MlaCache(2, 8, 1, H, 13) + try: + md = cache.metadata() + cache.prepare(md) + view = cache.view(md, 0) + q, w_vb, gate = _inputs(14, 16, H) + tp4 = K3MlaAttnWorkspace.create(torch.device("cuda"), 4) + right = K3MlaAttnWorkspace.create(torch.device("cuda"), 1) + out = torch.zeros(16, H * V, dtype=torch.bfloat16, device="cuda") + before = (_bits(tp4.buffer).clone(), _bits(right.buffer).clone()) + for bad in ( + tp4, + K3MlaAttnWorkspace(buffer=right.buffer.float(), groups=1), + K3MlaAttnWorkspace(buffer=right.buffer[:-8], groups=1), + ): + with pytest.raises(ValueError, match="workspace"): + k3_mla_attn_vb_out( + q, view["pool"], DQK, view["page_table"], view["page_offset"], view["seq_len"], SCALE, + w_vb, out, bad, gate, GATE_COL0, + ) # fmt: skip + torch.cuda.synchronize() + assert not bool(out.any()) + assert torch.equal(_bits(tp4.buffer), before[0]) and torch.equal( + _bits(right.buffer), before[1] + ) + finally: + cache.shutdown() + + +def test_create_refuses_capture(): + """A workspace is made eagerly: K3MlaAttnWorkspace.create raises under CUDA-graph capture.""" + graph = torch.cuda.CUDAGraph() + with pytest.raises(RuntimeError, match="capture"): + with torch.cuda.graph(graph): + K3MlaAttnWorkspace.create(torch.device("cuda"), 1) diff --git a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_mla_qkv.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_mla_qkv.py new file mode 100644 index 000000000000..d0f9ab96b4ce --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_mla_qkv.py @@ -0,0 +1,366 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the k3_mla_qkv catalog entry (and its k3_mla_q / k3_mla_qkv_out forms). + +The op writes the paged latent cache through addressing tensors the runtime derives from a KVCacheManager and a +prepared TrtllmAttentionMetadata. The test builds that state for real: an MLA (SELFKONLY, kv_factor 1) KVCacheManager +of three layers at Kimi K3's latent width (576 = 512 + 64) and page size (64), a prepared TrtllmAttentionMetadata of R +generation requests of T tokens, and one TRTLLM MLA attention object per layer, whose k3_mla_decode_view gives the +op's pool, row stride, page table, layer slot (page offset) and lengths. Every pool of the manager starts filled with +a sentinel; the expected image of each layer is kept from the manager's own block ids (get_batch_cache_indices), not +from the op's addressing, and compared with the whole pool after every step. + +Covered, at Kimi K3's per-rank shapes (6 heads at TP16, 24 at TP4; q_lora 1536, latent 512, rope 64): + +1. Cells: fused_q against a reference with the model's bf16 roundings, the cache rows against an fp32 reference + (rope columns bit-exact), k3_mla_q's and k3_mla_qkv_out's results bit-identical to k3_mla_qkv's. +2. Call sequences on real cache objects: layers x decode steps eagerly, then the same steps captured once as a CUDA + graph and replayed with rewritten inputs and advanced lengths; two managers' calls interleaved. +3. Negative control: a call given another layer's slot. Nothing raises; it overwrites that layer's rows at the + step's positions and leaves its own slot unwritten, which the pool comparison sees. +""" + +from typing import Dict + +import pytest +import torch + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.k3_mla_qkv import ( + k3_mla_q, + k3_mla_qkv, + k3_mla_qkv_out, +) +from tensorrt_llm._torch.attention.backends.fmha.cute_dsl_mla import k3_mla_decode_view +from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata +from tensorrt_llm._torch.attention.backends.utils import create_attention +from tensorrt_llm._torch.metadata import KVCacheParams +from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager +from tensorrt_llm.bindings import DataType +from tensorrt_llm.bindings.internal.batch_manager import CacheType +from tensorrt_llm.llmapi.llm_args import KvCacheConfig +from tensorrt_llm.mapping import Mapping + + +def _is_sm100() -> bool: + if not torch.cuda.is_available(): + return False + major, minor = torch.cuda.get_device_capability() + return major * 10 + minor in (100, 103) + + +pytestmark = pytest.mark.skipif(not _is_sm100(), reason="needs an SM100 / SM103 GPU") + +H, NOPE, PE, QK, LAT, QL, V = 6, 128, 64, 192, 512, 1536, 128 +DQK, PAGE = LAT + PE, 64 +EPS = KV_EPS = 1e-6 +LAYERS = 3 +MAX_REQUESTS = 8 +SENTINEL = 0x7F7F # bf16 bits of the largest finite value: no cache row of the test holds it +# Request lengths before the first step, assigned in turn: rows crossing a page boundary (579, 1989), > 2048 rows, +# a short context, exactly one page. +LENGTHS = (1100, 64 * 9 + 3, 2049, 127, 4100, 64 * 31 + 5, 300, 64) + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def _weights(seed: int, heads: int): + gen = torch.Generator(device="cuda").manual_seed(seed) + w_qa = (1.0 + 0.1 * torch.randn(QL, generator=gen, device="cuda")).bfloat16() + w_qb = (torch.randn(heads * QK, QL, generator=gen, device="cuda") * 0.03).bfloat16() + w_kb = (torch.randn(heads, LAT, NOPE, generator=gen, device="cuda") * 0.08).bfloat16() + w_kv = (1.0 + 0.1 * torch.randn(LAT, generator=gen, device="cuda")).bfloat16() + return w_qa, w_qb, w_kb, w_kv + + +def _ag(seed: int, m: int, heads: int) -> torch.Tensor: + """The fused projection's rows: [q_a 1536 | kv_a latent 512 | rope 64 | gate heads * 128].""" + gen = torch.Generator(device="cuda").manual_seed(seed) + return (torch.randn(m, QL + DQK + heads * V, generator=gen, device="cuda") * 0.7).bfloat16() + + +def _reference_q(ag, w_qa, w_qb, w_kb): + """The norm in fp32, the GEMMs in float64, bf16 rounding where the model rounds (norm output, q_b output, q_abs).""" + m, heads = ag.shape[0], w_kb.shape[0] + x = ag[:, :QL].float() + qn = (x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + EPS) * w_qa.float()).bfloat16() + q = (qn.double() @ w_qb.double().t()).bfloat16().view(m, heads, QK) + q_abs = torch.einsum("thd,hcd->thc", q[..., :NOPE].double(), w_kb.double()).bfloat16() + return torch.cat([q_abs, q[..., NOPE:]], dim=-1).reshape(m, heads * DQK) + + +def _reference_kv(ag, w_kv): + """The cache rows in fp32 ((x * rrms) * w, one bf16 rounding); the rope columns copied.""" + x = ag[:, QL : QL + LAT].float() + r = torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + KV_EPS) + return torch.cat([((x * r) * w_kv.float()).bfloat16(), ag[:, QL + LAT : QL + DQK]], dim=-1) + + +def _max_rel(a, b) -> float: + return (a.float() - b.float()).abs().max().item() / b.float().abs().max().item() + + +class _MlaCache: + """Real op state: an MLA (SELFKONLY) paged KV cache manager of LAYERS layers holding R requests that grow by T + tokens per decode step, a TrtllmAttentionMetadata over it, one TRTLLM MLA attention object per layer for + k3_mla_decode_view, and the expected image of every layer's pool.""" + + def __init__(self, num_requests: int, tokens: int, steps: int, heads: int, seed: int): + self.tokens, self.heads = tokens, heads + self.request_ids = list(range(num_requests)) + self.cached = [LENGTHS[(i + seed) % len(LENGTHS)] for i in range(num_requests)] + final = [n + steps * tokens for n in self.cached] + pages = sum((n + PAGE - 1) // PAGE for n in final) + self.mgr = KVCacheManager( + KvCacheConfig(max_tokens=(pages + 16) * PAGE, enable_block_reuse=False), + CacheType.SELFKONLY, + num_layers=LAYERS, + num_kv_heads=1, + head_dim=DQK, + tokens_per_block=PAGE, + max_seq_len=((max(final) + PAGE - 1) // PAGE + 1) * PAGE, + max_batch_size=MAX_REQUESTS, + mapping=Mapping(world_size=1, tp_size=1, rank=0), + dtype=DataType.BF16, + ) + self.mgr.add_dummy_requests(self.request_ids, token_nums=final) + self.attn = [ + create_attention( + "TRTLLM", + layer_idx=layer, + num_heads=heads, + head_dim=DQK, + num_kv_heads=1, + is_mla_enable=True, + q_lora_rank=QL, + kv_lora_rank=LAT, + qk_nope_head_dim=NOPE, + qk_rope_head_dim=PE, + v_head_dim=V, + ) + for layer in range(LAYERS) + ] + self.image: Dict[int, torch.Tensor] = {} + for layer in range(LAYERS): + pool = self.mgr.get_buffers(layer) + assert pool.dtype == torch.bfloat16 and tuple(pool.shape[1:]) == (1, PAGE, 1, DQK) + pool.view(torch.int16).fill_(SENTINEL) + self.image[layer] = torch.full(pool.shape, SENTINEL, dtype=torch.int16, device="cuda") + self.blocks = { + layer: self.mgr.get_batch_cache_indices(self.request_ids, layer) + for layer in range(LAYERS) + } + + def metadata(self) -> TrtllmAttentionMetadata: + return TrtllmAttentionMetadata( + max_num_requests=MAX_REQUESTS, max_num_tokens=8192, kv_cache_manager=self.mgr + ) + + def prepare(self, md: TrtllmAttentionMetadata) -> None: + """The metadata of the next decode step: every request's cached rows plus its T new tokens.""" + md.seq_lens = torch.tensor([self.tokens] * len(self.request_ids), dtype=torch.int) + md.num_contexts = 0 + md.request_ids = self.request_ids + md.prompt_lens = [c + self.tokens for c in self.cached] + md.kv_cache_params = KVCacheParams( + use_cache=True, num_cached_tokens_per_seq=list(self.cached) + ) + md.prepare() + + def view(self, md: TrtllmAttentionMetadata, layer: int) -> dict: + view = k3_mla_decode_view(self.attn[layer], md, len(self.request_ids) * self.tokens) + assert isinstance(view, dict), view + assert view["row_stride"] == DQK and view["page_offset"] == self.mgr.layer_offsets[layer] + return view + + def record(self, layer: int, rows: torch.Tensor) -> None: + """Expect `rows` (the step's cache rows, request-major) at each token's position of `layer`.""" + for i in range(len(self.request_ids)): + for u in range(self.tokens): + pos = self.cached[i] + u + block = self.blocks[layer][i][pos // PAGE] + self.image[layer][block, 0, pos % PAGE, 0] = _bits(rows[i * self.tokens + u]) + + def advance(self) -> None: + self.cached = [c + self.tokens for c in self.cached] + + def pool_matches(self, layer: int) -> bool: + return torch.equal(self.mgr.get_buffers(layer).view(torch.int16), self.image[layer]) + + def shutdown(self) -> None: + self.mgr.shutdown() + + +def _call(view: dict, ag, weights) -> torch.Tensor: + w_qa, w_qb, w_kb, w_kv = weights + return k3_mla_qkv( + ag, + w_qa, + EPS, + w_qb, + w_kb, + w_kv, + KV_EPS, + view["pool"], + view["row_stride"], + view["page_table"], + view["page_offset"], + view["seq_len"], + ) + + +def _dense(ag, weights): + """fused_q and the step's cache rows from the dense form: the bits k3_mla_qkv stores.""" + w_qa, w_qb, w_kb, w_kv = weights + rows = torch.empty(ag.shape[0], DQK, dtype=torch.bfloat16, device="cuda") + y = k3_mla_qkv_out(ag, w_qa, EPS, w_qb, w_kb, w_kv, KV_EPS, rows) + return y, rows + + +def _check_cell(y, rows, ag, weights) -> None: + """fused_q and the stored rows against the references; the q-only form's fused_q bit-identical.""" + w_qa, w_qb, w_kb, w_kv = weights + m = ag.shape[0] + y_q = k3_mla_q(ag, w_qa, EPS, w_qb, w_kb) + ref_q = _reference_q(ag, w_qa, w_qb, w_kb).view(m, -1, DQK) + ref_kv = _reference_kv(ag, w_kv) + torch.cuda.synchronize() + for part in (slice(0, LAT), slice(LAT, DQK)): + assert _max_rel(y.view(m, -1, DQK)[..., part], ref_q[..., part]) <= 2e-2 + assert _max_rel(rows[:, :LAT], ref_kv[:, :LAT]) <= 1e-2 + assert torch.equal(_bits(rows[:, LAT:]), _bits(ag[:, QL + LAT : QL + DQK])) + assert torch.equal(_bits(y), _bits(y_q)) + + +# (R, T, heads): T 1 (decode without speculation) and 8 (a DSpark verify step), up to 8 requests, the TP4 head count. +CASES = [(3, 1, H), (8, 1, H), (2, 8, H), (8, 8, H), (2, 8, 4 * H)] +CASE_IDS = [f"{r}x{t}-h{h}" for r, t, h in CASES] +EAGER_STEPS, GRAPH_STEPS = 2, 2 + + +@pytest.mark.parametrize("num_requests,tokens,heads", CASES, ids=CASE_IDS) +def test_layers_by_steps_eager_then_graph(num_requests, tokens, heads): + """LAYERS layers x EAGER_STEPS decode steps called eagerly, then one step captured as a CUDA graph and replayed + for GRAPH_STEPS steps with rewritten inputs and the metadata prepared for each step: after every step each + layer's whole pool equals its expected image (the step's rows at the manager's blocks, nothing else written), + and every call's outputs pass the cell checks.""" + seed = 31 * num_requests + tokens + heads + weights = [_weights(seed + layer, heads) for layer in range(LAYERS)] + cache = _MlaCache(num_requests, tokens, EAGER_STEPS + GRAPH_STEPS, heads, seed) + m = num_requests * tokens + try: + md = cache.metadata() + for step in range(EAGER_STEPS): + cache.prepare(md) + for layer in range(LAYERS): + ag = _ag(1000 * step + 10 * layer + seed, m, heads) + y = _call(cache.view(md, layer), ag, weights[layer]) + y_dense, rows = _dense(ag, weights[layer]) + _check_cell(y, rows, ag, weights[layer]) + assert torch.equal(_bits(y), _bits(y_dense)) + cache.record(layer, rows) + for layer in range(LAYERS): + assert cache.pool_matches(layer), (step, layer) + cache.advance() + + graph_md = cache.metadata().create_cuda_graph_metadata( + num_requests, max_draft_tokens=tokens - 1 + ) + cache.prepare(graph_md) + views = [cache.view(graph_md, layer) for layer in range(LAYERS)] + static_ag = [torch.zeros(m, QL + DQK + heads * V, dtype=torch.bfloat16, device="cuda") + for _ in range(LAYERS)] # fmt: skip + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + static_y = [ + _call(views[layer], static_ag[layer], weights[layer]) for layer in range(LAYERS) + ] + for step in range(EAGER_STEPS, EAGER_STEPS + GRAPH_STEPS): + if step > EAGER_STEPS: + cache.prepare(graph_md) + inputs = [_ag(1000 * step + 10 * layer + seed, m, heads) for layer in range(LAYERS)] + for layer in range(LAYERS): + static_ag[layer].copy_(inputs[layer]) + graph.replay() + for layer in range(LAYERS): + y_dense, rows = _dense(inputs[layer], weights[layer]) + torch.cuda.synchronize() + assert torch.equal(_bits(static_y[layer]), _bits(y_dense)), (step, layer) + cache.record(layer, rows) + for layer in range(LAYERS): + assert cache.pool_matches(layer), (step, layer) + cache.advance() + finally: + cache.shutdown() + + +def test_two_managers_interleaved(): + """Two managers (two models' caches) with their calls alternating step by step and layer by layer: each pool + holds exactly its own rows.""" + caches = [_MlaCache(3, 8, 3, H, seed) for seed in (5, 6)] + weights = _weights(77, H) + try: + mds = [cache.metadata() for cache in caches] + for step in range(3): + for cache, md in zip(caches, mds): + cache.prepare(md) + for layer in range(LAYERS): + for k, (cache, md) in enumerate(zip(caches, mds)): + ag = _ag(100 * step + 10 * layer + k, 24, H) + _call(cache.view(md, layer), ag, weights) + cache.record(layer, _dense(ag, weights)[1]) + for cache in caches: + for layer in range(LAYERS): + assert cache.pool_matches(layer), (step, layer) + cache.advance() + finally: + for cache in caches: + cache.shutdown() + + +def test_wrong_layer_slot_overwrites_its_neighbour(): + """Negative control. Layer 1's call given layer 0's slot (page offset): nothing raises; layer 0's rows at the + step's positions now hold layer 1's values and layer 1's slot is unwritten. The op trusts the page offset, and the + pool comparison detects the misplacement.""" + cache = _MlaCache(3, 8, 1, H, 9) + weights = _weights(78, H) + try: + md = cache.metadata() + cache.prepare(md) + ag0, ag1 = _ag(1, 24, H), _ag(2, 24, H) + _call(cache.view(md, 0), ag0, weights) + cache.record(0, _dense(ag0, weights)[1]) + assert cache.pool_matches(0) and cache.pool_matches(1) + wrong = dict(cache.view(md, 1), page_offset=cache.view(md, 0)["page_offset"]) + _call(wrong, ag1, weights) + torch.cuda.synchronize() + assert not cache.pool_matches(0) + assert cache.pool_matches(1) # nothing written into layer 1's slot + cache.record(0, _dense(ag1, weights)[1]) + assert cache.pool_matches(0) # layer 0's step rows are now layer 1's + finally: + cache.shutdown() + + +def test_rejects_out_of_contract(): + """More than 64 tokens, or page-table rows / lengths that do not describe R requests of the call's tokens, raise + ValueError before any launch; the pool is left as it was.""" + cache = _MlaCache(2, 4, 1, H, 3) + weights = _weights(79, H) + try: + md = cache.metadata() + cache.prepare(md) + view = cache.view(md, 0) + with pytest.raises(ValueError): # 7 tokens over 2 requests + _call(view, _ag(4, 7, H), weights) + with pytest.raises(ValueError): # 1 page-table row, 2 lengths + _call(dict(view, page_table=view["page_table"][:1]), _ag(4, 8, H), weights) + with pytest.raises(ValueError): # an int64 page table + _call(dict(view, page_table=view["page_table"].long()), _ag(4, 8, H), weights) + with pytest.raises(ValueError): # 65 tokens + k3_mla_q(_ag(5, 65, H), weights[0], EPS, weights[1], weights[2]) + torch.cuda.synchronize() + assert all(cache.pool_matches(layer) for layer in range(LAYERS)) + finally: + cache.shutdown() diff --git a/tests/unittest/_torch/modeling_v2/ssm/_kda_cells.py b/tests/unittest/_torch/modeling_v2/ssm/_kda_cells.py new file mode 100644 index 000000000000..14762d56fe39 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/ssm/_kda_cells.py @@ -0,0 +1,329 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""What the ssm/ entry tests share: a real ``MambaHybridCacheManagerV2`` at Kimi K3's per-rank KDA shapes, its +per-layer pools, and the references of the ops' own tests (k3_kda_qkvg's rows decoded from its Lamport buffers; a +float64 plain decode; a float64 verify that keeps each request's committed history). + +Not a test file (no ``test_`` prefix): the four entry tests import it by name, as ``comm/`` imports ``_rank_job``. + +Shapes: 6 heads, K = V = 128, conv width 4 (the TP16 rank slice); conv states bf16 in the ``[q | k | v]`` layout, +SSM states fp32. The manager coalesces each slot's per-layer states, so every per-layer view is a strided slot view, +the layout the model hands these ops. +""" + +from __future__ import annotations + +import torch + +H = 6 +K = V = 128 +HK = H * K +W = 4 +CONV_DIM = 3 * HK +PROJ = 4 * HK + K + H + 2 # the fused [q | k | v | og | f_a | b | pad] rows (3208) +K_IN = 7168 +NUM_SPEC = 7 +NT = NUM_SPEC + 1 +LOWER_BOUND = -5.0 +EPS = 1e-5 +SCALE = K**-0.5 +TOL_OUT = 2e-2 # bf16 outputs +TOL_STATE = 1e-3 # fp32 states: approximate exp / rcp in the kernels and the bf16 f_b gate + + +def sm100() -> bool: + return torch.cuda.is_available() and torch.cuda.get_device_capability()[0] == 10 + + +def load_ops(): + """Register the C++ ops and the K3 KDA CuTe DSL ops; return the k3_kda_attn op module (its constants).""" + import tensorrt_llm # noqa: F401 (registers the trtllm ops) + from tensorrt_llm._torch.cute_dsl_kernels.k3_kda_attn import op + from tensorrt_llm._torch.cute_dsl_kernels.k3_kda_verify import op as _verify # noqa: F401 + + return op + + +# --------------------------------------------------------------------------------------------------------------- +# The cache objects +# --------------------------------------------------------------------------------------------------------------- +def build_manager(num_layers: int, num_spec: int | None = None, max_batch_size: int = 8): + """A real MambaHybridCacheManagerV2 with ``num_layers`` KDA layers (and one attention layer). With ``num_spec``: + MTP-style speculation of ``num_spec`` drafts, the KDA replay caches and the per-token states + (``kda_token_states``) that k3_kda_verify / k3_kda_attn read and write.""" + from tensorrt_llm._torch.pyexecutor.kv_cache.mamba_cache_manager import ( + MambaHybridCacheManagerV2, + ) + from tensorrt_llm._torch.pyexecutor.resource_manager import CacheTypeCpp, DataType + from tensorrt_llm.llmapi.llm_args import KvCacheConfig, MTPDecodingConfig + from tensorrt_llm.mapping import Mapping + + spec = MTPDecodingConfig(max_draft_len=num_spec) if num_spec else None + return MambaHybridCacheManagerV2( + mamba_d_state=K, + mamba_d_conv=W, + mamba_num_heads=H, + mamba_n_groups=H, + mamba_head_dim=V, + mamba_num_layers=num_layers, + mamba_layer_mask=[True] * num_layers + [False], + mamba_cache_dtype=torch.bfloat16, + mamba_ssm_cache_dtype=torch.float32, + kv_cache_config=KvCacheConfig(max_tokens=512, enable_block_reuse=False), + kv_cache_type=CacheTypeCpp.SELF, + num_layers=1, + num_kv_heads=4, + head_dim=64, + tokens_per_block=32, + max_seq_len=128, + max_batch_size=max_batch_size, + mapping=Mapping(world_size=1, rank=0, tp_size=1, pp_size=1), + dtype=DataType.HALF, + spec_config=spec, + layer_mask=[False] * num_layers + [True], + vocab_size=1024, + conv_state_layout="q_k_v", + kda_replay_num_spec=num_spec, + kda_token_states=num_spec is not None, + ) + + +def request_slots(mgr, count: int, first_id: int) -> torch.Tensor: + """State slots of ``count`` new requests (the manager's own slot assignment), int32 on the GPU.""" + ids = list(range(first_id, first_id + count)) + mgr.add_dummy_requests(ids, token_nums=[8] * count, is_gen=False) + slots = mgr.get_state_indices(ids, [False] * count) + assert len(set(slots)) == count, slots + return torch.tensor(slots, dtype=torch.int32, device="cuda") + + +def layer_pools(mgr, layer: int) -> dict: + """The manager's views of one KDA layer's pools: ``conv`` bf16 [slots, 3 HK, W - 1] (q | k | v channels) and + ``ssm`` fp32 [slots, H, V, K] (plain decode); with the replay caches also ``cs_q`` / ``cs_k`` / ``cs_v`` fp32 + [slots, HK, W - 1 + num_spec] (dim-contiguous), ``state_tok`` fp32 [slots, num_spec, H, V, K] and ``pending``, + the accepted-draft record every layer shares (``prev_num_accepted_tokens``).""" + pools = {"conv": mgr.get_conv_states(layer), "ssm": mgr.get_ssm_states(layer)} + if mgr.use_kda_replay_update: + cache = mgr.mamba_layer_cache(layer) + pools.update( + cs_q=cache.kda_conv_q, + cs_k=cache.kda_conv_k, + cs_v=cache.kda_conv_v, + state_tok=cache.kda_state_tok, + pending=cache.prev_num_accepted_tokens, + ) + return pools + + +def fill_pools(pools: dict, seed: int) -> None: + """Random finite contents in every slot of every pool (the pending record zeroed).""" + g = torch.Generator(device="cuda").manual_seed(seed) + + def rnd(t, scale): + return torch.randn(t.shape, generator=g, device="cuda") * scale + + pools["conv"].copy_(rnd(pools["conv"], 0.5).bfloat16()) + pools["ssm"].copy_(rnd(pools["ssm"], 0.05)) + for name in ("cs_q", "cs_k", "cs_v"): + if name in pools: + pools[name].copy_(rnd(pools[name], 0.5)) + if "state_tok" in pools: + pools["state_tok"].zero_() + if "pending" in pools: + pools["pending"].zero_() + + +def clone_pools(pools: dict) -> dict: + """Standalone copies with the same strides (so the same layout constraints hold), e.g. for a reference path.""" + out = {} + for name, t in pools.items(): + c = torch.empty_strided(t.shape, t.stride(), dtype=t.dtype, device=t.device) + c.copy_(t) + out[name] = c + return out + + +def snapshot(pools: dict) -> dict: + return {name: t.clone() for name, t in pools.items()} + + +def same_rows(a: dict, b: dict, rows, names) -> bool: + """Bit equality of the given slots (``rows``: list of ints) of the given pools.""" + return all(torch.equal(a[n][rows], b[n][rows]) for n in names) + + +def other_slots(num_slots: int, used) -> list: + used = set(int(s) for s in used) + return [s for s in range(num_slots) if s not in used] + + +def rel(a: torch.Tensor, b: torch.Tensor) -> float: + return ((a.float() - b.float()).abs().max() / b.float().abs().max().clamp_min(1e-6)).item() + + +# --------------------------------------------------------------------------------------------------------------- +# Weights and the projection rows +# --------------------------------------------------------------------------------------------------------------- +def make_weights(seed: int) -> dict: + """One KDA layer's decode weights at the TP16 shapes; the conv taps are bf16 values (kda_decode reads them as + bf16, the K3 kernels as fp32).""" + g = torch.Generator(device="cuda").manual_seed(seed) + + def rnd(*s, scale=1.0): + return torch.randn(*s, generator=g, device="cuda") * scale + + conv = rnd(3, HK, W, scale=0.3).bfloat16().float() + return { + "w": rnd(PROJ, K_IN, scale=0.02).bfloat16(), + "w_fb": rnd(HK, K, scale=0.05).bfloat16(), + "w_q": conv[0].contiguous(), + "w_k": conv[1].contiguous(), + "w_v": conv[2].contiguous(), + "w_t": [conv[i].t().bfloat16().contiguous() for i in range(3)], + "a_log": rnd(H, scale=0.5), + "dt_bias": rnd(HK, scale=0.5), + "onorm_w": (1 + 0.1 * rnd(V)).float(), + } + + +def decode_rows(p1: torch.Tensor, part: torch.Tensor, buf: int, n: int) -> torch.Tensor: + """The ``[q | k | v | og | f_a | b | pad]`` rows of tokens 0..n-1 from buffer ``buf`` of a k3_kda_qkvg launch: + q, k and f_a are bf16 bits; v, og and b the bf16 sum of the two fp32 K-half partials.""" + qkfa = p1.view(3, 8, 2 * HK + K)[buf, :n].view(torch.bfloat16) + parts = part.view(3, 3, 2, 8, HK)[buf].view(torch.float32) + v, og, b = ((parts[r, 0, :n] + parts[r, 1, :n]).bfloat16() for r in range(3)) + rows = torch.zeros(n, PROJ, dtype=torch.bfloat16, device=p1.device) + rows[:, : 2 * HK] = qkfa[:, : 2 * HK] + rows[:, 2 * HK : 3 * HK] = v + rows[:, 3 * HK : 4 * HK] = og + rows[:, 4 * HK : 4 * HK + K] = qkfa[:, 2 * HK :] + rows[:, 4 * HK + K : 4 * HK + K + H] = b[:, :H] + return rows + + +# --------------------------------------------------------------------------------------------------------------- +# float64 references +# --------------------------------------------------------------------------------------------------------------- +def _gate_decay(a_log, dt_bias, g): + """The lower-bound gate: decay = exp(lower_bound * sigmoid(exp(A_log) * (g + dt_bias))), per key channel.""" + xg = torch.exp(a_log.double()).unsqueeze(-1) * (g + dt_bias.double().view(H, K)) + return torch.exp(LOWER_BOUND * torch.sigmoid(xg)) + + +class F64Decode: + """The plain KDA decode in float64 torch over copies of a layer's ``conv`` / ``ssm`` pools: conv4 + SiLU, q / k + L2 norm (q scaled), beta sigmoid, the lower-bound gate, S <- S d; S <- S + beta (v - S k) k^T; o = S q, the + gated RMSNorm.""" + + def __init__(self, wt: dict, pools: dict): + self.wt = wt + self.conv = pools["conv"].double().clone() + self.state = pools["ssm"].double().clone() + + def step(self, raw_qkv, g, beta_raw, gate_raw, slots) -> torch.Tensor: + """``raw_qkv`` [n, 3 HK] (the projection's q | k | v columns), ``g`` [n, HK] (f_b's output, before dt_bias), + ``beta_raw`` [n, H], ``gate_raw`` [n, HK] (the output gate); returns [n, H, V].""" + wt = self.wt + conv_w = torch.stack([wt["w_q"], wt["w_k"], wt["w_v"]]).double() # [3, HK, W] + outs = torch.empty(raw_qkv.shape[0], H, V, dtype=torch.float64, device="cuda") + for i, s in enumerate(slots.tolist()): + new = raw_qkv[i].double().view(3, HK) + win = self.conv[s].view(3, HK, W - 1) + u = torch.cat([win, new.unsqueeze(-1)], dim=-1) # oldest first + act = (u * conv_w).sum(-1) + act = act * torch.sigmoid(act) + self.conv[s] = u[:, :, 1:].reshape(3 * HK, W - 1) + q, k, v = (act[j].view(H, K) for j in range(3)) + q = q / torch.sqrt((q * q).sum(-1, keepdim=True) + 1e-6) * SCALE + k = k / torch.sqrt((k * k).sum(-1, keepdim=True) + 1e-6) + beta = torch.sigmoid(beta_raw[i].double()) + decay = _gate_decay(wt["a_log"], wt["dt_bias"], g[i].double().view(H, K)) + st = self.state[s] * decay.unsqueeze(1) + res = (v - (st * k.unsqueeze(1)).sum(-1)) * beta.unsqueeze(-1) + st = st + res.unsqueeze(-1) * k.unsqueeze(1) + self.state[s] = st + o = (st * q.unsqueeze(1)).sum(-1) + rms = torch.rsqrt((o * o).mean(-1, keepdim=True) + EPS) + gate = torch.sigmoid(gate_raw[i].double().view(H, V)) + outs[i] = o * rms * wt["onorm_w"].double() * gate + return outs + + def step_rows(self, rows: torch.Tensor, slots) -> torch.Tensor: + """As :meth:`step` from fused projection rows (f_b in float64, rounded to bf16 as the GEMM's output is).""" + r = rows.double() + g = (r[:, 4 * HK : 4 * HK + K] @ self.wt["w_fb"].double().t()).bfloat16().double() + return self.step( + r[:, : 3 * HK], g, r[:, 4 * HK + K : 4 * HK + K + H], r[:, 3 * HK : 4 * HK], slots + ) + + +def _conv_silu(win, raw, c, w): + x = win[0][c] * w[:, 0] + win[1][c] * w[:, 1] + win[2][c] * w[:, 2] + raw[c] * w[:, 3] + return x * torch.sigmoid(x) + + +class F64Verify: + """The speculative verify in float64 over each request's committed history (raw conv inputs and the state after + the last committed token): every request's 1 + num_spec tokens from the state after its accepted drafts. The + conv caches start with pending 0, so their window columns 0..2 are the committed raw inputs.""" + + def __init__(self, wt: dict, pools: dict, slots, num_spec: int): + self.wt, self.num_spec = wt, num_spec + self.slots = slots.tolist() + self.seq = [ + [ + torch.stack([pools[c][s, :, i].double() for c in ("cs_q", "cs_k", "cs_v")]) + for i in range(3) + ] + for s in self.slots + ] + self.state = [pools["ssm"][s].double().clone() for s in self.slots] + self.last = None + + def __call__(self, rows: torch.Tensor) -> torch.Tensor: + """Outputs [N (1 + num_spec), H, V] for the fused projection rows of N requests.""" + wt, steps = self.wt, self.num_spec + 1 + r = rows.double() + g_all = (r[:, 4 * HK : 4 * HK + K] @ wt["w_fb"].double().t()).bfloat16().double() + wq, wk, wv = (wt[n].double() for n in ("w_q", "w_k", "w_v")) + out = torch.empty(rows.shape[0], H, V, dtype=torch.float64, device="cuda") + self.last = [] + for n in range(len(self.slots)): + seq, s_cur = list(self.seq[n]), self.state[n] + states, raws = [], [] + for t in range(steps): + row = n * steps + t + raw = r[row, : 3 * HK].view(3, HK) + win = seq[-3:] + q = _conv_silu(win, raw, 0, wq).view(H, K) + k = _conv_silu(win, raw, 1, wk).view(H, K) + v = _conv_silu(win, raw, 2, wv).view(H, V) + q = q * torch.rsqrt((q * q).sum(-1, keepdim=True) + 1e-6) * SCALE + k = k * torch.rsqrt((k * k).sum(-1, keepdim=True) + 1e-6) + beta = torch.sigmoid(r[row, 4 * HK + K : 4 * HK + K + H]) + decay = _gate_decay(wt["a_log"], wt["dt_bias"], g_all[row].view(H, K)) + sd = s_cur * decay[:, None, :] + vn = v - torch.einsum("hvk,hk->hv", sd, k) + s_cur = sd + beta[:, None, None] * vn[:, :, None] * k[:, None, :] + o = torch.einsum("hvk,hk->hv", s_cur, q) + rms = torch.rsqrt((o * o).mean(-1, keepdim=True) + EPS) + gate = torch.sigmoid(r[row, 3 * HK : 4 * HK].view(H, V)) + out[row] = o * rms * wt["onorm_w"].double() * gate + states.append(s_cur) + raws.append(raw) + seq.append(raw) + self.last.append((states, raws)) + return out + + def commit(self, accepted) -> None: + """The sampler accepted ``accepted[n]`` drafts: the golden token and those drafts commit.""" + for n, p in enumerate(accepted): + states, raws = self.last[n] + self.seq[n] = (self.seq[n] + raws[: p + 1])[-3:] + self.state[n] = states[p] + + +def pending_schedule(num_requests: int, num_spec: int, rnd: int) -> list: + """Accepted drafts per request after round ``rnd``: every count 0..num_spec over the rounds, different per + request.""" + return [(3 * n + 5 * rnd + 1 + (rnd * n) % 3) % (num_spec + 1) for n in range(num_requests)] diff --git a/tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_k3_kda_attn.py b/tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_k3_kda_attn.py new file mode 100644 index 000000000000..f05c0dbe78ca --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_k3_kda_attn.py @@ -0,0 +1,292 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the ssm/k3_kda_attn catalog entry (and its sibling k3_kda_qkvg), on a real cache object and a +caller-owned K3KdaBuffers. + +Kimi K3's fused KDA projection + speculative verify of one request's golden token and 7 drafts, at the TP16 rank +slice, on the verify pools of a real ``MambaHybridCacheManagerV2`` built with MTP-style speculation of 7 drafts, the +KDA replay caches and per-token states (``_kda_cells.build_manager``): the layer's conv caches (fp32, dim-contiguous), +its fp32 state (strided by the manager's per-slot coalescing), its per-draft states and the accepted-draft record +every layer shares (``prev_num_accepted_tokens``), at the slots the manager assigned. + +The reference, as the op's own test: the projection stream alone (``k3_kda_qkvg``, its Lamport buffers decoded into +rows) and the ``ssm/k3_kda_verify`` entry on them, on a copy of the pools: outputs and every pool bit for bit. Layer 0 +is also checked against a float64 verify over the request's committed history. Call sequences: layers x rounds on one +shared set against a set per layer and a captured round replayed, the per-CTA index across the int32 wrap; negative +controls. + +Every launch here follows a plain copy kernel (its input), as each KDA launch in the model follows other layers' +kernels: the head CTAs read the slot's pools before their grid-dependency wait, so a launch must not directly follow +another launch on the same pools. +""" + +import _kda_cells as kc +import pytest +import torch + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.ssm.k3_kda_attn import ( + k3_kda_attn, + k3_kda_qkvg, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.ssm.k3_kda_buffers import ( + CTAS, + K3KdaBuffers, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.ssm.k3_kda_verify import k3_kda_verify + +pytestmark = pytest.mark.skipif(not kc.sm100(), reason="needs SM100 (tcgen05, TMA, clusters)") + +LAYERS = 2 +SLOT_ROUNDS = 9 # pending 0, then the schedule's 1, 6, 3, 0, 5, 2, 7, 4 +POOL_NAMES = ("cs_q", "cs_k", "cs_v", "ssm", "state_tok") + + +@pytest.fixture(scope="module") +def mgr(): + kc.load_ops() + m = kc.build_manager(LAYERS, num_spec=kc.NUM_SPEC) + try: + yield m + finally: + m.shutdown() + + +@pytest.fixture(scope="module") +def slots(mgr): + return kc.request_slots(mgr, 4, first_id=400) + + +def _layers(mgr) -> list: + return [kc.layer_pools(mgr, layer) for layer in range(LAYERS)] + + +def _fused(wt, pools, x_src, slot, bufs) -> torch.Tensor: + """One launch, after a plain copy kernel (the input).""" + x = x_src.clone() + return k3_kda_attn( + x, wt["w"], wt["w_fb"], wt["w_q"], wt["w_k"], wt["w_v"], wt["a_log"], wt["dt_bias"], wt["onorm_w"], + pools["cs_q"], pools["cs_k"], pools["cs_v"], pools["ssm"], pools["state_tok"], slot, pools["pending"], bufs, + kc.NUM_SPEC, kc.LOWER_BOUND, kc.SCALE, kc.EPS, + ) # fmt: skip + + +class _Unfused: + """k3_kda_qkvg's rows (its Lamport buffers decoded), then the ssm/k3_kda_verify entry on them.""" + + def __init__(self): + self.bufs = K3KdaBuffers.create("cuda", ctas=CTAS) + self.last_rows = None + + def rows(self, w, x) -> torch.Tensor: + buf = int(self.bufs.epoch[0].item()) + k3_kda_qkvg(x, w, self.bufs) + return kc.decode_rows(self.bufs.p1, self.bufs.part, buf, x.shape[0]) + + def __call__(self, wt, pools, x, slot) -> torch.Tensor: + self.last_rows = self.rows(wt["w"], x) + return k3_kda_verify( + self.last_rows, wt["w_fb"], wt["w_q"], wt["w_k"], wt["w_v"], wt["a_log"], wt["dt_bias"], wt["onorm_w"], + pools["cs_q"], pools["cs_k"], pools["cs_v"], pools["ssm"], pools["state_tok"], slot, pools["pending"], + kc.NUM_SPEC, kc.LOWER_BOUND, kc.SCALE, kc.EPS, None, + ) # fmt: skip + + +def _x(gen: torch.Generator) -> torch.Tensor: + return torch.randn(kc.NT, kc.K_IN, generator=gen, device="cuda").bfloat16() + + +def _accept(pools_list, slot: int, accepted: int) -> None: + """The sampler's record of the drafts it accepted: one count per slot, shared by every layer of a pool set.""" + for pools in pools_list: + pools["pending"][slot] = accepted + + +def test_attn_on_manager_pools(mgr, slots) -> None: + """Two requests, one after the other, each on its slot for SLOT_ROUNDS rounds of every layer (pending through + 0..7), all launches on one set: outputs and pools bit for bit the unfused path's on a copy; the other slots + untouched; layer 0's outputs also within fp32 tolerance of float64 over the committed history.""" + pools = _layers(mgr) + with torch.inference_mode(): + for layer, p in enumerate(pools): + kc.fill_pools(p, 1000 + layer) + copies = [kc.clone_pools(p) for p in pools] + # The record is one tensor for every layer in the manager; keep the copies' one shared too. + for c in copies[1:]: + c["pending"] = copies[0]["pending"] + wts = [kc.make_weights(500 + layer) for layer in range(LAYERS)] + bufs = K3KdaBuffers.create("cuda") + ref = _Unfused() + gen = torch.Generator(device="cuda").manual_seed(50) + for req in range(2): + s = int(slots[req]) + slot = slots[req : req + 1] + f64 = kc.F64Verify(wts[0], pools[0], slot, kc.NUM_SPEC) + for rnd in range(SLOT_ROUNDS): + for layer, (p, c) in enumerate(zip(pools, copies)): + before = kc.snapshot(p) + x = _x(gen) + got = _fused(wts[layer], p, x, slot, bufs) + want = ref(wts[layer], c, x, slot) + if layer == 0: + out_f64 = f64(ref.last_rows) + torch.cuda.synchronize() + assert torch.equal(got, want), (req, rnd, layer) + assert kc.same_rows(p, c, [s], POOL_NAMES), (req, rnd, layer) + others = kc.other_slots(p["ssm"].shape[0], [s]) + assert kc.same_rows(p, before, others, POOL_NAMES), (req, rnd, layer) + if layer == 0: + assert kc.rel(got, out_f64) <= kc.TOL_OUT, (req, rnd, kc.rel(got, out_f64)) + accepted = kc.pending_schedule(1, kc.NUM_SPEC, rnd)[0] + _accept([pools[0], copies[0]], s, accepted) + f64.commit([accepted]) + assert bool((bufs.epoch == (2 * SLOT_ROUNDS * LAYERS) % 3).all()), ( + bufs.epoch.unique().tolist() + ) + + +def test_rounds_share_one_buffer_set_and_replay(mgr, slots) -> None: + """LAYERS layers x ROUNDS rounds of one request: one set for every launch, a set per layer, and one round of every + layer captured once and replayed (the inputs rewritten in place, the record written between replays) give the + same outputs and pools bit for bit.""" + rounds = 5 + s = int(slots[2]) + slot = slots[2:3] + pools = _layers(mgr) + with torch.inference_mode(): + for layer, p in enumerate(pools): + kc.fill_pools(p, 1100 + layer) + own_pools = [kc.clone_pools(p) for p in pools] + graph_pools = [kc.clone_pools(p) for p in pools] + for copies in (own_pools, graph_pools): + for c in copies[1:]: + c["pending"] = copies[0]["pending"] + wts = [kc.make_weights(510 + layer) for layer in range(LAYERS)] + gen = torch.Generator(device="cuda").manual_seed(51) + xs = [[_x(gen) for _ in range(LAYERS)] for _ in range(rounds)] + shared = K3KdaBuffers.create("cuda") + own = [K3KdaBuffers.create("cuda") for _ in range(LAYERS)] + got, want = [], [] + for rnd in range(rounds): + for layer in range(LAYERS): + got.append(_fused(wts[layer], pools[layer], xs[rnd][layer], slot, shared).clone()) + want.append( + _fused(wts[layer], own_pools[layer], xs[rnd][layer], slot, own[layer]).clone() + ) + accepted = kc.pending_schedule(1, kc.NUM_SPEC, rnd)[0] + _accept([pools[0], own_pools[0]], s, accepted) + torch.cuda.synchronize() + # Compiles outside capture (already compiled above; kept for running this test alone). + bufs = K3KdaBuffers.create("cuda") + static = [x.clone() for x in xs[0]] + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + graph = torch.cuda.CUDAGraph() + with torch.cuda.stream(stream), torch.cuda.graph(graph, stream=stream): + outs = [ + _fused(wts[li], graph_pools[li], static[li], slot, bufs) for li in range(LAYERS) + ] + torch.cuda.current_stream().wait_stream(stream) + replayed = [] + for rnd in range(rounds): + for layer in range(LAYERS): + static[layer].copy_(xs[rnd][layer]) + graph.replay() + replayed.extend(o.clone() for o in outs) + _accept([graph_pools[0]], s, kc.pending_schedule(1, kc.NUM_SPEC, rnd)[0]) + torch.cuda.synchronize() + assert all(torch.equal(a, b) for a, b in zip(got, want)) + assert all(torch.equal(a, b) for a, b in zip(replayed, want)) + for p, q, r in zip(pools, own_pools, graph_pools): + assert all(torch.equal(p[n], q[n]) and torch.equal(r[n], q[n]) for n in POOL_NAMES) + + +@pytest.mark.parametrize("tokens", [1, 3, 8]) +def test_qkvg_rows_match_the_projection(tokens) -> None: + """The sibling k3_kda_qkvg: the rows decoded from its Lamport buffers are x w^T rounded to bf16 (q, k and f_a the + cluster's four K-partials summed in rank order, v, og and b the bf16 sum of two fp32 K-halves), against float64.""" + kc.load_ops() + with torch.inference_mode(): + wt = kc.make_weights(520) + gen = torch.Generator(device="cuda").manual_seed(52) + x = torch.randn(tokens, kc.K_IN, generator=gen, device="cuda").bfloat16() + bufs = K3KdaBuffers.create("cuda", ctas=CTAS) + k3_kda_qkvg(x, wt["w"], bufs) + rows = kc.decode_rows(bufs.p1, bufs.part, 0, tokens) + want = (x.double() @ wt["w"].double().t())[:, : kc.PROJ - 2] + torch.cuda.synchronize() + err = kc.rel(rows[:, : kc.PROJ - 2], want) + print(f"k3_kda_qkvg rows vs float64, T = {tokens}: rel {err:.3e}") + assert err <= 1e-2, err + assert bool((bufs.epoch == 1).all()), bufs.epoch.unique().tolist() + + +def test_epoch_wrap_on_the_object(mgr, slots) -> None: + """The set's per-CTA indices preset to 2^31 - 2, as after 2^31 launches on a device (one set serves every KDA + layer): four launches give the bits of a fresh set's and leave every index in 0..2.""" + s = int(slots[3]) + slot = slots[3:4] + p = kc.layer_pools(mgr, 0) + with torch.inference_mode(): + kc.fill_pools(p, 1200) + fresh_pools = kc.clone_pools(p) + wt = kc.make_weights(530) + gen = torch.Generator(device="cuda").manual_seed(53) + xs = [_x(gen) for _ in range(4)] + fresh = K3KdaBuffers.create("cuda") + wrapped = K3KdaBuffers.create("cuda") + wrapped.epoch.fill_(2**31 - 2) + got, want = [], [] + for rnd, x in enumerate(xs): + want.append(_fused(wt, fresh_pools, x, slot, fresh).clone()) + torch.cuda.synchronize() # the same pools again next: one launch at a time + got.append(_fused(wt, p, x, slot, wrapped).clone()) + torch.cuda.synchronize() + accepted = kc.pending_schedule(1, kc.NUM_SPEC, rnd)[0] + _accept([p, fresh_pools], s, accepted) + assert all(torch.equal(a, b) for a, b in zip(got, want)) + assert all(torch.equal(p[n], fresh_pools[n]) for n in POOL_NAMES) + assert bool(((wrapped.epoch >= 0) & (wrapped.epoch < 3)).all()), wrapped.epoch.unique().tolist() + + +def test_swapped_rounds_are_silently_wrong(mgr, slots) -> None: + """Negative control: two rounds of one request in swapped order raise nothing and give the second round's + outputs and the slot's state of a different history.""" + s = int(slots[3]) + slot = slots[3:4] + p = kc.layer_pools(mgr, 0) + with torch.inference_mode(): + kc.fill_pools(p, 1300) + swapped = kc.clone_pools(p) + wt = kc.make_weights(540) + gen = torch.Generator(device="cuda").manual_seed(54) + a, b = _x(gen), _x(gen) + bufs = K3KdaBuffers.create("cuda") + _fused(wt, p, a, slot, bufs) + torch.cuda.synchronize() + _accept([p], s, 2) + right = _fused(wt, p, b, slot, bufs).clone() + torch.cuda.synchronize() + _fused(wt, swapped, b, slot, bufs) + torch.cuda.synchronize() + _accept([swapped], s, 2) + wrong = _fused(wt, swapped, a, slot, bufs).clone() + torch.cuda.synchronize() + print(f"swapped rounds: second output rel diff {kc.rel(wrong, right):.3e}, " + f"state rel diff {kc.rel(swapped['ssm'][s], p['ssm'][s]):.3e}") # fmt: skip + assert kc.rel(swapped["ssm"][s], p["ssm"][s]) > kc.TOL_STATE + assert not torch.equal(wrong, right) + + +def test_qkvg_set_is_refused(mgr, slots) -> None: + """A set made for k3_kda_qkvg's 104 CTAs is refused by k3_kda_attn before anything is written.""" + slot = slots[3:4] + p = kc.layer_pools(mgr, 0) + with torch.inference_mode(): + kc.fill_pools(p, 1400) + before = kc.snapshot(p) + wt = kc.make_weights(550) + x = _x(torch.Generator(device="cuda").manual_seed(55)) + with pytest.raises(ValueError, match=f"epoch {CTAS}"): + _fused(wt, p, x, slot, K3KdaBuffers.create("cuda", ctas=CTAS)) + torch.cuda.synchronize() + assert all(torch.equal(p[n], before[n]) for n in POOL_NAMES) diff --git a/tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_k3_kda_decode_attn.py b/tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_k3_kda_decode_attn.py new file mode 100644 index 000000000000..319c8b4918af --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_k3_kda_decode_attn.py @@ -0,0 +1,358 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the ssm/k3_kda_decode_attn catalog entry, on a real cache object and a caller-owned K3KdaBuffers. + +Kimi K3's fused KDA projection + plain decode of R <= 8 requests of one token, at the TP16 rank slice, on the +plain-decode pools of a real ``MambaHybridCacheManagerV2`` (``_kda_cells.build_manager``: conv bf16 +``[slots, 2304, 3]``, fp32 state, strided by the manager's per-slot coalescing) at the slots the manager assigned. + +References, as the op's own test: the projection stream alone (``k3_kda_qkvg`` on a 104-CTA set, its Lamport buffers +decoded into rows), f_b as a bf16 ``F.linear``, then the ``ssm/kda_decode`` entry on a copy of the pools; and a +float64 decode on the same rows. Call sequences: layers x steps on one shared buffer set against a set per layer and +two sets alternating (bit for bit), one step captured and replayed, launches interleaved with ``ssm/k3_kda_attn`` on +one set, the per-CTA index across the int32 wrap; negative controls. + +Every launch here follows a plain copy kernel (its input), as each KDA launch in the model follows other layers' +kernels: the head CTAs read the slots' pools before their grid-dependency wait, so a launch must not directly follow +another launch on the same pools. +""" + +import _kda_cells as kc +import pytest +import torch +import torch.nn.functional as F + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.ssm.k3_kda_attn import ( + k3_kda_attn, + k3_kda_qkvg, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.ssm.k3_kda_buffers import ( + CTAS, + FUSED_CTAS, + K3KdaBuffers, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.ssm.k3_kda_decode_attn import ( + k3_kda_decode_attn, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.ssm.kda_decode import kda_decode + +pytestmark = pytest.mark.skipif(not kc.sm100(), reason="needs SM100 (tcgen05, TMA, clusters)") + +LAYERS = 3 +STEPS = 3 + + +@pytest.fixture(scope="module") +def mgr(): + kc.load_ops() + m = kc.build_manager(LAYERS) + try: + yield m + finally: + m.shutdown() + + +@pytest.fixture(scope="module") +def slots(mgr): + return kc.request_slots(mgr, 8, first_id=200) + + +def _layers(mgr) -> list: + return [kc.layer_pools(mgr, layer) for layer in range(LAYERS)] + + +def _launch(wt, pools, x_src, slots, bufs) -> torch.Tensor: + """One launch, after a plain copy kernel (the input).""" + x = x_src.clone() + return k3_kda_decode_attn( + x, wt["w"], wt["w_fb"], wt["w_q"], wt["w_k"], wt["w_v"], wt["a_log"], wt["dt_bias"], wt["onorm_w"], + pools["conv"], pools["ssm"], slots, bufs, kc.LOWER_BOUND, kc.SCALE, kc.EPS, + ) # fmt: skip + + +class _Rows: + """The projection rows the fused kernel computes: the stream alone on a 104-CTA set, its buffers decoded.""" + + def __init__(self): + self.bufs = K3KdaBuffers.create("cuda", ctas=CTAS) + + def __call__(self, w, x) -> torch.Tensor: + buf = int(self.bufs.epoch[0].item()) + k3_kda_qkvg(x, w, self.bufs) + return kc.decode_rows(self.bufs.p1, self.bufs.part, buf, x.shape[0]) + + +def _native(wt, pools, rows, slots) -> torch.Tensor: + """The model's plain decode on those rows: f_b as a bf16 F.linear, then the ssm/kda_decode entry.""" + num, hk = rows.shape[0], kc.HK + + def heads(cols): + return cols.unflatten(-1, (kc.H, kc.K)).unsqueeze(0) + + g = F.linear(rows[:, 4 * hk : 4 * hk + kc.K], wt["w_fb"]) + out = torch.empty(num, 1, kc.H, kc.V, dtype=torch.bfloat16, device="cuda") + zeros = torch.zeros(hk, dtype=torch.bfloat16, device="cuda") + conv = pools["conv"] + kda_decode( + heads(rows[:, :hk]), heads(rows[:, hk : 2 * hk]), heads(rows[:, 2 * hk : 3 * hk]), + wt["w_t"][0], wt["w_t"][1], wt["w_t"][2], zeros, zeros, zeros, + conv[:, :hk], conv[:, hk : 2 * hk], conv[:, 2 * hk :], wt["a_log"], heads(g), wt["dt_bias"], + rows[:, 4 * hk + kc.K : 4 * hk + kc.K + kc.H].unsqueeze(0), heads(rows[:, 3 * hk : 4 * hk]), wt["onorm_w"], + slots, pools["ssm"], True, True, True, True, kc.LOWER_BOUND, kc.SCALE, kc.EPS, out, + ) # fmt: skip + return out.view(num, kc.H, kc.V) + + +def _x(num: int, gen: torch.Generator) -> torch.Tensor: + return torch.randn(num, kc.K_IN, generator=gen, device="cuda").bfloat16() + + +@pytest.mark.parametrize("num", range(1, 9)) +def test_decode_attn_on_manager_pools(mgr, slots, num) -> None: + """Every layer's step, all layers on one buffer set: outputs against the native path and float64, the state rows + against both (fp32 tolerance), the conv windows against the native path bit for bit; the other slots of the + layer and every slot of the other layers untouched.""" + used = slots[:num] + rows_idx = used.long() + pools = _layers(mgr) + with torch.inference_mode(): + for layer, p in enumerate(pools): + kc.fill_pools(p, 30 * num + layer) + bufs = K3KdaBuffers.create("cuda") + rows_of = _Rows() + gen = torch.Generator(device="cuda").manual_seed(40 + num) + for layer, p in enumerate(pools): + wt = kc.make_weights(100 + layer) + native = kc.clone_pools(p) + ref = kc.F64Decode(wt, p) + before = [kc.snapshot(q) for q in pools] + x = _x(num, gen) + out = _launch(wt, p, x, used, bufs) + rows = rows_of(wt["w"], x) + out_n = _native(wt, native, rows, used) + out_r = ref.step_rows(rows, used) + torch.cuda.synchronize() + assert kc.rel(out, out_n) <= kc.TOL_OUT, (layer, kc.rel(out, out_n)) + assert kc.rel(out, out_r) <= kc.TOL_OUT, (layer, kc.rel(out, out_r)) + assert kc.rel(p["ssm"][rows_idx], ref.state[rows_idx]) <= kc.TOL_STATE, layer + assert kc.rel(p["ssm"][rows_idx], native["ssm"][rows_idx]) <= kc.TOL_STATE, layer + assert torch.equal(p["conv"][rows_idx], native["conv"][rows_idx]), layer + others = kc.other_slots(p["ssm"].shape[0], used.tolist()) + assert kc.same_rows(p, before[layer], others, ("conv", "ssm")), layer + for o, q in enumerate(pools): + if o != layer: + assert all(torch.equal(q[n], before[o][n]) for n in ("conv", "ssm")), (layer, o) + assert bool((bufs.epoch == LAYERS % 3).all()), bufs.epoch.unique().tolist() + + +def _schedule(gen, num) -> list: + return [[_x(num, gen) for _ in range(LAYERS)] for _ in range(STEPS)] + + +def test_steps_share_one_buffer_set(mgr, slots) -> None: + """LAYERS layers x STEPS steps on four requests: one set for every launch (the model's layout), a set per layer, + and two sets alternating between launches give the same outputs and pools bit for bit; on the shared set every + CTA's index has moved once per launch.""" + used = slots[:4] + pools = _layers(mgr) + with torch.inference_mode(): + for layer, p in enumerate(pools): + kc.fill_pools(p, 700 + layer) + own_pools = [kc.clone_pools(p) for p in pools] + alt_pools = [kc.clone_pools(p) for p in pools] + wts = [kc.make_weights(110 + layer) for layer in range(LAYERS)] + xs = _schedule(torch.Generator(device="cuda").manual_seed(41), 4) + shared = K3KdaBuffers.create("cuda") + own = [K3KdaBuffers.create("cuda") for _ in range(LAYERS)] + alt = [K3KdaBuffers.create("cuda") for _ in range(2)] + got, want, both = [], [], [] + for s in range(STEPS): + for layer in range(LAYERS): + got.append(_launch(wts[layer], pools[layer], xs[s][layer], used, shared).clone()) + want.append( + _launch(wts[layer], own_pools[layer], xs[s][layer], used, own[layer]).clone() + ) + pick = alt[(s * LAYERS + layer) % 2] + both.append(_launch(wts[layer], alt_pools[layer], xs[s][layer], used, pick).clone()) + torch.cuda.synchronize() + assert all(torch.equal(a, b) for a, b in zip(got, want)) + assert all(torch.equal(a, b) for a, b in zip(both, want)) + for p, q, r in zip(pools, own_pools, alt_pools): + assert all(torch.equal(p[n], q[n]) and torch.equal(r[n], q[n]) for n in ("conv", "ssm")) + assert bool((shared.epoch == (STEPS * LAYERS) % 3).all()), shared.epoch.unique().tolist() + + +def test_graph_replay(mgr, slots) -> None: + """One step of every layer captured once (the inputs rewritten in place before every replay) on the shared set: + STEPS replays give the bits of the same steps run eagerly on a copy of the pools and a set of their own.""" + used = slots[:3] + pools = _layers(mgr) + with torch.inference_mode(): + for layer, p in enumerate(pools): + kc.fill_pools(p, 800 + layer) + copies = [kc.clone_pools(p) for p in pools] + wts = [kc.make_weights(120 + layer) for layer in range(LAYERS)] + xs = _schedule(torch.Generator(device="cuda").manual_seed(42), 3) + eager_bufs = K3KdaBuffers.create("cuda") + want = [[_launch(wts[li], copies[li], xs[s][li], used, eager_bufs).clone() for li in range(LAYERS)] + for s in range(STEPS)] # fmt: skip + # Compiles outside capture, on scratch pools and a scratch set. + _launch(wts[0], kc.clone_pools(pools[0]), xs[0][0], used, K3KdaBuffers.create("cuda")) + bufs = K3KdaBuffers.create("cuda") + static = [x.clone() for x in xs[0]] + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + graph = torch.cuda.CUDAGraph() + with torch.cuda.stream(stream), torch.cuda.graph(graph, stream=stream): + outs = [_launch(wts[li], pools[li], static[li], used, bufs) for li in range(LAYERS)] + torch.cuda.current_stream().wait_stream(stream) + got = [] + for s in range(STEPS): + for layer in range(LAYERS): + static[layer].copy_(xs[s][layer]) + graph.replay() + got.append([o.clone() for o in outs]) + torch.cuda.synchronize() + assert all(torch.equal(a, b) for s in range(STEPS) for a, b in zip(got[s], want[s])) + for p, c in zip(pools, copies): + assert all(torch.equal(p[n], c[n]) for n in ("conv", "ssm")) + + +@pytest.fixture(scope="module") +def spec_mgr(): + kc.load_ops() + m = kc.build_manager(1, num_spec=kc.NUM_SPEC) + try: + yield m + finally: + m.shutdown() + + +def test_shared_with_k3_kda_attn(mgr, slots, spec_mgr) -> None: + """Plain-decode launches and ``ssm/k3_kda_attn`` verify launches interleaved in stream order on one set (the + model's single set per device) give the bits of the same launches on sets of their own, one at a time. Both run + the same stream role and move every CTA's index once per launch, so either may follow the other.""" + vslot = kc.request_slots(spec_mgr, 1, first_id=300) + dpools = kc.layer_pools(mgr, 0) + vpools = kc.layer_pools(spec_mgr, 0) + with torch.inference_mode(): + kc.fill_pools(dpools, 900) + kc.fill_pools(vpools, 901) + own_d, own_v = kc.clone_pools(dpools), kc.clone_pools(vpools) + wt = kc.make_weights(130) + gen = torch.Generator(device="cuda").manual_seed(43) + calls = [] + for num in (1, 4, 8, 3): + calls.append(("decode", _x(num, gen), slots[:num])) + calls.append(("verify", _x(kc.NT, gen), vslot)) + + def verify(pools, x, slot, bufs): + return k3_kda_attn( + x.clone(), wt["w"], wt["w_fb"], wt["w_q"], wt["w_k"], wt["w_v"], wt["a_log"], wt["dt_bias"], + wt["onorm_w"], pools["cs_q"], pools["cs_k"], pools["cs_v"], pools["ssm"], pools["state_tok"], slot, + pools["pending"], bufs, kc.NUM_SPEC, kc.LOWER_BOUND, kc.SCALE, kc.EPS, + ) # fmt: skip + + own_bufs = {"decode": K3KdaBuffers.create("cuda"), "verify": K3KdaBuffers.create("cuda")} + want = [] + for kind, x, s in calls: + if kind == "decode": + want.append(_launch(wt, own_d, x, s, own_bufs[kind]).clone()) + else: + want.append(verify(own_v, x, s, own_bufs[kind]).clone()) + torch.cuda.synchronize() + shared = K3KdaBuffers.create("cuda") + got = [ + ( + _launch(wt, dpools, x, s, shared) + if kind == "decode" + else verify(vpools, x, s, shared) + ).clone() + for kind, x, s in calls + ] + torch.cuda.synchronize() + assert all(torch.equal(a, b) for a, b in zip(got, want)) + assert all(torch.equal(dpools[n], own_d[n]) for n in ("conv", "ssm")) + assert all( + torch.equal(vpools[n], own_v[n]) for n in ("cs_q", "cs_k", "cs_v", "ssm", "state_tok") + ) + assert bool((shared.epoch == len(calls) % 3).all()), shared.epoch.unique().tolist() + + +def test_epoch_wrap_on_the_object(mgr, slots) -> None: + """The set's per-CTA indices preset to 2^31 - 2, as after 2^31 launches on a device: four launches give the bits of + a fresh set's and leave every index in 0..2 (the kernel keeps the index, not a raw count).""" + used = slots[:2] + p = kc.layer_pools(mgr, 0) + with torch.inference_mode(): + kc.fill_pools(p, 950) + fresh_pools = kc.clone_pools(p) + wt = kc.make_weights(140) + gen = torch.Generator(device="cuda").manual_seed(44) + xs = [_x(2, gen) for _ in range(4)] + fresh = K3KdaBuffers.create("cuda") + want = [] + for x in xs: + want.append(_launch(wt, fresh_pools, x, used, fresh).clone()) + torch.cuda.synchronize() # the same pools again next: one launch at a time + wrapped = K3KdaBuffers.create("cuda") + wrapped.epoch.fill_(2**31 - 2) + got = [] + for x in xs: + got.append(_launch(wt, p, x, used, wrapped).clone()) + torch.cuda.synchronize() + assert all(torch.equal(a, b) for a, b in zip(got, want)) + assert all(torch.equal(p[n], fresh_pools[n]) for n in ("conv", "ssm")) + assert bool(((wrapped.epoch >= 0) & (wrapped.epoch < 3)).all()), wrapped.epoch.unique().tolist() + + +def test_swapped_steps_are_silently_wrong(mgr, slots) -> None: + """Negative control: two steps of one request in swapped order raise nothing and give the second step's output and + the final state of a different history; the pools carry the order, and nothing in a call can detect it.""" + used = slots[:1] + rows_idx = used.long() + p = kc.layer_pools(mgr, 0) + with torch.inference_mode(): + kc.fill_pools(p, 960) + swapped = kc.clone_pools(p) + wt = kc.make_weights(150) + gen = torch.Generator(device="cuda").manual_seed(45) + a, b = _x(1, gen), _x(1, gen) + bufs = K3KdaBuffers.create("cuda") + _launch(wt, p, a, used, bufs) + torch.cuda.synchronize() + right = _launch(wt, p, b, used, bufs).clone() + torch.cuda.synchronize() + _launch(wt, swapped, b, used, bufs) + torch.cuda.synchronize() + wrong = _launch(wt, swapped, a, used, bufs).clone() + torch.cuda.synchronize() + print(f"swapped steps: second output rel diff {kc.rel(wrong, right):.3e}, " + f"state rel diff {kc.rel(swapped['ssm'][rows_idx], p['ssm'][rows_idx]):.3e}") # fmt: skip + assert kc.rel(swapped["ssm"][rows_idx], p["ssm"][rows_idx]) > kc.TOL_STATE + assert not torch.equal(wrong, right) + + +def test_buffer_set_misuse_raises(mgr, slots) -> None: + """A set made for k3_kda_qkvg's 104 CTAs is refused before anything is written; a set cannot be made under + CUDA-graph capture (it allocates); a set for any other CTA count cannot be made.""" + used = slots[:2] + p = kc.layer_pools(mgr, 0) + with torch.inference_mode(): + kc.fill_pools(p, 970) + before = kc.snapshot(p) + wt = kc.make_weights(160) + x = _x(2, torch.Generator(device="cuda").manual_seed(46)) + with pytest.raises(ValueError, match=f"epoch {CTAS}"): + _launch(wt, p, x, used, K3KdaBuffers.create("cuda", ctas=CTAS)) + torch.cuda.synchronize() + assert all(torch.equal(p[n], before[n]) for n in ("conv", "ssm")) + with pytest.raises(ValueError, match="ctas must be"): + K3KdaBuffers.create("cuda", ctas=FUSED_CTAS - 1) + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + graph = torch.cuda.CUDAGraph() + with pytest.raises(RuntimeError, match="before CUDA-graph capture"): + with torch.cuda.stream(stream), torch.cuda.graph(graph, stream=stream): + K3KdaBuffers.create("cuda") + torch.cuda.current_stream().wait_stream(stream) diff --git a/tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_k3_kda_verify.py b/tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_k3_kda_verify.py new file mode 100644 index 000000000000..247ed41d29f5 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_k3_kda_verify.py @@ -0,0 +1,202 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the ssm/k3_kda_verify catalog entry, on a real cache object. + +Kimi K3's KDA speculative verify of R requests of 1 + num_spec tokens, at the TP16 rank slice (6 heads, K = V = 128, +conv width 4), on the verify pools of a real ``MambaHybridCacheManagerV2`` built with MTP-style speculation, the KDA +replay caches and per-token states (``_kda_cells.build_manager``): every layer's conv caches (fp32, dim-contiguous), +its fp32 state (strided by the manager's per-slot coalescing) and per-draft states, and the accepted-draft record all +layers share (``prev_num_accepted_tokens``), at the slots the manager assigned. + +The reference is a float64 verify over each request's committed history (``_kda_cells.F64Verify``): every round's +outputs and the state committed after each golden token against it; the per-draft states and the conv caches are +checked through the next round, which starts from the drafts the sampler accepted. Call sequences: layers x rounds, +a captured round replayed against the same rounds run eagerly, the schedule twice; negative controls. + +Every launch here follows a plain copy kernel (its input), as each KDA launch in the model follows other layers' +kernels: the head CTAs read the slots' pools before their grid-dependency wait. +""" + +import _kda_cells as kc +import pytest +import torch + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.ssm.k3_kda_verify import k3_kda_verify + +pytestmark = pytest.mark.skipif(not kc.sm100(), reason="needs SM100 (tcgen05, TMA, clusters)") + +LAYERS = 2 +ROUNDS = 6 +POOL_NAMES = ("cs_q", "cs_k", "cs_v", "ssm", "state_tok") +# (requests, 1 + num_spec): the model's 7 drafts, and a 3-draft manager. +CELLS = [(1, 8), (4, 8), (8, 8), (4, 4)] + + +@pytest.fixture(scope="module") +def managers(): + kc.load_ops() + built = {} + try: + for num_spec in sorted({steps - 1 for _, steps in CELLS}): + mgr = kc.build_manager(LAYERS, num_spec=num_spec) + built[num_spec] = (mgr, kc.request_slots(mgr, 8, first_id=600 + 10 * num_spec)) + yield built + finally: + for mgr, _ in built.values(): + mgr.shutdown() + + +def _layers(mgr) -> list: + return [kc.layer_pools(mgr, layer) for layer in range(LAYERS)] + + +def _verify(wt, pools, proj_src, slots, num_spec) -> torch.Tensor: + """One launch, after a plain copy kernel (its projection rows).""" + proj = proj_src.clone() + return k3_kda_verify( + proj, wt["w_fb"], wt["w_q"], wt["w_k"], wt["w_v"], wt["a_log"], wt["dt_bias"], wt["onorm_w"], + pools["cs_q"], pools["cs_k"], pools["cs_v"], pools["ssm"], pools["state_tok"], slots, pools["pending"], + num_spec, kc.LOWER_BOUND, kc.SCALE, kc.EPS, + ) # fmt: skip + + +def _proj(num: int, steps: int, gen: torch.Generator) -> torch.Tensor: + return torch.randn(num * steps, kc.PROJ, generator=gen, device="cuda").bfloat16() + + +def _accept(pools_list, slots, accepted) -> None: + for pools in pools_list: + pools["pending"][slots.long()] = torch.tensor(accepted, dtype=torch.int32, device="cuda") + + +@pytest.mark.parametrize("num,steps", CELLS, ids=[f"{r}x{t}" for r, t in CELLS]) +def test_verify_on_manager_pools(managers, num, steps) -> None: + """ROUNDS rounds of every layer with a pending count per request that changes every round: outputs and the + committed states against float64 (fp32 tolerance); the other slots of each layer untouched.""" + num_spec = steps - 1 + mgr, all_slots = managers[num_spec] + slots = all_slots[:num] + pools = _layers(mgr) + with torch.inference_mode(): + for layer, p in enumerate(pools): + kc.fill_pools(p, 2000 + 10 * num + layer) + wts = [kc.make_weights(600 + layer) for layer in range(LAYERS)] + refs = [kc.F64Verify(wts[layer], p, slots, num_spec) for layer, p in enumerate(pools)] + gen = torch.Generator(device="cuda").manual_seed(60 + num) + for rnd in range(ROUNDS): + for layer, p in enumerate(pools): + before = kc.snapshot(p) + proj = _proj(num, steps, gen) + out = _verify(wts[layer], p, proj, slots, num_spec) + want = refs[layer](proj) + torch.cuda.synchronize() + assert kc.rel(out, want) <= kc.TOL_OUT, (rnd, layer, kc.rel(out, want)) + for n, s in enumerate(slots.tolist()): + committed = refs[layer].last[n][0][0] # the state after the golden token + assert kc.rel(p["ssm"][s], committed) <= kc.TOL_STATE, (rnd, layer, n) + others = kc.other_slots(p["ssm"].shape[0], slots.tolist()) + assert kc.same_rows(p, before, others, POOL_NAMES), (rnd, layer) + accepted = kc.pending_schedule(num, num_spec, rnd) + _accept([pools[0]], slots, accepted) + for ref in refs: + ref.commit(accepted) + + +def test_layers_rounds_replay_and_repeat(managers) -> None: + """Four requests, LAYERS layers x ROUNDS rounds: one round of every layer captured once and replayed (the rows + rewritten in place, the record written between replays), and the same schedule run eagerly twice from the same + pools, give the same outputs and pools bit for bit.""" + mgr, all_slots = managers[kc.NUM_SPEC] + slots = all_slots[:4] + steps = kc.NT + pools = _layers(mgr) + with torch.inference_mode(): + for layer, p in enumerate(pools): + kc.fill_pools(p, 2100 + layer) + again = [kc.clone_pools(p) for p in pools] + graph_pools = [kc.clone_pools(p) for p in pools] + for copies in (again, graph_pools): + for c in copies[1:]: + c["pending"] = copies[0]["pending"] + wts = [kc.make_weights(610 + layer) for layer in range(LAYERS)] + gen = torch.Generator(device="cuda").manual_seed(61) + projs = [[_proj(4, steps, gen) for _ in range(LAYERS)] for _ in range(ROUNDS)] + first, second = [], [] + for rnd in range(ROUNDS): + for layer in range(LAYERS): + first.append( + _verify(wts[layer], pools[layer], projs[rnd][layer], slots, kc.NUM_SPEC).clone() + ) + second.append( + _verify(wts[layer], again[layer], projs[rnd][layer], slots, kc.NUM_SPEC).clone() + ) + accepted = kc.pending_schedule(4, kc.NUM_SPEC, rnd) + _accept([pools[0], again[0]], slots, accepted) + torch.cuda.synchronize() + static = [p.clone() for p in projs[0]] + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + graph = torch.cuda.CUDAGraph() + with torch.cuda.stream(stream), torch.cuda.graph(graph, stream=stream): + outs = [ + _verify(wts[li], graph_pools[li], static[li], slots, kc.NUM_SPEC) + for li in range(LAYERS) + ] + torch.cuda.current_stream().wait_stream(stream) + replayed = [] + for rnd in range(ROUNDS): + for layer in range(LAYERS): + static[layer].copy_(projs[rnd][layer]) + graph.replay() + replayed.extend(o.clone() for o in outs) + _accept([graph_pools[0]], slots, kc.pending_schedule(4, kc.NUM_SPEC, rnd)) + torch.cuda.synchronize() + assert all(torch.equal(a, b) for a, b in zip(first, second)), "the schedule twice differs" + assert all(torch.equal(a, b) for a, b in zip(replayed, first)) + for p, q, r in zip(pools, again, graph_pools): + assert all(torch.equal(p[n], q[n]) and torch.equal(r[n], q[n]) for n in POOL_NAMES) + + +def test_swapped_rounds_are_silently_wrong(managers) -> None: + """Negative control: two rounds of one request in swapped order raise nothing and give the second round's outputs + and the slot's state of a different history.""" + mgr, all_slots = managers[kc.NUM_SPEC] + slots = all_slots[:1] + s = int(slots[0]) + p = kc.layer_pools(mgr, 0) + with torch.inference_mode(): + kc.fill_pools(p, 2200) + swapped = kc.clone_pools(p) + wt = kc.make_weights(620) + gen = torch.Generator(device="cuda").manual_seed(62) + a, b = _proj(1, kc.NT, gen), _proj(1, kc.NT, gen) + _verify(wt, p, a, slots, kc.NUM_SPEC) + _accept([p], slots, [3]) + right = _verify(wt, p, b, slots, kc.NUM_SPEC).clone() + _verify(wt, swapped, b, slots, kc.NUM_SPEC) + _accept([swapped], slots, [3]) + wrong = _verify(wt, swapped, a, slots, kc.NUM_SPEC).clone() + torch.cuda.synchronize() + print(f"swapped rounds: second output rel diff {kc.rel(wrong, right):.3e}, " + f"state rel diff {kc.rel(swapped['ssm'][s], p['ssm'][s]):.3e}") # fmt: skip + assert kc.rel(swapped["ssm"][s], p["ssm"][s]) > kc.TOL_STATE + assert not torch.equal(wrong, right) + + +def test_rejects_out_of_contract(managers) -> None: + """Per-draft states for another draft count, and an int64 record, are refused before anything is written.""" + mgr, all_slots = managers[kc.NUM_SPEC] + slots = all_slots[:2] + p = kc.layer_pools(mgr, 0) + with torch.inference_mode(): + kc.fill_pools(p, 2300) + before = kc.snapshot(p) + wt = kc.make_weights(630) + proj = _proj(2, kc.NT, torch.Generator(device="cuda").manual_seed(63)) + short = dict(p, state_tok=p["state_tok"][:, : kc.NUM_SPEC - 1].contiguous()) + with pytest.raises(ValueError, match="unsupported call"): + _verify(wt, short, proj, slots, kc.NUM_SPEC) + with pytest.raises(ValueError, match="unsupported call"): + _verify(wt, dict(p, pending=p["pending"].long()), proj, slots, kc.NUM_SPEC) + torch.cuda.synchronize() + assert all(torch.equal(p[n], before[n]) for n in POOL_NAMES) diff --git a/tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_kda_decode.py b/tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_kda_decode.py new file mode 100644 index 000000000000..07cc6952b52f --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_kda_decode.py @@ -0,0 +1,194 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the ssm/kda_decode catalog entry, on a real cache object. + +Kimi K3's plain decode shape (6 heads, K = V = 128, conv width 4, the lower-bound gate, the output norm, beta sigmoid +in the kernel) on the pools of a real ``MambaHybridCacheManagerV2`` (``_kda_cells.build_manager``): its per-layer +conv views (``[q | k | v]`` sections of one bf16 ``[slots, 2304, 3]`` pool) and fp32 state views, strided by the +manager's per-slot coalescing, at the slots the manager assigned. The reference is a float64 torch decode +(``_kda_cells.F64Decode``). B runs 1..8: on sm_100 B <= 5 selects the four-CTA cluster kernel and 6 <= B <= 24 the +legacy compact-heads kernel (the one #19830 fixes). +""" + +import _kda_cells as kc +import pytest +import torch + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.ssm.kda_decode import kda_decode + +assert torch.cuda.is_available(), "kda_decode requires a CUDA device" + +LAYERS = 3 +STEPS = 3 + + +@pytest.fixture(scope="module") +def mgr(): + kc.load_ops() + m = kc.build_manager(LAYERS) + try: + yield m + finally: + m.shutdown() + + +@pytest.fixture(scope="module") +def slots(mgr): + return kc.request_slots(mgr, 8, first_id=100) + + +def _heads(cols: torch.Tensor) -> torch.Tensor: + return cols.unflatten(-1, (kc.H, kc.K)).unsqueeze(0) + + +def _inputs(num: int, gen: torch.Generator) -> dict: + def rnd(*s): + return torch.randn(*s, generator=gen, device="cuda").bfloat16() + + return { + "raw": rnd(num, 3 * kc.HK), + "g": rnd(num, kc.HK), + "beta": rnd(num, kc.H), + "gate": rnd(num, kc.HK), + } + + +def _call(wt, pools, inp, slots, out=None, use_lower_bound=True) -> torch.Tensor: + """The entry on one layer's pools, as the model calls it: packed conv sections updated in place, indexed state.""" + num = slots.numel() + raw, conv = inp["raw"], pools["conv"] + zeros = torch.zeros(kc.HK, dtype=torch.bfloat16, device="cuda") + if out is None: + out = torch.empty(num, 1, kc.H, kc.V, dtype=torch.bfloat16, device="cuda") + kda_decode( + _heads(raw[:, : kc.HK]), _heads(raw[:, kc.HK : 2 * kc.HK]), _heads(raw[:, 2 * kc.HK :]), + wt["w_t"][0], wt["w_t"][1], wt["w_t"][2], zeros, zeros, zeros, + conv[:, : kc.HK], conv[:, kc.HK : 2 * kc.HK], conv[:, 2 * kc.HK :], + wt["a_log"], _heads(inp["g"]), wt["dt_bias"], inp["beta"].unsqueeze(0), _heads(inp["gate"]), + wt["onorm_w"], slots, pools["ssm"], True, True, use_lower_bound, True, kc.LOWER_BOUND, kc.SCALE, kc.EPS, out, + ) # fmt: skip + return out.view(num, kc.H, kc.V) + + +def _layers(mgr) -> list: + return [kc.layer_pools(mgr, layer) for layer in range(LAYERS)] + + +@pytest.mark.parametrize("num", range(1, 9)) +def test_decode_on_manager_pools(mgr, slots, num) -> None: + """Every layer's step against float64: outputs and the state rows at fp32 tolerance, the conv windows bit for + bit (raw bf16 inputs), the other slots of the layer and every slot of the other layers untouched.""" + used = slots[:num] + pools = _layers(mgr) + with torch.inference_mode(): + for layer, p in enumerate(pools): + kc.fill_pools(p, 10 * num + layer) + gen = torch.Generator(device="cuda").manual_seed(num) + for layer, p in enumerate(pools): + wt = kc.make_weights(100 + layer) + ref = kc.F64Decode(wt, p) + before = [kc.snapshot(q) for q in pools] + inp = _inputs(num, gen) + out = _call(wt, p, inp, used) + want = ref.step(inp["raw"], inp["g"], inp["beta"], inp["gate"], used) + torch.cuda.synchronize() + rows = used.long() + assert kc.rel(out, want) <= kc.TOL_OUT, (layer, kc.rel(out, want)) + assert kc.rel(p["ssm"][rows], ref.state[rows]) <= kc.TOL_STATE, layer + assert torch.equal(p["conv"][rows], ref.conv[rows].bfloat16()), layer + others = kc.other_slots(p["ssm"].shape[0], used.tolist()) + assert kc.same_rows(p, before[layer], others, ("conv", "ssm")), layer + for o, q in enumerate(pools): + if o != layer: + assert all(torch.equal(q[n], before[o][n]) for n in ("conv", "ssm")), (layer, o) + + +def test_steps_on_layers_and_graph_replay(mgr, slots) -> None: + """LAYERS layers x STEPS decode steps on four requests: the steps captured once in a CUDA graph (inputs and slots + rewritten in place before each replay) give the bits of the same steps run eagerly on a copy of the pools.""" + used = slots[:4] + pools = _layers(mgr) + with torch.inference_mode(): + for layer, p in enumerate(pools): + kc.fill_pools(p, 500 + layer) + copies = [kc.clone_pools(p) for p in pools] + wts = [kc.make_weights(200 + layer) for layer in range(LAYERS)] + gen = torch.Generator(device="cuda").manual_seed(7) + steps = [[_inputs(4, gen) for _ in range(LAYERS)] for _ in range(STEPS)] + want = [ + [_call(wts[li], copies[li], steps[s][li], used).clone() for li in range(LAYERS)] + for s in range(STEPS) + ] + static = [{n: t.clone() for n, t in steps[0][li].items()} for li in range(LAYERS)] + s_in = used.clone() + outs = [ + torch.empty(4, 1, kc.H, kc.V, dtype=torch.bfloat16, device="cuda") + for _ in range(LAYERS) + ] + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + graph = torch.cuda.CUDAGraph() + with torch.cuda.stream(stream), torch.cuda.graph(graph, stream=stream): + for li in range(LAYERS): + _call(wts[li], pools[li], static[li], s_in, outs[li]) + torch.cuda.current_stream().wait_stream(stream) + got = [] + for s in range(STEPS): + for li in range(LAYERS): + for n, t in static[li].items(): + t.copy_(steps[s][li][n]) + graph.replay() + got.append([o.view(4, kc.H, kc.V).clone() for o in outs]) + torch.cuda.synchronize() + assert all(torch.equal(a, b) for s in range(STEPS) for a, b in zip(got[s], want[s])) + for p, c in zip(pools, copies): + assert torch.equal(p["ssm"], c["ssm"]) and torch.equal(p["conv"], c["conv"]) + + +def test_swapped_steps_are_silently_wrong(mgr, slots) -> None: + """Negative control: two decode steps of one request applied in swapped order raise nothing and give the second + step's output and the final state of a different history (the pools carry the order; nothing in the call can + detect it).""" + used = slots[:1] + p = kc.layer_pools(mgr, 0) + with torch.inference_mode(): + kc.fill_pools(p, 900) + start = kc.clone_pools(p) + wt = kc.make_weights(300) + gen = torch.Generator(device="cuda").manual_seed(11) + a, b = _inputs(1, gen), _inputs(1, gen) + _call(wt, p, a, used) + right = _call(wt, p, b, used).clone() + right_state = p["ssm"][used.long()].clone() + swapped = kc.clone_pools(start) + _call(wt, swapped, b, used) + wrong = _call(wt, swapped, a, used).clone() + torch.cuda.synchronize() + rows = used.long() + print(f"swapped steps: second output rel diff {kc.rel(wrong, right):.3e}, " + f"state rel diff {kc.rel(swapped['ssm'][rows], right_state):.3e}") # fmt: skip + assert kc.rel(swapped["ssm"][rows], right_state) > kc.TOL_STATE + assert not torch.equal(wrong, right) + + +def test_rejects_out_of_contract(mgr, slots) -> None: + """A state view that is not 16-byte aligned, int64 slot indices, and (on sm_100 / sm_103, where small batches + run the optimized kernels) the gate without its lower bound: each raises before the pools are written.""" + used = slots[:2] + p = kc.layer_pools(mgr, 0) + with torch.inference_mode(): + kc.fill_pools(p, 950) + before = kc.snapshot(p) + wt = kc.make_weights(400) + inp = _inputs(2, torch.Generator(device="cuda").manual_seed(13)) + ssm = p["ssm"] + shifted = dict(p, ssm=ssm.as_strided(ssm.shape, ssm.stride(), ssm.storage_offset() + 1)) + with pytest.raises(RuntimeError, match="16B-aligned"): + _call(wt, shifted, inp, used) + with pytest.raises(RuntimeError, match="int32"): + _call(wt, p, inp, used.long()) + if torch.cuda.get_device_capability() in ((10, 0), (10, 3)): + with pytest.raises(RuntimeError, match="lower-bound gate"): + _call(wt, p, inp, used, use_lower_bound=False) + torch.cuda.synchronize() + assert all(torch.equal(p[n], before[n]) for n in ("conv", "ssm")) From 373b10b4e23550d9058125fd5f12131bdc592852 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 18:16:37 -0700 Subject: [PATCH 022/161] [None][chore] Kimi K3 KDA / MLA decode: format with main's hooks ruff-format and ruff's import sort on the two files the K3 stack wrote in its 120-column style (k3_kda_attn_kernel.py, k3_mla/op.py). The AST of every Python file of this PR is unchanged; the full pre-commit passes on all of them. Signed-off-by: Vasanth Sabavat --- .../k3_kda_attn/k3_kda_attn_kernel.py | 15 ++++++++++++--- tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/op.py | 3 ++- 2 files changed, 14 insertions(+), 4 deletions(-) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_attn/k3_kda_attn_kernel.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_attn/k3_kda_attn_kernel.py index c61c40dd9760..3ff125d1c3b4 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_attn/k3_kda_attn_kernel.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_kda_attn/k3_kda_attn_kernel.py @@ -57,7 +57,13 @@ except ImportError: # older DSL layout from cutlass.utils import SmemAllocator -from ..k3_kda_verify.k3_kda_verify_kernel import _bf16, _butterfly, _st_async_f32, _store8, _test_wait_cluster +from ..k3_kda_verify.k3_kda_verify_kernel import ( + _bf16, + _butterfly, + _st_async_f32, + _store8, + _test_wait_cluster, +) K_IN = 7168 HK = 768 # 6 local heads x 128 @@ -952,7 +958,9 @@ def _head_role( tmem_acc = cutlass.inttoptr(tmem_holder.load(), 6, cutlass.Int32) while not cute.arch.mbarrier_test_wait(w_full.data_ptr(), 0): pass - prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) # f_a came from the other threads' smem stores + prims.tcgen05_fence( + prims.Tcgen05Fence.AFTER_THREAD_SYNC + ) # f_a came from the other threads' smem stores for kb in cutlass.range_constexpr(2 * (BOX_K // MMA_K)): box = kb // (BOX_K // MMA_K) within = kb % (BOX_K // MMA_K) @@ -1218,7 +1226,8 @@ def _head_role( xq_last = _rows_out([r_st.load(idx=R_Q + r) for r in range(REC_ROWS)], lane) if (lane & cutlass.Int32(7)) == cutlass.Int32(0): s_o.store( - _bf16(xq_last), idx=(NT - 1) * V_CTA + warp * REC_ROWS + (lane >> 4) * 2 + ((lane >> 3) & 1) + _bf16(xq_last), + idx=(NT - 1) * V_CTA + warp * REC_ROWS + (lane >> 4) * 2 + ((lane >> 3) & 1), ) # ---- Phase 3 (the output gate), then the gated RMSNorm over V: every CTA sends its two 16-row sums of squares diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/op.py index 96cb022aa777..0c3269ba5118 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_mla/op.py @@ -478,7 +478,8 @@ def _launch_attn( clusters = num_requests * groups no_cluster = ( clusters > kernel.CLUSTER_WAVE - and clusters * kernel.CLUSTER <= torch.cuda.get_device_properties(q.device).multi_processor_count + and clusters * kernel.CLUSTER + <= torch.cuda.get_device_properties(q.device).multi_processor_count ) if out is None: out = torch.empty(num_tokens, total_heads * width, dtype=torch.bfloat16, device=q.device) From ce2b3c32fe928dd3feb62ff9fe72dc0ee0e4a10d Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 18:34:33 -0700 Subject: [PATCH 023/161] [None][fix] modeling_v2 ssm/kda_decode: the op supports only the production flags kda_decode checks apply_onorm, use_lower_bound and apply_beta_sigmoid before anything else and refuses any of them off, on every architecture and batch size ("KDA decode only supports apply_onorm=true, use_lower_bound=true, and apply_beta_sigmoid=true"). The contract described the off branches and a kernel-level refusal on sm_100 / sm_103 only; its semantics now show the one supported combination, and the test expects the op's message on every architecture. Signed-off-by: Vasanth Sabavat --- .../modeling_v2/catalog/ssm/kda_decode.md | 17 ++++++++++------- .../ssm/test_modeling_v2_kda_decode.py | 12 +++++++----- 2 files changed, 17 insertions(+), 12 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/kda_decode.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/kda_decode.md index 07bca5fd177f..f1c94f995764 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/kda_decode.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/kda_decode.md @@ -18,13 +18,15 @@ For request b (slot `s = ssm_state_indices[b]`), head h, with `x_q`, `x_k`, `x_v ``` q, k, v = SiLU(conv4(window[s], new raw) + bias) # taps w_*_t[j], oldest input first; bias_* added before SiLU q = q / sqrt(sum(q^2) + 1e-6) * scale; k = k / sqrt(sum(k^2) + 1e-6) -beta = sigmoid(beta) if apply_beta_sigmoid else beta -decay = exp(lower_bound * sigmoid(exp(a_log) * (g + dt_bias))) if use_lower_bound - exp(-exp(a_log) * softplus(g + dt_bias)) otherwise +beta = sigmoid(beta) +decay = exp(lower_bound * sigmoid(exp(a_log) * (g + dt_bias))) S = state[s] * decay (per key); S += beta (v - S k) k^T; state[s] = S; o = S q -output = o * rsqrt(mean(o^2) + onorm_eps) * onorm_weight * sigmoid(onorm_g) if apply_onorm, else o +output = o * rsqrt(mean(o^2) + onorm_eps) * onorm_weight * sigmoid(onorm_g) ``` +`apply_onorm`, `use_lower_bound` and `apply_beta_sigmoid` must all be True: the op supports only that combination +(*Preconditions*), which is the one above and Kimi K3's. + With `update_conv_cache` the conv windows at the slot shift by one (the last two raw inputs, then the new one). The float64 reference of the test is this arithmetic; the op matches it within fp32 tolerance (outputs 2e-2 relative, bf16; state rows 1e-3), and the conv windows bit for bit (raw bf16 inputs). @@ -122,9 +124,10 @@ TBD(tray: test_modeling_v2_kda_decode.py::test_swapped_steps_are_silently_wrong, - Every check in the table is the op's (`TORCH_CHECK`): a violation raises `RuntimeError` before the launch, with the pools unchanged. Among them: the state base must be 16-byte aligned and its slot stride a multiple of 4 floats (the kernels move state with 16-byte accesses at `slot * stride(0)`); `ssm_state_indices` must be int32. -- On sm_100 and sm_103, the optimized kernels the dispatcher picks for small and large workloads require - `apply_onorm`, `use_lower_bound` and `apply_beta_sigmoid`; with any of them off such a call raises `RuntimeError` - ("Optimized KDA decode requires ...") before the launch. Certified with the lower bound off at B = 2, H = 6. +- `apply_onorm`, `use_lower_bound` and `apply_beta_sigmoid` must all be True, on every architecture and batch size: + the op checks them first and otherwise raises `RuntimeError` ("KDA decode only supports apply_onorm=true, + use_lower_bound=true, and apply_beta_sigmoid=true") before the launch. Certified with the lower bound off at B = 2, + H = 6, the pools unchanged. - The conv inputs, gates and outputs are bf16 and the state fp32; K = V = 128 and conv width 4 only. ## Notes diff --git a/tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_kda_decode.py b/tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_kda_decode.py index 07cc6952b52f..f767a1c4e2c9 100644 --- a/tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_kda_decode.py +++ b/tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_kda_decode.py @@ -172,8 +172,9 @@ def test_swapped_steps_are_silently_wrong(mgr, slots) -> None: def test_rejects_out_of_contract(mgr, slots) -> None: - """A state view that is not 16-byte aligned, int64 slot indices, and (on sm_100 / sm_103, where small batches - run the optimized kernels) the gate without its lower bound: each raises before the pools are written.""" + """A state view that is not 16-byte aligned, int64 slot indices, and the gate without its lower bound (the op + supports only apply_onorm, use_lower_bound and apply_beta_sigmoid all on): each raises before the pools are + written.""" used = slots[:2] p = kc.layer_pools(mgr, 0) with torch.inference_mode(): @@ -187,8 +188,9 @@ def test_rejects_out_of_contract(mgr, slots) -> None: _call(wt, shifted, inp, used) with pytest.raises(RuntimeError, match="int32"): _call(wt, p, inp, used.long()) - if torch.cuda.get_device_capability() in ((10, 0), (10, 3)): - with pytest.raises(RuntimeError, match="lower-bound gate"): - _call(wt, p, inp, used, use_lower_bound=False) + with pytest.raises( + RuntimeError, match="only supports apply_onorm=true, use_lower_bound=true" + ): + _call(wt, p, inp, used, use_lower_bound=False) torch.cuda.synchronize() assert all(torch.equal(p[n], before[n]) for n in ("conv", "ssm")) From 95fbc35740559029e2ad8cd1993f800bb30c0d1b Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 18:38:56 -0700 Subject: [PATCH 024/161] [None][test] modeling_v2 K3 KDA / MLA entries: sm_100 only, measured negative controls The entries' tests ran on SM 10.x; their receipts are sm_100's, so they now skip on any other architecture (a missing receipt reads as unknown). The KDA contracts state what their negative controls measured on sm_100: two steps or rounds of one request in swapped order leave the second output and the slot's state off by 0.8-1.4 of their largest magnitude, with nothing raised; and k3_kda_qkvg's rows are within 2.5e-3 of a float64 projection. Signed-off-by: Vasanth Sabavat --- .../_experimental/modeling_v2/catalog/ssm/k3_kda_attn.md | 6 +++--- .../modeling_v2/catalog/ssm/k3_kda_decode_attn.md | 4 ++-- .../_experimental/modeling_v2/catalog/ssm/k3_kda_verify.md | 4 ++-- .../_experimental/modeling_v2/catalog/ssm/kda_decode.md | 4 ++-- .../attention/test_modeling_v2_k3_mla_attn_vb_out.py | 5 +++-- .../modeling_v2/attention/test_modeling_v2_k3_mla_qkv.py | 5 +++-- tests/unittest/_torch/modeling_v2/ssm/_kda_cells.py | 3 ++- .../_torch/modeling_v2/ssm/test_modeling_v2_kda_decode.py | 1 + 8 files changed, 18 insertions(+), 14 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_attn.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_attn.md index 505145f1ad39..d937baa9da10 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_attn.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_attn.md @@ -41,7 +41,7 @@ same pools (certified, every pool word). Repeated runs are bit-identical. **`k3_kda_qkvg`** runs phase 1 alone on `x` bf16 `[T <= 8, 7168]` and publishes the rows: buffer `e = buffers.epoch[0]` (before the call) holds q, k and f_a as bf16 bits in `p1`, and v, og and b as the fp32 bits of the two K-half partials in `part`. The projection is bf16 of their sum, the consumer's job. Measured against a float64 -`x @ w^T`: TBD(tray: test_modeling_v2_k3_kda_attn.py::test_qkvg_rows_match_the_projection). +`x @ w^T`: the rows are within 2.5e-3 of the projection's largest magnitude at T = 1, and 2.4e-3 at T = 3 and 8. ## Signature @@ -160,8 +160,8 @@ committed history; then layers x rounds on one shared set against a set per laye once and replayed with rewritten inputs and records. **What a wrong order does.** Two rounds of one request in swapped order (the test's negative control): nothing raises, -and the second round's output and the slot's state are those of a different history. Measured on sm_100: -TBD(tray: test_modeling_v2_k3_kda_attn.py::test_swapped_rounds_are_silently_wrong, its printed rel diffs). +and the second round's output and the slot's state are those of a different history. Measured on sm_100: the second round's output is off by 1.00 and +the state by 1.01 (the largest absolute difference over the in-order result's largest magnitude). ## Metadata consumed diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_decode_attn.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_decode_attn.md index 11cc7085fa88..19e3671134e4 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_decode_attn.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_decode_attn.md @@ -99,8 +99,8 @@ runs layers x steps of a real manager, and captures one step of every layer once inputs: the replays give the bits of the same steps run eagerly on a copy of the pools. **What a wrong order does.** Two steps of one request in swapped order (the negative control): nothing raises, and the -second step's output and the final state are those of a different history. Measured on sm_100: -TBD(tray: test_modeling_v2_k3_kda_decode_attn.py::test_swapped_steps_are_silently_wrong, its printed rel diffs). +second step's output and the final state are those of a different history. Measured on sm_100: the second step's output is off by 0.81 and the state by +0.89 (the largest absolute difference over the in-order result's largest magnitude). ## Metadata consumed diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_verify.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_verify.md index 7a819e96b565..568124e74aaf 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_verify.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_verify.md @@ -114,8 +114,8 @@ captures a round of every layer once, replaying it with rewritten rows and recor same rounds run eagerly, and the schedule run twice gives the same bits. **What a wrong order does.** Two rounds of one request in swapped order (the negative control): nothing raises, and -the second round's outputs and the slot's state are those of a different history. Measured on sm_100: -TBD(tray: test_modeling_v2_k3_kda_verify.py::test_swapped_rounds_are_silently_wrong, its printed rel diffs). +the second round's outputs and the slot's state are those of a different history. Measured on sm_100: the second round's outputs are off by 0.98 and the +state by 0.98 (the largest absolute difference over the in-order result's largest magnitude). ## Metadata consumed diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/kda_decode.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/kda_decode.md index f1c94f995764..479e6f1cc469 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/kda_decode.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/kda_decode.md @@ -112,8 +112,8 @@ on a real manager and captures a step of every layer once, replaying it with rew bits of the same steps run eagerly on a copy of the pools. **What a wrong order does.** Two steps of one request in swapped order (the negative control): nothing raises, and the -second step's output and the final state are those of a different history. Measured at Kimi K3's shape on sm_100: -TBD(tray: test_modeling_v2_kda_decode.py::test_swapped_steps_are_silently_wrong, its printed rel diffs). +second step's output and the final state are those of a different history. Measured at Kimi K3's shape on sm_100: the second step's output is off by 1.38 and the state +by 1.27 (the largest absolute difference over the in-order result's largest magnitude). ## Metadata consumed diff --git a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_mla_attn_vb_out.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_mla_attn_vb_out.py index c966efaa08ae..065b142ed32c 100644 --- a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_mla_attn_vb_out.py +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_mla_attn_vb_out.py @@ -53,10 +53,11 @@ def _is_sm100() -> bool: if not torch.cuda.is_available(): return False major, minor = torch.cuda.get_device_capability() - return major * 10 + minor in (100, 103) + # sm_100 exactly: the receipts' architecture (a missing receipt reads as unknown). + return (major, minor) == (10, 0) -pytestmark = pytest.mark.skipif(not _is_sm100(), reason="needs an SM100 / SM103 GPU") +pytestmark = pytest.mark.skipif(not _is_sm100(), reason="certified on SM100 only") H, NOPE, LAT, PE, QL, V = 6, 128, 512, 64, 1536, 128 DQK, PAGE = LAT + PE, 64 diff --git a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_mla_qkv.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_mla_qkv.py index d0f9ab96b4ce..feadc66a35e6 100644 --- a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_mla_qkv.py +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_mla_qkv.py @@ -45,10 +45,11 @@ def _is_sm100() -> bool: if not torch.cuda.is_available(): return False major, minor = torch.cuda.get_device_capability() - return major * 10 + minor in (100, 103) + # sm_100 exactly: the receipts' architecture (a missing receipt reads as unknown). + return (major, minor) == (10, 0) -pytestmark = pytest.mark.skipif(not _is_sm100(), reason="needs an SM100 / SM103 GPU") +pytestmark = pytest.mark.skipif(not _is_sm100(), reason="certified on SM100 only") H, NOPE, PE, QK, LAT, QL, V = 6, 128, 64, 192, 512, 1536, 128 DQK, PAGE = LAT + PE, 64 diff --git a/tests/unittest/_torch/modeling_v2/ssm/_kda_cells.py b/tests/unittest/_torch/modeling_v2/ssm/_kda_cells.py index 14762d56fe39..9c5d0e705b68 100644 --- a/tests/unittest/_torch/modeling_v2/ssm/_kda_cells.py +++ b/tests/unittest/_torch/modeling_v2/ssm/_kda_cells.py @@ -32,7 +32,8 @@ def sm100() -> bool: - return torch.cuda.is_available() and torch.cuda.get_device_capability()[0] == 10 + """sm_100 exactly: the architecture these entries are certified on (a missing receipt reads as unknown).""" + return torch.cuda.is_available() and torch.cuda.get_device_capability() == (10, 0) def load_ops(): diff --git a/tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_kda_decode.py b/tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_kda_decode.py index f767a1c4e2c9..035eb7ee3869 100644 --- a/tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_kda_decode.py +++ b/tests/unittest/_torch/modeling_v2/ssm/test_modeling_v2_kda_decode.py @@ -17,6 +17,7 @@ from tensorrt_llm._torch._experimental.modeling_v2.catalog.ssm.kda_decode import kda_decode assert torch.cuda.is_available(), "kda_decode requires a CUDA device" +pytestmark = pytest.mark.skipif(not kc.sm100(), reason="certified on SM100 only") LAYERS = 3 STEPS = 3 From 8ce079253b0fd7f8da528a8ece9eaf6e694f8704 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 18:46:30 -0700 Subject: [PATCH 025/161] [None][test] Kimi K3 KDA / MLA decode: sm_100 receipts and l0_b200 entries Receipts of the six entries on GB200 (sm_100), from one run of every test file of this PR at its previous commit: kda_decode 11, k3_kda_attn 8, k3_kda_decode_attn 14, k3_kda_verify 7, k3_mla_qkv 8, k3_mla_attn_vb_out 9 tests passed. The op tests and the entry tests join l0_b200's pre-merge list; each file ran in under a minute. Signed-off-by: Vasanth Sabavat --- .../catalog/attention/k3_mla_attn_vb_out.md | 3 ++- .../modeling_v2/catalog/attention/k3_mla_qkv.md | 3 ++- .../modeling_v2/catalog/ssm/k3_kda_attn.md | 3 ++- .../modeling_v2/catalog/ssm/k3_kda_decode_attn.md | 3 ++- .../modeling_v2/catalog/ssm/k3_kda_verify.md | 3 ++- .../modeling_v2/catalog/ssm/kda_decode.md | 3 ++- tests/integration/test_lists/test-db/l0_b200.yml | 11 +++++++++++ 7 files changed, 23 insertions(+), 6 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_attn_vb_out.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_attn_vb_out.md index d6cf18aa0050..2b10e316f27b 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_attn_vb_out.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_attn_vb_out.md @@ -1,5 +1,6 @@ --- -receipts: {} +receipts: + sm_100: {status: passed, tests: 9} --- # k3_mla_attn_vb_out diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_qkv.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_qkv.md index 114d1d4d8a86..fb766b8d41e1 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_qkv.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_mla_qkv.md @@ -1,5 +1,6 @@ --- -receipts: {} +receipts: + sm_100: {status: passed, tests: 8} --- # k3_mla_qkv diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_attn.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_attn.md index d937baa9da10..67648d10ae0e 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_attn.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_attn.md @@ -1,5 +1,6 @@ --- -receipts: {} +receipts: + sm_100: {status: passed, tests: 8} --- # k3_kda_attn diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_decode_attn.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_decode_attn.md index 19e3671134e4..a65d9032c18d 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_decode_attn.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_decode_attn.md @@ -1,5 +1,6 @@ --- -receipts: {} +receipts: + sm_100: {status: passed, tests: 14} --- # k3_kda_decode_attn diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_verify.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_verify.md index 568124e74aaf..a65197c0634f 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_verify.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/k3_kda_verify.md @@ -1,5 +1,6 @@ --- -receipts: {} +receipts: + sm_100: {status: passed, tests: 7} --- # k3_kda_verify diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/kda_decode.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/kda_decode.md index 479e6f1cc469..8974bb788a54 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/kda_decode.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/ssm/kda_decode.md @@ -1,5 +1,6 @@ --- -receipts: {} +receipts: + sm_100: {status: passed, tests: 11} --- # kda_decode diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index 2ca039299f9c..2996d86ea3cb 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -375,6 +375,17 @@ l0_b200: - examples/visual_gen/test_minimax_h3_e2e.py::test_minimax_h3_public_visual_gen_api_smoke TIMEOUT (30) - examples/visual_gen/test_minimax_h3_e2e.py::test_minimax_h3_diffusers_lpips_and_audio_reference[t2va] TIMEOUT (30) - examples/visual_gen/test_minimax_h3_e2e.py::test_minimax_h3_diffusers_lpips_and_audio_reference[fl2va] TIMEOUT (30) + # ------------- Kimi K3 KDA / MLA decode kernels (sm_100) and their modeling_v2 entries --------------- + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_kda_attn.py + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_kda_decode_attn.py + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_kda_pools_past_2g.py + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_kda_verify.py + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_attn.py + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_q.py + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_decode_view.py + - unittest/_torch/modeling_v2/ssm + - unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_mla_qkv.py + - unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_mla_attn_vb_out.py # ------------- Host perf module regression tests (6 representative scenarios) --------------- - perf/host_perf/test_module_scheduler.py::test_scheduler_production[production_gen_only_bs8] - perf/host_perf/test_module_scheduler.py::test_scheduler_production[production_mixed_32gen_4ctx] From 0ca3a9f1b8baf098366d0414497bc07b31182490 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 18:46:32 -0700 Subject: [PATCH 026/161] [None][feat] Kimi K3 decode kernels (CuTe DSL): CTM GEMVs, decode GEMV, head GEMV, embedding The single-GPU decode kernels of the Kimi K3 port, with their op tests: - k3_ctm_gemv: trtllm::k3_ctm_gemv, _swiglu, _long, _wide (bf16 GEMVs for up to 8 tokens, 64 for _wide, on tcgen05) and trtllm::k3_situ_mul, the dense MLP's SiTU-and-mul; - k3_decode_gemv: trtllm::k3_decode_gemv and its row-parallel MoE tail; - k3_head_gemv: trtllm::k3_head_gemv, the vocabulary-shard head GEMV on a persistent stream-K kernel of min(SMs, tiles x k-tiles) CTAs; - k3_embed: trtllm::k3_embed and k3_embed_norm (the decode step's embedding and layer 0's input RMSNorm in one launch). Two source lints read these kernels: test_k3_cluster_waits (a wait on a mailbox that other CTAs complete with st.async acquires at cluster scope) and test_k3_tcgen05_fences (the tcgen05 fences around thread syncs). Every kernel and test file is as on the K3 stack's tip; the decode GEMV test checks its tail against a torch reference. Signed-off-by: Vasanth Sabavat --- .../cute_dsl_kernels/k3_ctm_gemv/__init__.py | 19 + .../k3_ctm_gemv/k3_ctm_gemv_kernel.py | 1681 +++++++++++++++++ .../_torch/cute_dsl_kernels/k3_ctm_gemv/op.py | 544 ++++++ .../k3_decode_gemv/__init__.py | 19 + .../k3_decode_gemv/k3_decode_gemv_kernel.py | 599 ++++++ .../cute_dsl_kernels/k3_decode_gemv/op.py | 209 ++ .../cute_dsl_kernels/k3_embed/__init__.py | 19 + .../k3_embed/k3_embed_kernel.py | 224 +++ .../_torch/cute_dsl_kernels/k3_embed/op.py | 181 ++ .../cute_dsl_kernels/k3_head_gemv/__init__.py | 19 + .../k3_head_gemv/k3_head_gemv_kernel.py | 1137 +++++++++++ .../cute_dsl_kernels/k3_head_gemv/op.py | 198 ++ .../kimi_k3/test_k3_cluster_waits.py | 97 + .../kimi_k3/test_k3_ctm_gemv.py | 281 +++ .../kimi_k3/test_k3_ctm_gemv_wide.py | 179 ++ .../kimi_k3/test_k3_decode_gemv.py | 167 ++ .../cute_dsl_kernels/kimi_k3/test_k3_embed.py | 574 ++++++ .../kimi_k3/test_k3_head_gemv.py | 176 ++ .../kimi_k3/test_k3_tcgen05_fences.py | 110 ++ 19 files changed, 6433 insertions(+) create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_ctm_gemv/__init__.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_ctm_gemv/k3_ctm_gemv_kernel.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_ctm_gemv/op.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_decode_gemv/__init__.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_decode_gemv/k3_decode_gemv_kernel.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_decode_gemv/op.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_embed/__init__.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_embed/k3_embed_kernel.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_embed/op.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_head_gemv/__init__.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_head_gemv/k3_head_gemv_kernel.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_head_gemv/op.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_cluster_waits.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctm_gemv.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctm_gemv_wide.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_decode_gemv.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_embed.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_head_gemv.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_ctm_gemv/__init__.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_ctm_gemv/__init__.py new file mode 100644 index 000000000000..71b0dc23eedd --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_ctm_gemv/__init__.py @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Kimi K3 CTM decode GEMV (``trtllm::k3_ctm_gemv``, ``_tail``, ``_gated``). + +Importing :mod:`.op` registers the torch ops; nothing is imported eagerly here so that +the CuTe DSL dependency stays optional for every other model. +""" diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_ctm_gemv/k3_ctm_gemv_kernel.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_ctm_gemv/k3_ctm_gemv_kernel.py new file mode 100644 index 000000000000..a405e008f725 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_ctm_gemv/k3_ctm_gemv_kernel.py @@ -0,0 +1,1681 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +# ============================================================================= +# Kimi K3 decode GEMV -- CTM (prims/cute) kernel, M <= 8 tokens, bf16 in/out +# ============================================================================= +# +# y[t, n] = sum_k B[t, k] * W[n, k] (A @ B^T, fp32 accumulation in TMEM) +# A = W (N, K <= 768) K-major bf16 (the streamed weight, 128 rows per cluster) +# B (M <= 8, K) K-major bf16 (8 token columns; rows past M are zero) +# +# Variants (compile-time): +# PLAIN B = x (KDA / MLA o_proj) +# TAIL y = rsqrt(mean(latent^2) + eps) * (latent_slice @ W_lat^T) + act @ W_act^T +# (two TMA maps, two TMEM accumulators; the row-parallel MoE tail) +# GATED B = bf16(a * bf16(sigmoid(g))) (MLA o_proj with its output gate) +# GATED_S B = bf16(a * s) with s = bf16(sigmoid(g)) precomputed (the long-K GEMV's sigmoid rows) +# SWIGLU B = bf16(g * sigmoid(g) * a) in fp32, g and a the two halves of a gate_up output (silu_and_mul folded +# into the down projection) +# a and g land by TMA in two identically swizzled tiles; the epilogue warps rewrite B in +# place (same offsets in both tiles), fence the generic writes to the async proxy and +# release the MMA through a 128-arrival barrier (the resident-B prolog of +# fuse_rmsnorm_qkv_rope). +# +# Geometry: one 128-row weight tile per cluster of SPLIT CTAs (1, 2 or 4). Rank r takes the interleaved +# k-tiles r, r + SPLIT, ... (when SPLIT does not divide the k-tiles, the first k-tiles % SPLIT ranks take one +# more) and owns output rows [R r, R r + R) of the tile, R = 128 / SPLIT (the TMEM +# lanes of epilogue warps r * 4 / SPLIT ...); every other rank's warps holding those rows send their +# fp32 partials to slot [source rank] of the owner's mailbox; the owner adds the SPLIT partials in rank +# order, then rounds once. Two ways to send (compile-time PUSH, the same sums): DSMEM stores and a release +# arrive at cluster scope on the owner's barrier, acquired by a try_wait at cluster scope; or 16-byte +# st.async stores that complete the owner's barrier by bytes, which the owner spins on with test_wait (acquire at +# cluster scope). +# +# Every k-tile of the CTA has its own shared-memory stage, and the whole weight slice is TMA'd +# (EVICT_FIRST) before griddepcontrol.wait: launched early under PDL the kernel streams its weight +# while the predecessor runs. Only the activation loads follow the wait. +# +# Warps: 0 weight TMA (+ early dependent trigger), 1 activation TMA after the grid dependency, +# 2 TMEM allocation + tcgen05 MMA (M 128, N 8, K 16), 3 idle, 4-7 epilogue (gate prologue, TMEM -> +# registers, split-K reduce, bf16 staging, 16-byte coalesced stores). +# ============================================================================= +"""CTM decode GEMV for Kimi K3 (``y = B @ W^T``, M <= 8): plain, MoE-tail and output-gated variants.""" + +from __future__ import annotations + +import cuda.bindings.driver as cuda_driver +import cutlass +import cutlass.cute as cute +import cutlass.experimental.cuda as cuda +from cutlass._mlir.dialects import llvm as _llvm +from cutlass._mlir.extras import types as _T +from cutlass.cutlass_dsl import dsl_user_op +from cutlass.experimental import primitives as prims + +CTA_M = 128 # weight rows (output features) per cluster = the MMA's M +MMA_N = 8 # token columns +CTA_K = 128 # one k-tile: two 64-element halves of the 128-byte swizzle +MMA_K = 16 +TMA_K_BOX = 64 +TMA_COPY_ITERS = CTA_K // TMA_K_BOX +K_BLOCKS_PER_HALF = TMA_K_BOX // MMA_K +MAX_K_TILES = 6 # 6 x (32 KB weight + 2 KB B [+ 2 KB gate]) of shared memory per CTA +THREADS = 256 +EPI_THREADS = 128 +TMEM_COLS = 32 +ELEM_BYTES = 2 +VEC = 8 # bf16 per 16-byte vector +EVICT_FIRST = 0x12F0000000000000 # createpolicy.fractional.L2::evict_first, fraction 1.0 (sm_100) + +# Shared-memory descriptor strides for the 128-byte swizzle, in 16-byte units. +LEADING = 16 +STRIDE = 8 * TMA_K_BOX * ELEM_BYTES +A_HALF_ELEMS = CTA_M * TMA_K_BOX +B_HALF_ELEMS = MMA_N * TMA_K_BOX +STEP = (MMA_K * ELEM_BYTES) >> 4 +A_BOX = A_HALF_ELEMS >> 3 +B_BOX = B_HALF_ELEMS >> 3 +STAGE_A = (CTA_M * CTA_K * ELEM_BYTES) >> 4 +STAGE_B = (MMA_N * CTA_K * ELEM_BYTES) >> 4 + +PLAIN = 0 +TAIL = 1 +GATED = 2 # B = bf16(a * bf16(sigmoid(g))), g given +GATED_S = 3 # B = bf16(a * s), s = bf16(sigmoid(g)) given (e.g. by k3_ctm_gemv_long's sigmoid rows) +SWIGLU = 4 # B = bf16(g * sigmoid(g) * a) in fp32: silu_and_mul of a gate_up output (g the first half, a the second) + +io_dtype = cutlass.BFloat16 + + +def num_k_tiles(k_in: int) -> int: + return (k_in + CTA_K - 1) // CTA_K + + +def supports(n_out: int, k_in: int, split: int = 1) -> bool: + """Shapes the kernel runs: whole 128-row tiles, whole k-tiles, at least one and at most MAX_K_TILES k-tiles per + CTA (the first k-tiles % split ranks take one more).""" + k_tiles = num_k_tiles(k_in) + return ( + n_out % CTA_M == 0 + and k_in % CTA_K == 0 + and split in (1, 2, 4) + and split <= k_tiles + and (k_tiles + split - 1) // split <= MAX_K_TILES + ) + + +@dsl_user_op +def _div_rn(a, b, *, loc=None, ip=None): + """IEEE fp32 division (div.rn.f32), as torch's ``one / (one + exp(-x))``.""" + return cutlass.Float32( + _llvm.inline_asm( + _T.f32(), [a.ir_value(loc=loc, ip=ip), b.ir_value(loc=loc, ip=ip)], + "div.rn.f32 $0, $1, $2;", "=f,f,f", has_side_effects=False, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _mapa_u32(smem_ptr, peer, *, loc=None, ip=None): + """The shared::cluster address of this CTA's shared-memory location in cluster CTA ``peer``.""" + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [smem_ptr.toint(loc=loc, ip=ip).ir_value(), cutlass.Int32(peer).ir_value(loc=loc, ip=ip)], + "mapa.shared::cluster.u32 $0, $1, $2;", "=r,r,r", has_side_effects=False, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _st_async_v4(dst, a, b, c, d, mbar, *, loc=None, ip=None): + """st.async of four fp32 (16 bytes) to a shared::cluster address, completing ``mbar`` (a shared::cluster address) + by 16 bytes.""" + _llvm.inline_asm( + None, + [cutlass.Int32(dst).ir_value(loc=loc, ip=ip), cutlass.Float32(a).ir_value(loc=loc, ip=ip), + cutlass.Float32(b).ir_value(loc=loc, ip=ip), cutlass.Float32(c).ir_value(loc=loc, ip=ip), + cutlass.Float32(d).ir_value(loc=loc, ip=ip), cutlass.Int32(mbar).ir_value(loc=loc, ip=ip)], + "st.async.shared::cluster.mbarrier::complete_tx::bytes.v4.f32 [$0], {$1, $2, $3, $4}, [$5];", "r,f,f,f,f,r", + has_side_effects=True, is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _test_wait_cluster(mbar, parity, *, loc=None, ip=None): + """mbarrier.test_wait.parity.acquire.cluster on a barrier of this CTA (``mbar``, its shared-memory pointer): whether + phase ``parity`` has completed, acquiring at cluster scope. The barriers it is used on are completed by other CTAs + (st.async complete_tx, remote arrives), which release at cluster scope.""" + done = cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [mbar.toint(loc=loc, ip=ip).ir_value(), cutlass.Int32(parity).ir_value(loc=loc, ip=ip)], + "{\n\t.reg .pred p;\n\tmbarrier.test_wait.parity.acquire.cluster.shared::cta.b64 p, [$1], $2;\n\t" + "selp.u32 $0, 1, 0, p;\n\t}", "=r,r,r,~{memory}", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + return done != cutlass.Int32(0) + + +def _bf16_rn(v): + """fp32 -> bf16 precision (round to nearest even) in the integer domain, kept as fp32. + + Bit operations, not a float cast, so no fptrunc/fpext pair exists that the compiler could + fold into the following multiply and skip this rounding. + """ + u = v.bitcast(cutlass.Int32) + u = u + (((u >> 16) & 1) + 0x7FFF) + u = (u >> 16) << 16 + return u.bitcast(cutlass.Float32) + + +def _sigmoid_bf16(g): + """bf16(sigmoid(g)) for a bf16 value g held in fp32, as torch computes it for bf16 tensors.""" + one = cutlass.Float32(1.0) + return _bf16_rn(_div_rn(one, one + cute.math.exp(-g, fastmath=False))) + + +@cute.kernel +def k3_ctm_gemv_kernel( + tma_desc_w: cutlass.GridConstant[cuda.TensorMap], # W [N, K] bf16, 5-D, one call per k-tile + tma_desc_x: cutlass.GridConstant[ + cuda.TensorMap + ], # plain: x; tail: latent; gated: a. Box 64 x 8 + tma_desc_x2: cutlass.GridConstant[ + cuda.TensorMap + ], # tail: act; gated: the tensor holding g. Box 64 x 8 + y: cutlass.Array, # [M * N] bf16, token-major + rms_src: cutlass.Array, # tail: int32 words of the bf16 [M, rms_cols] latent rows + num_tokens: cutlass.Int32, + x_col0: cutlass.Int32, # tail: column of the latent where k-tile 0 starts + x2_col0: cutlass.Int32, # gated: column of g in its tensor + eps: cutlass.Float32, + n_out: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + variant: cutlass.Constexpr[int], + x_tiles: cutlass.Constexpr[ + int + ], # tail: latent k-tiles (the rest accumulate from x2 into accumulator 1) + rms_cols: cutlass.Constexpr[int], + split: cutlass.Constexpr[int], + trigger_early: cutlass.Constexpr[bool], + push: cutlass.Constexpr[ + bool + ], # split-K partials by st.async (test_wait owner) instead of stores + arrive +): + """One weight tile of 128 rows per cluster; the CTA's k-tiles all resident; B by TMA (gated: rewritten).""" + k_tiles = num_k_tiles(k_in) + extra = k_tiles % split + my_tiles = (k_tiles + split - 1) // split # stages (the last one unused on ranks >= extra) + rows_owned = CTA_M // split # output rows reduced and stored by this CTA + tx, _, _ = cute.arch.thread_idx() + bx, _, _ = cute.arch.block_idx() + warp_id = cute.arch.warp_idx() + tma_ptr_w = tma_desc_w.get_ptr() + tma_ptr_x = tma_desc_x.get_ptr() + tma_ptr_x2 = tma_desc_x2.get_ptr() + rank = cutlass.Int32(0) + if cutlass.const_expr(split > 1): + rank = cute.arch.block_idx_in_cluster() + m_offset = (bx // cutlass.Int32(split)) * cutlass.Int32(CTA_M) + my_count = cutlass.Int32(my_tiles) + if cutlass.const_expr(extra > 0): + my_count = cutlass.Int32(k_tiles // split) + cutlass.Int32( + cutlass.select_(rank < cutlass.Int32(extra), cutlass.Int32(1), cutlass.Int32(0)) + ) + + # Allocation order is the same in every CTA, so mapa() addresses the peer's mailbox and barrier. + smem_a = cutlass.Array( + io_dtype, my_tiles * CTA_M * CTA_K, space=cutlass.AddressSpace.smem, alignment=1024 + ) + smem_b = cutlass.Array( + io_dtype, my_tiles * MMA_N * CTA_K, space=cutlass.AddressSpace.smem, alignment=1024 + ) + smem_g = smem_b + if cutlass.const_expr(variant >= GATED): + smem_g = cutlass.Array( + io_dtype, my_tiles * MMA_N * CTA_K, space=cutlass.AddressSpace.smem, alignment=1024 + ) + stage = cutlass.Array( + io_dtype, MMA_N * rows_owned, space=cutlass.AddressSpace.smem, alignment=16 + ) + weight_full = cutlass.Array( + cutlass.Int64, my_tiles, space=cutlass.AddressSpace.smem, alignment=8 + ) + act_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + b_ready = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + acc_done = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + mail_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + tmem_ptr_i32 = cutlass.Array(cutlass.Int32, 1, space=cutlass.AddressSpace.smem) + if cutlass.const_expr(variant == TAIL): + row_scale = cutlass.Array( + cutlass.Float32, MMA_N, space=cutlass.AddressSpace.smem, alignment=16 + ) + if cutlass.const_expr(split > 1): + # [source rank][row within the owned block][token] fp32 partials (the own rank's slot stays unused). + mailbox = cutlass.Array( + cutlass.Float32, + split * rows_owned * MMA_N, + space=cutlass.AddressSpace.smem, + alignment=16, + ) + + if warp_id == 0: + prims.prefetch_tensormap(tma_ptr_w) + prims.prefetch_tensormap(tma_ptr_x) + if cutlass.const_expr(variant != PLAIN): + prims.prefetch_tensormap(tma_ptr_x2) + if prims.elect_sync(): + for i in cutlass.range_constexpr(my_tiles): + prims.mbarrier_init(weight_full.subview(i), 1) + prims.mbarrier_init(act_full, 1) + prims.mbarrier_init(acc_done, 1) + if cutlass.const_expr(variant >= GATED): + prims.mbarrier_init(b_ready, EPI_THREADS) + if cutlass.const_expr(split > 1): + if cutlass.const_expr(push): + # The other ranks' partials of this CTA's rows arrive by st.async (16-byte stores completing + # this barrier's transaction count); expected here, before cluster formation. + prims.mbarrier_init(mail_full, 1) + prims.mbarrier_arrive_expect_tx(mail_full, (split - 1) * rows_owned * MMA_N * 4) + else: + # Every lane of the other ranks' epilogue warps holding this CTA's rows arrives once. + prims.mbarrier_init(mail_full, (split - 1) * rows_owned) + if warp_id == 2: + prims.tcgen05_alloc(tmem_ptr_i32, TMEM_COLS) + prims.tcgen05_relinquish_alloc_permit() + prims.fence_mbarrier_init() + if cutlass.const_expr(split > 1): + # Cluster formation: the peer's shared memory and barriers are addressable from here on. + prims.barrier_cluster_arrive_relaxed() + prims.barrier_cluster_wait() + prims.barrier_cta_sync(0) + tmem_ptr = cutlass.inttoptr(tmem_ptr_i32.load(), 6, cutlass.Int32) + + if warp_id == 0: + # ===================================================================== + # Weight TMA: every k-tile of this CTA, ahead of the grid dependency. + # ===================================================================== + if prims.elect_sync(): + for i in cutlass.range_constexpr(my_tiles): + k = rank + cutlass.Int32(i * split) + if cutlass.Int32(i) < my_count: + prims.mbarrier_arrive_expect_tx( + weight_full.subview(i), CTA_M * CTA_K * ELEM_BYTES + ) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a.subview(i * CTA_M * CTA_K), + tma_ptr_w, + ( + cutlass.Int32(0), + m_offset, + k * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), + cutlass.Int32(0), + ), + weight_full.subview(i), + l2_cache_hint=EVICT_FIRST, + ) + if cutlass.const_expr(trigger_early): + # Dependents may launch now; they wait for this whole grid before reading y. + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + elif warp_id == 1: + # ===================================================================== + # Activation TMA: B (and g), written by predecessors -> after the wait. + # ===================================================================== + prims.griddepcontrol(prims.GridDepAction.WAIT) + if prims.elect_sync(): + tile_bytes = MMA_N * CTA_K * ELEM_BYTES + if cutlass.const_expr(variant >= GATED): + tile_bytes = 2 * tile_bytes + prims.mbarrier_arrive_expect_tx(act_full, my_count * cutlass.Int32(tile_bytes)) + for i in cutlass.range_constexpr(my_tiles): + k = rank + cutlass.Int32(i * split) + k_c = cutlass.Int32(cutlass.select_(cutlass.Int32(i) < my_count, k, rank)) + for half in cutlass.range_constexpr(TMA_COPY_ITERS): + b_off = i * MMA_N * CTA_K + half * B_HALF_ELEMS + col = k_c * cutlass.Int32(CTA_K) + cutlass.Int32(half * TMA_K_BOX) + if cutlass.const_expr(variant == TAIL): + if cutlass.const_expr(i < x_tiles): + prims.cp_async_bulk_tensor_shared_cta_global( + smem_b.subview(b_off), + tma_ptr_x, + (x_col0 + col, cutlass.Int32(0)), + act_full, + ) + else: + prims.cp_async_bulk_tensor_shared_cta_global( + smem_b.subview(b_off), + tma_ptr_x2, + ( + cutlass.Int32((i - x_tiles) * CTA_K + half * TMA_K_BOX), + cutlass.Int32(0), + ), + act_full, + ) + else: + if cutlass.Int32(i) < my_count: + prims.cp_async_bulk_tensor_shared_cta_global( + smem_b.subview(b_off), + tma_ptr_x, + (x_col0 + col, cutlass.Int32(0)), + act_full, + ) + if cutlass.const_expr(variant >= GATED): + prims.cp_async_bulk_tensor_shared_cta_global( + smem_g.subview(b_off), + tma_ptr_x2, + (x2_col0 + col, cutlass.Int32(0)), + act_full, + ) + elif warp_id == 2: + # ===================================================================== + # MMA: this CTA's k-tiles in ascending order into one TMEM accumulator + # (tail: the act k-tiles into a second one). + # ===================================================================== + idesc = prims.Tcgen05InstrDesc.build( + c_dtype=cutlass.Float32, a_dtype=io_dtype, b_dtype=io_dtype, n_dim=MMA_N, m_dim=CTA_M + ) + desc_a_base = prims.Tcgen05SmemDesc.build( + start_address=smem_a, leading_byte_offset=LEADING, stride_byte_offset=STRIDE, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + desc_b_base = prims.Tcgen05SmemDesc.build( + start_address=smem_b, leading_byte_offset=LEADING, stride_byte_offset=STRIDE, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + split_acc = variant == TAIL + tmem_acc1 = tmem_ptr + if cutlass.const_expr(split_acc): + tmem_acc1 = cutlass.inttoptr( + tmem_ptr_i32.load() + cutlass.Int32(MMA_N), 6, cutlass.Int32 + ) + # B is ready when its TMA lands (gated: when the epilogue warps have rewritten it). + if cutlass.const_expr(variant >= GATED): + while not cute.arch.mbarrier_try_wait(b_ready.data_ptr(), 0): + pass + else: + while not cute.arch.mbarrier_try_wait(act_full.data_ptr(), 0): + pass + for i in cutlass.range_constexpr(my_tiles): + second = split_acc and i >= x_tiles + first_tile = i == 0 or (split_acc and i == x_tiles) + if cutlass.Int32(i) < my_count: + while not cute.arch.mbarrier_try_wait(weight_full.subview(i).data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + for kb in cutlass.range_constexpr(TMA_COPY_ITERS * K_BLOCKS_PER_HALF): + box = kb // K_BLOCKS_PER_HALF + within = kb % K_BLOCKS_PER_HALF + desc_a = desc_a_base + (i * STAGE_A + box * A_BOX + within * STEP) + desc_b = desc_b_base + (i * STAGE_B + box * B_BOX + within * STEP) + if prims.elect_sync(): + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, prims.CTAGroup.CTA_1, tmem_acc1 if second else tmem_ptr, + desc_a, desc_b, idesc, not (first_tile and kb == 0), + ) # fmt: skip + if prims.elect_sync(): + prims.tcgen05_commit(acc_done) + elif warp_id >= 4: + # ===================================================================== + # Epilogue warps: [gate prologue] -> TMEM -> registers -> [split-K + # reduce] -> bf16 staging -> 16-byte stores. + # warp 4 + w reads TMEM lanes 32 w .. 32 w + 31 = tile rows 32 w + lane. + # ===================================================================== + tid = tx - cutlass.Int32(EPI_THREADS) + lane = tx % 32 + w = warp_id - 4 + if cutlass.const_expr(variant >= GATED): + # B = bf16(a * bf16(sigmoid(g))): a and g share one swizzled layout, so chunk c of B is + # chunk c of both inputs. 16-byte chunks, consecutive threads -> no bank conflicts. + while not cute.arch.mbarrier_try_wait(act_full.data_ptr(), 0): + pass + for j in cutlass.range_constexpr(my_tiles * MMA_N * CTA_K // (VEC * EPI_THREADS)): + c = (tid + cutlass.Int32(j * EPI_THREADS)) * cutlass.Int32(VEC) + av = smem_b.load(idx=c, vector_size=VEC, alignment=16) + gv = smem_g.load(idx=c, vector_size=VEC, alignment=16) + outs = [] + for e in cutlass.range_constexpr(VEC): + if cutlass.const_expr(variant == SWIGLU): + # silu_and_mul's order in fp32: g * sigmoid(g), then times the up value, one rounding. + g = cutlass.Float32(gv[e]) + one = cutlass.Float32(1.0) + sig = _div_rn(one, one + cute.math.exp(-g, fastmath=False)) + outs.append(((g * sig) * cutlass.Float32(av[e])).to(io_dtype)) + else: + if cutlass.const_expr(variant == GATED): + gate = _sigmoid_bf16(cutlass.Float32(gv[e])) + else: + gate = cutlass.Float32(gv[e]) + outs.append((cutlass.Float32(av[e]) * gate).to(io_dtype)) + smem_b.store( + cutlass.Vector.from_elements(tuple(outs), io_dtype), + idx=c, + vector_size=VEC, + alignment=16, + ) + # Generic-proxy writes of B -> the tensor core's async-proxy reads, then release the MMA. + prims.fence_proxy(prims.Proxy.ASYNC_SHARED, space=prims.SharedSpace.shared_cta) + prims.mbarrier_arrive(b_ready) + if cutlass.const_expr(variant == TAIL): + # Per-token RMS of the latent rows while the MMA runs: epilogue warp w owns tokens 2w, 2w+1. + prims.griddepcontrol(prims.GridDepAction.WAIT) + for j in cutlass.range_constexpr(MMA_N // 4): + t = w * cutlass.Int32(MMA_N // 4) + cutlass.Int32(j) + sum_sq = cutlass.Float32(0.0) + if t < num_tokens: + for v in cutlass.range_constexpr(rms_cols // (8 * 32)): + words = rms_src.load( + idx=t * cutlass.Int32(rms_cols // 2) + + cutlass.Int32(v * 32 * 4) + + lane * cutlass.Int32(4), + vector_size=4, + alignment=16, + ) + for q in cutlass.range_constexpr(4): + lo = (words[q] << cutlass.Int32(16)).bitcast(cutlass.Float32) + hi = (words[q] & cutlass.Int32(-65536)).bitcast(cutlass.Float32) + sum_sq = sum_sq + lo * lo + hi * hi + for offset in (16, 8, 4, 2, 1): + sum_sq = sum_sq + cute.arch.shuffle_sync_bfly(sum_sq, offset=offset) + if lane == 0: + row_scale.store( + cute.math.rsqrt(sum_sq * cutlass.Float32(1.0 / rms_cols) + eps), idx=t + ) + prims.barrier_cta_sync(1, thread_count=EPI_THREADS) + + while not cute.arch.mbarrier_try_wait(acc_done.data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc0 = prims.tcgen05_ld( + "32x32b", cutlass.inttoptr(tmem_ptr_i32.load(), 6, cutlass.Float32), num=MMA_N + ) + acc1 = acc0 + if cutlass.const_expr(variant == TAIL): + acc1 = prims.tcgen05_ld( + "32x32b", + cutlass.inttoptr(tmem_ptr_i32.load() + cutlass.Int32(MMA_N), 6, cutlass.Float32), + num=MMA_N, + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + vals = [] + for t in cutlass.range_constexpr(MMA_N): + value = cutlass.Float32(acc0[t]) + if cutlass.const_expr(variant == TAIL): + value = value * row_scale.load(idx=t) + cutlass.Float32(acc1[t]) + vals.append(value) + + row = w * cutlass.Int32(32) + lane # tile row of this thread + local_row = row + if cutlass.const_expr(split > 1): + # Rows [R r, R r + R) belong to rank r = w // (4 / SPLIT); the other ranks push, the owner adds. + owner = w // cutlass.Int32(rows_owned // 32) + local_row = row - owner * cutlass.Int32(rows_owned) + if owner != rank: + slot = (rank * cutlass.Int32(rows_owned) + local_row) * cutlass.Int32(MMA_N) + if cutlass.const_expr(push): + peer_slot = _mapa_u32(mailbox.subview(slot).data_ptr(), owner) + mbar_peer = _mapa_u32(mail_full.data_ptr(), owner) + _st_async_v4(peer_slot, vals[0], vals[1], vals[2], vals[3], mbar_peer) + _st_async_v4( + peer_slot + cutlass.Int32(16), vals[4], vals[5], vals[6], vals[7], mbar_peer + ) + else: + for t in cutlass.range_constexpr(MMA_N): + prims.mapa(mailbox.subview(slot + cutlass.Int32(t)), owner).store(vals[t]) + # Release at cluster scope: the partials above are visible to the owner's acquire. + prims.mbarrier_arrive( + prims.mapa(mail_full, owner), scope=prims.MemScope.CLUSTER + ) + else: + if cutlass.const_expr(push): + # test_wait spin: a warp suspended in try_wait on a barrier completed by peers wakes late. + # Acquire at cluster scope: the bytes come from other CTAs' st.async (complete_tx releases at + # cluster scope). + while not _test_wait_cluster(mail_full.data_ptr(), 0): + pass + else: + while not prims.mbarrier_try_wait_parity( + mail_full, 0, scope=prims.MBarrierScope.CLUSTER + ): + pass + # The SPLIT partials in rank order (this rank's own from registers). + summed = [cutlass.Float32(0.0)] * MMA_N + for q in cutlass.range_constexpr(split): + slot = (cutlass.Int32(q * rows_owned) + local_row) * cutlass.Int32(MMA_N) + for t in cutlass.range_constexpr(MMA_N): + part = mailbox.load(idx=slot + cutlass.Int32(t)) + summed[t] = summed[t] + cutlass.Float32( + cutlass.select_(rank == cutlass.Int32(q), vals[t], part) + ) + for t in cutlass.range_constexpr(MMA_N): + stage.store( + summed[t].to(io_dtype), idx=cutlass.Int32(t * rows_owned) + local_row + ) + else: + for t in cutlass.range_constexpr(MMA_N): + stage.store(vals[t].to(io_dtype), idx=cutlass.Int32(t * rows_owned) + local_row) + # Every TMEM reader has waited for its load; the staging tile is complete. + prims.barrier_cta_sync(1, thread_count=EPI_THREADS) + if warp_id == 4: + prims.tcgen05_dealloc(tmem_ptr, TMEM_COLS) + # [token][rows_owned] bf16 -> y[token, rows]: 16 bytes per thread, rows contiguous per token. + chunks_per_token = rows_owned // VEC + ct = tid // cutlass.Int32(chunks_per_token) + cc = tid % cutlass.Int32(chunks_per_token) + if ct < num_tokens: + if ct < cutlass.Int32(MMA_N): + v = stage.load( + idx=ct * cutlass.Int32(rows_owned) + cc * cutlass.Int32(VEC), + vector_size=VEC, + alignment=16, + ) + y.store( + v, + idx=ct * cutlass.Int32(n_out) + + m_offset + + rank * cutlass.Int32(rows_owned) + + cc * cutlass.Int32(VEC), + vector_size=VEC, + alignment=16, + ) + + +# ============================================================================= +# Long-K variant (K > 6 k-tiles, e.g. the MLA [W_a; W_g] projection 2880 x 7168): one 128-row tile per cluster +# of SPLIT >= 4 CTAs, rank r streaming the interleaved k-tiles r, r + SPLIT, ... (the first `extra` ranks one +# more when SPLIT does not divide the k-tiles) through a RING-stage ring filled before griddepcontrol.wait, +# the rest of its k-tiles prefetched into L2 at the same time, so a latency-bound predecessor (an all-reduce) +# hides the whole weight read. B (all of the rank's k-tiles of x) is resident. Rows [32 w, 32 w + 32) of the +# tile belong to rank w < 4 (epilogue warp w's TMEM lanes; with SPLIT 2, rank w // 2 owns two such blocks); every +# other rank sends its fp32 partials of them to slot [rank] of the owner's mailbox (PUSH as in the short kernel: +# DSMEM stores + a release arrive at cluster scope, or st.async completing the barrier by bytes); the owner adds +# the SPLIT partials in rank order and rounds once. Output rows >= sig_row0 store bf16(sigmoid(bf16(acc))) instead +# of bf16(acc): the MLA output gate, ready for GATED_S. +# ============================================================================= +LONG_ROWS_PER_WARP = 32 + + +def long_supports(n_out: int, k_in: int, split: int, ring: int) -> bool: + """K in whole 64-column halves: a last k-tile of 64 columns has its second half past K, which the TMA fills with + zeros in both W and x (e.g. the dense MLP's down projection, K = 2112 at TP16).""" + k_tiles = num_k_tiles(k_in) + return ( + k_in % TMA_K_BOX == 0 + and (split == 2 or 4 <= split <= 8) + and 1 <= ring <= k_tiles // split + and ring * CTA_M * CTA_K * ELEM_BYTES + (k_tiles // split + 1) * MMA_N * CTA_K * ELEM_BYTES + <= 216 * 1024 + and n_out > 0 + ) + + +@cute.kernel +def k3_ctm_gemv_long_kernel( + tma_desc_w: cutlass.GridConstant[cuda.TensorMap], # W [N, K] bf16, 5-D, one call per k-tile + tma_desc_x: cutlass.GridConstant[cuda.TensorMap], # x [M, K] bf16, box 64 x 8 + y: cutlass.Array, # [M * N] bf16, token-major + num_tokens: cutlass.Int32, + sig_row0: cutlass.Int32, # rows >= sig_row0 store bf16(sigmoid(bf16(acc))); n_out: none + n_out: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + split: cutlass.Constexpr[int], + ring: cutlass.Constexpr[int], + trigger_early: cutlass.Constexpr[bool], + push: cutlass.Constexpr[ + bool + ], # split-K partials by st.async (test_wait owner) instead of stores + arrive +): + """One 128-row tile per cluster of `split` CTAs; weight ring + L2 prefetch before the grid wait.""" + k_tiles = num_k_tiles(k_in) + extra = k_tiles % split + max_tiles = k_tiles // split + (1 if extra > 0 else 0) + tx, _, _ = cute.arch.thread_idx() + bx, _, _ = cute.arch.block_idx() + warp_id = cute.arch.warp_idx() + rank = cute.arch.block_idx_in_cluster() + tma_ptr_w = tma_desc_w.get_ptr() + tma_ptr_x = tma_desc_x.get_ptr() + m_offset = (bx // cutlass.Int32(split)) * cutlass.Int32(CTA_M) + my_tiles = cutlass.Int32(k_tiles // split) + if cutlass.const_expr(extra > 0): + if rank < cutlass.Int32(extra): + my_tiles = my_tiles + cutlass.Int32(1) + + smem_a = cutlass.Array( + io_dtype, ring * CTA_M * CTA_K, space=cutlass.AddressSpace.smem, alignment=1024 + ) + smem_b = cutlass.Array( + io_dtype, max_tiles * MMA_N * CTA_K, space=cutlass.AddressSpace.smem, alignment=1024 + ) + tma_full = cutlass.Array(cutlass.Int64, ring, space=cutlass.AddressSpace.smem, alignment=8) + mma_done = cutlass.Array(cutlass.Int64, ring, space=cutlass.AddressSpace.smem, alignment=8) + act_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + acc_done = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + mail_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + tmem_ptr_i32 = cutlass.Array(cutlass.Int32, 1, space=cutlass.AddressSpace.smem) + # [source rank][owned block][row in the block][token] fp32 partials (the own rank's slots stay unused). A rank owns + # one 32-row block with SPLIT >= 4, two with SPLIT 2. + blocks = 4 // split if split < 4 else 1 + mailbox = cutlass.Array( + cutlass.Float32, + split * blocks * LONG_ROWS_PER_WARP * MMA_N, + space=cutlass.AddressSpace.smem, + alignment=16, + ) + + if warp_id == 0: + prims.prefetch_tensormap(tma_ptr_w) + prims.prefetch_tensormap(tma_ptr_x) + if prims.elect_sync(): + for s in cutlass.range_constexpr(ring): + prims.mbarrier_init(tma_full.subview(s), 1) + prims.mbarrier_init(mma_done.subview(s), 1) + prims.mbarrier_init(act_full, 1) + prims.mbarrier_init(acc_done, 1) + if cutlass.const_expr(push): + # Owner ranks (< 4): the other ranks' partials of the owned rows arrive by st.async (completing + # the transaction count expected here, before cluster formation); other ranks never wait on it. + prims.mbarrier_init(mail_full, 1) + prims.mbarrier_arrive_expect_tx( + mail_full, (split - 1) * blocks * LONG_ROWS_PER_WARP * MMA_N * 4 + ) + else: + # Owner ranks: every lane of the other ranks' epilogue warps of the owned blocks arrives once. + prims.mbarrier_init(mail_full, (split - 1) * blocks * LONG_ROWS_PER_WARP) + if warp_id == 2: + prims.tcgen05_alloc(tmem_ptr_i32, TMEM_COLS) + prims.tcgen05_relinquish_alloc_permit() + prims.fence_mbarrier_init() + # Cluster formation: the peers' shared memory and barriers are addressable from here on. + prims.barrier_cluster_arrive_relaxed() + prims.barrier_cluster_wait() + prims.barrier_cta_sync(0) + tmem_ptr = cutlass.inttoptr(tmem_ptr_i32.load(), 6, cutlass.Int32) + + if warp_id == 0: + # ===================================================================== + # Weight TMA: fill the ring and prefetch the rest into L2 before the + # grid dependency; then refill each stage when its MMAs are done. + # ===================================================================== + if prims.elect_sync(): + for i in cutlass.range_constexpr(ring): + k = rank + cutlass.Int32(i * split) + prims.mbarrier_arrive_expect_tx(tma_full.subview(i), CTA_M * CTA_K * ELEM_BYTES) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a.subview(i * CTA_M * CTA_K), + tma_ptr_w, + ( + cutlass.Int32(0), + m_offset, + k * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), + cutlass.Int32(0), + ), + tma_full.subview(i), + l2_cache_hint=EVICT_FIRST, + ) + for i in range(ring, my_tiles): + k = rank + i * cutlass.Int32(split) + prims.cp_async_bulk_tensor_prefetch( + tma_ptr_w, + [ + cutlass.Int32(0), + m_offset, + k * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), + cutlass.Int32(0), + ], + [], # tile mode: no im2col offsets + ) + if cutlass.const_expr(trigger_early): + # Dependents may launch now; they wait for this whole grid before reading y. + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + if prims.elect_sync(): + stage = cutlass.Int32(0) + phase = cutlass.Int32(0) + for i in range(ring, my_tiles): + while not cute.arch.mbarrier_try_wait(mma_done.subview(stage).data_ptr(), phase): + pass + k = rank + i * cutlass.Int32(split) + prims.mbarrier_arrive_expect_tx(tma_full.subview(stage), CTA_M * CTA_K * ELEM_BYTES) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a.subview(stage * cutlass.Int32(CTA_M * CTA_K)), + tma_ptr_w, + ( + cutlass.Int32(0), + m_offset, + k * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), + cutlass.Int32(0), + ), + tma_full.subview(stage), + l2_cache_hint=EVICT_FIRST, + ) + stage = stage + cutlass.Int32(1) + if stage == cutlass.Int32(ring): + stage = cutlass.Int32(0) + phase = phase ^ cutlass.Int32(1) + elif warp_id == 1: + # ===================================================================== + # Activation TMA: the rank's k-tiles of x, resident, after the wait. + # ===================================================================== + prims.griddepcontrol(prims.GridDepAction.WAIT) + if prims.elect_sync(): + prims.mbarrier_arrive_expect_tx( + act_full, my_tiles * cutlass.Int32(MMA_N * CTA_K * ELEM_BYTES) + ) + for i in range(my_tiles): + k = rank + i * cutlass.Int32(split) + for half in cutlass.range_constexpr(TMA_COPY_ITERS): + prims.cp_async_bulk_tensor_shared_cta_global( + smem_b.subview( + i * cutlass.Int32(MMA_N * CTA_K) + cutlass.Int32(half * B_HALF_ELEMS) + ), + tma_ptr_x, + ( + k * cutlass.Int32(CTA_K) + cutlass.Int32(half * TMA_K_BOX), + cutlass.Int32(0), + ), + act_full, + ) + elif warp_id == 2: + # ===================================================================== + # MMA: the rank's k-tiles in order, one TMEM accumulator (partial sum). + # ===================================================================== + idesc = prims.Tcgen05InstrDesc.build( + c_dtype=cutlass.Float32, a_dtype=io_dtype, b_dtype=io_dtype, n_dim=MMA_N, m_dim=CTA_M + ) + desc_a_base = prims.Tcgen05SmemDesc.build( + start_address=smem_a, leading_byte_offset=LEADING, stride_byte_offset=STRIDE, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + desc_b_base = prims.Tcgen05SmemDesc.build( + start_address=smem_b, leading_byte_offset=LEADING, stride_byte_offset=STRIDE, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + while not cute.arch.mbarrier_try_wait(act_full.data_ptr(), 0): + pass + stage = cutlass.Int32(0) + phase = cutlass.Int32(0) + for i in range(my_tiles): + while not cute.arch.mbarrier_try_wait(tma_full.subview(stage).data_ptr(), phase): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + for kb in cutlass.range_constexpr(TMA_COPY_ITERS * K_BLOCKS_PER_HALF): + box = kb // K_BLOCKS_PER_HALF + within = kb % K_BLOCKS_PER_HALF + desc_a = desc_a_base + ( + stage * cutlass.Int32(STAGE_A) + cutlass.Int32(box * A_BOX + within * STEP) + ) + desc_b = desc_b_base + ( + i * cutlass.Int32(STAGE_B) + cutlass.Int32(box * B_BOX + within * STEP) + ) + accumulate = cutlass.Boolean(True) + if cutlass.const_expr(kb == 0): + accumulate = i > cutlass.Int32(0) + if prims.elect_sync(): + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, + prims.CTAGroup.CTA_1, + tmem_ptr, + desc_a, + desc_b, + idesc, + accumulate, + ) + if prims.elect_sync(): + prims.tcgen05_commit(mma_done.subview(stage)) + stage = stage + cutlass.Int32(1) + if stage == cutlass.Int32(ring): + stage = cutlass.Int32(0) + phase = phase ^ cutlass.Int32(1) + if prims.elect_sync(): + prims.tcgen05_commit(acc_done) + elif warp_id >= 4: + # ===================================================================== + # Epilogue: TMEM -> registers, push to / reduce at the row owner, store. + # ===================================================================== + lane = tx % 32 + w = warp_id - 4 # TMEM lanes 32 w .. 32 w + 31: tile rows owned by rank w + while not cute.arch.mbarrier_try_wait(acc_done.data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc = prims.tcgen05_ld( + "32x32b", cutlass.inttoptr(tmem_ptr_i32.load(), 6, cutlass.Float32), num=MMA_N + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + slot_row = lane * cutlass.Int32(MMA_N) + owner = w // cutlass.Int32(blocks) + block_slot = (w % cutlass.Int32(blocks)) * cutlass.Int32( + LONG_ROWS_PER_WARP * MMA_N + ) + slot_row + if owner != rank: + base = rank * cutlass.Int32(blocks * LONG_ROWS_PER_WARP * MMA_N) + block_slot + if cutlass.const_expr(push): + peer_slot = _mapa_u32(mailbox.subview(base).data_ptr(), owner) + mbar_peer = _mapa_u32(mail_full.data_ptr(), owner) + _st_async_v4(peer_slot, acc[0], acc[1], acc[2], acc[3], mbar_peer) + _st_async_v4( + peer_slot + cutlass.Int32(16), acc[4], acc[5], acc[6], acc[7], mbar_peer + ) + else: + for t in cutlass.range_constexpr(MMA_N): + prims.mapa(mailbox.subview(base + cutlass.Int32(t)), owner).store( + cutlass.Float32(acc[t]) + ) + # Release at cluster scope: this lane's partials are visible to the owner's acquire. + prims.mbarrier_arrive(prims.mapa(mail_full, owner), scope=prims.MemScope.CLUSTER) + else: + if cutlass.const_expr(push): + # test_wait spin: a warp suspended in try_wait on a barrier completed by peers wakes late. + # Acquire at cluster scope: the bytes come from other CTAs' st.async (complete_tx releases at + # cluster scope). + while not _test_wait_cluster(mail_full.data_ptr(), 0): + pass + else: + while not prims.mbarrier_try_wait_parity( + mail_full, 0, scope=prims.MBarrierScope.CLUSTER + ): + pass + n = m_offset + w * cutlass.Int32(LONG_ROWS_PER_WARP) + lane + for t in cutlass.range_constexpr(MMA_N): + total = cutlass.Float32(0.0) + for q in cutlass.range_constexpr(split): + part = mailbox.load( + idx=cutlass.Int32(q * blocks * LONG_ROWS_PER_WARP * MMA_N) + + block_slot + + cutlass.Int32(t) + ) + total = total + cutlass.Float32( + cutlass.select_(rank == cutlass.Int32(q), cutlass.Float32(acc[t]), part) + ) + out = total.to(io_dtype) + if n >= sig_row0: + out = _sigmoid_bf16(_bf16_rn(total)).to(io_dtype) + if cutlass.Int32(t) < num_tokens: + if n < cutlass.Int32(n_out): + y.store(out, idx=cutlass.Int32(t * n_out) + n) + prims.barrier_cta_sync(1, thread_count=EPI_THREADS) + if warp_id == 4: + prims.tcgen05_dealloc(tmem_ptr, TMEM_COLS) + + +def _weight_tensor_map(w, n_out, k_in): + """W as five TMA dimensions (64-element column chunk, row, 64-element chunk index, 1, 1) so one call + per k-tile lands both 128-byte-swizzled halves; strides in 16-byte units.""" + return cuda.create_tensor_map_tiled( + global_address=w.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[TMA_K_BOX, n_out, k_in // TMA_K_BOX, 1, 1], + global_strides=[ + (k_in * ELEM_BYTES) // 16, + (TMA_K_BOX * ELEM_BYTES) // 16, + (n_out * k_in * ELEM_BYTES) // 16, + (n_out * k_in * ELEM_BYTES) // 16, + ], + box_dims=[TMA_K_BOX, CTA_M, TMA_COPY_ITERS, 1, 1], + swizzle=cuda.TensorMapSwizzle.s128b, + ) + + +def _activation_tensor_map(x, cols, num_tokens): + """x [M, cols] (rows dense) as (cols, M) with an 8-row box: rows past num_tokens arrive as zeros.""" + return cuda.create_tensor_map_tiled( + global_address=x.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[cols, num_tokens], + global_strides=[(cols * ELEM_BYTES) // 16], + box_dims=[TMA_K_BOX, MMA_N], + swizzle=cuda.TensorMapSwizzle.s128b, + ) + + +def _launch(kernel, n_out, split, use_pdl, stream): + kernel.launch( + grid=((n_out // CTA_M) * split, 1, 1), + block=(THREADS, 1, 1), + cluster=(split, 1, 1) if split > 1 else None, + stream=stream, + use_pdl=use_pdl, + ) + + +@cute.jit +def k3_ctm_gemv( + w: cute.Tensor, # [N, K] bf16, K contiguous + x: cute.Tensor, # [M, K] bf16, K contiguous, M <= 8 + y: cute.Tensor, # [M * N] bf16 + num_tokens: cutlass.Int32, + n_out: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + split: cutlass.Constexpr[int], + trigger_early: cutlass.Constexpr[bool], + push: cutlass.Constexpr[bool], + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + """``y = x @ w^T``.""" + tma_desc_w = _weight_tensor_map(w, n_out, k_in) + tma_desc_x = _activation_tensor_map(x, k_in, num_tokens) + kernel = k3_ctm_gemv_kernel( + tma_desc_w, tma_desc_x, tma_desc_x, y, y, num_tokens, cutlass.Int32(0), cutlass.Int32(0), + cutlass.Float32(0.0), n_out, k_in, PLAIN, num_k_tiles(k_in), 0, split, trigger_early, push, + ) # fmt: skip + _launch(kernel, n_out, split, use_pdl, stream) + + +@cute.jit +def k3_ctm_gemv_tail( + w: cute.Tensor, # [N, K_lat + K_act] bf16: latent-up columns of this rank's slice (zero-padded) | shared down + latent: cute.Tensor, # [M, rms_cols] bf16, the whole reduced latent row + latent_words: cute.Tensor, # the same memory as int32 words, for the RMS + act: cute.Tensor, # [M, K_act] bf16, the shared-expert activation + y: cute.Tensor, # [M * N] bf16 + num_tokens: cutlass.Int32, + lat_col0: cutlass.Int32, # first latent column of this rank's slice + eps: cutlass.Float32, + n_out: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + lat_tiles: cutlass.Constexpr[int], + rms_cols: cutlass.Constexpr[int], + trigger_early: cutlass.Constexpr[bool], + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + """``y = rmsnorm(latent)[:, slice] @ W_lat^T + act @ W_act^T`` (the RMS applied to the latent accumulator).""" + tma_desc_w = _weight_tensor_map(w, n_out, k_in) + tma_desc_lat = _activation_tensor_map(latent, rms_cols, num_tokens) + tma_desc_act = _activation_tensor_map(act, k_in - lat_tiles * CTA_K, num_tokens) + kernel = k3_ctm_gemv_kernel( + tma_desc_w, tma_desc_lat, tma_desc_act, y, latent_words, num_tokens, lat_col0, cutlass.Int32(0), eps, + n_out, k_in, TAIL, lat_tiles, rms_cols, 1, trigger_early, False, + ) # fmt: skip + _launch(kernel, n_out, 1, use_pdl, stream) + + +@cute.jit +def k3_ctm_gemv_gated( + w: cute.Tensor, # [N, K] bf16 (the o_proj weight), K contiguous + a: cute.Tensor, # [M, K] bf16, the attention output (o_proj's input before the gate) + gsrc: cute.Tensor, # [M, gsrc_cols] bf16, dense rows; g = gsrc[:, g_col0 : g_col0 + K] + y: cute.Tensor, # [M * N] bf16 + num_tokens: cutlass.Int32, + g_col0: cutlass.Int32, + n_out: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + gsrc_cols: cutlass.Constexpr[int], + split: cutlass.Constexpr[int], + gate_sigmoid: cutlass.Constexpr[bool], # True: gsrc holds g; False: it holds bf16(sigmoid(g)) + trigger_early: cutlass.Constexpr[bool], + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + """``y = (a * sigmoid(g)) @ w^T`` with torch's bf16 roundings of the sigmoid and of the product.""" + tma_desc_w = _weight_tensor_map(w, n_out, k_in) + tma_desc_a = _activation_tensor_map(a, k_in, num_tokens) + tma_desc_g = _activation_tensor_map(gsrc, gsrc_cols, num_tokens) + kernel = k3_ctm_gemv_kernel( + tma_desc_w, tma_desc_a, tma_desc_g, y, y, num_tokens, cutlass.Int32(0), g_col0, cutlass.Float32(0.0), + n_out, k_in, GATED if gate_sigmoid else GATED_S, num_k_tiles(k_in), 0, split, trigger_early, False, + ) # fmt: skip + _launch(kernel, n_out, split, use_pdl, stream) + + +@cute.jit +def k3_ctm_gemv_swiglu( + w: cute.Tensor, # [N, K] bf16 (the down projection), K contiguous + gu: cute.Tensor, # [M, 2 K] bf16, dense rows: the gate_up output (gate columns first) + y: cute.Tensor, # [M * N] bf16 + num_tokens: cutlass.Int32, + n_out: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + split: cutlass.Constexpr[int], + trigger_early: cutlass.Constexpr[bool], + push: cutlass.Constexpr[bool], + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + """``y = (silu(gu[:, :K]) * gu[:, K:]) @ w^T`` with silu_and_mul's fp32 arithmetic and bf16 rounding.""" + tma_desc_w = _weight_tensor_map(w, n_out, k_in) + tma_desc_gu = _activation_tensor_map(gu, 2 * k_in, num_tokens) + kernel = k3_ctm_gemv_kernel( + tma_desc_w, tma_desc_gu, tma_desc_gu, y, y, num_tokens, cutlass.Int32(k_in), cutlass.Int32(0), + cutlass.Float32(0.0), n_out, k_in, SWIGLU, num_k_tiles(k_in), 0, split, trigger_early, push, + ) # fmt: skip + _launch(kernel, n_out, split, use_pdl, stream) + + +@cute.jit +def k3_ctm_gemv_long( + w: cute.Tensor, # [N, K] bf16, K contiguous + x: cute.Tensor, # [M, K] bf16, K contiguous, M <= 8 + y: cute.Tensor, # [M * N] bf16 + num_tokens: cutlass.Int32, + sig_row0: cutlass.Int32, + n_out: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + split: cutlass.Constexpr[int], + ring: cutlass.Constexpr[int], + trigger_early: cutlass.Constexpr[bool], + push: cutlass.Constexpr[bool], + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + """``y = x @ w^T``; output rows >= sig_row0 hold bf16(sigmoid(bf16(.))) instead.""" + tma_desc_w = _weight_tensor_map(w, n_out, k_in) + tma_desc_x = _activation_tensor_map(x, k_in, num_tokens) + k3_ctm_gemv_long_kernel( + tma_desc_w, + tma_desc_x, + y, + num_tokens, + sig_row0, + n_out, + k_in, + split, + ring, + trigger_early, + push, + ).launch( + grid=(((n_out + CTA_M - 1) // CTA_M) * split, 1, 1), + block=(THREADS, 1, 1), + cluster=(split, 1, 1), + stream=stream, + use_pdl=use_pdl, + ) + + +# ============================================================================= +# Wide variant (up to 64 tokens: DSpark verify of several requests): the long kernel's geometry, weight ring and L2 +# prefetch, with all the call's token columns in one MMA of N = N_TILE (16, 32 or 64). x no longer fits resident +# beside the ring, so each ring stage holds a weight k-tile and the matching x k-tile: the weight half is loaded +# before griddepcontrol.wait (the first RING stages; the rest prefetched into L2), the x half after it, each refilled +# once the stage's MMAs are done. +# Split-K reduction by token: rank r owns tokens [r T, r T + T) of the tile (T = N_TILE / SPLIT rounded up to 4) for +# all 128 rows. Every epilogue thread (one row) sends each other owner its partials of that owner's tokens by st.async +# (16-byte chunks of 4 tokens, chunk j at j ^ f(lane) so a quarter-warp's 8 rows fall in 8 bank groups), writes its +# own into its own slot, then sums its row's T tokens over the SPLIT slots in rank order: every token's sums are the +# long kernel's (same k-tile order per rank, same rank order). Each warp stages its 32 rows as [token][row] and stores +# them with one TMA store (tokens >= M clipped; the block holding row N - 1 when N is not a multiple of 32 has its own +# box): bf16 in the x ring (free once the MMAs are done; rows >= sig_row0 as bf16(sigmoid(bf16(acc)))), or fp32 +# (OUT_FP32: the MoE head's router logits) in place of its own slot. +# ============================================================================= +WIDE_TILES = (16, 32, 64) +WIDE_SMEM_BYTES = 220 * 1024 # the ring, the x stages and the mailbox; the barriers fit in the rest + + +def wide_tile(num_tokens: int) -> int: + """The token tile (MMA N) that holds ``num_tokens`` columns.""" + for n_tile in WIDE_TILES: + if num_tokens <= n_tile: + return n_tile + raise ValueError(f"k3_ctm_gemv_wide: {num_tokens} tokens exceed {WIDE_TILES[-1]}") + + +def wide_owned_tokens(n_tile: int, split: int) -> int: + """Tokens of the tile a rank reduces and stores (the last owner may have fewer, later ranks none).""" + return (n_tile + split - 1) // split + 3 & ~3 + + +def wide_smem_bytes(split: int, ring: int, n_tile: int, x_ring: int) -> int: + return ( + ring * CTA_M + x_ring * n_tile + ) * CTA_K * ELEM_BYTES + split * CTA_M * wide_owned_tokens(n_tile, split) * 4 + + +def wide_supports(n_out: int, k_in: int, split: int, ring: int, n_tile: int, x_ring: int) -> bool: + """K in whole 64-column halves (a last half k-tile reads zeros past K, as in the long kernel); N a multiple of 8 + (the output rows of a token are a whole number of 16-byte units for the TMA store).""" + return ( + k_in % TMA_K_BOX == 0 + and n_out % 8 == 0 + and n_tile in WIDE_TILES + and (split == 2 or 4 <= split <= 8) + and 1 <= ring <= num_k_tiles(k_in) // split + and 1 <= x_ring <= ring + and wide_smem_bytes(split, ring, n_tile, x_ring) <= WIDE_SMEM_BYTES + and n_out > 0 + ) + + +def _chunk_swizzle(lane, chunks: int): + """XOR for the 16-byte chunk index of a row of ``chunks`` chunks: the 8 rows of a quarter-warp in 8 bank groups.""" + if chunks == 2: + return (lane >> cutlass.Int32(2)) & cutlass.Int32(1) + if chunks == 4: + return (lane >> cutlass.Int32(1)) & cutlass.Int32(3) + if chunks >= 8: + return lane & cutlass.Int32(7) + return cutlass.Int32(0) # 1 or 3 chunks: rows already in distinct bank groups + + +@cute.kernel +def k3_ctm_gemv_wide_kernel( + tma_desc_w: cutlass.GridConstant[cuda.TensorMap], # W [N, K] bf16, 5-D, one call per k-tile + tma_desc_x: cutlass.GridConstant[cuda.TensorMap], # x [M, K] bf16, box 64 x n_tile + tma_desc_y: cutlass.GridConstant[ + cuda.TensorMap + ], # y [M, N] bf16 (fp32 with out_fp32), box 32 x T + tma_desc_y_tail: cutlass.GridConstant[cuda.TensorMap], # the same, box (N % 32 or 32) x T + sig_row0: cutlass.Int32, # bf16 output: rows >= sig_row0 store bf16(sigmoid(bf16(acc))); n_out: none + n_out: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + split: cutlass.Constexpr[int], + ring: cutlass.Constexpr[int], + x_ring: cutlass.Constexpr[int], + n_tile: cutlass.Constexpr[int], + out_fp32: cutlass.Constexpr[bool], + trigger_early: cutlass.Constexpr[bool], +): + """One 128-row tile per cluster of `split` CTAs; a weight ring filled before the grid wait, an x ring after it.""" + k_tiles = num_k_tiles(k_in) + extra = k_tiles % split + tmem_cols = max(TMEM_COLS, n_tile) + b_stage = n_tile * CTA_K # elements of one x k-tile + b_half = n_tile * TMA_K_BOX + tail_rows = n_out % LONG_ROWS_PER_WARP or LONG_ROWS_PER_WARP + owned = wide_owned_tokens(n_tile, split) # T + slot_elems = CTA_M * owned # one source rank's partials of the owner's tokens: [row][T] + tx, _, _ = cute.arch.thread_idx() + bx, _, _ = cute.arch.block_idx() + warp_id = cute.arch.warp_idx() + rank = cute.arch.block_idx_in_cluster() + tma_ptr_w = tma_desc_w.get_ptr() + tma_ptr_x = tma_desc_x.get_ptr() + tma_ptr_y = tma_desc_y.get_ptr() + tma_ptr_y_tail = tma_desc_y_tail.get_ptr() + m_offset = (bx // cutlass.Int32(split)) * cutlass.Int32(CTA_M) + my_tiles = cutlass.Int32(k_tiles // split) + if cutlass.const_expr(extra > 0): + if rank < cutlass.Int32(extra): + my_tiles = my_tiles + cutlass.Int32(1) + # The tokens this rank owns: [rank T, rank T + my_owned). + my_owned = cutlass.Int32(n_tile) - rank * cutlass.Int32(owned) + if my_owned > cutlass.Int32(owned): + my_owned = cutlass.Int32(owned) + + smem_a = cutlass.Array( + io_dtype, ring * CTA_M * CTA_K, space=cutlass.AddressSpace.smem, alignment=1024 + ) + smem_b = cutlass.Array( + io_dtype, x_ring * b_stage, space=cutlass.AddressSpace.smem, alignment=1024 + ) + a_full = cutlass.Array(cutlass.Int64, ring, space=cutlass.AddressSpace.smem, alignment=8) + b_full = cutlass.Array(cutlass.Int64, x_ring, space=cutlass.AddressSpace.smem, alignment=8) + mma_done = cutlass.Array(cutlass.Int64, ring, space=cutlass.AddressSpace.smem, alignment=8) + x_done = cutlass.Array(cutlass.Int64, x_ring, space=cutlass.AddressSpace.smem, alignment=8) + acc_done = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + mail_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + tmem_ptr_i32 = cutlass.Array(cutlass.Int32, 1, space=cutlass.AddressSpace.smem) + # [source rank][row][T tokens, chunks swizzled] fp32 partials of this rank's tokens (its own slot written locally). + mailbox = cutlass.Array( + cutlass.Float32, split * slot_elems, space=cutlass.AddressSpace.smem, alignment=128 + ) + + if warp_id == 0: + prims.prefetch_tensormap(tma_ptr_w) + prims.prefetch_tensormap(tma_ptr_x) + prims.prefetch_tensormap(tma_ptr_y) + prims.prefetch_tensormap(tma_ptr_y_tail) + if prims.elect_sync(): + for s in cutlass.range_constexpr(ring): + prims.mbarrier_init(a_full.subview(s), 1) + prims.mbarrier_init(mma_done.subview(s), 1) + for s in cutlass.range_constexpr(x_ring): + prims.mbarrier_init(b_full.subview(s), 1) + prims.mbarrier_init(x_done.subview(s), 1) + prims.mbarrier_init(acc_done, 1) + # The other ranks' partials of this rank's tokens arrive by st.async, completing the transaction count + # expected here, before cluster formation (a rank owning no tokens never waits on it). + prims.mbarrier_init(mail_full, 1) + if my_owned > cutlass.Int32(0): + prims.mbarrier_arrive_expect_tx( + mail_full, cutlass.Int32((split - 1) * CTA_M * 4) * my_owned + ) + if warp_id == 2: + prims.tcgen05_alloc(tmem_ptr_i32, tmem_cols) + prims.tcgen05_relinquish_alloc_permit() + prims.fence_mbarrier_init() + # Cluster formation: the peers' shared memory and barriers are addressable from here on. + prims.barrier_cluster_arrive_relaxed() + prims.barrier_cluster_wait() + prims.barrier_cta_sync(0) + tmem_ptr = cutlass.inttoptr(tmem_ptr_i32.load(), 6, cutlass.Int32) + + if warp_id == 0: + # ===================================================================== + # Weight TMA: fill the ring and prefetch the rest into L2 before the + # grid dependency; then refill each stage when its MMAs are done. + # ===================================================================== + if prims.elect_sync(): + for i in cutlass.range_constexpr(ring): + k = rank + cutlass.Int32(i * split) + prims.mbarrier_arrive_expect_tx(a_full.subview(i), CTA_M * CTA_K * ELEM_BYTES) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a.subview(i * CTA_M * CTA_K), + tma_ptr_w, + ( + cutlass.Int32(0), + m_offset, + k * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), + cutlass.Int32(0), + ), + a_full.subview(i), + l2_cache_hint=EVICT_FIRST, + ) + for i in range(ring, my_tiles): + k = rank + i * cutlass.Int32(split) + prims.cp_async_bulk_tensor_prefetch( + tma_ptr_w, + [ + cutlass.Int32(0), + m_offset, + k * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), + cutlass.Int32(0), + ], + [], # tile mode: no im2col offsets + ) + if cutlass.const_expr(trigger_early): + # Dependents may launch now; they wait for this whole grid before reading y. + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + if prims.elect_sync(): + stage = cutlass.Int32(0) + phase = cutlass.Int32(0) + for i in range(ring, my_tiles): + while not cute.arch.mbarrier_try_wait(mma_done.subview(stage).data_ptr(), phase): + pass + k = rank + i * cutlass.Int32(split) + prims.mbarrier_arrive_expect_tx(a_full.subview(stage), CTA_M * CTA_K * ELEM_BYTES) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a.subview(stage * cutlass.Int32(CTA_M * CTA_K)), + tma_ptr_w, + ( + cutlass.Int32(0), + m_offset, + k * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), + cutlass.Int32(0), + ), + a_full.subview(stage), + l2_cache_hint=EVICT_FIRST, + ) + stage = stage + cutlass.Int32(1) + if stage == cutlass.Int32(ring): + stage = cutlass.Int32(0) + phase = phase ^ cutlass.Int32(1) + elif warp_id == 1: + # ===================================================================== + # Activation TMA after the wait: the x k-tile of each weight k-tile + # into the x ring, refilled once the stage's MMAs are done. + # ===================================================================== + prims.griddepcontrol(prims.GridDepAction.WAIT) + if prims.elect_sync(): + for i in cutlass.range_constexpr(x_ring): + k = rank + cutlass.Int32(i * split) + prims.mbarrier_arrive_expect_tx(b_full.subview(i), b_stage * ELEM_BYTES) + for half in cutlass.range_constexpr(TMA_COPY_ITERS): + prims.cp_async_bulk_tensor_shared_cta_global( + smem_b.subview(i * b_stage + half * b_half), + tma_ptr_x, + ( + k * cutlass.Int32(CTA_K) + cutlass.Int32(half * TMA_K_BOX), + cutlass.Int32(0), + ), + b_full.subview(i), + ) + stage = cutlass.Int32(0) + phase = cutlass.Int32(0) + for i in range(x_ring, my_tiles): + while not cute.arch.mbarrier_try_wait(x_done.subview(stage).data_ptr(), phase): + pass + k = rank + i * cutlass.Int32(split) + prims.mbarrier_arrive_expect_tx(b_full.subview(stage), b_stage * ELEM_BYTES) + for half in cutlass.range_constexpr(TMA_COPY_ITERS): + prims.cp_async_bulk_tensor_shared_cta_global( + smem_b.subview( + stage * cutlass.Int32(b_stage) + cutlass.Int32(half * b_half) + ), + tma_ptr_x, + ( + k * cutlass.Int32(CTA_K) + cutlass.Int32(half * TMA_K_BOX), + cutlass.Int32(0), + ), + b_full.subview(stage), + ) + stage = stage + cutlass.Int32(1) + if stage == cutlass.Int32(x_ring): + stage = cutlass.Int32(0) + phase = phase ^ cutlass.Int32(1) + elif warp_id == 2: + # ===================================================================== + # MMA: the rank's k-tiles in order, one TMEM accumulator (partial sum) + # of N = n_tile token columns. + # ===================================================================== + idesc = prims.Tcgen05InstrDesc.build( + c_dtype=cutlass.Float32, a_dtype=io_dtype, b_dtype=io_dtype, n_dim=n_tile, m_dim=CTA_M + ) + desc_a_base = prims.Tcgen05SmemDesc.build( + start_address=smem_a, leading_byte_offset=LEADING, stride_byte_offset=STRIDE, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + desc_b_base = prims.Tcgen05SmemDesc.build( + start_address=smem_b, leading_byte_offset=LEADING, stride_byte_offset=STRIDE, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + stage = cutlass.Int32(0) + phase = cutlass.Int32(0) + x_stage = cutlass.Int32(0) + x_phase = cutlass.Int32(0) + for i in range(my_tiles): + while not cute.arch.mbarrier_try_wait(a_full.subview(stage).data_ptr(), phase): + pass + while not cute.arch.mbarrier_try_wait(b_full.subview(x_stage).data_ptr(), x_phase): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + for kb in cutlass.range_constexpr(TMA_COPY_ITERS * K_BLOCKS_PER_HALF): + box = kb // K_BLOCKS_PER_HALF + within = kb % K_BLOCKS_PER_HALF + desc_a = desc_a_base + ( + stage * cutlass.Int32(STAGE_A) + cutlass.Int32(box * A_BOX + within * STEP) + ) + desc_b = desc_b_base + ( + x_stage * cutlass.Int32(b_stage >> 3) + + cutlass.Int32(box * (b_half >> 3) + within * STEP) + ) + accumulate = cutlass.Boolean(True) + if cutlass.const_expr(kb == 0): + accumulate = i > cutlass.Int32(0) + if prims.elect_sync(): + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, + prims.CTAGroup.CTA_1, + tmem_ptr, + desc_a, + desc_b, + idesc, + accumulate, + ) + if prims.elect_sync(): + prims.tcgen05_commit(mma_done.subview(stage)) + prims.tcgen05_commit(x_done.subview(x_stage)) + stage = stage + cutlass.Int32(1) + if stage == cutlass.Int32(ring): + stage = cutlass.Int32(0) + phase = phase ^ cutlass.Int32(1) + x_stage = x_stage + cutlass.Int32(1) + if x_stage == cutlass.Int32(x_ring): + x_stage = cutlass.Int32(0) + x_phase = x_phase ^ cutlass.Int32(1) + if prims.elect_sync(): + prims.tcgen05_commit(acc_done) + elif warp_id >= 4: + # ===================================================================== + # Epilogue: TMEM -> registers; partials of each owner's tokens to that + # owner (own into the own slot); sum this rank's tokens in rank order; + # stage the warp's rows, one TMA store. + # ===================================================================== + lane = tx % 32 + w = warp_id - 4 # TMEM lanes 32 w .. 32 w + 31: tile rows 32 w + lane + row = w * cutlass.Int32(LONG_ROWS_PER_WARP) + lane + while not cute.arch.mbarrier_try_wait(acc_done.data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc = prims.tcgen05_ld( + "32x32b", cutlass.inttoptr(tmem_ptr_i32.load(), 6, cutlass.Float32), num=n_tile + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + swz = _chunk_swizzle(lane, owned // 4) + row_slot = row * cutlass.Int32(owned) + peer_slots = [] + peer_mbars = [] + for r in cutlass.range_constexpr(split): + peer_slots.append( + _mapa_u32( + mailbox.subview(rank * cutlass.Int32(slot_elems) + row_slot).data_ptr(), + cutlass.Int32(r), + ) + ) + peer_mbars.append(_mapa_u32(mail_full.data_ptr(), cutlass.Int32(r))) + # Chunk j for every owner in turn, so the owners' incoming traffic interleaves; the own chunks to the own slot. + for j in cutlass.range_constexpr(owned // 4): + for r in cutlass.range_constexpr(split): + if cutlass.const_expr(r * owned + 4 * j < n_tile): + t = r * owned + 4 * j + if rank == cutlass.Int32(r): + own = ( + cutlass.Int32(r * slot_elems) + + row_slot + + ((cutlass.Int32(j) ^ swz) << cutlass.Int32(2)) + ) + mailbox.store( + cutlass.Vector.from_elements( + tuple(cutlass.Float32(acc[t + e]) for e in range(4)), cutlass.Float32 + ), + idx=own, + vector_size=4, + alignment=16, + ) # fmt: skip + else: + _st_async_v4( + peer_slots[r] + ((cutlass.Int32(j) ^ swz) << cutlass.Int32(4)), + acc[t], acc[t + 1], acc[t + 2], acc[t + 3], peer_mbars[r], + ) # fmt: skip + if my_owned > cutlass.Int32(0): + # test_wait spin: a warp suspended in try_wait on a barrier completed by peers wakes late. + # Acquire at cluster scope: the bytes come from other CTAs' st.async (complete_tx releases at + # cluster scope). + while not _test_wait_cluster(mail_full.data_ptr(), 0): + pass + # This row's owned tokens over the SPLIT slots in rank order, as the long kernel. + totals = [cutlass.Float32(0.0)] * owned + for q in cutlass.range_constexpr(split): + for j in cutlass.range_constexpr(owned // 4): + part = mailbox.load( + idx=cutlass.Int32(q * slot_elems) + row_slot + ((cutlass.Int32(j) ^ swz) << cutlass.Int32(2)), + vector_size=4, + alignment=16, + ) # fmt: skip + for e in cutlass.range_constexpr(4): + totals[4 * j + e] = totals[4 * j + e] + cutlass.Float32(part[e]) + # [token][rows] staging of the warp's rows below N, one TMA store (tokens >= M clipped). + n0 = m_offset + w * cutlass.Int32(LONG_ROWS_PER_WARP) + tail_block = n0 + cutlass.Int32(LONG_ROWS_PER_WARP) > cutlass.Int32(n_out) + pitch = cutlass.Int32(LONG_ROWS_PER_WARP) + if tail_block: + pitch = cutlass.Int32(tail_rows) + t0 = rank * cutlass.Int32(owned) + if cutlass.const_expr(out_fp32): + # In place of this rank's own slot, rows of this warp only (read above by this warp alone). + staged = mailbox.subview( + rank * cutlass.Int32(slot_elems) + w * cutlass.Int32(LONG_ROWS_PER_WARP * owned) + ) + cute.arch.sync_warp() + else: + staged = smem_b.subview(w * cutlass.Int32(LONG_ROWS_PER_WARP * owned)) + if n0 < cutlass.Int32(n_out): + if lane < pitch: + if cutlass.const_expr(out_fp32): + for i in cutlass.range_constexpr(owned): + staged.store(totals[i], idx=cutlass.Int32(i) * pitch + lane) + else: + if n0 + lane >= sig_row0: + for i in cutlass.range_constexpr(owned): + staged.store( + _sigmoid_bf16(_bf16_rn(totals[i])).to(io_dtype), idx=cutlass.Int32(i) * pitch + lane + ) # fmt: skip + else: + for i in cutlass.range_constexpr(owned): + staged.store( + totals[i].to(io_dtype), idx=cutlass.Int32(i) * pitch + lane + ) + # Generic-proxy writes of the staged rows -> the TMA store's async-proxy reads. + prims.fence_proxy(prims.Proxy.ASYNC_SHARED, space=prims.SharedSpace.shared_cta) + cute.arch.sync_warp() + if lane == 0: + if tail_block: + prims.cp_async_bulk_tensor_global_shared_cta( + tma_ptr_y_tail, staged, [n0, t0] + ) + else: + prims.cp_async_bulk_tensor_global_shared_cta(tma_ptr_y, staged, [n0, t0]) + prims.cp_async_bulk_commit_group() + prims.cp_async_bulk_wait_group(0, read=True) + prims.barrier_cta_sync(1, thread_count=EPI_THREADS) + if warp_id == 4: + prims.tcgen05_dealloc(tmem_ptr, tmem_cols) + + +def _activation_tensor_map_rows(x, cols, num_tokens, rows): + """x [M, cols] (rows dense) as (cols, M) with a ``rows``-row box: rows past num_tokens arrive as zeros.""" + return cuda.create_tensor_map_tiled( + global_address=x.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[cols, num_tokens], + global_strides=[(cols * ELEM_BYTES) // 16], + box_dims=[TMA_K_BOX, rows], + swizzle=cuda.TensorMapSwizzle.s128b, + ) + + +def _output_tensor_map(y, n_out, num_tokens, tokens, out_fp32, rows=LONG_ROWS_PER_WARP): + """y [M, N] (rows dense) as (N, M) with a rows x tokens box, no swizzle: a warp's staged [token][rows].""" + elem_bytes = 4 if out_fp32 else ELEM_BYTES + return cuda.create_tensor_map_tiled( + global_address=y.iterator.toint(), + dtype=cutlass.Float32 if out_fp32 else cutlass.BFloat16, + global_dims=[n_out, num_tokens], + global_strides=[(n_out * elem_bytes) // 16], + box_dims=[rows, tokens], + swizzle=cuda.TensorMapSwizzle.none, + ) + + +@cute.jit +def k3_ctm_gemv_wide( + w: cute.Tensor, # [N, K] bf16, K contiguous + x: cute.Tensor, # [M, K] bf16, K contiguous, M <= n_tile + y: cute.Tensor, # [M * N] bf16 (fp32 with out_fp32) + num_tokens: cutlass.Int32, + sig_row0: cutlass.Int32, + n_out: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + split: cutlass.Constexpr[int], + ring: cutlass.Constexpr[int], + x_ring: cutlass.Constexpr[int], + n_tile: cutlass.Constexpr[int], + out_fp32: cutlass.Constexpr[bool], + trigger_early: cutlass.Constexpr[bool], + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + """``y = x @ w^T`` for up to ``n_tile`` tokens; bf16 output rows >= sig_row0 hold bf16(sigmoid(bf16(.))).""" + owned = wide_owned_tokens(n_tile, split) + tma_desc_w = _weight_tensor_map(w, n_out, k_in) + tma_desc_x = _activation_tensor_map_rows(x, k_in, num_tokens, n_tile) + tma_desc_y = _output_tensor_map(y, n_out, num_tokens, owned, out_fp32) + tma_desc_y_tail = _output_tensor_map( + y, n_out, num_tokens, owned, out_fp32, n_out % LONG_ROWS_PER_WARP or LONG_ROWS_PER_WARP + ) + k3_ctm_gemv_wide_kernel( + tma_desc_w, tma_desc_x, tma_desc_y, tma_desc_y_tail, sig_row0, n_out, k_in, split, ring, x_ring, n_tile, + out_fp32, trigger_early, + ).launch( + grid=(((n_out + CTA_M - 1) // CTA_M) * split, 1, 1), + block=(THREADS, 1, 1), + cluster=(split, 1, 1), + stream=stream, + use_pdl=use_pdl, + ) # fmt: skip + + +# ============================================================================= +# SiTU-and-mul with PDL: a = bf16(situ(g) * situ_lin(u)) for a gate_up output [M, 2K] (g its first K columns), in +# SituAndMul's fp32 order (beta tanh(g / beta) sigmoid(g), times linear_beta tanh(u / linear_beta) or u), one 16-byte +# vector of 8 columns per thread. The dependents launch at once, so a following projection streams its weights while +# this kernel and its predecessor run; the activation reads wait for the grid. +# ============================================================================= +SITU_THREADS = 128 +SITU_COLS_PER_CTA = SITU_THREADS * VEC + + +def _situ_mul(g, u, beta, linear_beta, has_linear: bool): + """bf16(situ(g) * situ_lin(u)) in fp32, SituAndMul's order.""" + one = cutlass.Float32(1.0) + sig = _div_rn(one, one + cute.math.exp(-g, fastmath=False)) + a = beta * cute.math.tanh(_div_rn(g, beta), fastmath=False) * sig + if cutlass.const_expr(has_linear): + u = linear_beta * cute.math.tanh(_div_rn(u, linear_beta), fastmath=False) + return (a * u).to(io_dtype) + + +@cute.kernel +def k3_situ_mul_kernel( + gu: cutlass.Array, # bf16 [M * 2K] + out: cutlass.Array, # bf16 [M * K] + num_tokens: cutlass.Int32, + beta: cutlass.Float32, + linear_beta: cutlass.Float32, + k_in: cutlass.Constexpr[int], + has_linear: cutlass.Constexpr[bool], +): + tx, _, _ = cute.arch.thread_idx() + bx, by, _ = cute.arch.block_idx() + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + prims.griddepcontrol(prims.GridDepAction.WAIT) + col = bx * cutlass.Int32(SITU_COLS_PER_CTA) + tx * cutlass.Int32(VEC) + if col < cutlass.Int32(k_in): + if by < num_tokens: + row = by * cutlass.Int32(2 * k_in) + gv = gu.load(idx=row + col, vector_size=VEC, alignment=16) + uv = gu.load(idx=row + cutlass.Int32(k_in) + col, vector_size=VEC, alignment=16) + outs = [ + _situ_mul( + cutlass.Float32(gv[e]), cutlass.Float32(uv[e]), beta, linear_beta, has_linear + ) + for e in range(VEC) + ] + out.store( + cutlass.Vector.from_elements(tuple(outs), io_dtype), + idx=by * cutlass.Int32(k_in) + col, vector_size=VEC, alignment=16, + ) # fmt: skip + + +@cute.jit +def k3_situ_mul( + gu: cute.Tensor, # bf16 [M * 2K] + out: cute.Tensor, # bf16 [M * K] + num_tokens: cutlass.Int32, + beta: cutlass.Float32, + linear_beta: cutlass.Float32, + max_tokens: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + has_linear: cutlass.Constexpr[bool], + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + """``out = SituAndMul(beta, linear_beta)(gu)`` for rows < num_tokens (grid over max_tokens rows).""" + k3_situ_mul_kernel(gu, out, num_tokens, beta, linear_beta, k_in, has_linear).launch( + grid=((k_in + SITU_COLS_PER_CTA - 1) // SITU_COLS_PER_CTA, max_tokens, 1), + block=(SITU_THREADS, 1, 1), + stream=stream, + use_pdl=use_pdl, + ) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_ctm_gemv/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_ctm_gemv/op.py new file mode 100644 index 000000000000..97ff2483e831 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_ctm_gemv/op.py @@ -0,0 +1,544 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Torch ops of the CTM decode GEMV (M <= 8, bf16, fp32 accumulation, K <= 768). + +``trtllm::k3_ctm_gemv``: ``x @ weight^T`` (same results as ``trtllm::k3_decode_gemv`` with ``split=1``). +``trtllm::k3_ctm_gemv_tail``: the row-parallel MoE tail (same results as ``trtllm::k3_decode_gemv_tail``). +``trtllm::k3_ctm_gemv_gated``: ``(a * sigmoid(g)) @ weight^T`` with torch's bf16 roundings, where ``g`` is a +column window of a dense bf16 matrix (the MLA output gate inside the fused q_a/kv_a/gate projection output), or +``a * s`` when that window already holds ``s = bf16(sigmoid(g))``. +``trtllm::k3_ctm_gemv_long``: ``x @ weight^T`` for long K (split-K over a 4-8 CTA cluster, weight ring and L2 +prefetch before the grid-dependency wait); output columns >= ``sig_col0`` hold ``bf16(sigmoid(bf16(.)))``. +``trtllm::k3_ctm_gemv_swiglu``: ``silu_and_mul(gu) @ weight^T`` with the activation in the B prologue. +``trtllm::k3_ctm_gemv_wide``: ``x @ weight^T`` for up to 64 tokens in one MMA of 16, 32 or 64 token columns, on the +long kernel's split-K clusters (bf16 output with optional sigmoid columns, or fp32 output). + +``push`` (the split-K ops) chooses how a cluster's ranks send their fp32 partials to the row owner: DSMEM +stores + a release arrive (False) or 16-byte st.async completing the owner's barrier by bytes (True). The +sums and their order are the same; each call site takes the one measured faster at its shape. + +Each kernel is compiled on the first call for its shape and flags, which must happen outside CUDA-graph +capture. The whole weight slice is loaded before the grid-dependency wait. +""" + +from __future__ import annotations + +import os +import threading +from typing import Dict, Optional + +import torch + +MAX_TOKENS = 8 +WIDE_MAX_TOKENS = 64 + +_lock = threading.Lock() +_compiled: Dict[tuple, object] = {} + + +def _arg(t: torch.Tensor): + from cutlass.cute.runtime import from_dlpack + + # detach(): DLPack refuses tensors that require grad (weights are parameters). + return from_dlpack(t.detach(), assumed_align=16).mark_layout_dynamic(leading_dim=t.dim() - 1) + + +def _use_pdl() -> bool: + return os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" + + +def _stream(t: torch.Tensor): + import cuda.bindings.driver as cuda_driver + + return cuda_driver.CUstream(torch.cuda.current_stream(t.device).cuda_stream) + + +def _compile(key, entry, *args): + fn = _compiled.get(key) + if fn is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + f"{key[0]}: run once per shape outside CUDA-graph capture first (it compiles its kernel)." + ) + import cutlass.cute as cute + + with _lock: + fn = _compiled.get(key) + if fn is None: + fn = _compiled[key] = cute.compile(entry, *args) + return fn + + +def supports(x: torch.Tensor, weight: torch.Tensor, split: int = 1) -> bool: + """Whether ``k3_ctm_gemv`` runs ``x @ weight^T``.""" + from . import k3_ctm_gemv_kernel as kernel + + return ( + x.is_cuda + and x.dtype == torch.bfloat16 + and weight.dtype == torch.bfloat16 + and x.dim() == 2 + and weight.dim() == 2 + and 0 < x.shape[0] <= MAX_TOKENS + and x.shape[1] == weight.shape[1] + and x.is_contiguous() + and weight.is_contiguous() + and kernel.supports(weight.shape[0], weight.shape[1], split) + ) + + +@torch.library.custom_op("trtllm::k3_ctm_gemv", mutates_args=()) +def k3_ctm_gemv( + x: torch.Tensor, + weight: torch.Tensor, + trigger_early: bool = True, + split: int = 1, + push: bool = False, +) -> torch.Tensor: + """``x @ weight.T`` for bf16 ``x`` [M <= 8, K <= 768] and ``weight`` [N, K]; returns bf16 [M, N]. + + ``split=2`` splits the k-tiles of each 128-row tile over a 2-CTA cluster (twice the CTAs streaming the + weight; the two fp32 partials are added before the bf16 rounding).""" + if not supports(x, weight, split): + raise ValueError( + f"k3_ctm_gemv: unsupported call x {tuple(x.shape)} {x.dtype}, weight {tuple(weight.shape)} " + f"{weight.dtype}, split {split}" + ) + from . import k3_ctm_gemv_kernel as kernel + + num_tokens, k_in = x.shape + n_out = weight.shape[0] + y = torch.empty(num_tokens, n_out, dtype=torch.bfloat16, device=x.device) + args = (_arg(weight), _arg(x), _arg(y.view(-1))) + stream = _stream(x) + use_pdl = _use_pdl() + fn = _compile( + ("k3_ctm_gemv", n_out, k_in, split, trigger_early, push, use_pdl), + kernel.k3_ctm_gemv, *args, num_tokens, n_out, k_in, split, trigger_early, push, use_pdl, stream, + ) # fmt: skip + fn(*args, num_tokens, stream) + return y + + +@k3_ctm_gemv.register_fake +def _(x, weight, trigger_early=True, split=1, push=False): + return x.new_empty((x.shape[0], weight.shape[0]), dtype=torch.bfloat16) + + +def supports_tail( + latent: torch.Tensor, act: torch.Tensor, weight: torch.Tensor, width: int +) -> bool: + """Whether ``k3_ctm_gemv_tail`` runs the call (the conditions of ``k3_decode_gemv``'s tail).""" + from . import k3_ctm_gemv_kernel as kernel + + k_act = act.shape[1] if act.dim() == 2 else -1 + k_lat = weight.shape[1] - k_act + return ( + latent.is_cuda + and latent.dtype == act.dtype == weight.dtype == torch.bfloat16 + and latent.dim() == 2 + and act.dim() == 2 + and weight.dim() == 2 + and 0 < latent.shape[0] <= MAX_TOKENS + and act.shape[0] == latent.shape[0] + and latent.is_contiguous() + and act.is_contiguous() + and weight.is_contiguous() + and k_lat % kernel.CTA_K == 0 + and k_act % kernel.CTA_K == 0 + and 0 < width <= k_lat + and latent.shape[1] % (8 * 32) == 0 + and kernel.supports(weight.shape[0], weight.shape[1]) + ) + + +@torch.library.custom_op("trtllm::k3_ctm_gemv_tail", mutates_args=()) +def k3_ctm_gemv_tail( + latent: torch.Tensor, + act: torch.Tensor, + weight: torch.Tensor, + lo: int, + width: int, + eps: float, + trigger_early: bool = True, +) -> torch.Tensor: + """``[rmsnorm(latent)[:, lo:lo+width] | act] @ weight.T`` with the RMS applied to the fp32 latent + accumulator (``weight`` = [latent-up columns of the slice, zero-padded to 128 | shared down]).""" + if not supports_tail(latent, act, weight, width): + raise ValueError( + f"k3_ctm_gemv_tail: unsupported call latent {tuple(latent.shape)}, act {tuple(act.shape)}, " + f"weight {tuple(weight.shape)}, width {width}" + ) + from . import k3_ctm_gemv_kernel as kernel + + num_tokens, rms_cols = latent.shape + n_out, k_in = weight.shape + lat_tiles = (k_in - act.shape[1]) // kernel.CTA_K + y = torch.empty(num_tokens, n_out, dtype=torch.bfloat16, device=latent.device) + args = ( + _arg(weight), + _arg(latent), + _arg(latent.view(-1).view(torch.int32)), + _arg(act), + _arg(y.view(-1)), + ) + stream = _stream(latent) + use_pdl = _use_pdl() + fn = _compile( + ("k3_ctm_gemv_tail", n_out, k_in, lat_tiles, rms_cols, trigger_early, use_pdl), + kernel.k3_ctm_gemv_tail, *args, num_tokens, lo, float(eps), n_out, k_in, lat_tiles, rms_cols, + trigger_early, use_pdl, stream, + ) # fmt: skip + fn(*args, num_tokens, lo, float(eps), stream) + return y + + +@k3_ctm_gemv_tail.register_fake +def _(latent, act, weight, lo, width, eps, trigger_early=True): + return latent.new_empty((latent.shape[0], weight.shape[0]), dtype=torch.bfloat16) + + +def supports_gated( + a: torch.Tensor, gsrc: torch.Tensor, g_col0: int, weight: torch.Tensor, split: int = 2 +) -> bool: + """Whether ``k3_ctm_gemv_gated`` runs ``(a * sigmoid(gsrc[:, g_col0 : g_col0 + K])) @ weight^T``.""" + from . import k3_ctm_gemv_kernel as kernel + + k_in = weight.shape[1] if weight.dim() == 2 else -1 + return ( + a.is_cuda + and a.dtype == gsrc.dtype == weight.dtype == torch.bfloat16 + and a.dim() == 2 + and gsrc.dim() == 2 + and weight.dim() == 2 + and 0 < a.shape[0] <= MAX_TOKENS + and gsrc.shape[0] == a.shape[0] + and a.shape[1] == k_in + and a.is_contiguous() + and gsrc.is_contiguous() + and weight.is_contiguous() + and g_col0 % 8 == 0 + and 0 <= g_col0 + and g_col0 + k_in <= gsrc.shape[1] + and gsrc.shape[1] % 8 == 0 + and kernel.supports(weight.shape[0], k_in, split) + ) + + +@torch.library.custom_op("trtllm::k3_ctm_gemv_gated", mutates_args=()) +def k3_ctm_gemv_gated( + a: torch.Tensor, + gsrc: torch.Tensor, + g_col0: int, + weight: torch.Tensor, + trigger_early: bool = True, + split: int = 2, + gate_sigmoid: bool = True, +) -> torch.Tensor: + """``(a * gsrc[:, g_col0:g_col0 + K].sigmoid()) @ weight.T``: bf16 ``a`` [M <= 8, K], ``gsrc`` [M, C] with dense + rows, ``weight`` [N, K]; the sigmoid and the product each round to bf16 as the unfused torch ops do. With + ``gate_sigmoid=False`` the window already holds the sigmoid: ``(a * gsrc[:, g_col0:g_col0 + K]) @ weight.T``.""" + if not supports_gated(a, gsrc, g_col0, weight, split): + raise ValueError( + f"k3_ctm_gemv_gated: unsupported call a {tuple(a.shape)}, gsrc {tuple(gsrc.shape)} at {g_col0}, " + f"weight {tuple(weight.shape)}, split {split}" + ) + from . import k3_ctm_gemv_kernel as kernel + + num_tokens, k_in = a.shape + n_out = weight.shape[0] + gsrc_cols = gsrc.shape[1] + y = torch.empty(num_tokens, n_out, dtype=torch.bfloat16, device=a.device) + args = (_arg(weight), _arg(a), _arg(gsrc), _arg(y.view(-1))) + stream = _stream(a) + use_pdl = _use_pdl() + fn = _compile( + ("k3_ctm_gemv_gated", n_out, k_in, gsrc_cols, split, gate_sigmoid, trigger_early, use_pdl), + kernel.k3_ctm_gemv_gated, *args, num_tokens, g_col0, n_out, k_in, gsrc_cols, split, gate_sigmoid, + trigger_early, use_pdl, stream, + ) # fmt: skip + fn(*args, num_tokens, g_col0, stream) + return y + + +@k3_ctm_gemv_gated.register_fake +def _(a, gsrc, g_col0, weight, trigger_early=True, split=2, gate_sigmoid=True): + return a.new_empty((a.shape[0], weight.shape[0]), dtype=torch.bfloat16) + + +def supports_swiglu(gu: torch.Tensor, weight: torch.Tensor, split: int = 2) -> bool: + """Whether ``k3_ctm_gemv_swiglu`` runs ``(silu(gu[:, :K]) * gu[:, K:]) @ weight^T``.""" + from . import k3_ctm_gemv_kernel as kernel + + k_in = weight.shape[1] if weight.dim() == 2 else -1 + return ( + gu.is_cuda + and gu.dtype == weight.dtype == torch.bfloat16 + and gu.dim() == 2 + and weight.dim() == 2 + and 0 < gu.shape[0] <= MAX_TOKENS + and gu.shape[1] == 2 * k_in + and gu.is_contiguous() + and weight.is_contiguous() + and kernel.supports(weight.shape[0], k_in, split) + ) + + +@torch.library.custom_op("trtllm::k3_ctm_gemv_swiglu", mutates_args=()) +def k3_ctm_gemv_swiglu( + gu: torch.Tensor, + weight: torch.Tensor, + trigger_early: bool = True, + split: int = 2, + push: bool = False, +) -> torch.Tensor: + """``silu_and_mul(gu) @ weight.T`` for a bf16 gate_up output ``gu`` [M <= 8, 2 K] (gate columns first) and + ``weight`` [N, K]: the activation is computed in the GEMV's B prologue with silu_and_mul's fp32 arithmetic and + one bf16 rounding. ``split`` k-tile ranks per 128-row tile (need not divide the k-tiles).""" + if not supports_swiglu(gu, weight, split): + raise ValueError( + f"k3_ctm_gemv_swiglu: unsupported call gu {tuple(gu.shape)} {gu.dtype}, weight {tuple(weight.shape)} " + f"{weight.dtype}, split {split}" + ) + from . import k3_ctm_gemv_kernel as kernel + + num_tokens = gu.shape[0] + n_out, k_in = weight.shape + y = torch.empty(num_tokens, n_out, dtype=torch.bfloat16, device=gu.device) + args = (_arg(weight), _arg(gu), _arg(y.view(-1))) + stream = _stream(gu) + use_pdl = _use_pdl() + fn = _compile( + ("k3_ctm_gemv_swiglu", n_out, k_in, split, trigger_early, push, use_pdl), + kernel.k3_ctm_gemv_swiglu, *args, num_tokens, n_out, k_in, split, trigger_early, push, use_pdl, stream, + ) # fmt: skip + fn(*args, num_tokens, stream) + return y + + +@k3_ctm_gemv_swiglu.register_fake +def _(gu, weight, trigger_early=True, split=2, push=False): + return gu.new_empty((gu.shape[0], weight.shape[0]), dtype=torch.bfloat16) + + +def supports_long(x: torch.Tensor, weight: torch.Tensor, split: int, ring: int) -> bool: + """Whether ``k3_ctm_gemv_long`` runs ``x @ weight^T`` with this split and ring.""" + from . import k3_ctm_gemv_kernel as kernel + + return ( + x.is_cuda + and x.dtype == torch.bfloat16 + and weight.dtype == torch.bfloat16 + and x.dim() == 2 + and weight.dim() == 2 + and 0 < x.shape[0] <= MAX_TOKENS + and x.shape[1] == weight.shape[1] + and x.is_contiguous() + and weight.is_contiguous() + and kernel.long_supports(weight.shape[0], weight.shape[1], split, ring) + ) + + +@torch.library.custom_op("trtllm::k3_ctm_gemv_long", mutates_args=()) +def k3_ctm_gemv_long( + x: torch.Tensor, + weight: torch.Tensor, + sig_col0: int = -1, + split: int = 6, + ring: int = 5, + trigger_early: bool = True, + push: bool = False, +) -> torch.Tensor: + """``x @ weight.T`` for bf16 ``x`` [M <= 8, K] and ``weight`` [N, K] (long K); columns >= ``sig_col0`` (if >= 0) + hold ``bf16(sigmoid(bf16(x @ weight.T)))``. Each 128-row weight tile is split over a cluster of ``split`` CTAs, + each streaming its k-tiles through a ``ring``-stage ring filled before the grid-dependency wait.""" + return _launch_long(x, weight, sig_col0, split, ring, trigger_early, push=push) + + +def _launch_long(x, weight, sig_col0, split, ring, trigger_early, push=False) -> torch.Tensor: + if not supports_long(x, weight, split, ring): + raise ValueError( + f"k3_ctm_gemv_long: unsupported call x {tuple(x.shape)} {x.dtype}, weight {tuple(weight.shape)} " + f"{weight.dtype}, split {split}, ring {ring}" + ) + from . import k3_ctm_gemv_kernel as kernel + + num_tokens, k_in = x.shape + n_out = weight.shape[0] + y = torch.empty(num_tokens, n_out, dtype=torch.bfloat16, device=x.device) + args = (_arg(weight), _arg(x), _arg(y.view(-1))) + stream = _stream(x) + use_pdl = _use_pdl() + sig_row0 = sig_col0 if sig_col0 >= 0 else n_out + fn = _compile( + ("k3_ctm_gemv_long", n_out, k_in, split, ring, trigger_early, push, use_pdl), + kernel.k3_ctm_gemv_long, *args, num_tokens, sig_row0, n_out, k_in, split, ring, trigger_early, push, use_pdl, + stream, + ) # fmt: skip + fn(*args, num_tokens, sig_row0, stream) + return y + + +@k3_ctm_gemv_long.register_fake +def _(x, weight, sig_col0=-1, split=6, ring=5, trigger_early=True, push=False): + return x.new_empty((x.shape[0], weight.shape[0]), dtype=torch.bfloat16) + + +def wide_config(n_out: int, k_in: int, n_tile: int, num_sms: int) -> Optional[tuple]: + """``(split, ring, x_ring)`` of a ``k3_ctm_gemv_wide`` call, or None: the largest cluster whose CTAs fit one wave, + then the deepest weight ring that fits shared memory beside an x ring of up to 3 stages. At the call sites that run + ``k3_ctm_gemv_long`` at most 8 tokens (MLA [W_a; W_g] 6, dense gate_up 4, dense down 2, drafter 8) this is their + split, so a token's sums are the same.""" + from . import k3_ctm_gemv_kernel as kernel + + tiles = (n_out + kernel.CTA_M - 1) // kernel.CTA_M + for split in (8, 7, 6, 5, 4, 2): + if tiles * split > num_sms: + continue + for ring in range(kernel.num_k_tiles(k_in) // split, 0, -1): + if kernel.wide_supports(n_out, k_in, split, ring, n_tile, min(ring, 3)): + return split, ring, min(ring, 3) + return None + + +def _num_sms(device: torch.device) -> int: + return torch.cuda.get_device_properties(device).multi_processor_count + + +def supports_wide( + x: torch.Tensor, weight: torch.Tensor, sig_col0: int = -1, out_fp32: bool = False +) -> bool: + """Whether ``k3_ctm_gemv_wide`` runs ``x @ weight^T``.""" + from . import k3_ctm_gemv_kernel as kernel + + if not ( + x.is_cuda + and x.dtype == torch.bfloat16 + and weight.dtype == torch.bfloat16 + and x.dim() == 2 + and weight.dim() == 2 + and 0 < x.shape[0] <= WIDE_MAX_TOKENS + and x.shape[1] == weight.shape[1] + and x.is_contiguous() + and weight.is_contiguous() + and x.data_ptr() % 16 == 0 + and weight.data_ptr() % 16 == 0 + and sig_col0 < weight.shape[0] + and (sig_col0 < 0 or not out_fp32) + ): + return False + n_tile = kernel.wide_tile(x.shape[0]) + return wide_config(weight.shape[0], weight.shape[1], n_tile, _num_sms(x.device)) is not None + + +@torch.library.custom_op("trtllm::k3_ctm_gemv_wide", mutates_args=()) +def k3_ctm_gemv_wide( + x: torch.Tensor, weight: torch.Tensor, sig_col0: int = -1, out_fp32: bool = False +) -> torch.Tensor: + """``x @ weight.T`` for bf16 ``x`` [M <= 64, K] and ``weight`` [N, K], all M tokens in one MMA of N = 16, 32 or 64 + columns (the long kernel's split-K clusters and weight ring, the partials reduced by token; split and ring from + ``wide_config``). Returns bf16 + [M, N], whose columns >= ``sig_col0`` (if >= 0) hold ``bf16(sigmoid(bf16(x @ weight.T)))``, or fp32 [M, N] with + ``out_fp32``.""" + if not supports_wide(x, weight, sig_col0, out_fp32): + raise ValueError( + f"k3_ctm_gemv_wide: unsupported call x {tuple(x.shape)} {x.dtype}, weight {tuple(weight.shape)} " + f"{weight.dtype}, sig_col0 {sig_col0}, out_fp32 {out_fp32}" + ) + from . import k3_ctm_gemv_kernel as kernel + + n_tile = kernel.wide_tile(x.shape[0]) + split, ring, x_ring = wide_config(weight.shape[0], weight.shape[1], n_tile, _num_sms(x.device)) + return _launch_wide(x, weight, sig_col0, out_fp32, split, ring, x_ring) + + +def _launch_wide( + x, weight, sig_col0, out_fp32, split, ring, x_ring, trigger_early=True +) -> torch.Tensor: + from . import k3_ctm_gemv_kernel as kernel + + num_tokens, k_in = x.shape + n_out = weight.shape[0] + n_tile = kernel.wide_tile(num_tokens) + if not kernel.wide_supports(n_out, k_in, split, ring, n_tile, x_ring): + raise ValueError( + f"k3_ctm_gemv_wide: split {split}, rings {ring} / {x_ring} do not fit weight {tuple(weight.shape)} at " + f"{n_tile} tokens" + ) + y = torch.empty( + num_tokens, n_out, dtype=torch.float32 if out_fp32 else torch.bfloat16, device=x.device + ) + args = (_arg(weight), _arg(x), _arg(y.view(-1))) + stream = _stream(x) + use_pdl = _use_pdl() + sig_row0 = sig_col0 if sig_col0 >= 0 else n_out + fn = _compile( + ("k3_ctm_gemv_wide", n_out, k_in, split, ring, x_ring, n_tile, out_fp32, trigger_early, use_pdl), + kernel.k3_ctm_gemv_wide, *args, num_tokens, sig_row0, n_out, k_in, split, ring, x_ring, n_tile, out_fp32, + trigger_early, use_pdl, stream, + ) # fmt: skip + fn(*args, num_tokens, sig_row0, stream) + return y + + +@k3_ctm_gemv_wide.register_fake +def _(x, weight, sig_col0=-1, out_fp32=False): + return x.new_empty( + (x.shape[0], weight.shape[0]), dtype=torch.float32 if out_fp32 else torch.bfloat16 + ) + + +def supports_situ_mul(gu: torch.Tensor) -> bool: + """Whether ``k3_situ_mul`` runs SituAndMul on ``gu`` [M <= 8, 2 K] (K a multiple of 8, rows dense).""" + return ( + gu.is_cuda + and gu.dtype == torch.bfloat16 + and gu.dim() == 2 + and 0 < gu.shape[0] <= MAX_TOKENS + and gu.shape[1] % 16 == 0 + and gu.is_contiguous() + and gu.data_ptr() % 16 == 0 + ) + + +@torch.library.custom_op("trtllm::k3_situ_mul", mutates_args=()) +def k3_situ_mul( + gu: torch.Tensor, beta: float = 1.0, linear_beta: Optional[float] = None +) -> torch.Tensor: + """``SituAndMul(beta, linear_beta)(gu)`` for a bf16 gate_up output ``gu`` [M <= 8, 2 K] (gate columns first), with + programmatic dependent launch: the next kernel launches at once and may stream its weights meanwhile.""" + if not supports_situ_mul(gu): + raise ValueError(f"k3_situ_mul: unsupported call gu {tuple(gu.shape)} {gu.dtype}") + from . import k3_ctm_gemv_kernel as kernel + + num_tokens, two_k = gu.shape + k_in = two_k // 2 + out = torch.empty(num_tokens, k_in, dtype=torch.bfloat16, device=gu.device) + args = (_arg(gu.view(-1)), _arg(out.view(-1))) + stream = _stream(gu) + use_pdl = _use_pdl() + has_linear = linear_beta is not None + fn = _compile( + ("k3_situ_mul", k_in, has_linear, use_pdl), + kernel.k3_situ_mul, *args, num_tokens, float(beta), float(linear_beta or 1.0), MAX_TOKENS, k_in, has_linear, + use_pdl, stream, + ) # fmt: skip + fn(*args, num_tokens, float(beta), float(linear_beta or 1.0), stream) + return out + + +@k3_situ_mul.register_fake +def _(gu, beta=1.0, linear_beta=None): + return gu.new_empty((gu.shape[0], gu.shape[1] // 2)) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_decode_gemv/__init__.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_decode_gemv/__init__.py new file mode 100644 index 000000000000..cb3670f60918 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_decode_gemv/__init__.py @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Kimi K3 decode GEMV in CuTe DSL (``trtllm::k3_decode_gemv``). + +Importing :mod:`.op` registers the torch op; nothing is imported eagerly here so that +the CuTe DSL dependency stays optional for every other model. +""" diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_decode_gemv/k3_decode_gemv_kernel.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_decode_gemv/k3_decode_gemv_kernel.py new file mode 100644 index 000000000000..430dd42d7cd3 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_decode_gemv/k3_decode_gemv_kernel.py @@ -0,0 +1,599 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Decode GEMV for Kimi K3 projections: ``y[M, N] = x[M, K] @ W[N, K]^T``, M <= 8, bf16 in and out. + +tcgen05 with the weight as the MMA's A operand (128 output features per CTA) and the activation as B +(8 token columns), fp32 accumulation in TMEM. For the short-K projections (K <= 6 * 128: the MoE +tail, the attention o_proj) one CTA owns one 128-row weight tile over the whole K extent, and every +weight k-tile of that CTA gets its own shared-memory stage. The weight TMA for all stages is issued +at launch, before ``griddepcontrol.wait``: launched early under PDL, the kernel streams its whole +weight while its predecessor still runs, and after the wait only the activation (8 x K) remains to +load. The activation box is always 8 rows; rows past M are zero-filled by the TMA, so one compiled +kernel serves every M <= 8. + +Warps: 0 weight TMA, 1 activation TMA (after the grid dependency), 2 TMEM allocation and MMA, 3 idle, +4-7 epilogue (TMEM -> registers -> bf16 stores; warp 4 + i reads TMEM lanes 32i..32i+31). +""" + +from __future__ import annotations + +import cuda.bindings.driver as cuda_driver +import cutlass +import cutlass.cute as cute +import cutlass.experimental.cuda as cuda +from cutlass.experimental import primitives as prims + +CTA_M = 128 # weight rows (output features) per CTA = the MMA's M +MMA_N = 8 # token columns +CTA_K = 128 # one k-tile: two 64-element halves of the 128-byte swizzle +MMA_K = 16 +TMA_K_BOX = 64 +TMA_COPY_ITERS = CTA_K // TMA_K_BOX +K_BLOCKS_PER_HALF = TMA_K_BOX // MMA_K +MAX_K_TILES = 6 # 6 x (32 KB weight + 2 KB activation) of shared memory +THREADS = 256 +TMEM_COLS = 32 +ELEM_BYTES = 2 +EVICT_FIRST = 0x12F0000000000000 # createpolicy.fractional.L2::evict_first, fraction 1.0 (sm_100) + +# Shared-memory descriptor strides for the 128-byte swizzle, in 16-byte units. +LEADING = 16 +STRIDE = 8 * TMA_K_BOX * ELEM_BYTES +A_HALF_ELEMS = CTA_M * TMA_K_BOX +B_HALF_ELEMS = MMA_N * TMA_K_BOX +STEP = (MMA_K * ELEM_BYTES) >> 4 +A_BOX = A_HALF_ELEMS >> 3 +B_BOX = B_HALF_ELEMS >> 3 +STAGE_A = (CTA_M * CTA_K * ELEM_BYTES) >> 4 +STAGE_B = (MMA_N * CTA_K * ELEM_BYTES) >> 4 + +io_dtype = cutlass.BFloat16 + + +def num_k_tiles(k_in: int) -> int: + return (k_in + CTA_K - 1) // CTA_K + + +def supports(n_out: int, k_in: int) -> bool: + """Shapes this kernel runs: whole 128-row tiles and at most MAX_K_TILES k-tiles.""" + return n_out % CTA_M == 0 and k_in % TMA_K_BOX == 0 and num_k_tiles(k_in) <= MAX_K_TILES + + +@cute.kernel +def k3_decode_gemv_kernel( + tma_desc_w: cutlass.GridConstant[cuda.TensorMap], # W [N, K] bf16, 5-D, one call per k-tile + tma_desc_x: cutlass.GridConstant[ + cuda.TensorMap + ], # x [M, *] bf16, box 8 x 64: k-tiles [0, x_tiles) + tma_desc_x2: cutlass.GridConstant[ + cuda.TensorMap + ], # x2 [M, *] bf16, box 8 x 64: k-tiles [x_tiles, K/128) + y: cutlass.Array, # [M * N] bf16, token-major + rms_src: cutlass.Array, # int32 words of the bf16 [M, rms_cols] rows whose RMS scales accumulator 0 + num_tokens: cutlass.Int32, + x_col0: cutlass.Int32, # column of x where k-tile 0 starts + eps: cutlass.Float32, + n_out: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + x_tiles: cutlass.Constexpr[int], + split_acc: cutlass.Constexpr[bool], # the x2 k-tiles accumulate into a second TMEM accumulator + rms_cols: cutlass.Constexpr[int], # > 0: y = rsqrt(mean(rms_src row^2) + eps) * acc0 + acc1 + trigger_early: cutlass.Constexpr[bool], +): + k_tiles = num_k_tiles(k_in) + tx, _, _ = cute.arch.thread_idx() + bx, _, _ = cute.arch.block_idx() + warp_id = cute.arch.warp_idx() + tma_ptr_w = tma_desc_w.get_ptr() + tma_ptr_x = tma_desc_x.get_ptr() + tma_ptr_x2 = tma_desc_x2.get_ptr() + m_offset = bx * cutlass.Int32(CTA_M) + + smem_a = cutlass.Array( + io_dtype, k_tiles * CTA_M * CTA_K, space=cutlass.AddressSpace.smem, alignment=1024 + ) + smem_b = cutlass.Array( + io_dtype, k_tiles * MMA_N * CTA_K, space=cutlass.AddressSpace.smem, alignment=1024 + ) + weight_full = cutlass.Array( + cutlass.Int64, k_tiles, space=cutlass.AddressSpace.smem, alignment=8 + ) + act_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + acc_done = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + tmem_ptr_i32 = cutlass.Array(cutlass.Int32, 1, space=cutlass.AddressSpace.smem) + row_scale = cutlass.Array(cutlass.Float32, MMA_N, space=cutlass.AddressSpace.smem, alignment=16) + + if warp_id == 0: + prims.prefetch_tensormap(tma_ptr_w) + prims.prefetch_tensormap(tma_ptr_x) + if cutlass.const_expr(x_tiles < k_tiles): + prims.prefetch_tensormap(tma_ptr_x2) + if prims.elect_sync(): + for k in cutlass.range_constexpr(k_tiles): + prims.mbarrier_init(weight_full.subview(k), 1) + prims.mbarrier_init(act_full, 1) + prims.mbarrier_init(acc_done, 1) + if warp_id == 2: + prims.tcgen05_alloc(tmem_ptr_i32, TMEM_COLS) + prims.tcgen05_relinquish_alloc_permit() + prims.fence_mbarrier_init() + prims.barrier_cta_sync(0) + tmem_ptr = cutlass.inttoptr(tmem_ptr_i32.load(), 6, cutlass.Int32) + + if warp_id == 0: + # The whole weight slice of this CTA, ahead of the grid dependency. + if prims.elect_sync(): + for k in cutlass.range_constexpr(k_tiles): + prims.mbarrier_arrive_expect_tx(weight_full.subview(k), CTA_M * CTA_K * ELEM_BYTES) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a.subview(k * CTA_M * CTA_K), + tma_ptr_w, + ( + cutlass.Int32(0), + m_offset, + cutlass.Int32(k * TMA_COPY_ITERS), + cutlass.Int32(0), + cutlass.Int32(0), + ), + weight_full.subview(k), + l2_cache_hint=EVICT_FIRST, + ) + if cutlass.const_expr(trigger_early): + # Dependents may launch now; they wait for this whole grid before reading y. + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + elif warp_id == 1: + # The activations are written by the predecessor. + prims.griddepcontrol(prims.GridDepAction.WAIT) + if prims.elect_sync(): + prims.mbarrier_arrive_expect_tx(act_full, k_tiles * MMA_N * CTA_K * ELEM_BYTES) + for k in cutlass.range_constexpr(k_tiles): + for half in cutlass.range_constexpr(TMA_COPY_ITERS): + if cutlass.const_expr(k < x_tiles): + prims.cp_async_bulk_tensor_shared_cta_global( + smem_b.subview(k * MMA_N * CTA_K + half * B_HALF_ELEMS), + tma_ptr_x, + ( + x_col0 + cutlass.Int32(k * CTA_K + half * TMA_K_BOX), + cutlass.Int32(0), + ), + act_full, + ) + else: + prims.cp_async_bulk_tensor_shared_cta_global( + smem_b.subview(k * MMA_N * CTA_K + half * B_HALF_ELEMS), + tma_ptr_x2, + ( + cutlass.Int32((k - x_tiles) * CTA_K + half * TMA_K_BOX), + cutlass.Int32(0), + ), + act_full, + ) + elif warp_id == 2: + idesc = prims.Tcgen05InstrDesc.build( + c_dtype=cutlass.Float32, a_dtype=io_dtype, b_dtype=io_dtype, n_dim=MMA_N, m_dim=CTA_M + ) + desc_a_base = prims.Tcgen05SmemDesc.build( + start_address=smem_a, leading_byte_offset=LEADING, stride_byte_offset=STRIDE, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + desc_b_base = prims.Tcgen05SmemDesc.build( + start_address=smem_b, leading_byte_offset=LEADING, stride_byte_offset=STRIDE, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + tmem_acc1 = tmem_ptr + if cutlass.const_expr(split_acc): + tmem_acc1 = cutlass.inttoptr( + tmem_ptr_i32.load() + cutlass.Int32(MMA_N), 6, cutlass.Int32 + ) + while not cute.arch.mbarrier_try_wait(act_full.data_ptr(), 0): + pass + for k in cutlass.range_constexpr(k_tiles): + second = split_acc and k >= x_tiles + first_tile = k == 0 or (split_acc and k == x_tiles) + while not cute.arch.mbarrier_try_wait(weight_full.subview(k).data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + for kb in cutlass.range_constexpr(TMA_COPY_ITERS * K_BLOCKS_PER_HALF): + box = kb // K_BLOCKS_PER_HALF + within = kb % K_BLOCKS_PER_HALF + desc_a = desc_a_base + (k * STAGE_A + box * A_BOX + within * STEP) + desc_b = desc_b_base + (k * STAGE_B + box * B_BOX + within * STEP) + if prims.elect_sync(): + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, prims.CTAGroup.CTA_1, tmem_acc1 if second else tmem_ptr, + desc_a, desc_b, idesc, not (first_tile and kb == 0), + ) # fmt: skip + if prims.elect_sync(): + prims.tcgen05_commit(acc_done) + elif warp_id >= 4: + lane = tx % 32 + if cutlass.const_expr(rms_cols > 0): + # Per-token RMS of the rms_src rows while the MMA runs: epilogue warp w owns tokens 2w, 2w+1. + prims.griddepcontrol(prims.GridDepAction.WAIT) + for j in cutlass.range_constexpr(MMA_N // 4): + t = (warp_id % 4) * cutlass.Int32(MMA_N // 4) + cutlass.Int32(j) + sum_sq = cutlass.Float32(0.0) + if t < num_tokens: + for v in cutlass.range_constexpr(rms_cols // (8 * 32)): + words = rms_src.load( + idx=t * cutlass.Int32(rms_cols // 2) + + cutlass.Int32(v * 32 * 4) + + lane * cutlass.Int32(4), + vector_size=4, + alignment=16, + ) + for i in cutlass.range_constexpr(4): + lo = (words[i] << cutlass.Int32(16)).bitcast(cutlass.Float32) + hi = (words[i] & cutlass.Int32(-65536)).bitcast(cutlass.Float32) + sum_sq = sum_sq + lo * lo + hi * hi + for offset in (16, 8, 4, 2, 1): + sum_sq = sum_sq + cute.arch.shuffle_sync_bfly(sum_sq, offset=offset) + if lane == 0: + row_scale.store( + cute.math.rsqrt(sum_sq * cutlass.Float32(1.0 / rms_cols) + eps), idx=t + ) + prims.barrier_cta_sync(1, thread_count=128) + while not cute.arch.mbarrier_try_wait(acc_done.data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc0 = prims.tcgen05_ld( + "32x32b", cutlass.inttoptr(tmem_ptr_i32.load(), 6, cutlass.Float32), num=MMA_N + ) + acc1 = acc0 + if cutlass.const_expr(split_acc): + acc1 = prims.tcgen05_ld( + "32x32b", + cutlass.inttoptr(tmem_ptr_i32.load() + cutlass.Int32(MMA_N), 6, cutlass.Float32), + num=MMA_N, + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + row = (warp_id % 4) * 32 + lane + n = m_offset + row + for t in cutlass.range_constexpr(MMA_N): + if cutlass.Int32(t) < num_tokens: + value = cutlass.Float32(acc0[t]) + if cutlass.const_expr(rms_cols > 0): + value = value * row_scale.load(idx=t) + if cutlass.const_expr(split_acc): + value = value + cutlass.Float32(acc1[t]) + y.store(cutlass.BFloat16(value), idx=cutlass.Int32(t * n_out) + n) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + prims.barrier_cta_sync(1, thread_count=128) + if warp_id == 4: + prims.tcgen05_dealloc(tmem_ptr, TMEM_COLS) + + +# ============================================================================= +# Long-K split-K variant (e.g. the KDA qkvg projection, 3208 x 7168): the fuse_o_linear_ra structure. +# A cluster of SPLIT CTAs shares one 128-row weight tile; rank r streams the interleaved k-tiles r, +# r + SPLIT, ... through a RING_STAGES ring. Before the grid dependency each CTA fills its ring from HBM and +# prefetches the rest of its k-tiles into L2. Rows [32r, 32r + 32) of the tile are owned by rank r: every lane +# of the other ranks' epilogue warp r stores its fp32 partials into rank r's shared memory and arrives on its +# barrier (release, cluster scope); the owner waits (acquire, cluster scope), adds the SPLIT partials in rank +# order and stores bf16. (fuse_o_linear_ra pushes with st.async, which this DSL build does not expose.) +# ============================================================================= +SPLIT = 4 +RING_STAGES = 6 +ROWS_PER_RANK = CTA_M // SPLIT + + +def splitk_supports(n_out: int, k_in: int) -> bool: + k_tiles = num_k_tiles(k_in) + return k_in % CTA_K == 0 and k_tiles > MAX_K_TILES and k_tiles % SPLIT == 0 and n_out > 0 + + +@cute.kernel +def k3_decode_gemv_splitk_kernel( + tma_desc_w: cutlass.GridConstant[cuda.TensorMap], + tma_desc_x: cutlass.GridConstant[cuda.TensorMap], + y: cutlass.Array, # [M * N] bf16, token-major + num_tokens: cutlass.Int32, + n_out: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + trigger_early: cutlass.Constexpr[bool], +): + k_tiles = num_k_tiles(k_in) + my_tiles = k_tiles // SPLIT + tx, _, _ = cute.arch.thread_idx() + bx, _, _ = cute.arch.block_idx() + warp_id = cute.arch.warp_idx() + rank = cute.arch.block_idx_in_cluster() + tma_ptr_w = tma_desc_w.get_ptr() + tma_ptr_x = tma_desc_x.get_ptr() + m_offset = (bx // cutlass.Int32(SPLIT)) * cutlass.Int32(CTA_M) + + smem_a = cutlass.Array( + io_dtype, RING_STAGES * CTA_M * CTA_K, space=cutlass.AddressSpace.smem, alignment=1024 + ) + smem_b = cutlass.Array( + io_dtype, RING_STAGES * MMA_N * CTA_K, space=cutlass.AddressSpace.smem, alignment=1024 + ) + tma_done = cutlass.Array( + cutlass.Int64, RING_STAGES, space=cutlass.AddressSpace.smem, alignment=8 + ) + mma_done = cutlass.Array( + cutlass.Int64, RING_STAGES, space=cutlass.AddressSpace.smem, alignment=8 + ) + acc_done = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + reduce_ready = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + tmem_ptr_i32 = cutlass.Array(cutlass.Int32, 1, space=cutlass.AddressSpace.smem) + # Mailbox of the rows this rank owns: [source rank][row][token] fp32 (its own slot stays unused). + mailbox = cutlass.Array( + cutlass.Float32, + SPLIT * ROWS_PER_RANK * MMA_N, + space=cutlass.AddressSpace.smem, + alignment=16, + ) + + if warp_id == 0: + prims.prefetch_tensormap(tma_ptr_w) + prims.prefetch_tensormap(tma_ptr_x) + if prims.elect_sync(): + for s in cutlass.range_constexpr(RING_STAGES): + prims.mbarrier_init(tma_done.subview(s), 2) # the weight and the activation TMA + prims.mbarrier_init(mma_done.subview(s), 1) + prims.mbarrier_init(acc_done, 1) + prims.mbarrier_init(reduce_ready, (SPLIT - 1) * 32) # every lane of the 3 pushing warps + if warp_id == 2: + prims.tcgen05_alloc(tmem_ptr_i32, TMEM_COLS) + prims.tcgen05_relinquish_alloc_permit() + prims.fence_mbarrier_init() + # Cluster formation: makes the peers' shared memory and barriers addressable. + prims.barrier_cluster_arrive_relaxed() + prims.barrier_cluster_wait() + prims.barrier_cta_sync(0) + tmem_ptr = cutlass.inttoptr(tmem_ptr_i32.load(), 6, cutlass.Int32) + + if warp_id == 0: + if prims.elect_sync(): + # k-tiles past the ring go to L2 now, so the refills after the dependency read L2. + for i in cutlass.range_constexpr(RING_STAGES, my_tiles): + k_global = rank + cutlass.Int32(i * SPLIT) + prims.cp_async_bulk_tensor_prefetch( + tma_ptr_w, + [ + cutlass.Int32(0), + m_offset, + k_global * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), + cutlass.Int32(0), + ], + [], # tile mode: no im2col offsets + ) + for i in cutlass.range_constexpr(my_tiles): + stage = i % RING_STAGES + if cutlass.const_expr(i >= RING_STAGES): + while not cute.arch.mbarrier_try_wait( + mma_done.subview(stage).data_ptr(), (i // RING_STAGES - 1) % 2 + ): + pass + k_global = rank + cutlass.Int32(i * SPLIT) + prims.mbarrier_arrive_expect_tx(tma_done.subview(stage), CTA_M * CTA_K * ELEM_BYTES) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a.subview(stage * CTA_M * CTA_K), + tma_ptr_w, + ( + cutlass.Int32(0), + m_offset, + k_global * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), + cutlass.Int32(0), + ), + tma_done.subview(stage), + l2_cache_hint=EVICT_FIRST, + ) + if cutlass.const_expr(trigger_early and i == RING_STAGES - 1): + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + if cutlass.const_expr(trigger_early and my_tiles < RING_STAGES): + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + elif warp_id == 1: + prims.griddepcontrol(prims.GridDepAction.WAIT) + if prims.elect_sync(): + for i in cutlass.range_constexpr(my_tiles): + stage = i % RING_STAGES + if cutlass.const_expr(i >= RING_STAGES): + while not cute.arch.mbarrier_try_wait( + mma_done.subview(stage).data_ptr(), (i // RING_STAGES - 1) % 2 + ): + pass + k_global = rank + cutlass.Int32(i * SPLIT) + prims.mbarrier_arrive_expect_tx(tma_done.subview(stage), MMA_N * CTA_K * ELEM_BYTES) + for half in cutlass.range_constexpr(TMA_COPY_ITERS): + prims.cp_async_bulk_tensor_shared_cta_global( + smem_b.subview(stage * MMA_N * CTA_K + half * B_HALF_ELEMS), + tma_ptr_x, + ( + k_global * cutlass.Int32(CTA_K) + cutlass.Int32(half * TMA_K_BOX), + cutlass.Int32(0), + ), + tma_done.subview(stage), + ) + elif warp_id == 2: + idesc = prims.Tcgen05InstrDesc.build( + c_dtype=cutlass.Float32, a_dtype=io_dtype, b_dtype=io_dtype, n_dim=MMA_N, m_dim=CTA_M + ) + desc_a_base = prims.Tcgen05SmemDesc.build( + start_address=smem_a, leading_byte_offset=LEADING, stride_byte_offset=STRIDE, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + desc_b_base = prims.Tcgen05SmemDesc.build( + start_address=smem_b, leading_byte_offset=LEADING, stride_byte_offset=STRIDE, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + for i in cutlass.range_constexpr(my_tiles): + stage = i % RING_STAGES + while not cute.arch.mbarrier_try_wait( + tma_done.subview(stage).data_ptr(), (i // RING_STAGES) % 2 + ): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + for kb in cutlass.range_constexpr(TMA_COPY_ITERS * K_BLOCKS_PER_HALF): + box = kb // K_BLOCKS_PER_HALF + within = kb % K_BLOCKS_PER_HALF + desc_a = desc_a_base + (stage * STAGE_A + box * A_BOX + within * STEP) + desc_b = desc_b_base + (stage * STAGE_B + box * B_BOX + within * STEP) + if prims.elect_sync(): + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, prims.CTAGroup.CTA_1, tmem_ptr, desc_a, desc_b, idesc, + i > 0 or kb > 0, + ) # fmt: skip + if prims.elect_sync(): + prims.tcgen05_commit(mma_done.subview(stage)) + if prims.elect_sync(): + prims.tcgen05_commit(acc_done) + elif warp_id >= 4: + lane = tx % 32 + owner = warp_id % 4 # TMEM lanes 32 * owner .. + 31 = the rows rank `owner` reduces + while not cute.arch.mbarrier_try_wait(acc_done.data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc = prims.tcgen05_ld( + "32x32b", cutlass.inttoptr(tmem_ptr_i32.load(), 6, cutlass.Float32), num=MMA_N + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + slot_row = lane * MMA_N + if owner != rank: + base = rank * cutlass.Int32(ROWS_PER_RANK * MMA_N) + slot_row + for t in cutlass.range_constexpr(MMA_N): + prims.mapa(mailbox.subview(base + cutlass.Int32(t)), owner).store( + cutlass.Float32(acc[t]) + ) + # Release at cluster scope: this lane's partials are visible to the owner's acquire. + prims.mbarrier_arrive(prims.mapa(reduce_ready, owner), scope=prims.MemScope.CLUSTER) + else: + while not prims.mbarrier_try_wait_parity( + reduce_ready, 0, scope=prims.MBarrierScope.CLUSTER + ): + pass + n = m_offset + owner * cutlass.Int32(ROWS_PER_RANK) + lane + for t in cutlass.range_constexpr(MMA_N): + total = cutlass.Float32(0.0) + for q in cutlass.range_constexpr(SPLIT): + part = mailbox.load( + idx=cutlass.Int32(q * ROWS_PER_RANK * MMA_N) + slot_row + cutlass.Int32(t) + ) + total = total + cutlass.Float32( + cutlass.select_(rank == cutlass.Int32(q), cutlass.Float32(acc[t]), part) + ) + if cutlass.Int32(t) < num_tokens: + if n < cutlass.Int32(n_out): + y.store(cutlass.BFloat16(total), idx=cutlass.Int32(t * n_out) + n) + prims.barrier_cta_sync(1, thread_count=128) + if warp_id == 4: + prims.tcgen05_dealloc(tmem_ptr, TMEM_COLS) + + +def _weight_tensor_map(w, n_out, k_in): + """W as five TMA dimensions (64-element column chunk, row, 64-element chunk index, 1, 1) so one call + per k-tile lands both 128-byte-swizzled halves; strides in 16-byte units.""" + return cuda.create_tensor_map_tiled( + global_address=w.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[TMA_K_BOX, n_out, k_in // TMA_K_BOX, 1, 1], + global_strides=[ + (k_in * ELEM_BYTES) // 16, + (TMA_K_BOX * ELEM_BYTES) // 16, + (n_out * k_in * ELEM_BYTES) // 16, + (n_out * k_in * ELEM_BYTES) // 16, + ], + box_dims=[TMA_K_BOX, CTA_M, TMA_COPY_ITERS, 1, 1], + swizzle=cuda.TensorMapSwizzle.s128b, + ) + + +def _activation_tensor_map(x, cols, num_tokens): + """x [M, cols] as (cols, M) with an 8-row box: rows past num_tokens arrive as zeros.""" + return cuda.create_tensor_map_tiled( + global_address=x.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[cols, num_tokens], + global_strides=[(cols * ELEM_BYTES) // 16], + box_dims=[TMA_K_BOX, MMA_N], + swizzle=cuda.TensorMapSwizzle.s128b, + ) + + +@cute.jit +def k3_decode_gemv( + w: cute.Tensor, # [N, K] bf16, K contiguous + x: cute.Tensor, # [M, K] bf16, K contiguous, M <= 8 + y: cute.Tensor, # [M * N] bf16 + num_tokens: cutlass.Int32, + n_out: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + trigger_early: cutlass.Constexpr[bool], + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + tma_desc_w = _weight_tensor_map(w, n_out, k_in) + tma_desc_x = _activation_tensor_map(x, k_in, num_tokens) + k3_decode_gemv_kernel( + tma_desc_w, tma_desc_x, tma_desc_x, y, y, num_tokens, cutlass.Int32(0), cutlass.Float32(0.0), + n_out, k_in, num_k_tiles(k_in), False, 0, trigger_early, + ).launch(grid=(n_out // CTA_M, 1, 1), block=(THREADS, 1, 1), stream=stream, use_pdl=use_pdl) # fmt: skip + + +@cute.jit +def k3_decode_gemv_tail( + w: cute.Tensor, # [N, K_lat + K_act] bf16: latent-up columns of this rank's slice (zero-padded) | shared down + latent: cute.Tensor, # [M, rms_cols] bf16, the whole reduced latent row + latent_words: cute.Tensor, # the same memory as int32 words, for the RMS + act: cute.Tensor, # [M, K_act] bf16, the shared-expert activation + y: cute.Tensor, # [M * N] bf16 + num_tokens: cutlass.Int32, + lat_col0: cutlass.Int32, # first latent column of this rank's slice + eps: cutlass.Float32, + n_out: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + lat_tiles: cutlass.Constexpr[int], + rms_cols: cutlass.Constexpr[int], + trigger_early: cutlass.Constexpr[bool], + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + """``y = rmsnorm(latent)[:, slice] @ W_lat^T + act @ W_act^T`` (the RMS applied to the latent accumulator).""" + tma_desc_w = _weight_tensor_map(w, n_out, k_in) + tma_desc_lat = _activation_tensor_map(latent, rms_cols, num_tokens) + tma_desc_act = _activation_tensor_map(act, k_in - lat_tiles * CTA_K, num_tokens) + k3_decode_gemv_kernel( + tma_desc_w, tma_desc_lat, tma_desc_act, y, latent_words, num_tokens, lat_col0, eps, + n_out, k_in, lat_tiles, True, rms_cols, trigger_early, + ).launch(grid=(n_out // CTA_M, 1, 1), block=(THREADS, 1, 1), stream=stream, use_pdl=use_pdl) # fmt: skip + + +@cute.jit +def k3_decode_gemv_splitk( + w: cute.Tensor, # [N, K] bf16, K contiguous, K / 128 a multiple of SPLIT + x: cute.Tensor, # [M, K] bf16 + y: cute.Tensor, # [M * N] bf16 + num_tokens: cutlass.Int32, + n_out: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + trigger_early: cutlass.Constexpr[bool], + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + tma_desc_w = _weight_tensor_map(w, n_out, k_in) + tma_desc_x = _activation_tensor_map(x, k_in, num_tokens) + k3_decode_gemv_splitk_kernel( + tma_desc_w, tma_desc_x, y, num_tokens, n_out, k_in, trigger_early + ).launch( + grid=(((n_out + CTA_M - 1) // CTA_M) * SPLIT, 1, 1), + block=(THREADS, 1, 1), + cluster=(SPLIT, 1, 1), + stream=stream, + use_pdl=use_pdl, + ) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_decode_gemv/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_decode_gemv/op.py new file mode 100644 index 000000000000..1e816ca8d964 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_decode_gemv/op.py @@ -0,0 +1,209 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``trtllm::k3_decode_gemv``: ``x[M, K] @ weight[N, K]^T`` for M <= 8 in CuTe DSL (bf16, fp32 accumulation). + +The kernel loads its whole weight slice before waiting for the producer of ``x``, so launched with +PDL behind a long predecessor it only pays for the activation load, the MMA and the store. It is +compiled on the first call for each (N, K, early trigger, PDL), which must happen outside CUDA-graph +capture. +""" + +from __future__ import annotations + +import os +import threading +from typing import Dict + +import torch + +MAX_TOKENS = 8 + +_lock = threading.Lock() +_compiled: Dict[tuple, object] = {} + + +def _arg(t: torch.Tensor): + from cutlass.cute.runtime import from_dlpack + + # detach(): DLPack refuses tensors that require grad (weights are parameters). + return from_dlpack(t.detach(), assumed_align=16).mark_layout_dynamic(leading_dim=t.dim() - 1) + + +def _use_pdl() -> bool: + return os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" + + +def supports(x: torch.Tensor, weight: torch.Tensor) -> bool: + """Whether the kernel runs this call (shapes, dtypes, layout).""" + from . import k3_decode_gemv_kernel as kernel + + return ( + x.is_cuda + and x.dtype == torch.bfloat16 + and weight.dtype == torch.bfloat16 + and x.dim() == 2 + and weight.dim() == 2 + and 0 < x.shape[0] <= MAX_TOKENS + and x.shape[1] == weight.shape[1] + and x.is_contiguous() + and weight.is_contiguous() + and ( + kernel.supports(weight.shape[0], weight.shape[1]) + or kernel.splitk_supports(weight.shape[0], weight.shape[1]) + ) + ) + + +@torch.library.custom_op("trtllm::k3_decode_gemv", mutates_args=()) +def k3_decode_gemv( + x: torch.Tensor, weight: torch.Tensor, trigger_early: bool = True +) -> torch.Tensor: + """``x @ weight.T`` for bf16 ``x`` [M <= 8, K] and ``weight`` [N, K]; returns bf16 [M, N]. + + ``trigger_early`` lets the dependent grid launch once every CTA has issued its weight loads + (for dependents that wait for this whole grid before reading the output).""" + import cuda.bindings.driver as cuda_driver + + if not supports(x, weight): + raise ValueError( + f"k3_decode_gemv: unsupported call x {tuple(x.shape)} {x.dtype}, weight {tuple(weight.shape)} " + f"{weight.dtype}" + ) + num_tokens, k_in = x.shape + n_out = weight.shape[0] + y = torch.empty(num_tokens, n_out, dtype=torch.bfloat16, device=x.device) + args = (_arg(weight), _arg(x), _arg(y.view(-1))) + stream = cuda_driver.CUstream(torch.cuda.current_stream(x.device).cuda_stream) + use_pdl = _use_pdl() + key = (n_out, k_in, trigger_early, use_pdl) + fn = _compiled.get(key) + if fn is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "trtllm::k3_decode_gemv must run once per shape outside CUDA-graph capture first " + "(it compiles its kernel on the first call)." + ) + import cutlass.cute as cute + + from . import k3_decode_gemv_kernel as kernel + + # Short K: one CTA per weight tile, whole slice resident; long K: split-K over a cluster. + entry = ( + kernel.k3_decode_gemv if kernel.supports(n_out, k_in) else kernel.k3_decode_gemv_splitk + ) + with _lock: + fn = _compiled.get(key) + if fn is None: + fn = _compiled[key] = cute.compile( + entry, *args, num_tokens, n_out, k_in, trigger_early, use_pdl, stream + ) + # The compiled function takes the runtime arguments only. + fn(*args, num_tokens, stream) + return y + + +@k3_decode_gemv.register_fake +def _(x, weight, trigger_early=True): + return x.new_empty((x.shape[0], weight.shape[0]), dtype=torch.bfloat16) + + +def supports_tail( + latent: torch.Tensor, act: torch.Tensor, weight: torch.Tensor, width: int +) -> bool: + """Whether the tail kernel runs this call: the latent part of the weight a whole number of k-tiles + covering the slice, the activation too, at most MAX_K_TILES in all, a latent row the RMS loop tiles.""" + from . import k3_decode_gemv_kernel as kernel + + k_act = act.shape[1] if act.dim() == 2 else -1 + k_lat = weight.shape[1] - k_act + return ( + latent.is_cuda + and latent.dtype == act.dtype == weight.dtype == torch.bfloat16 + and latent.dim() == 2 + and act.dim() == 2 + and weight.dim() == 2 + and 0 < latent.shape[0] <= MAX_TOKENS + and act.shape[0] == latent.shape[0] + and latent.is_contiguous() + and act.is_contiguous() + and weight.is_contiguous() + and k_lat % kernel.CTA_K == 0 + and k_act % kernel.CTA_K == 0 + and 0 < width <= k_lat + and latent.shape[1] % (8 * 32) == 0 + and kernel.supports(weight.shape[0], weight.shape[1]) + ) + + +@torch.library.custom_op("trtllm::k3_decode_gemv_tail", mutates_args=()) +def k3_decode_gemv_tail( + latent: torch.Tensor, + act: torch.Tensor, + weight: torch.Tensor, + lo: int, + width: int, + eps: float, + trigger_early: bool = True, +) -> torch.Tensor: + """The row-parallel MoE tail, ``[rmsnorm(latent)[:, lo:lo+width] | act] @ weight.T``: ``latent`` + bf16 [M, H] is the whole reduced latent row, whose RMS scales the latent part; ``weight`` is + ``[latent up columns of the slice, zero-padded to a multiple of 128 | shared down]``. The RMS is + applied to the fp32 latent accumulator, not to bf16 inputs.""" + import cuda.bindings.driver as cuda_driver + + if not supports_tail(latent, act, weight, width): + raise ValueError( + f"k3_decode_gemv_tail: unsupported call latent {tuple(latent.shape)}, act {tuple(act.shape)}, " + f"weight {tuple(weight.shape)}, width {width}" + ) + from . import k3_decode_gemv_kernel as kernel + + num_tokens, rms_cols = latent.shape + n_out, k_in = weight.shape + lat_tiles = (k_in - act.shape[1]) // kernel.CTA_K + y = torch.empty(num_tokens, n_out, dtype=torch.bfloat16, device=latent.device) + args = ( + _arg(weight), + _arg(latent), + _arg(latent.view(-1).view(torch.int32)), + _arg(act), + _arg(y.view(-1)), + ) + stream = cuda_driver.CUstream(torch.cuda.current_stream(latent.device).cuda_stream) + use_pdl = _use_pdl() + key = ("tail", n_out, k_in, lat_tiles, rms_cols, trigger_early, use_pdl) + fn = _compiled.get(key) + if fn is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "trtllm::k3_decode_gemv_tail must run once per shape outside CUDA-graph capture first " + "(it compiles its kernel on the first call)." + ) + import cutlass.cute as cute + + with _lock: + fn = _compiled.get(key) + if fn is None: + fn = _compiled[key] = cute.compile( + kernel.k3_decode_gemv_tail, *args, num_tokens, lo, float(eps), n_out, k_in, lat_tiles, + rms_cols, trigger_early, use_pdl, stream, + ) # fmt: skip + fn(*args, num_tokens, lo, float(eps), stream) + return y + + +@k3_decode_gemv_tail.register_fake +def _(latent, act, weight, lo, width, eps, trigger_early=True): + return latent.new_empty((latent.shape[0], weight.shape[0]), dtype=torch.bfloat16) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_embed/__init__.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_embed/__init__.py new file mode 100644 index 000000000000..b34a787db3b4 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_embed/__init__.py @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Decode-size embedding row gather in CuTe DSL (``trtllm::k3_embed``). + +Importing :mod:`.op` registers the torch op; nothing is imported eagerly here so that +the CuTe DSL dependency stays optional for every other model. +""" diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_embed/k3_embed_kernel.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_embed/k3_embed_kernel.py new file mode 100644 index 000000000000..ece6836f32cc --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_embed/k3_embed_kernel.py @@ -0,0 +1,224 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +# ============================================================================= +# Decode-size embedding row gather -- CTM (prims/cute) kernel +# ============================================================================= +# +# out[t, :] = table[ids[t], :] t < N (a decode step's tokens), bf16 rows of H elements +# +# One 16-byte vector per thread: the grid covers N x H / 8 vectors, so every row is read by one wave of loads (one +# round trip after the grid dependency). The ids are the predecessor's output, read after griddepcontrol.wait; the +# dependents are launched at entry (they wait for this grid before reading out). An id outside [0, V) gives a zero +# row (the table is not read out of bounds). +# +# Norm mode (k3_embed_norm): the rows and the first layer's input RMSNorm in one launch +# +# raw[t, :] = table[ids[t], :] (the attention-residual snapshot bank's slot 0) +# out[t, :] = raw[t, :] * rsqrt(mean(raw[t, :]^2) + eps) * weight +# +# Bit-identical to k3_embed followed by flashinfer.norm.rmsnorm, whose GB200 kernel is the CuTe DSL RMSNormKernel +# (flashinfer/norm/kernels/rmsnorm.py). For a row of 6144 < H <= 16384 elements that kernel runs one 128-thread CTA +# per row, and thread t holds columns 8 t + v + 1024 k. This kernel keeps that geometry, the same tiled copies and +# fragments, and the same DSL expressions in the same order: +# - x * x, then the fragment's TensorSSA sum; +# - a butterfly over offsets 1 .. 16, the 4 warp sums through shared memory, a butterfly again; +# - sum / H, rsqrt(mean + eps) (fast-math), x * rstd * (w + 0). +# One CTA per token; no CTA reads another's writes. +# ============================================================================= +"""CTM decode embedding: the rows of a replicated embedding table for a step's token ids, one launch (optionally +with the first layer's input RMSNorm).""" + +from __future__ import annotations + +import cuda.bindings.driver as cuda_driver +import cutlass +import cutlass.cute as cute +from cutlass.experimental import primitives as prims + +THREADS = 256 +VEC_WORDS = 4 # int32 words per 16-byte vector (8 bf16) +NORM_THREADS = 128 # flashinfer's threads per row for 6144 < H <= 16384 (one row per CTA) +NORM_VEC = 8 # bf16 per 16-byte copy +NORM_WARPS = NORM_THREADS // 32 + + +def norm_supports_hidden(hidden: int) -> bool: + """Rows this norm mode reproduces: flashinfer's one-CTA 128-thread geometry, and a tile that covers the row.""" + return 6144 < hidden <= 16384 and hidden % (NORM_VEC * NORM_THREADS) == 0 + + +@cute.kernel +def k3_embed_kernel( + ids: cutlass.Array, # int32 or int64 [N] + table: cutlass.Array, # int32 words of the bf16 table [V * H / 2] + out: cutlass.Array, # int32 words of the bf16 output [N * H / 2] + vocab: cutlass.Int32, + n_tokens: cutlass.Constexpr[int], + row_vecs: cutlass.Constexpr[int], # H / 8 +): + tx, _, _ = cute.arch.thread_idx() + bx, _, _ = cute.arch.block_idx() + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + prims.griddepcontrol(prims.GridDepAction.WAIT) + v = bx * cutlass.Int32(THREADS) + tx + if v < cutlass.Int32(n_tokens * row_vecs): + t = v // cutlass.Int32(row_vecs) + c = v - t * cutlass.Int32(row_vecs) + tok = cutlass.Int64(ids.load(idx=t)) + valid = (tok >= cutlass.Int64(0)) & (tok < cutlass.Int64(vocab)) + row = cutlass.Int64(cutlass.select_(valid, tok, cutlass.Int64(0))) + e = table.load(idx=(row * cutlass.Int64(row_vecs) + cutlass.Int64(c)) * cutlass.Int64(VEC_WORDS), + vector_size=VEC_WORDS, alignment=16) # fmt: skip + zero = cutlass.Int32(0) + out.store( + (cutlass.Int32(cutlass.select_(valid, cutlass.Int32(e[0]), zero)), + cutlass.Int32(cutlass.select_(valid, cutlass.Int32(e[1]), zero)), + cutlass.Int32(cutlass.select_(valid, cutlass.Int32(e[2]), zero)), + cutlass.Int32(cutlass.select_(valid, cutlass.Int32(e[3]), zero))), + idx=v * cutlass.Int32(VEC_WORDS), + alignment=16, + ) # fmt: skip + + +@cute.jit +def k3_embed( + ids: cute.Tensor, + table: cute.Tensor, + out: cute.Tensor, + vocab: cutlass.Int32, + n_tokens: cutlass.Constexpr[int], + row_vecs: cutlass.Constexpr[int], + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + grid = (n_tokens * row_vecs + THREADS - 1) // THREADS + k3_embed_kernel(ids, table, out, vocab, n_tokens, row_vecs).launch( + grid=[grid, 1, 1], + block=[THREADS, 1, 1], + stream=stream, + use_pdl=use_pdl, + ) + + +@cute.kernel +def k3_embed_norm_kernel( + ids: cute.Tensor, # int32 or int64 [N] + table: cute.Tensor, # bf16 [V, H] + weight: cute.Tensor, # bf16 [H] + raw: cute.Tensor, # bf16 [N, H], out: the rows + out: cute.Tensor, # bf16 [N, H], out: the normed rows + vocab: cutlass.Int64, + eps: cutlass.Float32, + tv_layout: cute.Layout, + tiler_mn: cute.Shape, + hidden: cutlass.Constexpr[int], +): + tidx, _, _ = cute.arch.thread_idx() + bidx, _, _ = cute.arch.block_idx() + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + prims.griddepcontrol(prims.GridDepAction.WAIT) + + smem = cutlass.utils.SmemAllocator() + reduction_buffer = smem.allocate_tensor( + cutlass.Float32, cute.make_layout((1, NORM_WARPS)), byte_alignment=4 + ) + + tok = cutlass.Int64(ids[bidx]) + valid = (tok >= cutlass.Int64(0)) & (tok < vocab) + row = cutlass.Int64(cutlass.select_(valid, tok, cutlass.Int64(0))) + + gX = cute.local_tile(table, tiler_mn, (row, 0)) + gR = cute.local_tile(raw, tiler_mn, (bidx, 0)) + gY = cute.local_tile(out, tiler_mn, (bidx, 0)) + w_layout = cute.prepend(weight.layout, cute.make_layout((tiler_mn[0],), stride=(0,))) + gW = cute.local_tile(cute.make_tensor(weight.iterator, w_layout), tiler_mn, (0, 0)) + + copy_atom_load = cute.make_copy_atom( + cute.nvgpu.CopyUniversalOp(), table.element_type, num_bits_per_copy=128 + ) + copy_atom_store = cute.make_copy_atom( + cute.nvgpu.CopyUniversalOp(), out.element_type, num_bits_per_copy=128 + ) + thr_copy_X = cute.make_tiled_copy(copy_atom_load, tv_layout, tiler_mn).get_slice(tidx) + thr_copy_W = cute.make_tiled_copy(copy_atom_load, tv_layout, tiler_mn).get_slice(tidx) + thr_copy_O = cute.make_tiled_copy(copy_atom_store, tv_layout, tiler_mn).get_slice(tidx) + + tXgX = thr_copy_X.partition_S(gX) + tXrX = cute.make_fragment_like(tXgX) + tWgW = thr_copy_W.partition_S(gW) + tWrW = cute.make_fragment_like(tWgW) + tXrW = thr_copy_X.retile(tWrW) + tXgR = thr_copy_O.partition_D(gR) + tXgO = thr_copy_O.partition_D(gY) + tXrO = cute.make_fragment_like(tXgO) + + # The row (zero for an id outside [0, V)), straight to the snapshot slot; the weight. + tXrX.store(cute.zeros_like(tXrX, dtype=table.element_type)) + if valid: + cute.copy(copy_atom_load, tXgX, tXrX) + cute.copy(copy_atom_load, tWgW, tWrW) + cute.copy(copy_atom_store, tXrX, tXgR) + + # flashinfer's RMSNormKernel arithmetic, in its order (see the header). + x = tXrX.load().to(cutlass.Float32) + x_sq = x * x + sum_sq = x_sq.reduce(cute.ReductionOp.ADD, init_val=cutlass.Float32(0.0), reduction_profile=0) + for i in cutlass.range_constexpr(5): + sum_sq = sum_sq + cute.arch.shuffle_sync_bfly(sum_sq, offset=1 << i) + lane = cute.arch.lane_idx() + warp = cute.arch.warp_idx() + if lane == 0: + reduction_buffer[0, warp] = sum_sq + cute.arch.barrier() + total = cutlass.Float32(0.0) + if lane < NORM_WARPS: + total = reduction_buffer[0, lane] + for i in cutlass.range_constexpr(5): + total = total + cute.arch.shuffle_sync_bfly(total, offset=1 << i) + mean_sq = total / cutlass.Float32(hidden) + rstd = cute.math.rsqrt(mean_sq + eps, fastmath=True) + w = tXrW.load().to(cutlass.Float32) + y = x * rstd * (w + cutlass.Float32(0.0)) + tXrO.store(y.to(out.element_type)) + cute.copy(copy_atom_store, tXrO, tXgO) + + +@cute.jit +def k3_embed_norm( + ids: cute.Tensor, + table: cute.Tensor, + weight: cute.Tensor, + raw: cute.Tensor, + out: cute.Tensor, + vocab: cutlass.Int64, + eps: cutlass.Float32, + n_tokens: cutlass.Constexpr[int], + hidden: cutlass.Constexpr[int], + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + # flashinfer's thread-value layout for one row: thread t, value (v, k) -> column 8 t + v + 1024 k. + blocks = hidden // (NORM_VEC * NORM_THREADS) + tv_layout = cute.make_layout(((NORM_THREADS, 1), (NORM_VEC, blocks)), + stride=((NORM_VEC, 1), (1, NORM_VEC * NORM_THREADS))) # fmt: skip + k3_embed_norm_kernel( + ids, table, weight, raw, out, vocab, eps, tv_layout, (1, hidden), hidden + ).launch( + grid=[n_tokens, 1, 1], + block=[NORM_THREADS, 1, 1], + smem=NORM_WARPS * 4, + stream=stream, + use_pdl=use_pdl, + ) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_embed/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_embed/op.py new file mode 100644 index 000000000000..420d619fc2ee --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_embed/op.py @@ -0,0 +1,181 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Torch ops of the decode-size embedding gather: ``trtllm::k3_embed`` (``table[ids]`` for up to ``MAX_TOKENS`` int32 +or int64 ids and a bf16 table whose rows are whole 16-byte vectors, one vector per thread) and ``trtllm::k3_embed_norm`` +(the same rows written into a caller's buffer, plus their RMSNorm, bit-identical to ``flashinfer.norm.rmsnorm``). +Compiled on the first call for its shape, which must happen outside CUDA-graph capture.""" + +from __future__ import annotations + +import os +import threading +from typing import Dict + +import torch + +MAX_TOKENS = 64 + +_lock = threading.Lock() +_compiled: Dict[tuple, object] = {} + + +def _arg(t: torch.Tensor): + from cutlass.cute.runtime import from_dlpack + + return from_dlpack(t.detach(), assumed_align=16).mark_layout_dynamic(leading_dim=t.dim() - 1) + + +def supports(ids: torch.Tensor, table: torch.Tensor) -> bool: + """Whether ``k3_embed`` gathers ``table[ids]``.""" + return ( + ids.is_cuda + and ids.dim() == 1 + and ids.dtype in (torch.int32, torch.int64) + and 0 < ids.numel() <= MAX_TOKENS + and table.dtype == torch.bfloat16 + and table.dim() == 2 + and table.is_contiguous() + and table.shape[1] % 8 == 0 + and table.data_ptr() % 16 == 0 + ) + + +@torch.library.custom_op("trtllm::k3_embed", mutates_args=()) +def k3_embed(ids: torch.Tensor, table: torch.Tensor) -> torch.Tensor: + """``table[ids]`` for int32 / int64 ``ids`` [N <= 64] and a contiguous bf16 ``table`` [V, H] (H % 8 == 0); ids + outside [0, V) give zero rows.""" + import cuda.bindings.driver as cuda_driver + + if not supports(ids, table): + raise ValueError( + f"k3_embed: unsupported call: ids {tuple(ids.shape)} {ids.dtype}, table {tuple(table.shape)} {table.dtype} " + f"(int32 / int64 [N <= {MAX_TOKENS}], contiguous bf16 [V, H % 8 == 0])" + ) + from . import k3_embed_kernel as kern + + n = ids.numel() + vocab, hidden = table.shape + out = torch.empty(n, hidden, dtype=torch.bfloat16, device=ids.device) + args = ( + _arg(ids.contiguous()), + _arg(table.view(-1).view(torch.int32)), + _arg(out.view(-1).view(torch.int32)), + ) + use_pdl = os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" + stream = cuda_driver.CUstream(torch.cuda.current_stream(ids.device).cuda_stream) + key = (n, hidden // 8, ids.dtype, use_pdl) + fn = _compiled.get(key) + if fn is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "trtllm::k3_embed must run once per shape outside CUDA-graph capture first" + ) + import cutlass.cute as cute + + with _lock: + fn = _compiled.get(key) + if fn is None: + fn = _compiled[key] = cute.compile( + kern.k3_embed, *args, int(vocab), n, hidden // 8, use_pdl, stream + ) + fn(*args, int(vocab), stream) + return out + + +@k3_embed.register_fake +def _(ids, table): + return table.new_empty((ids.numel(), table.shape[1])) + + +def _plain_arg(t: torch.Tensor, align: int = 16): + from cutlass.cute.runtime import from_dlpack + + return from_dlpack(t.detach(), assumed_align=align) + + +def norm_supports_hidden(hidden: int) -> bool: + """Row widths ``k3_embed_norm`` reproduces flashinfer's RMSNorm for.""" + from . import k3_embed_kernel as kern + + return kern.norm_supports_hidden(hidden) + + +def supports_norm( + ids: torch.Tensor, table: torch.Tensor, weight: torch.Tensor, raw: torch.Tensor +) -> bool: + """Whether ``k3_embed_norm`` gathers ``table[ids]`` into ``raw`` and norms it with ``weight``.""" + hidden = table.shape[1] if table.dim() == 2 else 0 + return ( + supports(ids, table) + and norm_supports_hidden(hidden) + and weight.dtype == torch.bfloat16 + and tuple(weight.shape) == (hidden,) + and weight.is_contiguous() + and weight.data_ptr() % 16 == 0 + and raw.dtype == torch.bfloat16 + and tuple(raw.shape) == (ids.numel(), hidden) + and raw.is_contiguous() + and raw.data_ptr() % 16 == 0 + ) + + +@torch.library.custom_op("trtllm::k3_embed_norm", mutates_args=("raw",)) +def k3_embed_norm( + ids: torch.Tensor, table: torch.Tensor, weight: torch.Tensor, eps: float, raw: torch.Tensor +) -> torch.Tensor: + """``raw[:] = table[ids]`` (ids outside [0, V) give zero rows) and returns ``rmsnorm(raw, weight, eps)``, bf16 + [N, H]: bit-identical to ``k3_embed`` followed by ``flashinfer.norm.rmsnorm`` (see ``k3_embed_kernel``).""" + import cuda.bindings.driver as cuda_driver + + if not supports_norm(ids, table, weight, raw): + raise ValueError( + f"k3_embed_norm: unsupported call: ids {tuple(ids.shape)} {ids.dtype}, table {tuple(table.shape)} " + f"{table.dtype}, weight {tuple(weight.shape)} {weight.dtype}, raw {tuple(raw.shape)} {raw.dtype} " + f"(int32 / int64 [N <= {MAX_TOKENS}], contiguous bf16 [V, H] with 6144 < H <= 16384 and H % 1024 == 0, " + f"weight [H], raw [N, H])" + ) + from . import k3_embed_kernel as kern + + n = ids.numel() + vocab, hidden = table.shape + out = torch.empty(n, hidden, dtype=torch.bfloat16, device=ids.device) + ids = ids.contiguous() + args = (_plain_arg(ids, ids.element_size()),) + tuple( + _plain_arg(t) for t in (table, weight, raw, out) + ) + use_pdl = os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" + stream = cuda_driver.CUstream(torch.cuda.current_stream(ids.device).cuda_stream) + key = ("norm", n, vocab, hidden, ids.dtype, use_pdl) + fn = _compiled.get(key) + if fn is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "trtllm::k3_embed_norm must run once per shape outside CUDA-graph capture first" + ) + import cutlass.cute as cute + + with _lock: + fn = _compiled.get(key) + if fn is None: + fn = _compiled[key] = cute.compile( + kern.k3_embed_norm, *args, int(vocab), float(eps), n, hidden, use_pdl, stream + ) + fn(*args, int(vocab), float(eps), stream) + return out + + +@k3_embed_norm.register_fake +def _(ids, table, weight, eps, raw): + return table.new_empty((ids.numel(), table.shape[1])) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_head_gemv/__init__.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_head_gemv/__init__.py new file mode 100644 index 000000000000..db044cfd9984 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_head_gemv/__init__.py @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Kimi K3 persistent CTM head GEMV (``trtllm::k3_head_gemv``). + +Importing :mod:`.op` registers the torch op; nothing is imported eagerly here so that +the CuTe DSL dependency stays optional for every other model. +""" diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_head_gemv/k3_head_gemv_kernel.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_head_gemv/k3_head_gemv_kernel.py new file mode 100644 index 000000000000..247b6c1f6ca8 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_head_gemv/k3_head_gemv_kernel.py @@ -0,0 +1,1137 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +# ============================================================================= +# Kimi K3 vocab-shard head GEMV -- CTM (prims/cute) kernel, M <= 8 tokens, bf16 in/out, persistent +# ============================================================================= +# +# y[t, n] = sum_k x[t, k] * W[n, k] (fp32 accumulation in TMEM, one bf16 rounding) +# A = W (N, K) K-major bf16 (the streamed weight: an lm_head shard, N ~ 10^4) +# B = x (M <= 8, K) K-major bf16 (8 token columns; rows past M arrive as zeros) +# +# Work units: the weight is cut into 128-row tiles and every tile's K into CHUNKS chunks of CH k-tiles (128 each); +# unit u is (tile u // CHUNKS, chunk u % CHUNKS). One CTA per SM (GRID CTAs), each running units until the pool is +# empty: unit bx is static (its first stages are loaded before griddepcontrol.wait), the rest are claimed from a +# global counter (every CTA makes exactly one failing claim; the last ticket rolls the counter back to 0, so every +# launch and graph replay starts from 0). Units are claimed in tile-major order, so a tile's chunks finish together. +# +# Split-K combine (CHUNKS > 1): a unit's epilogue stores its fp32 partial [128 rows][8 tokens] to ws[u], then +# thread 0 of the epilogue warps counts the unit on cnt[tile] with an acq_rel atomic after a barrier of those warps +# (the barrier orders their stores before the release). The unit that brings the count to CHUNKS finalizes the +# tile: it adds the CHUNKS partials in chunk order (its own from registers), rounds once to bf16 and stores y's +# [M][128] block, and resets cnt[tile] to 0. A unit's partial does not depend on which CTA computes it and the +# combine order is fixed, so y is bit-identical from run to run whatever the claim order. +# +# Stage = one k-tile: A 32 KB (both 64-element halves of the 128-byte swizzle, one 5-D TMA call) + B 2 KB (two TMA +# calls from x). RING stages; the phases run straight through unit boundaries. TMEM: two 8-column fp32 +# accumulators, so the epilogue of one unit overlaps the MMAs of the next. +# +# L2: the weight loads are EVICT_FIRST, except for tiles < keep_tiles (normal priority), so that a later reader of +# the same shard (the drafter head after the target head) can find them in L2. +# +# Warps: 0 claimer + A/B TMA (one elected lane), 2 TMEM allocation + tcgen05 MMA (M 128, N 8, K 16), 4-7 epilogue +# (TMEM -> registers -> partial / fixup -> bf16 staging -> 16-byte stores), 1 and 3 idle. +# +# PDL: before the grid-dependency wait only barrier init, TMEM allocation, the tensormap prefetches and the static +# unit's A loads (a weight: nothing in the graph writes it). After it: B loads, claims, every global write. The +# dependents are released after the CTA's failing claim, so they launch only once every CTA of this grid has +# started (no dependent CTA can hold an SM that a static unit still needs). +# ============================================================================= +"""Persistent CTM head GEMV for Kimi K3 (``y = x @ W^T``, M <= 8, a large vocab shard): dynamic (tile, k-chunk) +units with a deterministic split-K combine, the weight streamed before the grid-dependency wait.""" + +from __future__ import annotations + +import cuda.bindings.driver as cuda_driver +import cutlass +import cutlass.cute as cute +import cutlass.experimental.cuda as cuda +from cutlass._mlir.dialects import llvm as _llvm +from cutlass._mlir.extras import types as _T +from cutlass.cutlass_dsl import dsl_user_op +from cutlass.experimental import primitives as prims + +CTA_M = 128 # weight rows per tile = the MMA's M +MMA_N = 8 # token columns +CTA_K = 128 # one k-tile: two 64-element halves of the 128-byte swizzle +MMA_K = 16 +TMA_K_BOX = 64 +TMA_COPY_ITERS = CTA_K // TMA_K_BOX +K_BLOCKS_PER_HALF = TMA_K_BOX // MMA_K +THREADS = 256 +EPI_THREADS = 128 +ELEM_BYTES = 2 +VEC = 8 # bf16 per 16-byte vector +TMEM_COLS = 32 +UNIT_RING = ( + 16 # claimed-unit broadcast slots (the claimer runs at most RING k-tiles + 2 units ahead) +) +EVICT_FIRST = 0x12F0000000000000 # createpolicy.fractional.L2::evict_first, fraction 1.0 (sm_100) +SMEM_BYTES = 227 * 1024 +MAX_RING = 6 +FLAG_BACKOFF_NS = ( + 256 # between polls of a split tile's piece flags (the pieces finish long before the finalizer) +) + +# Shared-memory descriptor strides for the 128-byte swizzle, in 16-byte units. +LEADING = 16 +STRIDE = 8 * TMA_K_BOX * ELEM_BYTES +A_HALF_ELEMS = CTA_M * TMA_K_BOX +B_HALF_ELEMS = MMA_N * TMA_K_BOX +STEP = (MMA_K * ELEM_BYTES) >> 4 +A_BOX = A_HALF_ELEMS >> 3 +B_BOX = B_HALF_ELEMS >> 3 +STAGE_A = (CTA_M * CTA_K * ELEM_BYTES) >> 4 +STAGE_B = (MMA_N * CTA_K * ELEM_BYTES) >> 4 +A_BYTES = CTA_M * CTA_K * ELEM_BYTES +B_BYTES = MMA_N * CTA_K * ELEM_BYTES + +io_dtype = cutlass.BFloat16 + + +def num_k_tiles(k_in: int) -> int: + return k_in // CTA_K + + +def num_units(n_out: int, k_in: int, chunk_tiles: int) -> int: + return (n_out // CTA_M) * (num_k_tiles(k_in) // chunk_tiles) + + +def smem_bytes(ring: int) -> int: + """Shared memory of a CTA: the A and B rings, the bf16 output staging tile, barriers and slots.""" + return ring * (A_BYTES + B_BYTES) + MMA_N * CTA_M * ELEM_BYTES + 1024 + + +def supports(n_out: int, k_in: int, chunk_tiles: int, ring: int) -> bool: + """Shapes the kernel runs: whole 128-row tiles, whole k-tiles, chunks that tile K, a ring that fits.""" + k_tiles = num_k_tiles(k_in) + return ( + n_out > 0 + and n_out % CTA_M == 0 + and k_in > 0 + and k_in % CTA_K == 0 + and 0 < chunk_tiles <= k_tiles + and k_tiles % chunk_tiles == 0 + and 1 <= ring <= MAX_RING + and smem_bytes(ring) <= SMEM_BYTES + ) + + +@dsl_user_op +def _atomic_add_acq_rel(addr_i64, val, *, loc=None, ip=None): + """atom.acq_rel.gpu.global.add.u32, returning the old value.""" + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [addr_i64.ir_value(loc=loc, ip=ip), val.ir_value(loc=loc, ip=ip)], + "atom.acq_rel.gpu.global.add.u32 $0, [$1], $2;", "=r,l,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _atomic_add_relaxed(addr_i64, val, *, loc=None, ip=None): + """atom.relaxed.gpu.global.add.u32, returning the old value (the unit claim).""" + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [addr_i64.ir_value(loc=loc, ip=ip), val.ir_value(loc=loc, ip=ip)], + "atom.relaxed.gpu.global.add.u32 $0, [$1], $2;", "=r,l,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _ld_acquire(addr_i64, *, loc=None, ip=None): + """ld.acquire.gpu.global.u32 (a flag written by another SM).""" + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [addr_i64.ir_value(loc=loc, ip=ip)], + "ld.acquire.gpu.global.u32 $0, [$1];", "=r,l", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _flag_release(addr_i64, *, loc=None, ip=None): + """fence.acq_rel.gpu, then st.relaxed.gpu 1: the release of everything the thread (and, through a preceding + barrier, its CTA) stored before.""" + _llvm.inline_asm( + None, [addr_i64.ir_value(loc=loc, ip=ip)], + "fence.acq_rel.gpu;\n\tst.relaxed.gpu.global.u32 [$0], 1;", "l", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _add_rn(a, b, *, loc=None, ip=None): + """add.rn.f32: an fp32 add the compiler cannot reassociate or contract.""" + return cutlass.Float32( + _llvm.inline_asm( + _T.f32(), [a.ir_value(loc=loc, ip=ip), b.ir_value(loc=loc, ip=ip)], + "add.rn.f32 $0, $1, $2;", "=f,f,f", has_side_effects=False, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +def _store_tile(stage_out, y, vals, row, tid, tile, n_out): + """The finalized [8 tokens] sums of tile row ``row`` -> bf16 -> y[token, tile rows] for all 8 token rows (y has 8 + rows; rows past M hold the zero-filled x rows' zeros): staged per token in shared memory, then 16-byte stores + (rows contiguous per token). The epilogue warps' barriers bracket the staging tile's reuse.""" + for t in range(MMA_N): + stage_out.store(vals[t].to(io_dtype), idx=cutlass.Int32(t * CTA_M) + row) + prims.barrier_cta_sync(1, thread_count=EPI_THREADS) + ct = tid // cutlass.Int32(CTA_M // VEC) + cc = tid % cutlass.Int32(CTA_M // VEC) + v = stage_out.load( + idx=ct * cutlass.Int32(CTA_M) + cc * cutlass.Int32(VEC), vector_size=VEC, alignment=16 + ) + y.store(v, idx=ct * cutlass.Int32(n_out) + tile * cutlass.Int32(CTA_M) + cc * cutlass.Int32(VEC), + vector_size=VEC, alignment=16) # fmt: skip + prims.barrier_cta_sync(1, thread_count=EPI_THREADS) + + +@cute.kernel +def k3_head_gemv_kernel( + tma_desc_w: cutlass.GridConstant[cuda.TensorMap], # W [N, K] bf16, 5-D, one call per k-tile + tma_desc_x: cutlass.GridConstant[cuda.TensorMap], # x [M, K] bf16, box 64 x 8 + y: cutlass.Array, # [8 * N] bf16, token-major (all 8 token rows are written) + ws: cutlass.Array, # fp32 [units * 128 * 8]: the units' partials (CHUNKS > 1) + cnt: cutlass.Array, # int32 [tiles]: units of the tile done (0 between launches) + claim: cutlass.Array, # int32 [1]: the unit pool's ticket counter (0 between launches) + num_tokens: cutlass.Int32, + keep_tiles: cutlass.Int32, # tiles [0, keep_tiles) load at normal L2 priority, the rest EVICT_FIRST + n_out: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + chunk_tiles: cutlass.Constexpr[int], # CH: k-tiles per unit + ring: cutlass.Constexpr[int], + grid: cutlass.Constexpr[int], +): + """Persistent: unit bx, then claimed units until the pool is empty; the last unit of a tile finalizes it.""" + k_tiles = num_k_tiles(k_in) + chunks = k_tiles // chunk_tiles + units = (n_out // CTA_M) * chunks + tickets = max( + units, grid + ) # claims made in one launch: units - grid succeed (if positive), grid fail + pre = min(ring, chunk_tiles) # stages of the static unit loaded before the grid-dependency wait + tx, _, _ = cute.arch.thread_idx() + bx, _, _ = cute.arch.block_idx() + warp_id = cute.arch.warp_idx() + tma_ptr_w = tma_desc_w.get_ptr() + tma_ptr_x = tma_desc_x.get_ptr() + + smem_a = cutlass.Array( + io_dtype, ring * CTA_M * CTA_K, space=cutlass.AddressSpace.smem, alignment=1024 + ) + smem_b = cutlass.Array( + io_dtype, ring * MMA_N * CTA_K, space=cutlass.AddressSpace.smem, alignment=1024 + ) + stage_out = cutlass.Array( + io_dtype, MMA_N * CTA_M, space=cutlass.AddressSpace.smem, alignment=16 + ) + full = cutlass.Array(cutlass.Int64, ring, space=cutlass.AddressSpace.smem, alignment=8) + empty = cutlass.Array(cutlass.Int64, ring, space=cutlass.AddressSpace.smem, alignment=8) + acc_full = cutlass.Array(cutlass.Int64, 2, space=cutlass.AddressSpace.smem, alignment=8) + acc_empty = cutlass.Array(cutlass.Int64, 2, space=cutlass.AddressSpace.smem, alignment=8) + unit_ready = cutlass.Array( + cutlass.Int64, UNIT_RING, space=cutlass.AddressSpace.smem, alignment=8 + ) + unit_slot = cutlass.Array( + cutlass.Int32, UNIT_RING, space=cutlass.AddressSpace.smem, alignment=16 + ) + s_last = cutlass.Array(cutlass.Int32, 4, space=cutlass.AddressSpace.smem, alignment=16) + tmem_ptr_i32 = cutlass.Array(cutlass.Int32, 1, space=cutlass.AddressSpace.smem) + + if warp_id == 0: + prims.prefetch_tensormap(tma_ptr_w) + prims.prefetch_tensormap(tma_ptr_x) + if prims.elect_sync(): + for s in cutlass.range_constexpr(ring): + prims.mbarrier_init(full.subview(s), 1) + prims.mbarrier_init(empty.subview(s), 1) + for b in cutlass.range_constexpr(2): + prims.mbarrier_init(acc_full.subview(b), 1) + prims.mbarrier_init(acc_empty.subview(b), 4) # one elected lane per epilogue warp + for s in cutlass.range_constexpr(UNIT_RING): + prims.mbarrier_init(unit_ready.subview(s), 1) + if warp_id == 2: + prims.tcgen05_alloc(tmem_ptr_i32, TMEM_COLS) + prims.tcgen05_relinquish_alloc_permit() + prims.fence_mbarrier_init() + prims.barrier_cta_sync(0) + + if warp_id == 0: + # ===================================================================== + # Claimer + TMA: the static unit's first stages before the wait, then + # every stage of every unit this CTA gets, A and B on one barrier. + # ===================================================================== + if prims.elect_sync(): + unit = cutlass.Int32(cutlass.select_(bx < cutlass.Int32(units), bx, cutlass.Int32(-1))) + unit_slot.store(unit, idx=0) + prims.mbarrier_arrive(unit_ready.subview(0)) + static_tile = unit // cutlass.Int32(chunks) + static_k0 = (unit % cutlass.Int32(chunks)) * cutlass.Int32(chunk_tiles) + if unit >= cutlass.Int32(0): + for jj in cutlass.range_constexpr(pre): + prims.mbarrier_arrive_expect_tx(full.subview(jj), A_BYTES + B_BYTES) + if static_tile < keep_tiles: + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a.subview(jj * CTA_M * CTA_K), + tma_ptr_w, + ( + cutlass.Int32(0), + static_tile * cutlass.Int32(CTA_M), + (static_k0 + cutlass.Int32(jj)) * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), + cutlass.Int32(0), + ), + full.subview(jj), + ) + else: + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a.subview(jj * CTA_M * CTA_K), + tma_ptr_w, + ( + cutlass.Int32(0), + static_tile * cutlass.Int32(CTA_M), + (static_k0 + cutlass.Int32(jj)) * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), + cutlass.Int32(0), + ), + full.subview(jj), + l2_cache_hint=EVICT_FIRST, + ) + prims.griddepcontrol(prims.GridDepAction.WAIT) + issued = cutlass.Int32(0) # k-tiles issued by this CTA (the ring position) + next_j = cutlass.Int32(0) # first k-tile of the current unit still to issue + if unit >= cutlass.Int32(0): + for jj in cutlass.range_constexpr(pre): + for half in cutlass.range_constexpr(TMA_COPY_ITERS): + prims.cp_async_bulk_tensor_shared_cta_global( + smem_b.subview(jj * MMA_N * CTA_K + half * B_HALF_ELEMS), + tma_ptr_x, + ((static_k0 + cutlass.Int32(jj)) * cutlass.Int32(CTA_K) + cutlass.Int32(half * TMA_K_BOX), + cutlass.Int32(0)), + full.subview(jj), + ) # fmt: skip + issued = cutlass.Int32(pre) + next_j = cutlass.Int32(pre) + n_units = cutlass.Int32(0) # units published so far - 1 + pending = cutlass.Boolean(True) + while pending: + if unit >= cutlass.Int32(0): + tile = unit // cutlass.Int32(chunks) + k0 = (unit % cutlass.Int32(chunks)) * cutlass.Int32(chunk_tiles) + j = next_j + while j < cutlass.Int32(chunk_tiles): + stage = issued % cutlass.Int32(ring) + # The stage's previous MMAs committed (a fresh barrier passes parity 1 at once). + parity = ( + (issued // cutlass.Int32(ring)) & cutlass.Int32(1) + ) ^ cutlass.Int32(1) + while not cute.arch.mbarrier_try_wait( + empty.subview(stage).data_ptr(), parity + ): + pass + k = k0 + j + prims.mbarrier_arrive_expect_tx(full.subview(stage), A_BYTES + B_BYTES) + if tile < keep_tiles: + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a.subview(stage * cutlass.Int32(CTA_M * CTA_K)), tma_ptr_w, + (cutlass.Int32(0), tile * cutlass.Int32(CTA_M), k * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), cutlass.Int32(0)), + full.subview(stage), + ) # fmt: skip + else: + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a.subview(stage * cutlass.Int32(CTA_M * CTA_K)), tma_ptr_w, + (cutlass.Int32(0), tile * cutlass.Int32(CTA_M), k * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), cutlass.Int32(0)), + full.subview(stage), l2_cache_hint=EVICT_FIRST, + ) # fmt: skip + for half in cutlass.range_constexpr(TMA_COPY_ITERS): + prims.cp_async_bulk_tensor_shared_cta_global( + smem_b.subview( + stage * cutlass.Int32(MMA_N * CTA_K) + + cutlass.Int32(half * B_HALF_ELEMS) + ), + tma_ptr_x, + ( + k * cutlass.Int32(CTA_K) + cutlass.Int32(half * TMA_K_BOX), + cutlass.Int32(0), + ), + full.subview(stage), + ) + issued = issued + cutlass.Int32(1) + j = j + cutlass.Int32(1) + # Claim the next unit; the holder of the last ticket rolls the counter back to 0. + ticket = _atomic_add_relaxed(claim.data_ptr().toint(), cutlass.Int32(1)) + if ticket == cutlass.Int32(tickets - 1): + _atomic_add_relaxed(claim.data_ptr().toint(), cutlass.Int32(-tickets)) + claimed = ticket + cutlass.Int32(grid) + unit = cutlass.Int32( + cutlass.select_(claimed < cutlass.Int32(units), claimed, cutlass.Int32(-1)) + ) + n_units = n_units + cutlass.Int32(1) + slot = n_units % cutlass.Int32(UNIT_RING) + unit_slot.store(unit, idx=slot) + prims.mbarrier_arrive(unit_ready.subview(slot)) + next_j = cutlass.Int32(0) + pending = unit >= cutlass.Int32(0) + # Every CTA of this grid has started and made its last claim: the dependents may launch. + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + elif warp_id == 2: + # ===================================================================== + # MMA: each unit's CH k-tiles into accumulator (unit count) % 2. + # ===================================================================== + idesc = prims.Tcgen05InstrDesc.build( + c_dtype=cutlass.Float32, a_dtype=io_dtype, b_dtype=io_dtype, n_dim=MMA_N, m_dim=CTA_M + ) + desc_a_base = prims.Tcgen05SmemDesc.build( + start_address=smem_a, leading_byte_offset=LEADING, stride_byte_offset=STRIDE, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + desc_b_base = prims.Tcgen05SmemDesc.build( + start_address=smem_b, leading_byte_offset=LEADING, stride_byte_offset=STRIDE, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + tmem_base = tmem_ptr_i32.load() + used = cutlass.Int32(0) # k-tiles consumed (the ring position) + count = cutlass.Int32(0) # units done by this CTA + while not cute.arch.mbarrier_try_wait(unit_ready.subview(0).data_ptr(), 0): + pass + unit = unit_slot.load(idx=0) + while unit >= cutlass.Int32(0): + buf = count & cutlass.Int32(1) + if count >= cutlass.Int32(2): + # The epilogue has read this accumulator's previous unit. + while not cute.arch.mbarrier_try_wait( + acc_empty.subview(buf).data_ptr(), + ((count >> cutlass.Int32(1)) - cutlass.Int32(1)) & cutlass.Int32(1), + ): + pass + tmem_acc = cutlass.inttoptr(tmem_base + buf * cutlass.Int32(MMA_N), 6, cutlass.Int32) + for j in cutlass.range(chunk_tiles, unroll=1): + stage = used % cutlass.Int32(ring) + while not cute.arch.mbarrier_try_wait( + full.subview(stage).data_ptr(), (used // cutlass.Int32(ring)) & cutlass.Int32(1) + ): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + for kb in cutlass.range_constexpr(TMA_COPY_ITERS * K_BLOCKS_PER_HALF): + box = kb // K_BLOCKS_PER_HALF + within = kb % K_BLOCKS_PER_HALF + desc_a = desc_a_base + ( + stage * cutlass.Int32(STAGE_A) + cutlass.Int32(box * A_BOX + within * STEP) + ) + desc_b = desc_b_base + ( + stage * cutlass.Int32(STAGE_B) + cutlass.Int32(box * B_BOX + within * STEP) + ) + accumulate = cutlass.Boolean(True) + if cutlass.const_expr(kb == 0): + accumulate = j > cutlass.Int32(0) + if prims.elect_sync(): + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, + prims.CTAGroup.CTA_1, + tmem_acc, + desc_a, + desc_b, + idesc, + accumulate, + ) + if prims.elect_sync(): + prims.tcgen05_commit(empty.subview(stage)) + used = used + cutlass.Int32(1) + if prims.elect_sync(): + prims.tcgen05_commit(acc_full.subview(buf)) + count = count + cutlass.Int32(1) + slot = count % cutlass.Int32(UNIT_RING) + while not cute.arch.mbarrier_try_wait( + unit_ready.subview(slot).data_ptr(), + (count // cutlass.Int32(UNIT_RING)) & cutlass.Int32(1), + ): + pass + unit = unit_slot.load(idx=slot) + elif warp_id >= 4: + # ===================================================================== + # Epilogue: TMEM -> registers -> partial + count (or the whole sum when + # K is one chunk) -> the finalizer's ordered sum -> bf16 -> y. + # warp 4 + w reads TMEM lanes 32 w .. 32 w + 31 = tile rows 32 w + lane. + # ===================================================================== + tid = tx - cutlass.Int32(EPI_THREADS) + lane = tx % 32 + w = warp_id - 4 + row = w * cutlass.Int32(32) + lane + tmem_base = tmem_ptr_i32.load() + prims.griddepcontrol(prims.GridDepAction.WAIT) + count = cutlass.Int32(0) + while not cute.arch.mbarrier_try_wait(unit_ready.subview(0).data_ptr(), 0): + pass + unit = unit_slot.load(idx=0) + while unit >= cutlass.Int32(0): + buf = count & cutlass.Int32(1) + while not cute.arch.mbarrier_try_wait( + acc_full.subview(buf).data_ptr(), (count >> cutlass.Int32(1)) & cutlass.Int32(1) + ): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc = prims.tcgen05_ld( + "32x32b", + cutlass.inttoptr( + tmem_base + + ((w * cutlass.Int32(32)) << cutlass.Int32(16)) + + buf * cutlass.Int32(MMA_N), + 6, + cutlass.Float32, + ), + num=MMA_N, + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + if prims.elect_sync(): + prims.mbarrier_arrive(acc_empty.subview(buf)) + vals = [cutlass.Float32(acc[t]) for t in range(MMA_N)] + tile = unit // cutlass.Int32(chunks) + if cutlass.const_expr(chunks == 1): + _store_tile(stage_out, y, vals, row, tid, tile, n_out) + else: + chunk = unit % cutlass.Int32(chunks) + part = (unit * cutlass.Int32(CTA_M) + row) * cutlass.Int32(MMA_N) + ws.store(cutlass.Vector.from_elements((vals[0], vals[1], vals[2], vals[3]), cutlass.Float32), idx=part, + vector_size=4, alignment=16) # fmt: skip + ws.store(cutlass.Vector.from_elements((vals[4], vals[5], vals[6], vals[7]), cutlass.Float32), + idx=part + cutlass.Int32(4), vector_size=4, alignment=16) # fmt: skip + # Every epilogue thread's partial is stored; thread 0's release covers them. + prims.barrier_cta_sync(1, thread_count=EPI_THREADS) + if tid == cutlass.Int32(0): + done = _atomic_add_acq_rel( + cnt.subview(tile).data_ptr().toint(), cutlass.Int32(1) + ) + last = cutlass.Int32( + cutlass.select_( + done == cutlass.Int32(chunks - 1), cutlass.Int32(1), cutlass.Int32(0) + ) + ) + s_last.store(last, idx=0) + prims.barrier_cta_sync(1, thread_count=EPI_THREADS) + if s_last.load(idx=0) != cutlass.Int32(0): + # The tile's partials in chunk order (this unit's own from registers). Volatile loads: other + # SMs wrote them; thread 0's acquire and the barrier above order these reads after the writes. + total = [] + for c in cutlass.range_constexpr(chunks): + src = ( + (tile * cutlass.Int32(chunks) + cutlass.Int32(c)) * cutlass.Int32(CTA_M) + + row + ) * cutlass.Int32(MMA_N) + lo = ws.load(idx=src, vector_size=4, alignment=16, is_volatile=True) + hi = ws.load( + idx=src + cutlass.Int32(4), + vector_size=4, + alignment=16, + is_volatile=True, + ) + mine = chunk == cutlass.Int32(c) + for t in cutlass.range_constexpr(MMA_N): + got = cutlass.Float32(lo[t]) if t < 4 else cutlass.Float32(hi[t - 4]) + got = cutlass.Float32(cutlass.select_(mine, vals[t], got)) + if cutlass.const_expr(c == 0): + total.append(got) + else: + total[t] = _add_rn(total[t], got) + if tid == cutlass.Int32(0): + cnt.store(cutlass.Int32(0), idx=tile) + _store_tile(stage_out, y, total, row, tid, tile, n_out) + count = count + cutlass.Int32(1) + slot = count % cutlass.Int32(UNIT_RING) + while not cute.arch.mbarrier_try_wait( + unit_ready.subview(slot).data_ptr(), + (count // cutlass.Int32(UNIT_RING)) & cutlass.Int32(1), + ): + pass + unit = unit_slot.load(idx=slot) + prims.barrier_cta_sync(1, thread_count=EPI_THREADS) + if warp_id == 4: + prims.tcgen05_dealloc(cutlass.inttoptr(tmem_base, 6, cutlass.Int32), TMEM_COLS) + + +# ============================================================================= +# Stream-K variant: the (tile, k-tile) space flattened tile-major and cut into GRID equal ranges, CTA c owning +# [c T / G, (c + 1) T / G). A CTA's range is a sequence of segments (one tile's contiguous k-tiles each); a tile +# split over several CTAs ("pieces", in k order = CTA order) is combined by the CTA of its piece 0 (the one holding +# k-tile 0), which is the end of that CTA's range: every other piece starts its CTA's range (or is all of it), so +# it is done early. A piece p > 0 stores its fp32 partial to ws[tile][p] and raises flag[tile][p] (fence.acq_rel + +# a relaxed store after a barrier of the epilogue warps); the finalizer acquires those flags and loads the partials +# while its own MMAs still run, then adds them after its accumulator in k order, rounds once, stores y and lowers +# the flags. The partition is fixed, so the sums are too. No claims: every warp walks the same static sequence, and +# the whole ring is loaded before the grid-dependency wait (optionally the next k-tiles are prefetched into L2). +# ============================================================================= +def streamk_cta_of(q: int, total: int, grid: int) -> int: + """The CTA whose range holds flat k-tile q (ranges [c T / G, (c + 1) T / G), some empty when G > T).""" + return ((q + 1) * grid - 1) // total + + +def streamk_max_pieces(n_out: int, k_in: int, grid: int) -> int: + """The most CTAs any tile is split over.""" + k_tiles = num_k_tiles(k_in) + tiles = n_out // CTA_M + total = tiles * k_tiles + return max( + streamk_cta_of((t + 1) * k_tiles - 1, total, grid) + - streamk_cta_of(t * k_tiles, total, grid) + + 1 + for t in range(tiles) + ) + + +def streamk_supports(n_out: int, k_in: int, ring: int) -> bool: + return ( + n_out > 0 + and n_out % CTA_M == 0 + and k_in > 0 + and k_in % CTA_K == 0 + and 1 <= ring <= MAX_RING + and (smem_bytes(ring) <= SMEM_BYTES) + ) + + +@cute.kernel +def k3_head_gemv_sk_kernel( + tma_desc_w: cutlass.GridConstant[cuda.TensorMap], # W [N, K] bf16, 5-D, one call per k-tile + tma_desc_x: cutlass.GridConstant[cuda.TensorMap], # x [M, K] bf16, box 64 x 8 + y: cutlass.Array, # [8 * N] bf16, token-major (all 8 token rows are written) + ws: cutlass.Array, # fp32 [tiles * max_pieces * 128 * 8]: the pieces' partials + flags: cutlass.Array, # int32 [tiles * max_pieces]: piece p > 0 of the tile stored its partial (0 between launches) + num_tokens: cutlass.Int32, + keep_tiles: cutlass.Int32, # tiles [0, keep_tiles) load at normal L2 priority, the rest EVICT_FIRST + n_out: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + ring: cutlass.Constexpr[int], + grid: cutlass.Constexpr[int], + max_pieces: cutlass.Constexpr[int], + prefetch: cutlass.Constexpr[ + int + ], # k-tiles after the ring prefetched into L2 before the grid-dependency wait +): + """Stream-K: this CTA's equal share of the flat (tile, k-tile) space; split tiles combined by their last piece.""" + k_tiles = num_k_tiles(k_in) + total = (n_out // CTA_M) * k_tiles + tx, _, _ = cute.arch.thread_idx() + bx, _, _ = cute.arch.block_idx() + warp_id = cute.arch.warp_idx() + tma_ptr_w = tma_desc_w.get_ptr() + tma_ptr_x = tma_desc_x.get_ptr() + lo = (bx * cutlass.Int32(total)) // cutlass.Int32(grid) + hi = ((bx + cutlass.Int32(1)) * cutlass.Int32(total)) // cutlass.Int32(grid) + count = hi - lo + + smem_a = cutlass.Array( + io_dtype, ring * CTA_M * CTA_K, space=cutlass.AddressSpace.smem, alignment=1024 + ) + smem_b = cutlass.Array( + io_dtype, ring * MMA_N * CTA_K, space=cutlass.AddressSpace.smem, alignment=1024 + ) + stage_out = cutlass.Array( + io_dtype, MMA_N * CTA_M, space=cutlass.AddressSpace.smem, alignment=16 + ) + full = cutlass.Array(cutlass.Int64, ring, space=cutlass.AddressSpace.smem, alignment=8) + empty = cutlass.Array(cutlass.Int64, ring, space=cutlass.AddressSpace.smem, alignment=8) + acc_full = cutlass.Array(cutlass.Int64, 2, space=cutlass.AddressSpace.smem, alignment=8) + acc_empty = cutlass.Array(cutlass.Int64, 2, space=cutlass.AddressSpace.smem, alignment=8) + tmem_ptr_i32 = cutlass.Array(cutlass.Int32, 1, space=cutlass.AddressSpace.smem) + + if warp_id == 0: + prims.prefetch_tensormap(tma_ptr_w) + prims.prefetch_tensormap(tma_ptr_x) + if prims.elect_sync(): + for s in cutlass.range_constexpr(ring): + prims.mbarrier_init(full.subview(s), 1) + prims.mbarrier_init(empty.subview(s), 1) + for b in cutlass.range_constexpr(2): + prims.mbarrier_init(acc_full.subview(b), 1) + prims.mbarrier_init(acc_empty.subview(b), 4) # one elected lane per epilogue warp + if warp_id == 2: + prims.tcgen05_alloc(tmem_ptr_i32, TMEM_COLS) + prims.tcgen05_relinquish_alloc_permit() + prims.fence_mbarrier_init() + prims.barrier_cta_sync(0) + + if warp_id == 0: + # ===================================================================== + # TMA: the range's first RING k-tiles' weight before the wait, their + # activation after it, then every later k-tile when its stage frees. + # ===================================================================== + if prims.elect_sync(): + for jj in cutlass.range_constexpr(ring): + if cutlass.Int32(jj) < count: + q0 = lo + cutlass.Int32(jj) + t0 = q0 // cutlass.Int32(k_tiles) + prims.mbarrier_arrive_expect_tx(full.subview(jj), A_BYTES + B_BYTES) + if t0 < keep_tiles: + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a.subview(jj * CTA_M * CTA_K), + tma_ptr_w, + ( + cutlass.Int32(0), + t0 * cutlass.Int32(CTA_M), + (q0 % cutlass.Int32(k_tiles)) * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), + cutlass.Int32(0), + ), + full.subview(jj), + ) + else: + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a.subview(jj * CTA_M * CTA_K), + tma_ptr_w, + ( + cutlass.Int32(0), + t0 * cutlass.Int32(CTA_M), + (q0 % cutlass.Int32(k_tiles)) * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), + cutlass.Int32(0), + ), + full.subview(jj), + l2_cache_hint=EVICT_FIRST, + ) + # The next k-tiles of the range into L2 while the predecessor runs (a weight; the loads below hit L2). + for jp in cutlass.range_constexpr(ring, ring + prefetch): + if cutlass.Int32(jp) < count: + qp = lo + cutlass.Int32(jp) + prims.cp_async_bulk_tensor_prefetch( + tma_ptr_w, + [ + cutlass.Int32(0), + (qp // cutlass.Int32(k_tiles)) * cutlass.Int32(CTA_M), + (qp % cutlass.Int32(k_tiles)) * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), + cutlass.Int32(0), + ], + [], # tile mode: no im2col offsets + ) + prims.griddepcontrol(prims.GridDepAction.WAIT) + for jj in cutlass.range_constexpr(ring): + if cutlass.Int32(jj) < count: + kb0 = (lo + cutlass.Int32(jj)) % cutlass.Int32(k_tiles) + for half in cutlass.range_constexpr(TMA_COPY_ITERS): + prims.cp_async_bulk_tensor_shared_cta_global( + smem_b.subview(jj * MMA_N * CTA_K + half * B_HALF_ELEMS), + tma_ptr_x, + (kb0 * cutlass.Int32(CTA_K) + cutlass.Int32(half * TMA_K_BOX), cutlass.Int32(0)), + full.subview(jj), + ) # fmt: skip + n = cutlass.Int32(ring) + while n < count: + stage = n % cutlass.Int32(ring) + # The stage's previous MMAs committed. + while not cute.arch.mbarrier_try_wait( + empty.subview(stage).data_ptr(), + ((n // cutlass.Int32(ring)) & cutlass.Int32(1)) ^ cutlass.Int32(1), + ): + pass + q = lo + n + tile = q // cutlass.Int32(k_tiles) + k = q % cutlass.Int32(k_tiles) + prims.mbarrier_arrive_expect_tx(full.subview(stage), A_BYTES + B_BYTES) + if tile < keep_tiles: + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a.subview(stage * cutlass.Int32(CTA_M * CTA_K)), tma_ptr_w, + (cutlass.Int32(0), tile * cutlass.Int32(CTA_M), k * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), cutlass.Int32(0)), + full.subview(stage), + ) # fmt: skip + else: + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a.subview(stage * cutlass.Int32(CTA_M * CTA_K)), tma_ptr_w, + (cutlass.Int32(0), tile * cutlass.Int32(CTA_M), k * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), cutlass.Int32(0)), + full.subview(stage), l2_cache_hint=EVICT_FIRST, + ) # fmt: skip + for half in cutlass.range_constexpr(TMA_COPY_ITERS): + prims.cp_async_bulk_tensor_shared_cta_global( + smem_b.subview(stage * cutlass.Int32(MMA_N * CTA_K) + cutlass.Int32(half * B_HALF_ELEMS)), + tma_ptr_x, + (k * cutlass.Int32(CTA_K) + cutlass.Int32(half * TMA_K_BOX), cutlass.Int32(0)), + full.subview(stage), + ) # fmt: skip + n = n + cutlass.Int32(1) + # All of this CTA's loads are issued: the dependents may launch once every CTA gets here. + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + elif warp_id == 2: + # ===================================================================== + # MMA: each segment's k-tiles into accumulator (segment count) % 2. + # ===================================================================== + idesc = prims.Tcgen05InstrDesc.build( + c_dtype=cutlass.Float32, a_dtype=io_dtype, b_dtype=io_dtype, n_dim=MMA_N, m_dim=CTA_M + ) + desc_a_base = prims.Tcgen05SmemDesc.build( + start_address=smem_a, leading_byte_offset=LEADING, stride_byte_offset=STRIDE, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + desc_b_base = prims.Tcgen05SmemDesc.build( + start_address=smem_b, leading_byte_offset=LEADING, stride_byte_offset=STRIDE, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + tmem_base = tmem_ptr_i32.load() + used = cutlass.Int32(0) + seg = cutlass.Int32(0) + q = lo + while q < hi: + seg_tile = q // cutlass.Int32(k_tiles) + seg_end = cutlass.Int32( + cutlass.select_( + (seg_tile + cutlass.Int32(1)) * cutlass.Int32(k_tiles) < hi, + (seg_tile + cutlass.Int32(1)) * cutlass.Int32(k_tiles), + hi, + ) + ) + buf = seg & cutlass.Int32(1) + if seg >= cutlass.Int32(2): + # The epilogue has read this accumulator's previous segment. + while not cute.arch.mbarrier_try_wait( + acc_empty.subview(buf).data_ptr(), + ((seg >> cutlass.Int32(1)) - cutlass.Int32(1)) & cutlass.Int32(1), + ): + pass + tmem_acc = cutlass.inttoptr(tmem_base + buf * cutlass.Int32(MMA_N), 6, cutlass.Int32) + j = q + while j < seg_end: + stage = used % cutlass.Int32(ring) + while not cute.arch.mbarrier_try_wait( + full.subview(stage).data_ptr(), (used // cutlass.Int32(ring)) & cutlass.Int32(1) + ): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + for kb in cutlass.range_constexpr(TMA_COPY_ITERS * K_BLOCKS_PER_HALF): + box = kb // K_BLOCKS_PER_HALF + within = kb % K_BLOCKS_PER_HALF + desc_a = desc_a_base + ( + stage * cutlass.Int32(STAGE_A) + cutlass.Int32(box * A_BOX + within * STEP) + ) + desc_b = desc_b_base + ( + stage * cutlass.Int32(STAGE_B) + cutlass.Int32(box * B_BOX + within * STEP) + ) + accumulate = cutlass.Boolean(True) + if cutlass.const_expr(kb == 0): + accumulate = j > q + if prims.elect_sync(): + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, + prims.CTAGroup.CTA_1, + tmem_acc, + desc_a, + desc_b, + idesc, + accumulate, + ) + if prims.elect_sync(): + prims.tcgen05_commit(empty.subview(stage)) + used = used + cutlass.Int32(1) + j = j + cutlass.Int32(1) + if prims.elect_sync(): + prims.tcgen05_commit(acc_full.subview(buf)) + seg = seg + cutlass.Int32(1) + q = seg_end + elif warp_id >= 4: + # ===================================================================== + # Epilogue: per segment, TMEM -> registers -> y (a whole tile), or piece + # p > 0's partial + flag, or piece 0's ordered sum of all pieces -> y. + # ===================================================================== + tid = tx - cutlass.Int32(EPI_THREADS) + lane = tx % 32 + w = warp_id - 4 + row = w * cutlass.Int32(32) + lane + tmem_base = tmem_ptr_i32.load() + prims.griddepcontrol(prims.GridDepAction.WAIT) + seg = cutlass.Int32(0) + q = lo + while q < hi: + seg_tile = q // cutlass.Int32(k_tiles) + seg_end = cutlass.Int32( + cutlass.select_( + (seg_tile + cutlass.Int32(1)) * cutlass.Int32(k_tiles) < hi, + (seg_tile + cutlass.Int32(1)) * cutlass.Int32(k_tiles), + hi, + ) + ) + buf = seg & cutlass.Int32(1) + acc_addr = ( + tmem_base + + ((w * cutlass.Int32(32)) << cutlass.Int32(16)) + + buf * cutlass.Int32(MMA_N) + ) + acc_parity = (seg >> cutlass.Int32(1)) & cutlass.Int32(1) + # The tile's pieces: the CTAs whose ranges hold its first and last k-tiles, and everything between. + first_cta = ( + (seg_tile * cutlass.Int32(k_tiles) + cutlass.Int32(1)) * cutlass.Int32(grid) + - cutlass.Int32(1) + ) // cutlass.Int32(total) + last_cta = ( + (seg_tile + cutlass.Int32(1)) * cutlass.Int32(k_tiles) * cutlass.Int32(grid) + - cutlass.Int32(1) + ) // cutlass.Int32(total) + pieces = last_cta - first_cta + cutlass.Int32(1) + piece = bx - first_cta + slot0 = seg_tile * cutlass.Int32(max_pieces) + if pieces == cutlass.Int32(1): + while not cute.arch.mbarrier_try_wait(acc_full.subview(buf).data_ptr(), acc_parity): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc = prims.tcgen05_ld( + "32x32b", cutlass.inttoptr(acc_addr, 6, cutlass.Float32), num=MMA_N + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + if prims.elect_sync(): + prims.mbarrier_arrive(acc_empty.subview(buf)) + _store_tile( + stage_out, + y, + [cutlass.Float32(acc[t]) for t in range(MMA_N)], + row, + tid, + seg_tile, + n_out, + ) + elif piece == cutlass.Int32(0): + # The finalizer. The other pieces' partials first, while this segment's MMAs run: thread 0 + # acquires the flags (backing off between polls, so the epilogue warps do not load the memory + # system while the weight streams), the barrier orders every thread's loads after its acquires, + # then each thread loads its row (a slot past the piece count reloads piece 1's and is not added). + if tid == cutlass.Int32(0): + for p in cutlass.range_constexpr(1, max_pieces): + if cutlass.Int32(p) < pieces: + while _ld_acquire( + flags.subview(slot0 + cutlass.Int32(p)).data_ptr().toint() + ) == cutlass.Int32(0): + prims.nanosleep(FLAG_BACKOFF_NS) + prims.barrier_cta_sync(1, thread_count=EPI_THREADS) + pre = [] + for p in cutlass.range_constexpr(1, max_pieces): + live = cutlass.Int32(p) < pieces + src = ( + ( + slot0 + + cutlass.Int32( + cutlass.select_(live, cutlass.Int32(p), cutlass.Int32(1)) + ) + ) + * cutlass.Int32(CTA_M) + + row + ) * cutlass.Int32(MMA_N) + pre.append( + ( + ws.load(idx=src, vector_size=4, alignment=16, is_volatile=True), + ws.load( + idx=src + cutlass.Int32(4), + vector_size=4, + alignment=16, + is_volatile=True, + ), + ) + ) + while not cute.arch.mbarrier_try_wait(acc_full.subview(buf).data_ptr(), acc_parity): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc = prims.tcgen05_ld( + "32x32b", cutlass.inttoptr(acc_addr, 6, cutlass.Float32), num=MMA_N + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + if prims.elect_sync(): + prims.mbarrier_arrive(acc_empty.subview(buf)) + # Piece 0 (k from 0), then pieces 1, 2, ... in k order. + total_v = [cutlass.Float32(acc[t]) for t in range(MMA_N)] + for p in cutlass.range_constexpr(1, max_pieces): + live = cutlass.Int32(p) < pieces + lo_v, hi_v = pre[p - 1] + for t in cutlass.range_constexpr(MMA_N): + got = cutlass.Float32(lo_v[t]) if t < 4 else cutlass.Float32(hi_v[t - 4]) + total_v[t] = cutlass.Float32( + cutlass.select_(live, _add_rn(total_v[t], got), total_v[t]) + ) + _store_tile(stage_out, y, total_v, row, tid, seg_tile, n_out) + # Every epilogue thread's acquires are behind the store helper's barriers: lower the flags. + if tid == cutlass.Int32(0): + for p in cutlass.range_constexpr(1, max_pieces): + if cutlass.Int32(p) < pieces: + flags.store(cutlass.Int32(0), idx=slot0 + cutlass.Int32(p)) + else: + while not cute.arch.mbarrier_try_wait(acc_full.subview(buf).data_ptr(), acc_parity): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc = prims.tcgen05_ld( + "32x32b", cutlass.inttoptr(acc_addr, 6, cutlass.Float32), num=MMA_N + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + if prims.elect_sync(): + prims.mbarrier_arrive(acc_empty.subview(buf)) + part = ((slot0 + piece) * cutlass.Int32(CTA_M) + row) * cutlass.Int32(MMA_N) + ws.store( + cutlass.Vector.from_elements( + ( + cutlass.Float32(acc[0]), + cutlass.Float32(acc[1]), + cutlass.Float32(acc[2]), + cutlass.Float32(acc[3]), + ), + cutlass.Float32, + ), + idx=part, + vector_size=4, + alignment=16, + ) + ws.store( + cutlass.Vector.from_elements( + ( + cutlass.Float32(acc[4]), + cutlass.Float32(acc[5]), + cutlass.Float32(acc[6]), + cutlass.Float32(acc[7]), + ), + cutlass.Float32, + ), + idx=part + cutlass.Int32(4), + vector_size=4, + alignment=16, + ) + # Every epilogue thread's partial is stored; thread 0's release covers them. + prims.barrier_cta_sync(1, thread_count=EPI_THREADS) + if tid == cutlass.Int32(0): + _flag_release(flags.subview(slot0 + piece).data_ptr().toint()) + seg = seg + cutlass.Int32(1) + q = seg_end + prims.barrier_cta_sync(1, thread_count=EPI_THREADS) + if warp_id == 4: + prims.tcgen05_dealloc(cutlass.inttoptr(tmem_base, 6, cutlass.Int32), TMEM_COLS) + + +def _weight_tensor_map(w, n_out, k_in): + """W as five TMA dimensions (64-element column chunk, row, 64-element chunk index, 1, 1) so one call + per k-tile lands both 128-byte-swizzled halves; strides in 16-byte units.""" + return cuda.create_tensor_map_tiled( + global_address=w.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[TMA_K_BOX, n_out, k_in // TMA_K_BOX, 1, 1], + global_strides=[ + (k_in * ELEM_BYTES) // 16, + (TMA_K_BOX * ELEM_BYTES) // 16, + (n_out * k_in * ELEM_BYTES) // 16, + (n_out * k_in * ELEM_BYTES) // 16, + ], + box_dims=[TMA_K_BOX, CTA_M, TMA_COPY_ITERS, 1, 1], + swizzle=cuda.TensorMapSwizzle.s128b, + ) + + +def _activation_tensor_map(x, cols, num_tokens): + """x [M, cols] (rows dense) as (cols, M) with an 8-row box: rows past num_tokens arrive as zeros.""" + return cuda.create_tensor_map_tiled( + global_address=x.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[cols, num_tokens], + global_strides=[(cols * ELEM_BYTES) // 16], + box_dims=[TMA_K_BOX, MMA_N], + swizzle=cuda.TensorMapSwizzle.s128b, + ) + + +@cute.jit +def k3_head_gemv( + w: cute.Tensor, # [N, K] bf16, K contiguous + x: cute.Tensor, # [M, K] bf16, K contiguous, M <= 8 + y: cute.Tensor, # [8 * N] bf16 + ws: cute.Tensor, # fp32 [units * 128 * 8] + cnt: cute.Tensor, # int32 [tiles], zero + claim: cute.Tensor, # int32 [1], zero + num_tokens: cutlass.Int32, + keep_tiles: cutlass.Int32, + n_out: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + chunk_tiles: cutlass.Constexpr[int], + ring: cutlass.Constexpr[int], + grid: cutlass.Constexpr[int], + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + """``y = x @ w^T`` over ``grid`` persistent CTAs.""" + tma_desc_w = _weight_tensor_map(w, n_out, k_in) + tma_desc_x = _activation_tensor_map(x, k_in, num_tokens) + k3_head_gemv_kernel( + tma_desc_w, + tma_desc_x, + y, + ws, + cnt, + claim, + num_tokens, + keep_tiles, + n_out, + k_in, + chunk_tiles, + ring, + grid, + ).launch( + grid=[grid, 1, 1], + block=[THREADS, 1, 1], + stream=stream, + use_pdl=use_pdl, + ) + + +@cute.jit +def k3_head_gemv_sk( + w: cute.Tensor, # [N, K] bf16, K contiguous + x: cute.Tensor, # [M, K] bf16, K contiguous, M <= 8 + y: cute.Tensor, # [8 * N] bf16 + ws: cute.Tensor, # fp32 [tiles * max_pieces * 128 * 8] + flags: cute.Tensor, # int32 [tiles * max_pieces], zero + num_tokens: cutlass.Int32, + keep_tiles: cutlass.Int32, + n_out: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + ring: cutlass.Constexpr[int], + grid: cutlass.Constexpr[int], + max_pieces: cutlass.Constexpr[int], + prefetch: cutlass.Constexpr[int], + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + """``y = x @ w^T`` over ``grid`` stream-K CTAs.""" + tma_desc_w = _weight_tensor_map(w, n_out, k_in) + tma_desc_x = _activation_tensor_map(x, k_in, num_tokens) + k3_head_gemv_sk_kernel( + tma_desc_w, + tma_desc_x, + y, + ws, + flags, + num_tokens, + keep_tiles, + n_out, + k_in, + ring, + grid, + max_pieces, + prefetch, + ).launch( + grid=[grid, 1, 1], + block=[THREADS, 1, 1], + stream=stream, + use_pdl=use_pdl, + ) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_head_gemv/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_head_gemv/op.py new file mode 100644 index 000000000000..0b2b5ada30b7 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_head_gemv/op.py @@ -0,0 +1,198 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``trtllm::k3_head_gemv``: ``x @ weight^T`` for a large weight (an lm_head vocab shard) and M <= 8 rows. + +One CTA per SM, and under stream-K at most one per (tile, k-tile), streams the weight, EVICT_FIRST except for the +first ``keep_tiles`` 128-row tiles, and adds each tile's k-split partials in a fixed order before one bf16 rounding: +the result is bit-identical from run to run (not bit-identical to cuBLAS or pdl_gemv, whose accumulation orders +differ). Two schedules: ``streamk`` (every CTA an equal share of the flat (tile, k-tile) space; the default) and +``dynamic`` ((tile, k-chunk) units claimed from a counter, ``chunk_tiles`` k-tiles each). + +The op keeps one workspace per device and shape (the partials, a per-tile count and the unit counter, the counters +zero between launches), so calls of one shape must be ordered on one stream. Each kernel compiles on the first +call for its shape, which must happen outside CUDA-graph capture. +""" + +from __future__ import annotations + +import os +import threading +from typing import Dict + +import torch + +MAX_TOKENS = 8 +UNITS_PER_CTA = 4 # the chunk choice keeps at least this many units per CTA for the dynamic balance + +_lock = threading.Lock() +_compiled: Dict[tuple, object] = {} +_workspaces: Dict[tuple, tuple] = {} + + +def _arg(t: torch.Tensor): + from cutlass.cute.runtime import from_dlpack + + # detach(): DLPack refuses tensors that require grad (weights are parameters). + return from_dlpack(t.detach(), assumed_align=16).mark_layout_dynamic(leading_dim=t.dim() - 1) + + +def _use_pdl() -> bool: + return os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" + + +def _stream(t: torch.Tensor): + import cuda.bindings.driver as cuda_driver + + return cuda_driver.CUstream(torch.cuda.current_stream(t.device).cuda_stream) + + +def _grid(device) -> int: + return torch.cuda.get_device_properties(device).multi_processor_count + + +def pick_chunk(n_out: int, k_in: int, grid: int) -> int: + """k-tiles per unit: the largest divisor of the k-tiles that leaves at least UNITS_PER_CTA units per CTA, or the + smallest divisor above 1 when none does (every divisor if the k-tiles are prime).""" + from . import k3_head_gemv_kernel as kernel + + k_tiles = kernel.num_k_tiles(k_in) + tiles = n_out // kernel.CTA_M + divisors = [d for d in range(1, k_tiles + 1) if k_tiles % d == 0] + fitting = [d for d in divisors if tiles * (k_tiles // d) >= UNITS_PER_CTA * grid] + if fitting: + return max(fitting) + return min([d for d in divisors if d > 1] or divisors) + + +def supports( + x: torch.Tensor, + weight: torch.Tensor, + chunk_tiles: int = 0, + ring: int = 6, + schedule: str = "streamk", +) -> bool: + """Whether ``k3_head_gemv`` runs ``x @ weight^T`` (``chunk_tiles`` 0: the op's choice).""" + from . import k3_head_gemv_kernel as kernel + + if not ( + x.is_cuda + and x.dtype == torch.bfloat16 + and weight.dtype == torch.bfloat16 + and x.dim() == 2 + and weight.dim() == 2 + and 0 < x.shape[0] <= MAX_TOKENS + and x.shape[1] == weight.shape[1] + and x.is_contiguous() + and weight.is_contiguous() + and weight.shape[0] % kernel.CTA_M == 0 + and weight.shape[1] % kernel.CTA_K == 0 + ): + return False + n_out, k_in = weight.shape + if schedule == "streamk": + return kernel.streamk_supports(n_out, k_in, ring) + if schedule != "dynamic": + return False + chunk = chunk_tiles or pick_chunk(n_out, k_in, _grid(x.device)) + return kernel.supports(n_out, k_in, chunk, ring) + + +def _workspace(device, units: int, tiles: int): + """(partials fp32 [units * 128 * 8], ``tiles`` count / flag words, claim counter), the counters zero.""" + key = (device.index, units, tiles) + ws = _workspaces.get(key) + if ws is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "k3_head_gemv: the workspace must be allocated outside CUDA-graph capture" + ) + ws = ( + torch.empty(max(units, 1) * 128 * 8, dtype=torch.float32, device=device), + torch.zeros(tiles, dtype=torch.int32, device=device), + torch.zeros(1, dtype=torch.int32, device=device), + ) + _workspaces[key] = ws + return ws + + +def _compiled_fn(key, entry, *args): + fn = _compiled.get(key) + if fn is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "k3_head_gemv: run once per shape outside CUDA-graph capture first (it compiles)." + ) + import cutlass.cute as cute + + with _lock: + fn = _compiled.get(key) + if fn is None: + fn = _compiled[key] = cute.compile(entry, *args) + return fn + + +@torch.library.custom_op("trtllm::k3_head_gemv", mutates_args=()) +def k3_head_gemv( + x: torch.Tensor, + weight: torch.Tensor, + keep_tiles: int = 0, + chunk_tiles: int = 0, + ring: int = 6, + schedule: str = "streamk", + prefetch: int = 16, +) -> torch.Tensor: + """``x @ weight.T`` for bf16 ``x`` [M <= 8, K] and ``weight`` [N, K] (N and K multiples of 128); returns bf16 + [M, N]. Weight tiles (128 rows) below ``keep_tiles`` are read at normal L2 priority, the rest EVICT_FIRST. + ``schedule``: ``streamk`` or ``dynamic``; ``chunk_tiles`` (dynamic; 0: the op's choice) sets the k-tiles per + work unit; ``prefetch`` (stream-K) the k-tiles after the ring that each CTA prefetches into L2 before the grid + dependency.""" + if not supports(x, weight, chunk_tiles, ring, schedule): + raise ValueError( + f"k3_head_gemv: unsupported call x {tuple(x.shape)} {x.dtype}, weight {tuple(weight.shape)} " + f"{weight.dtype}, chunk_tiles {chunk_tiles}, ring {ring}, schedule {schedule}" + ) + from . import k3_head_gemv_kernel as kernel + + num_tokens, k_in = x.shape + n_out = weight.shape[0] + tiles = n_out // kernel.CTA_M + grid = _grid(x.device) + # All 8 token rows are written (the rows past M from TMA's zero-filled x rows); the result is the first M. + y = torch.empty(kernel.MMA_N, n_out, dtype=torch.bfloat16, device=x.device) + stream = _stream(x) + use_pdl = _use_pdl() + if schedule == "streamk": + # At most one CTA per (tile, k-tile): every CTA's share is then non-empty, so each CTA between a split tile's + # first and last piece holds part of that tile and raises the flag its finalizer waits for. + grid = min(grid, tiles * kernel.num_k_tiles(k_in)) + pieces = kernel.streamk_max_pieces(n_out, k_in, grid) + ws, flags, _ = _workspace(x.device, tiles * pieces, tiles * pieces) + args = (_arg(weight), _arg(x), _arg(y.view(-1)), _arg(ws), _arg(flags)) + fn = _compiled_fn(("k3_head_gemv_sk", n_out, k_in, ring, grid, pieces, prefetch, use_pdl), + kernel.k3_head_gemv_sk, *args, num_tokens, keep_tiles, n_out, k_in, ring, grid, pieces, + prefetch, use_pdl, stream) # fmt: skip + else: + chunk = chunk_tiles or pick_chunk(n_out, k_in, grid) + ws, cnt, claim = _workspace(x.device, kernel.num_units(n_out, k_in, chunk), tiles) + args = (_arg(weight), _arg(x), _arg(y.view(-1)), _arg(ws), _arg(cnt), _arg(claim)) + fn = _compiled_fn(("k3_head_gemv", n_out, k_in, chunk, ring, grid, use_pdl), kernel.k3_head_gemv, *args, + num_tokens, keep_tiles, n_out, k_in, chunk, ring, grid, use_pdl, stream) # fmt: skip + fn(*args, num_tokens, keep_tiles, stream) + return y[:num_tokens] + + +@k3_head_gemv.register_fake +def _(x, weight, keep_tiles=0, chunk_tiles=0, ring=6, schedule="streamk", prefetch=16): + return x.new_empty((x.shape[0], weight.shape[0]), dtype=torch.bfloat16) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_cluster_waits.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_cluster_waits.py new file mode 100644 index 000000000000..f9433e11a165 --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_cluster_waits.py @@ -0,0 +1,97 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Source check of the mailbox waits in Kimi K3 kernels whose barriers other CTAs complete with st.async: the +complete_tx releases at cluster scope, so the waiting CTA must acquire at cluster scope +(``mbarrier.{test,try}_wait.parity.acquire.cluster``, the kernels' ``_test_wait_cluster`` / ``_try_wait_cluster``); +the DSL's ``mbarrier_test_wait`` / ``mbarrier_try_wait`` acquire at CTA scope, which the PTX memory model does not let +synchronize with another CTA's release. + + pytest test_k3_cluster_waits.py +""" + +import importlib.util +import re + +import pytest + +# Kernel module -> the barriers in it that other CTAs complete with st.async. +MAILBOXES = { + "tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv.k3_ctm_gemv_kernel": ("mail_full",), +} +# Kernel module -> the header of the blocks whose mailbox waits are on this CTA's own arrival and keep CTA scope (a +# kernel that arrives on its mailbox itself and acquires the other CTAs' data through a counter). +OWN_ARRIVAL = {} +CTA_WAIT = re.compile(r"mbarrier_(test|try)_wait\(") +CLUSTER_WAIT = re.compile(r"_(test|try)_wait_cluster\(") + + +def block_header(code, i): + """The header of the block that line i is in: the nearest earlier code line indented less.""" + indent = len(code[i]) - len(code[i].lstrip()) + for ln in reversed(code[:i]): + if ln.strip() and len(ln) - len(ln.lstrip()) < indent: + return ln.strip() + return "" + + +def cta_scope_mailbox_waits(lines, mailboxes, own_arrival=None): + """Line numbers of CTA-scope waits on a mailbox barrier (the barrier named on the wait's line or the next; + a name is a regular expression), except those directly in a block whose header matches ``own_arrival``.""" + code = [ln.split("#", 1)[0] for ln in lines] + found = [] + for i, ln in enumerate(code): + if CTA_WAIT.search(ln): + window = ln + (code[i + 1] if i + 1 < len(code) else "") + if any(re.search(rf"\b{name}\b", window) for name in mailboxes): + if own_arrival is None or not re.fullmatch(own_arrival, block_header(code, i)): + found.append(i + 1) + return found + + +def waited_at_cluster_scope(lines, name): + """Whether a cluster-scope wait names the barrier (on the wait's line or the next).""" + code = [ln.split("#", 1)[0] for ln in lines] + return any( + CLUSTER_WAIT.search(ln) + and re.search(rf"\b{name}\b", ln + (code[i + 1] if i + 1 < len(code) else "")) + for i, ln in enumerate(code) + ) + + +def test_checker_catches_the_pattern(): + bad = [" while not cute.arch.mbarrier_test_wait(mail_full.data_ptr(), 0):", " pass"] + good = [" while not _test_wait_cluster(mail_full.data_ptr(), 0):", " pass"] + assert cta_scope_mailbox_waits(bad, ("mail_full",)) == [1] + assert cta_scope_mailbox_waits(good, ("mail_full",)) == [] + assert waited_at_cluster_scope(good, "mail_full") + assert not waited_at_cluster_scope(bad, "mail_full") + # Only the own-arrival block keeps CTA scope; the same wait before it or in its else branch is flagged. + own = r"if cutlass\.const_expr\(no_cluster\):" + branches = ( + bad + ["if cutlass.const_expr(no_cluster):", " # own arrival"] + bad + ["else:"] + bad + ) + assert cta_scope_mailbox_waits(branches, ("mail_full",), own) == [1, 8] + + +@pytest.mark.parametrize("module", list(MAILBOXES), ids=[m.rsplit(".", 1)[1] for m in MAILBOXES]) +def test_mailbox_waits_acquire_at_cluster_scope(module): + path = importlib.util.find_spec(module).origin + with open(path) as f: + lines = f.read().split("\n") + found = cta_scope_mailbox_waits(lines, MAILBOXES[module], OWN_ARRIVAL.get(module)) + assert not found, f"{path}: CTA-scope waits on st.async mailboxes at lines {found}" + # Every listed mailbox is still waited on, so a renamed barrier or helper cannot pass unchecked. + missing = [name for name in MAILBOXES[module] if not waited_at_cluster_scope(lines, name)] + assert not missing, f"{path}: no cluster-scope wait on {missing}" diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctm_gemv.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctm_gemv.py new file mode 100644 index 000000000000..757fbba95655 --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctm_gemv.py @@ -0,0 +1,281 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""The CTM decode GEMVs (trtllm::k3_ctm_gemv, _gated, _swiglu, _long, _tail) and trtllm::k3_situ_mul at the Kimi K3 +TP16 per-rank shapes and the call sites' split / ring / push flags, at every M in 1..8: error against an fp64 product +of the unfused activation (torch's bf16 roundings) and against the stock path (cuBLAS F.linear; the Triton SiTU), +run-to-run identical bits, each M's rows bit-identical to the same rows of the 8-row call; split 1 the bits of +k3_decode_gemv, the long GEMV's sigmoid columns torch.sigmoid of its own plain output. 0 and 9 rows are refused.""" + +import functools + +import pytest +import torch +import torch.nn.functional as F + + +def _is_sm100() -> bool: + if not torch.cuda.is_available(): + return False + major, _ = torch.cuda.get_device_capability() + return major == 10 + + +pytestmark = pytest.mark.skipif(not _is_sm100(), reason="needs an SM100-family GPU") + +M_ALL = list(range(1, 9)) +TOL = 8e-3 # max |y - ref| / max |ref| + +# k3_ctm_gemv: (N, K, split, push). The MLA o_proj shape at splits 1 and 2, the drafter o_proj (tuned and synthetic). +PLAIN = { + "o_proj_s1": (7168, 768, 1, False), + "o_proj_s2": (7168, 768, 2, False), + "drafter_o_proj": (7168, 384, 1, True), + "drafter_o_proj_dummy": (7168, 256, 1, True), +} +# k3_ctm_gemv_long: (N, K, sig_col0, split, ring, push), as each call site passes them. +LONG = { + "mla_qkv_a_gate": (2880, 7168, 2112, 6, 6, True), # [W_a; W_g], the gate columns as bf16(sigmoid) + "dense_gate_up": (4224, 7168, -1, 4, 5, False), + "dense_down": (7168, 2112, -1, 2, 6, False), # K 2112 ends in a half k-tile + "drafter_qkv": (512, 7168, -1, 8, 6, True), + "drafter_gate_up": (1792, 7168, -1, 8, 6, True), + "drafter_gate_up_dummy": (1536, 7168, -1, 8, 6, True), + "kda_qkvg": (3208, 7168, -1, 5, 6, True), # KDA q/k/v/g/f_a/b: 25 whole tiles and an 8-row last one +} +# k3_ctm_gemv_swiglu: (N, K, split, push), the drafter down projection (tuned and synthetic). +SWIGLU = {"drafter_down": (7168, 896, 2, True), "drafter_down_dummy": (7168, 768, 2, True)} +AG_COLS, G_COL0, K_O = 2880, 2112, 768 # MLA: [q_a 1536 | kv_a 576 | gate 768] rows of the fused projection +HIDDEN, LATENT, WIDTH, PAD, ACT = 7168, 3584, 224, 256, 384 +EPS = 1e-6 +SITU = {"k3": (4.0, 25.0), "plain": (1.0, None)} # (beta, linear_beta): the K3 checkpoint's, and the defaults + + +def _ops(): + import tensorrt_llm # noqa: F401 (registers the trtllm ops) + from tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv import op # noqa: F401 + from tensorrt_llm._torch.cute_dsl_kernels.k3_decode_gemv import op as _op # noqa: F401 + + return torch.ops.trtllm + + +def _ctm(): + from tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv import op + + return op + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def _rel(y: torch.Tensor, ref: torch.Tensor) -> float: + return (y.double() - ref.double()).abs().max().item() / ref.double().abs().max().item() + + +def _report(op, case, m, y, ref, stock=None, **flags): + extra = "" if stock is None else f" vs_stock={_rel(y, stock):.2e}" + marks = " ".join(f"{k}={v}" for k, v in flags.items()) + abs_err = (y.double() - ref.double()).abs().max().item() + print(f"OPCHECK op={op} case={case} M={m} abs={abs_err:.3e} rel={_rel(y, ref):.3e}{extra} {marks}") + + +@functools.lru_cache(maxsize=None) +def _weight(n: int, k: int, seed: int, scale: float = 0.02) -> torch.Tensor: + gen = torch.Generator(device="cuda").manual_seed(seed) + return (torch.randn(n, k, generator=gen, device="cuda") * scale).bfloat16() + + +@functools.lru_cache(maxsize=None) +def _rows(k: int, seed: int, scale: float = 1.0) -> torch.Tensor: + gen = torch.Generator(device="cuda").manual_seed(seed) + return (torch.randn(8, k, generator=gen, device="cuda") * scale).bfloat16() + + +def _checks(call, args8, m): + """The M-row call (the first M rows of each 8-row input), a rerun and the 8-row call.""" + args = [a[:m].contiguous() for a in args8] + y = call(*args) + again = call(*args) + y8 = call(*args8) + return args, y, torch.equal(_bits(y), _bits(again)), torch.equal(_bits(y), _bits(y8[:m])) + + +@pytest.mark.parametrize("m", M_ALL) +@pytest.mark.parametrize("name", list(PLAIN)) +def test_k3_ctm_gemv(name, m): + ops = _ops() + n, k, split, push = PLAIN[name] + w = _weight(n, k, 1, 0.03) + assert _ctm().supports(_rows(k, 2)[:m], w, split) + (x,), y, det, minv = _checks(lambda x_: ops.k3_ctm_gemv(x_, w, True, split, push), [_rows(k, 2)], m) + ref = x.double() @ w.double().t() + decode = ops.k3_decode_gemv(x, w, True) + same_decode = torch.equal(_bits(y), _bits(decode)) + _report("k3_ctm_gemv", name, m, y, ref, F.linear(x, w), det=det, rows_as_m8=minv, eq_k3_decode_gemv=same_decode) + assert _rel(y, ref) <= TOL and _rel(y, F.linear(x, w)) <= TOL + assert det and minv + if split == 1: + assert same_decode + + +@pytest.mark.parametrize("m", M_ALL) +@pytest.mark.parametrize("gate_sigmoid", [False, True], ids=["sigmoid_given", "sigmoid_in_kernel"]) +@pytest.mark.parametrize("split", [2, 1]) +def test_k3_ctm_gemv_gated(split, gate_sigmoid, m): + """The MLA o_proj with its output gate: (a * sigmoid(g)) @ W^T, g the gate columns of the fused projection rows + (or a * s with s = bf16(sigmoid(g)) already there, the in-model form), against the unfused torch roundings.""" + ops = _ops() + w = _weight(HIDDEN, K_O, 3, 0.03) + a8 = _rows(K_O, 4) + if gate_sigmoid: + ag8 = _rows(AG_COLS, 5, 2.0) + else: + gen = torch.Generator(device="cuda").manual_seed(6) + ag8 = torch.rand(8, AG_COLS, generator=gen, device="cuda").bfloat16() + assert _ctm().supports_gated(a8, ag8, G_COL0, w, split) + + def call(a_, ag_): + return ops.k3_ctm_gemv_gated(a_, ag_, G_COL0, w, True, split, gate_sigmoid) + + (a, ag), y, det, minv = _checks(call, [a8, ag8], m) + g = ag[:, G_COL0 : G_COL0 + K_O] + b = a * (g.sigmoid() if gate_sigmoid else g) # bf16, as the unfused model path + ref = b.double() @ w.double().t() + decode = ops.k3_decode_gemv(b, w, True) + same_decode = torch.equal(_bits(y), _bits(decode)) + _report("k3_ctm_gemv_gated", f"s{split}_{'sig' if gate_sigmoid else 'given'}", m, y, ref, F.linear(b, w), + det=det, rows_as_m8=minv, eq_unfused_k3_decode_gemv=same_decode) # fmt: skip + assert _rel(y, ref) <= TOL and _rel(y, F.linear(b, w)) <= TOL + assert det and minv + if split == 1: + assert same_decode + + +def _silu_and_mul(gu: torch.Tensor) -> torch.Tensor: + """silu_and_mul's fp32 arithmetic and one bf16 rounding.""" + k = gu.shape[1] // 2 + return (F.silu(gu[:, :k].float()) * gu[:, k:].float()).bfloat16() + + +@pytest.mark.parametrize("m", M_ALL) +@pytest.mark.parametrize("name", list(SWIGLU)) +def test_k3_ctm_gemv_swiglu(name, m): + ops = _ops() + n, k, split, push = SWIGLU[name] + w = _weight(n, k, 7) + gu8 = _rows(2 * k, 8, 2.0) + assert _ctm().supports_swiglu(gu8, w, split) + (gu,), y, det, minv = _checks(lambda g_: ops.k3_ctm_gemv_swiglu(g_, w, True, split, push), [gu8], m) + act = _silu_and_mul(gu) + ref = act.double() @ w.double().t() + _report("k3_ctm_gemv_swiglu", name, m, y, ref, F.linear(act, w), det=det, rows_as_m8=minv) + assert _rel(y, ref) <= TOL and _rel(y, F.linear(act, w)) <= TOL + assert det and minv + + +@pytest.mark.parametrize("m", M_ALL) +@pytest.mark.parametrize("name", list(LONG)) +def test_k3_ctm_gemv_long(name, m): + ops = _ops() + n, k, sig, split, ring, push = LONG[name] + w = _weight(n, k, 9) + x8 = _rows(k, 10) + assert _ctm().supports_long(x8, w, split, ring) + + def call(x_, sig_col0): + return ops.k3_ctm_gemv_long(x_, w, sig_col0, split, ring, True, push) + + (x,), y, det, minv = _checks(lambda x_: call(x_, sig), [x8], m) + plain = call(x, -1) if sig >= 0 else y + ref = x.double() @ w.double().t() + cols = slice(0, sig if sig >= 0 else n) + # Columns >= sig_col0 hold bf16(sigmoid(bf16(x @ W^T))): torch.sigmoid of the op's own plain output. + sigmoid_ok = sig < 0 or ( + torch.equal(_bits(y[:, :sig]), _bits(plain[:, :sig])) + and torch.equal(_bits(y[:, sig:]), _bits(plain[:, sig:].sigmoid())) + ) + _report("k3_ctm_gemv_long", name, m, plain, ref, F.linear(x, w), det=det, rows_as_m8=minv, + sigmoid_cols=sigmoid_ok) # fmt: skip + assert _rel(plain, ref) <= TOL and _rel(plain, F.linear(x, w)) <= TOL + assert _rel(y[:, cols], ref[:, cols]) <= TOL + assert det and minv and sigmoid_ok + + +@pytest.mark.parametrize("m", M_ALL) +@pytest.mark.parametrize("rank", [0, 7, 15]) +def test_k3_ctm_gemv_tail(rank, m): + """The row-parallel MoE tail on the CTM kernel: the bits of k3_decode_gemv_tail.""" + ops = _ops() + w = _weight(HIDDEN, PAD + ACT, 11, 0.03).clone() + w[:, WIDTH:PAD] = 0 + lo = rank * WIDTH + + def call(lat_, act_): + return ops.k3_ctm_gemv_tail(lat_, act_, w, lo, WIDTH, EPS, True) + + (lat, act), y, det, minv = _checks(call, [_rows(LATENT, 12, 0.8), _rows(ACT, 13, 0.5)], m) + lat64 = lat.double() + normed = lat64 * torch.rsqrt(lat64.pow(2).mean(dim=1, keepdim=True) + EPS) + ref = torch.cat([normed[:, lo : lo + WIDTH], act.double()], dim=1) @ torch.cat( + [w[:, :WIDTH], w[:, PAD:]], dim=1 + ).double().t() + decode = ops.k3_decode_gemv_tail(lat, act, w, lo, WIDTH, EPS, True) + same_decode = torch.equal(_bits(y), _bits(decode)) + _report("k3_ctm_gemv_tail", f"rank{rank}", m, y, ref, det=det, rows_as_m8=minv, eq_k3_decode_gemv_tail=same_decode) + assert _rel(y, ref) <= TOL + assert det and minv and same_decode + + +def _situ_ref(gu, beta, linear_beta): + """SituAndMul's eager fp32 path (in fp64).""" + k = gu.shape[1] // 2 + g, u = gu[:, :k].double(), gu[:, k:].double() + a = beta * torch.tanh(g / beta) * torch.sigmoid(g) + if linear_beta is not None: + u = linear_beta * torch.tanh(u / linear_beta) + return a * u + + +@pytest.mark.parametrize("m", M_ALL) +@pytest.mark.parametrize("situ", list(SITU)) +def test_k3_situ_mul(situ, m): + """The dense MLP's SiTU-and-mul at TP16 (gate_up [M, 4224] -> [M, 2112]) against the Triton SituAndMul.""" + ops = _ops() + from tensorrt_llm._torch.modules.situ import SituAndMul + + beta, linear_beta = SITU[situ] + gu8 = _rows(2 * 2112, 14, 2.0) + (gu,), y, det, minv = _checks(lambda g_: ops.k3_situ_mul(g_, beta, linear_beta), [gu8], m) + stock = SituAndMul(beta=beta, linear_beta=linear_beta, use_fused_activation=True)(gu) + ref = _situ_ref(gu, beta, linear_beta) + identical = (_bits(y) == _bits(stock)).float().mean().item() + _report("k3_situ_mul", situ, m, y, ref, stock, det=det, rows_as_m8=minv, frac_eq_triton=f"{identical:.4f}") + assert y.shape == (m, 2112) + assert _rel(y, ref) <= TOL and _rel(y, stock) <= TOL + assert det and minv + + +@pytest.mark.parametrize("m", [0, 9, 16]) +def test_token_limit(m): + op = _ctm() + _ops() + x = torch.zeros(m, 7168, dtype=torch.bfloat16, device="cuda") + for n, k, _, split, ring, _ in LONG.values(): + assert not op.supports_long(torch.zeros(m, k, dtype=torch.bfloat16, device="cuda"), _weight(n, k, 9), split, ring) + with pytest.raises(ValueError): + torch.ops.trtllm.k3_ctm_gemv_long(x, _weight(2880, 7168, 9), 2112, 6, 6, True, True) + assert not op.supports(torch.zeros(m, 768, dtype=torch.bfloat16, device="cuda"), _weight(7168, 768, 1, 0.03), 1) + assert not op.supports_situ_mul(torch.zeros(m, 4224, dtype=torch.bfloat16, device="cuda")) + assert not op.supports_swiglu(torch.zeros(m, 1792, dtype=torch.bfloat16, device="cuda"), _weight(7168, 896, 7), 2) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctm_gemv_wide.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctm_gemv_wide.py new file mode 100644 index 000000000000..0158e2be99d3 --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctm_gemv_wide.py @@ -0,0 +1,179 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""trtllm::k3_ctm_gemv_wide (up to 64 tokens in one MMA of 16 / 32 / 64 token columns) at the Kimi K3 TP16 per-rank +projection shapes, at every M in 1..64: error against an fp64 product and against cuBLAS (F.linear), run-to-run +identical bits, each M's rows bit-identical to the same rows of the 64-row call, the sigmoid columns torch.sigmoid of +the op's own plain output, and, where the split matches the call site's k3_ctm_gemv_long at most 8 tokens, the same +bits as that kernel. Both split-K transports give the same bits. 0 and 65 rows are refused.""" + +import functools + +import pytest +import torch +import torch.nn.functional as F + + +def _is_sm100() -> bool: + if not torch.cuda.is_available(): + return False + major, _ = torch.cuda.get_device_capability() + return major == 10 + + +pytestmark = pytest.mark.skipif(not _is_sm100(), reason="needs an SM100-family GPU") + +M_ALL = list(range(1, 65)) +TOL = 8e-3 # bf16 output: max |y - ref| / max |ref| +TOL_FP32 = 1e-4 # fp32 output + +# (N, K, sig_col0, out_fp32): the per-rank TP16 shapes of the projections of a wide decode step. +WIDE = { + "kda_qkvg": (3208, 7168, -1, False), # KDA q/k/v/g/f_a/b: 25 whole tiles and an 8-row last one + "mla_qkv_a_gate": (2880, 7168, 2112, False), # [W_a; W_g], the gate rows as bf16(sigmoid) + "o_proj": (7168, 768, -1, False), # KDA / MLA o_proj + "moe_head": (280, 7168, -1, True), # [latent down slice; router rows], fp32 (router logits) + "moe_tail": (7168, 640, -1, False), # [latent up slice | padding | shared down], 5 k-tiles + "shared_gate_up": (768, 7168, -1, False), + "dense_gate_up": (4224, 7168, -1, False), + "dense_down": (7168, 2112, -1, False), # K 2112 ends in a half k-tile + "drafter_qkv": (512, 7168, -1, False), + "drafter_gate_up": (1792, 7168, -1, False), + "drafter_o_proj": (7168, 384, -1, False), +} +# The k3_ctm_gemv_long call (split, ring, push) of the shapes that run it at most 8 tokens. +LONG_SITES = { + "mla_qkv_a_gate": (6, 6, True), + "dense_gate_up": (4, 5, False), + "dense_down": (2, 6, False), + "drafter_qkv": (8, 6, True), + "drafter_gate_up": (8, 6, True), +} + + +def _ops(): + import tensorrt_llm # noqa: F401 (registers the trtllm ops) + from tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv import op # noqa: F401 + + return torch.ops.trtllm + + +def _ctm(): + from tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv import op + + return op + + +def _bits(t: torch.Tensor) -> torch.Tensor: + t = t.contiguous() + return t.view(torch.int32) if t.dtype == torch.float32 else t.view(torch.int16) + + +def _rel(y: torch.Tensor, ref: torch.Tensor) -> float: + return (y.double() - ref.double()).abs().max().item() / ref.double().abs().max().item() + + +@functools.lru_cache(maxsize=None) +def _weight(n: int, k: int, seed: int) -> torch.Tensor: + gen = torch.Generator(device="cuda").manual_seed(seed) + return (torch.randn(n, k, generator=gen, device="cuda") * 0.02).bfloat16() + + +@functools.lru_cache(maxsize=None) +def _rows(k: int, seed: int) -> torch.Tensor: + gen = torch.Generator(device="cuda").manual_seed(seed) + return torch.randn(64, k, generator=gen, device="cuda").bfloat16() + + +@functools.lru_cache(maxsize=None) +def _full_call(name: str) -> torch.Tensor: + """The 64-row call of a case (its rows are what every M's call must reproduce).""" + n, k, sig, fp32 = WIDE[name] + return _ops().k3_ctm_gemv_wide(_rows(k, 2), _weight(n, k, 1), sig, fp32) + + +@pytest.mark.parametrize("m", M_ALL) +@pytest.mark.parametrize("name", list(WIDE)) +def test_k3_ctm_gemv_wide(name, m): + ops = _ops() + n, k, sig, fp32 = WIDE[name] + w = _weight(n, k, 1) + x = _rows(k, 2)[:m].contiguous() + assert _ctm().supports_wide(x, w, sig, fp32) + y = ops.k3_ctm_gemv_wide(x, w, sig, fp32) + det = torch.equal(_bits(y), _bits(ops.k3_ctm_gemv_wide(x, w, sig, fp32))) + rows_as_m64 = torch.equal(_bits(y), _bits(_full_call(name)[:m])) + plain = ops.k3_ctm_gemv_wide(x, w, -1, fp32) if sig >= 0 else y + ref = x.double() @ w.double().t() + cols = slice(0, sig if sig >= 0 else n) + # Columns >= sig_col0 hold bf16(sigmoid(bf16(x @ W^T))): torch.sigmoid of the op's own plain output. + sigmoid_ok = sig < 0 or ( + torch.equal(_bits(y[:, :sig]), _bits(plain[:, :sig])) + and torch.equal(_bits(y[:, sig:]), _bits(plain[:, sig:].sigmoid())) + ) + eq_long = None + if name in LONG_SITES and m <= 8: + split, ring, push = LONG_SITES[name] + from tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv import k3_ctm_gemv_kernel as kernel + + wide_split = _ctm().wide_config( + n, k, kernel.wide_tile(m), torch.cuda.get_device_properties(0).multi_processor_count + )[0] + long_y = ops.k3_ctm_gemv_long(x, w, sig, split, ring, True, push) + eq_long = torch.equal(_bits(y), _bits(long_y)) if wide_split == split else None + stock = F.linear(x.float(), w.float()) if fp32 else F.linear(x, w) + print( + f"OPCHECK op=k3_ctm_gemv_wide case={name} M={m} dtype={str(y.dtype)[6:]} " + f"abs={(plain.double() - ref).abs().max().item():.3e} rel={_rel(plain, ref):.3e} " + f"vs_stock={_rel(plain, stock):.2e} det={det} rows_as_m64={rows_as_m64} sigmoid_cols={sigmoid_ok} " + f"eq_long={eq_long}" + ) # fmt: skip + assert y.shape == (m, n) and y.dtype == (torch.float32 if fp32 else torch.bfloat16) + tol = TOL_FP32 if fp32 else TOL + assert _rel(plain, ref) <= tol and _rel(plain, stock) <= tol + assert _rel(y[:, cols], ref[:, cols]) <= tol + assert det and rows_as_m64 and sigmoid_ok + assert eq_long is not False + + +@pytest.mark.parametrize("m", list(range(16, 65, 8))) +@pytest.mark.parametrize("name", list(LONG_SITES)) +def test_k3_ctm_gemv_wide_rows_as_long(name, m): + """A wide step of R x 8 tokens: every token's row is the bits k3_ctm_gemv_long gives that token at <= 8 tokens (the + call sites' split), so wide steps and batch 1 agree on these projections.""" + ops = _ops() + n, k, sig, _ = WIDE[name] + split, ring, push = LONG_SITES[name] + w = _weight(n, k, 1) + x = _rows(k, 2)[:m].contiguous() + y = ops.k3_ctm_gemv_wide(x, w, sig) + + def long_rows(r): + return ops.k3_ctm_gemv_long(x[r : r + 8].contiguous(), w, sig, split, ring, True, push) + + same = [torch.equal(_bits(y[r : r + 8]), _bits(long_rows(r))) for r in range(0, m, 8)] + print( + f"OPCHECK op=k3_ctm_gemv_wide case={name} M={m} rows_as_long_per_8={sum(same)}/{len(same)}" + ) + assert all(same) + + +@pytest.mark.parametrize("m", [0, 65]) +def test_token_limit(m): + _ops() + x = torch.zeros(m, 7168, dtype=torch.bfloat16, device="cuda") + w = _weight(3208, 7168, 1) + assert not _ctm().supports_wide(x, w) + with pytest.raises(ValueError): + torch.ops.trtllm.k3_ctm_gemv_wide(x, w) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_decode_gemv.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_decode_gemv.py new file mode 100644 index 000000000000..b75aa0ae4396 --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_decode_gemv.py @@ -0,0 +1,167 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""trtllm::k3_decode_gemv / k3_decode_gemv_tail (CuTe DSL decode GEMV, M <= 8) at the Kimi K3 TP16 per-rank shapes, +at every M in 1..8: error against an fp64 product and against the stock path (cuBLAS F.linear; RMSNorm -> slice -> +concat -> linear for the tail), for the tail also against torch with the kernel's arithmetic (fp32 accumulators, the +RMS on the latent one, one bf16 rounding), run-to-run identical bits and each M's rows bit-identical to the same rows +of the 8-row call. 0 and 9 rows are refused.""" + +import functools + +import pytest +import torch +import torch.nn.functional as F + + +def _is_sm100() -> bool: + if not torch.cuda.is_available(): + return False + major, _ = torch.cuda.get_device_capability() + return major == 10 + + +pytestmark = pytest.mark.skipif(not _is_sm100(), reason="needs an SM100-family GPU") + +M_ALL = list(range(1, 9)) +TOL = 8e-3 # max |y - ref| / max |ref| +# (N, K) per rank at TP16: the KDA o_proj (one CTA per weight tile) and the KDA input projection (split-K). +SHAPES = {"kda_o_proj": (7168, 768), "kda_qkvg": (3208, 7168)} +HIDDEN, LATENT, WIDTH, PAD, ACT = 7168, 3584, 224, 256, 384 +EPS = 1e-6 + + +def _ops(): + import tensorrt_llm # noqa: F401 (registers the trtllm ops) + from tensorrt_llm._torch.cute_dsl_kernels.k3_decode_gemv import op # noqa: F401 + + return torch.ops.trtllm + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def _rel(y: torch.Tensor, ref: torch.Tensor) -> float: + return (y.double() - ref.double()).abs().max().item() / ref.double().abs().max().item() + + +def _report(op, case, m, y, ref, stock=None, **flags): + extra = "" if stock is None else f" vs_stock={_rel(y, stock):.2e}" + marks = " ".join(f"{k}={v}" for k, v in flags.items()) + abs_err = (y.double() - ref.double()).abs().max().item() + print(f"OPCHECK op={op} case={case} M={m} abs={abs_err:.3e} rel={_rel(y, ref):.3e}{extra} {marks}") + + +@functools.lru_cache(maxsize=None) +def _weight(n: int, k: int, seed: int) -> torch.Tensor: + gen = torch.Generator(device="cuda").manual_seed(seed) + return (torch.randn(n, k, generator=gen, device="cuda") * 0.03).bfloat16() + + +@functools.lru_cache(maxsize=None) +def _rows(k: int, seed: int, scale: float = 1.0) -> torch.Tensor: + gen = torch.Generator(device="cuda").manual_seed(seed) + return (torch.randn(8, k, generator=gen, device="cuda") * scale).bfloat16() + + +@pytest.mark.parametrize("m", M_ALL) +@pytest.mark.parametrize("name", list(SHAPES)) +def test_k3_decode_gemv(name, m): + ops = _ops() + n, k = SHAPES[name] + w = _weight(n, k, 1) + x8 = _rows(k, 2) + x = x8[:m].contiguous() + y = ops.k3_decode_gemv(x, w, True) + again = ops.k3_decode_gemv(x, w, True) + y8 = ops.k3_decode_gemv(x8, w, True) + ref = x.double() @ w.double().t() + deterministic = torch.equal(_bits(y), _bits(again)) + m_invariant = torch.equal(_bits(y), _bits(y8[:m])) + _report("k3_decode_gemv", name, m, y, ref, F.linear(x, w), det=deterministic, rows_as_m8=m_invariant) + assert y.shape == (m, n) + assert _rel(y, ref) <= TOL and _rel(y, F.linear(x, w)) <= TOL + assert deterministic and m_invariant + + +def _tail_weight() -> torch.Tensor: + w = _weight(HIDDEN, PAD + ACT, 3).clone() + w[:, WIDTH:PAD] = 0 + return w + + +def _tail_ref(latent, act, w, lo): + lat = latent.double() + normed = lat * torch.rsqrt(lat.pow(2).mean(dim=1, keepdim=True) + EPS) + x = torch.cat([normed[:, lo : lo + WIDTH], act.double()], dim=1) + return x @ torch.cat([w[:, :WIDTH], w[:, PAD:]], dim=1).double().t() + + +def _tail_stock(latent, act, w, lo): + from tensorrt_llm._torch.modules.rms_norm import RMSNorm + + norm = RMSNorm(hidden_size=LATENT, eps=EPS, dtype=torch.bfloat16).cuda() + norm.weight.data.fill_(1.0) + x = torch.cat([norm(latent)[:, lo : lo + WIDTH], act], dim=1) + return F.linear(x, torch.cat([w[:, :WIDTH], w[:, PAD:]], dim=1)) + + +def _tail_fp32(latent, act, w, lo): + """The tail with the kernel's arithmetic: fp32 accumulators of the latent slice and of the activation, the latent + one scaled by the RMS of the whole latent row, one bf16 rounding.""" + lat = latent.float() + scale = torch.rsqrt(lat.pow(2).mean(dim=1, keepdim=True) + EPS) + acc_lat = lat[:, lo : lo + WIDTH] @ w[:, :WIDTH].float().t() + return (acc_lat * scale + act.float() @ w[:, PAD:].float().t()).bfloat16() + + +@pytest.mark.parametrize("m", M_ALL) +@pytest.mark.parametrize("rank", [0, 7, 15]) +def test_k3_decode_gemv_tail(rank, m): + ops = _ops() + w = _tail_weight() + lat8, act8 = _rows(LATENT, 4, 0.8), _rows(ACT, 5, 0.5) + lat, act, lo = lat8[:m].contiguous(), act8[:m].contiguous(), rank * WIDTH + + def call(lat_, act_): + return ops.k3_decode_gemv_tail(lat_, act_, w, lo, WIDTH, EPS, True) + + y = call(lat, act) + again = call(lat, act) + y8 = call(lat8, act8) + fp32 = _tail_fp32(lat, act, w, lo) + ref = _tail_ref(lat, act, w, lo) + stock = _tail_stock(lat, act, w, lo) + deterministic = torch.equal(_bits(y), _bits(again)) + m_invariant = torch.equal(_bits(y), _bits(y8[:m])) + _report("k3_decode_gemv_tail", f"rank{rank}", m, y, ref, stock, det=deterministic, rows_as_m8=m_invariant, + vs_fp32=f"{_rel(y, fp32):.2e}") # fmt: skip + assert _rel(y, ref) <= TOL and _rel(y, stock) <= TOL and _rel(y, fp32) <= TOL + assert deterministic and m_invariant + + +@pytest.mark.parametrize("m", [0, 9, 16]) +def test_token_limit(m): + ops = _ops() + from tensorrt_llm._torch.cute_dsl_kernels.k3_decode_gemv import op + + w = _weight(*SHAPES["kda_o_proj"], 1) + x = torch.zeros(m, w.shape[1], dtype=torch.bfloat16, device="cuda") + assert not op.supports(x, w) + with pytest.raises(ValueError): + ops.k3_decode_gemv(x, w, True) + lat = torch.zeros(m, LATENT, dtype=torch.bfloat16, device="cuda") + act = torch.zeros(m, ACT, dtype=torch.bfloat16, device="cuda") + assert not op.supports_tail(lat, act, _tail_weight(), WIDTH) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_embed.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_embed.py new file mode 100644 index 000000000000..408bb5b1ef35 --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_embed.py @@ -0,0 +1,574 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``trtllm::k3_embed`` (a decode step's embedding rows, one 16-byte load per thread) and ``trtllm::k3_embed_norm`` +(the rows into a snapshot-bank slot plus layer 0's input RMSNorm, one launch) at the in-model shape: Kimi K3's +replicated bf16 embedding table [163840, 7168], M = 1 .. 64 tokens (every M; each compiles its own kernel). + +The reference is what the model runs without them (KimiLinearModel._embed and layer 0's input_layernorm): +nn.Embedding's ``F.embedding`` (an index_select) for the rows, and the RMSNorm module, whose bf16 path is +``trtllm::flashinfer_rmsnorm`` (``flashinfer.norm.rmsnorm``), for the norm. Both ops claim bit-identity (bf16 bits). +k3_embed_norm reproduces flashinfer's CuTe DSL RMSNormKernel (one 128-thread CTA per row for 6144 < H <= 16384, see +k3_embed_kernel), so its checks need flashinfer (the model's gate for the op) on that kernel. The table is replicated +(no vocabulary shard), so the model never passes an id outside [0, V); the ops define one as a zero row, checked +against the masked reference. + +Checks (in the first three, every call runs twice and the reruns must be bit-identical): + +* k3_embed: int32 ids (the engine's) at every M, int64 ids at a few; ids with 0, V - 1, repeats, one id M times; +* k3_embed_norm at every M, called as the model calls it (the module's weight and eps, a bank slot as ``raw``): the + normed rows against the module on the F.embedding rows, the raw rows, the bank's other slots untouched; eps 1e-6 + and 1e-5, bank slots 0 (the model's) and 1; +* ids outside [0, V) (negative, V, the int32 extremes, int64 ids past 2^32): zero rows, normed as the module norms + them; +* every table row through both ops (64 consecutive ids per call); +* CUDA graphs of both ops at M = 1 .. 16, 17 .. 32, 33 .. 48 and 49 .. 64, captured once and replayed with the ids + rewritten in place, bit-identical to the eager ops and the reference on every replay; the same ids replayed again + give the same bits. + +Table: ``python3 test_k3_embed.py report``. Timing: ``python3 test_k3_embed.py time`` (CUDA graphs of back-to-back +calls, each on its own random ids, median us per call over 15 replays at every M: F.embedding vs k3_embed, and the +rows plus the snapshot copy plus the RMSNorm module vs k3_embed_norm). +""" + +import math +import os +import statistics +import sys + +import pytest +import torch +import torch.nn.functional as F + +VOCAB = 163840 # Kimi K3's vocabulary (a replicated nn.Embedding) +HIDDEN = 7168 +EPS = 1e-6 +EPS_ALT = 1e-5 +MAX_TOKENS = 64 # k3_embed's MAX_TOKENS +TOKENS = list(range(1, MAX_TOKENS + 1)) +INT64_TOKENS = [1, 7, 8, 16, 33, 64] +OUTSIDE_TOKENS = [1, 8, 64] +GRAPH_FAMILIES = [list(range(lo, lo + 16)) for lo in range(1, MAX_TOKENS + 1, 16)] +BANK_SLOTS = 4 # attention-residual snapshots; the model passes slot 0 as raw +SENTINEL = 7.0 +ID_KINDS = ("0, V - 1, repeats", "V - 1, 0", "random", "one id x M") +GRAPH_KINDS = ("random", "0, V - 1, repeats", "outside [0, V)", "one id x M", "V - 1, 0") +OUTSIDE_IDS = (-1, VOCAB, VOCAB + 7, -VOCAB, 2**31 - 1, -(2**31)) +OUTSIDE_IDS_64 = OUTSIDE_IDS + (2**32, 2**32 + 5, 3 - 2**32, 2**40) + + +def _sm90() -> bool: + return torch.cuda.is_available() and torch.cuda.get_device_capability() >= (9, 0) + + +pytestmark = pytest.mark.skipif( + not _sm90(), reason="needs SM90 or newer (griddepcontrol, programmatic dependent launch)" +) + + +def _op(): + import tensorrt_llm # noqa: F401 (registers the trtllm ops) + from tensorrt_llm._torch.cute_dsl_kernels.k3_embed import op + + return op + + +def _flashinfer_norm() -> bool: + """Whether layer 0's RMSNorm runs flashinfer's CuTe DSL RMSNormKernel: the model takes k3_embed_norm only with + flashinfer, and FLASHINFER_USE_CUDA_NORM=1 selects flashinfer's CUDA kernel, which is not the one the op + reproduces.""" + from tensorrt_llm._torch.flashinfer_utils import IS_FLASHINFER_AVAILABLE + + return IS_FLASHINFER_AVAILABLE and os.environ.get("FLASHINFER_USE_CUDA_NORM", "0") != "1" + + +def _need_norm(): + if not _flashinfer_norm(): + pytest.skip("needs flashinfer's CuTe DSL RMSNorm (the kernel k3_embed_norm reproduces)") + + +# ---------------------------------------------------------------------------------------------------------------- +# Inputs and the model's path +# ---------------------------------------------------------------------------------------------------------------- + +_state = {} + + +def table() -> torch.Tensor: + """The embedding table, built once: N(0, 1) rows at log-uniform scales 1e-4 .. 1e2, row 1 all zero.""" + if "table" not in _state: + with torch.inference_mode(False): + gen = torch.Generator(device="cuda").manual_seed(20260930) + t = torch.randn(VOCAB, HIDDEN, generator=gen, device="cuda", dtype=torch.bfloat16) + scale = torch.empty(VOCAB, 1, device="cuda") + scale.uniform_(math.log(1e-4), math.log(1e2), generator=gen) + t.mul_(scale.exp_().bfloat16()) + t[1].zero_() + _state["table"] = t + return _state["table"] + + +def layer0_norm(eps): + """Layer 0's input_layernorm as KimiLinearDecoderLayer builds it (the RMSNorm module, bf16): weights around 1, + every 101st negative.""" + key = ("norm", eps) + if key not in _state: + from tensorrt_llm._torch.modules.rms_norm import RMSNorm + + with torch.inference_mode(False): + gen = torch.Generator(device="cuda").manual_seed(7) + w = 1.0 + 0.2 * torch.randn(HIDDEN, generator=gen, device="cuda") + w[::101] *= -1.0 + norm = RMSNorm(hidden_size=HIDDEN, eps=eps, dtype=torch.bfloat16, device="cuda") + norm.requires_grad_(False) + norm.weight.copy_(w) + _state[key] = norm + return _state[key] + + +def make_ids(m, kind, gen, dtype=torch.int32): + """``m`` ids in [0, V) of one of ID_KINDS.""" + ids = torch.randint(0, VOCAB, (m,), generator=gen, device="cuda") + if kind == "0, V - 1, repeats": + ids[0] = 0 + if m >= 2: + ids[-1] = VOCAB - 1 + if m >= 4: + ids[1] = VOCAB - 1 + ids[2] = 0 + if m >= 6: + ids[4] = ids[3] + elif kind == "V - 1, 0": + ids[0] = VOCAB - 1 + if m >= 2: + ids[-1] = 0 + elif kind == "one id x M": + ids = ids[:1].repeat(m) + return ids.to(dtype) + + +def make_outside_ids(m, dtype, gen): + """``m`` ids: the even positions outside [0, V) (cycling through the dtype's out-of-range values), the others + random in [0, V).""" + bad = OUTSIDE_IDS_64 if dtype == torch.int64 else OUTSIDE_IDS + ids = torch.randint(0, VOCAB, (m,), generator=gen, device="cuda") + ids[0::2] = torch.tensor([bad[i % len(bad)] for i in range((m + 1) // 2)], device="cuda") + return ids.to(dtype) + + +def rows_ref(ids): + """The model's rows: F.embedding (nn.Embedding's index_select); an id outside [0, V) gives the ops' zero row.""" + valid = (ids >= 0) & (ids < VOCAB) + if bool(valid.all()): + return F.embedding(ids, table()) + rows = F.embedding(torch.where(valid, ids, torch.zeros_like(ids)), table()) + return rows.masked_fill_(~valid[:, None], 0.0) + + +def embed(ids): + """k3_embed as KimiLinearModel._embed calls it.""" + return torch.ops.trtllm.k3_embed(ids, table()) + + +def new_bank(m): + """A snapshot bank [BANK_SLOTS, m, H] filled with SENTINEL.""" + return torch.full((BANK_SLOTS, m, HIDDEN), SENTINEL, dtype=torch.bfloat16, device="cuda") + + +def embed_norm_into(ids, norm, raw): + """k3_embed_norm as KimiLinearModel._embed_norm calls it: the norm module's weight and eps, ``raw`` a bank slot.""" + return torch.ops.trtllm.k3_embed_norm(ids, table(), norm.weight, norm.variance_epsilon, raw) + + +def embed_norm(ids, eps=EPS, slot=0): + """k3_embed_norm into ``slot`` of a fresh bank: (normed rows, bank).""" + bank = new_bank(ids.numel()) + return embed_norm_into(ids, layer0_norm(eps), bank[slot]), bank + + +def same(a, b) -> bool: + """Bit-identical bf16 tensors.""" + return a.shape == b.shape and torch.equal(a.view(torch.int16), b.view(torch.int16)) + + +def differ(a, b): + """The number of rows of ``a`` that differ from ``b`` in any bit (on the device).""" + return (a.view(torch.int16) != b.view(torch.int16)).any(dim=1).sum() + + +def check_embed(ids, kind): + got = embed(ids) + return {f"{kind}: rows": same(got, rows_ref(ids)), f"{kind}: rerun": same(embed(ids), got)} + + +def check_norm(ids, kind, eps=EPS, slot=0): + normed, bank = embed_norm(ids, eps, slot) + normed_again, bank_again = embed_norm(ids, eps, slot) + rows = rows_ref(ids) + others = [s for s in range(BANK_SLOTS) if s != slot] + return { + f"{kind}: normed": same(normed, layer0_norm(eps)(rows)), + f"{kind}: raw": same(bank[slot], rows), + f"{kind}: other slots": bool((bank[others] == SENTINEL).all()), + f"{kind}: rerun": same(normed_again, normed) and same(bank_again, bank), + } + + +def result(op_name, m, case, checks): + """A table row: ``checks`` maps each comparison to whether it held.""" + failed = [k for k, ok in checks.items() if not ok] + identical = "yes" if not failed else "no: " + ", ".join(failed) + return dict(op=op_name, m=m, case=case, identical=identical, ok=not failed) + + +# ---------------------------------------------------------------------------------------------------------------- +# Measurements +# ---------------------------------------------------------------------------------------------------------------- + + +def measure_embed(m): + """k3_embed at M tokens: int32 ids of every kind, int64 ids at INT64_TOKENS.""" + gen = torch.Generator(device="cuda").manual_seed(1000 + m) + checks = {} + for kind in ID_KINDS: + checks.update(check_embed(make_ids(m, kind, gen), kind)) + out = [result("k3_embed", m, "int32 ids: " + " / ".join(ID_KINDS), checks)] + if m in INT64_TOKENS: + ids = make_ids(m, ID_KINDS[0], gen, torch.int64) + out.append(result("k3_embed", m, f"int64 ids: {ID_KINDS[0]}", check_embed(ids, "int64"))) + return out + + +def measure_norm(m): + """k3_embed_norm at M tokens: int32 ids of every kind into slot 0 at EPS, random ids into slot 1 at EPS_ALT, + int64 ids at INT64_TOKENS.""" + gen = torch.Generator(device="cuda").manual_seed(2000 + m) + checks = {} + for kind in ID_KINDS: + checks.update(check_norm(make_ids(m, kind, gen), kind)) + case = f"int32 ids: {' / '.join(ID_KINDS)}; eps {EPS:g}, slot 0" + out = [result("k3_embed_norm", m, case, checks)] + checks = check_norm(make_ids(m, "random", gen), "random", EPS_ALT, 1) + out.append(result("k3_embed_norm", m, f"int32 ids: random; eps {EPS_ALT:g}, slot 1", checks)) + if m in INT64_TOKENS: + checks = check_norm(make_ids(m, ID_KINDS[0], gen, torch.int64), "int64") + case = f"int64 ids: {ID_KINDS[0]}; eps {EPS:g}, slot 0" + out.append(result("k3_embed_norm", m, case, checks)) + return out + + +def measure_outside(m, dtype, with_norm): + """Both ops on ids outside [0, V) at the even positions.""" + gen = torch.Generator(device="cuda").manual_seed(3000 + m) + ids = make_outside_ids(m, dtype, gen) + values = ", ".join(map(str, OUTSIDE_IDS_64 if dtype == torch.int64 else OUTSIDE_IDS)) + case = f"{str(dtype).split('.')[-1]} ids, even positions outside [0, V): {values}" + out = [result("k3_embed", m, case, check_embed(ids, "outside"))] + if with_norm: + out.append(result("k3_embed_norm", m, case, check_norm(ids, "outside"))) + return out + + +def measure_every_row(with_norm): + """Every table row through k3_embed (and k3_embed_norm), MAX_TOKENS consecutive ids per call.""" + tab = table() + norm = layer0_norm(EPS) if with_norm else None + bank = new_bank(MAX_TOKENS) + every_id = torch.arange(VOCAB, dtype=torch.int32, device="cuda") + bad = torch.zeros(3, dtype=torch.int64, device="cuda") + for lo in range(0, VOCAB, MAX_TOKENS): + ids = every_id[lo : lo + MAX_TOKENS] + rows = F.embedding(ids, tab) + bad[0] += differ(embed(ids), rows) + if with_norm: + normed = embed_norm_into(ids, norm, bank[0]) + bad[1] += differ(bank[0], rows) + bad[2] += differ(normed, norm(rows)) + bad = bad.tolist() + case = f"all {VOCAB} rows, {MAX_TOKENS} consecutive int32 ids per call" + out = [result("k3_embed", "all", case, {f"rows ({bad[0]} differ)": bad[0] == 0})] + if with_norm: + checks = {f"raw ({bad[1]} rows differ)": bad[1] == 0} + checks[f"normed ({bad[2]} rows differ)"] = bad[2] == 0 + out.append(result("k3_embed_norm", "all", f"{case}, eps {EPS:g}", checks)) + return out + + +def measure_graph(ms, with_norm): + """One CUDA graph with both ops at every M of ``ms``, captured once and replayed with the ids rewritten in place + (one GRAPH_KINDS kind per replay), each replay against the eager ops and the reference; then the last ids + replayed again against the first replay of them.""" + tab = table() + norm = layer0_norm(EPS) if with_norm else None + gen = torch.Generator(device="cuda").manual_seed(4000 + ms[0]) + ids_buf = {m: torch.zeros(m, dtype=torch.int32, device="cuda") for m in ms} + banks = {m: new_bank(m) for m in ms} + outs = {} + + def body(): + for m in ms: + outs["k3_embed", m] = torch.ops.trtllm.k3_embed(ids_buf[m], tab) + if with_norm: + outs["k3_embed_norm", m] = embed_norm_into(ids_buf[m], norm, banks[m][0]) + + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(stream): + body() # every M compiles on its first call, outside capture + torch.cuda.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + body() + torch.cuda.synchronize() + family = f"{ms[0]}-{ms[-1]}" + out = [] + for rep, kind in enumerate(GRAPH_KINDS): + for m in ms: + if kind == "outside [0, V)": + ids_buf[m].copy_(make_outside_ids(m, torch.int32, gen)) + else: + ids_buf[m].copy_(make_ids(m, kind, gen)) + banks[m].fill_(SENTINEL) + graph.replay() + torch.cuda.synchronize() + checks = {"k3_embed": {}, "k3_embed_norm": {}} + for m in ms: + ids = ids_buf[m] + rows = rows_ref(ids) + got = outs["k3_embed", m] + checks["k3_embed"][f"M {m} = reference"] = same(got, rows) + checks["k3_embed"][f"M {m} = eager"] = same(got, embed(ids)) + if with_norm: + got = outs["k3_embed_norm", m] + eager, eager_bank = embed_norm(ids) + checks["k3_embed_norm"].update({ + f"M {m} normed = reference": same(got, norm(rows)), + f"M {m} raw = reference": same(banks[m][0], rows), + f"M {m} = eager": same(got, eager) and same(banks[m], eager_bank), + }) # fmt: skip + for name, op_checks in checks.items(): + if op_checks: + out.append(result(name, family, f"replay {rep}: {kind} ids rewritten", op_checks)) + first = {k: v.clone() for k, v in outs.items()} + first_banks = {m: b.clone() for m, b in banks.items()} + for b in banks.values(): + b.fill_(SENTINEL) + graph.replay() + torch.cuda.synchronize() + for name in ("k3_embed", "k3_embed_norm") if with_norm else ("k3_embed",): + rerun = {f"M {m}": same(outs[name, m], first[name, m]) for m in ms} + if name == "k3_embed_norm": + rerun.update({f"M {m} bank": same(banks[m], first_banks[m]) for m in ms}) + out.append(result(name, family, "the last ids replayed again", rerun)) + return out + + +# ---------------------------------------------------------------------------------------------------------------- +# Tests +# ---------------------------------------------------------------------------------------------------------------- + + +@pytest.mark.parametrize("m", TOKENS) +def test_embed(m): + _op() + with torch.inference_mode(): + bad = [r for r in measure_embed(m) if not r["ok"]] + assert not bad, bad + + +@pytest.mark.parametrize("m", TOKENS) +def test_embed_norm(m): + _op() + _need_norm() + with torch.inference_mode(): + bad = [r for r in measure_norm(m) if not r["ok"]] + assert not bad, bad + + +@pytest.mark.parametrize("dtype", [torch.int32, torch.int64], ids=["int32", "int64"]) +@pytest.mark.parametrize("m", OUTSIDE_TOKENS) +def test_ids_outside_vocab(m, dtype): + _op() + with torch.inference_mode(): + bad = [r for r in measure_outside(m, dtype, _flashinfer_norm()) if not r["ok"]] + assert not bad, bad + + +def test_every_row(): + _op() + with torch.inference_mode(): + bad = [r for r in measure_every_row(_flashinfer_norm()) if not r["ok"]] + assert not bad, bad + + +@pytest.mark.parametrize("ms", GRAPH_FAMILIES, ids=[f"M{f[0]}-{f[-1]}" for f in GRAPH_FAMILIES]) +def test_graph_replay(ms): + _op() + with torch.inference_mode(): + bad = [r for r in measure_graph(ms, _flashinfer_norm()) if not r["ok"]] + assert not bad, bad + + +def test_supports(): + """The model's calls are supported at every M in 1 .. 64 (hidden 7168, int32 or int64 ids); M = 0 or 65, other id + dtypes and widths outside the norm's geometry are refused.""" + op = _op() + with torch.inference_mode(): + tab = table() + norm = layer0_norm(EPS) + bank = torch.empty(2, MAX_TOKENS, HIDDEN, dtype=torch.bfloat16, device="cuda") + + def ids(n, dtype=torch.int32): + return torch.zeros(n, dtype=dtype, device="cuda") + + for m in TOKENS: + raw = bank[0, :m] + assert op.supports(ids(m), tab) and op.supports(ids(m, torch.int64), tab), m + assert op.supports_norm(ids(m), tab, norm.weight, raw), m + assert op.norm_supports_hidden(HIDDEN) + assert not op.norm_supports_hidden(6144) and not op.norm_supports_hidden(HIDDEN + 512) + assert not op.supports(ids(0), tab) and not op.supports(ids(MAX_TOKENS + 1), tab) + assert not op.supports(ids(8, torch.int16), tab) + assert not op.supports_norm(ids(8), tab, norm.weight, bank[0]) # raw [64, H] for 8 ids + with pytest.raises((ValueError, RuntimeError)): + torch.ops.trtllm.k3_embed(ids(MAX_TOKENS + 1), tab) + + +# ---------------------------------------------------------------------------------------------------------------- +# Error table (python3 test_k3_embed.py report) and timing (python3 test_k3_embed.py time) +# ---------------------------------------------------------------------------------------------------------------- + + +def report() -> int: + _op() + with_norm = _flashinfer_norm() + rows = [] + with torch.inference_mode(): + print(f"{torch.cuda.get_device_name()}; table [{VOCAB}, {HIDDEN}] bf16") + for m in TOKENS: + rows += measure_embed(m) + if with_norm: + rows += measure_norm(m) + for m in OUTSIDE_TOKENS: + for dtype in (torch.int32, torch.int64): + rows += measure_outside(m, dtype, with_norm) + rows += measure_every_row(with_norm) + for ms in GRAPH_FAMILIES: + rows += measure_graph(ms, with_norm) + what = { + "k3_embed": "identical: rows = F.embedding (zero rows for ids outside [0, V)); rerun = first run; in a " + "graph replay also = the eager op", + "k3_embed_norm": "identical: normed = the RMSNorm module on the reference rows; raw (the bank slot) = the " + "reference rows; the bank's other slots untouched; rerun = first run; in a graph replay also = the eager op", + } + for name in ("k3_embed", "k3_embed_norm"): + print(f"\n## {name}\n\n{what[name]}\n") + if name == "k3_embed_norm" and not with_norm: + print( + "skipped: needs flashinfer's CuTe DSL RMSNorm (the kernel k3_embed_norm reproduces)" + ) + continue + print("| M | case | identical | result |") + print("| :-- | :-- | :-- | :-- |") + for r in (r for r in rows if r["op"] == name): + verdict = "PASS" if r["ok"] else "FAIL" + print(f"| {r['m']} | {r['case']} | {r['identical']} | {verdict} |") + ok_all = all(r["ok"] for r in rows) + print("\nALL PASS" if ok_all else "\nFAIL") + return 0 if ok_all else 1 + + +def time_graph(body, calls, replays=15): + """Per-call us of a CUDA graph of ``calls`` back-to-back ``body(i)``: median (min, max) over ``replays`` replays. + ``body(-1)`` runs first, outside capture.""" + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(stream): + body(-1) + torch.cuda.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + for i in range(calls): + body(i) + torch.cuda.synchronize() + for _ in range(3): + graph.replay() + torch.cuda.synchronize() + per_call = [] + for _ in range(replays): + start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + start.record() + graph.replay() + end.record() + torch.cuda.synchronize() + per_call.append(start.elapsed_time(end) * 1e3 / calls) + return statistics.median(per_call), min(per_call), max(per_call) + + +def timing() -> None: + _op() + with_norm = _flashinfer_norm() + calls = 32 + with torch.inference_mode(): + tab = table() + norm = layer0_norm(EPS) if with_norm else None + gen = torch.Generator(device="cuda").manual_seed(11) + print(f"{torch.cuda.get_device_name()}; table [{VOCAB}, {HIDDEN}] bf16; graphs of {calls} back-to-back calls, " + "each on its own random int32 ids, 15 replays: median (min-max) us per call") # fmt: skip + print("| M | F.embedding | k3_embed | F.embedding + snapshot + RMSNorm | k3_embed + snapshot + RMSNorm " + "| k3_embed_norm |") # fmt: skip + print("| --: | --: | --: | --: | --: | --: |") + for m in TOKENS: + idss = [ + torch.randint(0, VOCAB, (m,), generator=gen, device="cuda", dtype=torch.int32) + for _ in range(calls) + ] + bank = new_bank(m) + + def unfused(rows, bank=bank): + """Layer 0 without k3_embed_norm: the bank's snapshot of the rows and the input RMSNorm.""" + bank[0].copy_(rows) + return norm(rows) + + arms = [ + lambda i: F.embedding(idss[i], tab), + lambda i: torch.ops.trtllm.k3_embed(idss[i], tab), + ] + if with_norm: + arms += [ + lambda i: unfused(F.embedding(idss[i], tab)), + lambda i: unfused(torch.ops.trtllm.k3_embed(idss[i], tab)), + lambda i: embed_norm_into(idss[i], norm, bank[0]), + ] + res = [[] for _ in arms] + for rep in range(3): # alternating order + order = range(len(arms)) if rep % 2 == 0 else reversed(range(len(arms))) + for a in order: + res[a].append(time_graph(arms[a], calls)) + cells = [] + for timings in res: + meds = sorted(x[0] for x in timings) + lo, hi = min(x[1] for x in timings), max(x[2] for x in timings) + cells.append(f"{meds[1]:.2f} ({lo:.2f}-{hi:.2f})") + cells += ["n/a"] * (5 - len(cells)) + print(f"| {m} | " + " | ".join(cells) + " |", flush=True) + + +if __name__ == "__main__": + if len(sys.argv) > 1 and sys.argv[1] == "time": + timing() + elif len(sys.argv) > 1 and sys.argv[1] == "report": + sys.exit(report()) + else: + sys.exit(pytest.main([__file__, "-q", "-p", "no:cacheprovider", *sys.argv[1:]])) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_head_gemv.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_head_gemv.py new file mode 100644 index 000000000000..f7a5afa2223d --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_head_gemv.py @@ -0,0 +1,176 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""trtllm::k3_head_gemv (the drafter's lm_head vocab shard, stream-K, M <= 8) at the Kimi K3 TP16 shard +[10240, 7168], at every M in 1..8 as the DSpark worker calls it (default schedule): error against an fp64 product and +cuBLAS F.linear, run-to-run identical bits, each M's rows bit-identical to the same rows of the 8-row call, and calls of +different M back to back leaving the workspace counters at zero (the next call's bits unchanged). Stream-K weights +with fewer 128 x 128 tiles than SMs, or exactly as many, complete and match an fp64 product at every M.""" + +import functools +import os +import subprocess +import sys + +import pytest +import torch +import torch.nn.functional as F + + +def _is_sm100() -> bool: + if not torch.cuda.is_available(): + return False + major, _ = torch.cuda.get_device_capability() + return major == 10 + + +pytestmark = pytest.mark.skipif(not _is_sm100(), reason="needs an SM100-family GPU") + +M_ALL = list(range(1, 9)) +TOL = 8e-3 +VOCAB_SHARD, HIDDEN = 10240, 7168 # 163840 / 16 rows per rank +# [N, K] weights with fewer 128 x 128 tiles than an SM100 GPU has SMs (56, 108, 120 and 144 tiles). +FEW_TILES = [(128, 7168), (1152, 1536), (1280, 1536), (2304, 1024)] +# For the few-tile child; import, compile and the calls take about a minute. +FEW_TILES_TIMEOUT_S = 300 + +# The few-tile shapes ("NxK" arguments) one after another at every M: prints each call's error against an fp64 product +# once the call has completed. +_FEW_TILES_CHILD = r""" +import sys +import torch +import tensorrt_llm # noqa: F401 +from tensorrt_llm._torch.cute_dsl_kernels.k3_head_gemv import op # noqa: F401 +gen = torch.Generator(device="cuda").manual_seed(20261002) +for shape in sys.argv[1:]: + n, k = map(int, shape.split("x")) + w = (torch.randn(n, k, generator=gen, device="cuda") * 0.02).bfloat16() + x8 = torch.randn(8, k, generator=gen, device="cuda").bfloat16() + for m in range(1, 9): + x = x8[:m].contiguous() + y = torch.ops.trtllm.k3_head_gemv(x, w) + torch.cuda.synchronize() + ref = x.double() @ w.double().t() + print("REL", n, k, m, ((y.double() - ref).abs().max() / ref.abs().max()).item(), flush=True) +""" + + +def _ops(): + import tensorrt_llm # noqa: F401 + from tensorrt_llm._torch.cute_dsl_kernels.k3_head_gemv import op # noqa: F401 + + return torch.ops.trtllm + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def _rel(y: torch.Tensor, ref: torch.Tensor) -> float: + return (y.double() - ref.double()).abs().max().item() / ref.double().abs().max().item() + + +@functools.lru_cache(maxsize=None) +def _inputs(): + gen = torch.Generator(device="cuda").manual_seed(20261001) + w = (torch.randn(VOCAB_SHARD, HIDDEN, generator=gen, device="cuda") * 0.02).bfloat16() + x8 = torch.randn(8, HIDDEN, generator=gen, device="cuda").bfloat16() + return w, x8 + + +@pytest.mark.parametrize("m", M_ALL) +def test_k3_head_gemv(m): + ops = _ops() + from tensorrt_llm._torch.cute_dsl_kernels.k3_head_gemv import op + + w, x8 = _inputs() + x = x8[:m].contiguous() + assert op.supports(x, w) + y = ops.k3_head_gemv(x, w) + y8 = ops.k3_head_gemv(x8, w) + again = ops.k3_head_gemv(x, w) + ref = x.double() @ w.double().t() + stock = F.linear(x, w) + det = torch.equal(_bits(y), _bits(again)) + minv = torch.equal(_bits(y), _bits(y8[:m])) + abs_err = (y.double() - ref).abs().max().item() + print(f"OPCHECK op=k3_head_gemv case=lm_head_shard M={m} abs={abs_err:.3e} rel={_rel(y, ref):.3e} " + f"vs_stock={_rel(y, stock):.2e} det={det} rows_as_m8={minv}") # fmt: skip + assert y.shape == (m, VOCAB_SHARD) + assert _rel(y, ref) <= TOL and _rel(y, stock) <= TOL + assert det and minv + + +def test_k3_head_gemv_mixed_m_sequence(): + """M 8, 1, 5, 8, 3 back to back on one stream (one workspace per shape): each the bits of its own call.""" + ops = _ops() + w, x8 = _inputs() + single = {m: ops.k3_head_gemv(x8[:m].contiguous(), w) for m in (1, 3, 5, 8)} + seq = [ops.k3_head_gemv(x8[:m].contiguous(), w) for m in (8, 1, 5, 8, 3)] + for m, y in zip((8, 1, 5, 8, 3), seq): + assert torch.equal(_bits(y), _bits(single[m])) + + +def _tiles_equal_to_sms(sms): + """An [N, K] weight with exactly ``sms`` 128 x 128 tiles, each 128-row tile split over at least 2 k-tiles.""" + k_tiles = max([d for d in range(2, 17) if sms % d == 0] or [1]) + return 128 * (sms // k_tiles), 128 * k_tiles + + +def test_k3_head_gemv_few_tiles(): + """Stream-K weights with fewer 128 x 128 tiles than SMs, and one with exactly as many: every call completes and + matches an fp64 product. A stuck call would block its process, so the shapes run in a child process with a + deadline.""" + sms = torch.cuda.get_device_properties(0).multi_processor_count + shapes = [(n, k) for n, k in FEW_TILES if (n // 128) * (k // 128) < sms] + assert shapes + shapes.append(_tiles_equal_to_sms(sms)) + # Without a compute-sanitizer target's injection variables: they would load the tool into the child too, and + # there CUDA initialization can stall. + env = { + k: v + for k, v in os.environ.items() + if not k.startswith(("NV_SANITIZER_", "NVIDIA_PROCESS_INJECTION_")) + and k != "NVTX_INJECTION64_PATH" + } + child = subprocess.Popen( + [sys.executable, "-c", _FEW_TILES_CHILD, *(f"{n}x{k}" for n, k in shapes)], + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + text=True, + env=env, + ) + try: + out = child.communicate(timeout=FEW_TILES_TIMEOUT_S)[0] + except subprocess.TimeoutExpired: + child.kill() + out = child.communicate()[0] + f"\nno result after {FEW_TILES_TIMEOUT_S} s" + rel = {shape: {} for shape in shapes} + for _, n, k, m, r in (ln.split() for ln in out.splitlines() if ln.startswith("REL ")): + rel[(int(n), int(k))][int(m)] = float(r) + for (n, k), by_m in rel.items(): + print(f"OPCHECK op=k3_head_gemv case=few_tiles N={n} K={k} tiles={(n // 128) * (k // 128)} sms={sms} " + f"rel={by_m}") # fmt: skip + done = all(set(by_m) == set(M_ALL) for by_m in rel.values()) + assert child.returncode == 0 and done, out[-3000:] + assert max(max(by_m.values()) for by_m in rel.values()) <= TOL, rel + + +@pytest.mark.parametrize("m", [0, 9, 16]) +def test_token_limit(m): + _ops() + from tensorrt_llm._torch.cute_dsl_kernels.k3_head_gemv import op + + w, _ = _inputs() + assert not op.supports(torch.zeros(m, HIDDEN, dtype=torch.bfloat16, device="cuda"), w) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py new file mode 100644 index 000000000000..05584abac209 --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py @@ -0,0 +1,110 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Source check of the tcgen05 thread-sync fences in Kimi K3 kernels (PTX ISA, tcgen05 memory consistency, canonical +sync patterns): a tcgen05 operation issued after an mbarrier wait follows ``tcgen05.fence::after_thread_sync``, and a +thread whose TMEM loads another thread's tcgen05 work must not overtake (an mbarrier arrive or a CTA barrier after +``tcgen05.wait::ld``) issues ``tcgen05.fence::before_thread_sync`` first. Without them ptxas may move the TMEM access +across the synchronization; no test of values can see it, so the kernels' sources are read. + + pytest test_k3_tcgen05_fences.py +""" + +import importlib.util +import re + +import pytest + +KERNELS = [ + "tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv.k3_ctm_gemv_kernel", + "tensorrt_llm._torch.cute_dsl_kernels.k3_decode_gemv.k3_decode_gemv_kernel", +] + +# Thread syncs that order other threads' work before a tcgen05 operation of this thread. Waits on a barrier only TMA +# completes (the weight / activation rings) are not: TMA writes and tcgen05 reads are both async proxy, ordered by the +# barrier's complete_tx, the CUTLASS mainloop pattern. +WAIT = re.compile( + r"mbarrier_try_wait\(|mbarrier_test_wait\(|_(try|test)_wait_cluster\(|(? Date: Fri, 2 Oct 2026 18:46:50 -0700 Subject: [PATCH 027/161] [None][perf] Kimi K3 attn_res decode RMSNorm: release the next kernel after the grid wait The decode kernels behind attn_res_rmsnorm_fwd and attn_res_add_rmsnorm_fwd (the single-CTA and split-K s1 kernels that launchAttnResDecodeRmsNorm starts) now trigger programmatic dependent launch right after their own grid-dependency wait instead of at their end. Under PDL the next kernel then launches and streams its weights while they run. It still waits for this grid to complete before it reads the output, so results do not change. The other launches of the s1 kernels keep the trigger at the end. Signed-off-by: Vasanth Sabavat --- .../kernels/kimiK3AttnRes/attnResFwd.cu | 61 ++++++++++++++----- 1 file changed, 47 insertions(+), 14 deletions(-) diff --git a/cpp/tensorrt_llm/kernels/kimiK3AttnRes/attnResFwd.cu b/cpp/tensorrt_llm/kernels/kimiK3AttnRes/attnResFwd.cu index 39a14df806b6..60baff165e2b 100644 --- a/cpp/tensorrt_llm/kernels/kimiK3AttnRes/attnResFwd.cu +++ b/cpp/tensorrt_llm/kernels/kimiK3AttnRes/attnResFwd.cu @@ -1047,12 +1047,18 @@ __global__ void __launch_bounds__(256, 1) attn_res_fwd_s1_single_cta_kernel(bf16 bf16_t const* __restrict__ layer_res, bf16_t const* __restrict__ layer_res_add, bf16_t const* __restrict__ res_w, bf16_t const* __restrict__ rms_w, bf16_t const* __restrict__ output_rms_w, bf16_t* __restrict__ updated_layer_res, bf16_t* __restrict__ output, float* __restrict__ rsigma_out, float* __restrict__ probs_out, - float* __restrict__ logits_out, float rms_eps, float output_rms_eps, int num_tokens) + float* __restrict__ logits_out, float rms_eps, float output_rms_eps, int num_tokens, bool early_trigger) { #if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 1000 if constexpr (ENABLE_PDL) { cudaGridDependencySynchronize(); + // The dependent grid still waits for this one to complete before reading its + // output; triggering here lets it launch and stream its weights while this kernel runs. + if (early_trigger) + { + cudaTriggerProgrammaticLaunchCompletion(); + } } constexpr int H = 7168; @@ -1322,7 +1328,10 @@ __global__ void __launch_bounds__(256, 1) attn_res_fwd_s1_single_cta_kernel(bf16 if constexpr (ENABLE_PDL) { - cudaTriggerProgrammaticLaunchCompletion(); + if (!early_trigger) + { + cudaTriggerProgrammaticLaunchCompletion(); + } } #else if (cute::thread0()) @@ -1332,11 +1341,17 @@ __global__ void __launch_bounds__(256, 1) attn_res_fwd_s1_single_cta_kernel(bf16 #endif } +// Decode handoff options for the s1 kernels (see AttnResFwdParams). +struct S1Handoff +{ + bool early_trigger = false; +}; + template static void launch_s1_single_cta(bf16_t const* block_residual, bf16_t const* layer_residual, bf16_t const* layer_residual_add, bf16_t const* res_weight, bf16_t const* rms_weight, bf16_t const* output_rms_weight, bf16_t* updated_layer_residual, bf16_t* output, float* rsigma, float* probs, - float* logits, float rms_eps, float output_rms_eps, int num_tokens, cudaStream_t stream) + float* logits, float rms_eps, float output_rms_eps, int num_tokens, cudaStream_t stream, S1Handoff handoff = {}) { if (tensorrt_llm::common::getEnvEnablePDL()) { @@ -1352,14 +1367,14 @@ static void launch_s1_single_cta(bf16_t const* block_residual, bf16_t const* lay config.numAttrs = 1; cudaLaunchKernelEx(&config, kernel, block_residual, layer_residual, layer_residual_add, res_weight, rms_weight, output_rms_weight, updated_layer_residual, output, rsigma, probs, logits, rms_eps, output_rms_eps, - num_tokens); + num_tokens, handoff.early_trigger); } else { attn_res_fwd_s1_single_cta_kernel <<>>(block_residual, layer_residual, layer_residual_add, res_weight, rms_weight, output_rms_weight, updated_layer_residual, output, rsigma, probs, logits, rms_eps, output_rms_eps, - num_tokens); + num_tokens, false); } } @@ -1381,13 +1396,22 @@ __global__ void __launch_bounds__(256, 1) attn_res_fwd_s1_splitk_kernel(bf16_t c bf16_t const* __restrict__ layer_res, bf16_t const* __restrict__ layer_res_add, bf16_t const* __restrict__ res_w, bf16_t const* __restrict__ rms_w, bf16_t const* __restrict__ output_rms_w, bf16_t* __restrict__ updated_layer_res, bf16_t* __restrict__ output, float* __restrict__ rsigma_out, float* __restrict__ probs_out, - float* __restrict__ logits_out, float rms_eps, float output_rms_eps, int num_tokens) + float* __restrict__ logits_out, float rms_eps, float output_rms_eps, int num_tokens, bool early_trigger) { #if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 1000 - if constexpr (ENABLE_PDL) + auto const wait_for_previous_grid = [&] { - cudaGridDependencySynchronize(); - } + if constexpr (ENABLE_PDL) + { + cudaGridDependencySynchronize(); + // See attn_res_fwd_s1_single_cta_kernel. + if (early_trigger) + { + cudaTriggerProgrammaticLaunchCompletion(); + } + } + }; + wait_for_previous_grid(); namespace cg = cooperative_groups; constexpr int H = 7168; @@ -1440,6 +1464,7 @@ __global__ void __launch_bounds__(256, 1) attn_res_fwd_s1_splitk_kernel(bf16_t c float sq[N] = {}; float dot[N] = {}; + // Every candidate in one pass, so all loads are in flight together. #pragma unroll for (int ki = tid; ki < K_PER_CTA; ki += THREADS) { @@ -1639,7 +1664,10 @@ __global__ void __launch_bounds__(256, 1) attn_res_fwd_s1_splitk_kernel(bf16_t c if constexpr (ENABLE_PDL) { - cudaTriggerProgrammaticLaunchCompletion(); + if (!early_trigger) + { + cudaTriggerProgrammaticLaunchCompletion(); + } } #else if (cute::thread0()) @@ -1653,13 +1681,14 @@ template : &attn_res_fwd_s1_splitk_kernel; static std::once_flag attrs_set[2][64]; @@ -1679,7 +1708,7 @@ static void launch_s1_splitk(bf16_t const* block_residual, bf16_t const* layer_r void* args[] = {const_cast(&block_residual), const_cast(&layer_residual), const_cast(&layer_residual_add), const_cast(&res_weight), const_cast(&rms_weight), const_cast(&output_rms_weight), &updated_layer_residual, &output, &rsigma, &probs, &logits, &rms_eps, - &output_rms_eps, &num_tokens}; + &output_rms_eps, &num_tokens, &early_trigger}; cudaLaunchConfig_t config{}; // One cluster per token. clusterDim stays GROUPS, so clusters are contiguous // spans of the grid and each token's DSM exchange stays within its own. @@ -1888,18 +1917,22 @@ static void launchAttnResDecodeRmsNorm(AttnResFwdParams const& params, cudaStrea auto const* layer_residual_add = FUSE_LAYER_ADD ? params.layerResidualAdd : nullptr; auto* updated_layer_residual = FUSE_LAYER_ADD ? params.updatedLayerResidual : nullptr; + // The dependent kernel launches (and streams its weights) while this one runs: it waits for + // this grid before reading the output. + constexpr bool early_trigger = true; + S1Handoff const handoff{early_trigger}; if constexpr (N <= 4) { launch_s1_single_cta(params.blockResidual, params.layerResidual, layer_residual_add, params.resWeight, params.rmsWeight, params.outputRmsWeight, updated_layer_residual, params.output, nullptr, - nullptr, nullptr, params.rmsEps, params.outputRmsEps, params.seqLen, stream); + nullptr, nullptr, params.rmsEps, params.outputRmsEps, params.seqLen, stream, handoff); } else { launch_s1_splitk(params.blockResidual, params.layerResidual, layer_residual_add, params.resWeight, params.rmsWeight, params.outputRmsWeight, updated_layer_residual, params.output, nullptr, - nullptr, nullptr, params.rmsEps, params.outputRmsEps, params.seqLen, stream); + nullptr, nullptr, params.rmsEps, params.outputRmsEps, params.seqLen, stream, handoff); } } From b6a8fb58cf711dd24a470b24315faeb46b458f6a Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 18:47:07 -0700 Subject: [PATCH 028/161] [None][feat] Kimi K3 head GEMV: a caller-owned workspace trtllm::k3_head_gemv kept its cross-launch state in a module-level dict keyed by device and size, allocated by the first call: the stream-K partials and flag words, or the dynamic schedule's per-tile counts and unit counter. That state is now a K3HeadGemvWorkspace that the caller creates once per weight shape and schedule, eagerly (create() refuses CUDA-graph capture), and passes to every call. The op takes the workspace's three buffers as arguments, names them in mutates_args, and refuses a workspace whose buffers do not fit the call. The op tests add two workspaces used alternately, a captured graph replayed between eager calls on one workspace, a workspace of another shape (refused) and create() under capture (refused), and check that every call leaves the words at zero. Signed-off-by: Vasanth Sabavat --- .../cute_dsl_kernels/k3_head_gemv/op.py | 141 +++++++++++++----- .../kimi_k3/test_k3_head_gemv.py | 100 +++++++++++-- 2 files changed, 187 insertions(+), 54 deletions(-) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_head_gemv/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_head_gemv/op.py index 0b2b5ada30b7..3434256477d8 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_head_gemv/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_head_gemv/op.py @@ -16,19 +16,21 @@ One CTA per SM, and under stream-K at most one per (tile, k-tile), streams the weight, EVICT_FIRST except for the first ``keep_tiles`` 128-row tiles, and adds each tile's k-split partials in a fixed order before one bf16 rounding: -the result is bit-identical from run to run (not bit-identical to cuBLAS or pdl_gemv, whose accumulation orders -differ). Two schedules: ``streamk`` (every CTA an equal share of the flat (tile, k-tile) space; the default) and +the result is bit-identical from run to run (not bit-identical to cuBLAS, whose accumulation order differs). Two +schedules: ``streamk`` (every CTA an equal share of the flat (tile, k-tile) space; the default) and ``dynamic`` ((tile, k-chunk) units claimed from a counter, ``chunk_tiles`` k-tiles each). -The op keeps one workspace per device and shape (the partials, a per-tile count and the unit counter, the counters -zero between launches), so calls of one shape must be ordered on one stream. Each kernel compiles on the first -call for its shape, which must happen outside CUDA-graph capture. +A call's cross-launch state is a caller-owned ``K3HeadGemvWorkspace`` (the partials, the per-tile flag or count words +and the unit counter, the counters zero between launches), created for one weight shape and schedule by its eager +``create()`` before any CUDA-graph capture. The calls that share one workspace must be ordered on one stream. Each +kernel compiles on the first call for its shape, which must also happen outside CUDA-graph capture. """ from __future__ import annotations import os import threading +from dataclasses import dataclass from typing import Dict import torch @@ -38,7 +40,6 @@ _lock = threading.Lock() _compiled: Dict[tuple, object] = {} -_workspaces: Dict[tuple, tuple] = {} def _arg(t: torch.Tensor): @@ -109,22 +110,68 @@ def supports( return kernel.supports(n_out, k_in, chunk, ring) -def _workspace(device, units: int, tiles: int): - """(partials fp32 [units * 128 * 8], ``tiles`` count / flag words, claim counter), the counters zero.""" - key = (device.index, units, tiles) - ws = _workspaces.get(key) - if ws is None: +def _launch_config(n_out: int, k_in: int, device, chunk_tiles: int, schedule: str): + """(grid, units, flag words, pieces or k-tiles per unit) of a launch: the workspace sizes follow from them.""" + from . import k3_head_gemv_kernel as kernel + + tiles = n_out // kernel.CTA_M + grid = _grid(device) + if schedule == "streamk": + # At most one CTA per (tile, k-tile): every CTA's share is then non-empty, so each CTA between a split tile's + # first and last piece holds part of that tile and raises the flag its finalizer waits for. + grid = min(grid, tiles * kernel.num_k_tiles(k_in)) + pieces = kernel.streamk_max_pieces(n_out, k_in, grid) + return grid, tiles * pieces, tiles * pieces, pieces + chunk = chunk_tiles or pick_chunk(n_out, k_in, grid) + return grid, kernel.num_units(n_out, k_in, chunk), tiles, chunk + + +@dataclass(frozen=True) +class K3HeadGemvWorkspace: + """The cross-launch state of ``trtllm::k3_head_gemv`` for one weight shape and schedule. + + ``partials`` fp32 [units * 128 * 8] holds the fp32 partials of a split tile's pieces after the first (stream-K) + or of each unit (dynamic). ``flags`` int32 holds one word per stream-K piece, raised by the piece and lowered by + its tile's finalizer, or one unit count per tile (dynamic). ``claim`` int32 [1] is the dynamic schedule's unit + counter (unused by stream-K). Every launch leaves the words at zero for the next one. ``create()`` allocates them + zeroed, outside CUDA-graph capture. Calls of the same shape and schedule may share one workspace when they are + ordered on one stream. + """ + + n_out: int + k_in: int + schedule: str + chunk_tiles: int + partials: torch.Tensor + flags: torch.Tensor + claim: torch.Tensor + + @classmethod + def create( + cls, + n_out: int, + k_in: int, + device: torch.device, + schedule: str = "streamk", + chunk_tiles: int = 0, + ) -> "K3HeadGemvWorkspace": + """The workspace of [n_out, k_in] weights on ``device`` (``chunk_tiles``: the dynamic schedule's k-tiles per + unit, 0 for the op's choice).""" + if schedule not in ("streamk", "dynamic"): + raise ValueError(f"k3_head_gemv: unknown schedule {schedule!r}") if torch.cuda.is_current_stream_capturing(): - raise RuntimeError( - "k3_head_gemv: the workspace must be allocated outside CUDA-graph capture" - ) - ws = ( - torch.empty(max(units, 1) * 128 * 8, dtype=torch.float32, device=device), - torch.zeros(tiles, dtype=torch.int32, device=device), - torch.zeros(1, dtype=torch.int32, device=device), + raise RuntimeError("k3_head_gemv: create the workspace outside CUDA-graph capture") + device = torch.device(device) + _, units, words, per = _launch_config(n_out, k_in, device, chunk_tiles, schedule) + return cls( + n_out=n_out, + k_in=k_in, + schedule=schedule, + chunk_tiles=per if schedule == "dynamic" else 0, + partials=torch.empty(max(units, 1) * 128 * 8, dtype=torch.float32, device=device), + flags=torch.zeros(words, dtype=torch.int32, device=device), + claim=torch.zeros(1, dtype=torch.int32, device=device), ) - _workspaces[key] = ws - return ws def _compiled_fn(key, entry, *args): @@ -143,10 +190,13 @@ def _compiled_fn(key, entry, *args): return fn -@torch.library.custom_op("trtllm::k3_head_gemv", mutates_args=()) +@torch.library.custom_op("trtllm::k3_head_gemv", mutates_args=("partials", "flags", "claim")) def k3_head_gemv( x: torch.Tensor, weight: torch.Tensor, + partials: torch.Tensor, + flags: torch.Tensor, + claim: torch.Tensor, keep_tiles: int = 0, chunk_tiles: int = 0, ring: int = 6, @@ -154,10 +204,11 @@ def k3_head_gemv( prefetch: int = 16, ) -> torch.Tensor: """``x @ weight.T`` for bf16 ``x`` [M <= 8, K] and ``weight`` [N, K] (N and K multiples of 128); returns bf16 - [M, N]. Weight tiles (128 rows) below ``keep_tiles`` are read at normal L2 priority, the rest EVICT_FIRST. - ``schedule``: ``streamk`` or ``dynamic``; ``chunk_tiles`` (dynamic; 0: the op's choice) sets the k-tiles per - work unit; ``prefetch`` (stream-K) the k-tiles after the ring that each CTA prefetches into L2 before the grid - dependency.""" + [M, N]. ``partials`` / ``flags`` / ``claim`` are a ``K3HeadGemvWorkspace``'s, created for this shape and + schedule (and ``chunk_tiles``). Weight tiles (128 rows) below ``keep_tiles`` are read at normal L2 priority, the + rest EVICT_FIRST. ``schedule``: ``streamk`` or ``dynamic``; ``chunk_tiles`` (dynamic; 0: the op's choice) sets the + k-tiles per work unit; ``prefetch`` (stream-K) the k-tiles after the ring that each CTA prefetches into L2 before + the grid dependency.""" if not supports(x, weight, chunk_tiles, ring, schedule): raise ValueError( f"k3_head_gemv: unsupported call x {tuple(x.shape)} {x.dtype}, weight {tuple(weight.shape)} " @@ -167,32 +218,40 @@ def k3_head_gemv( num_tokens, k_in = x.shape n_out = weight.shape[0] - tiles = n_out // kernel.CTA_M - grid = _grid(x.device) + grid, units, words, per = _launch_config(n_out, k_in, x.device, chunk_tiles, schedule) + # A workspace of another shape or schedule is too small, or keeps its counters elsewhere: refused, not raced. + if not ( + partials.dtype == torch.float32 + and flags.dtype == torch.int32 + and claim.dtype == torch.int32 + and partials.device == x.device + and flags.device == x.device + and claim.device == x.device + and partials.numel() >= units * 128 * 8 + and flags.numel() == words + and claim.numel() == 1 + ): + raise ValueError( + f"k3_head_gemv: the workspace (partials {tuple(partials.shape)}, flags {tuple(flags.shape)}) is not " + f"one for weight {tuple(weight.shape)}, schedule {schedule}, chunk_tiles {chunk_tiles}" + ) # All 8 token rows are written (the rows past M from TMA's zero-filled x rows); the result is the first M. y = torch.empty(kernel.MMA_N, n_out, dtype=torch.bfloat16, device=x.device) stream = _stream(x) use_pdl = _use_pdl() if schedule == "streamk": - # At most one CTA per (tile, k-tile): every CTA's share is then non-empty, so each CTA between a split tile's - # first and last piece holds part of that tile and raises the flag its finalizer waits for. - grid = min(grid, tiles * kernel.num_k_tiles(k_in)) - pieces = kernel.streamk_max_pieces(n_out, k_in, grid) - ws, flags, _ = _workspace(x.device, tiles * pieces, tiles * pieces) - args = (_arg(weight), _arg(x), _arg(y.view(-1)), _arg(ws), _arg(flags)) - fn = _compiled_fn(("k3_head_gemv_sk", n_out, k_in, ring, grid, pieces, prefetch, use_pdl), - kernel.k3_head_gemv_sk, *args, num_tokens, keep_tiles, n_out, k_in, ring, grid, pieces, + args = (_arg(weight), _arg(x), _arg(y.view(-1)), _arg(partials), _arg(flags)) + fn = _compiled_fn(("k3_head_gemv_sk", n_out, k_in, ring, grid, per, prefetch, use_pdl), + kernel.k3_head_gemv_sk, *args, num_tokens, keep_tiles, n_out, k_in, ring, grid, per, prefetch, use_pdl, stream) # fmt: skip else: - chunk = chunk_tiles or pick_chunk(n_out, k_in, grid) - ws, cnt, claim = _workspace(x.device, kernel.num_units(n_out, k_in, chunk), tiles) - args = (_arg(weight), _arg(x), _arg(y.view(-1)), _arg(ws), _arg(cnt), _arg(claim)) - fn = _compiled_fn(("k3_head_gemv", n_out, k_in, chunk, ring, grid, use_pdl), kernel.k3_head_gemv, *args, - num_tokens, keep_tiles, n_out, k_in, chunk, ring, grid, use_pdl, stream) # fmt: skip + args = (_arg(weight), _arg(x), _arg(y.view(-1)), _arg(partials), _arg(flags), _arg(claim)) + fn = _compiled_fn(("k3_head_gemv", n_out, k_in, per, ring, grid, use_pdl), kernel.k3_head_gemv, *args, + num_tokens, keep_tiles, n_out, k_in, per, ring, grid, use_pdl, stream) # fmt: skip fn(*args, num_tokens, keep_tiles, stream) return y[:num_tokens] @k3_head_gemv.register_fake -def _(x, weight, keep_tiles=0, chunk_tiles=0, ring=6, schedule="streamk", prefetch=16): +def _(x, weight, partials, flags, claim, keep_tiles=0, chunk_tiles=0, ring=6, schedule="streamk", prefetch=16): return x.new_empty((x.shape[0], weight.shape[0]), dtype=torch.bfloat16) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_head_gemv.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_head_gemv.py index f7a5afa2223d..b97d9a9f5e91 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_head_gemv.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_head_gemv.py @@ -12,11 +12,13 @@ # 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. -"""trtllm::k3_head_gemv (the drafter's lm_head vocab shard, stream-K, M <= 8) at the Kimi K3 TP16 shard -[10240, 7168], at every M in 1..8 as the DSpark worker calls it (default schedule): error against an fp64 product and -cuBLAS F.linear, run-to-run identical bits, each M's rows bit-identical to the same rows of the 8-row call, and calls of -different M back to back leaving the workspace counters at zero (the next call's bits unchanged). Stream-K weights -with fewer 128 x 128 tiles than SMs, or exactly as many, complete and match an fp64 product at every M.""" +"""trtllm::k3_head_gemv (an lm_head vocab shard, stream-K, M <= 8) at the Kimi K3 TP16 shard [10240, 7168], at every +M in 1..8 (default schedule): error against an fp64 product and cuBLAS F.linear, run-to-run identical bits, each M's +rows bit-identical to the same rows of the 8-row call, and calls of different M back to back leaving the workspace +counters at zero (the next call's bits unchanged). The caller-owned workspace: two of one shape interleaved, a CUDA +graph replayed between eager calls on the same workspace, a workspace of another shape refused, and creation refused +under capture. Stream-K weights with fewer 128 x 128 tiles than SMs, or exactly as many, complete and match an fp64 +product at every M.""" import functools import os @@ -57,9 +59,10 @@ def _is_sm100() -> bool: n, k = map(int, shape.split("x")) w = (torch.randn(n, k, generator=gen, device="cuda") * 0.02).bfloat16() x8 = torch.randn(8, k, generator=gen, device="cuda").bfloat16() + ws = op.K3HeadGemvWorkspace.create(n, k, w.device) for m in range(1, 9): x = x8[:m].contiguous() - y = torch.ops.trtllm.k3_head_gemv(x, w) + y = torch.ops.trtllm.k3_head_gemv(x, w, ws.partials, ws.flags, ws.claim) torch.cuda.synchronize() ref = x.double() @ w.double().t() print("REL", n, k, m, ((y.double() - ref).abs().max() / ref.abs().max()).item(), flush=True) @@ -89,17 +92,30 @@ def _inputs(): return w, x8 +@functools.lru_cache(maxsize=None) +def _workspace(): + _ops() + from tensorrt_llm._torch.cute_dsl_kernels.k3_head_gemv import op + + return op.K3HeadGemvWorkspace.create(VOCAB_SHARD, HIDDEN, torch.device("cuda")) + + +def _call(x, w, ws): + return torch.ops.trtllm.k3_head_gemv(x, w, ws.partials, ws.flags, ws.claim) + + @pytest.mark.parametrize("m", M_ALL) def test_k3_head_gemv(m): ops = _ops() from tensorrt_llm._torch.cute_dsl_kernels.k3_head_gemv import op w, x8 = _inputs() + ws = _workspace() x = x8[:m].contiguous() assert op.supports(x, w) - y = ops.k3_head_gemv(x, w) - y8 = ops.k3_head_gemv(x8, w) - again = ops.k3_head_gemv(x, w) + y = _call(x, w, ws) + y8 = _call(x8, w, ws) + again = _call(x, w, ws) ref = x.double() @ w.double().t() stock = F.linear(x, w) det = torch.equal(_bits(y), _bits(again)) @@ -110,18 +126,76 @@ def test_k3_head_gemv(m): assert y.shape == (m, VOCAB_SHARD) assert _rel(y, ref) <= TOL and _rel(y, stock) <= TOL assert det and minv + assert not ws.flags.any() and not ws.claim.any() def test_k3_head_gemv_mixed_m_sequence(): - """M 8, 1, 5, 8, 3 back to back on one stream (one workspace per shape): each the bits of its own call.""" - ops = _ops() + """M 8, 1, 5, 8, 3 back to back on one stream and one workspace: each the bits of its own call.""" w, x8 = _inputs() - single = {m: ops.k3_head_gemv(x8[:m].contiguous(), w) for m in (1, 3, 5, 8)} - seq = [ops.k3_head_gemv(x8[:m].contiguous(), w) for m in (8, 1, 5, 8, 3)] + ws = _workspace() + single = {m: _call(x8[:m].contiguous(), w, ws) for m in (1, 3, 5, 8)} + seq = [_call(x8[:m].contiguous(), w, ws) for m in (8, 1, 5, 8, 3)] for m, y in zip((8, 1, 5, 8, 3), seq): assert torch.equal(_bits(y), _bits(single[m])) +def test_k3_head_gemv_two_workspaces_interleaved(): + """Two workspaces of one shape (e.g. two heads of the same shard size), their calls alternating on one stream: + every call the bits of its own single call, both workspaces' counters at zero after.""" + _ops() + from tensorrt_llm._torch.cute_dsl_kernels.k3_head_gemv import op + + w, x8 = _inputs() + w2 = (w.float() * -0.5).bfloat16() + wa, wb = _workspace(), op.K3HeadGemvWorkspace.create(VOCAB_SHARD, HIDDEN, w.device) + single = {(m, i): _call(x8[:m].contiguous(), (w, w2)[i], wa) for m in (1, 4, 8) for i in (0, 1)} + for m in (8, 1, 4, 1, 8): + for i, ws in ((0, wa), (1, wb)): + assert torch.equal(_bits(_call(x8[:m].contiguous(), (w, w2)[i], ws)), _bits(single[(m, i)])) + assert not wa.flags.any() and not wb.flags.any() + + +def test_k3_head_gemv_graph_replays_between_eager_calls(): + """A CUDA graph of M 1, 8 and 3 calls on one workspace, replayed three times with eager calls on the same + workspace between the replays: every result the bits of its eager call, the counters at zero after.""" + w, x8 = _inputs() + ws = _workspace() + xs = {m: x8[:m].contiguous() for m in (1, 3, 8)} + want = {m: _call(xs[m], w, ws) for m in xs} # also compiles before the capture + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + out = {m: _call(xs[m], w, ws) for m in (1, 8, 3)} + for _ in range(3): + graph.replay() + between = _call(xs[8], w, ws) + torch.cuda.synchronize() + assert all(torch.equal(_bits(out[m]), _bits(want[m])) for m in xs) + assert torch.equal(_bits(between), _bits(want[8])) + assert not ws.flags.any() + + +def test_k3_head_gemv_refuses_a_workspace_of_another_shape(): + """A workspace created for another weight shape is refused (too small, or its flag words laid out for another + tile count) instead of running the kernel over it.""" + _ops() + from tensorrt_llm._torch.cute_dsl_kernels.k3_head_gemv import op + + w, x8 = _inputs() + other = op.K3HeadGemvWorkspace.create(VOCAB_SHARD // 2, HIDDEN, w.device) + with pytest.raises(ValueError, match="workspace"): + _call(x8[:1].contiguous(), w, other) + + +def test_k3_head_gemv_workspace_not_created_under_capture(): + _ops() + from tensorrt_llm._torch.cute_dsl_kernels.k3_head_gemv import op + + graph = torch.cuda.CUDAGraph() + with pytest.raises(RuntimeError, match="capture"): + with torch.cuda.graph(graph): + op.K3HeadGemvWorkspace.create(VOCAB_SHARD, HIDDEN, torch.device("cuda")) + + def _tiles_equal_to_sms(sms): """An [N, K] weight with exactly ``sms`` 128 x 128 tiles, each 128-row tile split over at least 2 k-tiles.""" k_tiles = max([d for d in range(2, 17) if sms % d == 0] or [1]) From 11cefbdf07ddfce1604c98b62913fd1c0683dfc0 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 18:47:49 -0700 Subject: [PATCH 029/161] [None][feat] modeling_v2 catalog: the Kimi K3 single-GPU decode entries Ten entries for the single-GPU Kimi K3 decode ops, each a contract, a one-call wrapper and a GPU test: - gemm: k3_ctm_gemv, k3_ctm_gemv_swiglu, k3_ctm_gemv_long, k3_ctm_gemv_wide, k3_decode_gemv, and k3_head_gemv over a caller-owned K3HeadGemvWorkspace (its contract has a State section); - activation: k3_situ_mul; - norm: k3_embed_norm, attn_res_fwd, attn_res_rmsnorm_fwd. They are certified on sm_100 (B200 / GB200), where their first caller runs; their tests skip on other architectures. l0_b200 lists them, the kernels' op tests and lints, and the cublas_mm and flashinfer_rmsnorm catalog tests, whose entries gain sm_100 receipts. Signed-off-by: Vasanth Sabavat --- .../_experimental/modeling_v2/README.md | 6 +- .../catalog/activation/k3_situ_mul.md | 104 +++++++++ .../catalog/activation/k3_situ_mul.py | 19 ++ .../modeling_v2/catalog/gemm/k3_ctm_gemv.md | 138 ++++++++++++ .../modeling_v2/catalog/gemm/k3_ctm_gemv.py | 20 ++ .../catalog/gemm/k3_ctm_gemv_long.md | 155 +++++++++++++ .../catalog/gemm/k3_ctm_gemv_long.py | 28 +++ .../catalog/gemm/k3_ctm_gemv_swiglu.md | 137 ++++++++++++ .../catalog/gemm/k3_ctm_gemv_swiglu.py | 20 ++ .../catalog/gemm/k3_ctm_gemv_wide.md | 155 +++++++++++++ .../catalog/gemm/k3_ctm_gemv_wide.py | 17 ++ .../catalog/gemm/k3_decode_gemv.md | 112 ++++++++++ .../catalog/gemm/k3_decode_gemv.py | 16 ++ .../modeling_v2/catalog/gemm/k3_head_gemv.md | 199 +++++++++++++++++ .../modeling_v2/catalog/gemm/k3_head_gemv.py | 40 ++++ .../modeling_v2/catalog/index.yaml | 47 +++- .../modeling_v2/catalog/norm/attn_res_fwd.md | 152 +++++++++++++ .../modeling_v2/catalog/norm/attn_res_fwd.py | 28 +++ .../catalog/norm/attn_res_rmsnorm_fwd.md | 161 +++++++++++++ .../catalog/norm/attn_res_rmsnorm_fwd.py | 34 +++ .../modeling_v2/catalog/norm/k3_embed_norm.md | 119 ++++++++++ .../modeling_v2/catalog/norm/k3_embed_norm.py | 20 ++ .../test_lists/test-db/l0_b200.yml | 23 ++ .../test_modeling_v2_k3_situ_mul.py | 121 ++++++++++ .../gemm/test_modeling_v2_k3_ctm_gemv.py | 129 +++++++++++ .../gemm/test_modeling_v2_k3_ctm_gemv_long.py | 154 +++++++++++++ .../test_modeling_v2_k3_ctm_gemv_swiglu.py | 135 +++++++++++ .../gemm/test_modeling_v2_k3_ctm_gemv_wide.py | 184 +++++++++++++++ .../gemm/test_modeling_v2_k3_decode_gemv.py | 123 ++++++++++ .../gemm/test_modeling_v2_k3_head_gemv.py | 183 +++++++++++++++ .../norm/test_modeling_v2_attn_res_fwd.py | 168 ++++++++++++++ .../test_modeling_v2_attn_res_rmsnorm_fwd.py | 181 +++++++++++++++ .../norm/test_modeling_v2_k3_embed_norm.py | 211 ++++++++++++++++++ 33 files changed, 3336 insertions(+), 3 deletions(-) create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/activation/k3_situ_mul.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/activation/k3_situ_mul.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_long.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_long.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_swiglu.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_swiglu.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_wide.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_wide.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_decode_gemv.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_decode_gemv.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_head_gemv.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_head_gemv.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_fwd.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_fwd.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_rmsnorm_fwd.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_rmsnorm_fwd.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/k3_embed_norm.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/k3_embed_norm.py create mode 100644 tests/unittest/_torch/modeling_v2/activation/test_modeling_v2_k3_situ_mul.py create mode 100644 tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_ctm_gemv.py create mode 100644 tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_ctm_gemv_long.py create mode 100644 tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_ctm_gemv_swiglu.py create mode 100644 tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_ctm_gemv_wide.py create mode 100644 tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_decode_gemv.py create mode 100644 tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_head_gemv.py create mode 100644 tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_attn_res_fwd.py create mode 100644 tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_attn_res_rmsnorm_fwd.py create mode 100644 tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_k3_embed_norm.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/README.md b/tensorrt_llm/_torch/_experimental/modeling_v2/README.md index 61be49f5be0d..5c7448160818 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/README.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/README.md @@ -181,8 +181,10 @@ Perf is measured, never gated. ## Status of every record in this tree -**The catalog is fully certified on sm_103. The targets construct but have -never executed.** +**The catalog is certified: 19 entries on sm_103, and 12 on sm_100 (B200 / +GB200): the 10 Kimi K3 entries, where their first caller runs, and +`cublas_mm` and `flashinfer_rmsnorm`. The targets construct but have never +executed.** Two things voided every receipt in the move: each catalog test file was rewritten, and the targets moved from sm_100 (B200) to sm_103 (GB300), where diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/activation/k3_situ_mul.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/activation/k3_situ_mul.md new file mode 100644 index 000000000000..febbd25c5bb5 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/activation/k3_situ_mul.md @@ -0,0 +1,104 @@ +--- +receipts: {} +--- + +# k3_situ_mul + +**Wraps** `torch.ops.trtllm.k3_situ_mul` (one call). + +## Semantics + +The SiTU-gated multiply of the Kimi K3 dense MLP (`SituAndMul` in `tensorrt_llm/_torch/modules/situ.py`) for at +most 8 tokens, as one CuTe DSL kernel that lets the next kernel launch at once under programmatic dependent launch. +The last dim of the input holds the gate half followed by the up half: + +``` +K = gu.shape[1] // 2; g = gu[m, k], u = gu[m, K + k] 0 <= m < M <= 8, 0 <= k < K, in fp32 +a = (beta * tanh(g / beta)) * sigmoid(g) sigmoid(g) = 1 / (1 + exp(-g)) +v = linear_beta * tanh(u / linear_beta) if linear_beta is not None, else v = u +out[m, k] = bf16(a * v) +``` + +The arithmetic is fp32 in this order, with IEEE division and the full-precision `exp` and `tanh`, and the result +is rounded to bf16 once, round-to-nearest-even. That is the order of `SituAndMul`'s eager path. The `exp` and +`tanh` implementations are not torch's or Triton's, so the result is not bit-identical to `SituAndMul` (eager or +fused); they agree within the tolerance below. Rows are independent: a token's output row has the same bits +whatever `M` is. + +Fusion boundary: the single call computes the activation only. The caller owns the gate_up projection that +produces `gu` and the down projection that consumes the output. There is no quantization of the output. + +## Signature + +```python +def k3_situ_mul(gu: torch.Tensor, beta: float = 1.0, linear_beta: float | None = None) -> torch.Tensor +``` + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `gu` | `[M, 2 * K]`, gate half first, `1 <= M <= 8`, `K % 8 == 0` | bf16 | contiguous, 16-byte-aligned | CUDA | +| `beta` | scalar | Python float | — | — | +| `linear_beta` | scalar or `None` | Python float, not 0 (the wrapper asserts) | — | — | +| returns | `[M, K]` | bf16 | newly allocated, contiguous | `gu`'s device | + +Every element of the returned tensor is written; `gu` is not mutated. + +### Certified arguments + +The dense MLP of Kimi K3 at TP16 (gate_up output `[M, 4224]`, activation `[M, 2112]`) at every `M` in 1..8, with +`(beta, linear_beta)` = `(4.0, 25.0)` (the Kimi K3 checkpoint's) and `(1.0, None)` (the defaults). Inputs are +`gu ~ 2 * N(0, 1)`, rounded to bf16. The gate per cell: `max|y - ref| / max|ref| <= 8e-3` against the formula +above evaluated in fp64 on the same bf16 input, identical bits on a repeated call, and each `M`'s rows +bit-identical to the same rows of the 8-row call. Also certified, with both settings at `M` 1 and 8: calls +captured in a CUDA graph after an eager call per key, replayed with `gu` rewritten in place, return the bits of +eager calls on the new `gu`. Other widths and `beta` / `linear_beta` values are accepted but not certified. + +## Metadata consumed + +None. The op reads no attention metadata, KV cache or module state; every input is an argument, and the kernel +runs on the current CUDA stream of `gu`'s device. + +One process-global cache sits behind it, and it is result-neutral: the compiled kernel per key +`(K, linear_beta is not None, PDL on/off)`. The first call with a new key compiles the kernel with the CuTe DSL, +which costs host time; `M`, `beta` and the value of `linear_beta` are runtime arguments and never recompile. That +first call refuses to run under CUDA-graph capture: it raises `RuntimeError` ("run once per shape outside +CUDA-graph capture first") before launching anything. Call every key once eagerly, then capture; captured and +later calls reuse the compiled kernel. The PDL setting is read from `TRTLLM_ENABLE_PDL` on every call, so +changing it within a process adds a key. + +## Preconditions + +The op checks these and raises `ValueError` ("k3_situ_mul: unsupported call ...") when one fails: + +- `gu` is a 2-D contiguous CUDA bf16 tensor with `1 <= M <= 8` rows. +- `gu.shape[1] % 16 == 0`, so each half is a whole number of 16-byte vectors. +- `gu` starts at a 16-byte-aligned address. + +Certified refusals: 0 and 9 rows, a width of 4216, a `gu` starting 2 bytes past a 16-byte boundary, fp16, a +row-strided `gu` and a 1-D `gu`. + +The wrapper adds one check, because the op does not fail on it: `linear_beta` must not be `0.0`. The op hands the +kernel `linear_beta or 1.0`, so `linear_beta=0.0` would run as `1.0` (`v = tanh(u)`) with no error, where the +formula gives `0 * tanh(u / 0)`. The wrapper raises `AssertionError` before calling the op (certified). + +Not checked by the op: + +- `K >= 8`; an empty `gu` (width 0) passes the op's check. +- The architecture. The kernel needs no SM 10.x feature (no tcgen05, no clusters; `griddepcontrol` needs SM 9.0 + or newer), but only sm_100 has been measured. +- The CuTe DSL (`cutlass`) and `cuda-python` (`cuda.bindings`) are importable. The op imports the kernel module + after its check, so without them a call raises `ImportError`. + +## Notes + +- Programmatic dependent launch (PDL). With `TRTLLM_ENABLE_PDL` unset or `1` the kernel is launched with PDL; any + other value launches it without. There is no `trigger_early` argument: every CTA executes + `griddepcontrol.launch_dependents` first, then `griddepcontrol.wait`, and only then reads `gu`. So the next + kernel on the stream, if launched with PDL, launches at once and can stream its own weights while this kernel + and its predecessor run (the CTM GEMVs do); it must execute `griddepcontrol.wait` + (`cudaGridDependencySynchronize`) before it reads the output, as those GEMVs do before they read their + activation. Kernels launched without PDL, torch's included, start after this one completes as usual. None of + this changes results. +- Two identical calls return identical bits. +- Grid: `ceil(K / 1024)` x 8 CTAs of 128 threads, 8 columns per thread; rows past `M` exit at once. +- Only sm_100 has been measured. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/activation/k3_situ_mul.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/activation/k3_situ_mul.py new file mode 100644 index 000000000000..6118c1fbc2ef --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/activation/k3_situ_mul.py @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""SiTU-gated multiply (SituAndMul) of a gate_up output for M <= 8 tokens via the CTM (CuTe DSL) kernel.""" + +import torch + +import tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv.op # noqa: F401 (registers torch.ops.trtllm.k3_situ_mul) + + +def k3_situ_mul( + gu: torch.Tensor, beta: float = 1.0, linear_beta: float | None = None +) -> torch.Tensor: + """Return `SituAndMul(beta, linear_beta)(gu)`, gate half first, in one k3_situ_mul call.""" + # The op hands the kernel `linear_beta or 1.0`, so 0.0 would run as linear_beta=1.0 (up half + # tanh(u)) with no error, where the formula gives 0 * tanh(u / 0). + assert linear_beta is None or linear_beta != 0.0, ( + "linear_beta=0.0 would silently run as 1.0; pass None to leave the up half unscaled" + ) + return torch.ops.trtllm.k3_situ_mul(gu, beta=beta, linear_beta=linear_beta) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv.md new file mode 100644 index 000000000000..3d706c5d732d --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv.md @@ -0,0 +1,138 @@ +--- +receipts: {} +--- + +# k3_ctm_gemv + +**Wraps** `torch.ops.trtllm.k3_ctm_gemv` (one call). + +## Semantics + +The decode GEMV of a Kimi K3 projection whose reduction dimension fits in shared memory: at most 8 bf16 tokens +times a bf16 weight in `nn.Linear` layout (the bias-free `F.linear(x, weight)`), as one CuTe DSL kernel on the +tensor cores. + +``` +y[m, n] = bf16( sum_k x[m, k] * weight[n, k] ) 0 <= m < M <= 8, 0 <= n < N, 0 <= k < K +``` + +The bf16 products are accumulated in fp32 by the tensor cores (tcgen05 MMAs of 128 weight rows x 8 token columns +x 16 k, into tensor memory) and the sum is rounded to bf16 once, round-to-nearest-even. The summation order is +fixed by `K` and `split`: + +- `K` is cut into 128-column k-tiles. Each block of 128 output rows is one CTA (`split=1`) or one cluster of + `split` CTAs. +- With `split=1` the CTA accumulates all k-tiles in ascending order. +- With `split=2` or `4`, rank `r` of the cluster accumulates k-tiles `r, r + split, r + 2 * split, ...` in + ascending order into its own fp32 partial (when `split` does not divide the k-tiles, the first + `k_tiles % split` ranks take one more). The rank that owns an output row adds the `split` partials in rank order + `0, 1, ...` in fp32 and rounds that sum once. + +`push` only chooses how partials travel to the owning rank: DSMEM stores plus a cluster-scope release arrive +(`False`), or 16-byte `st.async` stores that complete the owner's barrier by bytes (`True`). The sums and their +order are the same either way, and `push` has no effect at `split=1`. Token rows are independent: rows past `M` +enter the MMA as zeros and are not stored, so a token's output row has the same bits whatever `M` is. + +Fusion boundary: the single call computes the GEMV only. There is no bias, activation, residual, all-reduce +(under tensor parallelism the per-rank partial is returned as is) or quantization; the caller owns those. + +## Signature + +```python +def k3_ctm_gemv( + x: torch.Tensor, + weight: torch.Tensor, + trigger_early: bool = True, + split: int = 1, + push: bool = False, +) -> torch.Tensor +``` + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `x` | `[M, K]`, `1 <= M <= 8` | bf16 | contiguous, 16-byte-aligned start | CUDA | +| `weight` | `[N, K]` (`nn.Linear` layout) | bf16 | contiguous, 16-byte-aligned start | CUDA, `x`'s device | +| `trigger_early` | scalar | bool: let PDL dependents launch early | — | — | +| `split` | scalar | int: 1, 2 or 4 CTAs per 128 output rows | — | — | +| `push` | scalar | bool: split-K transport | — | — | +| returns | `[M, N]` | bf16 | newly allocated, contiguous | `x`'s device | + +Every element of the returned tensor is written; the inputs are not mutated. + +### Certified arguments + +The Kimi K3 TP16 per-rank shapes with the `split` and `push` their call sites pass, `trigger_early=True`, at +every `M` in 1..8: + +| Cell | `N` | `K` | `split` | `push` | +|---|---|---|---|---| +| MLA o_proj | 7168 | 768 | 1 | False | +| MLA o_proj | 7168 | 768 | 2 | False | +| drafter o_proj | 7168 | 384 | 1 | True | +| drafter o_proj, synthetic drafter | 7168 | 256 | 1 | True | + +Inputs are `x ~ N(0, 1)` and `weight ~ 0.03 * N(0, 1)`, rounded to bf16. The gate per cell: +`max|y - ref| / max|ref| <= 8e-3` against the fp64 product of the same bf16 inputs, identical bits on a repeated +call, and each `M`'s rows bit-identical to the same rows of the 8-row call. Also certified, at the MLA o_proj: + +- with `split=2`, `push=True` and `trigger_early=False` return the bits of the `push=False`, + `trigger_early=True` call, at every `M`; +- with `split` 1 and 2, `M` 1 and 8: calls captured in a CUDA graph after an eager call per key, replayed with + `x` rewritten in place, return the bits of eager calls on the new `x`. + +`split=4` and other shapes the preconditions admit are accepted but not certified. + +## Metadata consumed + +None. The op reads no attention metadata, KV cache or module state; every input is an argument, and the kernel +runs on the current CUDA stream of `x`'s device. + +One process-global cache sits behind it, and it is result-neutral: the compiled kernel per key +`(N, K, split, trigger_early, push, PDL on/off)`. The first call with a new key compiles the kernel with the CuTe +DSL, which costs host time; `M` is a runtime argument and never recompiles. That first call refuses to run under +CUDA-graph capture: it raises `RuntimeError` ("run once per shape outside CUDA-graph capture first") before +launching anything. Call every key once eagerly, then capture; captured and later calls reuse the compiled +kernel. The PDL setting is read from `TRTLLM_ENABLE_PDL` on every call, so changing it within a process adds a +key. + +## Preconditions + +The op checks these and raises `ValueError` ("k3_ctm_gemv: unsupported call ...") when one fails: + +- `x` is a 2-D contiguous CUDA bf16 tensor with `1 <= M <= 8` rows. +- `weight` is a 2-D contiguous bf16 tensor with `weight.shape[1] == x.shape[1]`. +- `N % 128 == 0` and `K % 128 == 0`. +- `split` is 1, 2 or 4, `split <= K / 128`, and no rank holds more than 6 k-tiles: + `ceil(K / 128 / split) <= 6`, i.e. `K <= 768` at `split=1`, `K <= 1536` at 2, `K <= 3072` at 4. + +Certified refusals: 0 and 9 rows, fp16 `x`, a row-strided `x`, `N = 7104`, `K = 704`, `split=3`, and `K = 896` +(7 k-tiles) at `split=1`. + +Not checked by the op: + +- `x` and `weight` start at 16-byte-aligned addresses. The op passes both to the kernel with a declared 16-byte + alignment and loads them by TMA; a view whose start is not a multiple of 8 elements past an aligned allocation + is outside the contract. Row slices `t[a:b]` of these shapes are aligned. +- `weight` is on `x`'s device; the op checks only `x.is_cuda`. +- The GPU has compute capability 10.x: the kernel uses tcgen05 MMA, tensor memory and, at `split > 1`, + thread-block clusters. +- The CuTe DSL (`cutlass`) and `cuda-python` (`cuda.bindings`) are importable. The op's check imports the kernel + module, so without them a call raises `ImportError`, not `ValueError`. + +## Notes + +- Programmatic dependent launch (PDL). With `TRTLLM_ENABLE_PDL` unset or `1` the kernel is launched with PDL; any + other value launches it without. Each CTA issues its whole weight read (TMA, L2 evict-first) without waiting + for the grid dependency; only the read of `x` follows `griddepcontrol.wait`. So `weight` must not be written by + work still running ahead of this call on the stream, while `x` may be. With `trigger_early=True` each CTA executes + `griddepcontrol.launch_dependents` right after issuing its weight loads: the next kernel on the stream, if + launched with PDL, can start while this one runs, and it must execute `griddepcontrol.wait` + (`cudaGridDependencySynchronize`) before it reads `y`. Kernels launched without PDL, torch's included, start + after this one completes as usual. With `trigger_early=False` dependents launch when this grid completes. + None of this changes results. +- Two identical calls return identical bits. The bits depend on `split`, whose summation orders can differ in the + last bf16 place, but not on `push`, `trigger_early`, the PDL setting or `M`. +- The result is not bit-identical to cuBLAS: `F.linear` sums in another order. The two agree within the + tolerance above. +- Grid: `(N / 128) * split` CTAs of 256 threads, in clusters of `split` when `split > 1`. +- Only sm_100 has been measured. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv.py new file mode 100644 index 000000000000..d8d1de17eab3 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv.py @@ -0,0 +1,20 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3 decode GEMV x @ weight^T for M <= 8 bf16 tokens via the CTM (CuTe DSL) kernel.""" + +import torch + +import tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv.op # noqa: F401 (registers torch.ops.trtllm.k3_ctm_gemv*) + + +def k3_ctm_gemv( + x: torch.Tensor, + weight: torch.Tensor, + trigger_early: bool = True, + split: int = 1, + push: bool = False, +) -> torch.Tensor: + """Return `bf16(x @ weight.T)`, fp32-accumulated, in one k3_ctm_gemv call.""" + return torch.ops.trtllm.k3_ctm_gemv( + x, weight, trigger_early=trigger_early, split=split, push=push + ) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_long.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_long.md new file mode 100644 index 000000000000..a455bf61744b --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_long.md @@ -0,0 +1,155 @@ +--- +receipts: {} +--- + +# k3_ctm_gemv_long + +**Wraps** `torch.ops.trtllm.k3_ctm_gemv_long` (one call). + +## Semantics + +The decode GEMV of a Kimi K3 projection with a long reduction dimension (the hidden size 7168, or any `K` too long +to hold in shared memory): at most 8 bf16 tokens times a bf16 weight in `nn.Linear` layout, as one CuTe DSL kernel +that splits each 128-row block of the weight over a cluster of `split` CTAs and streams it through a `ring` of +shared-memory stages. Output columns from `sig_col0` on can hold the sigmoid of the product instead. + +``` +acc[m, n] = sum_k x[m, k] * weight[n, k] fp32 accumulation, 0 <= m < M <= 8, 0 <= n < N +y[m, n] = bf16(acc[m, n]) n < sig_col0, or every n when sig_col0 < 0 +y[m, n] = bf16(sigmoid(bf16(acc[m, n]))) n >= sig_col0 >= 0; sigmoid(v) = 1 / (1 + exp(-v)) in fp32 +``` + +The bf16 products are accumulated in fp32 by the tensor cores (tcgen05 MMAs of 128 weight rows x 8 token columns +x 16 k, into tensor memory). Summation order: `K` is cut into 128-column k-tiles; the last one may be a 64-column +half, for which the TMA reads zeros past `K` in both operands. Rank `r` of a block's cluster accumulates k-tiles +`r, r + split, r + 2 * split, ...` in ascending order into its own fp32 partial (when `split` does not divide the +k-tiles, the first `k_tiles % split` ranks take one more). Each 32-row quarter `q` of the block has an owning rank +(`q`, or `q // 2` at `split=2`), which adds the `split` partials of its rows in rank order `0, 1, ...` in fp32 and +rounds the sum once, round-to-nearest-even. A sigmoid column rounds that sum to bf16, takes the sigmoid of the +bf16 value in fp32 (IEEE division, full-precision `exp`) and rounds again: its bits are exactly `torch.sigmoid` of +the bf16 value the same call returns with `sig_col0=-1`. That is the form of the MLA output gate +`bf16(sigmoid(g))` in the fused `[W_a; W_g]` projection. + +`ring` (the weight stages each CTA keeps in shared memory) and `push` (how partials travel to the owning rank: +DSMEM stores plus a cluster-scope release arrive, or 16-byte `st.async` stores completing the owner's barrier by +bytes) change neither the sums nor their order. Rows past `M` enter the MMA as zeros and are not stored, and rows +past `N` are not stored, so a token's output row has the same bits whatever `M` is. + +Fusion boundary: the single call computes the GEMV and, from `sig_col0` on, the sigmoid. There is no bias, other +activation, residual, all-reduce (under tensor parallelism the per-rank partial is returned as is) or +quantization; the caller owns those. + +## Signature + +```python +def k3_ctm_gemv_long( + x: torch.Tensor, + weight: torch.Tensor, + sig_col0: int = -1, + split: int = 6, + ring: int = 5, + trigger_early: bool = True, + push: bool = False, +) -> torch.Tensor +``` + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `x` | `[M, K]`, `1 <= M <= 8` | bf16 | contiguous, 16-byte-aligned start | CUDA | +| `weight` | `[N, K]` (`nn.Linear` layout) | bf16 | contiguous, 16-byte-aligned start | CUDA, `x`'s device | +| `sig_col0` | scalar | int: first sigmoid column; `< 0` or `>= N`: none | — | — | +| `split` | scalar | int: 2, 4, 5, 6, 7 or 8 CTAs per 128 output rows | — | — | +| `ring` | scalar | int: weight k-tiles staged per CTA (see Preconditions) | — | — | +| `trigger_early` | scalar | bool: let PDL dependents launch early | — | — | +| `push` | scalar | bool: split-K transport | — | — | +| returns | `[M, N]` | bf16 | newly allocated, contiguous | `x`'s device | + +Every element of the returned tensor is written; the inputs are not mutated. + +### Certified arguments + +The Kimi K3 TP16 per-rank shapes with the `sig_col0`, `split`, `ring` and `push` their call sites pass, +`trigger_early=True`, at every `M` in 1..8: + +| Cell | `N` | `K` | `sig_col0` | `split` | `ring` | `push` | +|---|---|---|---|---|---|---| +| MLA `[W_a; W_g]`, gate columns as sigmoid | 2880 | 7168 | 2112 | 6 | 6 | True | +| dense MLP gate_up | 4224 | 7168 | -1 | 4 | 5 | False | +| dense MLP down (last k-tile a half) | 7168 | 2112 | -1 | 2 | 6 | False | +| drafter qkv | 512 | 7168 | -1 | 8 | 6 | True | +| drafter gate_up | 1792 | 7168 | -1 | 8 | 6 | True | +| drafter gate_up, synthetic drafter | 1536 | 7168 | -1 | 8 | 6 | True | +| KDA q/k/v/g/f_a/b (last block 8 rows) | 3208 | 7168 | -1 | 5 | 6 | True | + +Inputs are `x ~ N(0, 1)` and `weight ~ 0.02 * N(0, 1)`, rounded to bf16. The gate per cell: the `sig_col0=-1` +output within `max|y - ref| / max|ref| <= 8e-3` of the fp64 product of the same bf16 inputs; with `sig_col0 >= 0`, +the columns before it bit-identical to that output and the columns from it bit-identical to `torch.sigmoid` of it; +identical bits on a repeated call; and each `M`'s rows bit-identical to the same rows of the 8-row call. Also +certified: + +- one flag changed from the call site's values returns the call site's bits, at every `M`: `push=True` and + `ring=3` at the dense MLP down, `push=True` and `trigger_early=False` at the dense MLP gate_up; +- at the MLA `[W_a; W_g]` and dense MLP down cells, `M` 1 and 8: calls captured in a CUDA graph after an eager + call per key, replayed with `x` rewritten in place, return the bits of eager calls on the new `x`. + +Other values the preconditions admit are accepted but not certified. + +## Metadata consumed + +None. The op reads no attention metadata, KV cache or module state; every input is an argument, and the kernel +runs on the current CUDA stream of `x`'s device. + +One process-global cache sits behind it, and it is result-neutral: the compiled kernel per key +`(N, K, split, ring, trigger_early, push, PDL on/off)`. The first call with a new key compiles the kernel with the +CuTe DSL, which costs host time; `M` and `sig_col0` are runtime arguments and never recompile. That first call +refuses to run under CUDA-graph capture: it raises `RuntimeError` ("run once per shape outside CUDA-graph capture +first") before launching anything. Call every key once eagerly, then capture; captured and later calls reuse the +compiled kernel. The PDL setting is read from `TRTLLM_ENABLE_PDL` on every call, so changing it within a process +adds a key. + +## Preconditions + +The op checks these and raises `ValueError` ("k3_ctm_gemv_long: unsupported call ...") when one fails: + +- `x` is a 2-D contiguous CUDA bf16 tensor with `1 <= M <= 8` rows. +- `weight` is a 2-D contiguous bf16 tensor with `weight.shape[1] == x.shape[1]` and `N >= 1`. `N` needs no + divisibility: the last 128-row block may be partial. +- `K % 64 == 0`. +- `split` is 2, 4, 5, 6, 7 or 8. +- `1 <= ring <= ceil(K / 128) // split` (so `ceil(K / 128) >= split`), and the ring and the resident activation + fit 216 KiB of shared memory: `ring * 32 KiB + (ceil(K / 128) // split + 1) * 2 KiB <= 216 KiB`. That caps + `ring` at 6, and at 5 for `K = 7168` with `split=4`. + +Certified refusals, at `N = 7168`, `K = 2112` (17 k-tiles): 0 and 9 rows, fp16 `x`, a row-strided `x`, +`K = 2080`, `split` 3 and 9, `ring=0`, `ring=3` at `split=8` (2 k-tiles per rank) and `ring=7` at `split=2` +(past shared memory). + +Not checked by the op: + +- `x` and `weight` start at 16-byte-aligned addresses. The op passes both to the kernel with a declared 16-byte + alignment and loads them by TMA; a view whose start is not a multiple of 8 elements past an aligned allocation + is outside the contract. Row slices `t[a:b]` are aligned, since `K % 64 == 0`. +- `weight` is on `x`'s device; the op checks only `x.is_cuda`. +- The GPU has compute capability 10.x: the kernel uses tcgen05 MMA, tensor memory and thread-block clusters. +- The CuTe DSL (`cutlass`) and `cuda-python` (`cuda.bindings`) are importable. The op's check imports the kernel + module, so without them a call raises `ImportError`, not `ValueError`. + +## Notes + +- Programmatic dependent launch (PDL). With `TRTLLM_ENABLE_PDL` unset or `1` the kernel is launched with PDL; any + other value launches it without. Each CTA fills its weight ring (TMA, L2 evict-first) and prefetches the rest of + its weight k-tiles into L2 without waiting for the grid dependency; only the read of `x` follows + `griddepcontrol.wait`. So `weight` must not be written by work still running ahead of this call on the stream, + while `x` may be. With `trigger_early=True` each CTA executes `griddepcontrol.launch_dependents` right after + issuing that first weight traffic: the next kernel on the stream, if launched with PDL, can start while this one + runs, and it must execute `griddepcontrol.wait` (`cudaGridDependencySynchronize`) before it reads `y`. Kernels + launched without PDL, torch's included, start after this one completes as usual. With `trigger_early=False` + dependents launch when this grid completes. None of this changes results. +- Two identical calls return identical bits. The bits depend on `split`, whose summation orders can differ in the + last bf16 place, but not on `ring`, `push`, `trigger_early`, the PDL setting or `M`. `k3_ctm_gemv_wide` at the + same split returns the same bits for each token. +- Grid: `ceil(N / 128) * split` CTAs of 256 threads in clusters of `split`; the op does not limit it to one + wave. +- The result is not bit-identical to cuBLAS: `F.linear` sums in another order. The two agree within the + tolerance above. +- Only sm_100 has been measured. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_long.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_long.py new file mode 100644 index 000000000000..a55c8bc21b7c --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_long.py @@ -0,0 +1,28 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3 long-K decode GEMV x @ weight^T for M <= 8 tokens via the split-K CTM (CuTe DSL) kernel.""" + +import torch + +import tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv.op # noqa: F401 (registers torch.ops.trtllm.k3_ctm_gemv*) + + +def k3_ctm_gemv_long( + x: torch.Tensor, + weight: torch.Tensor, + sig_col0: int = -1, + split: int = 6, + ring: int = 5, + trigger_early: bool = True, + push: bool = False, +) -> torch.Tensor: + """Return `bf16(x @ weight.T)`, columns >= `sig_col0` (if >= 0) as `bf16(sigmoid(.))`, in one call.""" + return torch.ops.trtllm.k3_ctm_gemv_long( + x, + weight, + sig_col0=sig_col0, + split=split, + ring=ring, + trigger_early=trigger_early, + push=push, + ) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_swiglu.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_swiglu.md new file mode 100644 index 000000000000..8ff9e8fb1a4f --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_swiglu.md @@ -0,0 +1,137 @@ +--- +receipts: {} +--- + +# k3_ctm_gemv_swiglu + +**Wraps** `torch.ops.trtllm.k3_ctm_gemv_swiglu` (one call). + +## Semantics + +The down projection of a SwiGLU MLP with the activation folded into it: the `k3_ctm_gemv` GEMV whose input is +`silu(gate) * up` of a gate_up output, computed inside the kernel and never written to memory. + +``` +K = weight.shape[1]; g = gu[m, k], u = gu[m, K + k] (gate columns first) +act[m, k] = bf16( (g * sigmoid(g)) * u ) in fp32; sigmoid(g) = 1 / (1 + exp(-g)) +y[m, n] = bf16( sum_k act[m, k] * weight[n, k] ) 0 <= m < M <= 8 +``` + +The activation is evaluated in fp32 in that order, with IEEE division and the full-precision `exp`, and rounded +to bf16 once, as `silu_and_mul` rounds its output. The GEMV is `k3_ctm_gemv`'s: the products are accumulated in +fp32 by the tensor cores (tcgen05 MMA into tensor memory) and the sum is rounded to bf16 once, round-to-nearest-even. +Summation order: `K` is cut into 128-column k-tiles; each block of 128 output rows runs on `split` CTAs (a +cluster when `split > 1`), rank `r` accumulating k-tiles `r, r + split, ...` in ascending order into its own fp32 +partial (when `split` does not divide the k-tiles, the first `k_tiles % split` ranks take one more); the rank that +owns an output row adds the `split` partials in rank order in fp32 and rounds once. `push` only chooses how +partials travel to the owner (DSMEM stores plus a release arrive, or `st.async` stores completing the owner's +barrier by bytes); the sums and their order are the same. Rows past `M` are zeros and are not stored, so a token's +output row has the same bits whatever `M` is. + +The activation is the formula above, not a call of another kernel. torch's +`(F.silu(gate.float()) * up.float()).bfloat16()` computes the same function, but not necessarily with the same fp32 +operations, so individual activation values can differ from it in the last bf16 place. + +Fusion boundary: the single call computes the activation and the GEMV. The caller owns the gate_up projection +that produces `gu`. There is no bias, all-reduce (under tensor parallelism the per-rank partial is returned as is) +or quantization. + +## Signature + +```python +def k3_ctm_gemv_swiglu( + gu: torch.Tensor, + weight: torch.Tensor, + trigger_early: bool = True, + split: int = 2, + push: bool = False, +) -> torch.Tensor +``` + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `gu` | `[M, 2 * K]`, gate columns first, `1 <= M <= 8` | bf16 | contiguous, 16-byte-aligned start | CUDA | +| `weight` | `[N, K]` (`nn.Linear` layout) | bf16 | contiguous, 16-byte-aligned start | CUDA, `gu`'s device | +| `trigger_early` | scalar | bool: let PDL dependents launch early | — | — | +| `split` | scalar | int: 1, 2 or 4 CTAs per 128 output rows | — | — | +| `push` | scalar | bool: split-K transport | — | — | +| returns | `[M, N]` | bf16 | newly allocated, contiguous | `gu`'s device | + +Every element of the returned tensor is written; the inputs are not mutated. + +### Certified arguments + +The Kimi K3 TP16 per-rank drafter down projection with the `split` and `push` its call site passes, +`trigger_early=True`, at every `M` in 1..8: + +| Cell | `N` | `K` | `split` | `push` | +|---|---|---|---|---| +| drafter down | 7168 | 896 | 2 | True | +| drafter down, synthetic drafter | 7168 | 768 | 2 | True | + +`K = 896` is 7 k-tiles, so the two ranks hold 4 and 3. Inputs are `gu ~ 2 * N(0, 1)` and +`weight ~ 0.02 * N(0, 1)`, rounded to bf16. The gate per cell: `max|y - ref| / max|ref| <= 8e-3` against the fp64 +product of `weight` and torch's bf16 activation `(F.silu(gate.float()) * up.float()).bfloat16()`, identical bits on +a repeated call, and each `M`'s rows bit-identical to the same rows of the 8-row call. Also certified, at the +drafter down: + +- `push=False` and `trigger_early=False` return the call site's bits, at every `M`; +- at `M` 1 and 8, calls captured in a CUDA graph after an eager call per key, replayed with `gu` rewritten in + place, return the bits of eager calls on the new `gu`. + +`split` 1 and 4 and other shapes the preconditions admit are accepted but not certified. + +## Metadata consumed + +None. The op reads no attention metadata, KV cache or module state; every input is an argument, and the kernel +runs on the current CUDA stream of `gu`'s device. + +One process-global cache sits behind it, and it is result-neutral: the compiled kernel per key +`(N, K, split, trigger_early, push, PDL on/off)`. The first call with a new key compiles the kernel with the CuTe +DSL, which costs host time; `M` is a runtime argument and never recompiles. That first call refuses to run under +CUDA-graph capture: it raises `RuntimeError` ("run once per shape outside CUDA-graph capture first") before +launching anything. Call every key once eagerly, then capture; captured and later calls reuse the compiled +kernel. The PDL setting is read from `TRTLLM_ENABLE_PDL` on every call, so changing it within a process adds a +key. + +## Preconditions + +The op checks these and raises `ValueError` ("k3_ctm_gemv_swiglu: unsupported call ...") when one fails: + +- `gu` is a 2-D contiguous CUDA bf16 tensor with `1 <= M <= 8` rows and `gu.shape[1] == 2 * weight.shape[1]`. +- `weight` is a 2-D contiguous bf16 tensor. +- `N % 128 == 0` and `K % 128 == 0`. +- `split` is 1, 2 or 4, `split <= K / 128`, and no rank holds more than 6 k-tiles: + `ceil(K / 128 / split) <= 6`, i.e. `K <= 768` at `split=1`, `K <= 1536` at 2, `K <= 3072` at 4. + +Certified refusals: 0 and 9 rows, fp16 `gu`, a row-strided `gu`, a `gu` of width `1664 != 2 * 896`, `N = 7104`, +`split=3`, and `K = 896` (7 k-tiles) at `split=1`. + +Not checked by the op: + +- `gu` and `weight` start at 16-byte-aligned addresses. The op passes both to the kernel with a declared 16-byte + alignment and loads them by TMA; a view whose start is not a multiple of 8 elements past an aligned allocation + is outside the contract. Row slices `t[a:b]` of these shapes are aligned. +- `weight` is on `gu`'s device; the op checks only `gu.is_cuda`. +- The GPU has compute capability 10.x: the kernel uses tcgen05 MMA, tensor memory and, at `split > 1`, + thread-block clusters. +- The CuTe DSL (`cutlass`) and `cuda-python` (`cuda.bindings`) are importable. The op's check imports the kernel + module, so without them a call raises `ImportError`, not `ValueError`. + +## Notes + +- Programmatic dependent launch (PDL). With `TRTLLM_ENABLE_PDL` unset or `1` the kernel is launched with PDL; any + other value launches it without. Each CTA issues its whole weight read (TMA, L2 evict-first) without waiting + for the grid dependency; only the read of `gu` follows `griddepcontrol.wait`. So `weight` must not be written by + work still running ahead of this call on the stream, while `gu` may be. With `trigger_early=True` each CTA + executes `griddepcontrol.launch_dependents` right after issuing its weight loads: the next kernel on the stream, + if launched with PDL, can start while this one runs, and it must execute `griddepcontrol.wait` + (`cudaGridDependencySynchronize`) before it reads `y`. Kernels launched without PDL, torch's included, start + after this one completes as usual. With `trigger_early=False` dependents launch when this grid completes. None + of this changes results. +- Two identical calls return identical bits. The bits depend on `split`, whose summation orders can differ in the + last bf16 place, but not on `push`, `trigger_early`, the PDL setting or `M`. +- The result is not bit-identical to `silu_and_mul` followed by cuBLAS: the activation can differ in the last bf16 + place (see Semantics) and `F.linear` sums in another order. They agree within the tolerance above. +- Grid: `(N / 128) * split` CTAs of 256 threads, in clusters of `split` when `split > 1`. +- Only sm_100 has been measured. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_swiglu.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_swiglu.py new file mode 100644 index 000000000000..ffb558ebe793 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_swiglu.py @@ -0,0 +1,20 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""SwiGLU down projection silu_and_mul(gu) @ weight^T for M <= 8 tokens via the CTM (CuTe DSL) kernel.""" + +import torch + +import tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv.op # noqa: F401 (registers torch.ops.trtllm.k3_ctm_gemv*) + + +def k3_ctm_gemv_swiglu( + gu: torch.Tensor, + weight: torch.Tensor, + trigger_early: bool = True, + split: int = 2, + push: bool = False, +) -> torch.Tensor: + """Return `bf16(silu_and_mul(gu) @ weight.T)`, gate columns first, in one k3_ctm_gemv_swiglu call.""" + return torch.ops.trtllm.k3_ctm_gemv_swiglu( + gu, weight, trigger_early=trigger_early, split=split, push=push + ) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_wide.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_wide.md new file mode 100644 index 000000000000..1c7a2c0347d7 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_wide.md @@ -0,0 +1,155 @@ +--- +receipts: {} +--- + +# k3_ctm_gemv_wide + +**Wraps** `torch.ops.trtllm.k3_ctm_gemv_wide` (one call). + +## Semantics + +The `k3_ctm_gemv_long` GEMV for up to 64 bf16 tokens (a wide decode step, e.g. speculative verification of several +requests): all `M` token columns go through one tensor-core MMA of 16, 32 or 64 columns, the smallest that holds +`M`, on the long kernel's split-K clusters and weight ring. Split and ring are not arguments; the op picks them +from the shape and the device. Output columns from `sig_col0` on can hold the sigmoid of the product, or the whole +output can be fp32. + +``` +acc[m, n] = sum_k x[m, k] * weight[n, k] fp32 accumulation, 0 <= m < M <= 64, 0 <= n < N +out_fp32=False: y[m, n] = bf16(acc[m, n]) n < sig_col0, or every n when sig_col0 < 0 + y[m, n] = bf16(sigmoid(bf16(acc[m, n]))) n >= sig_col0 >= 0; sigmoid in fp32 +out_fp32=True: y[m, n] = acc[m, n] fp32, not rounded +``` + +The bf16 products are accumulated in fp32 by the tensor cores (tcgen05 MMA into tensor memory). `split` and `ring` +come from `wide_config(N, K, token tile, SM count)` in the op module: the largest split in (8, 7, 6, 5, 4, 2) whose +`ceil(N / 128) * split` CTAs fit one wave of the device's SMs (with `ceil(K / 128) >= split`), then the deepest +weight ring that fits shared memory beside an activation ring of up to 3 stages. The summation order is +`k3_ctm_gemv_long`'s at that split: `K` in 128-column k-tiles (the last may be a 64-column half, read as zeros past +`K`), rank `r` of a 128-row block's cluster accumulating k-tiles `r, r + split, ...` in ascending order into its +own fp32 partial (the first `k_tiles % split` ranks one more when `split` does not divide them), and the `split` +partials of a token added in rank order in fp32 by the rank that reduces that token, then rounded once, +round-to-nearest-even. So a token's row is bit-identical to the row `k3_ctm_gemv_long` returns for that token at +the same split, with any ring and either transport, and it does not depend on `M`. A sigmoid column rounds the sum +to bf16, takes the sigmoid in fp32 (IEEE division, full-precision `exp`) and rounds again: its bits are exactly +`torch.sigmoid` of the bf16 value the same call returns with `sig_col0=-1`. + +The split follows the device's SM count, so the bits can differ between devices with different SM counts. On +148- and 152-SM devices the certified shapes run at the splits listed under *Certified arguments*. + +Fusion boundary: the single call computes the GEMV and either the sigmoid from `sig_col0` on or an fp32 output. +There is no bias, other activation, residual, all-reduce (under tensor parallelism the per-rank partial is +returned as is) or quantization; the caller owns those. + +## Signature + +```python +def k3_ctm_gemv_wide( + x: torch.Tensor, + weight: torch.Tensor, + sig_col0: int = -1, + out_fp32: bool = False, +) -> torch.Tensor +``` + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `x` | `[M, K]`, `1 <= M <= 64` | bf16 | contiguous, 16-byte-aligned start | CUDA | +| `weight` | `[N, K]` (`nn.Linear` layout) | bf16 | contiguous, 16-byte-aligned start | CUDA, `x`'s device | +| `sig_col0` | scalar | int: first sigmoid column, `< N`; `< 0`: none | — | — | +| `out_fp32` | scalar | bool: fp32 output, no rounding (needs `sig_col0 < 0`) | — | — | +| returns | `[M, N]` | bf16, or fp32 with `out_fp32` | newly allocated, contiguous | `x`'s device | + +Every element of the returned tensor is written; the inputs are not mutated. + +### Certified arguments + +The Kimi K3 TP16 per-rank shapes of the projections of a wide decode step, at every `M` in 1..64 (all three token +tiles): + +| Cell | `N` | `K` | `sig_col0` | `out_fp32` | split at 148 / 152 SMs | +|---|---|---|---|---|---| +| KDA q/k/v/g/f_a/b (last block 8 rows) | 3208 | 7168 | -1 | False | 5 | +| MLA `[W_a; W_g]`, gate columns as sigmoid | 2880 | 7168 | 2112 | False | 6 | +| KDA / MLA o_proj | 7168 | 768 | -1 | False | 2 | +| MoE head (latent down slice; router rows) | 280 | 7168 | -1 | True | 8 | +| MoE tail (5 k-tiles) | 7168 | 640 | -1 | False | 2 | +| shared expert gate_up | 768 | 7168 | -1 | False | 8 | +| dense MLP gate_up | 4224 | 7168 | -1 | False | 4 | +| dense MLP down (last k-tile a half) | 7168 | 2112 | -1 | False | 2 | +| drafter qkv | 512 | 7168 | -1 | False | 8 | +| drafter gate_up | 1792 | 7168 | -1 | False | 8 | +| drafter o_proj | 7168 | 384 | -1 | False | 2 | + +Inputs are `x ~ N(0, 1)` and `weight ~ 0.02 * N(0, 1)`, rounded to bf16. The gate per cell: the `sig_col0=-1` +output within `max|y - ref| / max|ref| <= 8e-3` (bf16 output) or `<= 1e-4` (fp32 output) of the fp64 product of +the same bf16 inputs; with `sig_col0 >= 0`, the columns before it bit-identical to that output and the columns +from it bit-identical to `torch.sigmoid` of it; identical bits on a repeated call; and each `M`'s rows +bit-identical to the same rows of the 64-row call. Also certified: + +- at the five shapes that `k3_ctm_gemv_long` runs at up to 8 tokens (MLA `[W_a; W_g]` with split 6, ring 6, + `push=True`; dense MLP gate_up 4, 5, False; dense MLP down 2, 6, False; drafter qkv and drafter gate_up 8, 6, + True), at `M` 1..8 and 16, 24, ..., 64: every token's row is bit-identical to `k3_ctm_gemv_long`'s for that + token, in calls of up to 8 tokens, wherever the device's split equals that call's (all five at 148 or 152 SMs); +- at the MLA `[W_a; W_g]` and MoE head cells, `M` 1 and 64: calls captured in a CUDA graph after an eager call per + key, replayed with `x` rewritten in place, return the bits of eager calls on the new `x`. + +Other shapes the preconditions admit are accepted but not certified. + +## Metadata consumed + +None. The op reads no attention metadata, KV cache or module state; every input is an argument, and the kernel +runs on the current CUDA stream of `x`'s device. The device's SM count enters through `wide_config`. + +One process-global cache sits behind it, and it is result-neutral: the compiled kernel per key +`(N, K, split, ring, activation ring, token tile, out_fp32, PDL on/off)`, where split and the rings follow from +`(N, K, token tile)` and the SM count. On one device that is one compile per `(N, K, out_fp32)` and token tile +(16 for `M <= 16`, 32 for `M <= 32`, 64 above). The first call with a new key compiles the kernel with the CuTe +DSL, which costs host time; `M` within a token tile and `sig_col0` are runtime arguments and never recompile. That +first call refuses to run under CUDA-graph capture: it raises `RuntimeError` ("run once per shape outside +CUDA-graph capture first") before launching anything. Call every key once eagerly, then capture; captured and +later calls reuse the compiled kernel. The PDL setting is read from `TRTLLM_ENABLE_PDL` on every call, so changing +it within a process adds a key. + +## Preconditions + +The op checks these and raises `ValueError` ("k3_ctm_gemv_wide: unsupported call ...") when one fails: + +- `x` is a 2-D contiguous CUDA bf16 tensor with `1 <= M <= 64` rows whose start address is 16-byte aligned. +- `weight` is a 2-D contiguous bf16 tensor with `weight.shape[1] == x.shape[1]` whose start address is 16-byte + aligned. +- `sig_col0 < N`, and `sig_col0 < 0` when `out_fp32`. +- `K % 64 == 0` and `N % 8 == 0`. +- A split exists: `2 * ceil(N / 128) <= SM count` and `ceil(K / 128) >= 2`. On a 148-SM device that is + `N <= 9472`. + +Certified refusals, at `N = 3208`, `K = 7168`: 0 and 65 rows, an `x` starting 2 bytes past a 16-byte boundary, +`sig_col0 = N`, `sig_col0 = 100` with `out_fp32`, `N = 3204`, `K = 7136`, and fp16 `x`. + +Not checked by the op: + +- `weight` is on `x`'s device; the op checks only `x.is_cuda`. +- The GPU has compute capability 10.x: the kernel uses tcgen05 MMA, tensor memory and thread-block clusters. +- The CuTe DSL (`cutlass`) and `cuda-python` (`cuda.bindings`) are importable. The op's check imports the kernel + module, so without them a call raises `ImportError`, not `ValueError`. + +## Notes + +- Programmatic dependent launch (PDL). With `TRTLLM_ENABLE_PDL` unset or `1` the kernel is launched with PDL; any + other value launches it without. Each CTA fills its weight ring (TMA, L2 evict-first) and prefetches the rest of + its weight k-tiles into L2 without waiting for the grid dependency; only the reads of `x` follow + `griddepcontrol.wait`. So `weight` must not be written by work still running ahead of this call on the stream, + while `x` may be. There is no `trigger_early` argument: each CTA always executes + `griddepcontrol.launch_dependents` right after issuing that first weight traffic, as `k3_ctm_gemv_long` does + with `trigger_early=True`. The next kernel on the stream, if launched with PDL, can start while this one runs, + and it must execute `griddepcontrol.wait` (`cudaGridDependencySynchronize`) before it reads `y`. Kernels + launched without PDL, torch's included, start after this one completes as usual. None of this changes results. +- There is no `push` argument either: the split-K partials always travel by 16-byte `st.async` stores, and each + rank reduces the partials of its own share of the tokens. +- Two identical calls return identical bits. The bits depend on the split, so on the device's SM count, but not + on `M` or the PDL setting. +- The result is not bit-identical to cuBLAS: `F.linear` sums in another order. The two agree within the + tolerances above. +- Grid: `ceil(N / 128) * split` CTAs of 256 threads in clusters of `split`; the choice of split keeps that at + most the SM count. +- Only sm_100 has been measured. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_wide.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_wide.py new file mode 100644 index 000000000000..4788c47145d9 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_wide.py @@ -0,0 +1,17 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3 wide decode GEMV x @ weight^T for M <= 64 tokens via the split-K CTM (CuTe DSL) kernel.""" + +import torch + +import tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv.op # noqa: F401 (registers torch.ops.trtllm.k3_ctm_gemv*) + + +def k3_ctm_gemv_wide( + x: torch.Tensor, + weight: torch.Tensor, + sig_col0: int = -1, + out_fp32: bool = False, +) -> torch.Tensor: + """Return `x @ weight.T` in bf16 (columns >= `sig_col0` as sigmoid) or fp32, in one call.""" + return torch.ops.trtllm.k3_ctm_gemv_wide(x, weight, sig_col0=sig_col0, out_fp32=out_fp32) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_decode_gemv.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_decode_gemv.md new file mode 100644 index 000000000000..f58fa5919dfa --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_decode_gemv.md @@ -0,0 +1,112 @@ +--- +receipts: {} +--- + +# k3_decode_gemv + +**Wraps** `torch.ops.trtllm.k3_decode_gemv` (one call). + +## Semantics + +The product of a decode step's activation rows with a projection weight, for at most 8 rows, in bf16: + +``` +y[m, n] = bf16(acc[m, n]) acc[m, n] = sum over k of x[m, k] * weight[n, k], accumulated in fp32 + m < M <= 8, n < N, k < K +``` + +The products are accumulated in fp32 by tcgen05 MMAs (M 128 x N 8 x K 16; the weight is the 128-row operand, the +activation the 8 token columns) into a TMEM accumulator, and each sum is rounded once, to nearest, to bf16. The op +picks one of two kernels from the weight's shape: + +- **Short K** (`N % 128 == 0`, `K % 64 == 0`, `K <= 768`): one CTA per 128-row weight tile accumulates the whole of + `K`, in k order, into one fp32 accumulator. Its whole weight slice is resident in shared memory, one stage per + 128-column k-tile. +- **Split K** (`K % 512 == 0`, `K > 768`, any `N`): a cluster of 4 CTAs per 128-row weight tile. Rank `r` + accumulates k-tiles `r, r + 4, r + 8, ...` in fp32; rank `r` also owns rows `32 r .. 32 r + 31` of the tile, the + other ranks push their fp32 partials of those rows into its shared memory, and it adds the four partials in rank + order, `((p0 + p1) + p2) + p3`, before the one bf16 rounding. A last tile that runs past `N` computes zero-filled + rows there and does not store them. + +The kernel always computes 8 token columns (rows of `x` past `M` arrive as zeros from the TMA), so row `m` of the +result depends on row `m` of `x` only: the result of an `M`-row call is bit-identical to the same rows of an 8-row +call (certified). The summation order is fixed, so the result is bit-identical from run to run (certified). It is not +bit-identical to cuBLAS (`F.linear`), whose accumulation order differs. + +Fusion boundary. Inside: the product and its bf16 rounding. Outside: any bias, activation, scaling or quantization, +the tensor-parallel reduction of a row-parallel projection's partial outputs, and flattening leading dims into `M`. +`x` and `weight` are read only; the result is a new tensor. + +## Signature + +```python +def k3_decode_gemv(x: torch.Tensor, weight: torch.Tensor, trigger_early: bool = True) -> torch.Tensor +``` + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `x` | `[M, K]`, `1 <= M <= 8` | bf16 | contiguous, 16-byte aligned | CUDA | +| `weight` | `[N, K]`, short K or split K (*Semantics*) | bf16 | contiguous, 16-byte aligned | CUDA, `x`'s device | +| `trigger_early` | scalar | Python bool | — | — | +| returns | `[M, N]` | bf16 | contiguous, newly allocated | `x`'s device | + +`trigger_early`: under programmatic dependent launch (PDL), release the next kernel on the stream as soon as every CTA +has issued its weight loads (short K: its whole slice; split K: the first 6 of its k-tiles, or all if it has fewer), +rather than when this grid ends. Scheduling only (*Notes*). + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `x` | `[M, K]`, `M` = 1, 2, ..., 8 | bf16 | contiguous | CUDA | +| `weight` | `[7168, 768]` (short K) and `[3208, 7168]` (split K) | bf16 | contiguous | CUDA | +| `trigger_early` | `True` and `False` | bool | — | — | +| returns | `[M, N]` | bf16 | contiguous | CUDA | + +The weights are Kimi K3's TP16 per-rank KDA o_proj slice (`[7168, 768]`) and KDA input projection slice +(`[3208, 7168]`, whose 26th tile has 8 rows). Every cell: the error is at most 8e-3 of `max |ref|` against an fp64 +torch product; a rerun is bit-identical; the rows are bit-identical to the same rows of the 8-row call; +`trigger_early=False` is bit-identical to `True`. Also certified: a CUDA graph of `M` 1 and `M` 8 calls on each +weight, captured after eager calls and replayed with `x` rewritten in place, bit-identical to eager calls on the same +rows; and the refusals under *Preconditions*. + +## Metadata consumed + +None besides the arguments and the current CUDA stream (the kernel is launched there), plus two pieces of process +state: + +- A per-process compile cache keyed by `(N, K, trigger_early, PDL)`. `M` is a runtime argument, so one compiled + kernel serves every `M`. The first call for a key compiles the kernel (seconds) and must be made eagerly: under + CUDA-graph capture it raises `RuntimeError` ("must run once per shape outside CUDA-graph capture first"). The cache + is result-neutral. +- `TRTLLM_ENABLE_PDL` (default `"1"`), read on every call: whether the kernel is launched with PDL. It is part of the + cache key and changes scheduling, not results. The test runs with the default. + +## Preconditions + +- `x` and `weight` bf16, 2-D and contiguous; `x.shape[1] == weight.shape[1]`; `1 <= M <= 8`; `x` on a CUDA device. +- `(N, K)` taken by one of the kernels: short K (`N % 128 == 0`, `K % 64 == 0`, `ceil(K / 128) <= 6`) or split K + (`K % 128 == 0`, `K / 128 > 6` and a multiple of 4, `N >= 1`). +- A call outside these raises `ValueError` ("k3_decode_gemv: unsupported call ...") before launching anything. + Certified: `M` = 0 and 9, `[128, 1152]` (9 k-tiles: too many for short K, not a multiple of 4 for split K) and + `[200, 768]` (short K needs whole 128-row tiles, split K more than 6 k-tiles). +- Not checked by the op: `weight` on `x`'s device; 16-byte-aligned data pointers (the op declares that alignment to + the kernel; row slices `x[a:b]` of these shapes are aligned); an SM 10.x GPU (the kernels use tcgen05); the CuTe DSL + (`cutlass`) and `cuda-python` (`cuda.bindings`), which the call imports, so without them it raises `ImportError`. +- Under PDL each CTA loads its weight before it waits for the kernel it follows on the stream, so `weight` must not + be written by that kernel; a projection weight is constant during inference. `x` is read, and `y` written, after + the wait. + +## Notes + +- PDL: with `TRTLLM_ENABLE_PDL` = `"1"` the kernel is launched with programmatic dependent launch. Each short-K CTA + issues the loads of its whole weight slice at launch; each split-K CTA loads up to 6 of its k-tiles into its + shared-memory ring and prefetches the rest into L2. Only the activation load waits for the preceding kernel + (`griddepcontrol.wait`), so behind a long predecessor the call pays for the activation load, the MMAs and the store. + A dependent released early by `trigger_early` must itself wait for this grid before reading `y`, as every PDL kernel + waits for its predecessor before reading its output. +- The weight's shared-memory loads carry an L2 EVICT_FIRST hint: a projection's weight streams once per step without + evicting the rest of L2. +- Grid: short K, `N / 128` CTAs of 256 threads; split K, `4 x ceil(N / 128)` CTAs in clusters of 4. +- `trtllm::k3_decode_gemv_tail` (the row-parallel MoE tail on the same kernel, with an RMS-scaled latent part) is a + separate op and not this entry. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_decode_gemv.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_decode_gemv.py new file mode 100644 index 000000000000..17b9d2b2f005 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_decode_gemv.py @@ -0,0 +1,16 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3 decode GEMV: ``x @ weight^T`` for at most 8 bf16 rows on a CuTe DSL kernel (short-K or split-K, chosen +by the weight's shape).""" + +import torch + +import tensorrt_llm._torch.cute_dsl_kernels.k3_decode_gemv.op # noqa: F401 — registers the op + + +def k3_decode_gemv( + x: torch.Tensor, weight: torch.Tensor, trigger_early: bool = True +) -> torch.Tensor: + """Return ``x @ weight^T`` as a new bf16 ``[M, N]`` tensor (fp32 accumulation, one bf16 rounding) in one + k3_decode_gemv call.""" + return torch.ops.trtllm.k3_decode_gemv(x, weight, trigger_early) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_head_gemv.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_head_gemv.md new file mode 100644 index 000000000000..3ea9ebc47f51 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_head_gemv.md @@ -0,0 +1,199 @@ +--- +receipts: {} +--- + +# k3_head_gemv + +**Wraps** `torch.ops.trtllm.k3_head_gemv` (one call). + +## Semantics + +The product of a decode step's activation rows with a large weight (an lm_head vocabulary shard), for at most 8 rows, +in bf16: + +``` +y[m, n] = bf16(acc[m, n]) acc[m, n] = sum over k of x[m, k] * weight[n, k], accumulated in fp32 + m < M <= 8, n < N, k < K +``` + +The weight is cut into 128-row tiles and 128-column k-tiles. The products are accumulated in fp32 by tcgen05 MMAs +(M 128 x N 8 x K 16; the weight is the 128-row operand, the activation the 8 token columns) into TMEM accumulators, +and each sum is rounded once, to nearest, to bf16. Under the stream-K schedule (certified): + +- The `T = (N / 128) x (K / 128)` (tile, k-tile) items, tile-major, are cut into `G = min(SMs, T)` equal contiguous + ranges, one persistent CTA per range: CTA `c` owns items `[c T / G, (c + 1) T / G)`. +- A tile whose items all lie in one range is accumulated by that CTA in k order and stored. +- A tile split over several CTAs has one piece per CTA, in k order. Each piece accumulates its k-tiles in k order in + fp32. Pieces 1, 2, ... store their fp32 partials in the workspace and raise their flag words; the CTA of piece 0 + (the finalizer) waits for those flags and adds the partials after its own accumulator in k order, + `((acc_0 + acc_1) + acc_2) + ...`, before the one bf16 rounding. + +The partition depends only on `(N, K, G)`, so every sum has a fixed order and the result is bit-identical from run to +run on one GPU (certified). A GPU with another SM count partitions the items differently, so its sums may round +differently. The kernel always computes 8 token columns (rows of `x` past `M` arrive as zeros from the TMA), so the +result of an `M`-row call is bit-identical to the same rows of an 8-row call (certified). It is not bit-identical to +cuBLAS (`F.linear`), whose accumulation order differs. + +The `dynamic` schedule (a workspace created with `schedule="dynamic"`; accepted, not certified) cuts each tile's K +into chunks of `chunk_tiles` k-tiles; one persistent CTA per SM takes one (tile, chunk) unit and claims the rest from +a counter, and the last unit of a tile to finish adds the tile's unit partials in chunk order before the one bf16 +rounding. Its order is fixed too. + +Fusion boundary. Inside: the product and its bf16 rounding. Outside: any bias or logit processing (soft-capping, +temperature, fp32 conversion) and the gather of the vocabulary shards across ranks. `x` and `weight` are read only; +the call writes the result and the workspace's buffers (*State*), nothing else. + +## Signature + +```python +def k3_head_gemv( + x: torch.Tensor, + weight: torch.Tensor, + workspace: K3HeadGemvWorkspace, + keep_tiles: int = 0, + ring: int = 6, + prefetch: int = 16, +) -> torch.Tensor +``` + +The wrapper passes `workspace.partials`, `workspace.flags` and `workspace.claim` as the op's state buffers and +`workspace.chunk_tiles` and `workspace.schedule` as its schedule, so a call always runs the schedule its workspace was +created for. `K3HeadGemvWorkspace` is importable from this entry's wrapper module. + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `x` | `[M, K]`, `1 <= M <= 8` | bf16 | contiguous, 16-byte aligned | CUDA | +| `weight` | `[N, K]`, `N % 128 == 0`, `K % 128 == 0` | bf16 | contiguous, 16-byte aligned | CUDA, `x`'s device | +| `workspace` | a `K3HeadGemvWorkspace` for `(N, K)` on `x`'s device (*State*) | — | — | — | +| `keep_tiles` | scalar | Python int | — | — | +| `ring` | scalar | Python int, 1-6 | — | — | +| `prefetch` | scalar | Python int, `>= 0` | — | — | +| returns | `[M, N]` | bf16 | contiguous, the first `M` rows of a new `[8, N]` buffer | `x`'s device | + +`keep_tiles`: weight tiles below it load at normal L2 priority, the rest with an EVICT_FIRST hint (*Notes*). `ring`: +the shared-memory pipeline stages. `prefetch` (stream-K): the k-tiles past the ring that each CTA prefetches into L2 +before its grid-dependency wait. + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `x` | `[M, 7168]`, `M` = 1, 2, ..., 8 | bf16 | contiguous | CUDA | +| `weight` | `[10240, 7168]`, two different weights | bf16 | contiguous | CUDA | +| `workspace` | `K3HeadGemvWorkspace.create(10240, 7168, device)` (stream-K) | — | — | — | +| `keep_tiles` | 0 (every cell); 40 and 80 at `M` 1 and 8 | int | — | — | +| `ring`, `prefetch` | 6 and 16 (the defaults) | int | — | — | +| returns | `[M, 10240]` | bf16 | contiguous | CUDA | + +`[10240, 7168]` is Kimi K3's TP16 per-rank LM-head vocabulary shard (163840 / 16 rows). The `dynamic` schedule is +accepted but not certified. Every cell: the error is at most 8e-3 of `max |ref|` against an fp64 torch product; a +rerun is bit-identical; the rows are bit-identical to the same rows of the 8-row call; the workspace's words are all +zero after the call; `keep_tiles` 40 and 80 are bit-identical to 0. The call sequences and refusals certified on top +of the cells are listed under *State* and *Preconditions*. + +## State + +**Object.** `K3HeadGemvWorkspace` (`tensorrt_llm/_torch/cute_dsl_kernels/k3_head_gemv/op.py`, re-exported by this +entry's wrapper module), owned by the caller: a frozen dataclass of the weight shape (`n_out`, `k_in`), `schedule`, +`chunk_tiles` (0 for stream-K) and three tensors. + +**Contents and size.** With `tiles = N / 128`, `G = min(SMs, tiles x K / 128)` and `P` the most CTAs any tile is +split over (`k3_head_gemv_kernel.streamk_max_pieces(N, K, G)`), a stream-K workspace holds: + +- `partials`, fp32 `[tiles x P x 128 x 8]`: slot `(t, p)` holds the fp32 partial `[128 rows][8 tokens]` of piece + `p > 0` of tile `t` (slot `p = 0` is not used); +- `flags`, int32 `[tiles x P]`: word `(t, p)` is raised to 1 by piece `p > 0` of tile `t` once its partial is stored, + and lowered to 0 by the tile's finalizer once it has read the partial; +- `claim`, int32 `[1]`: not used by stream-K (the dynamic schedule's unit counter). + +At the certified shape on a GPU with 148 or 152 SMs: 80 tiles, 56 k-tiles, `P` = 3, so 240 flag words (960 B) and +245,760 partial floats (960 KiB). The sizes do not depend on `M`: one workspace serves every `M`. A dynamic workspace +holds the fp32 partial of every (tile, chunk) unit, one unit count per tile in `flags` and the unit counter in `claim`. + +**Who creates it, and when.** The target, once per weight shape and schedule, in `post_load_weights`, with +`K3HeadGemvWorkspace.create(n_out, k_in, device, schedule="streamk", chunk_tiles=0)`. It is single-GPU, not +collective, and eager: it allocates `partials` uninitialized and zeroes `flags` and `claim` on the current stream (the +first call must be ordered after that), and it raises `RuntimeError` under CUDA-graph capture (certified). An unknown +schedule raises `ValueError` (certified). The sizes follow the SM count of `device`, so the workspace is created for +the device its calls run on. No environment variable is read. + +**Which ops may share one object.** Only `k3_head_gemv` calls of the workspace's weight shape and schedule on its +device: any `M`, any weight of that shape, any `keep_tiles` (and any `ring` or `prefetch`, which do not change the +layout). At TP16 the target's LM head and DSpark's draft-logits head both have the `[10240, 7168]` shape. Certified: +`M` = 8, 8, 2, 7, 8, 1, 1, 8, 3, 8 back to back on one workspace, and two weights alternating on one workspace, every +call bit-identical to the same call on a fresh workspace. Two workspaces are independent: calls alternating between +two, each with its own weight, are all correct (certified). The op refuses, with `ValueError` before launching, a +workspace whose buffers do not fit the call: another dtype or device, another number of `flags` words, too few +`partials` (certified with a workspace created for `[5120, 7168]`). It checks nothing else, so the sharing rule is the +caller's. + +**Call-order invariant.** A launch on a workspace must start after the previous launch on it has completed: they use +the same flag words and partial slots. The calls that share a workspace are therefore ordered on one stream, which +holds under PDL too: a launch reads and writes the workspace only after its grid-dependency wait +(`griddepcontrol.wait`), that is, after the work before it on the stream has completed. Eager calls from a second +stream must not share the workspace unless the caller orders them (a synchronization or an event between, e.g., a +load-time call on one stream and the decode steps on another). Captured calls replay in capture order on the +replaying stream, so a graph holding calls on a workspace may be replayed on the stream that issues the eager calls +on it, between those calls. Certified: a graph of `M` 1, 8 and 3 calls on one workspace, replayed three times with `x` +rewritten in place and an eager `M` 8 call on the same workspace between replays, every result bit-identical to the +same call on a fresh workspace. Not exercised: two launches running at once on one workspace raise and lower each +other's words, so a finalizer can read the other launch's partial (a silently wrong result) or wait for a word the +other launch has already lowered (a hang). + +**What a later launch reads.** Every word of `flags`, which must be 0 when the launch starts: a finalizer waits until +the word of each later piece of its tile is non-zero and then reads that piece's slot of `partials`. `partials` are +written before they are read within one launch, so their contents between launches do not matter. A dynamic launch +also reads the per-tile unit counts and the unit counter, which must be 0 as well. + +**How it is re-armed.** By the launch itself. Each finalizer lowers its tile's words after reading the partials (in the +dynamic schedule the finalizing unit resets its tile's count, and the last claim rolls the counter back to 0), so a +launch that completes leaves every word at 0 for the next one, whatever its `M` (certified: all words zero after every +sequence above). The layout depends only on `(N, K, schedule, SMs)`; no call depends on an earlier call's `M`. A +launch that never completes can leave words raised; the workspace must then be recreated (`create()` zeroes them). + +## Metadata consumed + +Besides `workspace`, which is an explicit argument: + +- The current CUDA stream: the kernel is launched there. +- A per-process compile cache: stream-K kernels keyed by `(N, K, ring, G, P, prefetch, PDL)`, dynamic ones by + `(N, K, chunk_tiles, ring, G, PDL)`. `M` and `keep_tiles` are runtime arguments. The first call for a key compiles + the kernel (seconds) and must be made eagerly: under CUDA-graph capture it raises `RuntimeError` ("run once per + shape outside CUDA-graph capture first"). The cache is result-neutral. +- The device's SM count, read on every call: it sets `G`, hence the partition and the workspace sizes. +- `TRTLLM_ENABLE_PDL` (default `"1"`), read on every call: whether the kernel is launched with PDL. It is part of the + cache key and changes scheduling, not results. The test runs with the default. + +## Preconditions + +- `x` and `weight` bf16, 2-D and contiguous; `x.shape[1] == weight.shape[1]`; `1 <= M <= 8`; `x` on a CUDA device; + `N % 128 == 0` and `K % 128 == 0`. +- Stream-K: `1 <= ring <= 6`. Dynamic: the workspace's `chunk_tiles` divides `K / 128`. +- A call outside these raises `ValueError` ("k3_head_gemv: unsupported call ...") before launching anything + (certified: `M` = 0 and 9, `N` = 200). +- `workspace` created for this weight shape and schedule on `x`'s device: one whose buffers do not fit the call + raises `ValueError` ("... is not one for weight ...") before launching anything (certified, see *State*). +- `create()` runs outside CUDA-graph capture (else `RuntimeError`), and every key is compiled by an eager call before + a capture uses it (*Metadata consumed*). +- Not checked by the op: `weight` on `x`'s device; 16-byte-aligned data pointers (the op declares that alignment to + the kernel; row slices `x[a:b]` of these shapes are aligned); `prefetch >= 0`; an SM 10.x GPU (the kernel uses + tcgen05); the CuTe DSL (`cutlass`) and `cuda-python` (`cuda.bindings`), which the call and `create()` import, so + without them they raise `ImportError`. +- Under PDL each CTA loads weight tiles (its first `ring` k-tiles, and `prefetch` more into L2) before it waits for + the kernel it follows on the stream, so `weight` must not be written by that kernel; a head weight is constant + during inference. `x` is read, and the workspace and `y` written, after the wait. + +## Notes + +- PDL: with `TRTLLM_ENABLE_PDL` = `"1"` the kernel is launched with programmatic dependent launch. Before + `griddepcontrol.wait` a stream-K CTA only loads weight tiles; it releases its dependents once it has issued all of + its loads (a dynamic CTA, after its last claim). A dependent must itself wait for this grid before reading `y`, as + every PDL kernel waits for its predecessor before reading its output. +- L2: weight tiles `>= keep_tiles` are loaded with an EVICT_FIRST hint, so the streamed shard does not evict the rest + of L2; tiles `< keep_tiles` load at normal priority, so a later reader of the same weight can find them in L2. + Cache policy only: `keep_tiles` 0, 40 and 80 are bit-identical (certified). +- Grid: `G` persistent CTAs of 256 threads (dynamic: one per SM). Stream-K launches at most one CTA per (tile, + k-tile) item, so every CTA's range is non-empty and raises the flag its finalizer waits for, also for a weight with + fewer items than the GPU has SMs. Such weights are not certified here. +- The result is a view of the first `M` rows of a new `[8, N]` buffer: the kernel writes all 8 token rows, the rows + past `M` from the zero-filled activation rows. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_head_gemv.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_head_gemv.py new file mode 100644 index 000000000000..2158f38c825b --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_head_gemv.py @@ -0,0 +1,40 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3 head GEMV: ``x @ weight^T`` for at most 8 bf16 rows and a large weight (an lm_head vocabulary shard) on a +persistent CuTe DSL kernel, over a caller-owned :class:`K3HeadGemvWorkspace`. + +``K3HeadGemvWorkspace`` is re-exported here so that a target creates it through the catalog: once per weight shape +and schedule, eagerly, before any CUDA-graph capture. The contract is the ``## State`` section of ``k3_head_gemv.md``. +""" + +import torch + +# Importing the op module also registers torch.ops.trtllm.k3_head_gemv. +from tensorrt_llm._torch.cute_dsl_kernels.k3_head_gemv.op import K3HeadGemvWorkspace + +__all__ = ["K3HeadGemvWorkspace", "k3_head_gemv"] + + +def k3_head_gemv( + x: torch.Tensor, + weight: torch.Tensor, + workspace: K3HeadGemvWorkspace, + keep_tiles: int = 0, + ring: int = 6, + prefetch: int = 16, +) -> torch.Tensor: + """Return ``x @ weight^T`` as bf16 ``[M, N]`` (fp32 accumulation, one bf16 rounding) in one k3_head_gemv call over + ``workspace``, whose schedule and ``chunk_tiles`` the call takes. The calls that share a workspace must be ordered + on one stream.""" + return torch.ops.trtllm.k3_head_gemv( + x, + weight, + workspace.partials, + workspace.flags, + workspace.claim, + keep_tiles, + workspace.chunk_tiles, + ring, + workspace.schedule, + prefetch, + ) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml index 7382f31d99a1..58d4d516c291 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml @@ -28,7 +28,7 @@ # limits, e.g. no float8 arithmetic, are noted in the entry docstring). # Entries wrapping trtllm ops carry all three. # -# ── RECEIPT STATUS: all 19 entries certified on sm_103 ──────────────────── +# ── RECEIPT STATUS: 19 entries certified on sm_103, 12 on sm_100 ────────── # # A receipt says this entry's test passed on a stated GPU architecture. The # architecture is the whole key: it is a real axis -- see the @@ -66,6 +66,11 @@ # exactly as written. They are true records of what was observed on another # device, and rewriting them would manufacture GB300 evidence that does not # exist. Read them as provenance; the frontmatter is the certification. +# +# The Kimi K3 entries (k3_*, attn_res_*) are certified on sm_100 (B200 / +# GB200), where their first caller runs, and their tests skip on any other +# architecture. cublas_mm and flashinfer_rmsnorm, which that caller also +# uses, keep their sm_103 receipts and add sm_100 ones. # ────────────────────────────────────────────────────────────────────────── # # 30 of the source catalog's 43 entries are here — the union of what the two @@ -133,11 +138,27 @@ entries: impl: torch.ops.trtllm.flashinfer_fused_add_rmsnorm summary: "In-place fused residual add + RMS norm: residual += x; x = rmsnorm(residual) * weight" + - path: norm/k3_embed_norm.py + impl: torch.ops.trtllm.k3_embed_norm + summary: "Kimi K3 decode-step embedding + layer 0's input RMSNorm in one CuTe DSL launch: raw[:] = table[ids] (int32 or int64 ids, N <= 64; zero rows for ids outside [0, V)) written bit for bit into the caller's bf16 [N,H] buffer (the model passes slot 0 of its attention-residual snapshot bank), and a fresh bf16 [N,H] raw * rsqrt(mean(raw^2) + eps) * weight with fp32 statistics, an approximate rsqrt and one rounding, H in (6144, 16384] a multiple of 1024; rows independent of N and bit-reproducible run to run; the op states bit-identity with the gather followed by flashinfer's CuTe DSL RMSNorm (the kernel behind norm/flashinfer_rmsnorm); compiles per (N, V, H, ids dtype, PDL) on its first, eager, call; under PDL it releases its dependents at entry, so a dependent waits on the grid dependency before reading raw or the output" + + - path: norm/attn_res_fwd.py + impl: torch.ops.trtllm.attn_res_fwd + summary: "Kimi K3 attention-residual selection: per token (B == 1), a softmax over N = K+1 candidates (the K bf16 [K,T,1,H] snapshots in order, then the bf16 [T,1,H] layer residual) of fp32 logits rsqrt(mean(v^2) + rms_eps) * , mixing the raw candidates into a fresh bf16 [T,1,H] output (one rounding) and returning the fp32 [N,T,1] rsigma / probs / logits too; N <= 12, T <= 16384, H a multiple of 1024 in [4096, 8192], sm_100 family only; kernel chosen by (T, N) -- a single-CTA or 8-CTA-cluster decode kernel at T == 1 for N in {1,2,4} / {8,12}, else a persistent TMEM online-softmax kernel -- each bit-reproducible run to run" + + - path: norm/attn_res_rmsnorm_fwd.py + impl: torch.ops.trtllm.attn_res_rmsnorm_fwd + summary: "Kimi K3 attention-residual selection fused with the RMSNorm that consumes it: attn_res_fwd's bf16 mixture (that rounding kept), normalized in fp32 with output_rms_eps, rounded to bf16, times the bf16 output_rms_weight and rounded again (KimiK3RMSNorm order), into a fresh bf16 [T,1,7168] tensor; H == 7168 only, N <= 12, T <= 16384, one CTA (N <= 4) or one 8-CTA cluster (N >= 5) per token so rows do not depend on T; with TRTLLM_ENABLE_PDL on, each CTA waits on its grid dependency and then releases its dependents before writing, so a PDL dependent must wait on the grid dependency before reading the output" + # ─── activation ──────────────────────────────────────────────── - path: activation/flashinfer_silu_and_mul.py impl: torch.ops.trtllm.flashinfer_silu_and_mul summary: "SwiGLU activation: silu(x[..., :d]) * x[..., d:] over a packed gate/up last dim" + - path: activation/k3_situ_mul.py + impl: torch.ops.trtllm.k3_situ_mul + summary: "SiTU-gated multiply (SituAndMul) of a packed [gate | up] bf16 gate_up output for M <= 8 tokens on a CuTe DSL kernel: bf16(beta*tanh(g/beta)*sigmoid(g) * v), v = linear_beta*tanh(u/linear_beta) or u, in fp32 with one rounding; launches its PDL dependents at once and reads gu after the grid wait; linear_beta=0.0 refused by the wrapper (the op would run it as 1.0)" + # ─── gemm ────────────────────────────────────────────────────── - path: gemm/cublas_mm.py impl: torch.ops.trtllm.cublas_mm @@ -151,6 +172,30 @@ entries: impl: torch.ops.trtllm.nvfp4_gemm summary: "NVFP4 x NVFP4 dense GEMM in nn.Linear layout with in-kernel block-scale dequantization: alpha * ([M,K/2] e2m1 + 128x4-swizzled e4m3 scales) @ ([N,K/2] e2m1 + swizzled scales)^T (+ per-column bias) -> fresh [M,N] bf16/fp16/fp32 buffer, fp32 accumulation, backend auto-selected (cutlass/cublaslt/cuda_core/cutedsl)" + - path: gemm/k3_ctm_gemv.py + impl: torch.ops.trtllm.k3_ctm_gemv + summary: "Kimi K3 decode GEMV for M <= 8 bf16 tokens on the CTM CuTe DSL kernel (SM 10.x): bf16(x @ W^T) for an [N, K] weight with N % 128 == 0 and at most 6 128-column k-tiles per CTA, fp32 tcgen05 accumulation with one rounding, one CTA or a 2- or 4-CTA split-K cluster per 128 rows (partials summed in rank order); the weight is read before the PDL grid wait" + + - path: gemm/k3_ctm_gemv_swiglu.py + impl: torch.ops.trtllm.k3_ctm_gemv_swiglu + summary: "SwiGLU down projection for M <= 8 tokens on the CTM CuTe DSL kernel (SM 10.x): bf16(bf16(silu(gate) * up) @ W^T) from a packed [gate | up] bf16 gate_up output, the activation computed in fp32 in the GEMV's shared memory and never stored, fp32 accumulation, 1/2/4-CTA split-K summed in rank order" + + - path: gemm/k3_ctm_gemv_long.py + impl: torch.ops.trtllm.k3_ctm_gemv_long + summary: "Long-K decode GEMV for M <= 8 bf16 tokens on the CTM CuTe DSL kernel (SM 10.x): bf16(x @ W^T) with each 128-row block split over a cluster of 2 or 4-8 CTAs, its weight streamed through a shared-memory ring filled before the PDL grid wait, partials summed in rank order; columns >= sig_col0 as bf16(sigmoid(bf16(.))) (torch.sigmoid of the plain output, bitwise); any N, K % 64 == 0" + + - path: gemm/k3_ctm_gemv_wide.py + impl: torch.ops.trtllm.k3_ctm_gemv_wide + summary: "Wide decode GEMV for M <= 64 bf16 tokens on the long CTM kernel (SM 10.x): one MMA of 16/32/64 token columns, split and ring picked from the shape and the SM count (one wave), each token's row bit-identical to k3_ctm_gemv_long's at the same split; bf16 output with optional sigmoid columns from sig_col0, or unrounded fp32; N % 8 == 0, K % 64 == 0" + + - path: gemm/k3_decode_gemv.py + impl: torch.ops.trtllm.k3_decode_gemv + summary: "Kimi K3 decode GEMV for at most 8 rows: bf16 x [M,K] @ weight [N,K]^T into a fresh bf16 [M,N], fp32 accumulation (tcgen05 into TMEM) and one rounding, on a CuTe DSL kernel the weight's shape selects -- short K (K <= 768, N a multiple of 128: one CTA per 128-row tile with the whole weight slice in shared memory) or split K (K a multiple of 512 above 768: 4-CTA clusters, the partials added in rank order); rows independent of M and bit-reproducible run to run, not bit-identical to cuBLAS; SM 10.x; compiles per (N, K, trigger_early, PDL) on its first, eager, call; under PDL the weight is loaded before the grid-dependency wait, so it must not be written by the preceding kernel" + + - path: gemm/k3_head_gemv.py + impl: torch.ops.trtllm.k3_head_gemv + summary: "Kimi K3 vocabulary-shard head GEMV for at most 8 rows: bf16 x [M,K] @ weight [N,K]^T (N, K multiples of 128) into bf16 [M,N] (the first M rows of a fresh [8,N]) on a persistent stream-K CuTe DSL kernel of min(SMs, tiles x k-tiles) CTAs, a split tile's fp32 piece partials added in k order before one rounding -- rows independent of M, bit-reproducible run to run on one GPU; over a caller-owned K3HeadGemvWorkspace (fp32 partials and one flag word per (tile, piece); created eagerly, before any capture, for one weight shape and schedule; one that does not fit the call is refused with ValueError) that every completed launch leaves at zero, so the calls sharing one are ordered on one stream, eager calls and graph replays alike; weight tiles past keep_tiles streamed EVICT_FIRST; the dynamic schedule accepted, not certified; SM 10.x; compiles per shape on its first, eager, call" + # ─── attention ───────────────────────────────────────────────── - path: attention/fused_qk_norm_rope.py impl: torch.ops.trtllm.fused_qk_norm_rope diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_fwd.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_fwd.md new file mode 100644 index 000000000000..30a43426bb80 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_fwd.md @@ -0,0 +1,152 @@ +--- +receipts: {} +--- + +# attn_res_fwd + +**Wraps** `torch.ops.trtllm.attn_res_fwd` (one call). + +## Semantics + +Attention-residual selection: Kimi K3's per-token mixing of residual-stream snapshots. For every token the op +scores `N = K + 1` candidate rows -- the `K` snapshots of `block_residual` in their stored order, then +`layer_residual` -- and returns the softmax-weighted sum of the raw candidates. + +For token `t` (`B == 1`) and candidate `n`, every step in fp32 (`H` is the hidden size): + +``` +V[n] = block_residual[n, t, 0, :] for n < K +V[K] = layer_residual[t, 0, :] +q = fp32(rms_weight) * fp32(res_weight) # elementwise; never rounded to bf16 +rsigma[n] = rsqrt(sum_h V[n, h]^2 / H + rms_eps) +logits[n] = rsigma[n] * sum_h V[n, h] * q[h] +probs = softmax(logits) # over the token's N candidates +output = bf16_rn(sum_n probs[n] * V[n]) # the only rounding +``` + +`logits[n]` is the dot product of `res_weight` with the RMSNorm of `V[n]` under weight `rms_weight`, the +normalized row kept in fp32, as in HF Kimi's `_apply_attn_res`. The RMSNorm only scores the candidates: `output` +mixes the raw rows, not their normalized form. `rsigma`, `probs` and `logits` are returned as the fp32 values the +formulas name. With `N == 1` the single probability is 1, and the formulas return `layer_residual` as `output`. + +The kernels are compiled with `--use_fast_math`: `rsqrt`, the softmax exponentials and the reciprocal of the +softmax denominator are hardware approximations, denormals flush to zero, and every sum runs in the kernel's own +order. The results are close to an fp32 torch evaluation of the formulas, not bit-identical to it. + +Fusion boundary: selection only. No residual add, no trailing RMSNorm (`attn_res_rmsnorm_fwd` fuses the norm that +follows), and no snapshot write: the caller owns the snapshot bank, including appending the running residual to +it. + +## Signature + +```python +def attn_res_fwd( + layer_residual: torch.Tensor, + block_residual: torch.Tensor, + res_weight: torch.Tensor, + rms_weight: torch.Tensor, + rms_eps: float, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor] +``` + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `layer_residual` | `[T, 1, H]` | bf16 | contiguous | CUDA | +| `block_residual` | `[K, T, 1, H]`, `K = N - 1 >= 0` | bf16 | contiguous | `layer_residual`'s (not checked) | +| `res_weight` | `H` elements (`[H]`, `[1, H]` or `[H, 1]`) | bf16 | contiguous | `layer_residual`'s (not checked) | +| `rms_weight` | `H` elements | bf16 | contiguous | `layer_residual`'s (not checked) | +| `rms_eps` | scalar | Python float, passed as fp32 | — | — | +| returns `[0]` (`output`) | `[T, 1, H]` | bf16 | contiguous, newly allocated | `layer_residual`'s | +| returns `[1]` (`rsigma`) | `[N, T, 1]` | fp32 | contiguous, newly allocated | `layer_residual`'s | +| returns `[2]` (`probs`) | `[N, T, 1]` | fp32 | contiguous, newly allocated | `layer_residual`'s | +| returns `[3]` (`logits`) | `[N, T, 1]` | fp32 | contiguous, newly allocated | `layer_residual`'s | + +On the candidate axis of `rsigma`, `probs` and `logits`, index `n < K` is snapshot `n` and index `K` is +`layer_residual`. `K = 0` is passed as a zero-size `[0, T, 1, H]` tensor. No input is mutated and no output +aliases an input. `res_weight` is the flattened `[1, H]` weight of the scoring `nn.Linear(H, 1)` projection; +`rms_weight` is the scoring RMSNorm's weight. + +### Certified arguments + +- `H = 7168` (Kimi K3), `B = 1`, bf16 inputs, `res_weight` and `rms_weight` of shape `[H]`. +- Every `N = 2 ... 9` -- the candidate counts Kimi K3 produces: its 93 layers append a snapshot every 12 layers, so + the bank holds at most `ceil(93 / 12) = 8` snapshots -- at `T = 1 ... 8` (decode, including multi-token + speculative steps) and at `T = 300` and `T = 2048` (prefill). +- The remaining dispatch branches and `N` extremes: `N = 1, 10, 11, 12` at `T = 1`; `N = 1, 12` at `T = 8`; + `N = 12` at `T = 1024`. +- `rms_eps = 1e-6` in every cell; `1e-5` and `1e-2` in four cells that cover the three kernels. `1e-2` exceeds the + test inputs' mean square of `2.5e-3`, so it pins where `rms_eps` enters. + +## Metadata consumed + +None. Stateless: no attention metadata, no workspace, no caller-visible state. The op caches per device the SM +count, the architecture check and the one-time raise of a kernel's dynamic shared-memory limit, and reads +`TRTLLM_ENABLE_PDL` once per process. These select launch parameters, never results. + +## Preconditions + +Each item is a `TORCH_CHECK` in `cpp/tensorrt_llm/thop/attnResOp.cpp`, evaluated in this order; a failure raises +`RuntimeError` with the quoted message. This entry's test does not exercise the rejections. + +1. `layer_residual.dim() == 3` and `block_residual.dim() == 4` (`attn_res_fwd: layer_residual must be [T, B, H]`, + `attn_res_fwd: block_residual must be [K, T, B, H]`). +2. All four tensors are CUDA tensors (`attn_res_fwd: all input tensors must be CUDA tensors`). +3. `layer_residual`'s device has compute capability `10.x` (`attn_res_fwd requires an sm_100-family (datacenter + Blackwell) GPU`). +4. `B == 1`; `1 <= N <= 12` with `N = block_residual.shape[0] + 1`; `1 <= T <= 16384`; `H` a multiple of 1024 in + `[4096, 8192]` (`attn_res_fwd: unsupported B=...`, `N=...`, `T=...`, `H=...`). +5. All four tensors are bf16 (`attn_res_fwd: must be bf16`) and contiguous (`attn_res_fwd: inputs must be + contiguous`). A strided view raises; it is never read as if it were dense. +6. `block_residual.shape == (N - 1, T, B, H)` (`attn_res_fwd: block_residual shape must match layer_residual`). +7. `res_weight.numel() == H` and `rms_weight.numel() == H` (`attn_res_fwd: must have H elements`). + +Not checked: that `block_residual`, `res_weight` and `rms_weight` are on `layer_residual`'s device. The kernels +dereference all four on that device, so the caller must keep them together (`attn_res_rmsnorm_fwd` checks it). + +## Notes + +- **Code paths.** The op picks a kernel from `(T, N)`. At `H = 7168`: + + | `T`, `N` | Kernel | Launch | + |---|---|---| + | `T == 1`, `N` in {1, 2, 4} | single-CTA decode kernel | 1 CTA of 256 threads | + | `T == 1`, `N` in {8, 12} | split-K decode kernel | 1 cluster of 8 CTAs of 256 threads | + | `T == 1024`, `N == 12` | online kernel, fixed-`N = 12` variant | persistent: (SM count - 1) CTAs of 288 threads | + | every other `(T, N)` | online kernel | persistent: one CTA of 288 threads per SM | + + The single-CTA kernel gives each of its 256 threads 28 of the 7168 elements and keeps every candidate's values + for them in registers. The split-K kernel gives each CTA of the cluster 896 of the 7168 elements, keeps them in + shared memory as fp32, and exchanges per-candidate sums through distributed shared memory. The online kernel is + warp-specialized: one warp streams candidate rows into shared memory with bulk asynchronous copies, eight warps + score them and keep the fp32 rows in Tensor Memory until the mixing pass, the softmax runs online over chunks of + four candidates, and each CTA loops over tokens. Other hidden sizes run the online kernel with other tilings, or + a row-tiled kernel at `N == 1` for `H` 4096 and 8192; none of them is certified here. +- **Architecture.** The op requires compute capability `10.x`, and the kernels are built only for the sm_100 + family target (`100f`). They use sm_100 instructions: paired fp32 arithmetic (`.f32x2`) in all three kernels, + Tensor Memory (`tcgen05`) and bulk asynchronous copies in the online kernel, which also takes 115,040 bytes of + dynamic shared memory per CTA at `H = 7168` (one CTA per SM), and an 8-CTA thread-block cluster with distributed + shared memory in the split-K kernel. +- **PDL.** `TRTLLM_ENABLE_PDL` (read once per process; unset or `1` enables it, `0` disables it) controls + programmatic dependent launch for the two decode kernels. When it is enabled they launch with programmatic + stream serialization, wait on the grid dependency (`cudaGridDependencySynchronize`) before reading any input, + and trigger their dependents (`cudaTriggerProgrammaticLaunchCompletion`) after their last store. The online + kernel launches without the attribute: it starts after its predecessor completes and releases its dependents + when it completes. A dependent kernel launched with PDL must wait on the grid dependency before it reads any + output: the trigger only lets it start, and only the wait makes this op's stores visible. PDL changes scheduling, + not arithmetic. +- **Determinism.** Each kernel reduces in a fixed order without atomics, so identical inputs give identical bits + run to run; the test asserts it in every cell. The kernels' summation orders differ from each other, so a token + that a decode kernel computes in a `T == 1` call can come out different in the last bits when the online kernel + computes it inside a larger call. The online kernel's per-token arithmetic depends on neither `T` nor the grid + size, except that its fixed-`N = 12` variant sums the softmax denominator in a different order. +- **Tolerance.** The test compares `output` and each fp32 output with an fp32 torch evaluation of *Semantics*, + using the metric of main's op test (`tests/unittest/_torch/modules/kimi_k3_attn_res/test_attn_res_op.py`): + cosine similarity above 0.999 and relative L2 error below 3e-2. +- **CUDA graphs.** The call does not synchronize with the host, and its launch configuration depends on shapes + only. The test captures one call per kernel after an eager warm-up and checks that replay reproduces the eager + bits. +- **Kimi K3 call sites.** `modeling_kimi_linear.py` calls this op, followed by a separate RMSNorm, where the fused + `attn_res_rmsnorm_fwd` is not taken: above `KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS` tokens (1 by default) unless + `KIMI_K3_ATTN_RES_TOPOLOGY` routes the call to the persistent fused op, and on the layers whose pre-norm mixture + the speculative drafter captures. The `register_fake` in `custom_ops/cpp_custom_ops.py` reports the shapes and + dtypes above. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_fwd.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_fwd.py new file mode 100644 index 000000000000..f123f0ba7930 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_fwd.py @@ -0,0 +1,28 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3 attention-residual selection: a per-token softmax mixture of residual snapshots.""" + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + + +def attn_res_fwd( + layer_residual: torch.Tensor, + block_residual: torch.Tensor, + res_weight: torch.Tensor, + rms_weight: torch.Tensor, + rms_eps: float, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Mix each token's N = K + 1 candidates: the K snapshots of block_residual, then layer_residual. + + A candidate's logit is its fp32 RMSNorm (weight `rms_weight`, `rms_eps`) dotted with `res_weight`; + the softmax of the logits weights the raw candidates. `layer_residual` is bf16 [T, 1, H] and + `block_residual` bf16 [K, T, 1, H], both contiguous. + + Returns `(output, rsigma, probs, logits)`: the bf16 [T, 1, H] mixture and the fp32 [N, T, 1] + per-candidate rsqrt(mean square + eps), softmax probabilities and logits, all newly allocated. + """ + return torch.ops.trtllm.attn_res_fwd( + layer_residual, block_residual, res_weight, rms_weight, rms_eps + ) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_rmsnorm_fwd.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_rmsnorm_fwd.md new file mode 100644 index 000000000000..71657064b3b7 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_rmsnorm_fwd.md @@ -0,0 +1,161 @@ +--- +receipts: {} +--- + +# attn_res_rmsnorm_fwd + +**Wraps** `torch.ops.trtllm.attn_res_rmsnorm_fwd` (one call). + +## Semantics + +Attention-residual selection followed by the RMSNorm that consumes it, in one kernel. The selection is +`attn_res_fwd`'s: for every token, a softmax over `N = K + 1` candidate rows -- the `K` snapshots of +`block_residual` in their stored order, then `layer_residual` -- of fp32 logits, mixing the raw candidates. The +fused norm then normalizes that mixture. + +For token `t` (`B == 1`), every step in fp32 unless rounded explicitly (`H = 7168`): + +``` +V[n] = block_residual[n, t, 0, :] for n < K +V[K] = layer_residual[t, 0, :] +q = fp32(rms_weight) * fp32(res_weight) # elementwise; never rounded to bf16 +rsigma[n] = rsqrt(sum_h V[n, h]^2 / H + rms_eps) +logits[n] = rsigma[n] * sum_h V[n, h] * q[h] +probs = softmax(logits) # over the token's N candidates +mixed = bf16_rn(sum_n probs[n] * V[n]) # attn_res_fwd's output +r = rsqrt(sum_h mixed[h]^2 / H + output_rms_eps) +normed = bf16_rn(mixed * r) +output = bf16_rn(normed * fp32(output_rms_weight)) +``` + +The op keeps two bf16 rounding boundaries: + +- `mixed` is rounded to bf16 before the norm reads it, as when `attn_res_fwd` and a separate RMSNorm run in turn. +- The norm rounds twice, in `KimiK3RMSNorm`'s order: the normalized value is rounded to bf16, then multiplied by + the bf16 weight and rounded again. An RMSNorm that rounds once (normalize and scale in fp32, then cast, as the + `flashinfer_rmsnorm` entry describes) differs from this op by one unit in the last bf16 place in a large share of + the elements: about a fifth of them when both formulas are evaluated in torch on this entry's test inputs. + +As in `attn_res_fwd`, the kernels are compiled with `--use_fast_math`: `rsqrt`, the softmax exponentials and the +reciprocal of the softmax denominator are hardware approximations, denormals flush to zero, and every sum runs in +the kernel's own order. The output is close to an fp32 torch evaluation of the formulas, not bit-identical to it. + +Fusion boundary: the selection and the RMSNorm that follows it. The op returns only the normalized output -- not +the pre-norm mixture and not `attn_res_fwd`'s `rsigma`, `probs` or `logits`. No residual add (the separate op +`attn_res_add_rmsnorm_fwd` adds one first). + +## Signature + +```python +def attn_res_rmsnorm_fwd( + layer_residual: torch.Tensor, + block_residual: torch.Tensor, + res_weight: torch.Tensor, + rms_weight: torch.Tensor, + output_rms_weight: torch.Tensor, + rms_eps: float, + output_rms_eps: float, +) -> torch.Tensor +``` + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `layer_residual` | `[T, 1, 7168]` | bf16 | contiguous | CUDA | +| `block_residual` | `[K, T, 1, 7168]`, `K = N - 1 >= 0` | bf16 | contiguous | `layer_residual`'s | +| `res_weight` | 7168 elements (`[H]`, `[1, H]` or `[H, 1]`) | bf16 | contiguous | `layer_residual`'s | +| `rms_weight` | 7168 elements | bf16 | contiguous | `layer_residual`'s | +| `output_rms_weight` | 7168 elements | bf16 | contiguous | `layer_residual`'s | +| `rms_eps` | scalar | Python float, passed as fp32 | — | — | +| `output_rms_eps` | scalar | Python float, passed as fp32 | — | — | +| returns | `[T, 1, 7168]` | bf16 | contiguous, newly allocated | `layer_residual`'s | + +`res_weight` and `rms_weight` score the candidates, with `rms_eps`; `output_rms_weight` and `output_rms_eps` +belong to the trailing norm. `K = 0` is passed as a zero-size `[0, T, 1, 7168]` tensor. No input is mutated and the +output aliases no input. + +### Certified arguments + +- `H = 7168`, `B = 1`, bf16 inputs, `res_weight`, `rms_weight` and `output_rms_weight` of shape `[H]`. +- Every `N = 2 ... 9` -- the candidate counts Kimi K3 produces: its 93 layers append a snapshot every 12 layers, so + the bank holds at most `ceil(93 / 12) = 8` snapshots -- at `T = 1 ... 8` (decode, including multi-token + speculative steps) and at `T = 32` and `T = 300`. +- `N = 1, 10, 11, 12` at `T = 1` and `T = 8`. +- `rms_eps = output_rms_eps = 1e-6` in every cell; the pairs `(1e-5, 1e-5)`, `(1e-2, 1e-6)` and `(1e-6, 1e-2)` in + three cells that cover both kernels. `1e-2` exceeds the mean square of both the test inputs and their mixture, + so the two unequal pairs pin which norm each eps feeds. + +## Metadata consumed + +None. Stateless: no attention metadata, no workspace, no caller-visible state. The op caches per device the +architecture check and the one-time raise of the split-K kernel's dynamic shared-memory limit, and reads +`TRTLLM_ENABLE_PDL` once per process. These select launch parameters, never results. + +## Preconditions + +Each item is a `TORCH_CHECK` in `cpp/tensorrt_llm/thop/attnResOp.cpp`, evaluated in this order; a failure raises +`RuntimeError` with the quoted message. This entry's test does not exercise the rejections. + +1. `layer_residual.dim() == 3` and `block_residual.dim() == 4` (`attn_res_rmsnorm_fwd: layer_residual must be + [T, B, H]`, `attn_res_rmsnorm_fwd: block_residual must be [K, T, B, H]`). +2. All five tensors are CUDA tensors on one device (`attn_res_rmsnorm_fwd: all input tensors must be CUDA + tensors`, `... must be on the same CUDA device`). +3. The device has compute capability `10.x`; `B == 1`; `1 <= N <= 12` with `N = block_residual.shape[0] + 1`; + `1 <= T <= 16384`; `H` a multiple of 1024 in `[4096, 8192]`. This is the contract check `attn_res_fwd` shares, + and its messages name that op: `attn_res_fwd requires an sm_100-family (datacenter Blackwell) GPU`, + `attn_res_fwd: unsupported B=...`, `N=...`, `T=...`, `H=...`. +4. `H == 7168` (`attn_res_rmsnorm_fwd: requires B=1 and H=7168, got T=... B=... H=...`): the kernels derive their + per-thread tiling from `H` at compile time. +5. All five tensors are bf16 (`attn_res_rmsnorm_fwd: must be bf16`) and contiguous + (`attn_res_rmsnorm_fwd: inputs must be contiguous`). A strided view raises; it is never read as if it were dense. +6. `block_residual.shape == (N - 1, T, B, H)` (`attn_res_rmsnorm_fwd: block_residual shape must match + layer_residual`). +7. `res_weight`, `rms_weight` and `output_rms_weight` each have `H` elements (`attn_res_rmsnorm_fwd: must + have H elements`). + +The kernels take any token count as a grid dimension, but the shared contract check still caps `T` at 16384. + +## Notes + +- **Code paths.** The kernel depends on `N` only, never on `T`: + + | `N` | Kernel | Launch | + |---|---|---| + | 1 ... 4 | single-CTA kernel | `T` CTAs of 256 threads, one per token | + | 5 ... 12 | split-K kernel | `T` clusters of 8 CTAs of 256 threads, one per token | + + They are the templates of `attn_res_fwd`'s two decode kernels, instantiated with the trailing norm for every `N` + (`attn_res_fwd` itself runs them only at `T == 1`, for five `N` values). The single-CTA kernel gives each of its + 256 threads 28 of the 7168 elements and keeps every candidate's values for them in registers as bf16. In the + split-K kernel each CTA of the cluster owns 896 of the 7168 elements, keeps them in shared memory as fp32, and + exchanges the per-candidate sums, then the mixture's sum of squares, through distributed shared memory. This op + has no persistent path; the persistent fused kernel is the separate op `attn_res_add_rmsnorm_persistent_fwd`. +- **Architecture.** The op requires compute capability `10.x`, and the kernels are built only for the sm_100 + family target (`100f`). Both use paired fp32 arithmetic (`.f32x2`); the split-K kernel also uses an 8-CTA + thread-block cluster with distributed shared memory and takes `3,652 * N` bytes of dynamic shared memory per CTA. +- **PDL: the output is written after the dependents are released.** `TRTLLM_ENABLE_PDL` (read once per process; + unset or `1` enables it, `0` disables it) controls programmatic dependent launch. When it is enabled both + kernels launch with programmatic stream serialization. Every CTA first waits on the grid dependency + (`cudaGridDependencySynchronize`), so the op reads its inputs only after the preceding kernel has completed, and + then immediately triggers its dependents (`cudaTriggerProgrammaticLaunchCompletion`), before it computes + anything. A dependent kernel launched with PDL can therefore start while this op is still writing `output`: it + must wait on the grid dependency before it reads `output` or overwrites any of this op's inputs. A kernel + launched without the PDL attribute is ordered by the stream as usual. With PDL disabled the kernels launch + normally and contain no grid-dependency instructions. PDL changes scheduling, not arithmetic. The test chains two + calls, the second reading the first's output as its `layer_residual`, and checks the chained result bit for bit + against the same call on a settled copy. +- **Determinism.** Every token is computed by its own CTA or cluster, with the same code at every `T`, reducing in + a fixed order without atomics. Identical inputs give identical bits run to run, which the test asserts in every + cell, and a token's result does not depend on `T` or on its position in the call. +- **Tolerance.** The test compares the output with an fp32 torch evaluation of *Semantics*, both rounding + boundaries included, using the metric of main's op test + (`tests/unittest/_torch/modules/kimi_k3_attn_res/test_attn_res_rmsnorm_op.py`): cosine similarity above 0.9999 + and relative L2 error below 5e-3. +- **CUDA graphs.** The call does not synchronize with the host, and its launch configuration depends on shapes + only. The test captures one call per kernel after an eager warm-up and checks that replay reproduces the eager + bits. +- **Kimi K3 call sites.** `modeling_kimi_linear.py` uses this op for the attention-residual + RMSNorm pairs before + attention, before the MLP on the layers that append a snapshot, and at the model output, for calls of at most + `KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS` tokens (1 by default, 32 when `KIMI_K3_ATTN_RES_TOPOLOGY` is on) that the + persistent topology does not take. Layers whose pre-norm mixture the speculative drafter captures run + `attn_res_fwd` and a separate norm instead. The `register_fake` in `custom_ops/cpp_custom_ops.py` returns + `empty_like(layer_residual)`. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_rmsnorm_fwd.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_rmsnorm_fwd.py new file mode 100644 index 000000000000..e5d9087477b3 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_rmsnorm_fwd.py @@ -0,0 +1,34 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3 attention-residual selection fused with the RMSNorm that consumes it.""" + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + + +def attn_res_rmsnorm_fwd( + layer_residual: torch.Tensor, + block_residual: torch.Tensor, + res_weight: torch.Tensor, + rms_weight: torch.Tensor, + output_rms_weight: torch.Tensor, + rms_eps: float, + output_rms_eps: float, +) -> torch.Tensor: + """Return rmsnorm(attn_res_fwd's bf16 mixture) * output_rms_weight as a new bf16 [T, 1, 7168] tensor. + + `res_weight`, `rms_weight` and `rms_eps` score the candidates as in `attn_res_fwd`; + `output_rms_weight` and `output_rms_eps` belong to the trailing norm, which rounds the normalized + value to bf16 before the weight multiply. With PDL enabled the kernel releases its dependents before + it writes the output, so a dependent PDL kernel must wait on the grid dependency before reading it. + """ + return torch.ops.trtllm.attn_res_rmsnorm_fwd( + layer_residual, + block_residual, + res_weight, + rms_weight, + output_rms_weight, + rms_eps, + output_rms_eps, + ) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/k3_embed_norm.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/k3_embed_norm.md new file mode 100644 index 000000000000..b1432ba151f1 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/k3_embed_norm.md @@ -0,0 +1,119 @@ +--- +receipts: {} +--- + +# k3_embed_norm + +**Wraps** `torch.ops.trtllm.k3_embed_norm` (one call). + +## Semantics + +A decode step's embedding rows and the first layer's input RMSNorm, in one launch. For `N <= 64` token ids, a bf16 +table `[V, H]` and a bf16 norm weight `[H]`: + +``` +row[t] = table[ids[t], :] if 0 <= ids[t] < V, else a zero row t < N +raw[t, :] = row[t] (written into the caller's raw) +ss[t] = sum over h of fp32(row[t, h])^2 fp32 +rstd[t] = rsqrt(ss[t] / H + eps) fp32, approximate (fast-math) rsqrt +out[t, h] = bf16((fp32(row[t, h]) * rstd[t]) * fp32(weight[h])) one bf16 rounding +``` + +`raw` receives the rows bit for bit. One 128-thread CTA computes one token's row: each thread squares and sums its +`H / 128` values, the sum is reduced across the warp's lanes, then across the 4 warps through shared memory, in a +fixed order; so the result is bit-identical from run to run (certified), and a row's result depends only on its id, +not on `N` or the other ids (certified: every row bit-identical to the same id's row of a 64-id call). An id outside +`[0, V)` reads nothing from the table: its `raw` row and its output row are zeros (certified). + +The op states that `out` is bit-identical to the gather (`trtllm::k3_embed`) followed by `flashinfer.norm.rmsnorm` on +flashinfer's CuTe DSL RMSNorm kernel, the kernel the catalog's `norm/flashinfer_rmsnorm` runs unless +`FLASHINFER_USE_CUDA_NORM=1`: it keeps that kernel's geometry for `6144 < H <= 16384` (one 128-thread CTA per row) +and its arithmetic order. The op's own test checks that bit for bit. This entry's test checks `out` against a native +fp64 torch RMSNorm, within 8e-3 of each row's largest magnitude, and `raw` bit for bit against a torch gather. + +Fusion boundary. Inside: the gather (zero rows for ids outside `[0, V)`), the copy of the rows into `raw`, the RMSNorm +and the weight scaling. Outside: producing the ids; allocating `raw` (the Kimi K3 model passes slot 0 of its +attention-residual snapshot bank `[S, N, H]`, so the rows become layer 0's first snapshot); any residual add, +`(1 + weight)` scaling or quantization. `ids`, `table` and `weight` are read only; every element of `raw` is written; +the result is a new tensor. + +## Signature + +```python +def k3_embed_norm( + ids: torch.Tensor, + table: torch.Tensor, + weight: torch.Tensor, + eps: float, + raw: torch.Tensor, +) -> torch.Tensor +``` + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `ids` | `[N]`, `1 <= N <= 64` | int32 or int64 | 1-D (the op copies a strided one) | CUDA | +| `table` | `[V, H]`, `6144 < H <= 16384`, `H % 1024 == 0` | bf16 | contiguous, 16-byte aligned | CUDA, `ids`' device | +| `weight` | `[H]` | bf16 | contiguous, 16-byte aligned | CUDA, `ids`' device | +| `eps` | scalar | Python float | — | — | +| `raw` | `[N, H]`, overwritten | bf16 | contiguous, 16-byte aligned | CUDA, `ids`' device | +| returns | `[N, H]` | bf16 | contiguous, newly allocated | `ids`' device | + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `ids` | `[N]`, `N` = 1, 2, ..., 8 and 16, 24, ..., 64 | int32 (every `N`), int64 (`N` 1, 8, 64) | contiguous | CUDA | +| `table` | `[163840, 7168]` | bf16 | contiguous | CUDA | +| `weight` | `[7168]` | bf16 | contiguous | CUDA | +| `eps` | `1e-6` (every cell) and `1e-5` | float | — | — | +| `raw` | `[N, 7168]`: slot 0 (every cell) or slot 1 of a `[4, N, 7168]` bank | bf16 | contiguous | CUDA | +| returns | `[N, 7168]` | bf16 | contiguous | CUDA | + +The `N` are the token counts of Kimi K3's decode steps: up to 8 tokens, and DSpark verify steps of up to 8 requests of +8 tokens. The table has Kimi K3's shape (a replicated vocabulary of 163840), with rows at log-uniform scales from 1e-4 +to 1e2 and one all-zero row; the norm weight is around 1 with every 101st element negative. The ids are random in +`[0, V)` with 0, `V - 1`, the zero row, the smallest-scale row and repeats; and, at `N` = 8, ids outside `[0, V)` +(-1, `V`, `V + 7`, `-V`, `2^31 - 1`, `-2^31`, and for int64 also `2^32`, `2^32 + 5`, `3 - 2^32`, `2^40`). + +Every cell: `raw` bit-identical to the torch gather (zero rows outside `[0, V)`) and the bank's other slots +untouched; `out` within 8e-3 of each row's `max |ref|` against an fp64 torch RMSNorm, and exactly zero on zero rows; +a rerun bit-identical; every row bit-identical to the same id's row of the 64-id call; int64 ids bit-identical to +int32 ids. Also certified: a CUDA graph of an `N` = 8 call replayed with the ids rewritten in place, bit-identical to +eager calls; and the refusals under *Preconditions*. + +## Metadata consumed + +None besides the arguments and the current CUDA stream (the kernel is launched there), plus two pieces of process +state: + +- A per-process compile cache keyed by `(N, V, H, ids dtype, PDL)`: `N`, `V` and `H` are compiled into the kernel, + `eps` and the tensors' addresses are runtime arguments. The first call for a key compiles the kernel (seconds) and + must be made eagerly: under CUDA-graph capture it raises `RuntimeError` ("must run once per shape outside CUDA-graph + capture first"). The cache is result-neutral. +- `TRTLLM_ENABLE_PDL` (default `"1"`), read on every call: whether the kernel is launched with PDL. It is part of the + cache key and changes scheduling, not results. The test runs with the default. + +## Preconditions + +- `ids` 1-D, int32 or int64, `1 <= N <= 64`, on a CUDA device. +- `table` bf16, 2-D, contiguous, 16-byte-aligned data pointer, `6144 < H <= 16384` and `H % 1024 == 0`. +- `weight` bf16 of shape `(H,)`, contiguous, 16-byte aligned; `raw` bf16 of shape `(N, H)`, contiguous, 16-byte + aligned. +- A call outside these raises `ValueError` ("k3_embed_norm: unsupported call ...") before launching anything, leaving + `raw` untouched. Certified: `N` = 0 and 65, int16 ids, `raw` `[64, H]` for 8 ids, `raw` 2 bytes past a 16-byte + boundary, `weight` `[H / 2]`, and tables with `H` = 6144 and 7680. +- Not checked by the op: `table`, `weight` and `raw` on `ids`' device; `raw` not overlapping `ids`, `table` or + `weight`; a GPU with programmatic dependent launch (SM 9.0 or newer: the kernel uses `griddepcontrol`; the receipts + name the architectures it is certified on); the CuTe DSL (`cutlass`) and `cuda-python` (`cuda.bindings`), which the + call imports, so without them it raises `ImportError`. + +## Notes + +- PDL: with `TRTLLM_ENABLE_PDL` = `"1"` the kernel is launched with programmatic dependent launch. It releases its + dependents at entry and waits for the kernel it follows (`griddepcontrol.wait`) before reading anything, so the ids + may be the predecessor's output. A dependent must itself wait for this grid before reading `raw` or the result, as + every PDL kernel waits for its predecessor before reading its output. +- Grid: `N` CTAs of 128 threads; no CTA reads another's writes. +- The Kimi K3 table is replicated (no vocabulary shard), so the model never passes an id outside `[0, V)`; the zero + rows define the op on every input. Where `weight` is negative, a zero row's output holds `-0.0`. +- `trtllm::k3_embed` (the gather alone, into a new tensor) is a separate op and not this entry. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/k3_embed_norm.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/k3_embed_norm.py new file mode 100644 index 000000000000..927ff24364cd --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/k3_embed_norm.py @@ -0,0 +1,20 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3's decode-step embedding and its first layer's input RMSNorm in one CuTe DSL launch: the rows ``table[ids]`` +written into a caller's buffer, and their RMSNorm.""" + +import torch + +import tensorrt_llm._torch.cute_dsl_kernels.k3_embed.op # noqa: F401 — registers the op + + +def k3_embed_norm( + ids: torch.Tensor, + table: torch.Tensor, + weight: torch.Tensor, + eps: float, + raw: torch.Tensor, +) -> torch.Tensor: + """Write ``table[ids]`` into ``raw`` (zero rows for ids outside ``[0, V)``) and return + ``raw * rsqrt(mean(raw^2) + eps) * weight`` as a new bf16 tensor, in one k3_embed_norm call.""" + return torch.ops.trtllm.k3_embed_norm(ids, table, weight, eps, raw) diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index 2ca039299f9c..951d183ea564 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -135,6 +135,29 @@ l0_b200: - unittest/_torch/modules/kimi_kda/test_kda_mtp_decode_cute_parity.py - unittest/_torch/modules/kimi_k3_attn_res/test_attn_res_op.py - unittest/_torch/modules/kimi_k3_attn_res/test_attn_res_rmsnorm_op.py + # Kimi K3 decode kernels (CuTe DSL): the GEMVs, the head GEMV and the + # embedding, and the cluster-wait and tcgen05-fence lints over them. + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_embed.py + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctm_gemv.py + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctm_gemv_wide.py + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_decode_gemv.py + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_head_gemv.py + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_cluster_waits.py + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py + # modeling_v2 catalog entries with sm_100 receipts (the K3 ones skip on + # other architectures). + - unittest/_torch/modeling_v2/norm/test_modeling_v2_k3_embed_norm.py + - unittest/_torch/modeling_v2/norm/test_modeling_v2_attn_res_fwd.py + - unittest/_torch/modeling_v2/norm/test_modeling_v2_attn_res_rmsnorm_fwd.py + - unittest/_torch/modeling_v2/norm/test_modeling_v2_flashinfer_rmsnorm.py + - unittest/_torch/modeling_v2/activation/test_modeling_v2_k3_situ_mul.py + - unittest/_torch/modeling_v2/gemm/test_modeling_v2_cublas_mm.py + - unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_ctm_gemv.py + - unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_ctm_gemv_swiglu.py + - unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_ctm_gemv_long.py + - unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_ctm_gemv_wide.py + - unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_decode_gemv.py + - unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_head_gemv.py # KDA runtime: host-derived prefill metadata and the bf16 state pool # round-trip. Both are single-device cases that use GPU 0 only. - unittest/_torch/modules/kimi_kda/test_kda_host_metadata.py diff --git a/tests/unittest/_torch/modeling_v2/activation/test_modeling_v2_k3_situ_mul.py b/tests/unittest/_torch/modeling_v2/activation/test_modeling_v2_k3_situ_mul.py new file mode 100644 index 000000000000..41976eaef202 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/activation/test_modeling_v2_k3_situ_mul.py @@ -0,0 +1,121 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the k3_situ_mul catalog entry. + +The Kimi K3 dense MLP's SiTU-and-mul at TP16 (gate_up [M, 4224] -> [M, 2112]) with the checkpoint's and the default +(beta, linear_beta), at every M in 1..8: error against SituAndMul's formula evaluated in fp64, identical bits on a +repeated call, and each M's rows bit-identical to the same rows of the 8-row call. A CUDA-graph replay returns the +eager bits, calls outside the op's preconditions raise ValueError, and the wrapper refuses linear_beta=0.0. +""" + +import pytest +import torch + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.activation.k3_situ_mul import k3_situ_mul + +assert torch.cuda.is_available(), "k3_situ_mul requires a CUDA device" +# The receipts are sm_100 ones; CI's other architectures skip this file. +pytestmark = pytest.mark.skipif( + torch.cuda.get_device_capability() != (10, 0), reason="certified on sm_100 only" +) + +M_ALL = range(1, 9) +TOL = 8e-3 # max |y - ref| / max |ref| against the formula evaluated in fp64 +K = 2112 # the dense MLP's intermediate size per rank at TP16 +SITU = {"k3": (4.0, 25.0), "defaults": (1.0, None)} # (beta, linear_beta) + + +def _randn(rows: int, cols: int, seed: int, scale: float) -> torch.Tensor: + gen = torch.Generator(device="cuda").manual_seed(seed) + return (torch.randn(rows, cols, generator=gen, device="cuda") * scale).bfloat16() + + +def _gate_up() -> torch.Tensor: + return _randn(8, 2 * K, 14, 2.0) + + +def _situ_ref(gu: torch.Tensor, beta: float, linear_beta: float | None) -> torch.Tensor: + """SituAndMul's eager fp32 formula, evaluated in fp64.""" + k = gu.shape[1] // 2 + g, u = gu[:, :k].double(), gu[:, k:].double() + a = beta * torch.tanh(g / beta) * torch.sigmoid(g) + if linear_beta is not None: + u = linear_beta * torch.tanh(u / linear_beta) + return a * u + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def _rel(y: torch.Tensor, ref: torch.Tensor) -> float: + """max |y - ref| / max |ref|, in fp64.""" + return (y.double() - ref).abs().max().item() / ref.abs().max().item() + + +def test_call_site_cells() -> None: + gu8 = _gate_up() + for name, (beta, linear_beta) in SITU.items(): + ref8 = _situ_ref(gu8, beta, linear_beta) + y8 = k3_situ_mul(gu8, beta, linear_beta) + for m in M_ALL: + cell = f"{name} M={m}" + gu = gu8[:m].contiguous() + y = k3_situ_mul(gu, beta, linear_beta) + assert y.shape == (m, K) and y.dtype == torch.bfloat16, cell + rel = _rel(y, ref8[:m]) + assert rel <= TOL, f"{cell}: max |y - ref| / max |ref| = {rel:.2e} > {TOL}" + again = k3_situ_mul(gu, beta, linear_beta) + assert torch.equal(_bits(y), _bits(again)), f"{cell}: a repeated call changed bits" + assert torch.equal(_bits(y), _bits(y8[:m])), f"{cell}: rows differ from the 8-row call" + + +def test_cuda_graph_replay_matches_eager() -> None: + """Captured after an eager call per key and replayed with gu rewritten in place: the eager bits.""" + calls = [] + for name, (beta, linear_beta) in SITU.items(): + for m in (1, 8): + gu = _gate_up()[:m].clone() + k3_situ_mul(gu, beta, linear_beta) # compiles the key outside capture + calls.append((f"{name} M={m}", gu, beta, linear_beta)) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + outs = [k3_situ_mul(gu, beta, linear_beta) for _, gu, beta, linear_beta in calls] + for seed, (_, gu, *_) in enumerate(calls, start=100): + gu.copy_(_randn(gu.shape[0], gu.shape[1], seed, 2.0)) + graph.replay() + for (cell, gu, beta, linear_beta), y in zip(calls, outs): + eager = k3_situ_mul(gu, beta, linear_beta) + assert torch.equal(_bits(y), _bits(eager)), f"{cell}: the replay differs from an eager call" + + +def test_unsupported_calls_raise_value_error() -> None: + """The op's own check refuses these before compiling or launching anything.""" + gu8 = _gate_up() + flat = torch.zeros(8 * 2 * K + 8, dtype=torch.bfloat16, device="cuda") + cases = { + "0 rows": gu8[:0], + "9 rows": torch.cat([gu8, gu8[:1]]), + "width % 16 != 0": _randn(8, 2 * K - 8, 14, 2.0), + "gu 2 bytes past a 16-byte boundary": flat[1 : 1 + 8 * 2 * K].view(8, 2 * K), + "fp16 gu": gu8.half(), + "row-strided gu": torch.cat([gu8, gu8], dim=1)[:, : 2 * K], + "1-D gu": gu8[0], + } + for name, gu in cases.items(): + try: + k3_situ_mul(gu, 4.0, 25.0) + except ValueError: + continue + raise AssertionError(f"{name}: expected ValueError, the op accepted the call") + + +def test_zero_linear_beta_is_refused_before_dispatch() -> None: + """The op would run linear_beta=0.0 as 1.0 with no error; the wrapper refuses it instead.""" + gu = _gate_up()[:1] + for linear_beta in (0.0, -0.0): + try: + k3_situ_mul(gu, 4.0, linear_beta) + except AssertionError: + continue + raise AssertionError(f"linear_beta={linear_beta} should have been refused before dispatch") diff --git a/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_ctm_gemv.py b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_ctm_gemv.py new file mode 100644 index 000000000000..4befd44bf4a4 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_ctm_gemv.py @@ -0,0 +1,129 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the k3_ctm_gemv catalog entry. + +The Kimi K3 TP16 per-rank shapes with their call sites' split and push, at every M in 1..8: error against an fp64 +product of the same bf16 inputs, identical bits on a repeated call, and each M's rows bit-identical to the same rows +of the 8-row call. The schedule-only flags (push, trigger_early) move no bits, a CUDA-graph replay returns the +eager bits, and calls outside the op's preconditions raise ValueError. +""" + +import pytest +import torch + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv import k3_ctm_gemv + +assert torch.cuda.is_available(), "k3_ctm_gemv requires a CUDA device" +# The receipts are sm_100 ones; CI's other architectures skip this file. +pytestmark = pytest.mark.skipif( + torch.cuda.get_device_capability() != (10, 0), reason="certified on sm_100 only" +) + +M_ALL = range(1, 9) +TOL = 8e-3 # max |y - ref| / max |ref| against an fp64 product of the same bf16 inputs + +# (N, K, split, push): the MLA o_proj at splits 1 and 2, the drafter o_proj (tuned and synthetic drafter). +CELLS = { + "mla_o_proj_s1": (7168, 768, 1, False), + "mla_o_proj_s2": (7168, 768, 2, False), + "drafter_o_proj": (7168, 384, 1, True), + "drafter_o_proj_synthetic": (7168, 256, 1, True), +} + + +def _randn(rows: int, cols: int, seed: int, scale: float) -> torch.Tensor: + gen = torch.Generator(device="cuda").manual_seed(seed) + return (torch.randn(rows, cols, generator=gen, device="cuda") * scale).bfloat16() + + +def _weight(n: int, k: int) -> torch.Tensor: + return _randn(n, k, 1, 0.03) + + +def _rows(k: int) -> torch.Tensor: + return _randn(8, k, 2, 1.0) + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def _rel(y: torch.Tensor, ref: torch.Tensor) -> float: + """max |y - ref| / max |ref|, in fp64.""" + return (y.double() - ref).abs().max().item() / ref.abs().max().item() + + +def test_call_site_cells() -> None: + for name, (n, k, split, push) in CELLS.items(): + w = _weight(n, k) + x8 = _rows(k) + ref8 = x8.double() @ w.double().t() + y8 = k3_ctm_gemv(x8, w, True, split, push) + for m in M_ALL: + cell = f"{name} M={m}" + x = x8[:m].contiguous() + y = k3_ctm_gemv(x, w, True, split, push) + assert y.shape == (m, n) and y.dtype == torch.bfloat16, cell + rel = _rel(y, ref8[:m]) + assert rel <= TOL, f"{cell}: max |y - ref| / max |ref| = {rel:.2e} > {TOL}" + again = k3_ctm_gemv(x, w, True, split, push) + assert torch.equal(_bits(y), _bits(again)), f"{cell}: a repeated call changed bits" + assert torch.equal(_bits(y), _bits(y8[:m])), f"{cell}: rows differ from the 8-row call" + + +def test_schedule_flags_move_no_bits() -> None: + """At the MLA o_proj with split 2, push=True and trigger_early=False return the call site's bits.""" + n, k, split = 7168, 768, 2 + w = _weight(n, k) + x8 = _rows(k) + for m in M_ALL: + x = x8[:m].contiguous() + base = _bits(k3_ctm_gemv(x, w, True, split, False)) + for trigger_early, push in ((True, True), (False, False)): + y = k3_ctm_gemv(x, w, trigger_early, split, push) + assert torch.equal(_bits(y), base), ( + f"M={m} trigger_early={trigger_early} push={push}: bits differ from the call site's" + ) + + +def test_cuda_graph_replay_matches_eager() -> None: + """Captured after an eager call per key and replayed with x rewritten in place: the eager bits.""" + calls = [] + for name in ("mla_o_proj_s1", "mla_o_proj_s2"): + n, k, split, push = CELLS[name] + w = _weight(n, k) + for m in (1, 8): + x = _rows(k)[:m].clone() + k3_ctm_gemv(x, w, True, split, push) # compiles the key outside capture + calls.append((f"{name} M={m}", x, w, split, push)) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + outs = [k3_ctm_gemv(x, w, True, split, push) for _, x, w, split, push in calls] + for seed, (_, x, *_) in enumerate(calls, start=100): + x.copy_(_randn(x.shape[0], x.shape[1], seed, 1.0)) + graph.replay() + for (cell, x, w, split, push), y in zip(calls, outs): + eager = k3_ctm_gemv(x, w, True, split, push) + assert torch.equal(_bits(y), _bits(eager)), f"{cell}: the replay differs from an eager call" + + +def test_unsupported_calls_raise_value_error() -> None: + """The op's own check refuses these before compiling or launching anything.""" + w = _weight(7168, 768) + x8 = _rows(768) + cases = { + "0 rows": (x8[:0], w, 1), + "9 rows": (torch.cat([x8, x8[:1]]), w, 1), + "fp16 x": (x8.half(), w, 1), + "row-strided x": (torch.cat([x8, x8], dim=1)[:, :768], w, 1), + "N % 128 != 0": (x8, w[:7104], 1), + "K % 128 != 0": (x8[:, :704].contiguous(), w[:, :704].contiguous(), 1), + "split 3": (x8, w, 3), + "7 k-tiles on one CTA": (_rows(896), _weight(7168, 896), 1), + } + for name, (x, weight, split) in cases.items(): + try: + k3_ctm_gemv(x, weight, True, split, False) + except ValueError: + continue + raise AssertionError(f"{name}: expected ValueError, the op accepted the call") diff --git a/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_ctm_gemv_long.py b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_ctm_gemv_long.py new file mode 100644 index 000000000000..b842a80f0c5a --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_ctm_gemv_long.py @@ -0,0 +1,154 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the k3_ctm_gemv_long catalog entry. + +The Kimi K3 TP16 per-rank shapes with their call sites' sig_col0, split, ring and push, at every M in 1..8: error of +the plain (sig_col0=-1) output against an fp64 product of the same bf16 inputs, the sigmoid columns torch.sigmoid of +that output bit for bit, identical bits on a repeated call, and each M's rows bit-identical to the same rows of the +8-row call. The schedule-only flags (ring, push, trigger_early) move no bits, a CUDA-graph replay returns the eager +bits, and calls outside the op's preconditions raise ValueError. +""" + +import pytest +import torch + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv_long import ( + k3_ctm_gemv_long, +) + +assert torch.cuda.is_available(), "k3_ctm_gemv_long requires a CUDA device" +# The receipts are sm_100 ones; CI's other architectures skip this file. +pytestmark = pytest.mark.skipif( + torch.cuda.get_device_capability() != (10, 0), reason="certified on sm_100 only" +) + +M_ALL = range(1, 9) +TOL = 8e-3 # max |y - ref| / max |ref| against an fp64 product of the same bf16 inputs + +# (N, K, sig_col0, split, ring, push), as each call site passes them. mla_qkv_a_gate is the fused +# [W_a; W_g] projection, its gate columns returned as bf16(sigmoid). +CELLS = { + "mla_qkv_a_gate": (2880, 7168, 2112, 6, 6, True), + "dense_gate_up": (4224, 7168, -1, 4, 5, False), + "dense_down": (7168, 2112, -1, 2, 6, False), # K 2112 ends in a half k-tile + "drafter_qkv": (512, 7168, -1, 8, 6, True), + "drafter_gate_up": (1792, 7168, -1, 8, 6, True), + "drafter_gate_up_synthetic": (1536, 7168, -1, 8, 6, True), + "kda_qkvg": (3208, 7168, -1, 5, 6, True), # 25 whole 128-row blocks and an 8-row last one +} + +# One flag changed from a call site's values; each must return the call site's bits. +FLAG_CHANGES = { + "dense_down": ({"push": True}, {"ring": 3}), + "dense_gate_up": ({"push": True}, {"trigger_early": False}), +} + + +def _randn(rows: int, cols: int, seed: int, scale: float) -> torch.Tensor: + gen = torch.Generator(device="cuda").manual_seed(seed) + return (torch.randn(rows, cols, generator=gen, device="cuda") * scale).bfloat16() + + +def _weight(n: int, k: int) -> torch.Tensor: + return _randn(n, k, 9, 0.02) + + +def _rows(k: int) -> torch.Tensor: + return _randn(8, k, 10, 1.0) + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def _rel(y: torch.Tensor, ref: torch.Tensor) -> float: + """max |y - ref| / max |ref|, in fp64.""" + return (y.double() - ref).abs().max().item() / ref.abs().max().item() + + +def test_call_site_cells() -> None: + for name, (n, k, sig, split, ring, push) in CELLS.items(): + w = _weight(n, k) + x8 = _rows(k) + ref8 = x8.double() @ w.double().t() + y8 = k3_ctm_gemv_long(x8, w, sig, split, ring, True, push) + for m in M_ALL: + cell = f"{name} M={m}" + x = x8[:m].contiguous() + y = k3_ctm_gemv_long(x, w, sig, split, ring, True, push) + assert y.shape == (m, n) and y.dtype == torch.bfloat16, cell + again = k3_ctm_gemv_long(x, w, sig, split, ring, True, push) + assert torch.equal(_bits(y), _bits(again)), f"{cell}: a repeated call changed bits" + assert torch.equal(_bits(y), _bits(y8[:m])), f"{cell}: rows differ from the 8-row call" + plain = k3_ctm_gemv_long(x, w, -1, split, ring, True, push) if sig >= 0 else y + rel = _rel(plain, ref8[:m]) + assert rel <= TOL, f"{cell}: max |y - ref| / max |ref| = {rel:.2e} > {TOL}" + if sig >= 0: + assert torch.equal(_bits(y[:, :sig]), _bits(plain[:, :sig])), ( + f"{cell}: columns before sig_col0 differ from the sig_col0=-1 call" + ) + assert torch.equal(_bits(y[:, sig:]), _bits(plain[:, sig:].sigmoid())), ( + f"{cell}: columns from sig_col0 on are not torch.sigmoid of the sig_col0=-1 call" + ) + + +def test_schedule_flags_move_no_bits() -> None: + """ring, push and trigger_early only schedule the call: one changed, the call site's bits come back.""" + for name, changes in FLAG_CHANGES.items(): + n, k, sig, split, ring, push = CELLS[name] + w = _weight(n, k) + x8 = _rows(k) + site = {"split": split, "ring": ring, "trigger_early": True, "push": push} + for m in M_ALL: + x = x8[:m].contiguous() + base = _bits(k3_ctm_gemv_long(x, w, sig, **site)) + for change in changes: + y = k3_ctm_gemv_long(x, w, sig, **{**site, **change}) + assert torch.equal(_bits(y), base), ( + f"{name} M={m} {change}: bits differ from the call site's" + ) + + +def test_cuda_graph_replay_matches_eager() -> None: + """Captured after an eager call per key and replayed with x rewritten in place: the eager bits.""" + calls = [] + for name in ("mla_qkv_a_gate", "dense_down"): + n, k, sig, split, ring, push = CELLS[name] + w = _weight(n, k) + for m in (1, 8): + x = _rows(k)[:m].clone() + k3_ctm_gemv_long(x, w, sig, split, ring, True, push) # compiles the key outside capture + calls.append((f"{name} M={m}", x, w, (sig, split, ring, True, push))) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + outs = [k3_ctm_gemv_long(x, w, *flags) for _, x, w, flags in calls] + for seed, (_, x, _, _) in enumerate(calls, start=100): + x.copy_(_randn(x.shape[0], x.shape[1], seed, 1.0)) + graph.replay() + for (cell, x, w, flags), y in zip(calls, outs): + eager = k3_ctm_gemv_long(x, w, *flags) + assert torch.equal(_bits(y), _bits(eager)), f"{cell}: the replay differs from an eager call" + + +def test_unsupported_calls_raise_value_error() -> None: + """The op's own check refuses these before compiling or launching anything.""" + w = _weight(7168, 2112) + x8 = _rows(2112) + cases = { + "0 rows": (x8[:0], w, 2, 6), + "9 rows": (torch.cat([x8, x8[:1]]), w, 2, 6), + "fp16 x": (x8.half(), w, 2, 6), + "row-strided x": (torch.cat([x8, x8], dim=1)[:, :2112], w, 2, 6), + "K % 64 != 0": (x8[:, :2080].contiguous(), w[:, :2080].contiguous(), 2, 6), + "split 3": (x8, w, 3, 5), + "split 9": (x8, w, 9, 1), + "ring 0": (x8, w, 2, 0), + "ring past the rank's k-tiles": (x8, w, 8, 3), + "ring past shared memory": (x8, w, 2, 7), + } + for name, (x, weight, split, ring) in cases.items(): + try: + k3_ctm_gemv_long(x, weight, -1, split, ring, True, False) + except ValueError: + continue + raise AssertionError(f"{name}: expected ValueError, the op accepted the call") diff --git a/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_ctm_gemv_swiglu.py b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_ctm_gemv_swiglu.py new file mode 100644 index 000000000000..e61be03e4683 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_ctm_gemv_swiglu.py @@ -0,0 +1,135 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the k3_ctm_gemv_swiglu catalog entry. + +The Kimi K3 TP16 per-rank drafter down projection with its call site's split and push, at every M in 1..8: error +against an fp64 product of torch's bf16 silu_and_mul activation, identical bits on a repeated call, and each M's rows +bit-identical to the same rows of the 8-row call. The schedule-only flags (push, trigger_early) move no bits, a +CUDA-graph replay returns the eager bits, and calls outside the op's preconditions raise ValueError. +""" + +import pytest +import torch +import torch.nn.functional as F + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv_swiglu import ( + k3_ctm_gemv_swiglu, +) + +assert torch.cuda.is_available(), "k3_ctm_gemv_swiglu requires a CUDA device" +# The receipts are sm_100 ones; CI's other architectures skip this file. +pytestmark = pytest.mark.skipif( + torch.cuda.get_device_capability() != (10, 0), reason="certified on sm_100 only" +) + +M_ALL = range(1, 9) +TOL = 8e-3 # max |y - ref| / max |ref| against an fp64 product of the bf16 activation + +# (N, K, split, push): the drafter down projection (tuned and synthetic drafter); K 896 is 7 k-tiles. +CELLS = { + "drafter_down": (7168, 896, 2, True), + "drafter_down_synthetic": (7168, 768, 2, True), +} + + +def _randn(rows: int, cols: int, seed: int, scale: float) -> torch.Tensor: + gen = torch.Generator(device="cuda").manual_seed(seed) + return (torch.randn(rows, cols, generator=gen, device="cuda") * scale).bfloat16() + + +def _weight(n: int, k: int) -> torch.Tensor: + return _randn(n, k, 7, 0.02) + + +def _gate_up(k: int) -> torch.Tensor: + return _randn(8, 2 * k, 8, 2.0) + + +def _silu_and_mul(gu: torch.Tensor) -> torch.Tensor: + """silu_and_mul's fp32 arithmetic and one bf16 rounding (gate columns first).""" + k = gu.shape[1] // 2 + return (F.silu(gu[:, :k].float()) * gu[:, k:].float()).bfloat16() + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def _rel(y: torch.Tensor, ref: torch.Tensor) -> float: + """max |y - ref| / max |ref|, in fp64.""" + return (y.double() - ref).abs().max().item() / ref.abs().max().item() + + +def test_call_site_cells() -> None: + for name, (n, k, split, push) in CELLS.items(): + w = _weight(n, k) + gu8 = _gate_up(k) + ref8 = _silu_and_mul(gu8).double() @ w.double().t() + y8 = k3_ctm_gemv_swiglu(gu8, w, True, split, push) + for m in M_ALL: + cell = f"{name} M={m}" + gu = gu8[:m].contiguous() + y = k3_ctm_gemv_swiglu(gu, w, True, split, push) + assert y.shape == (m, n) and y.dtype == torch.bfloat16, cell + rel = _rel(y, ref8[:m]) + assert rel <= TOL, f"{cell}: max |y - ref| / max |ref| = {rel:.2e} > {TOL}" + again = k3_ctm_gemv_swiglu(gu, w, True, split, push) + assert torch.equal(_bits(y), _bits(again)), f"{cell}: a repeated call changed bits" + assert torch.equal(_bits(y), _bits(y8[:m])), f"{cell}: rows differ from the 8-row call" + + +def test_schedule_flags_move_no_bits() -> None: + """At the drafter down, push=False and trigger_early=False return the call site's bits.""" + n, k, split, push = CELLS["drafter_down"] + w = _weight(n, k) + gu8 = _gate_up(k) + for m in M_ALL: + gu = gu8[:m].contiguous() + base = _bits(k3_ctm_gemv_swiglu(gu, w, True, split, push)) + for trigger_early, push_flag in ((True, not push), (False, push)): + y = k3_ctm_gemv_swiglu(gu, w, trigger_early, split, push_flag) + assert torch.equal(_bits(y), base), ( + f"M={m} trigger_early={trigger_early} push={push_flag}: bits differ from the call site's" + ) + + +def test_cuda_graph_replay_matches_eager() -> None: + """Captured after an eager call per key and replayed with gu rewritten in place: the eager bits.""" + n, k, split, push = CELLS["drafter_down"] + w = _weight(n, k) + calls = [] + for m in (1, 8): + gu = _gate_up(k)[:m].clone() + k3_ctm_gemv_swiglu(gu, w, True, split, push) # compiles the key outside capture + calls.append((m, gu)) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + outs = [k3_ctm_gemv_swiglu(gu, w, True, split, push) for _, gu in calls] + for seed, (m, gu) in enumerate(calls, start=100): + gu.copy_(_randn(m, 2 * k, seed, 2.0)) + graph.replay() + for (m, gu), y in zip(calls, outs): + eager = k3_ctm_gemv_swiglu(gu, w, True, split, push) + assert torch.equal(_bits(y), _bits(eager)), f"M={m}: the replay differs from an eager call" + + +def test_unsupported_calls_raise_value_error() -> None: + """The op's own check refuses these before compiling or launching anything.""" + w = _weight(7168, 896) + gu8 = _gate_up(896) + cases = { + "0 rows": (gu8[:0], w, 2), + "9 rows": (torch.cat([gu8, gu8[:1]]), w, 2), + "fp16 gu": (gu8.half(), w, 2), + "row-strided gu": (torch.cat([gu8, gu8], dim=1)[:, : 2 * 896], w, 2), + "gu width != 2 K": (_randn(8, 1664, 8, 2.0), w, 2), + "N % 128 != 0": (gu8, w[:7104], 2), + "split 3": (gu8, w, 3), + "7 k-tiles on one CTA": (gu8, w, 1), + } + for name, (gu, weight, split) in cases.items(): + try: + k3_ctm_gemv_swiglu(gu, weight, True, split, True) + except ValueError: + continue + raise AssertionError(f"{name}: expected ValueError, the op accepted the call") diff --git a/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_ctm_gemv_wide.py b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_ctm_gemv_wide.py new file mode 100644 index 000000000000..af81feb06f7b --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_ctm_gemv_wide.py @@ -0,0 +1,184 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the k3_ctm_gemv_wide catalog entry. + +The Kimi K3 TP16 per-rank shapes of the projections of a wide decode step, at every M in 1..64: error of the plain +(sig_col0=-1) output against an fp64 product of the same bf16 inputs, the sigmoid columns torch.sigmoid of that +output bit for bit, identical bits on a repeated call, and each M's rows bit-identical to the same rows of the +64-row call. Where the device's split equals a k3_ctm_gemv_long call site's, every token's row carries that op's +bits; a CUDA-graph replay returns the eager bits; calls outside the op's preconditions raise ValueError. +""" + +import pytest +import torch + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv_long import ( + k3_ctm_gemv_long, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv_wide import ( + k3_ctm_gemv_wide, +) + +assert torch.cuda.is_available(), "k3_ctm_gemv_wide requires a CUDA device" +# The receipts are sm_100 ones; CI's other architectures skip this file. +pytestmark = pytest.mark.skipif( + torch.cuda.get_device_capability() != (10, 0), reason="certified on sm_100 only" +) + +M_ALL = range(1, 65) +TOL = 8e-3 # bf16 output: max |y - ref| / max |ref| against an fp64 product of the same bf16 inputs +TOL_FP32 = 1e-4 # fp32 output + +# (N, K, sig_col0, out_fp32): the per-rank TP16 shapes of the projections of a wide decode step. kda_qkvg is 25 +# whole 128-row blocks and an 8-row last one; mla_qkv_a_gate is [W_a; W_g] with its gate columns as bf16(sigmoid); +# moe_head is [latent down slice; router rows] in fp32 (router logits); moe_tail is 5 k-tiles; dense_down's K 2112 +# ends in a half k-tile. +CELLS = { + "kda_qkvg": (3208, 7168, -1, False), + "mla_qkv_a_gate": (2880, 7168, 2112, False), + "o_proj": (7168, 768, -1, False), + "moe_head": (280, 7168, -1, True), + "moe_tail": (7168, 640, -1, False), + "shared_gate_up": (768, 7168, -1, False), + "dense_gate_up": (4224, 7168, -1, False), + "dense_down": (7168, 2112, -1, False), + "drafter_qkv": (512, 7168, -1, False), + "drafter_gate_up": (1792, 7168, -1, False), + "drafter_o_proj": (7168, 384, -1, False), +} +# The k3_ctm_gemv_long call (split, ring, push) of the shapes that run it at up to 8 tokens. +LONG_SITES = { + "mla_qkv_a_gate": (6, 6, True), + "dense_gate_up": (4, 5, False), + "dense_down": (2, 6, False), + "drafter_qkv": (8, 6, True), + "drafter_gate_up": (8, 6, True), +} + + +def _randn(rows: int, cols: int, seed: int, scale: float) -> torch.Tensor: + gen = torch.Generator(device="cuda").manual_seed(seed) + return (torch.randn(rows, cols, generator=gen, device="cuda") * scale).bfloat16() + + +def _weight(n: int, k: int) -> torch.Tensor: + return _randn(n, k, 1, 0.02) + + +def _rows(k: int) -> torch.Tensor: + return _randn(64, k, 2, 1.0) + + +def _bits(t: torch.Tensor) -> torch.Tensor: + t = t.contiguous() + return t.view(torch.int32) if t.dtype == torch.float32 else t.view(torch.int16) + + +def _rel(y: torch.Tensor, ref: torch.Tensor) -> float: + """max |y - ref| / max |ref|, in fp64.""" + return (y.double() - ref).abs().max().item() / ref.abs().max().item() + + +def test_projection_cells() -> None: + for name, (n, k, sig, fp32) in CELLS.items(): + w = _weight(n, k) + x64 = _rows(k) + ref64 = x64.double() @ w.double().t() + y64 = k3_ctm_gemv_wide(x64, w, sig, fp32) + dtype = torch.float32 if fp32 else torch.bfloat16 + tol = TOL_FP32 if fp32 else TOL + for m in M_ALL: + cell = f"{name} M={m}" + x = x64[:m].contiguous() + y = k3_ctm_gemv_wide(x, w, sig, fp32) + assert y.shape == (m, n) and y.dtype == dtype, cell + again = k3_ctm_gemv_wide(x, w, sig, fp32) + assert torch.equal(_bits(y), _bits(again)), f"{cell}: a repeated call changed bits" + assert torch.equal(_bits(y), _bits(y64[:m])), ( + f"{cell}: rows differ from the 64-row call" + ) + plain = k3_ctm_gemv_wide(x, w, -1, fp32) if sig >= 0 else y + rel = _rel(plain, ref64[:m]) + assert rel <= tol, f"{cell}: max |y - ref| / max |ref| = {rel:.2e} > {tol}" + if sig >= 0: + assert torch.equal(_bits(y[:, :sig]), _bits(plain[:, :sig])), ( + f"{cell}: columns before sig_col0 differ from the sig_col0=-1 call" + ) + assert torch.equal(_bits(y[:, sig:]), _bits(plain[:, sig:].sigmoid())), ( + f"{cell}: columns from sig_col0 on are not torch.sigmoid of the sig_col0=-1 call" + ) + + +def test_rows_match_long_at_the_same_split() -> None: + """A token's row is k3_ctm_gemv_long's for it (calls of up to 8 tokens) where the splits agree.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv import op as ctm_op + from tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv.k3_ctm_gemv_kernel import wide_tile + + sms = torch.cuda.get_device_properties(torch.cuda.current_device()).multi_processor_count + compared = 0 + for name, (split, ring, push) in LONG_SITES.items(): + n, k, sig, _ = CELLS[name] + w = _weight(n, k) + x64 = _rows(k) + groups = torch.cat( + [ + k3_ctm_gemv_long(x64[r : r + 8], w, sig, split, ring, True, push) + for r in range(0, 64, 8) + ] + ) + for m in [*range(1, 9), *range(16, 65, 8)]: + if ctm_op.wide_config(n, k, wide_tile(m), sms)[0] != split: + continue + x = x64[:m].contiguous() + y = k3_ctm_gemv_wide(x, w, sig) + want = k3_ctm_gemv_long(x, w, sig, split, ring, True, push) if m <= 8 else groups[:m] + assert torch.equal(_bits(y), _bits(want)), ( + f"{name} M={m}: rows differ from k3_ctm_gemv_long's at split {split}" + ) + compared += 1 + assert compared > 0, f"no call site's split is this device's ({sms} SMs); nothing was compared" + + +def test_cuda_graph_replay_matches_eager() -> None: + """Captured after an eager call per key and replayed with x rewritten in place: the eager bits.""" + calls = [] + for name in ("mla_qkv_a_gate", "moe_head"): + n, k, sig, fp32 = CELLS[name] + w = _weight(n, k) + for m in (1, 64): + x = _rows(k)[:m].clone() + k3_ctm_gemv_wide(x, w, sig, fp32) # compiles the key outside capture + calls.append((f"{name} M={m}", x, w, sig, fp32)) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + outs = [k3_ctm_gemv_wide(x, w, sig, fp32) for _, x, w, sig, fp32 in calls] + for seed, (_, x, *_) in enumerate(calls, start=100): + x.copy_(_randn(x.shape[0], x.shape[1], seed, 1.0)) + graph.replay() + for (cell, x, w, sig, fp32), y in zip(calls, outs): + eager = k3_ctm_gemv_wide(x, w, sig, fp32) + assert torch.equal(_bits(y), _bits(eager)), f"{cell}: the replay differs from an eager call" + + +def test_unsupported_calls_raise_value_error() -> None: + """The op's own check refuses these before compiling or launching anything.""" + n, k = 3208, 7168 + w = _weight(n, k) + x8 = _rows(k)[:8] + flat = torch.zeros(8 * k + 8, dtype=torch.bfloat16, device="cuda") + cases = { + "0 rows": (x8[:0], w, -1, False), + "65 rows": (torch.cat([_rows(k), x8[:1]]), w, -1, False), + "x 2 bytes past a 16-byte boundary": (flat[1 : 1 + 8 * k].view(8, k), w, -1, False), + "sig_col0 = N": (x8, w, n, False), + "sigmoid columns with fp32 output": (x8, w, 100, True), + "N % 8 != 0": (x8, w[:3204], -1, False), + "K % 64 != 0": (x8[:, :7136].contiguous(), w[:, :7136].contiguous(), -1, False), + "fp16 x": (x8.half(), w, -1, False), + } + for name, (x, weight, sig, fp32) in cases.items(): + try: + k3_ctm_gemv_wide(x, weight, sig, fp32) + except ValueError: + continue + raise AssertionError(f"{name}: expected ValueError, the op accepted the call") diff --git a/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_decode_gemv.py b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_decode_gemv.py new file mode 100644 index 000000000000..b21f4c21c4ce --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_decode_gemv.py @@ -0,0 +1,123 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the k3_decode_gemv catalog entry. + +Both kernels at Kimi K3's TP16 per-rank shapes, at every M in 1..8, against an fp64 torch product (error at most 8e-3 +of max |ref|), with run-to-run bits, each M's rows equal to the same rows of the 8-row call, trigger_early False equal +to True, a CUDA graph replayed with its input rewritten, and the refused calls. +""" + +import functools + +import pytest +import torch + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_decode_gemv import k3_decode_gemv + +assert torch.cuda.is_available(), "k3_decode_gemv requires a CUDA device" +# The receipts are sm_100 ones; CI's other architectures skip this file. +pytestmark = pytest.mark.skipif( + torch.cuda.get_device_capability() != (10, 0), reason="certified on sm_100 only" +) + +TOL = 8e-3 # max |y - ref| / max |ref|, ref an fp64 product +M_ALL = range(1, 9) +# (N, K) per rank at TP16. The KDA o_proj: short K, one CTA per 128-row weight tile with the whole K resident. +SHORT_K = (7168, 768) +# The KDA input projection: split K over clusters of 4 CTAs; N is not a multiple of 128 (the last tile has 8 rows). +SPLIT_K = (3208, 7168) + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def _same(a: torch.Tensor, b: torch.Tensor) -> bool: + return a.shape == b.shape and torch.equal(_bits(a), _bits(b)) + + +def _rel(y: torch.Tensor, ref: torch.Tensor) -> float: + return (y.double() - ref).abs().max().item() / ref.abs().max().item() + + +@functools.lru_cache(maxsize=None) +def _inputs(n: int, k: int) -> tuple[torch.Tensor, torch.Tensor]: + """The weight [N, K] and 8 activation rows [8, K].""" + gen = torch.Generator(device="cuda").manual_seed(n * 100003 + k) + weight = (torch.randn(n, k, generator=gen, device="cuda") * 0.03).bfloat16() + x8 = torch.randn(8, k, generator=gen, device="cuda").bfloat16() + # The kernels read the weight before their PDL grid-dependency wait: it is written well before any call. + torch.cuda.synchronize() + return weight, x8 + + +def _check_cells(shape: tuple[int, int]) -> None: + n, k = shape + weight, x8 = _inputs(n, k) + first = {} + for trigger_early in (True, False): + y8 = k3_decode_gemv(x8, weight, trigger_early) + for m in M_ALL: + x = x8[:m] + y = k3_decode_gemv(x, weight, trigger_early) + again = k3_decode_gemv(x, weight, trigger_early) + ref = x.double() @ weight.double().t() + assert y.shape == (m, n) and y.dtype == torch.bfloat16 and y.device == x.device + assert y.is_contiguous() + err = _rel(y, ref) + assert err <= TOL, f"{shape} M={m} trigger_early={trigger_early}: rel {err:.3e}" + assert _same(y, again), f"{shape} M={m}: a rerun differs" + assert _same(y, y8[:m]), f"{shape} M={m}: rows differ from the 8-row call" + if trigger_early: + first[m] = y + else: + assert _same(y, first[m]), f"{shape} M={m}: trigger_early changes the result" + + +def test_short_k_cells() -> None: + """[7168, 768] (one CTA per 128-row tile, the whole K resident) at every M, trigger_early True and False.""" + _check_cells(SHORT_K) + + +def test_split_k_cells() -> None: + """[3208, 7168] (clusters of 4 CTAs over K) at every M, trigger_early True and False.""" + _check_cells(SPLIT_K) + + +def test_graph_replay() -> None: + """M 1 and M 8 calls on each weight in one CUDA graph, captured after eager calls compiled them and replayed with + the activation rewritten in place: every replay bit-identical to eager calls on the same rows.""" + gen = torch.Generator(device="cuda").manual_seed(5) + weights, x_bufs = [], [] + for n, k in (SHORT_K, SPLIT_K): + weight, x8 = _inputs(n, k) + weights.append(weight) + x_bufs.append(x8.clone()) + for m in (1, 8): + k3_decode_gemv(x_bufs[-1][:m], weight) # compiles outside the capture + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + out = [{m: k3_decode_gemv(x[:m], w) for m in (1, 8)} for w, x in zip(weights, x_bufs)] + for _ in range(3): + for x in x_bufs: + x.copy_(torch.randn(x.shape, generator=gen, device="cuda").bfloat16()) + graph.replay() + for i, (w, x) in enumerate(zip(weights, x_bufs)): + for m in (1, 8): + assert _same(out[i][m], k3_decode_gemv(x[:m], w)), f"replay of weight {i}, M={m}" + + +def test_refused_calls() -> None: + """M 0 and 9, and weights neither kernel takes, raise ValueError before launching anything.""" + weight, _ = _inputs(*SHORT_K) + for m in (0, 9): + x = torch.zeros(m, SHORT_K[1], dtype=torch.bfloat16, device="cuda") + with pytest.raises(ValueError, match="unsupported call"): + k3_decode_gemv(x, weight) + # [128, 1152]: 9 k-tiles, more than short K holds and not a multiple of 4 for split K. + # [200, 768]: short K needs whole 128-row tiles, split K more than 6 k-tiles. + for n, k in ((128, 1152), (200, 768)): + x = torch.zeros(1, k, dtype=torch.bfloat16, device="cuda") + w = torch.zeros(n, k, dtype=torch.bfloat16, device="cuda") + with pytest.raises(ValueError, match="unsupported call"): + k3_decode_gemv(x, w) diff --git a/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_head_gemv.py b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_head_gemv.py new file mode 100644 index 000000000000..e3e82157fdfc --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_head_gemv.py @@ -0,0 +1,183 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the k3_head_gemv catalog entry. + +Kimi K3's TP16 per-rank LM-head vocabulary shard [10240, 7168] on the stream-K schedule, at every M in 1..8, against an +fp64 torch product (error at most 8e-3 of max |ref|), with run-to-run bits and each M's rows equal to the same rows of +the 8-row call. The caller-owned workspace is driven through real call sequences: M dipping and growing back on one +workspace, two workspaces interleaved, two weights on one workspace, and a CUDA graph replayed between eager calls on +the same workspace, every call against the same call on a fresh workspace and every flag / count word back at zero +after. Negative controls: a workspace of another shape is refused, and so is creating one under capture. +""" + +import functools + +import pytest +import torch + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_head_gemv import ( + K3HeadGemvWorkspace, + k3_head_gemv, +) + +assert torch.cuda.is_available(), "k3_head_gemv requires a CUDA device" +# The receipts are sm_100 ones; CI's other architectures skip this file. +pytestmark = pytest.mark.skipif( + torch.cuda.get_device_capability() != (10, 0), reason="certified on sm_100 only" +) + +TOL = 8e-3 # max |y - ref| / max |ref|, ref an fp64 product +M_ALL = range(1, 9) +VOCAB_SHARD, HIDDEN = 10240, 7168 # Kimi K3's 163840-row LM head over 16 ranks +TILES = VOCAB_SHARD // 128 + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def _same(a: torch.Tensor, b: torch.Tensor) -> bool: + return a.shape == b.shape and torch.equal(_bits(a), _bits(b)) + + +def _rel(y: torch.Tensor, ref: torch.Tensor) -> float: + return (y.double() - ref).abs().max().item() / ref.abs().max().item() + + +@functools.lru_cache(maxsize=None) +def _inputs() -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Two shard weights (e.g. the target's LM head and DSpark's draft head) and 8 activation rows.""" + gen = torch.Generator(device="cuda").manual_seed(20261001) + weight = (torch.randn(VOCAB_SHARD, HIDDEN, generator=gen, device="cuda") * 0.02).bfloat16() + weight2 = (torch.randn(VOCAB_SHARD, HIDDEN, generator=gen, device="cuda") * 0.02).bfloat16() + x8 = torch.randn(8, HIDDEN, generator=gen, device="cuda").bfloat16() + # The kernel reads the weight before its PDL grid-dependency wait: it is written well before any call. + torch.cuda.synchronize() + return weight, weight2, x8 + + +def _workspace() -> K3HeadGemvWorkspace: + return K3HeadGemvWorkspace.create(VOCAB_SHARD, HIDDEN, torch.device("cuda")) + + +def _at_rest(ws: K3HeadGemvWorkspace) -> bool: + """Whether every flag / count word of the workspace is back at zero.""" + return not bool(ws.flags.any()) and not bool(ws.claim.any()) + + +def test_certified_cells() -> None: + """Every M on one workspace: shape, dtype, error, rerun and 8-row-call bits, the words at zero after each call; + keep_tiles 40 and 80 bit-identical to 0.""" + weight, _, x8 = _inputs() + ws = _workspace() + assert ws.schedule == "streamk" and ws.chunk_tiles == 0 + assert ws.partials.dtype == torch.float32 and ws.flags.dtype == torch.int32 + assert ws.claim.dtype == torch.int32 and ws.claim.numel() == 1 + # One flag word and one [128 x 8] fp32 partial slot per (tile, piece). + assert ws.flags.numel() % TILES == 0 and ws.partials.numel() == ws.flags.numel() * 128 * 8 + assert _at_rest(ws) + y8 = k3_head_gemv(x8, weight, ws) + for m in M_ALL: + x = x8[:m] + y = k3_head_gemv(x, weight, ws) + again = k3_head_gemv(x, weight, ws) + ref = x.double() @ weight.double().t() + assert y.shape == (m, VOCAB_SHARD) and y.dtype == torch.bfloat16 and y.device == x.device + assert y.is_contiguous() + err = _rel(y, ref) + assert err <= TOL, f"M={m}: rel {err:.3e}" + assert _same(y, again), f"M={m}: a rerun differs" + assert _same(y, y8[:m]), f"M={m}: rows differ from the 8-row call" + assert _at_rest(ws), f"M={m}: words left raised" + for m in (1, 8): + for keep in (TILES // 2, TILES): + y = k3_head_gemv(x8[:m], weight, ws, keep_tiles=keep) + assert _same(y, y8[:m]), f"M={m} keep_tiles={keep}: result differs" + assert _at_rest(ws) + + +def test_m_dipping_and_growing_on_one_workspace() -> None: + """M 8, 8, 2, 7, 8, 1, 1, 8, 3, 8 back to back on one stream and one workspace: every call bit-identical to the + same call on a fresh workspace, the words at zero after.""" + weight, _, x8 = _inputs() + seq = (8, 8, 2, 7, 8, 1, 1, 8, 3, 8) + want = {m: k3_head_gemv(x8[:m], weight, _workspace()) for m in set(seq)} + ws = _workspace() + got = [k3_head_gemv(x8[:m], weight, ws) for m in seq] + for i, (m, y) in enumerate(zip(seq, got)): + assert _same(y, want[m]), f"call {i} (M={m}) differs" + assert _at_rest(ws) + + +def test_two_workspaces_interleaved_and_two_weights_on_one() -> None: + """Two workspaces of the shard's shape, each with its own weight, their calls alternating on one stream; then both + weights alternating on one workspace: every call bit-identical to the same call on a fresh workspace, every + workspace's words at zero after.""" + weight, weight2, x8 = _inputs() + weights = (weight, weight2) + want = {} + for m in (1, 4, 8): + for i in (0, 1): + want[(m, i)] = k3_head_gemv(x8[:m], weights[i], _workspace()) + ws_a, ws_b = _workspace(), _workspace() + for m in (8, 1, 4, 1, 8): + for i, ws in ((0, ws_a), (1, ws_b)): + y = k3_head_gemv(x8[:m], weights[i], ws) + assert _same(y, want[(m, i)]), f"interleaved: M={m} weight {i}" + for m, i in ((8, 1), (8, 0), (1, 1), (4, 0), (4, 1), (1, 0)): + y = k3_head_gemv(x8[:m], weights[i], ws_a) + assert _same(y, want[(m, i)]), f"one workspace: M={m} weight {i}" + assert _at_rest(ws_a) and _at_rest(ws_b) + + +def test_graph_replays_between_eager_calls() -> None: + """A CUDA graph of M 1, 8 and 3 calls on one workspace, captured once and replayed three times on the stream of the + eager calls, with the activation rewritten in place before each replay and an eager M 8 call on the same workspace + after it: every replayed and eager result bit-identical to the same call on a fresh workspace, the words at zero + after each round.""" + weight, _, x8 = _inputs() + ws = _workspace() + x_buf = x8.clone() + for m in (1, 3, 8): + k3_head_gemv(x_buf[:m], weight, ws) # compiles outside the capture + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + out = {m: k3_head_gemv(x_buf[:m], weight, ws) for m in (1, 8, 3)} + gen = torch.Generator(device="cuda").manual_seed(4) + for rep in range(3): + x_buf.copy_(torch.randn(8, HIDDEN, generator=gen, device="cuda").bfloat16()) + graph.replay() + between = k3_head_gemv(x_buf, weight, ws) + fresh = _workspace() + for m in (1, 3, 8): + assert _same(out[m], k3_head_gemv(x_buf[:m], weight, fresh)), f"replay {rep}, M={m}" + assert _same(between, out[8]), f"replay {rep}: the eager call differs" + assert _at_rest(ws), f"replay {rep}: words left raised" + + +def test_refusals() -> None: + """Negative controls, none of which launches anything: a workspace created for another weight shape ([5120, 7168]: + other flag-word and partial sizes) is refused and stays at rest, and a call right after on a right workspace is + unaffected; creating a workspace under CUDA-graph capture, an unknown schedule, M 0 or 9, and N not a multiple of + 128 are refused.""" + weight, _, x8 = _inputs() + other = K3HeadGemvWorkspace.create(VOCAB_SHARD // 2, HIDDEN, torch.device("cuda")) + with pytest.raises(ValueError, match="workspace"): + k3_head_gemv(x8[:1], weight, other) + assert _at_rest(other) + ws = _workspace() + assert _same(k3_head_gemv(x8[:1], weight, ws), k3_head_gemv(x8[:1], weight, _workspace())) + graph = torch.cuda.CUDAGraph() + with pytest.raises(RuntimeError, match="capture"): + with torch.cuda.graph(graph): + _workspace() + with pytest.raises(ValueError, match="unknown schedule"): + K3HeadGemvWorkspace.create(VOCAB_SHARD, HIDDEN, torch.device("cuda"), schedule="splitk") + for m in (0, 9): + x = torch.zeros(m, HIDDEN, dtype=torch.bfloat16, device="cuda") + with pytest.raises(ValueError, match="unsupported call"): + k3_head_gemv(x, weight, ws) + odd = torch.zeros(200, HIDDEN, dtype=torch.bfloat16, device="cuda") + with pytest.raises(ValueError, match="unsupported call"): + k3_head_gemv(x8[:1], odd, ws) + assert _at_rest(ws) diff --git a/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_attn_res_fwd.py b/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_attn_res_fwd.py new file mode 100644 index 000000000000..f4a4a8bf3f50 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_attn_res_fwd.py @@ -0,0 +1,168 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the attn_res_fwd catalog entry. + +A cell is (T, N): T tokens and N = K + 1 candidates, K snapshots plus the layer residual. Every cell +checks the four outputs' shapes and dtypes, bit-identical results from two identical calls, untouched +inputs, and closeness to an fp32 torch evaluation of the contract's Semantics under the metric of +main's op test (tests/unittest/_torch/modules/kimi_k3_attn_res/test_attn_res_op.py). + +The op picks its kernel from (T, N): at T == 1 the single-CTA decode kernel for N in {1, 2, 4} and the +split-K cluster kernel for N in {8, 12}; the fixed-N = 12 online variant at (T, N) == (1024, 12); the +persistent online kernel everywhere else. +""" + +import pytest +import torch + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.norm.attn_res_fwd import attn_res_fwd + +assert torch.cuda.is_available(), "attn_res_fwd requires a CUDA device" +# The receipts are sm_100 ones; CI's other architectures skip this file. +pytestmark = pytest.mark.skipif( + torch.cuda.get_device_capability() != (10, 0), reason="certified on sm_100 only" +) + +HIDDEN_SIZE = 7168 +RMS_EPS = 1e-6 +# Kimi K3's 93 layers append a snapshot every 12 layers: at most 8 snapshots, so N runs 2 ... 9. +K3_NUM_CANDIDATES = tuple(range(2, 10)) +DECODE_NUM_TOKENS = tuple(range(1, 9)) +PREFILL_NUM_TOKENS = (300, 2048) +# The thresholds of main's op test. +MIN_COSINE = 0.999 +MAX_RELATIVE_L2 = 3e-2 +INPUT_NAMES = ("layer_residual", "block_residual", "res_weight", "rms_weight") +OUTPUT_NAMES = ("output", "rsigma", "probs", "logits") + + +def _make_inputs(num_tokens: int, num_candidates: int) -> tuple[torch.Tensor, ...]: + """Inputs at the scales of main's op tests, seeded by the cell.""" + torch.manual_seed(97 * num_tokens + num_candidates) + layer_residual = ( + torch.randn(num_tokens, 1, HIDDEN_SIZE, dtype=torch.bfloat16, device="cuda") * 0.05 + ) + block_residual = ( + torch.randn( + num_candidates - 1, num_tokens, 1, HIDDEN_SIZE, dtype=torch.bfloat16, device="cuda" + ) + * 0.05 + ) + res_weight = torch.randn(HIDDEN_SIZE, dtype=torch.bfloat16, device="cuda") * 0.02 + rms_weight = 1 + torch.randn(HIDDEN_SIZE, dtype=torch.bfloat16, device="cuda") * 0.02 + return layer_residual, block_residual, res_weight, rms_weight + + +def _reference( + layer_residual: torch.Tensor, + block_residual: torch.Tensor, + res_weight: torch.Tensor, + rms_weight: torch.Tensor, + rms_eps: float, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """fp32 evaluation of the contract's Semantics; snapshots first, the layer residual last.""" + values = torch.cat((block_residual, layer_residual.unsqueeze(0)), dim=0).float() + rsigma = torch.rsqrt(values.square().mean(dim=-1) + rms_eps) + score_weight = rms_weight.float() * res_weight.float() + logits = (values * rsigma.unsqueeze(-1) * score_weight).sum(dim=-1) + probs = torch.softmax(logits, dim=0) + output = (probs.unsqueeze(-1) * values).sum(dim=0).to(torch.bfloat16) + return output, rsigma, probs, logits + + +def _bits(tensor: torch.Tensor) -> torch.Tensor: + """Reinterpret as integers, so that equality means identical bits.""" + return tensor.view(torch.int16 if tensor.element_size() == 2 else torch.int32) + + +def _similarity(actual: torch.Tensor, expected: torch.Tensor) -> tuple[float, float]: + """Cosine similarity and relative L2 error, the metric of main's op tests.""" + actual_float = actual.float().flatten() + expected_float = expected.float().flatten() + cosine = torch.nn.functional.cosine_similarity(actual_float, expected_float, dim=0).item() + relative_l2 = ((actual_float - expected_float).norm() / (expected_float.norm() + 1e-12)).item() + return cosine, relative_l2 + + +def _check(num_tokens: int, num_candidates: int, rms_eps: float = RMS_EPS) -> None: + inputs = _make_inputs(num_tokens, num_candidates) + before = [tensor.clone() for tensor in inputs] + outputs = attn_res_fwd(*inputs, rms_eps) + repeat = attn_res_fwd(*inputs, rms_eps) + expected = _reference(*inputs, rms_eps) + + cell = f"T={num_tokens} N={num_candidates} rms_eps={rms_eps}" + stats_shape = (num_candidates, num_tokens, 1) + shapes = ((num_tokens, 1, HIDDEN_SIZE), stats_shape, stats_shape, stats_shape) + dtypes = (torch.bfloat16, torch.float32, torch.float32, torch.float32) + assert len(outputs) == len(OUTPUT_NAMES), f"{cell}: {len(outputs)} outputs" + for name, actual, again, reference, shape, dtype in zip( + OUTPUT_NAMES, outputs, repeat, expected, shapes, dtypes + ): + assert actual.shape == shape, f"{cell}: {name} shape {tuple(actual.shape)}" + assert actual.dtype == dtype, f"{cell}: {name} dtype {actual.dtype}" + assert actual.is_contiguous(), f"{cell}: {name} is not contiguous" + assert actual.device == inputs[0].device, f"{cell}: {name} on {actual.device}" + assert torch.equal(_bits(actual), _bits(again)), ( + f"{cell}: two identical calls disagree in {name}" + ) + cosine, relative_l2 = _similarity(actual, reference) + assert cosine > MIN_COSINE and relative_l2 < MAX_RELATIVE_L2, ( + f"{cell}: {name} cosine {cosine:.6f}, relative L2 {relative_l2:.3e}" + ) + for name, old, new in zip(INPUT_NAMES, before, inputs): + assert torch.equal(_bits(old), _bits(new)), f"{cell}: {name} was mutated" + + +def test_k3_decode_cells() -> None: + """Kimi K3's N = 2 ... 9 at T = 1 ... 8: every T == 1 kernel, then the online kernel.""" + for num_tokens in DECODE_NUM_TOKENS: + for num_candidates in K3_NUM_CANDIDATES: + _check(num_tokens, num_candidates) + + +def test_k3_prefill_cells() -> None: + """Kimi K3's N = 2 ... 9 at prefill token counts: the online kernel's CTAs loop over tokens.""" + for num_tokens in PREFILL_NUM_TOKENS: + for num_candidates in K3_NUM_CANDIDATES: + _check(num_tokens, num_candidates) + + +def test_other_dispatch_branches() -> None: + """The kernels and N extremes outside Kimi K3's range, down to the fixed-N = 12 variant.""" + for num_tokens, num_candidates in ( + (1, 1), # single-CTA kernel, no snapshot + (1, 10), # online kernel at T == 1 + (1, 11), # online kernel at T == 1 + (1, 12), # split-K kernel + (8, 1), # online kernel, a single candidate + (8, 12), # online kernel, three full chunks of four candidates + (1024, 12), # online kernel, fixed-N = 12 variant + ): + _check(num_tokens, num_candidates) + + +def test_rms_eps_enters_under_the_root() -> None: + """rms_eps is added to the mean square inside the rsqrt; 1e-2 dominates the inputs' 2.5e-3.""" + for num_tokens, num_candidates in ((1, 2), (1, 8), (1, 9), (300, 9)): + for rms_eps in (1e-5, 1e-2): + _check(num_tokens, num_candidates, rms_eps) + + +def test_cuda_graph_replay_matches_eager() -> None: + """One cell per kernel: a call captured in a CUDA graph replays the eager call's bits.""" + for num_tokens, num_candidates in ((1, 4), (1, 8), (4, 9)): + inputs = _make_inputs(num_tokens, num_candidates) + eager = attn_res_fwd(*inputs, RMS_EPS) + torch.cuda.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + captured = attn_res_fwd(*inputs, RMS_EPS) + for tensor in captured: + tensor.fill_(float("nan")) + graph.replay() + torch.cuda.synchronize() + for name, expected, actual in zip(OUTPUT_NAMES, eager, captured): + assert torch.equal(_bits(actual), _bits(expected)), ( + f"T={num_tokens} N={num_candidates}: replayed {name} differs from eager" + ) diff --git a/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_attn_res_rmsnorm_fwd.py b/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_attn_res_rmsnorm_fwd.py new file mode 100644 index 000000000000..48ef9887ccac --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_attn_res_rmsnorm_fwd.py @@ -0,0 +1,181 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the attn_res_rmsnorm_fwd catalog entry. + +A cell is (T, N): T tokens and N = K + 1 candidates, K snapshots plus the layer residual. Every cell +checks the output's shape and dtype, bit-identical results from two identical calls, untouched inputs, +and closeness to an fp32 torch evaluation of the contract's Semantics, with both bf16 rounding +boundaries, under the metric of main's op test +(tests/unittest/_torch/modules/kimi_k3_attn_res/test_attn_res_rmsnorm_op.py). + +The kernel depends on N only: one CTA per token for N <= 4, one cluster of 8 CTAs per token for +N >= 5. With PDL enabled both release their dependents right after their own grid-dependency wait. +""" + +import pytest +import torch + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.norm.attn_res_rmsnorm_fwd import ( + attn_res_rmsnorm_fwd, +) + +assert torch.cuda.is_available(), "attn_res_rmsnorm_fwd requires a CUDA device" +# The receipts are sm_100 ones; CI's other architectures skip this file. +pytestmark = pytest.mark.skipif( + torch.cuda.get_device_capability() != (10, 0), reason="certified on sm_100 only" +) + +HIDDEN_SIZE = 7168 +RMS_EPS = 1e-6 +# Kimi K3's 93 layers append a snapshot every 12 layers: at most 8 snapshots, so N runs 2 ... 9. +K3_NUM_CANDIDATES = tuple(range(2, 10)) +DECODE_NUM_TOKENS = tuple(range(1, 9)) +# 32 is the fused path's token ceiling when KIMI_K3_ATTN_RES_TOPOLOGY is on. +LARGER_NUM_TOKENS = (32, 300) +# The thresholds of main's op test. +MIN_COSINE = 0.9999 +MAX_RELATIVE_L2 = 5e-3 +INPUT_NAMES = ("layer_residual", "block_residual", "res_weight", "rms_weight", "output_rms_weight") + + +def _make_inputs(num_tokens: int, num_candidates: int) -> tuple[torch.Tensor, ...]: + """Inputs at the scales of main's op test, seeded by the cell.""" + torch.manual_seed(97 * num_tokens + num_candidates) + layer_residual = ( + torch.randn(num_tokens, 1, HIDDEN_SIZE, dtype=torch.bfloat16, device="cuda") * 0.05 + ) + block_residual = ( + torch.randn( + num_candidates - 1, num_tokens, 1, HIDDEN_SIZE, dtype=torch.bfloat16, device="cuda" + ) + * 0.05 + ) + res_weight = torch.randn(HIDDEN_SIZE, dtype=torch.bfloat16, device="cuda") * 0.02 + rms_weight = 1 + torch.randn(HIDDEN_SIZE, dtype=torch.bfloat16, device="cuda") * 0.02 + output_rms_weight = 1 + torch.randn(HIDDEN_SIZE, dtype=torch.bfloat16, device="cuda") * 0.02 + return layer_residual, block_residual, res_weight, rms_weight, output_rms_weight + + +def _reference( + layer_residual: torch.Tensor, + block_residual: torch.Tensor, + res_weight: torch.Tensor, + rms_weight: torch.Tensor, + output_rms_weight: torch.Tensor, + rms_eps: float, + output_rms_eps: float, +) -> torch.Tensor: + """fp32 evaluation of the contract's Semantics, rounding where the kernel rounds.""" + values = torch.cat((block_residual, layer_residual.unsqueeze(0)), dim=0).float() + rsigma = torch.rsqrt(values.square().mean(dim=-1, keepdim=True) + rms_eps) + logits = (values * rsigma * (rms_weight.float() * res_weight.float())).sum(dim=-1) + probs = torch.softmax(logits, dim=0) + mixed = (probs.unsqueeze(-1) * values).sum(dim=0).to(torch.bfloat16).float() + normed = mixed * torch.rsqrt(mixed.square().mean(dim=-1, keepdim=True) + output_rms_eps) + # bf16 times bf16: the exact product, rounded to bf16 once more. + return output_rms_weight * normed.to(torch.bfloat16) + + +def _bits(tensor: torch.Tensor) -> torch.Tensor: + """Reinterpret as integers, so that equality means identical bits.""" + return tensor.view(torch.int16) + + +def _similarity(actual: torch.Tensor, expected: torch.Tensor) -> tuple[float, float]: + """Cosine similarity and relative L2 error, the metric of main's op tests.""" + actual_float = actual.float().flatten() + expected_float = expected.float().flatten() + cosine = torch.nn.functional.cosine_similarity(actual_float, expected_float, dim=0).item() + relative_l2 = ((actual_float - expected_float).norm() / (expected_float.norm() + 1e-12)).item() + return cosine, relative_l2 + + +def _check( + num_tokens: int, + num_candidates: int, + rms_eps: float = RMS_EPS, + output_rms_eps: float = RMS_EPS, +) -> None: + inputs = _make_inputs(num_tokens, num_candidates) + before = [tensor.clone() for tensor in inputs] + output = attn_res_rmsnorm_fwd(*inputs, rms_eps, output_rms_eps) + repeat = attn_res_rmsnorm_fwd(*inputs, rms_eps, output_rms_eps) + expected = _reference(*inputs, rms_eps, output_rms_eps) + + cell = f"T={num_tokens} N={num_candidates} eps=({rms_eps}, {output_rms_eps})" + assert output.shape == (num_tokens, 1, HIDDEN_SIZE), f"{cell}: shape {tuple(output.shape)}" + assert output.dtype == torch.bfloat16, f"{cell}: dtype {output.dtype}" + assert output.is_contiguous(), f"{cell}: output is not contiguous" + assert output.device == inputs[0].device, f"{cell}: output on {output.device}" + assert torch.equal(_bits(output), _bits(repeat)), f"{cell}: two identical calls disagree" + cosine, relative_l2 = _similarity(output, expected) + assert cosine > MIN_COSINE and relative_l2 < MAX_RELATIVE_L2, ( + f"{cell}: cosine {cosine:.6f}, relative L2 {relative_l2:.3e}" + ) + for name, old, new in zip(INPUT_NAMES, before, inputs): + assert torch.equal(_bits(old), _bits(new)), f"{cell}: {name} was mutated" + + +def test_k3_decode_cells() -> None: + """Kimi K3's N = 2 ... 9 at T = 1 ... 8, on both kernels.""" + for num_tokens in DECODE_NUM_TOKENS: + for num_candidates in K3_NUM_CANDIDATES: + _check(num_tokens, num_candidates) + + +def test_larger_token_counts() -> None: + """Kimi K3's N = 2 ... 9 at more tokens: more CTAs or clusters, the same per-token kernel.""" + for num_tokens in LARGER_NUM_TOKENS: + for num_candidates in K3_NUM_CANDIDATES: + _check(num_tokens, num_candidates) + + +def test_other_candidate_counts() -> None: + """N = 1 (single-CTA kernel, no snapshot) and N = 10 ... 12 (split-K), outside Kimi K3's range.""" + for num_tokens in (1, 8): + for num_candidates in (1, 10, 11, 12): + _check(num_tokens, num_candidates) + + +def test_each_eps_feeds_its_own_norm() -> None: + """rms_eps scores the candidates and output_rms_eps scales the output; 1e-2 dominates either.""" + for num_tokens, num_candidates in ((1, 4), (1, 9), (8, 2)): + for rms_eps, output_rms_eps in ((1e-5, 1e-5), (1e-2, 1e-6), (1e-6, 1e-2)): + _check(num_tokens, num_candidates, rms_eps, output_rms_eps) + + +def test_cuda_graph_replay_matches_eager() -> None: + """One cell per kernel: a call captured in a CUDA graph replays the eager call's bits.""" + for num_tokens, num_candidates in ((1, 4), (1, 9), (8, 5)): + inputs = _make_inputs(num_tokens, num_candidates) + eager = attn_res_rmsnorm_fwd(*inputs, RMS_EPS, RMS_EPS) + torch.cuda.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + captured = attn_res_rmsnorm_fwd(*inputs, RMS_EPS, RMS_EPS) + captured.fill_(float("nan")) + graph.replay() + torch.cuda.synchronize() + assert torch.equal(_bits(captured), _bits(eager)), ( + f"T={num_tokens} N={num_candidates}: replay differs from eager" + ) + + +def test_chained_calls_wait_for_their_input() -> None: + """A call reading the previous call's output reads it settled, although that call triggers early. + + With PDL enabled the first call releases its dependents before it writes its output, so the second + call can start while the first still runs; its grid-dependency wait must hold its reads back. + """ + for num_tokens, num_candidates in ((1, 4), (1, 9), (8, 4), (8, 9)): + layer_residual, block_residual, res_weight, rms_weight, output_rms_weight = _make_inputs( + num_tokens, num_candidates + ) + weights = (res_weight, rms_weight, output_rms_weight, RMS_EPS, RMS_EPS) + first = attn_res_rmsnorm_fwd(layer_residual, block_residual, *weights) + chained = attn_res_rmsnorm_fwd(first, block_residual, *weights) + torch.cuda.synchronize() + settled = attn_res_rmsnorm_fwd(first.clone(), block_residual, *weights) + assert torch.equal(_bits(chained), _bits(settled)), ( + f"T={num_tokens} N={num_candidates}: the chained call read an unsettled input" + ) diff --git a/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_k3_embed_norm.py b/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_k3_embed_norm.py new file mode 100644 index 000000000000..89b1016e6d70 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_k3_embed_norm.py @@ -0,0 +1,211 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the k3_embed_norm catalog entry. + +Kimi K3's replicated embedding table [163840, 7168] at every token count of the K3 decode steps (N = 1..8, and +16..64 for DSpark verify steps of up to 8 requests x 8 tokens), against a native torch reference: the rows written into +`raw` bit for bit against a torch gather, the normed rows against an fp64 RMSNorm (within 8e-3 of each row's max +|ref|). Also: run-to-run bits, each row's bits independent of N, int64 ids equal to int32 ids, another eps and bank +slot, ids outside [0, V), a CUDA graph replayed with its ids rewritten, and the refused calls. +""" + +import functools +import math + +import pytest +import torch + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.norm.k3_embed_norm import k3_embed_norm + +assert torch.cuda.is_available(), "k3_embed_norm requires a CUDA device" +# The receipts are sm_100 ones; CI's other architectures skip this file. +pytestmark = pytest.mark.skipif( + torch.cuda.get_device_capability() != (10, 0), reason="certified on sm_100 only" +) + +VOCAB, HIDDEN = 163840, 7168 +EPS, EPS_ALT = 1e-6, 1e-5 +TOL = 8e-3 # per row: max |out - ref| / max |ref|, ref an fp64 RMSNorm +TOKENS = list(range(1, 9)) + list(range(16, 65, 8)) +INT64_TOKENS = (1, 8, 64) +BANK_SLOTS = 4 # the attention-residual snapshot bank; the model passes slot 0 as raw +SENTINEL = 7.0 +OUTSIDE_IDS = (-1, VOCAB, VOCAB + 7, -VOCAB, 2**31 - 1, -(2**31)) +OUTSIDE_IDS_INT64 = (2**32, 2**32 + 5, 3 - 2**32, 2**40) + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def _same(a: torch.Tensor, b: torch.Tensor) -> bool: + return a.shape == b.shape and torch.equal(_bits(a), _bits(b)) + + +@functools.lru_cache(maxsize=None) +def _table() -> tuple[torch.Tensor, int]: + """The table, built once: N(0, 1) rows at log-uniform scales 1e-4..1e2, row 1 all zero; and its smallest-scale + row's id (a row whose norm eps dominates).""" + gen = torch.Generator(device="cuda").manual_seed(20260930) + table = torch.randn(VOCAB, HIDDEN, generator=gen, device="cuda", dtype=torch.bfloat16) + scale = torch.empty(VOCAB, 1, device="cuda") + scale.uniform_(math.log(1e-4), math.log(1e2), generator=gen).exp_() + table.mul_(scale.bfloat16()) + table[1].zero_() + scale[1] = math.inf + return table, int(scale.argmin()) + + +@functools.lru_cache(maxsize=None) +def _weight() -> torch.Tensor: + """Layer 0's input-norm weight: around 1, every 101st element negative.""" + gen = torch.Generator(device="cuda").manual_seed(7) + weight = 1.0 + 0.2 * torch.randn(HIDDEN, generator=gen, device="cuda") + weight[::101] *= -1.0 + return weight.bfloat16() + + +@functools.lru_cache(maxsize=None) +def _ids64() -> torch.Tensor: + """64 int32 ids in [0, V): random, with 0, V - 1, the zero row, the smallest-scale row and repeats in front.""" + _, small = _table() + gen = torch.Generator(device="cuda").manual_seed(11) + ids = torch.randint(0, VOCAB, (64,), generator=gen, device="cuda", dtype=torch.int32) + ids[0], ids[1], ids[2], ids[4] = 0, VOCAB - 1, 1, small + ids[3] = ids[1] + ids[6] = ids[5] + torch.cuda.synchronize() + return ids + + +def _bank(n: int) -> torch.Tensor: + """A snapshot bank [BANK_SLOTS, n, H] filled with SENTINEL.""" + return torch.full((BANK_SLOTS, n, HIDDEN), SENTINEL, dtype=torch.bfloat16, device="cuda") + + +def _rows_ref(ids: torch.Tensor) -> torch.Tensor: + """The torch gather: table[ids], zero rows for ids outside [0, V).""" + table, _ = _table() + valid = (ids >= 0) & (ids < VOCAB) + rows = table.index_select(0, torch.where(valid, ids, torch.zeros_like(ids))) + return rows.masked_fill(~valid[:, None], 0.0) + + +def _norm_ref(rows: torch.Tensor, eps: float) -> torch.Tensor: + """fp64 RMSNorm: rows * rsqrt(mean(rows^2) + eps) * weight.""" + x = rows.double() + return x * torch.rsqrt(x.pow(2).mean(dim=1, keepdim=True) + eps) * _weight().double() + + +def _call(ids: torch.Tensor, eps: float = EPS, slot: int = 0) -> tuple[torch.Tensor, torch.Tensor]: + """k3_embed_norm into ``slot`` of a fresh bank: (out, bank).""" + table, _ = _table() + bank = _bank(ids.numel()) + return k3_embed_norm(ids, table, _weight(), eps, bank[slot]), bank + + +def _check(ids: torch.Tensor, eps: float = EPS, slot: int = 0) -> torch.Tensor: + """One call's checks against the reference and a rerun; returns the call's output.""" + n = ids.numel() + what = f"N={n} {ids.dtype} eps={eps:g} slot={slot}" + out, bank = _call(ids, eps, slot) + again, bank_again = _call(ids, eps, slot) + rows = _rows_ref(ids) + ref = _norm_ref(rows, eps) + assert out.shape == (n, HIDDEN) and out.dtype == torch.bfloat16 and out.device == ids.device + assert out.is_contiguous() + assert _same(bank[slot], rows), f"{what}: raw is not the gathered rows" + others = [s for s in range(BANK_SLOTS) if s != slot] + assert bool((bank[others] == SENTINEL).all()), f"{what}: wrote outside raw" + scale = ref.abs().amax(dim=1) + live = scale > 0 + err = ((out.double() - ref).abs().amax(dim=1)[live] / scale[live]).max().item() + assert err <= TOL, f"{what}: rel {err:.3e}" + assert bool((out[~live] == 0).all()), f"{what}: a zero row normed to non-zero" + assert _same(again, out) and _same(bank_again, bank), f"{what}: a rerun differs" + return out + + +def test_certified_cells() -> None: + """Every N with int32 ids, eps 1e-6, raw = slot 0; every row bit-identical to the same id's row of the 64-id + call.""" + ids64 = _ids64() + full = _check(ids64) + for n in TOKENS: + out = _check(ids64[:n]) + assert _same(out, full[:n]), f"N={n}: rows differ from the 64-id call" + + +def test_int64_ids_other_eps_and_slot() -> None: + """int64 ids at N 1, 8 and 64, bit-identical to int32 ids; eps 1e-5 into bank slot 1 at N 8.""" + ids64 = _ids64() + for n in INT64_TOKENS: + ids = ids64[:n] + assert _same(_check(ids.long()), _call(ids)[0]), f"N={n}: int64 ids differ from int32" + _check(ids64[:8], eps=EPS_ALT, slot=1) + + +def test_ids_outside_vocab() -> None: + """N 8, every out-of-range value at the even positions (int32 and int64): zero raw rows, zero normed rows, the other + rows as the reference.""" + gen = torch.Generator(device="cuda").manual_seed(3008) + for dtype, bad in ( + (torch.int32, OUTSIDE_IDS), + (torch.int64, OUTSIDE_IDS + OUTSIDE_IDS_INT64), + ): + for start in range(0, len(bad), 4): + ids = torch.randint(0, VOCAB, (8,), generator=gen, device="cuda", dtype=dtype) + values = [bad[(start + i) % len(bad)] for i in range(4)] + ids[0::2] = torch.tensor(values, dtype=dtype, device="cuda") + out = _check(ids) + assert bool((out[0::2] == 0).all()), f"{dtype} {values}: non-zero rows" + + +def test_graph_replay() -> None: + """An N = 8 call captured once after an eager call compiled it, replayed with the ids rewritten in place (random, + then some outside [0, V)): every replay bit-identical to an eager call on the same ids, raw the gathered rows.""" + table, _ = _table() + ids_buf = _ids64()[:8].clone() + bank = _bank(8) + k3_embed_norm(ids_buf, table, _weight(), EPS, bank[0]) # compiles outside the capture + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + out = k3_embed_norm(ids_buf, table, _weight(), EPS, bank[0]) + gen = torch.Generator(device="cuda").manual_seed(4008) + for rep in range(3): + ids = torch.randint(0, VOCAB, (8,), generator=gen, device="cuda", dtype=torch.int32) + if rep == 2: + ids[0::2] = -1 + ids_buf.copy_(ids) + bank.fill_(SENTINEL) + graph.replay() + eager, eager_bank = _call(ids_buf) + assert _same(out, eager) and _same(bank, eager_bank), f"replay {rep} differs from eager" + assert _same(bank[0], _rows_ref(ids_buf)), f"replay {rep}: raw is not the gathered rows" + + +def test_refused_calls() -> None: + """Calls outside the op's support raise ValueError before launching anything: raw stays untouched.""" + table, _ = _table() + weight = _weight() + ids8 = _ids64()[:8] + bank = _bank(8) + + def refused(ids, tab, w, raw): + with pytest.raises(ValueError, match="unsupported call"): + k3_embed_norm(ids, tab, w, EPS, raw) + + refused(ids8[:0], table, weight, bank[0, :0]) # N = 0 + ids65 = torch.zeros(65, dtype=torch.int32, device="cuda") + refused(ids65, table, weight, torch.empty(65, HIDDEN, dtype=torch.bfloat16, device="cuda")) + refused(ids8.to(torch.int16), table, weight, bank[0]) # id dtype + refused(ids8, table, weight, _bank(64)[0]) # raw [64, H] for 8 ids + flat = torch.empty(8 * HIDDEN + 8, dtype=torch.bfloat16, device="cuda") + refused(ids8, table, weight, flat[1 : 1 + 8 * HIDDEN].view(8, HIDDEN)) # raw 2 bytes off + refused(ids8, table, weight[: HIDDEN // 2], bank[0]) # weight [H / 2] + for hidden in (6144, 7680): # outside the norm's geometry + small = torch.zeros(16, hidden, dtype=torch.bfloat16, device="cuda") + w = torch.ones(hidden, dtype=torch.bfloat16, device="cuda") + raw = torch.empty(8, hidden, dtype=torch.bfloat16, device="cuda") + refused(ids8 % 16, small, w, raw) + assert bool((bank == SENTINEL).all()) From c69ed9c9d4ca8a077550cadcaba40ff2c0ea4fd9 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 18:48:36 -0700 Subject: [PATCH 030/161] [None][chore] Kimi K3 decode kernels: format with the repo's pre-commit hooks ruff-format reflows k3_head_gemv/op.py and the op tests test_k3_ctm_gemv, test_k3_decode_gemv, test_k3_head_gemv and test_k3_tcgen05_fences to the repo's 100-column code style, which also clears an E501 in test_k3_ctm_gemv. test_k3_head_gemv drops an unused assignment (F841). Apart from that assignment, every file's AST is unchanged. Signed-off-by: Vasanth Sabavat --- .../cute_dsl_kernels/k3_head_gemv/op.py | 13 ++- .../kimi_k3/test_k3_ctm_gemv.py | 95 +++++++++++++++---- .../kimi_k3/test_k3_decode_gemv.py | 8 +- .../kimi_k3/test_k3_head_gemv.py | 6 +- .../kimi_k3/test_k3_tcgen05_fences.py | 1 + 5 files changed, 102 insertions(+), 21 deletions(-) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_head_gemv/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_head_gemv/op.py index 3434256477d8..91e0bd132da1 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_head_gemv/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_head_gemv/op.py @@ -253,5 +253,16 @@ def k3_head_gemv( @k3_head_gemv.register_fake -def _(x, weight, partials, flags, claim, keep_tiles=0, chunk_tiles=0, ring=6, schedule="streamk", prefetch=16): +def _( + x, + weight, + partials, + flags, + claim, + keep_tiles=0, + chunk_tiles=0, + ring=6, + schedule="streamk", + prefetch=16, +): return x.new_empty((x.shape[0], weight.shape[0]), dtype=torch.bfloat16) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctm_gemv.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctm_gemv.py index 757fbba95655..972631c8d5c1 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctm_gemv.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctm_gemv.py @@ -46,20 +46,41 @@ def _is_sm100() -> bool: } # k3_ctm_gemv_long: (N, K, sig_col0, split, ring, push), as each call site passes them. LONG = { - "mla_qkv_a_gate": (2880, 7168, 2112, 6, 6, True), # [W_a; W_g], the gate columns as bf16(sigmoid) + "mla_qkv_a_gate": ( + 2880, + 7168, + 2112, + 6, + 6, + True, + ), # [W_a; W_g], the gate columns as bf16(sigmoid) "dense_gate_up": (4224, 7168, -1, 4, 5, False), "dense_down": (7168, 2112, -1, 2, 6, False), # K 2112 ends in a half k-tile "drafter_qkv": (512, 7168, -1, 8, 6, True), "drafter_gate_up": (1792, 7168, -1, 8, 6, True), "drafter_gate_up_dummy": (1536, 7168, -1, 8, 6, True), - "kda_qkvg": (3208, 7168, -1, 5, 6, True), # KDA q/k/v/g/f_a/b: 25 whole tiles and an 8-row last one + "kda_qkvg": ( + 3208, + 7168, + -1, + 5, + 6, + True, + ), # KDA q/k/v/g/f_a/b: 25 whole tiles and an 8-row last one } # k3_ctm_gemv_swiglu: (N, K, split, push), the drafter down projection (tuned and synthetic). SWIGLU = {"drafter_down": (7168, 896, 2, True), "drafter_down_dummy": (7168, 768, 2, True)} -AG_COLS, G_COL0, K_O = 2880, 2112, 768 # MLA: [q_a 1536 | kv_a 576 | gate 768] rows of the fused projection +AG_COLS, G_COL0, K_O = ( + 2880, + 2112, + 768, +) # MLA: [q_a 1536 | kv_a 576 | gate 768] rows of the fused projection HIDDEN, LATENT, WIDTH, PAD, ACT = 7168, 3584, 224, 256, 384 EPS = 1e-6 -SITU = {"k3": (4.0, 25.0), "plain": (1.0, None)} # (beta, linear_beta): the K3 checkpoint's, and the defaults +SITU = { + "k3": (4.0, 25.0), + "plain": (1.0, None), +} # (beta, linear_beta): the K3 checkpoint's, and the defaults def _ops(): @@ -88,7 +109,9 @@ def _report(op, case, m, y, ref, stock=None, **flags): extra = "" if stock is None else f" vs_stock={_rel(y, stock):.2e}" marks = " ".join(f"{k}={v}" for k, v in flags.items()) abs_err = (y.double() - ref.double()).abs().max().item() - print(f"OPCHECK op={op} case={case} M={m} abs={abs_err:.3e} rel={_rel(y, ref):.3e}{extra} {marks}") + print( + f"OPCHECK op={op} case={case} M={m} abs={abs_err:.3e} rel={_rel(y, ref):.3e}{extra} {marks}" + ) @functools.lru_cache(maxsize=None) @@ -119,11 +142,23 @@ def test_k3_ctm_gemv(name, m): n, k, split, push = PLAIN[name] w = _weight(n, k, 1, 0.03) assert _ctm().supports(_rows(k, 2)[:m], w, split) - (x,), y, det, minv = _checks(lambda x_: ops.k3_ctm_gemv(x_, w, True, split, push), [_rows(k, 2)], m) + (x,), y, det, minv = _checks( + lambda x_: ops.k3_ctm_gemv(x_, w, True, split, push), [_rows(k, 2)], m + ) ref = x.double() @ w.double().t() decode = ops.k3_decode_gemv(x, w, True) same_decode = torch.equal(_bits(y), _bits(decode)) - _report("k3_ctm_gemv", name, m, y, ref, F.linear(x, w), det=det, rows_as_m8=minv, eq_k3_decode_gemv=same_decode) + _report( + "k3_ctm_gemv", + name, + m, + y, + ref, + F.linear(x, w), + det=det, + rows_as_m8=minv, + eq_k3_decode_gemv=same_decode, + ) assert _rel(y, ref) <= TOL and _rel(y, F.linear(x, w)) <= TOL assert det and minv if split == 1: @@ -177,7 +212,9 @@ def test_k3_ctm_gemv_swiglu(name, m): w = _weight(n, k, 7) gu8 = _rows(2 * k, 8, 2.0) assert _ctm().supports_swiglu(gu8, w, split) - (gu,), y, det, minv = _checks(lambda g_: ops.k3_ctm_gemv_swiglu(g_, w, True, split, push), [gu8], m) + (gu,), y, det, minv = _checks( + lambda g_: ops.k3_ctm_gemv_swiglu(g_, w, True, split, push), [gu8], m + ) act = _silu_and_mul(gu) ref = act.double() @ w.double().t() _report("k3_ctm_gemv_swiglu", name, m, y, ref, F.linear(act, w), det=det, rows_as_m8=minv) @@ -228,12 +265,22 @@ def call(lat_, act_): (lat, act), y, det, minv = _checks(call, [_rows(LATENT, 12, 0.8), _rows(ACT, 13, 0.5)], m) lat64 = lat.double() normed = lat64 * torch.rsqrt(lat64.pow(2).mean(dim=1, keepdim=True) + EPS) - ref = torch.cat([normed[:, lo : lo + WIDTH], act.double()], dim=1) @ torch.cat( - [w[:, :WIDTH], w[:, PAD:]], dim=1 - ).double().t() + ref = ( + torch.cat([normed[:, lo : lo + WIDTH], act.double()], dim=1) + @ torch.cat([w[:, :WIDTH], w[:, PAD:]], dim=1).double().t() + ) decode = ops.k3_decode_gemv_tail(lat, act, w, lo, WIDTH, EPS, True) same_decode = torch.equal(_bits(y), _bits(decode)) - _report("k3_ctm_gemv_tail", f"rank{rank}", m, y, ref, det=det, rows_as_m8=minv, eq_k3_decode_gemv_tail=same_decode) + _report( + "k3_ctm_gemv_tail", + f"rank{rank}", + m, + y, + ref, + det=det, + rows_as_m8=minv, + eq_k3_decode_gemv_tail=same_decode, + ) assert _rel(y, ref) <= TOL assert det and minv and same_decode @@ -261,7 +308,17 @@ def test_k3_situ_mul(situ, m): stock = SituAndMul(beta=beta, linear_beta=linear_beta, use_fused_activation=True)(gu) ref = _situ_ref(gu, beta, linear_beta) identical = (_bits(y) == _bits(stock)).float().mean().item() - _report("k3_situ_mul", situ, m, y, ref, stock, det=det, rows_as_m8=minv, frac_eq_triton=f"{identical:.4f}") + _report( + "k3_situ_mul", + situ, + m, + y, + ref, + stock, + det=det, + rows_as_m8=minv, + frac_eq_triton=f"{identical:.4f}", + ) assert y.shape == (m, 2112) assert _rel(y, ref) <= TOL and _rel(y, stock) <= TOL assert det and minv @@ -273,9 +330,15 @@ def test_token_limit(m): _ops() x = torch.zeros(m, 7168, dtype=torch.bfloat16, device="cuda") for n, k, _, split, ring, _ in LONG.values(): - assert not op.supports_long(torch.zeros(m, k, dtype=torch.bfloat16, device="cuda"), _weight(n, k, 9), split, ring) + assert not op.supports_long( + torch.zeros(m, k, dtype=torch.bfloat16, device="cuda"), _weight(n, k, 9), split, ring + ) with pytest.raises(ValueError): torch.ops.trtllm.k3_ctm_gemv_long(x, _weight(2880, 7168, 9), 2112, 6, 6, True, True) - assert not op.supports(torch.zeros(m, 768, dtype=torch.bfloat16, device="cuda"), _weight(7168, 768, 1, 0.03), 1) + assert not op.supports( + torch.zeros(m, 768, dtype=torch.bfloat16, device="cuda"), _weight(7168, 768, 1, 0.03), 1 + ) assert not op.supports_situ_mul(torch.zeros(m, 4224, dtype=torch.bfloat16, device="cuda")) - assert not op.supports_swiglu(torch.zeros(m, 1792, dtype=torch.bfloat16, device="cuda"), _weight(7168, 896, 7), 2) + assert not op.supports_swiglu( + torch.zeros(m, 1792, dtype=torch.bfloat16, device="cuda"), _weight(7168, 896, 7), 2 + ) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_decode_gemv.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_decode_gemv.py index b75aa0ae4396..e3d12396fb67 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_decode_gemv.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_decode_gemv.py @@ -61,7 +61,9 @@ def _report(op, case, m, y, ref, stock=None, **flags): extra = "" if stock is None else f" vs_stock={_rel(y, stock):.2e}" marks = " ".join(f"{k}={v}" for k, v in flags.items()) abs_err = (y.double() - ref.double()).abs().max().item() - print(f"OPCHECK op={op} case={case} M={m} abs={abs_err:.3e} rel={_rel(y, ref):.3e}{extra} {marks}") + print( + f"OPCHECK op={op} case={case} M={m} abs={abs_err:.3e} rel={_rel(y, ref):.3e}{extra} {marks}" + ) @functools.lru_cache(maxsize=None) @@ -90,7 +92,9 @@ def test_k3_decode_gemv(name, m): ref = x.double() @ w.double().t() deterministic = torch.equal(_bits(y), _bits(again)) m_invariant = torch.equal(_bits(y), _bits(y8[:m])) - _report("k3_decode_gemv", name, m, y, ref, F.linear(x, w), det=deterministic, rows_as_m8=m_invariant) + _report( + "k3_decode_gemv", name, m, y, ref, F.linear(x, w), det=deterministic, rows_as_m8=m_invariant + ) assert y.shape == (m, n) assert _rel(y, ref) <= TOL and _rel(y, F.linear(x, w)) <= TOL assert deterministic and m_invariant diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_head_gemv.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_head_gemv.py index b97d9a9f5e91..7e72d01c8481 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_head_gemv.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_head_gemv.py @@ -106,7 +106,7 @@ def _call(x, w, ws): @pytest.mark.parametrize("m", M_ALL) def test_k3_head_gemv(m): - ops = _ops() + _ops() from tensorrt_llm._torch.cute_dsl_kernels.k3_head_gemv import op w, x8 = _inputs() @@ -151,7 +151,9 @@ def test_k3_head_gemv_two_workspaces_interleaved(): single = {(m, i): _call(x8[:m].contiguous(), (w, w2)[i], wa) for m in (1, 4, 8) for i in (0, 1)} for m in (8, 1, 4, 1, 8): for i, ws in ((0, wa), (1, wb)): - assert torch.equal(_bits(_call(x8[:m].contiguous(), (w, w2)[i], ws)), _bits(single[(m, i)])) + assert torch.equal( + _bits(_call(x8[:m].contiguous(), (w, w2)[i], ws)), _bits(single[(m, i)]) + ) assert not wa.flags.any() and not wb.flags.any() diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py index 05584abac209..613e5d559a89 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py @@ -45,6 +45,7 @@ LOAD_WAIT = "Tcgen05Wait.LOAD" SYNC = re.compile(r"mbarrier_arrive\(|barrier_cta_sync\(|cute\.arch\.barrier\(") + def violations(lines): """(line number, rule) of every tcgen05 operation reached from a wait without the after-fence, and every arrive / barrier reached from a TMEM load wait without the before-fence (scanning back within the function).""" From 2700b9d4c2c2f0d85d035eb9280bb187789b4683 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 18:49:40 -0700 Subject: [PATCH 031/161] [None][test] Kimi K3 MLA op tests: drop the batch-1 identity cases test_k3_mla_q.py and test_k3_mla_attn.py compared one request's outputs with the kernel of a second, unmodified tensorrt_llm package named by K3_BASE_TRTLLM, and skipped without it. CI has no such package, so those 9 cases could only ever skip; every other case of both files stays. Signed-off-by: Vasanth Sabavat --- .../kimi_k3/test_k3_mla_attn.py | 114 +--------------- .../cute_dsl_kernels/kimi_k3/test_k3_mla_q.py | 124 +----------------- 2 files changed, 2 insertions(+), 236 deletions(-) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_attn.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_attn.py index 843ca267e52b..c359a444cbc8 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_attn.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_attn.py @@ -16,14 +16,9 @@ heads per rank, latent 512 + rope 64, bf16 pool of 64-row pages) for decode steps of R requests x T tokens: against a torch float64 reference with per-request bottom-right causal masks and against the stock CuTe DSL MLA decode (trtllm::cute_dsl_mla_decode_fp16_blackwell); request i's rows bit-identical to the one-request call on its own -rows, pages and length; reruns bit-identical. +rows, pages and length; reruns bit-identical.""" -Batch-1 identity against the unmodified kernel: set ``K3_BASE_TRTLLM`` to an unmodified ``tensorrt_llm`` package -directory (one request's outputs must match its single-request kernel bit for bit).""" - -import importlib.util import math -import os import pytest import torch @@ -388,110 +383,3 @@ def test_attn_counter_wrap(num_requests, tokens, heads): for start in (2**31 - 16, -16): for got, want in zip(outs[start], outs[0]): assert torch.equal(_bits(got), _bits(want)) - - -# ---------------------------------------------------------------------------------------------------------------- -# Batch-1 identity against the unmodified kernel (K3_BASE_TRTLLM: an unmodified tensorrt_llm package directory). -# ---------------------------------------------------------------------------------------------------------------- - -_base = {} - - -def _base_kernel(): - """The unmodified single-request kernel module, loaded from K3_BASE_TRTLLM.""" - if "mod" not in _base: - path = os.path.join( - os.environ["K3_BASE_TRTLLM"], - "_torch", - "cute_dsl_kernels", - "k3_mla", - "k3_mla_attn_kernel.py", - ) - spec = importlib.util.spec_from_file_location("k3_mla_attn_kernel_base", path) - mod = importlib.util.module_from_spec(spec) - spec.loader.exec_module(mod) - _base["mod"] = mod - return _base["mod"] - - -def _base_attn(q, pool, row_stride, page_row, page_offset, seq_len, out, w_vb=None, gate=None): - """The unmodified op's launch (one request: one 16-byte aligned page-table row, seq_len [1]); q [M, heads * 576].""" - import cuda.bindings.driver as cuda_driver - import cutlass.cute as cute - from cutlass.cute.runtime import from_dlpack - - kern = _base_kernel() - - def arg(t): - return from_dlpack(t.detach(), assumed_align=16).mark_layout_dynamic( - leading_dim=t.dim() - 1 - ) - - heads = q.shape[1] // kern.QK - groups = heads // kern.HEADS - if ("ws", groups) not in _base: - _base["ws", groups] = torch.empty( - groups * kern.CLUSTER * kern.WS_SLOT_ELEMS, dtype=torch.float16, device=q.device - ) - fuse_vb, apply_gate = w_vb is not None, gate is not None - gate_flat = gate.as_strided((gate.numel(),), (1,)) if apply_gate else q.view(-1) - args = (arg(q.view(-1)), arg(pool.view(-1)[: PAGE * row_stride]), arg(page_row.reshape(-1)), - arg(seq_len.reshape(-1)), arg(_base["ws", groups]), arg(out.view(-1)), - arg((w_vb if fuse_vb else q).view(-1)), arg(gate_flat)) # fmt: skip - scalars = (q.shape[0], SCALE * kern.LOG2E, pool.numel() // row_stride, page_offset, GATE_COL0, - gate.stride(0) if apply_gate else 0) # fmt: skip - use_pdl = os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" - stream = cuda_driver.CUstream(torch.cuda.current_stream().cuda_stream) - key = (row_stride, heads, fuse_vb, apply_gate, use_pdl) - fn = _base.get(key) - if fn is None: - fn = _base[key] = cute.compile(kern.k3_mla_attn, *args, *scalars, row_stride, heads, fuse_vb, apply_gate, - use_pdl, stream) # fmt: skip - fn(*args, *scalars, stream) - return out - - -@pytest.mark.skipif( - not os.environ.get("K3_BASE_TRTLLM"), reason="K3_BASE_TRTLLM (unmodified package) not set" -) -@pytest.mark.parametrize("tokens", [8, 1, 2, 4, 7]) -def test_batch1_identity(tokens): - """One request: k3_mla_attn_out and k3_mla_attn_vb_out (plain and gated) bit-identical to the unmodified kernel - for every length in LENGTHS, rows of 576 and interleaved rows of 640 with a page offset, the page-table row given - flat, as a [1, W] row, and at a 4-byte offset (a row of a wider table).""" - from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import op # noqa: F401 - - for i, length in enumerate(LENGTHS): - layers, slot, row_stride = (2, 1, 640) if i % 2 else (1, 0, DQK) - q, pool, table, off, seq_len = _make_case( - 900 + 10 * tokens + i, - 1, - tokens, - H, - row_stride, - layers, - slot, - lens=[max(tokens, length)], - ) - gen = torch.Generator(device="cuda").manual_seed(1900 + 10 * tokens + i) - w_vb = (torch.randn(H, V, LAT, generator=gen, device="cuda") * 0.05).bfloat16() - ag = torch.rand(tokens, GATE_COL0 + H * V, generator=gen, device="cuda").bfloat16() - shifted = torch.zeros(table.shape[1] + 1, dtype=torch.int32, device="cuda") - shifted[1:] = table[0] - q2 = q.view(tokens, -1) - want = _base_attn(q2, pool, row_stride, table[0], off, seq_len, - torch.empty(tokens, H * LAT, dtype=torch.bfloat16, device="cuda")) # fmt: skip - for form, row in ( - ("flat", table[0]), - ("[1, W]", table[:1]), - ("4-byte offset", shifted[1:]), - ): - got = _attn_out(q, pool, row_stride, row, off, seq_len).view(tokens, -1) - assert torch.equal(_bits(got), _bits(want)), f"attn_out L {length} {form}" - for gate in (None, ag): - want = _base_attn(q2, pool, row_stride, table[0], off, seq_len, - torch.empty(tokens, H * V, dtype=torch.bfloat16, device="cuda"), w_vb, gate) # fmt: skip - got = _attn_vb(q, pool, row_stride, table[:1], off, seq_len, w_vb, gate) - assert torch.equal(_bits(got), _bits(want)), ( - f"attn_vb L {length} gate {gate is not None}" - ) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_q.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_q.py index 6dfad07ca19b..f3105d565e83 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_q.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_q.py @@ -16,13 +16,7 @@ KV half into the paged latent cache) at the TP16 shape (6 heads per rank, q_lora 1536, latent 512 + rope 64), M <= 64 tokens: fused_q against the unfused chain (the model's RMSNorm, the q_b GEMM, the k_b bmm) and a reference with the same bf16 roundings, every 8-token chunk bit-identical to the 8-token call on its rows; the KV rows of R requests -x T tokens at each request's positions in a sentinel-filled pool, nothing else written. - -Batch-1 identity against the unmodified kernel: set ``K3_BASE_TRTLLM`` to an unmodified ``tensorrt_llm`` package -directory (one request's M <= 8 outputs, and the pool, must match its kernel bit for bit).""" - -import importlib.util -import os +x T tokens at each request's positions in a sentinel-filled pool, nothing else written.""" import pytest import torch @@ -278,119 +272,3 @@ def test_q_rejects(): assert not op.supports_kv(ag, w_kv, pool, DQK, page_table[:1], seq_len) # 1 row, 2 lengths assert not op.supports_kv(ag, w_kv, pool, DQK, page_table[0], seq_len) # a flat row, 2 lengths assert not op.supports_kv(ag, w_kv, pool, DQK, page_table.long(), seq_len) # int64 table - - -# ---------------------------------------------------------------------------------------------------------------- -# Batch-1 identity against the unmodified kernel (K3_BASE_TRTLLM: an unmodified tensorrt_llm package directory). -# ---------------------------------------------------------------------------------------------------------------- - -_base = {} - - -def _base_kernel(): - """The unmodified kernel module, loaded from K3_BASE_TRTLLM.""" - if "mod" not in _base: - path = os.path.join( - os.environ["K3_BASE_TRTLLM"], - "_torch", - "cute_dsl_kernels", - "k3_mla", - "k3_mla_q_kernel.py", - ) - spec = importlib.util.spec_from_file_location("k3_mla_q_kernel_base", path) - mod = importlib.util.module_from_spec(spec) - spec.loader.exec_module(mod) - _base["mod"] = mod - return _base["mod"] - - -def _base_q(ag, w_qa, w_qb, w_kb, w_kv=None, kv=None): - """The unmodified op's launch (M <= 8, one request): ``kv`` None (k3_mla_q), dict(pool, row_stride, page_row, - page_offset, seq_len) (k3_mla_qkv) or dict(out) (k3_mla_qkv_out). Returns fused_q.""" - import cuda.bindings.driver as cuda_driver - import cutlass.cute as cute - from cutlass.cute.runtime import from_dlpack - - kern = _base_kernel() - - def arg(t): - return from_dlpack(t.detach(), assumed_align=16).mark_layout_dynamic( - leading_dim=t.dim() - 1 - ) - - m, ag_cols = ag.shape - heads = w_kb.shape[0] - out = torch.empty(m, heads * kern.FUSED, dtype=torch.bfloat16, device=ag.device) - if "dummies" not in _base: - _base["dummies"] = (torch.zeros(8, dtype=torch.bfloat16, device="cuda"), - torch.zeros(8, dtype=torch.int32, device="cuda")) # fmt: skip - bf16_dummy, i32_dummy = _base["dummies"] - kv_mode, row_stride, kv_eps, page_offset = 0, kern.FUSED, 0.0, 0 - w_kv_t, pool, page_row, seq_len, kv_out = ( - bf16_dummy, - bf16_dummy, - i32_dummy, - i32_dummy, - bf16_dummy, - ) - if kv is not None: - w_kv_t, kv_eps = w_kv, KV_EPS - if "out" in kv: - kv_mode, kv_out = 2, kv["out"] - else: - kv_mode, row_stride, page_offset = 1, kv["row_stride"], kv["page_offset"] - pool = kv["pool"].view(-1)[: PAGE * row_stride] - page_row, seq_len = kv["page_row"], kv["seq_len"] - args = (arg(w_qb), arg(w_kb.view(heads * kern.LATENT, kern.NOPE)), arg(ag.view(-1)), arg(w_qa), - arg(out.view(-1)), arg(w_kv_t), arg(pool), arg(page_row.reshape(-1)), arg(seq_len.reshape(-1)), - arg(kv_out.view(-1))) # fmt: skip - scalars = (m, m, 0, EPS, kv_eps, page_offset) # M, T (one request), page-table row stride - use_pdl = os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" - stream = cuda_driver.CUstream(torch.cuda.current_stream().cuda_stream) - key = (ag_cols, heads, kv_mode, row_stride, use_pdl) - fn = _base.get(key) - if fn is None: - # trigger_early True, single_hop False, cluster_rms True: the unmodified op's defaults. - fn = _base[key] = cute.compile(kern.k3_mla_q, *args, *scalars, ag_cols, heads, True, False, kv_mode, - row_stride, True, use_pdl, stream) # fmt: skip - fn(*args, *scalars, stream) - return out - - -@pytest.mark.skipif( - not os.environ.get("K3_BASE_TRTLLM"), reason="K3_BASE_TRTLLM (unmodified package) not set" -) -@pytest.mark.parametrize("m", [1, 2, 5, 8]) -def test_batch1_identity(m): - """One request of M <= 8 tokens: k3_mla_q, k3_mla_qkv (fused_q and the whole pool: rows of 576, and interleaved - rows of 640 with a page offset; the page-table row flat and as [1, W]; the new rows crossing a page) and - k3_mla_qkv_out (fused_q and the dense rows) bit-identical to the unmodified kernel.""" - from tensorrt_llm._torch.cute_dsl_kernels.k3_mla import op # noqa: F401 - - weights = _weights(700 + m) - ag = _ag(800 + m, m) - assert torch.equal(_bits(_q(ag, *weights[:3])), _bits(_base_q(ag, *weights[:3]))) - for layers, slot, row_stride in ((1, 0, DQK), (2, 1, 640)): - pool, table, seq_len = _kv_case(900 + m, 1, m, row_stride, layers, lens=[64 * 9 + 3]) - want_pool = pool.clone() - kv = dict( - pool=want_pool, - row_stride=row_stride, - page_row=table[0], - page_offset=slot, - seq_len=seq_len, - ) - want = _base_q(ag, *weights[:3], weights[3], kv) - for form, row in (("flat", table[0]), ("[1, W]", table[:1])): - got_pool = pool.clone() - got = _qkv(ag, weights, got_pool, row_stride, row, slot, seq_len) - assert torch.equal(_bits(got), _bits(want)), f"qkv fused_q rows of {row_stride} {form}" - assert torch.equal(_bits(got_pool), _bits(want_pool)), ( - f"qkv pool rows of {row_stride} {form}" - ) - want_rows = torch.empty(m, DQK, dtype=torch.bfloat16, device="cuda") - want = _base_q(ag, *weights[:3], weights[3], dict(out=want_rows)) - got_rows = torch.empty_like(want_rows) - got = torch.ops.trtllm.k3_mla_qkv_out(ag, weights[0], EPS, weights[1], weights[2], weights[3], KV_EPS, got_rows, - True) # fmt: skip - assert torch.equal(_bits(got), _bits(want)) and torch.equal(_bits(got_rows), _bits(want_rows)) From fe734863e6db3b56d886c51d9e4dcc2ad99215f5 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 19:04:50 -0700 Subject: [PATCH 032/161] [None][feat] modeling_v2 Kimi K3 target: classify each step for the decode kernels The target decides once per step, on the host, which Kimi K3 decode kernels take the step (DecodeStep, from the K3 stack's step predicate): - small: at most 8 tokens, context requests included (the token-count kernels); - decode: R <= 8 generation requests of the same T <= 8 tokens and no context request (the request-aware kernels: MLA attention and its KV store, the KDA verify, the drafter's attention); - wide: a decode step above one token tile, which keeps the decode layout's MoE head and tail on M-general ops. decode_step replaces the placeholder _step_path. Every other step runs the generic path, and so does every step until the fused decode path exists. The predicate reads num_generations and the host sequence lengths, now declared in REQUIRED_ENGINE_FIELDS. Signed-off-by: Vasanth Sabavat --- .../modeling.py | 136 +++++++++++++----- .../test_modeling_v2_kimi_k3_decode_step.py | 69 +++++++++ 2 files changed, 169 insertions(+), 36 deletions(-) create mode 100644 tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_step.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index e02a387828b3..dc74b550a372 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -12,15 +12,24 @@ `moe_expert_parallel_size: 4`, no attention data parallelism, all 16 GPUs in one NVLink domain. Attention is head-split (6 MLA query heads per rank), and every rank holds a quarter of the width of a quarter of the experts. -**Each step takes one of two paths, chosen on the host from the step's shape** (`_step_path`): - -* **The fused decode path**: pure decode steps of at most 8 tokens without speculation, or at most 64 with DSpark - (8 requests x 1 + 7 drafts), on the K3 decode kernels' catalog entries. The state those kernels share (MNNVL - workspace, sandwich and MoE Lamport buffers, KDA / MLA scratch) lives in typed objects this target creates - collectively in `post_load_weights`, before any graph capture. Until those entries exist `_fused_decode` stays - None, and every step takes the generic path. -* **The generic path**: prefill, mixed steps, and decode steps above those bounds, on the built-in Kimi K3 text - model, whose modules and ops have no catalog entries yet. `UNCERTIFIED_GENERIC_CALLS` names them. +**Each step is classified once, on the host, from its shape** (`decode_step`, a `DecodeStep` or None), and the +classification decides which kernels each module runs: + +* **small**: at most 8 tokens (`DECODE_MAX_TOKENS`, one token tile), context requests included. The token-count + kernels take it: decode GEMVs, the MoE front and routed experts, the sandwiches, the embedding and residual + epilogues. +* **decode**: a pure decode step of R <= 8 generation requests with the same T <= 8 tokens each (one token without + speculation, 1 + 7 drafts with DSpark) and no context request, so at most 64 tokens. The request-aware kernels + take it: MLA attention and its KV store, the KDA verify, the drafter's attention. +* **wide**: a decode step of more than one token tile (DSpark verify of several requests). It keeps the decode + layout's MoE head and tail, on M-general ops. + +Every other step (prefill, mixed steps, decode steps above those bounds) runs the **generic path**: the built-in Kimi +K3 text model, whose modules and ops have no catalog entries yet. `UNCERTIFIED_GENERIC_CALLS` names them. The **fused +decode path** runs the steps `decode_step` classifies on the K3 decode kernels' catalog entries. The state those +kernels share (MNNVL workspace, sandwich and MoE Lamport buffers, KDA / MLA scratch) lives in typed objects this +target creates collectively in `post_load_weights`, before any graph capture. Until those entries exist +`_fused_decode` stays None, and every step takes the generic path. **What this target asserts rather than adapts**: SM 10.0; the topology above; bf16 weights and a bf16 KV pool; tokens_per_block 64 (the MLA generation kernels K3's 96 heads reach exist only at 64); the V2 hybrid KV / state @@ -35,6 +44,7 @@ """ import copy +from dataclasses import dataclass from typing import Any, Literal, Optional import torch @@ -72,8 +82,10 @@ REQUIRED_ENGINE_FIELDS = { "attn_metadata": ( "num_contexts", + "num_generations", "num_seqs", "num_tokens", + "seq_lens", "tokens_per_block", "kv_cache_manager", ), @@ -86,10 +98,11 @@ "tensorrt_llm._torch.models.modeling_kimi_linear.KimiLinearForCausalLM", ) -# The fused decode path's token bounds per step: one token per request without speculation (the K3 decode kernels -# are built for up to 8 rows), and 1 + 7 drafts per request with DSpark at batch 8. -_FUSED_MAX_TOKENS = 8 -_FUSED_MAX_TOKENS_SPEC = 64 +# The K3 decode kernels' bounds: the token-count kernels take one tile of DECODE_MAX_TOKENS rows; the request-aware +# kernels take MAX_REQUESTS generation requests of at most MAX_TOKENS_PER_REQUEST tokens (1 + 7 drafts with DSpark). +DECODE_MAX_TOKENS = 8 +MAX_REQUESTS = 8 +MAX_TOKENS_PER_REQUEST = 8 # The MLA generation kernels for K3's 96 query heads exist only at a 64-token page (the built-in model's own # get_model_defaults sets it for the same reason). @@ -124,6 +137,67 @@ def _text_model_config(model_config: ModelConfig) -> ModelConfig: return text +@dataclass(frozen=True) +class DecodeStep: + """A step the Kimi K3 decode kernels take: ``num_tokens`` rows and, on a pure decode step, ``num_requests`` + generation requests of ``tokens_per_request`` tokens each (None on a step with context requests).""" + + num_tokens: int + num_requests: Optional[int] = None + tokens_per_request: Optional[int] = None + + @property + def small(self) -> bool: + """Whether the step fits one token tile of the token-count kernels.""" + return self.num_tokens <= DECODE_MAX_TOKENS + + @property + def decode(self) -> bool: + """Whether the step is a pure decode step the request-aware kernels take.""" + return self.num_requests is not None + + @property + def wide(self) -> bool: + """Whether the step is a pure decode step of more than one token tile: its token-count work keeps the decode + layout's MoE head and tail, on M-general ops.""" + return self.decode and not self.small + + +def decode_step(attn_metadata: AttentionMetadata, num_tokens: int) -> Optional[DecodeStep]: + """The step's shape if any Kimi K3 decode kernel takes it, else None (the generic path runs). + + ``num_tokens`` is the step's token count (the rows of the model input). Read on the host from per-step integers + and the host copy of the sequence lengths only. A CUDA graph is captured per decode batch shape, and every input + here is fixed by that shape, so a captured step and its replays are classified alike. + """ + if num_tokens <= 0: + return None + requests = _decode_requests(attn_metadata, num_tokens) + if requests is not None: + return DecodeStep(num_tokens, requests, num_tokens // requests) + if num_tokens <= DECODE_MAX_TOKENS: + return DecodeStep(num_tokens) + return None + + +def _decode_requests(attn_metadata: AttentionMetadata, num_tokens: int) -> Optional[int]: + """R when the step is R <= 8 generation requests of the same T <= 8 tokens and no context request, else None.""" + if attn_metadata.num_contexts != 0: + return None + requests = attn_metadata.num_generations + if not 0 < requests <= MAX_REQUESTS or num_tokens % requests != 0: + return None + tokens = num_tokens // requests + if not 0 < tokens <= MAX_TOKENS_PER_REQUEST: + return None + seq_lens = getattr(attn_metadata, "seq_lens", None) + if seq_lens is not None and seq_lens.device.type == "cpu": + lens = seq_lens[:requests] + if lens.numel() != requests or bool((lens != tokens).any()): + return None + return requests + + def _check_construction(model_config: ModelConfig) -> None: """The settings this target is built for that are fixed before the first step.""" capability = torch.cuda.get_device_capability() @@ -215,17 +289,6 @@ def _check_step_contract(self, attn_metadata: AttentionMetadata) -> None: ) self._step_checked = True - def _step_path(self, attn_metadata: AttentionMetadata, spec_metadata) -> str: - """`"fused"` for a pure decode step within the fused path's bounds, else `"generic"`. - - Read on the host from per-step integers only. A CUDA graph is captured per decode batch shape, and every - input here is fixed by that shape, so a captured step and its replays take the same path. - """ - if attn_metadata.num_contexts: - return "generic" - bound = _FUSED_MAX_TOKENS if spec_metadata is None else _FUSED_MAX_TOKENS_SPEC - return "fused" if attn_metadata.num_tokens <= bound else "generic" - def forward( self, attn_metadata: AttentionMetadata, @@ -243,18 +306,19 @@ def forward( ) if not self._step_checked: self._check_step_contract(attn_metadata) - if ( - self._fused_decode is not None - and self._step_path(attn_metadata, spec_metadata) == "fused" - ): - return self._fused_decode( - attn_metadata=attn_metadata, - input_ids=input_ids, - position_ids=position_ids, - spec_metadata=spec_metadata, - resource_manager=resource_manager, - **kwargs, - ) + if self._fused_decode is not None: + rows = input_ids if input_ids is not None else inputs_embeds + step = None if rows is None else decode_step(attn_metadata, rows.shape[0]) + if step is not None: + return self._fused_decode( + step, + attn_metadata=attn_metadata, + input_ids=input_ids, + position_ids=position_ids, + spec_metadata=spec_metadata, + resource_manager=resource_manager, + **kwargs, + ) return super().forward( attn_metadata, input_ids, diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_step.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_step.py new file mode 100644 index 000000000000..345f803fac8f --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_step.py @@ -0,0 +1,69 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""How the Kimi K3 target classifies a step for its decode kernels (host-side, no GPU): ``decode_step`` of +``kimi_k3_mxfp4__sm_100__tp16_moetp4ep4``.""" + +import types + +import pytest +import torch + +from tensorrt_llm._torch._experimental.modeling_v2.models.kimi_k3_vl.kimi_k3_mxfp4__sm_100__tp16_moetp4ep4.modeling import ( # noqa: E501 + DecodeStep, + decode_step, +) + +pytestmark = pytest.mark.cpu_only + + +def _metadata(seq_lens, num_contexts=0): + return types.SimpleNamespace( + num_contexts=num_contexts, + num_generations=len(seq_lens) - num_contexts, + seq_lens=torch.tensor(seq_lens, dtype=torch.int32), + ) + + +@pytest.mark.parametrize( + "requests,tokens", + [(1, 1), (1, 8), (8, 1), (2, 4), (4, 2), (3, 1), (5, 1), (2, 8), (4, 8), (8, 8), (8, 2)], +) +def test_pure_decode_steps(requests, tokens): + """R <= 8 generation requests of the same T <= 8 tokens: a decode step, small up to one token tile, wide above.""" + step = decode_step(_metadata([tokens] * requests), requests * tokens) + assert step == DecodeStep(requests * tokens, requests, tokens) + assert step.decode + assert step.small == (requests * tokens <= 8) + assert step.wide == (requests * tokens > 8) + + +@pytest.mark.parametrize( + "seq_lens,num_contexts,rows", + [ + ([5], 1, 5), # a short prefill + ([3, 1], 1, 4), # a short mixed step + ([4, 4], 0, 4), # rows that are not the step's tokens + ([3, 2], 0, 5), # a ragged decode step + ], +) +def test_short_steps_are_small_only(seq_lens, num_contexts, rows): + """Other steps of at most 8 tokens take the token-count kernels only.""" + step = decode_step(_metadata(seq_lens, num_contexts), rows) + assert step == DecodeStep(rows) + assert step.small and not step.decode and not step.wide + + +@pytest.mark.parametrize( + "seq_lens,num_contexts,rows", + [ + ([1024], 1, 1024), # prefill + ([100, 8], 1, 108), # mixed step + ([1] * 9, 0, 9), # more requests than the kernels take + ([16], 0, 16), # more tokens per request than the kernels take + ([8, 4], 0, 12), # ragged decode step + ([4, 4], 0, 16), # rows that are not the step's tokens + ([], 0, 0), # empty step + ], +) +def test_other_steps_take_the_generic_path(seq_lens, num_contexts, rows): + assert decode_step(_metadata(seq_lens, num_contexts), rows) is None From abb2564b563e9a089580eb662ae27a8c05ab6e46 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 19:09:32 -0700 Subject: [PATCH 033/161] [None][feat] modeling_v2 Kimi K3 target: its own text model, on the generic path The target builds its own text model instead of the built-in one: the Kimi K3 decoder (KimiLinearModel, KimiLinearDecoderLayer, the MoE and MLA runtimes, the router gate, the attention-residual helpers) now lives in the target's modeling.py, as the K3 stack carries it on its generic path. That path is the built-in model's: every definition is AST-identical to modeling_kimi_linear.py's, except one import that names the same module absolutely. The fused decode branches go into these classes in later commits; until then every step computes exactly as before. The registration shell keeps the built-in causal LM's checkpoint load and engine hooks, which walk the model by its module names, and constructs SpecDecOneEngineForCausalLM directly with the target's model (no helix, no fp8 weight-read conversion). UNCERTIFIED_GENERIC_CALLS now lists the stock modules the text model calls. Signed-off-by: Vasanth Sabavat --- .../modeling.py | 1574 ++++++++++++++++- 1 file changed, 1566 insertions(+), 8 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index dc74b550a372..1b7a5cacbd5e 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -24,9 +24,11 @@ * **wide**: a decode step of more than one token tile (DSpark verify of several requests). It keeps the decode layout's MoE head and tail, on M-general ops. -Every other step (prefill, mixed steps, decode steps above those bounds) runs the **generic path**: the built-in Kimi -K3 text model, whose modules and ops have no catalog entries yet. `UNCERTIFIED_GENERIC_CALLS` names them. The **fused -decode path** runs the steps `decode_step` classifies on the K3 decode kernels' catalog entries. The state those +Every other step (prefill, mixed steps, decode steps above those bounds) runs the **generic path**: this target's +text model (`KimiLinearModel` below: decoder layers, attention residuals, the MLA / KDA / MoE runtimes), computed +exactly as the built-in Kimi K3 text model computes it, on stock modules and ops that have no catalog entries yet. +`UNCERTIFIED_GENERIC_CALLS` names them. The **fused decode path** runs the steps `decode_step` classifies on the K3 +decode kernels' catalog entries. The state those kernels share (MNNVL workspace, sandwich and MoE Lamport buffers, KDA / MLA scratch) lives in typed objects this target creates collectively in `post_load_weights`, before any graph capture. Until those entries exist `_fused_decode` stays None, and every step takes the generic path. @@ -43,21 +45,49 @@ checkpoint, and SA. The worker and its kernels stay upstream code; this target does not own a worker. """ +from __future__ import annotations + import copy +import math +import os from dataclasses import dataclass -from typing import Any, Literal, Optional +from typing import TYPE_CHECKING, Any, Callable, Dict, Literal, NamedTuple, Optional, Tuple import torch +from torch import nn -from tensorrt_llm._torch.attention.backends.interface import AttentionMetadata +from tensorrt_llm._torch.attention.backends import AttentionMetadata +from tensorrt_llm._torch.distributed import AllReduce from tensorrt_llm._torch.model_config import ModelConfig from tensorrt_llm._torch.models.modeling_kimi_linear import KimiLinearForCausalLM -from tensorrt_llm._torch.models.modeling_utils import register_auto_model +from tensorrt_llm._torch.models.modeling_speculative import SpecDecOneEngineForCausalLM +from tensorrt_llm._torch.models.modeling_utils import DecoderModel, register_auto_model +from tensorrt_llm._torch.modules.gated_mlp import GatedMLP +from tensorrt_llm._torch.modules.kimi_kda import KimiKDALinearAttention +from tensorrt_llm._torch.modules.multi_stream_utils import maybe_execute_in_parallel +from tensorrt_llm._torch.modules.rms_norm import RMSNorm +from tensorrt_llm._torch.modules.situ import SituAndMul +from tensorrt_llm._torch.moe.fused_moe import ( + ConfigurableMoE, + SiTuActivation, + TRTLLMGenFusedMoE, + create_moe, +) +from tensorrt_llm._torch.moe.fused_moe.interface import MoESchedulerKind +from tensorrt_llm._torch.moe.fused_moe.routing import DeepSeekV3MoeRoutingMethod from tensorrt_llm._torch.pyexecutor.kv_cache.mamba_cache_manager import MambaHybridCacheManagerV2 +from tensorrt_llm._torch.utils import AuxStreamType from tensorrt_llm.functional import AllReduceStrategy +from tensorrt_llm.logger import logger +from tensorrt_llm.mapping import Mapping +from tensorrt_llm.models.modeling_utils import QuantAlgo, QuantConfig from . import weights as _weights +if TYPE_CHECKING: + from transformers import PretrainedConfig + + # The GPU architecture this target IS. Routing will not send another one here, but a direct instantiation could, # and the certification is per arch. _SM = (10, 0) @@ -95,7 +125,17 @@ #: Calls the generic path makes outside the catalog, declared so they are not consumed silently. A call leaves this #: list when a catalog entry replaces it. UNCERTIFIED_GENERIC_CALLS = ( + # The checkpoint load and the engine hooks this target inherits. "tensorrt_llm._torch.models.modeling_kimi_linear.KimiLinearForCausalLM", + # The text model's stock modules. + "tensorrt_llm._torch.modules.kimi_kda.KimiKDALinearAttention", + "tensorrt_llm._torch.modules.kimi_k3_mla.KimiK3MLAAttention", + "tensorrt_llm._torch.moe.fused_moe.create_moe", + "tensorrt_llm._torch.modules.gated_mlp.GatedMLP", + "tensorrt_llm._torch.modules.situ.SituAndMul", + "tensorrt_llm._torch.modules.rms_norm.RMSNorm", + "tensorrt_llm._torch.distributed.AllReduce", + "tensorrt_llm._torch.modules.multi_stream_utils.maybe_execute_in_parallel", ) # The K3 decode kernels' bounds: the token-count kernels take one tile of DECODE_MAX_TOKENS rows; the request-aware @@ -111,6 +151,1499 @@ _LANG_PREFIX = "language_model." +# ---------------------------------------------------------------------------------------------------------------------- +# The text model: Kimi K3's decoder (93 layers: KDA / MLA attention, attention residuals, the dense layer-0 MLP and +# the latent MoE), its generic path. +# ---------------------------------------------------------------------------------------------------------------------- + +# A/B escape hatch: restore nn.Linear for the K3 latent MoE projections +# instead of the min-latency fused GEMM op (read once at import). +_K3_DISABLE_MIN_LATENCY_LATENT_PROJ = ( + os.environ.get("TLLM_K3_DISABLE_MIN_LATENCY_LATENT_PROJ", "0") == "1" +) + + +# Identity-RoPE table positions for the MLA backends. K3 is NoPE (the table +# holds cos=1/sin=0), but the chunked-context path indexes the table by +# absolute position, so it must cover max_position_embeddings (~512MB per +# backend for the 1M-position checkpoint); a smaller table is read out of +# bounds. KIMI_K3_MLA_MAX_POSITIONS overrides the size for short-context +# deployments. +_KIMI_K3_MLA_MAX_POSITIONS_ENV = "KIMI_K3_MLA_MAX_POSITIONS" + + +class KimiK3MoEGate(nn.Module): + """Kimi K3 gate weights and routing method for ``ConfigurableMoE``.""" + + def __init__( + self, + config: Any, + *, + logits_gemm_dtype: torch.dtype | None = None, + device: torch.device | None = None, + ) -> None: + super().__init__() + self.config = config + self.top_k = config.num_experts_per_token + self.num_experts = config.num_experts + self.routed_scaling_factor = config.routed_scaling_factor + self.moe_router_activation_func = config.moe_router_activation_func + self.num_expert_group = getattr(config, "num_expert_group", 1) + self.topk_group = getattr(config, "topk_group", 1) + self.moe_renormalize = config.moe_renormalize + self.gating_dim = config.hidden_size + + assert self.moe_router_activation_func in ("sigmoid", "softmax"), ( + "K3 MoE gate supports sigmoid or softmax scoring only" + ) + + # The checkpoint stores the gate weight in bf16. Storing it in bf16 + # permits the single bf16xbf16 router GEMM while retaining fp32 output. + weight_dtype = logits_gemm_dtype or torch.float32 + self.weight = nn.Parameter( + torch.empty((self.num_experts, self.gating_dim), dtype=weight_dtype, device=device) + ) + self.e_score_correction_bias = nn.Parameter( + torch.empty(self.num_experts, dtype=torch.float32, device=device) + ) + + def compute_logits(self, hidden_states: torch.Tensor) -> torch.Tensor: + """Compute fp32 routing logits shaped ``[num_tokens, num_experts]``.""" + hidden_2d = hidden_states.reshape(-1, self.gating_dim) + if self.weight.dtype == torch.bfloat16 and hidden_2d.dtype == torch.bfloat16: + return torch.ops.trtllm.dsv3_router_gemm_op( + hidden_2d.contiguous(), + self.weight.t(), + bias=None, + out_dtype=torch.float32, + ) + return torch.nn.functional.linear( + hidden_2d.type(torch.float32), + self.weight.type(torch.float32), + None, + ) + + @property + def routing_method(self) -> DeepSeekV3MoeRoutingMethod: + """Return the shared DeepSeek-V3 router used by ``ConfigurableMoE``.""" + if self.moe_router_activation_func != "sigmoid": + raise ValueError("Kimi K3 ConfigurableMoE routing requires sigmoid scores.") + if not self.moe_renormalize: + raise ValueError( + "Kimi K3 ConfigurableMoE routing requires top-k weight renormalization." + ) + return DeepSeekV3MoeRoutingMethod( + top_k=self.top_k, + n_group=self.num_expert_group, + topk_group=self.topk_group, + routed_scaling_factor=self.routed_scaling_factor, + callable_e_score_correction_bias=lambda: self.e_score_correction_bias, + is_fused=True, + ) + + +class KimiK3RMSNorm(nn.Module): + """RMSNorm matching the Kimi checkpoint implementation's rounding.""" + + def __init__( + self, + hidden_size: int, + eps: float = 1e-6, + dtype: torch.dtype = torch.float32, + device: Optional[torch.device] = None, + ) -> None: + super().__init__() + self.weight = nn.Parameter(torch.ones(hidden_size, dtype=dtype, device=device)) + self.eps = eps + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + input_dtype = hidden_states.dtype + hidden_states_float = hidden_states.to(torch.float32) + variance = hidden_states_float.pow(2).mean(-1, keepdim=True) + hidden_states_float = hidden_states_float * torch.rsqrt(variance + self.eps) + return self.weight * hidden_states_float.to(input_dtype) + + +def _resolve_kimi_situ_betas(cfg: Any) -> tuple[float, float]: + """Return the finite SiTu betas required by the routed-expert kernels.""" + config_situ_beta = getattr(cfg, "activation_situ_beta", None) + situ_beta = 1.0 if config_situ_beta is None else config_situ_beta + situ_linear_beta = getattr(cfg, "activation_situ_linear_beta", None) + if situ_linear_beta is None: + raise ValueError( + "Kimi K3 routed SiTu experts require activation_situ_linear_beta; " + "None means an identity linear branch that the fused kernels cannot represent." + ) + if situ_beta <= 0 or situ_linear_beta <= 0: + raise ValueError( + f"Kimi K3 SiTu betas must be positive; got {situ_beta} and {situ_linear_beta}." + ) + return float(situ_beta), float(situ_linear_beta) + + +def _get_text_config(pretrained_config: "PretrainedConfig"): + """Return the Kimi text config, unwrapping a composite kimi_k3 config.""" + if getattr(pretrained_config, "model_type", None) == "kimi_k3" or ( + not hasattr(pretrained_config, "linear_attn_config") + and hasattr(pretrained_config, "text_config") + ): + return pretrained_config.text_config + return pretrained_config + + +def _is_kda_layer(cfg, layer_idx: int) -> bool: + return (layer_idx + 1) in cfg.linear_attn_config["kda_layers"] + + +def _is_mla_layer(cfg, layer_idx: int) -> bool: + return (layer_idx + 1) in cfg.linear_attn_config["full_attn_layers"] + + +KIMI_K3_AUX_ATTN_RES_STREAM_ENV = "KIMI_K3_AUX_ATTN_RES_STREAM" + + +_AUX_ATTN_RES_STREAM_ENABLED = os.environ.get(KIMI_K3_AUX_ATTN_RES_STREAM_ENV, "1") == "1" + + +KIMI_K3_FUSED_ATTN_RES_ENV = "KIMI_K3_FUSED_ATTN_RES" + + +_FUSED_ATTN_RES_ENABLED = os.environ.get(KIMI_K3_FUSED_ATTN_RES_ENV, "1") == "1" + + +KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS_ENV = "KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS" + + +KIMI_K3_ATTN_RES_TOPOLOGY_ENV = "KIMI_K3_ATTN_RES_TOPOLOGY" + + +_ATTN_RES_TOPOLOGIES = ("per_token", "persistent", "split") + + +# Turning the port on should not require knowing which of the three topologies +# is the right one: "1" selects the measured policy (``split``), so enabling the +# feature and choosing the policy are one action through one variable. +_ATTN_RES_TOPOLOGY_ON = "1" + + +def _read_attn_res_topology() -> str: + """Default ``per_token``: the persistent kernel is opt-in. + + ``1`` is the only accepted on-value and resolves to ``split``. The named + topologies stay available for measurement: ``persistent`` uses the + persistent kernel at every shape it implements, ``per_token`` at none. + """ + raw = os.environ.get(KIMI_K3_ATTN_RES_TOPOLOGY_ENV, "per_token") + if raw == _ATTN_RES_TOPOLOGY_ON: + return "split" + if raw not in _ATTN_RES_TOPOLOGIES: + # Loudly, for the same reason as the token ceiling below: a mistyped A/B + # arm that silently fell back to the default would measure one side + # twice and report no difference. + raise ValueError( + f"{KIMI_K3_ATTN_RES_TOPOLOGY_ENV} must be one of " + f"{_ATTN_RES_TOPOLOGIES} or {_ATTN_RES_TOPOLOGY_ON!r} " + f"(which means 'split'), got {raw!r}" + ) + return raw + + +def _read_fused_attn_res_max_tokens() -> int: + """Resolved after the topology, because its default follows it. + + With the persistent port off -- the default -- the ceiling is 1, the + pre-existing gate: the fused epilogue is taken at the single-token decode + shape and nowhere else. Enabling the port raises it to 32, the top of the + measured range, so that "off" keeps meaning "unchanged". + """ + default = "1" if _ATTN_RES_TOPOLOGY == "per_token" else "32" + raw = os.environ.get(KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS_ENV, default) + try: + value = int(raw) + except ValueError: + # Failing loudly matters more than usual here: a mistyped A/B arm that + # silently fell back to the default would measure the candidate twice + # and report no difference. + raise ValueError( + f"{KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS_ENV} must be a positive integer, got {raw!r}" + ) from None + if value < 1: + raise ValueError(f"{KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS_ENV} must be >= 1, got {value}") + return value + + +_ATTN_RES_TOPOLOGY = _read_attn_res_topology() + + +_FUSED_ATTN_RES_MAX_TOKENS = _read_fused_attn_res_max_tokens() + + +def _persistent_attn_res_applicable(M: int, H: int, N: int) -> bool: + """Shape gate for the persistent kernel: H == 7168 and 2 <= N <= 9. + + No token ceiling: the persistent grid is sized by the SM count, not by the + token count, so prefill is the case it exists for. + """ + del M # deliberately unused; see above + return H == 7168 and 2 <= N <= 9 + + +def _use_persistent_attn_res(M: int, H: int, N: int) -> bool: + """Pick between the two fused kernels for this call site. + + ``persistent`` takes the persistent kernel at every shape it implements; + ``split`` takes it only above ``KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS`` tokens, + which stands in for the prefill/decode boundary. Shapes the persistent + kernel does not implement fall through to the caller's existing gate and + land on the unfused path. + """ + if not _persistent_attn_res_applicable(M, H, N): + return False + if _ATTN_RES_TOPOLOGY == "persistent": + return True + return _ATTN_RES_TOPOLOGY == "split" and M > _FUSED_ATTN_RES_MAX_TOKENS + + +def _apply_attn_res_fused( + prefix_sum: torch.Tensor, block_residual: torch.Tensor, proj: nn.Linear, norm: KimiK3RMSNorm +) -> Optional[torch.Tensor]: + """Fused attn_res via the in-tree ``trtllm::attn_res_fwd`` op. + + Returns ``None`` when the call falls outside the fused kernel's + contract (dtype/shape/arch) so the caller can use the exact fp32 reference + instead. ``block_residual`` is kept in the kernel-native ``[K, M, H]`` + layout. Candidate order matches the reference: snapshots first, the + running prefix sum last. + """ + if ( + prefix_sum.dtype is not torch.bfloat16 + or not prefix_sum.is_cuda + or not block_residual.is_cuda + ): + return None + M, H = prefix_sum.shape + K = int(block_residual.shape[0]) + if K + 1 > 12 or M > 16384 or not (4096 <= H <= 8192 and H % 1024 == 0): + return None + try: + attn_res_op = torch.ops.trtllm.attn_res_fwd + except (AttributeError, RuntimeError): + return None + layer_kernel = prefix_sum.reshape(M, 1, H).contiguous() + block_kernel = block_residual.reshape(K, M, 1, H).contiguous() + output, _rsigma, _probs, _logits = attn_res_op( + layer_kernel, + block_kernel, + proj.weight.reshape(-1).to(torch.bfloat16).contiguous(), + norm.weight.to(torch.bfloat16).contiguous(), + float(norm.eps), + ) + return output.reshape(M, H) + + +def _rms_norm_eps(norm: nn.Module) -> float: + if hasattr(norm, "eps"): + return float(norm.eps) + return float(norm.variance_epsilon) + + +def _note_attn_res_fusion(site: str, fused: bool, M: int, H: int, N: int) -> None: + """Report whether the fused path was actually reached, once per shape. + + ``_FUSED_ATTN_RES_ENABLED`` only says the feature is switched on, not that + the shape gate let the call through, and a rejected call looks exactly like + a disabled one in the logs. Emitted at debug level: it is once per distinct + shape, not once per process, so it is a diagnostic rather than a summary. + """ + logger.debug_once( + f"Kimi K3 attn-res fusion [{site}]: " + f"{'FUSED' if fused else 'fallback'} (M={M}, H={H}, N={N})", + key=f"kimi_k3_attn_res_fusion_{site}_{fused}_{M}_{H}_{N}", + ) + + +def _apply_attn_res_rmsnorm_fused( + prefix_sum: torch.Tensor, + block_residual: torch.Tensor, + proj: nn.Linear, + norm: KimiK3RMSNorm, + output_norm: nn.Module, +) -> Optional[torch.Tensor]: + """Fuse attention-residual mixing with its immediately following norm.""" + if ( + prefix_sum.dtype is not torch.bfloat16 + or not prefix_sum.is_cuda + or not block_residual.is_cuda + ): + return None + M, H = prefix_sum.shape + K = int(block_residual.shape[0]) + N = K + 1 + # The fused path is taken for M <= _FUSED_ATTN_RES_MAX_TOKENS, H == 7168 and + # N <= 12, which is the measured window; larger token counts have not been + # measured and fall back to the unfused add + attn_res_fwd + RMSNorm path. + if _use_persistent_attn_res(M, H, N): + try: + persistent_op = torch.ops.trtllm.attn_res_add_rmsnorm_persistent_fwd + except (AttributeError, RuntimeError): + return None + _, output = persistent_op( + prefix_sum.reshape(M, 1, H).contiguous(), + None, + block_residual.reshape(K, M, 1, H).contiguous(), + proj.weight.reshape(-1).to(torch.bfloat16).contiguous(), + norm.weight.to(torch.bfloat16).contiguous(), + output_norm.weight.to(torch.bfloat16).contiguous(), + float(norm.eps), + _rms_norm_eps(output_norm), + ) + _note_attn_res_fusion("attn_res+norm/persistent", True, M, H, N) + return output.reshape(M, H) + + if M > _FUSED_ATTN_RES_MAX_TOKENS or H != 7168 or N > 12: + _note_attn_res_fusion("attn_res+norm", False, M, H, N) + return None + try: + attn_res_rmsnorm_op = torch.ops.trtllm.attn_res_rmsnorm_fwd + except (AttributeError, RuntimeError): + return None + layer_kernel = prefix_sum.reshape(M, 1, H).contiguous() + block_kernel = block_residual.reshape(K, M, 1, H).contiguous() + output = attn_res_rmsnorm_op( + layer_kernel, + block_kernel, + proj.weight.reshape(-1).to(torch.bfloat16).contiguous(), + norm.weight.to(torch.bfloat16).contiguous(), + output_norm.weight.to(torch.bfloat16).contiguous(), + float(norm.eps), + _rms_norm_eps(output_norm), + ) + _note_attn_res_fusion("attn_res+norm", True, M, H, N) + return output.reshape(M, H) + + +def _apply_attn_res_add_rmsnorm_fused( + prefix_sum: torch.Tensor, + addend: torch.Tensor, + block_residual: torch.Tensor, + proj: nn.Linear, + norm: KimiK3RMSNorm, + output_norm: nn.Module, +) -> Optional[Tuple[torch.Tensor, torch.Tensor]]: + """Fuse ``prefix_sum + addend``, attention-residual, and trailing norm. + + The production residual add produces a BF16 tensor that remains live across + the following MLP. The kernel therefore returns that materialized, + BF16-rounded prefix sum alongside the normalized attention-residual output, + while avoiding a separate add launch and a re-read of the intermediate by + attention-residual selection. + """ + if ( + prefix_sum.dtype is not torch.bfloat16 + or addend.dtype is not torch.bfloat16 + or not prefix_sum.is_cuda + or not addend.is_cuda + or not block_residual.is_cuda + or prefix_sum.shape != addend.shape + ): + return None + M, H = prefix_sum.shape + K = int(block_residual.shape[0]) + N = K + 1 + # Same measured window as _apply_attn_res_rmsnorm_fused above. + if _use_persistent_attn_res(M, H, N): + try: + persistent_op = torch.ops.trtllm.attn_res_add_rmsnorm_persistent_fwd + except (AttributeError, RuntimeError): + return None + updated_prefix_sum, output = persistent_op( + prefix_sum.reshape(M, 1, H).contiguous(), + addend.reshape(M, 1, H).contiguous(), + block_residual.reshape(K, M, 1, H).contiguous(), + proj.weight.reshape(-1).to(torch.bfloat16).contiguous(), + norm.weight.to(torch.bfloat16).contiguous(), + output_norm.weight.to(torch.bfloat16).contiguous(), + float(norm.eps), + _rms_norm_eps(output_norm), + ) + _note_attn_res_fusion("add+attn_res+norm/persistent", True, M, H, N) + return updated_prefix_sum.reshape(M, H), output.reshape(M, H) + + if M > _FUSED_ATTN_RES_MAX_TOKENS or H != 7168 or N > 12: + _note_attn_res_fusion("add+attn_res+norm", False, M, H, N) + return None + try: + attn_res_add_rmsnorm_op = torch.ops.trtllm.attn_res_add_rmsnorm_fwd + except (AttributeError, RuntimeError): + return None + layer_kernel = prefix_sum.reshape(M, 1, H).contiguous() + addend_kernel = addend.reshape(M, 1, H).contiguous() + block_kernel = block_residual.reshape(K, M, 1, H).contiguous() + updated_prefix_sum, output = attn_res_add_rmsnorm_op( + layer_kernel, + addend_kernel, + block_kernel, + proj.weight.reshape(-1).to(torch.bfloat16).contiguous(), + norm.weight.to(torch.bfloat16).contiguous(), + output_norm.weight.to(torch.bfloat16).contiguous(), + float(norm.eps), + _rms_norm_eps(output_norm), + ) + _note_attn_res_fusion("add+attn_res+norm", True, M, H, N) + return updated_prefix_sum.reshape(M, H), output.reshape(M, H) + + +def _apply_attn_res( + prefix_sum: torch.Tensor, block_residual: torch.Tensor, proj: nn.Linear, norm: KimiK3RMSNorm +) -> torch.Tensor: + """Exact port of HF ``modeling_kimi._apply_attn_res`` (fp32 math). + + prefix_sum: ``[num_tokens, hidden_size]`` + block_residual: ``[num_snapshots, num_tokens, hidden_size]`` + + Unless ``KIMI_K3_FUSED_ATTN_RES=0``, inputs fitting the fused kernel's + contract dispatch directly to the in-tree ``trtllm::attn_res_fwd`` op. + Only the fallback boundary restores the HF ``[M, K, H]`` layout. + """ + if _FUSED_ATTN_RES_ENABLED: + fused = _apply_attn_res_fused(prefix_sum, block_residual, proj, norm) + if fused is not None: + return fused + block_residual_hf = block_residual.transpose(0, 1) + v = torch.cat((block_residual_hf, prefix_sum.unsqueeze(1)), dim=1) + v_float = v.float() + variance = v_float.pow(2).mean(-1, keepdim=True) + k = v_float * torch.rsqrt(variance + norm.eps) + score_weight = norm.weight.float() * proj.weight.squeeze(0).float() + scores = (k * score_weight).sum(-1) + probs = scores.softmax(-1).unsqueeze(1) + hidden_states = torch.matmul(probs, v_float).squeeze(1) + return hidden_states.to(v.dtype) + + +def _apply_attn_res_and_rmsnorm( + prefix_sum: torch.Tensor, + block_residual: torch.Tensor, + proj: nn.Linear, + norm: KimiK3RMSNorm, + output_norm: nn.Module, +) -> torch.Tensor: + """Apply attention-residual selection and the next RMSNorm.""" + if _FUSED_ATTN_RES_ENABLED: + fused = _apply_attn_res_rmsnorm_fused(prefix_sum, block_residual, proj, norm, output_norm) + if fused is not None: + return fused + return output_norm(_apply_attn_res(prefix_sum, block_residual, proj, norm)) + + +def _apply_attn_res_add_and_rmsnorm( + prefix_sum: torch.Tensor, + addend: torch.Tensor, + block_residual: torch.Tensor, + proj: nn.Linear, + norm: KimiK3RMSNorm, + output_norm: nn.Module, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Add an attention output to the running residual, then select and norm.""" + if _FUSED_ATTN_RES_ENABLED: + fused = _apply_attn_res_add_rmsnorm_fused( + prefix_sum, addend, block_residual, proj, norm, output_norm + ) + if fused is not None: + return fused + updated_prefix_sum = prefix_sum + addend + return updated_prefix_sum, _apply_attn_res_and_rmsnorm( + updated_prefix_sum, block_residual, proj, norm, output_norm + ) + + +# Routed-expert key spellings that ModelOpt emits for Kimi K3. The NVFP4 +# checkpoint (``nvidia/Kimi-K3-NVFP4``) lists every prefix x module-name +# combination in ``quantized_layers``, so a lookup over this product finds it +# without needing the MiniMax-M3-style prefix normalization in ``ModelConfig``. +_K3_ROUTED_EXPERT_KEY_PREFIXES = ("language_model.model.", "model.", "") + + +_K3_ROUTED_EXPERT_KEY_SUFFIXES = ("block_sparse_moe.experts", "mlp.experts") + + +# The subset of the above that can be a real module path. ``exclude_modules`` +# matches with wildcards and walks ancestor prefixes, so an empty prefix would +# widen what matches instead of just missing, as it does in the dict lookup. +_K3_ROUTED_EXPERT_MODULE_PREFIXES = ("language_model.model.", "model.") + + +# Routed-expert quantization used when the checkpoint declares nothing per +# layer. The original ``moonshotai/Kimi-K3`` ships a compressed-tensors +# ``mxfp4-pack-quantized`` config with no ModelOpt per-layer entries, and that +# checkpoint is what this default has always served. +_K3_DEFAULT_ROUTED_QUANT_ALGO = QuantAlgo.W4A8_MXFP4_MXFP8 + + +def _load_packed_mxfp4_expert(backend, base, expert_idx, local_slot_id, get_tensor) -> None: + backend.quant_method.load_packed_mxfp4_expert( + backend, + global_expert_id=expert_idx, + local_slot_id=local_slot_id, + w1_weight=get_tensor(f"{base}.{expert_idx}.w1.weight_packed"), + w1_weight_scale=get_tensor(f"{base}.{expert_idx}.w1.weight_scale"), + w2_weight=get_tensor(f"{base}.{expert_idx}.w2.weight_packed"), + w2_weight_scale=get_tensor(f"{base}.{expert_idx}.w2.weight_scale"), + w3_weight=get_tensor(f"{base}.{expert_idx}.w3.weight_packed"), + w3_weight_scale=get_tensor(f"{base}.{expert_idx}.w3.weight_scale"), + ) + + +def _load_nvfp4_expert(backend, base, expert_idx, local_slot_id, get_tensor) -> None: + backend.quant_method.load_streaming_nvfp4_expert( + backend, + global_expert_id=expert_idx, + local_slot_id=local_slot_id, + **{ + f"{w}_{kind}": get_tensor(f"{base}.{expert_idx}.{w}.{kind}") + for w in ("w1", "w2", "w3") + for kind in ("weight", "weight_scale", "weight_scale_2", "input_scale") + }, + ) + + +class _K3ExpertCkptSpec(NamedTuple): + """How one routed-expert quantization is spelled and loaded.""" + + # Per-``w{1,2,3}`` checkpoint tensor suffixes this layout stores. + kinds: Tuple[str, ...] + loader: Callable[..., None] + # Set of filled slots the loader maintains, checked after the load. + loaded_slots_attr: str + # NVFP4 defers cat/pad/interleave and the alpha computation to + # ``process_weights_after_loading``; the MXFP4 loaders write through. + needs_layer_finalize: bool + + +_K3_EXPERT_CKPT_SPECS = { + QuantAlgo.W4A8_MXFP4_MXFP8: _K3ExpertCkptSpec( + kinds=("weight_packed", "weight_scale"), + loader=_load_packed_mxfp4_expert, + loaded_slots_attr="_packed_mxfp4_loaded_slots", + needs_layer_finalize=False, + ), + QuantAlgo.NVFP4: _K3ExpertCkptSpec( + kinds=("weight", "weight_scale", "weight_scale_2", "input_scale"), + loader=_load_nvfp4_expert, + loaded_slots_attr="_streamed_expert_slots", + needs_layer_finalize=True, + ), +} + + +def _k3_expert_ckpt_spec(quant_algo: Optional[QuantAlgo]) -> _K3ExpertCkptSpec: + spec = _K3_EXPERT_CKPT_SPECS.get(quant_algo) + if spec is None: + raise NotImplementedError( + f"Kimi K3 routed experts are quantized as {quant_algo}, for which " + "no per-expert checkpoint layout is known. Supported: " + f"{sorted(a.name for a in _K3_EXPERT_CKPT_SPECS)}." + ) + return spec + + +class KimiK3MoERuntime(nn.Module): + """Kimi K3 latent MoE block backed by ConfigurableMoE.""" + + def __init__( + self, + model_config: ModelConfig, + cfg, + layer_idx: int, + aux_stream_dict: Dict[AuxStreamType, torch.cuda.Stream], + ): + """Build the routed experts and the shared expert for one MoE layer. + + ``cfg`` is the raw ``PretrainedConfig`` rather than anything derived: + the SiTU soft-caps and the routed-expert geometry are Kimi K3 fields + that ``ModelConfig`` does not carry. + + ``aux_stream_dict`` is shared across every layer of the model, so the + streams reached through it are borrowed and must not be synchronized + or reassigned here. + """ + super().__init__() + self.layer_idx = layer_idx + self.hidden_size = cfg.hidden_size + self.num_experts = cfg.num_experts + self.top_k = cfg.num_experts_per_token + self.moe_hidden_size = cfg.routed_expert_hidden_size + # ValueError (not assert): these guard unsupported checkpoint + # configurations and must stay active under ``python -O``. + if self.moe_hidden_size is None: + raise ValueError("Kimi K3 runtime expects the latent MoE (routed_expert_hidden_size)") + if not getattr(cfg, "latent_moe_use_norm", False): + raise ValueError("Kimi K3 runtime expects latent_moe_use_norm=True") + + situ_beta, situ_linear_beta = _resolve_kimi_situ_betas(cfg) + dtype = torch.bfloat16 + + # Routing scores stay fp32; with attention-DP off the gate GEMM runs + # bf16xbf16 with fp32 accumulate/output (checkpoint stores the gate + # weight in bf16; saves a per-layer input cast + fp32 splitK pair on + # the bs1 decode path). Under attention-DP the legacy upcast-to-fp32 + # GEMM is kept: the bf16-input min-latency GEMM's different reduction + # order flips borderline top-16 picks (GSM8K 96.7 -> 96.1/96.4, + # 3-run bisect on 62b20dd868), and the bs1-latency win is irrelevant + # at DEP batch sizes. KIMI_K3_ROUTER_BF16=1/0 forces either path. + _router_bf16_env = os.environ.get("KIMI_K3_ROUTER_BF16") + _router_bf16 = ( + _router_bf16_env == "1" + if _router_bf16_env is not None + else not model_config.mapping.enable_attention_dp + ) + self.gate = KimiK3MoEGate(cfg, logits_gemm_dtype=torch.bfloat16 if _router_bf16 else None) + + routed_moe_model_config = self._routed_moe_model_config(model_config) + routed_quant_config = self._resolve_routed_quant_config(model_config, layer_idx) + # Resolved here so ``load_weights`` reads the checkpoint layout off the + # module instead of re-deriving it at each of its three call sites. + self.expert_ckpt_spec = _k3_expert_ckpt_spec(routed_quant_config.quant_algo) + routed_moe_kwargs = dict( + routing_method=self.gate.routing_method, + num_experts=self.num_experts, + hidden_size=self.moe_hidden_size, + intermediate_size=cfg.moe_intermediate_size, + dtype=dtype, + # Kimi owns the latent reduction so it can order that collective + # after the shared expert's auxiliary-stream reduction. + reduce_results=False, + model_config=routed_moe_model_config, + override_quant_config=routed_quant_config, + layer_idx=layer_idx, + aux_stream_dict=aux_stream_dict, + # Let CommunicationFactory select the best available strategy. + communication_method=None, + activation=SiTuActivation( + gate_softcap=situ_beta, + linear_softcap=situ_linear_beta, + ), + # A request that silently degraded to CUTLASS would be benchmarked + # as if it were the backend that was asked for, and the decline is + # easy to trigger: MegaMoE has its own token / top-k limits and is + # EP-only, and CuteDSL declines on activation shape, SM version and + # the CuTe DSL dependency. Measured 2026-09-08: a CUTEDSL request + # was turned down on every one of the 92 MoE layers, on all 16 + # ranks, and still produced correct text and a zero exit -- the + # only trace was a warning line per layer. Fail in the resolver + # instead, which reports the rejection trail. + # + # CUTLASS is absent on purpose: it is the fallback target, so + # "degraded to CUTLASS" is not a thing that can happen to it. + allow_backend_degradation=routed_moe_model_config.moe_backend + not in ("MEGAMOE_DEEPGEMM", "MEGAMOE_CUTEDSL", "CUTEDSL"), + ) + self._check_trtllm_situ_quant( + routed_moe_model_config.moe_backend, routed_quant_config.quant_algo + ) + + self.routed_experts = create_moe(**routed_moe_kwargs) + if not isinstance(self.routed_experts, ConfigurableMoE): + raise RuntimeError( + "Kimi K3 requires ConfigurableMoE; ENABLE_CONFIGURABLE_MOE must not be disabled." + ) + if self.routed_experts.layer_load_balancer is not None: + raise NotImplementedError( + "Kimi K3 packed-checkpoint streaming does not yet support " + "dynamic EPLB or replicated expert slots." + ) + local_expert_ids = list(self.routed_experts.backend.initial_local_expert_ids) + if local_expert_ids != list( + range(local_expert_ids[0], local_expert_ids[0] + len(local_expert_ids)) + ): + raise NotImplementedError( + "Kimi K3 packed-checkpoint streaming currently requires a " + "contiguous static expert partition." + ) + self.local_expert_ids = tuple(local_expert_ids) + self.experts_per_rank = len(local_expert_ids) + self.expert_lo = local_expert_ids[0] + self.expert_hi = self.expert_lo + self.experts_per_rank + + shared_intermediate = cfg.moe_intermediate_size * cfg.num_shared_experts + attention_dp = model_config.mapping.enable_attention_dp + shared_model_config = copy.copy(model_config) + shared_model_config.quant_config = QuantConfig() + # Under attention DP each rank owns different tokens, so the shared + # expert is replicated (TP size 1) and must not reduce across ranks. + # Direct MoE-TP leaves both branches as partials for one concatenated + # all-reduce. + use_shared_tp = not attention_dp and model_config.mapping.tp_size > 1 + self._reduce_routed_output = ( + use_shared_tp + and self.routed_experts.backend.scheduler_kind != MoESchedulerKind.FUSED_COMM + ) + if self._reduce_routed_output and self.routed_experts.all_reduce is None: + raise RuntimeError( + "Kimi K3 direct MoE tensor parallelism requires the " + "ConfigurableMoE all-reduce even when reduce_results=False." + ) + self.shared_experts = GatedMLP( + hidden_size=cfg.hidden_size, + intermediate_size=shared_intermediate, + bias=False, + activation=SituAndMul( + beta=situ_beta, + linear_beta=situ_linear_beta, + use_fused_activation=True, + ), + dtype=dtype, + config=shared_model_config, + overridden_tp_size=1 if attention_dp else None, + reduce_output=use_shared_tp, + layer_idx=layer_idx, + is_shared_expert=True, + ) + # Side stream (+ fork/join events) for overlapping shared-expert + # compute with the routed chain. Only engaged when multi-stream is + # active (CUDA graphs on); otherwise both run in order on the default + # stream. + self.shared_expert_stream = aux_stream_dict[AuxStreamType.MoeShared] + self.moe_main_event = torch.cuda.Event() + self.moe_shared_event = torch.cuda.Event() + self.routed_expert_down_proj = nn.Linear( + cfg.hidden_size, self.moe_hidden_size, bias=False, dtype=dtype + ) + self.routed_expert_up_proj = nn.Linear( + self.moe_hidden_size, cfg.hidden_size, bias=False, dtype=dtype + ) + # Stock fused RMSNorm (flashinfer kernel; the no-flashinfer + # fallback is the same fp32-variance eager math as KimiK3RMSNorm). + self.routed_expert_norm = RMSNorm( + hidden_size=self.moe_hidden_size, eps=cfg.rms_norm_eps, dtype=dtype + ) + + @staticmethod + def _routed_projection(hidden_states: torch.Tensor, projection: nn.Module) -> torch.Tensor: + if _K3_DISABLE_MIN_LATENCY_LATENT_PROJ or not isinstance(projection, nn.Linear): + return projection(hidden_states) + return torch.ops.trtllm.dsv3_fused_a_gemm_op( + hidden_states, projection.weight.t(), None, None + ) + + @staticmethod + def _select_moe_tp_ep(mapping: Mapping) -> Tuple[int, int]: + """Resolve the routed-expert ``(moe_tp, moe_ep)`` split. + + Precedence: + + 1. Explicit ``moe_tensor_parallel_size`` / ``moe_expert_parallel_size`` + from the user config. Detected via + ``mapping.moe_tp_ep_user_specified`` so the auto-resolved mapping + default (``moe_tp=tp_size, moe_ep=1``) is NOT mistaken for a TP + request. + 2. Default: EP-only (``moe_tp=1, moe_ep=tp_size``), the historical + K3 layout. + """ + tp_size = mapping.tp_size + if getattr(mapping, "moe_tp_ep_user_specified", False): + return mapping.moe_tp_size, mapping.moe_ep_size + return 1, tp_size + + @staticmethod + def _resolve_routed_quant_config(model_config: ModelConfig, layer_idx: int) -> QuantConfig: + """Routed-expert quantization for ``layer_idx``, taken from the checkpoint. + + ``nvidia/Kimi-K3-NVFP4`` declares the routed experts per layer as + ``NVFP4`` with ``group_size=16``; the original ``moonshotai/Kimi-K3`` + declares nothing per layer and keeps the historical + ``W4A8_MXFP4_MXFP8`` default. Reading the checkpoint instead of + hardcoding is what lets one code path serve both. + + An exclusion outranks the per-layer entry and the default below: + ``create_weights`` treats an override as authoritative over anything + ``__post_init__`` wrote, so this return value stands in for both + quantization passes and exclusion is the one that runs second. It is + matched as a pattern, so it is asked only about real module names. + """ + quant_config = model_config.quant_config + if quant_config is not None and any( + quant_config.is_module_excluded_from_quantization( + f"{prefix}layers.{layer_idx}.{suffix}" + ) + for prefix in _K3_ROUTED_EXPERT_MODULE_PREFIXES + for suffix in _K3_ROUTED_EXPERT_KEY_SUFFIXES + ): + logger.debug( + "Kimi K3 layer %d routed experts: excluded from quantization, " + "keeping them unquantized", + layer_idx, + ) + return QuantConfig(kv_cache_quant_algo=quant_config.kv_cache_quant_algo) + + per_layer = getattr(model_config, "quant_config_dict", None) + if per_layer: + for prefix in _K3_ROUTED_EXPERT_KEY_PREFIXES: + for suffix in _K3_ROUTED_EXPERT_KEY_SUFFIXES: + cfg = per_layer.get(f"{prefix}layers.{layer_idx}.{suffix}") + if cfg is not None and cfg.quant_algo is not None: + # Logged once per layer: the routed-expert format decides + # which MoE backends can serve this checkpoint at all. + logger.debug( + "Kimi K3 layer %d routed experts: %s (group_size=%s) " + "from the checkpoint", + layer_idx, + cfg.quant_algo, + cfg.group_size, + ) + return cfg + logger.debug( + "Kimi K3 layer %d routed experts: no per-layer quant config in the " + "checkpoint, defaulting to %s", + layer_idx, + _K3_DEFAULT_ROUTED_QUANT_ALGO, + ) + return QuantConfig(quant_algo=_K3_DEFAULT_ROUTED_QUANT_ALGO) + + @staticmethod + def _check_trtllm_situ_quant(moe_backend: str, quant_algo: Optional[QuantAlgo]) -> None: + """Reject a routed-expert format trtllm-gen has no fused SiTu cubin for. + + trtllm-gen has fused SiTu FC1 cubins for two input formats and no + standalone SiTu activation kernel, so anything else has to die here + rather than in a cubin lookup deep inside the runner. Checked against + the resolved backend, not the K3 architecture branch, because the + generic FP8_BLOCK_SCALES fallback in ``resolve_moe_backend`` can also + land on TRTLLM. + + The admitted set is read off the backend rather than restated here, + because restating it is what broke. This guard was written in #17865 + when MXFP4 was the only fused SiTu drop; #17940 then added the NVFP4 + (group-16 ``Bmm_E2m1_E2m1E2m1_..._siTuGlu_*``) cubins and updated + ``TRTLLMGenFusedMoE``'s set without touching this copy. For the week + in between, an NVFP4 K3 checkpoint could not start at all -- and not + only when TRTLLM was asked for by name, because + ``ModelConfig.resolve_moe_backend`` sends every K3 architecture to + TRTLLM, so the default AUTO configuration hit this raise too. The unit + tests did not catch it: they call ``create_moe`` directly and never + reach this guard, so the kernel path stayed green while the model path + was closed. + + A staticmethod, not an inline block, so that the invariant is + reachable from a test without constructing the whole runtime. + """ + situ_supported = TRTLLMGenFusedMoE.situ_supported_quant_algos() + if moe_backend != "TRTLLM" or quant_algo in situ_supported: + return + supported = ", ".join(sorted(algo.name for algo in situ_supported)) + raise ValueError( + f"Kimi K3 routed experts are quantized as {quant_algo}, which the " + "TRTLLM (trtllm-gen) MoE backend cannot serve: fused SiTu cubins " + f"exist only for {supported}. Set moe_config.backend to CUTLASS " + "or MEGAMOE_CUTEDSL." + ) + + @staticmethod + def _routed_moe_model_config(model_config: ModelConfig) -> ModelConfig: + """Build a private routed-expert mapping without mutating the shared + config. Default split is EP-only; see ``_select_moe_tp_ep``.""" + # Every backend here declares ``ActivationType.SiTu`` in its + # ``activation_support``; the list is not a preference order. CUTEDSL + # joined once its act-fusion kernel grew the SiTU epilogue. + supported_backends = { + "CUTLASS", + "TRTLLM", + "CUTEDSL", + "MEGAMOE_DEEPGEMM", + "MEGAMOE_CUTEDSL", + } + if model_config.moe_backend not in supported_backends: + raise ValueError( + "Kimi K3 SiTU routed experts only support the CUTLASS, TRTLLM, " + "CUTEDSL, MEGAMOE_DEEPGEMM, and MEGAMOE_CUTEDSL backends; " + f"got {model_config.moe_backend!r}." + ) + if model_config.moe_load_balancer is not None: + raise NotImplementedError( + "Kimi K3 packed-checkpoint streaming does not yet support " + "EPLB or replicated expert slots." + ) + mapping = model_config.mapping + if getattr(mapping, "_dwdp_size", 0) > 1: + raise NotImplementedError("Kimi K3 packed-checkpoint streaming does not support DWDP.") + + moe_tp, moe_ep = KimiK3MoERuntime._select_moe_tp_ep(mapping) + if moe_tp < 1 or moe_ep < 1 or moe_tp * moe_ep != mapping.tp_size: + raise ValueError( + f"Kimi K3 routed MoE split moe_tp={moe_tp} x moe_ep={moe_ep} " + f"must multiply to tp_size={mapping.tp_size}." + ) + if moe_tp > 1 and mapping.enable_attention_dp: + raise NotImplementedError( + "Kimi K3 MoE tensor parallelism requires " + "enable_attention_dp=false (the attention-DP dispatch/combine " + "path is validated for EP-only splits)." + ) + logger.info_once( + f"Kimi K3 routed MoE parallelism: moe_tp={moe_tp}, " + f"moe_ep={moe_ep} (tp_size={mapping.tp_size})", + key="kimi_k3_moe_tp_ep_split", + ) + + mapping_dict = mapping.to_dict() + mapping_dict["moe_cluster_size"] = 1 + mapping_dict["moe_tp_size"] = moe_tp + mapping_dict["moe_ep_size"] = moe_ep + routed_mapping = Mapping.from_dict(mapping_dict) + + routed_model_config = copy.copy(model_config) + routed_model_config._frozen = False + routed_model_config.extra_attrs = copy.copy(model_config.extra_attrs) + routed_model_config.mapping = routed_mapping + routed_model_config.moe_backend = model_config.moe_backend + # MegaMoE uses this value as global DP SymmBuffer capacity, then + # divides it by EP size for the per-rank allocation. Other backends + # keep the user-configured value as their MoE chunking bound. + # Preserve an explicitly larger capacity. + if routed_model_config.moe_backend in { + "MEGAMOE_DEEPGEMM", + "MEGAMOE_CUTEDSL", + }: + default_moe_max_num_tokens = routed_model_config.max_num_tokens * routed_mapping.dp_size + configured_moe_max_num_tokens = int(routed_model_config.moe_max_num_tokens or 0) + if configured_moe_max_num_tokens < default_moe_max_num_tokens: + logger.info_once( + "Kimi K3 MegaMoE raises moe_max_num_tokens from " + f"{configured_moe_max_num_tokens} to {default_moe_max_num_tokens} " + "because the global DP SymmBuffer requires capacity for " + "max_num_tokens * dp_size.", + key=( + "kimi_k3_megamoe_capacity_override_" + f"{configured_moe_max_num_tokens}_{default_moe_max_num_tokens}" + ), + ) + routed_model_config.moe_max_num_tokens = max( + configured_moe_max_num_tokens, + default_moe_max_num_tokens, + ) + routed_model_config._frozen = True + return routed_model_config + + def forward(self, hidden_states: torch.Tensor, all_rank_num_tokens=None) -> torch.Tensor: + """``hidden_states``: ``[num_tokens, hidden_size]`` bf16.""" + identity = hidden_states + router_logits = self.gate.compute_logits(hidden_states) + moe_all_reduce = self.routed_experts.all_reduce if self._reduce_routed_output else None + + def _routed_output(): + # Latent down/up projections via the min-latency fused GEMM op: + # at <=16 tokens (decode graphs) it runs a single pipelined + # bf16 kernel per projection instead of cuBLAS's split-K GEMV + + # splitKreduce pair (~17+3.6us -> ~8us for 7168->3584 at M=1); + # for larger token counts the op falls back to cuBLAS internally. + # TLLM_K3_DISABLE_MIN_LATENCY_LATENT_PROJ=1 restores nn.Linear + # (A/B escape hatch). When the FP8 weight-read conversion has + # replaced the projection module, call it directly: its weight is + # an e4m3 buffer the bf16 dsv3 op must not read, and its forward + # is already a single fused GEMM (fp8_swap_ab_gemm). + routed_in = self._routed_projection(hidden_states, self.routed_expert_down_proj) + y = self.routed_experts( + routed_in, + router_logits, + all_rank_num_tokens=all_rank_num_tokens, + ) + if self._reduce_routed_output: + return y + # Communication-backed paths return a complete routed result. + y = self.routed_expert_norm(y) + return self._routed_projection(y, self.routed_expert_up_proj) + + # Shared experts depend only on the block input, so overlap their GEMMs + # with the routed dispatch/expert/combine chain. Multi-stream engages + # only under CUDA graphs; otherwise both branches run in order on the + # default stream. The shared GatedMLP includes its output all-reduce on + # the auxiliary stream. The join below must precede the routed + # all-reduce: concurrent collectives on different streams can corrupt + # SYMM_MEM all-reduce state. + routed_out, shared_out = maybe_execute_in_parallel( + _routed_output, + lambda: self.shared_experts(identity), + self.moe_main_event, + self.moe_shared_event, + self.shared_expert_stream, + disable_on_compile=True, + ) + if self._reduce_routed_output: + routed_latent = moe_all_reduce(routed_out) + routed_latent = self.routed_expert_norm(routed_latent) + routed_out = self._routed_projection(routed_latent, self.routed_expert_up_proj) + return routed_out + shared_out + + +def resolve_attention_quant_config( + config: ModelConfig | None, layer_idx: int, projection: str +) -> QuantConfig: + """Resolve a checkpoint projection, including mixed-precision exclusions.""" + if config is None: + return QuantConfig() + global_config = config.quant_config or QuantConfig() + names = [ + f"{prefix}layers.{layer_idx}.self_attn.{projection}" + for prefix in ("language_model.model.", "model.", "") + ] + if any(global_config.is_module_excluded_from_quantization(name) for name in names): + return QuantConfig(kv_cache_quant_algo=global_config.kv_cache_quant_algo) + declarations = config.quant_config_dict or {} + matches = [declarations[name] for name in names if name in declarations] + if matches: + selected = matches[0] + if any(match.quant_algo != selected.quant_algo for match in matches[1:]): + raise ValueError(f"Conflicting Kimi K3 quantization aliases for {names[0]}") + elif global_config.quant_algo == QuantAlgo.MIXED_PRECISION: + selected = QuantConfig() + else: + selected = global_config + if selected.quant_algo not in (None, QuantAlgo.FP8_BLOCK_SCALES): + raise ValueError( + f"Kimi K3 attention projection {names[0]} has unsupported checkpoint " + f"quantization {selected.quant_algo}" + ) + if selected.quant_algo == QuantAlgo.FP8_BLOCK_SCALES and selected.group_size not in (None, 128): + raise ValueError(f"Kimi K3 attention requires 128x128 FP8 blocks for {names[0]}") + result = copy.copy(selected) + result.kv_cache_quant_algo = global_config.kv_cache_quant_algo + return result + + +class KimiMLARuntime(nn.Module): + """Wraps K3 MLA and applies its external TP output reduction.""" + + def __init__( + self, + cfg: "PretrainedConfig", + layer_idx: int, + model_config: ModelConfig, + aux_stream_dict: Dict[AuxStreamType, torch.cuda.Stream], + mapping_with_cp: Optional[Mapping] = None, + ) -> None: + super().__init__() + + from tensorrt_llm._torch.modules.kimi_k3_mla import KimiK3MLAAttention + + max_positions = int( + os.environ.get( + _KIMI_K3_MLA_MAX_POSITIONS_ENV, + cfg.max_position_embeddings, + ) + ) + self.layer_idx = layer_idx + # KimiK3MLAAttention owns MLA projection/head sharding. Keep only the + # final output reduction in this wrapper so the output gate remains + # between attention and the row-parallel o_proj. + # Helix: mapping_with_cp (the CP original) activates the base MLA's + # helix machinery; this wrapper's allreduce over the repurposed + # mapping sums the base o_proj's tp*cp partials. + mapping = model_config.mapping + reduce_output = not mapping.enable_attention_dp and mapping.tp_size > 1 + self._o_allreduce = ( + AllReduce( + mapping=mapping, + strategy=model_config.allreduce_strategy, + dtype=torch.bfloat16, + ) + if reduce_output + else None + ) + attention_config = copy.copy(model_config) + attention_config._frozen = False + attention_config.quant_config_dict = { + name: resolve_attention_quant_config(model_config, layer_idx, name) + for name in ( + "q_a_proj", + "kv_a_proj_with_mqa", + "q_b_proj", + "kv_b_proj", + "g_proj", + "o_proj", + ) + } + attention_config.quant_config = QuantConfig( + kv_cache_quant_algo=model_config.quant_config.kv_cache_quant_algo + if model_config.quant_config is not None + else None + ) + attention_config._frozen = model_config._frozen + self.mixer = KimiK3MLAAttention( + hidden_size=cfg.hidden_size, + num_heads=cfg.num_attention_heads, + q_lora_rank=cfg.q_lora_rank, + kv_lora_rank=cfg.kv_lora_rank, + qk_nope_head_dim=cfg.qk_nope_head_dim, + qk_rope_head_dim=cfg.qk_rope_head_dim, + v_head_dim=cfg.v_head_dim, + rms_norm_eps=cfg.rms_norm_eps, + dtype=torch.bfloat16, + layer_idx=layer_idx, + use_output_gate=cfg.mla_use_output_gate, + max_position_embeddings=max_positions, + model_config=attention_config, + aux_stream_dict=aux_stream_dict, + mapping_with_cp=mapping_with_cp, + ) + + def forward( + self, hidden_states: torch.Tensor, attn_metadata: AttentionMetadata + ) -> torch.Tensor: + # MLA.forward takes position_ids first; K3 is NoPE, so pass None. + out = self.mixer(None, hidden_states, attn_metadata) + if self._o_allreduce is not None: + # Head-sharded TP: sum the row-sharded o_proj partials across + # the head-shard group. + out = self._o_allreduce(out) + return out + + +class KimiLinearDecoderLayer(nn.Module): + def __init__( + self, + model_config: ModelConfig, + cfg, + layer_idx: int, + aux_stream_dict: Dict[AuxStreamType, torch.cuda.Stream], + ): + super().__init__() + self.layer_idx = layer_idx + self.hidden_size = cfg.hidden_size + dtype = torch.bfloat16 + + self.is_kda = _is_kda_layer(cfg, layer_idx) + is_mla = _is_mla_layer(cfg, layer_idx) + if self.is_kda == is_mla: + raise ValueError(f"Kimi K3 layer {layer_idx} must be exactly one of KDA/MLA") + + if self.is_kda: + projection_names = ("q_proj", "k_proj", "v_proj", "g_proj", "o_proj") + attention_config = copy.copy(model_config) + attention_config._frozen = False + attention_config.quant_config_dict = { + name: resolve_attention_quant_config(model_config, layer_idx, name) + for name in projection_names + } + attention_config._frozen = model_config._frozen + self.linear_attn = KimiKDALinearAttention( + cfg, + layer_idx, + mapping=model_config.mapping, + allreduce_strategy=model_config.allreduce_strategy, + aux_stream=aux_stream_dict[AuxStreamType.Attention], + model_config=attention_config, + ) + else: + self.self_attn = KimiMLARuntime( + cfg, + layer_idx, + model_config=model_config, + aux_stream_dict=aux_stream_dict, + # CP original stashed by _setup_helix_mappings; None outside helix. + mapping_with_cp=getattr(model_config, "_helix_mapping_with_cp", None), + ) + + self.is_moe = ( + cfg.num_experts is not None + and layer_idx >= cfg.first_k_dense_replace + and layer_idx % getattr(cfg, "moe_layer_freq", 1) == 0 + ) + if self.is_moe: + self.block_sparse_moe = KimiK3MoERuntime(model_config, cfg, layer_idx, aux_stream_dict) + else: + situ_beta = getattr(cfg, "activation_situ_beta", None) or 1.0 + situ_linear_beta = getattr(cfg, "activation_situ_linear_beta", None) + attention_dp = model_config.mapping.enable_attention_dp + if attention_dp: + self.mlp_tp_size = 1 + else: + self.mlp_tp_size = math.gcd(cfg.intermediate_size, model_config.mapping.tp_size) + if self.mlp_tp_size > model_config.mapping.gpus_per_node: + self.mlp_tp_size = math.gcd( + self.mlp_tp_size, model_config.mapping.gpus_per_node + ) + mlp_model_config = copy.copy(model_config) + mlp_model_config.quant_config = QuantConfig() + # K3's dense layer is BF16, so a unit block size gives the same + # subgroup selection as DeepSeek-V3. Attention DP replicates the + # MLP because ranks own different tokens; otherwise the subgroup + # is block-aligned and stays within one node. + self.mlp = GatedMLP( + hidden_size=cfg.hidden_size, + intermediate_size=cfg.intermediate_size, + bias=False, + activation=SituAndMul( + beta=situ_beta, + linear_beta=situ_linear_beta, + use_fused_activation=True, + ), + dtype=dtype, + config=mlp_model_config, + overridden_tp_size=self.mlp_tp_size, + reduce_output=self.mlp_tp_size > 1, + layer_idx=layer_idx, + ) + + # Stock fused RMSNorm for the plain (whole-tensor) norms; numerics + # are drop-in for KimiK3RMSNorm (fp32 variance, weight applied + # after downcast, use_gemma=False). + self.input_layernorm = RMSNorm( + hidden_size=cfg.hidden_size, eps=cfg.rms_norm_eps, dtype=dtype + ) + self.post_attention_layernorm = RMSNorm( + hidden_size=cfg.hidden_size, eps=cfg.rms_norm_eps, dtype=dtype + ) + + # Attention residual scheme (always on for K3). The res norms stay + # KimiK3RMSNorm: they are consumed field-wise (.weight/.eps) by + # _apply_attn_res and the fused attn_res op, never called as + # modules. + self.attn_res_block_size = cfg.attn_res_block_size + assert self.attn_res_block_size is not None, ( + "Kimi K3 runtime expects attn_res_block_size to be set" + ) + self.self_attention_res_norm = KimiK3RMSNorm( + cfg.hidden_size, eps=cfg.rms_norm_eps, dtype=dtype + ) + self.mlp_res_norm = KimiK3RMSNorm(cfg.hidden_size, eps=cfg.rms_norm_eps, dtype=dtype) + self.self_attention_res_proj = nn.Linear(cfg.hidden_size, 1, bias=False, dtype=dtype) + self.mlp_res_proj = nn.Linear(cfg.hidden_size, 1, bias=False, dtype=dtype) + + def forward( + self, + hidden_states: torch.Tensor, + block_residual: torch.Tensor, + num_snapshots: int, + attn_metadata: AttentionMetadata, + capture: Optional[Tuple[Any, int]] = None, + ) -> Tuple[torch.Tensor, int]: + """Port of HF ``KimiDecoderLayer._forward_attn_residual`` (per token). + + ``block_residual`` is a preallocated snapshot bank in kernel-native + ``[K_max, M, H]`` layout. Returns the running prefix sum and the + number of valid bank rows. + + ``capture`` is ``(spec_metadata, layer_id)`` and taps the DSpark aux + stream for the layer BEFORE this one: the aggregated stream for layer j + is by definition what its next consumer sees, so the mixture computed + below already is it. Reading it here beats recomputing it, and is only + possible because K3 asserts pp_size == 1 -- layer j+1 is always local. + PP support would need a recompute at the rank boundary. + """ + prefix_sum = hidden_states + valid_block_residual = block_residual[:num_snapshots] + + if capture is not None: + # The mixture tap needs the PRE-norm value, which the fused + # attn-res + RMSNorm kernel does not expose. Keep the two steps + # split on captured layers only and fuse everywhere else. + if num_snapshots > 0: + hidden_states = _apply_attn_res( + prefix_sum, + valid_block_residual, + self.self_attention_res_proj, + self.self_attention_res_norm, + ) + # A property of the DRAFTER checkpoint, not a knob: a mismatch only lowers + # acceptance, silently. hidden_states is the pre-norm attn_res mixture; + # prefix_only wants the running prefix, already in hand as prefix_sum. + tapped = hidden_states if _AUX_ATTN_RES_STREAM_ENABLED else prefix_sum + capture[0].maybe_capture_hidden_states(capture[1], tapped, None) + hidden_states = self.input_layernorm(hidden_states) + elif num_snapshots > 0: + hidden_states = _apply_attn_res_and_rmsnorm( + prefix_sum, + valid_block_residual, + self.self_attention_res_proj, + self.self_attention_res_norm, + self.input_layernorm, + ) + else: + hidden_states = self.input_layernorm(hidden_states) + + if self.layer_idx % self.attn_res_block_size == 0: + block_residual[num_snapshots].copy_(prefix_sum) + num_snapshots += 1 + valid_block_residual = block_residual[:num_snapshots] + prefix_sum = None + if self.is_kda: + hidden_states = self.linear_attn(hidden_states, attn_metadata) + else: + hidden_states = self.self_attn(hidden_states, attn_metadata) + + if prefix_sum is None: + prefix_sum = hidden_states + hidden_states = _apply_attn_res_and_rmsnorm( + prefix_sum, + valid_block_residual, + self.mlp_res_proj, + self.mlp_res_norm, + self.post_attention_layernorm, + ) + else: + prefix_sum, hidden_states = _apply_attn_res_add_and_rmsnorm( + prefix_sum, + hidden_states, + valid_block_residual, + self.mlp_res_proj, + self.mlp_res_norm, + self.post_attention_layernorm, + ) + if self.is_moe: + hidden_states = self.block_sparse_moe( + hidden_states, getattr(attn_metadata, "all_rank_num_tokens", None) + ) + else: + hidden_states = self.mlp(hidden_states) + + prefix_sum = prefix_sum + hidden_states + return prefix_sum, num_snapshots + + def skip_forward( + self, + hidden_states: torch.Tensor, + block_residual: torch.Tensor, + attn_metadata: AttentionMetadata, + ) -> Tuple[torch.Tensor, torch.Tensor]: + """No-op stand-in for ``forward``, matching ``DecoderLayer.skip_forward``. + + ``modeling_utils.skip_forward()`` only drops a module's weights when it + finds this attribute, so without it the layer-wise benchmarks would + allocate all 93 layers instead of the profiled slice. + """ + return hidden_states, block_residual + + +class KimiLinearModel(DecoderModel): + def __init__(self, model_config: ModelConfig): + super().__init__(model_config) + cfg = _get_text_config(model_config.pretrained_config) + self._text_cfg = cfg + dtype = torch.bfloat16 + + # Attention and MoE phases are sequential, so their branch-overlap + # roles share one stream; MoE-internal overlap roles remain separate. + aux_stream_list = [torch.cuda.Stream() for _ in range(4)] + self.aux_stream_dict = { + AuxStreamType.Attention: aux_stream_list[0], + AuxStreamType.MoeShared: aux_stream_list[0], + AuxStreamType.MoeChunkingOverlap: aux_stream_list[1], + AuxStreamType.MoeBalancer: aux_stream_list[2], + AuxStreamType.MoeOutputMemset: aux_stream_list[3], + } + + self.embed_tokens = nn.Embedding(cfg.vocab_size, cfg.hidden_size, dtype=dtype) + self.layers = nn.ModuleList( + [ + KimiLinearDecoderLayer(model_config, cfg, layer_idx, self.aux_stream_dict) + for layer_idx in range(cfg.num_hidden_layers) + ] + ) + self.norm = RMSNorm(hidden_size=cfg.hidden_size, eps=cfg.rms_norm_eps, dtype=dtype) + + # KimiK3RMSNorm (not RMSNorm): consumed field-wise (.weight/.eps) + # by _apply_attn_res and the fused attn_res op. + self.output_attn_res_norm = KimiK3RMSNorm( + cfg.hidden_size, eps=cfg.rms_norm_eps, dtype=dtype + ) + self.output_attn_res_proj = nn.Linear(cfg.hidden_size, 1, bias=False, dtype=dtype) + self.num_attn_res_snapshots = ( + cfg.num_hidden_layers + cfg.attn_res_block_size - 1 + ) // cfg.attn_res_block_size + + # Which convention the drafter tap is on is not recoverable from the + # served output -- a mismatch only lowers acceptance -- so state it once + # at construction rather than leaving it to be inferred from an AL. + logger.info_once( + "Kimi K3 aux hidden capture: mode=" + f"{'attn_res_stream' if _AUX_ATTN_RES_STREAM_ENABLED else 'prefix_only'} " + f"({KIMI_K3_AUX_ATTN_RES_STREAM_ENV}={int(_AUX_ATTN_RES_STREAM_ENABLED)})", + key="kimi_k3_aux_capture_mode", + ) + + def forward( + self, + attn_metadata: AttentionMetadata, + input_ids: Optional[torch.IntTensor] = None, + position_ids: Optional[torch.IntTensor] = None, + inputs_embeds: Optional[torch.FloatTensor] = None, + spec_metadata=None, + **kwargs, + ) -> torch.Tensor: + if (input_ids is None) ^ (inputs_embeds is not None): + raise ValueError("You must specify exactly one of input_ids or inputs_embeds") + + if inputs_embeds is None: + inputs_embeds = self.embed_tokens(input_ids) + hidden_states = inputs_embeds + + block_residual = hidden_states.new_empty( + self.num_attn_res_snapshots, + hidden_states.shape[0], + hidden_states.shape[1], + ) + num_snapshots = 0 + capture_set = ( + getattr(spec_metadata, "_capture_layer_set", None) + if spec_metadata is not None + else None + ) + for i, layer in enumerate(self.layers): + # DFlash/DSpark hidden-state capture. The drafter is distilled on + # the aggregated stream value -- the pre-norm softmax mixture its + # next consumer sees -- not on the raw prefix sum a layer returns, + # which is SGLang's fallback for models without the + # attention-residual scheme. Capturing the prefix sum costs 4.5pt + # of draft acceptance on K3 + RadixArk DSpark (AR 66.9% -> 71.4%). + # The tap fires inside layer i+1, which computes that tensor + # anyway; see its forward docstring. Ground truth: SGLang + # kimi_k3.py:2697 _dspark_capture_stream, attn_residual.py:313 + # aggregate_stream_torch. + capture = None + if ( + spec_metadata is not None + and i > 0 + and (capture_set is None or self.layers[i - 1].layer_idx in capture_set) + ): + capture = (spec_metadata, self.layers[i - 1].layer_idx) + hidden_states, num_snapshots = layer( + hidden_states, block_residual, num_snapshots, attn_metadata, capture=capture + ) + + # The last layer has no successor, so this one recompute is + # unavoidable -- output-side score weights, matching SGLang's + # layer_idx + 1 >= end_layer branch. Unreachable for K3's capture set + # against 93 layers; kept so a set that does include the final layer + # gets the right tensor rather than the raw prefix sum. + if spec_metadata is not None and len(self.layers) > 0: + last = self.layers[-1] + if capture_set is None or last.layer_idx in capture_set: + tail = ( + _apply_attn_res( + hidden_states, + block_residual[:num_snapshots], + self.output_attn_res_proj, + self.output_attn_res_norm, + ) + if num_snapshots > 0 and _AUX_ATTN_RES_STREAM_ENABLED + else hidden_states + ) + spec_metadata.maybe_capture_hidden_states(last.layer_idx, tail, None) + + return _apply_attn_res_and_rmsnorm( + hidden_states, + block_residual[:num_snapshots], + self.output_attn_res_proj, + self.output_attn_res_norm, + self.norm, + ) + + +# ---------------------------------------------------------------------------------------------------------------------- +# The target: step classification, the construction checks and the registration shell. +# ---------------------------------------------------------------------------------------------------------------------- + + def _text_model_config(model_config: ModelConfig) -> ModelConfig: """The language model's ModelConfig: the checkpoint's text_config, with quant exclusions renamed to match. @@ -233,7 +1766,12 @@ def _check_construction(model_config: ModelConfig) -> None: @register_auto_model("ModelingV2KimiK3Mxfp4Sm100Tp16Moetp4ep4") class ModelingV2KimiK3Mxfp4Sm100Tp16Moetp4ep4(KimiLinearForCausalLM): - """The registration shell: the built-in Kimi K3 text model as the generic path, behind this target's checks.""" + """The registration shell: this target's text model (`KimiLinearModel` above) behind its checks. + + It inherits the built-in Kimi K3 causal LM for the checkpoint load and the engine hooks (`load_weights` through + weights.py, the KDA metadata class, the model defaults), which walk the model by its module names; the text model + keeps the built-in one's. + """ @classmethod def get_preferred_kv_cache_manager_version(cls, pretrained_config: Any = None) -> Literal["V2"]: @@ -247,7 +1785,27 @@ def __init__(self, model_config: ModelConfig): "text_config" ) _check_construction(model_config) - super().__init__(_text_model_config(model_config)) + text = _text_model_config(model_config) + spec_config = getattr(text, "spec_config", None) + assert ( + spec_config is None + or spec_config.spec_dec_mode.is_sa() + or spec_config.spec_dec_mode.is_dflash() + or spec_config.spec_dec_mode.is_dspark() + ), "Kimi K3 supports speculative decoding only with SA, DFlash or DSpark" + # The inherited loader reads these: this target has neither helix context parallelism nor the fp8 + # weight-read conversion of the shared / latent MLPs. + self._fp8_weight_read_moe_mlp = False + self.mapping_with_cp = None + self._repurposed_tp_mapping = None + cfg = text.pretrained_config + SpecDecOneEngineForCausalLM.__init__( + self, + KimiLinearModel(text), + text, + hidden_size=cfg.hidden_size, + vocab_size=cfg.vocab_size, + ) self._step_checked = False # The fused decode path and the state its kernels share, built in post_load_weights once the catalog entries # it calls exist. None: every step takes the generic path. From 9bc4dae6b7d00293712c1b2b6b21e05c15db39fa Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 19:14:05 -0700 Subject: [PATCH 034/161] [None][test] modeling_v2 attn_res entries: replay once before poisoning the captured outputs The CUDA-graph tests of norm/attn_res_fwd and norm/attn_res_rmsnorm_fwd filled the captured call's outputs with NaN before the graph had ever run. Under PYTORCH_CUDA_ALLOC_CONF=backend:cudaMallocAsync those outputs are graph allocations with no memory behind them until the first launch, so the fill was an illegal address. They now replay once, poison the outputs, replay again and compare. Each iteration also frees its outputs before the next capture starts. Signed-off-by: Vasanth Sabavat --- .../_torch/modeling_v2/norm/test_modeling_v2_attn_res_fwd.py | 5 +++++ .../norm/test_modeling_v2_attn_res_rmsnorm_fwd.py | 5 +++++ 2 files changed, 10 insertions(+) diff --git a/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_attn_res_fwd.py b/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_attn_res_fwd.py index f4a4a8bf3f50..c74c9e24671a 100644 --- a/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_attn_res_fwd.py +++ b/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_attn_res_fwd.py @@ -158,6 +158,9 @@ def test_cuda_graph_replay_matches_eager() -> None: graph = torch.cuda.CUDAGraph() with torch.cuda.graph(graph): captured = attn_res_fwd(*inputs, RMS_EPS) + # One replay before the poison: under cudaMallocAsync the captured outputs are graph allocations, backed by + # memory only once the graph has run. + graph.replay() for tensor in captured: tensor.fill_(float("nan")) graph.replay() @@ -166,3 +169,5 @@ def test_cuda_graph_replay_matches_eager() -> None: assert torch.equal(_bits(actual), _bits(expected)), ( f"T={num_tokens} N={num_candidates}: replayed {name} differs from eager" ) + # Freed here rather than inside the next capture, where cudaMallocAsync's free would be part of it. + del captured diff --git a/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_attn_res_rmsnorm_fwd.py b/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_attn_res_rmsnorm_fwd.py index 48ef9887ccac..d6a94e5555e4 100644 --- a/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_attn_res_rmsnorm_fwd.py +++ b/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_attn_res_rmsnorm_fwd.py @@ -153,12 +153,17 @@ def test_cuda_graph_replay_matches_eager() -> None: graph = torch.cuda.CUDAGraph() with torch.cuda.graph(graph): captured = attn_res_rmsnorm_fwd(*inputs, RMS_EPS, RMS_EPS) + # One replay before the poison: under cudaMallocAsync the captured output is a graph allocation, backed by + # memory only once the graph has run. + graph.replay() captured.fill_(float("nan")) graph.replay() torch.cuda.synchronize() assert torch.equal(_bits(captured), _bits(eager)), ( f"T={num_tokens} N={num_candidates}: replay differs from eager" ) + # Freed here rather than inside the next capture, where cudaMallocAsync's free would be part of it. + del captured def test_chained_calls_wait_for_their_input() -> None: From 8bfb169a4ca571237838d3dda4488af76033bb5f Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 19:14:20 -0700 Subject: [PATCH 035/161] [None][test] modeling_v2 claims: a target's generic path declares every stock import A target whose generic path runs stock code lists it in UNCERTIFIED_GENERIC_CALLS. The claims test now checks the list both ways, from the target's source (no import): every tensorrt_llm name the target imports outside the catalog, module level or function local, is declared unless it computes nothing (types, configs, enums, metadata, registration, logging), and nothing is declared that the target does not import. The Kimi K3 target's list gains the stock code its copied text model runs: the causal-LM bases, ConfigurableMoE / TRTLLMGenFusedMoE and the DeepSeek-V3 routing method. Signed-off-by: Vasanth Sabavat --- .../modeling.py | 12 +++- .../modeling_v2/test_modeling_v2_claims.py | 68 +++++++++++++++++++ 2 files changed, 77 insertions(+), 3 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index 1b7a5cacbd5e..6cf3cd599c97 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -122,15 +122,21 @@ "kv_cache_manager": ("enable_block_reuse",), } -#: Calls the generic path makes outside the catalog, declared so they are not consumed silently. A call leaves this -#: list when a catalog entry replaces it. +#: Stock code the generic path runs outside the catalog, declared so it is not consumed silently: every +#: tensorrt_llm import of this module that computes (test_modeling_v2_claims.py checks the list both ways). An entry +#: leaves the list when a catalog entry replaces it. UNCERTIFIED_GENERIC_CALLS = ( - # The checkpoint load and the engine hooks this target inherits. + # The checkpoint load and the engine hooks this target inherits, and the causal LM around the text model. "tensorrt_llm._torch.models.modeling_kimi_linear.KimiLinearForCausalLM", + "tensorrt_llm._torch.models.modeling_speculative.SpecDecOneEngineForCausalLM", + "tensorrt_llm._torch.models.modeling_utils.DecoderModel", # The text model's stock modules. "tensorrt_llm._torch.modules.kimi_kda.KimiKDALinearAttention", "tensorrt_llm._torch.modules.kimi_k3_mla.KimiK3MLAAttention", "tensorrt_llm._torch.moe.fused_moe.create_moe", + "tensorrt_llm._torch.moe.fused_moe.ConfigurableMoE", + "tensorrt_llm._torch.moe.fused_moe.TRTLLMGenFusedMoE", + "tensorrt_llm._torch.moe.fused_moe.routing.DeepSeekV3MoeRoutingMethod", "tensorrt_llm._torch.modules.gated_mlp.GatedMLP", "tensorrt_llm._torch.modules.situ.SituAndMul", "tensorrt_llm._torch.modules.rms_norm.RMSNorm", diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py index 1f7e0435d5b3..bd8376b2dcea 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py @@ -17,6 +17,7 @@ from __future__ import annotations +import ast import re from pathlib import Path @@ -215,3 +216,70 @@ def test_targets_do_not_share_files(): f"directory for {tail!r}; targets may import the " f"catalog and nothing else" ) + + +#: tensorrt_llm imports that compute nothing, so a target's generic path need not declare them: types and configs +#: it reads, enums it passes, the metadata it is handed, its registration and its logger. +_NON_COMPUTE_IMPORTS = frozenset( + { + "AllReduceStrategy", + "AttentionMetadata", + "AuxStreamType", + "MambaHybridCacheManagerV2", + "Mapping", + "ModelConfig", + "MoESchedulerKind", + "QuantAlgo", + "QuantConfig", + "SiTuActivation", + "logger", + "register_auto_model", + } +) + + +def _declared_tuple(tree: ast.Module, name: str): + """The string tuple assigned to ``name`` at module level, or None when the module does not assign it.""" + for node in tree.body: + if isinstance(node, ast.Assign) and any( + isinstance(target, ast.Name) and target.id == name for target in node.targets + ): + return tuple(ast.literal_eval(node.value)) + return None + + +def test_uncertified_generic_calls_name_every_stock_import(): + """A target whose generic path runs stock code (``UNCERTIFIED_GENERIC_CALLS``) declares every tensorrt_llm + name it imports outside the catalog, module level or function local, unless the name computes nothing; and it + declares nothing it does not import, so the list cannot go stale as entries replace stock calls.""" + checked = 0 + for arch in _ARCHS: + routing = routing_module(arch) + for name, dotted in routing.TARGET_MODULES.items(): + tree = ast.parse(_module_path(dotted).read_text()) + declared = _declared_tuple(tree, "UNCERTIFIED_GENERIC_CALLS") + if declared is None: + continue + checked += 1 + imported = set() + for node in ast.walk(tree): + if ( + isinstance(node, ast.ImportFrom) + and node.level == 0 + and node.module is not None + and node.module.split(".")[0] == "tensorrt_llm" + and ".modeling_v2.catalog" not in node.module + ): + imported.update( + f"{node.module}.{alias.name}" + for alias in node.names + if alias.name not in _NON_COMPUTE_IMPORTS + ) + assert imported <= set(declared), ( + f"{name}: imported but not in UNCERTIFIED_GENERIC_CALLS: {sorted(imported - set(declared))}" + ) + assert set(declared) <= imported, ( + f"{name}: UNCERTIFIED_GENERIC_CALLS names what it does not import: " + f"{sorted(set(declared) - imported)}" + ) + assert checked, "no target declares UNCERTIFIED_GENERIC_CALLS" From b28044256ddf95b484ed3d0910b3f4dbc5272652 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 19:21:57 -0700 Subject: [PATCH 036/161] [None][test] modeling_v2 Kimi K3 decode entries: sm_100 receipts The ten Kimi K3 entries get their sm_100 receipts. gemm/cublas_mm and norm/flashinfer_rmsnorm, which the Kimi K3 target also calls, get sm_100 keys; their sm_103 receipts stay, since they were taken on the same, unchanged test files. Each count is the entry test's pytest count on GB200 on this branch's files. Signed-off-by: Vasanth Sabavat --- .../modeling_v2/catalog/activation/k3_situ_mul.md | 3 ++- .../_torch/_experimental/modeling_v2/catalog/gemm/cublas_mm.md | 2 ++ .../_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv.md | 3 ++- .../_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_long.md | 3 ++- .../modeling_v2/catalog/gemm/k3_ctm_gemv_swiglu.md | 3 ++- .../_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_wide.md | 3 ++- .../_experimental/modeling_v2/catalog/gemm/k3_decode_gemv.md | 3 ++- .../_experimental/modeling_v2/catalog/gemm/k3_head_gemv.md | 3 ++- .../_experimental/modeling_v2/catalog/norm/attn_res_fwd.md | 3 ++- .../modeling_v2/catalog/norm/attn_res_rmsnorm_fwd.md | 3 ++- .../modeling_v2/catalog/norm/flashinfer_rmsnorm.md | 2 ++ .../_experimental/modeling_v2/catalog/norm/k3_embed_norm.md | 3 ++- 12 files changed, 24 insertions(+), 10 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/activation/k3_situ_mul.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/activation/k3_situ_mul.md index febbd25c5bb5..14f89ee0710f 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/activation/k3_situ_mul.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/activation/k3_situ_mul.md @@ -1,5 +1,6 @@ --- -receipts: {} +receipts: + sm_100: {status: passed, tests: 4} --- # k3_situ_mul diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/cublas_mm.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/cublas_mm.md index 7b32f301dc29..7044e9449ced 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/cublas_mm.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/cublas_mm.md @@ -1,6 +1,8 @@ --- receipts: + # sm_103 was certified before the sm_100 key was added, on the same test file. sm_103: {status: passed, tests: 5} + sm_100: {status: passed, tests: 5} --- # cublas_mm diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv.md index 3d706c5d732d..13497756a6b3 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv.md @@ -1,5 +1,6 @@ --- -receipts: {} +receipts: + sm_100: {status: passed, tests: 4} --- # k3_ctm_gemv diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_long.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_long.md index a455bf61744b..2af2cbcde62c 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_long.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_long.md @@ -1,5 +1,6 @@ --- -receipts: {} +receipts: + sm_100: {status: passed, tests: 4} --- # k3_ctm_gemv_long diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_swiglu.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_swiglu.md index 8ff9e8fb1a4f..560710673676 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_swiglu.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_swiglu.md @@ -1,5 +1,6 @@ --- -receipts: {} +receipts: + sm_100: {status: passed, tests: 4} --- # k3_ctm_gemv_swiglu diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_wide.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_wide.md index 1c7a2c0347d7..d2e27a39d893 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_wide.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_ctm_gemv_wide.md @@ -1,5 +1,6 @@ --- -receipts: {} +receipts: + sm_100: {status: passed, tests: 4} --- # k3_ctm_gemv_wide diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_decode_gemv.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_decode_gemv.md index f58fa5919dfa..2ee95e8c7c17 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_decode_gemv.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_decode_gemv.md @@ -1,5 +1,6 @@ --- -receipts: {} +receipts: + sm_100: {status: passed, tests: 4} --- # k3_decode_gemv diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_head_gemv.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_head_gemv.md index 3ea9ebc47f51..85ace6ac6cb2 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_head_gemv.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/k3_head_gemv.md @@ -1,5 +1,6 @@ --- -receipts: {} +receipts: + sm_100: {status: passed, tests: 5} --- # k3_head_gemv diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_fwd.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_fwd.md index 30a43426bb80..26a0f9da0ab0 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_fwd.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_fwd.md @@ -1,5 +1,6 @@ --- -receipts: {} +receipts: + sm_100: {status: passed, tests: 5} --- # attn_res_fwd diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_rmsnorm_fwd.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_rmsnorm_fwd.md index 71657064b3b7..b07775ef94e8 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_rmsnorm_fwd.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/attn_res_rmsnorm_fwd.md @@ -1,5 +1,6 @@ --- -receipts: {} +receipts: + sm_100: {status: passed, tests: 6} --- # attn_res_rmsnorm_fwd diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/flashinfer_rmsnorm.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/flashinfer_rmsnorm.md index 490f1d9f71ac..05373dea5da6 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/flashinfer_rmsnorm.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/flashinfer_rmsnorm.md @@ -1,6 +1,8 @@ --- receipts: + # sm_103 was certified before the sm_100 key was added, on the same test file. sm_103: {status: passed, tests: 4} + sm_100: {status: passed, tests: 4} --- # flashinfer_rmsnorm diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/k3_embed_norm.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/k3_embed_norm.md index b1432ba151f1..bda71f5dba06 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/k3_embed_norm.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/k3_embed_norm.md @@ -1,5 +1,6 @@ --- -receipts: {} +receipts: + sm_100: {status: passed, tests: 5} --- # k3_embed_norm From e9b35956d8bb9a41f061f72afb9c42a7d1202334 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 19:42:06 -0700 Subject: [PATCH 037/161] [None][feat] modeling_v2: Kimi K3 MXFP4 target skeleton, sm_100, tp16 + moe tp16 x ep1 Route KimiK3ForConditionalGeneration to a second target, kimi_k3_mxfp4__sm_100__tp16_moetp16ep1: the same checkpoint, SM and attention layout, with the routed experts split 16 ways by tensor (every rank holds all 896 experts at a sixteenth of their width). The split counts only when the configuration sets it: an unset split resolves to the same sizes, but the built-in Kimi K3 model reads that default as expert parallelism over the 16 ranks (moe_tp_ep_user_specified), and so does the routing tree. No target serves that layout. The target is route A's skeleton with the route blocks changed. It asserts the explicit split and no speculative decoding at construction (DSpark runs on the tp16_moetp4ep4 target), and its fused decode path is bounded at 8 tokens. Every step still runs the built-in Kimi K3 text model; the fused path's engines (k3_moe_m1 / k3_moe_m2 at one and two tokens, k3_moe pushing into the latent exchange at 3-8) come with their catalog entries. Routing tests: route B matches with the split set, in auto and require, from a text_config object or dict. Both targets register as external. The near misses add the unset split, an explicit EP16 split, attention DP on route B's split, and pipeline parallelism. Signed-off-by: Vasanth Sabavat --- .../__init__.py | 3 + .../modeling.py | 277 ++++++++++++++++++ .../weights.py | 42 +++ .../modeling_v2/models/kimi_k3_vl/routing.py | 34 ++- .../modeling_v2/test_modeling_v2_routing.py | 37 ++- 5 files changed, 376 insertions(+), 17 deletions(-) create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/__init__.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/weights.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/__init__.py new file mode 100644 index 000000000000..2b24a35badfd --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/__init__.py @@ -0,0 +1,3 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3 (MXFP4) / sm_100 / tp16 attention, routed experts moe_tp 16 x moe_ep 1.""" diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py new file mode 100644 index 000000000000..3c86aabb0b0e --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py @@ -0,0 +1,277 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""ModelingV2 target: Kimi K3 (MXFP4) / sm_100 / tp16 attention, routed experts moe_tp 16 x moe_ep 1. + +Kimi K3's language model: 93 layers, hidden 7168. Every fourth layer (3, 7, ..., 91) is MLA attention with 96 +query heads; the others are Kimi Delta Attention (KDA), a gated linear-attention recurrence with per-request state. +Layer 0's MLP is dense. Every other layer's MLP is a latent MoE: 896 routed experts, top-16, expert width 3584. +Attention residuals mix each sublayer's input from a bank of earlier outputs. The routed experts are MXFP4 (the +checkpoint's compressed-tensors default, read as W4A8_MXFP4_MXFP8); everything else is bf16, and so is the KV pool. + +`tp16_moetp16ep1` is `tensor_parallel_size: 16` with `moe_tensor_parallel_size: 16` and +`moe_expert_parallel_size: 1`, both set explicitly, no attention data parallelism, all 16 GPUs in one NVLink domain. +Attention is head-split (6 MLA query heads per rank), and every rank holds all 896 experts at a sixteenth of their +width (192 of 3072 intermediate values, zero-padded to 256 by the loader). With the expert split left unset the +built-in model runs the experts expert-parallel over the 16 ranks instead, and no target serves that layout. + +**Each step takes one of two paths, chosen on the host from the step's shape** (`_step_path`): + +* **The fused decode path**: pure decode steps of at most 8 tokens, on the K3 decode kernels' catalog entries. The + routed experts of one and two tokens run as `moe/k3_moe_m1` and `moe/k3_moe_m2`, and of 3-8 tokens as k3_moe, each + pushing its partial into the latent exchange that `trtllm::k3_latent_reduce` sums. The state those kernels share + (MNNVL workspace, sandwich and MoE Lamport buffers, the latent exchange, the engines' workspaces, KDA / MLA scratch) + lives in typed objects this target creates in `post_load_weights`, before any graph capture. Until those entries + are wired `_fused_decode` stays None, and every step takes the generic path. +* **The generic path**: prefill, mixed steps, and decode steps above those bounds, on the built-in Kimi K3 text + model, whose modules and ops have no catalog entries yet. `UNCERTIFIED_GENERIC_CALLS` names them. + +**What this target asserts rather than adapts**: SM 10.0; the topology above, with its expert split explicit; no +speculative decoding; bf16 weights and a bf16 KV pool; +tokens_per_block 64 (the MLA generation kernels K3's 96 heads reach exist only at 64); the V2 hybrid KV / state +manager, which holds the KDA states, with block reuse off; an all-reduce strategy of AUTO or MNNVL. The +construction-time ones fail in `__init__`, the per-engine ones on the first forward, each naming the setting. + +**Text only.** The checkpoint is the vision-language wrapper. This target builds and loads no vision tower (its +weights are a predicted non-load, `weights.py`), and a step carrying multimodal input raises. + +**No speculative decoding.** This layout is the low-latency route for decoding without a drafter; DSpark runs on +the `tp16_moetp4ep4` target. A configuration with a speculative decoding config stops at construction. +""" + +import copy +from typing import Any, Literal, Optional + +import torch + +from tensorrt_llm._torch.attention.backends.interface import AttentionMetadata +from tensorrt_llm._torch.model_config import ModelConfig +from tensorrt_llm._torch.models.modeling_kimi_linear import KimiLinearForCausalLM +from tensorrt_llm._torch.models.modeling_utils import register_auto_model +from tensorrt_llm._torch.pyexecutor.kv_cache.mamba_cache_manager import MambaHybridCacheManagerV2 +from tensorrt_llm.functional import AllReduceStrategy + +from . import weights as _weights + +# The GPU architecture this target IS. Routing will not send another one here, but a direct instantiation could, +# and the certification is per arch. +_SM = (10, 0) + +#: Every trtllm op this target reaches for, in its forward and in the weight load. Declared here, asserted in +#: tests/unittest/_torch/modeling_v2. Today these are the K3-specific ops of the generic path (attention residuals, +#: KDA, the router and fused-A GEMMs); the fused decode path adds its own. +REQUIRED_TRTLLM_OPS = ( + "attn_res_fwd", + "attn_res_rmsnorm_fwd", + "attn_res_add_rmsnorm_fwd", + "attn_res_add_rmsnorm_persistent_fwd", + "kda_prefill", + "kda_decode", + "kda_mtp_decode", + "dsv3_router_gemm_op", + "dsv3_fused_a_gemm_op", +) + +#: The engine surface the first forward checks before this target relies on it: per object, the attributes read. +#: A renamed field upstream fails here, loudly, instead of reading as a default. +REQUIRED_ENGINE_FIELDS = { + "attn_metadata": ( + "num_contexts", + "num_seqs", + "num_tokens", + "tokens_per_block", + "kv_cache_manager", + ), + "kv_cache_manager": ("enable_block_reuse",), +} + +#: Calls the generic path makes outside the catalog, declared so they are not consumed silently. A call leaves this +#: list when a catalog entry replaces it. +UNCERTIFIED_GENERIC_CALLS = ( + "tensorrt_llm._torch.models.modeling_kimi_linear.KimiLinearForCausalLM", +) + +# The fused decode path's token bound per step: one token per request (the K3 decode kernels are built for up to 8 +# rows). +_FUSED_MAX_TOKENS = 8 + +# The MLA generation kernels for K3's 96 query heads exist only at a 64-token page (the built-in model's own +# get_model_defaults sets it for the same reason). +_TOKENS_PER_BLOCK = 64 + +_LANG_PREFIX = "language_model." + + +def _text_model_config(model_config: ModelConfig) -> ModelConfig: + """The language model's ModelConfig: the checkpoint's text_config, with quant exclusions renamed to match. + + The checkpoint names its language-model modules `language_model.`; the text model's are ``, with + `layers.*` under `model.`. + """ + config = model_config.pretrained_config + text = copy.copy(model_config) + text._frozen = False + text.pretrained_config = config.text_config + excluded = text.quant_config.exclude_modules + if excluded: + text.quant_config = copy.copy(text.quant_config) + renamed = [] + for name in excluded: + if name.startswith(_LANG_PREFIX): + name = name[len(_LANG_PREFIX) :] + if name.startswith("layers."): + name = "model." + name + renamed.append(name) + text.quant_config.exclude_modules = renamed + text.skip_create_weights_in_init = True + text._frozen = True + return text + + +def _check_construction(model_config: ModelConfig) -> None: + """The settings this target is built for that are fixed before the first step.""" + capability = torch.cuda.get_device_capability() + assert capability == _SM, ( + f"this target is certified on sm_{_SM[0]}{_SM[1]}, running on sm_{capability[0]}{capability[1]}" + ) + mapping = model_config.mapping + topology = ( + mapping.world_size, + mapping.tp_size, + mapping.pp_size, + mapping.moe_tp_size, + mapping.moe_ep_size, + mapping.enable_attention_dp, + ) + assert topology == (16, 16, 1, 16, 1, False), ( + "the tp16_moetp16ep1 target needs world_size 16, tensor_parallel_size 16, pipeline_parallel_size 1, " + "moe_tensor_parallel_size 16, moe_expert_parallel_size 1 and enable_attention_dp false; the engine built " + f"(world, tp, pp, moe_tp, moe_ep, attention_dp) = {topology}" + ) + assert getattr(mapping, "moe_tp_ep_user_specified", False), ( + "the tp16_moetp16ep1 target needs moe_tensor_parallel_size 16 and moe_expert_parallel_size 1 set explicitly; " + "with the split unset the built-in Kimi K3 model runs the experts expert-parallel over the 16 ranks" + ) + assert model_config.spec_config is None, ( + "the tp16_moetp16ep1 target decodes without speculation; DSpark runs on the tp16_moetp4ep4 target " + f"(moe_tensor_parallel_size 4, moe_expert_parallel_size 4); got {type(model_config.spec_config).__name__}" + ) + assert model_config.torch_dtype == torch.bfloat16, ( + f"this target computes in bf16; the engine resolved dtype {model_config.torch_dtype}" + ) + kv_algo = model_config.quant_config.kv_cache_quant_algo + assert kv_algo is None, ( + f"this target's MLA kernels read a bf16 KV pool; kv_cache_config.dtype resolved to {kv_algo}" + ) + strategy = model_config.allreduce_strategy + assert strategy in (AllReduceStrategy.AUTO, AllReduceStrategy.MNNVL), ( + f"this target runs its all-reduces over MNNVL; allreduce_strategy is {strategy.name}" + ) + + +@register_auto_model("ModelingV2KimiK3Mxfp4Sm100Tp16Moetp16ep1") +class ModelingV2KimiK3Mxfp4Sm100Tp16Moetp16ep1(KimiLinearForCausalLM): + """The registration shell: the built-in Kimi K3 text model as the generic path, behind this target's checks.""" + + @classmethod + def get_preferred_kv_cache_manager_version(cls, pretrained_config: Any = None) -> Literal["V2"]: + """The V2 hybrid manager holds the KDA states; the step contract requires it.""" + return "V2" + + def __init__(self, model_config: ModelConfig): + config = model_config.pretrained_config + assert getattr(config, "text_config", None) is not None, ( + "this target loads the KimiK3ForConditionalGeneration checkpoint, whose language model is its " + "text_config" + ) + _check_construction(model_config) + super().__init__(_text_model_config(model_config)) + self._step_checked = False + # The fused decode path and the state its kernels share, built in post_load_weights once the catalog entries + # it calls exist. None: every step takes the generic path. + self._fused_decode = None + # The executor reads generation settings (eos_token_id, ...) off the model config the engine holds, which + # must therefore be the text config, as the built-in wrapper leaves it. + model_config._frozen = False + model_config.pretrained_config = self.config + model_config._frozen = True + + def load_weights(self, weights, *args, **kwargs): + _weights.load(self, weights) + + def _check_step_contract(self, attn_metadata: AttentionMetadata) -> None: + """First-forward checks of the engine surface and the per-engine settings.""" + objects = { + "attn_metadata": attn_metadata, + "kv_cache_manager": attn_metadata.kv_cache_manager, + } + missing = [ + f"{owner}.{name}" + for owner, names in REQUIRED_ENGINE_FIELDS.items() + for name in names + if not hasattr(objects[owner], name) + ] + assert not missing, f"engine fields this target reads are missing: {missing}" + assert attn_metadata.tokens_per_block == _TOKENS_PER_BLOCK, ( + f"this target needs kv_cache_config.tokens_per_block {_TOKENS_PER_BLOCK}; the engine built " + f"{attn_metadata.tokens_per_block}" + ) + manager = attn_metadata.kv_cache_manager + assert isinstance(manager, MambaHybridCacheManagerV2), ( + "this target needs the V2 hybrid KV / state cache manager " + "(kv_cache_config.use_kv_cache_manager_v2); the engine built " + f"{type(manager).__name__}" + ) + assert not manager.enable_block_reuse, ( + "this target runs with kv_cache_config.enable_block_reuse false; the engine enabled it" + ) + self._step_checked = True + + def _step_path(self, attn_metadata: AttentionMetadata, spec_metadata) -> str: + """`"fused"` for a pure decode step within the fused path's bounds, else `"generic"`. + + Read on the host from per-step integers only. A CUDA graph is captured per decode batch shape, and every + input here is fixed by that shape, so a captured step and its replays take the same path. + """ + if attn_metadata.num_contexts: + return "generic" + return "fused" if attn_metadata.num_tokens <= _FUSED_MAX_TOKENS else "generic" + + def forward( + self, + attn_metadata: AttentionMetadata, + input_ids: Optional[torch.IntTensor] = None, + position_ids: Optional[torch.IntTensor] = None, + inputs_embeds: Optional[torch.FloatTensor] = None, + return_context_logits: bool = False, + spec_metadata=None, + resource_manager=None, + **kwargs, + ) -> torch.Tensor: + if kwargs.pop("multimodal_params", None): + raise ValueError( + "this Kimi K3 target is text only: it loads no vision tower, and a request carried image input" + ) + if not self._step_checked: + self._check_step_contract(attn_metadata) + if ( + self._fused_decode is not None + and self._step_path(attn_metadata, spec_metadata) == "fused" + ): + return self._fused_decode( + attn_metadata=attn_metadata, + input_ids=input_ids, + position_ids=position_ids, + spec_metadata=spec_metadata, + resource_manager=resource_manager, + **kwargs, + ) + return super().forward( + attn_metadata, + input_ids, + position_ids, + inputs_embeds, + return_context_logits, + spec_metadata, + resource_manager, + **kwargs, + ) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/weights.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/weights.py new file mode 100644 index 000000000000..e4c24eaaea22 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/weights.py @@ -0,0 +1,42 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Weight loading: Kimi K3 (MXFP4) / sm_100 / tp16_moetp16ep1. + +The checkpoint is the vision-language wrapper's. `language_model.*` holds the language model; `vision_tower.*` and +`mm_projector.*` hold the vision tower and its projector. This target is text only, so the keys split three ways: + +* `language_model.*` goes to the built-in text model's loader with the prefix stripped. That loader streams the + routed experts one at a time, keeps this rank's slice (a sixteenth of the width of every expert under moe_tp 16 x + moe_ep 1: 192 of 3072, zero-padded to 256 for the MoE kernels' tiles), and checks its own key coverage. +* The vision tower and the projector are a predicted non-load: listed here, never read. +* Any other key fails the load, naming it, rather than being dropped. +""" + +from tensorrt_llm._torch.models.checkpoints.base_weight_loader import ConsumableWeightsDict +from tensorrt_llm._torch.models.modeling_kimi_linear import KimiLinearForCausalLM +from tensorrt_llm._torch.models.modeling_utils import filter_weights + +_LANG_PREFIX = "language_model." + +# The vision tower's and the projector's key families: in the checkpoint, not loaded by this text-only target. +PREDICTED_NON_LOAD = ("vision_tower.", "mm_projector.") + + +def load(model, weights) -> None: + """Load the language model's weights into `model`, the target shell, and check every other key is predicted.""" + unknown = sorted( + k + for k in weights.keys() + if not k.startswith(_LANG_PREFIX) and not k.startswith(PREDICTED_NON_LOAD) + ) + assert not unknown, ( + f"{len(unknown)} checkpoint key(s) are neither language-model weights nor a predicted non-load, " + f"first {unknown[:5]}" + ) + lm_weights = ConsumableWeightsDict(filter_weights(_LANG_PREFIX[:-1], weights)) + assert len(lm_weights), f"the checkpoint has no {_LANG_PREFIX}* keys" + checkpoint_dir = getattr(weights, "checkpoint_dir", None) + if checkpoint_dir is not None: + lm_weights.checkpoint_dir = checkpoint_dir + lm_weights.checkpoint_prefix = _LANG_PREFIX + KimiLinearForCausalLM.load_weights(model, lm_weights) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/routing.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/routing.py index 048c790a0b84..18ea3225aeaa 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/routing.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/routing.py @@ -35,6 +35,7 @@ _TARGETS = { ("kimi_k3_mxfp4", "tp16_moetp4ep4"): "ModelingV2KimiK3Mxfp4Sm100Tp16Moetp4ep4", + ("kimi_k3_mxfp4", "tp16_moetp16ep1"): "ModelingV2KimiK3Mxfp4Sm100Tp16Moetp16ep1", } # Synthetic architecture name -> the module whose import registers it. @@ -42,6 +43,9 @@ "ModelingV2KimiK3Mxfp4Sm100Tp16Moetp4ep4": ( "models.kimi_k3_vl.kimi_k3_mxfp4__sm_100__tp16_moetp4ep4.modeling" ), + "ModelingV2KimiK3Mxfp4Sm100Tp16Moetp16ep1": ( + "models.kimi_k3_vl.kimi_k3_mxfp4__sm_100__tp16_moetp16ep1.modeling" + ), } @@ -67,20 +71,25 @@ def _mxfp4(quant_config: Any) -> bool: def _parallel(m) -> Optional[str]: """Name the parallel topology, or None if no target implements it. - Attention and the dense layers are sharded 16 ways; the routed experts are - split 4 ways by tensor and 4 ways by expert. The expert split decides which - expert weights each rank loads and into which shapes, so it selects a - target rather than a runtime branch. + Attention and the dense layers are sharded 16 ways. The routed experts are + split either 4 ways by tensor and 4 ways by expert, or 16 ways by tensor + with every expert on every rank. The expert split decides which expert + weights each rank loads and into which shapes, so it selects a target + rather than a runtime branch. + + The 16-way tensor split counts only when the configuration asks for it. A + mapping with no expert split resolves to the same sizes, but the built-in + Kimi K3 model reads that default as expert parallelism over the 16 ranks + (`moe_tp_ep_user_specified`), and so does this tree: no target serves it. """ - if ( - m.world_size == 16 - and m.tp_size == 16 - and m.pp_size == 1 - and m.moe_tp_size == 4 - and m.moe_ep_size == 4 - and not m.enable_attention_dp + if not ( + m.world_size == 16 and m.tp_size == 16 and m.pp_size == 1 and not m.enable_attention_dp ): + return None + if m.moe_tp_size == 4 and m.moe_ep_size == 4: return "tp16_moetp4ep4" + if m.moe_tp_size == 16 and m.moe_ep_size == 1 and getattr(m, "moe_tp_ep_user_specified", False): + return "tp16_moetp16ep1" return None @@ -114,7 +123,8 @@ def route(ctx: ModelingV2Context, trace: Trace = NULL_TRACE) -> Optional[str]: parallel = trace.resolve( "parallel", f"ws={m.world_size} tp={m.tp_size} pp={m.pp_size} moe_tp={m.moe_tp_size} " - f"moe_ep={m.moe_ep_size} attention_dp={m.enable_attention_dp}", + f"moe_ep={m.moe_ep_size} split_set={getattr(m, 'moe_tp_ep_user_specified', False)} " + f"attention_dp={m.enable_attention_dp}", _parallel(m), ) if parallel is None: diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py index 9728ba55d985..b785587fb973 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py @@ -239,7 +239,9 @@ def test_an_unknown_mode_raises_rather_than_falling_back(monkeypatch): _SM100 = (10, 0) _TP16_MOETP4EP4 = dict(world_size=16, tp_size=16, moe_tp_size=4, moe_ep_size=4) +_TP16_MOETP16EP1 = dict(world_size=16, tp_size=16, moe_tp_size=16, moe_ep_size=1) _K3_TARGET = "ModelingV2KimiK3Mxfp4Sm100Tp16Moetp4ep4" +_K3_TARGET_B = "ModelingV2KimiK3Mxfp4Sm100Tp16Moetp16ep1" def _k3_config(text_as_dict=False, **text_overrides): @@ -276,10 +278,25 @@ def test_kimi_k3_tp16_moetp4ep4_matches(monkeypatch, mode, text_as_dict): @pytest.mark.usefixtures("_on_sm100") -def test_kimi_k3_target_registers_and_counts_as_external(): - config = _model_config(_k3_config(), **_TP16_MOETP4EP4) +@pytest.mark.parametrize("mode", ["auto", "require"]) +@pytest.mark.parametrize("text_as_dict", [False, True], ids=["text-config", "text-dict"]) +def test_kimi_k3_tp16_moetp16ep1_matches(monkeypatch, mode, text_as_dict): + """Route B: the experts split 16 ways by tensor, set explicitly.""" + _set_mode(monkeypatch, mode) + config = _model_config(_k3_config(text_as_dict=text_as_dict), **_TP16_MOETP16EP1) + assert modeling_v2_resolve(config) == _K3_TARGET_B + + +@pytest.mark.usefixtures("_on_sm100") +@pytest.mark.parametrize( + "mapping_kwargs, target", + [(_TP16_MOETP4EP4, _K3_TARGET), (_TP16_MOETP16EP1, _K3_TARGET_B)], + ids=["tp16_moetp4ep4", "tp16_moetp16ep1"], +) +def test_kimi_k3_target_registers_and_counts_as_external(mapping_kwargs, target): + config = _model_config(_k3_config(), **mapping_kwargs) cls = get_registered_model_class(modeling_v2_resolve(config)) - assert cls is not None and cls.__name__ == _K3_TARGET + assert cls is not None and cls.__name__ == target assert not _is_builtin_model_class(cls) @@ -308,10 +325,20 @@ def test_kimi_k3_nvfp4_requant_does_not_match(monkeypatch): [ # a Kimi-family checkpoint of another depth (dict(num_hidden_layers=61), _TP16_MOETP4EP4, "shape"), - # route B's expert split: experts tensor-parallel 16 ways, no target yet - (dict(), dict(world_size=16, tp_size=16, moe_tp_size=16, moe_ep_size=1), "parallel"), + # the expert split left unset: the mapping resolves to moe_tp 16 x moe_ep 1, but the built-in model runs + # that default as expert parallelism over the 16 ranks, a layout no target serves + (dict(), dict(world_size=16, tp_size=16), "parallel"), + # experts split 16 ways by expert, set explicitly + (dict(), dict(world_size=16, tp_size=16, moe_tp_size=1, moe_ep_size=16), "parallel"), # attention data parallelism splits the requests, not the heads (dict(), dict(_TP16_MOETP4EP4, enable_attention_dp=True), "parallel"), + (dict(), dict(_TP16_MOETP16EP1, enable_attention_dp=True), "parallel"), + # pipeline parallelism over the 16 ranks + ( + dict(), + dict(world_size=16, tp_size=8, pp_size=2, moe_tp_size=8, moe_ep_size=1), + "parallel", + ), # one tray instead of four (dict(), dict(world_size=4, tp_size=4, moe_tp_size=1, moe_ep_size=4), "parallel"), ], From 75d8a597ece078011378243fa2a3177b0ff91e6d Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 19:48:10 -0700 Subject: [PATCH 038/161] [None][feat] modeling_v2 Kimi K3 target: KDA and MLA attention on the decode kernels The text model hands each step's classification (decode_step) to its attention modules, which run a decode step on the KDA / MLA decode kernels' catalog entries: - KDA (K3DecodeKDA, the built-in module subclassed): a decode step of one token per request runs the fused projection and the plain decode in one ssm/k3_kda_decode_attn launch. With DFlash / DSpark at an even verify width up to 8 the model asks the cache manager for the per-token KDA states (kda_token_states), and every verify of a KDA layer, on any step, runs the kernels that keep them: ssm/k3_kda_attn for one request of 8 tokens, else ssm/k3_kda_verify on the projection's rows. The [q | k | v | g | f_a | b] weight they read is built at the checkpoint load; the projections' weights and the built-in fused ones become views of it, so no weight is stored twice. - MLA (K3DecodeMLA): x [W_a; W_g]^T, attention/k3_mla_qkv and attention/k3_mla_attn_vb_out (the attention, v_b and the output gate in one launch), then the module's o_proj. [W_a; W_g] is built in post_load_weights. The target's post_load_weights creates the state the kernels share (K3KdaBuffers, K3MlaAttnWorkspace) once per device and hands it to every layer of its kind. Steps decode_step does not classify, and every step under a breakable CUDA graph, run the built-in paths. A layer the kernels do not take fails the load, and the first forward asserts an fp32 KDA state pool. The projections around the attention kernels stay torch / stock code until the GEMV entries take them. The claims test counts is_in_breakable_cuda_graph and k3_mla_decode_view as imports that compute nothing. Signed-off-by: Vasanth Sabavat --- .../modeling.py | 470 ++++++++++++++++-- .../modeling_v2/test_modeling_v2_claims.py | 5 +- 2 files changed, 438 insertions(+), 37 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index 6cf3cd599c97..6a2c09f7159a 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -27,16 +27,29 @@ Every other step (prefill, mixed steps, decode steps above those bounds) runs the **generic path**: this target's text model (`KimiLinearModel` below: decoder layers, attention residuals, the MLA / KDA / MoE runtimes), computed exactly as the built-in Kimi K3 text model computes it, on stock modules and ops that have no catalog entries yet. -`UNCERTIFIED_GENERIC_CALLS` names them. The **fused decode path** runs the steps `decode_step` classifies on the K3 -decode kernels' catalog entries. The state those -kernels share (MNNVL workspace, sandwich and MoE Lamport buffers, KDA / MLA scratch) lives in typed objects this -target creates collectively in `post_load_weights`, before any graph capture. Until those entries exist -`_fused_decode` stays None, and every step takes the generic path. +`UNCERTIFIED_GENERIC_CALLS` names them. + +The text model hands each step's classification to its attention modules, which run a **decode step** on the K3 +decode kernels' catalog entries: + +* KDA (`K3DecodeKDA`): one token per request, the fused input projection and the plain decode in one + `ssm/k3_kda_decode_attn` launch. Verify tokens (DFlash / DSpark at an even verify width up to 8): the cache manager + then keeps the KDA state after every verify token, and every verify of the layer, on any step, runs the kernels + that keep it: `ssm/k3_kda_attn` for one request of 8 tokens (the projection fused in), else `ssm/k3_kda_verify`. +* MLA (`K3DecodeMLA`): `attention/k3_mla_qkv` (the query path and the step's latent KV rows into the paged cache), + then `attention/k3_mla_attn_vb_out` (the attention, v_b and the output gate in one launch). + +The state those kernels share (the KDA projection's Lamport buffers, the MLA attention workspace) lives in typed +objects this target creates in `post_load_weights`, before any graph capture. The token-count kernels (the decode +GEMVs, the MoE front and routed experts, the sandwiches, the embedding and residual epilogues) come with their own +entries; until then every module but the attention runs the generic path on every step, as do the projections around +the attention kernels (the [W_a; W_g] and verify-row GEMMs, `o_proj` and its all-reduce). **What this target asserts rather than adapts**: SM 10.0; the topology above; bf16 weights and a bf16 KV pool; tokens_per_block 64 (the MLA generation kernels K3's 96 heads reach exist only at 64); the V2 hybrid KV / state -manager, which holds the KDA states, with block reuse off; an all-reduce strategy of AUTO or MNNVL. The -construction-time ones fail in `__init__`, the per-engine ones on the first forward, each naming the setting. +manager, which holds the KDA states, with block reuse off and fp32 recurrent states; an all-reduce strategy of AUTO or +MNNVL. The construction-time ones fail in `__init__`, the per-engine ones on the first forward, each naming the +setting. A layer the decode kernels do not take fails the weight load. **Text only.** The checkpoint is the vision-language wrapper. This target builds and loads no vision tower (its weights are a predicted non-load, `weights.py`), and a step carrying multimodal input raises. @@ -56,13 +69,28 @@ import torch from torch import nn +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.k3_mla_attn_vb_out import ( + k3_mla_attn_vb_out, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.k3_mla_attn_workspace import ( + K3MlaAttnWorkspace, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.k3_mla_qkv import k3_mla_qkv +from tensorrt_llm._torch._experimental.modeling_v2.catalog.ssm.k3_kda_attn import k3_kda_attn +from tensorrt_llm._torch._experimental.modeling_v2.catalog.ssm.k3_kda_buffers import K3KdaBuffers +from tensorrt_llm._torch._experimental.modeling_v2.catalog.ssm.k3_kda_decode_attn import ( + k3_kda_decode_attn, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.ssm.k3_kda_verify import k3_kda_verify from tensorrt_llm._torch.attention.backends import AttentionMetadata +from tensorrt_llm._torch.attention.backends.fmha.cute_dsl_mla import k3_mla_decode_view from tensorrt_llm._torch.distributed import AllReduce from tensorrt_llm._torch.model_config import ModelConfig from tensorrt_llm._torch.models.modeling_kimi_linear import KimiLinearForCausalLM from tensorrt_llm._torch.models.modeling_speculative import SpecDecOneEngineForCausalLM from tensorrt_llm._torch.models.modeling_utils import DecoderModel, register_auto_model from tensorrt_llm._torch.modules.gated_mlp import GatedMLP +from tensorrt_llm._torch.modules.kimi_k3_mla import KimiK3MLAAttention from tensorrt_llm._torch.modules.kimi_kda import KimiKDALinearAttention from tensorrt_llm._torch.modules.multi_stream_utils import maybe_execute_in_parallel from tensorrt_llm._torch.modules.rms_norm import RMSNorm @@ -75,6 +103,7 @@ ) from tensorrt_llm._torch.moe.fused_moe.interface import MoESchedulerKind from tensorrt_llm._torch.moe.fused_moe.routing import DeepSeekV3MoeRoutingMethod +from tensorrt_llm._torch.pyexecutor.breakable_cuda_graph import is_in_breakable_cuda_graph from tensorrt_llm._torch.pyexecutor.kv_cache.mamba_cache_manager import MambaHybridCacheManagerV2 from tensorrt_llm._torch.utils import AuxStreamType from tensorrt_llm.functional import AllReduceStrategy @@ -93,8 +122,8 @@ _SM = (10, 0) #: Every trtllm op this target reaches for, in its forward and in the weight load. Declared here, asserted in -#: tests/unittest/_torch/modeling_v2. Today these are the K3-specific ops of the generic path (attention residuals, -#: KDA, the router and fused-A GEMMs); the fused decode path adds its own. +#: tests/unittest/_torch/modeling_v2: the K3-specific ops of the generic path (attention residuals, KDA, the router +#: and fused-A GEMMs), then the decode kernels'. REQUIRED_TRTLLM_OPS = ( "attn_res_fwd", "attn_res_rmsnorm_fwd", @@ -105,6 +134,11 @@ "kda_mtp_decode", "dsv3_router_gemm_op", "dsv3_fused_a_gemm_op", + "k3_kda_decode_attn", + "k3_kda_attn", + "k3_kda_verify", + "k3_mla_qkv", + "k3_mla_attn_vb_out", ) #: The engine surface the first forward checks before this target relies on it: per object, the attributes read. @@ -118,8 +152,9 @@ "seq_lens", "tokens_per_block", "kv_cache_manager", + "mamba_metadata", ), - "kv_cache_manager": ("enable_block_reuse",), + "kv_cache_manager": ("enable_block_reuse", "mamba_layer_cache"), } #: Stock code the generic path runs outside the catalog, declared so it is not consumed silently: every @@ -159,7 +194,8 @@ # ---------------------------------------------------------------------------------------------------------------------- # The text model: Kimi K3's decoder (93 layers: KDA / MLA attention, attention residuals, the dense layer-0 MLP and -# the latent MoE), its generic path. +# the latent MoE), its generic path. Each step's classification goes to the attention modules (`K3DecodeKDA`, +# `K3DecodeMLA` below). # ---------------------------------------------------------------------------------------------------------------------- # A/B escape hatch: restore nn.Linear for the K3 latent MoE projections @@ -1229,8 +1265,6 @@ def __init__( ) -> None: super().__init__() - from tensorrt_llm._torch.modules.kimi_k3_mla import KimiK3MLAAttention - max_positions = int( os.environ.get( _KIMI_K3_MLA_MAX_POSITIONS_ENV, @@ -1274,7 +1308,7 @@ def __init__( else None ) attention_config._frozen = model_config._frozen - self.mixer = KimiK3MLAAttention( + self.mixer = K3DecodeMLA( hidden_size=cfg.hidden_size, num_heads=cfg.num_attention_heads, q_lora_rank=cfg.q_lora_rank, @@ -1293,10 +1327,13 @@ def __init__( ) def forward( - self, hidden_states: torch.Tensor, attn_metadata: AttentionMetadata + self, + hidden_states: torch.Tensor, + attn_metadata: AttentionMetadata, + step: Optional[DecodeStep] = None, ) -> torch.Tensor: # MLA.forward takes position_ids first; K3 is NoPE, so pass None. - out = self.mixer(None, hidden_states, attn_metadata) + out = self.mixer(None, hidden_states, attn_metadata, step=step) if self._o_allreduce is not None: # Head-sharded TP: sum the row-sharded o_proj partials across # the head-shard group. @@ -1331,7 +1368,7 @@ def __init__( for name in projection_names } attention_config._frozen = model_config._frozen - self.linear_attn = KimiKDALinearAttention( + self.linear_attn = K3DecodeKDA( cfg, layer_idx, mapping=model_config.mapping, @@ -1422,6 +1459,7 @@ def forward( num_snapshots: int, attn_metadata: AttentionMetadata, capture: Optional[Tuple[Any, int]] = None, + step: Optional[DecodeStep] = None, ) -> Tuple[torch.Tensor, int]: """Port of HF ``KimiDecoderLayer._forward_attn_residual`` (per token). @@ -1435,6 +1473,9 @@ def forward( below already is it. Reading it here beats recomputing it, and is only possible because K3 asserts pp_size == 1 -- layer j+1 is always local. PP support would need a recompute at the rank boundary. + + ``step`` is the step's classification (``decode_step``), handed to the + attention module. """ prefix_sum = hidden_states valid_block_residual = block_residual[:num_snapshots] @@ -1473,9 +1514,9 @@ def forward( valid_block_residual = block_residual[:num_snapshots] prefix_sum = None if self.is_kda: - hidden_states = self.linear_attn(hidden_states, attn_metadata) + hidden_states = self.linear_attn(hidden_states, attn_metadata, step=step) else: - hidden_states = self.self_attn(hidden_states, attn_metadata) + hidden_states = self.self_attn(hidden_states, attn_metadata, step=step) if prefix_sum is None: prefix_sum = hidden_states @@ -1567,6 +1608,20 @@ def __init__(self, model_config: ModelConfig): key="kimi_k3_aux_capture_mode", ) + @property + def kda_token_states(self) -> bool: + """Whether the hybrid cache manager keeps the KDA state after every verify token, the protocol of + ``ssm/k3_kda_verify`` and ``ssm/k3_kda_attn``. The engine reads it once the weights are loaded, to build the + manager: DFlash / DSpark drafts of an even verify width up to 8, every KDA layer taking the K3 kernels. + Otherwise the KDA verify replays the accepted drafts (the built-in verify).""" + spec_config = getattr(self.model_config, "spec_config", None) + return bool( + spec_config is not None + and (spec_config.spec_dec_mode.is_dflash() or spec_config.spec_dec_mode.is_dspark()) + and spec_config.tokens_per_gen_step in (2, 4, 6, 8) + and all(layer.linear_attn.takes_k3_kernels for layer in self.layers if layer.is_kda) + ) + def forward( self, attn_metadata: AttentionMetadata, @@ -1582,6 +1637,7 @@ def forward( if inputs_embeds is None: inputs_embeds = self.embed_tokens(input_ids) hidden_states = inputs_embeds + step = decode_step(attn_metadata, hidden_states.shape[0]) block_residual = hidden_states.new_empty( self.num_attn_res_snapshots, @@ -1613,7 +1669,12 @@ def forward( ): capture = (spec_metadata, self.layers[i - 1].layer_idx) hidden_states, num_snapshots = layer( - hidden_states, block_residual, num_snapshots, attn_metadata, capture=capture + hidden_states, + block_residual, + num_snapshots, + attn_metadata, + capture=capture, + step=step, ) # The last layer has no successor, so this one recompute is @@ -1645,6 +1706,332 @@ def forward( ) +# ---------------------------------------------------------------------------------------------------------------------- +# The attention modules: the built-in KDA and MLA modules, with the steps the decode kernels take on those kernels. +# ---------------------------------------------------------------------------------------------------------------------- + + +class K3DecodeKDA(KimiKDALinearAttention): + """Kimi K3's KDA attention: the built-in module, with the decode kernels on the steps they take. + + * A decode step of one token per request runs the fused input projection and the plain decode in one + ``ssm/k3_kda_decode_attn`` launch, then the module's ``o_proj`` and all-reduce. + * With the cache manager's per-token states (``KimiLinearModel.kda_token_states``), every verify of the layer, on + any step, runs the kernels that keep them: ``ssm/k3_kda_attn`` for one request of 8 tokens (the projection fused + in), else ``ssm/k3_kda_verify`` on the projection's rows. The built-in verify replays drafts from caches these + kernels do not fill, so the two never run on one manager. + + Every other step runs the built-in module. The kernels read one ``[q | k | v | g | f_a | b]`` weight, built at the + checkpoint load from the module's own, and the device's ``K3KdaBuffers``, which the target sets in + ``post_load_weights``. + """ + + def __init__(self, *args, **kwargs) -> None: + super().__init__(*args, **kwargs) + # The fused [q | k | v | g | f_a | b | pad] projection weight; the six projections' weights and the built-in + # fused [q | k | v | g] and [f_a | b | pad] ones are views of it. + self.k3_proj_weight: Optional[torch.Tensor] = None + # The fused projection's Lamport buffers: one set per device, shared by every KDA layer. + self.k3_buffers: Optional[K3KdaBuffers] = None + + @property + def takes_k3_kernels(self) -> bool: + """Whether the decode kernels can run this layer: its fused projection weight is built.""" + return self.k3_proj_weight is not None + + def finalize_decode_weights(self) -> None: + """The built-in fused weights, then one ``[q | k | v | g | f_a | b | pad]`` weight of both.""" + super().finalize_decode_weights() + assert ( + self.use_full_rank_gate + and self.gate_lower_bound is not None + and self._qkvg_proj_weight is not None + and self._bfa_proj_weight is not None + and self._qkvg_proj_weight.dtype == self._bfa_proj_weight.dtype == torch.bfloat16 + ), ( + f"Kimi K3 KDA layer {self.layer_idx}: the decode kernels read the bf16 fused projections the built-in " + "module builds on CUDA at head dim 128, with a full-rank output gate and a gate lower bound" + ) + rows = self._qkvg_proj_weight.shape[0] + with torch.no_grad(): + fused = self._merge_projection_weights( + (self.q_proj, self.k_proj, self.v_proj, self.g_proj, self.f_a_proj, self.b_proj), + pad_rows_to=8, + ) + self.k3_proj_weight = fused + self._qkvg_proj_weight, self._bfa_proj_weight = fused[:rows], fused[rows:] + + def forward( + self, + hidden_states: torch.Tensor, + attn_metadata: AttentionMetadata, + step: Optional[DecodeStep] = None, + ) -> torch.Tensor: + """The built-in forward, except on a decode step of one token per request.""" + if ( + step is not None + and step.decode + and step.tokens_per_request == 1 + and self.takes_k3_kernels + and self.k3_buffers is not None + and not is_in_breakable_cuda_graph() + ): + return self._project_output( + self._k3_decode(hidden_states[: step.num_tokens], attn_metadata) + ) + return super().forward(hidden_states, attn_metadata) + + def _k3_decode(self, x: torch.Tensor, attn_metadata: AttentionMetadata) -> torch.Tensor: + """``ssm/k3_kda_decode_attn``: the core output ``[R, H, 128]`` of one token of each of the step's R requests; + each slot's conv window and state advance in place.""" + mamba_metadata = attn_metadata.mamba_metadata + slots = getattr(mamba_metadata, "generation_state_indices", None) + if slots is None: + slots = mamba_metadata.state_indices[: x.shape[0]] + layer_cache = attn_metadata.kv_cache_manager.mamba_layer_cache(self.layer_idx) + w_q, w_k, w_v = self._get_mtp_conv_weights() + core = k3_kda_decode_attn( + x.contiguous(), + self.k3_proj_weight, + self.f_b_proj.weight, + w_q, + w_k, + w_v, + self._A_log_f32, + self._dt_bias_f32, + self._onorm_w_f32, + layer_cache.conv, + layer_cache.temporal, + slots, + self.k3_buffers, + float(self.gate_lower_bound), + self.head_k_dim**-0.5, + float(self.o_norm.eps), + ) + # Speculative decoding's replay caches keep their committed conv window in step with the pool's. + self._sync_kda_replay_conv_window(layer_cache, slots, layer_cache.conv) + return core + + def forward_verify( + self, + x2d, + num_steps, + layer_cache, + conv_pool, + ssm_pool, + slot_indices, + output: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + """The built-in verify, unless the cache manager keeps the per-token states; then the decode kernels'.""" + if not (layer_cache.has_kda_replay_caches and layer_cache.kda_state_tok is not None): + return super().forward_verify( + x2d, num_steps, layer_cache, conv_pool, ssm_pool, slot_indices, output=output + ) + assert self.takes_k3_kernels and self.k3_buffers is not None, ( + f"Kimi K3 KDA layer {self.layer_idx}: the cache manager keeps the per-token verify states, which only " + "the decode kernels write, and this layer has no fused projection weight or buffers" + ) + core = self._k3_verify(x2d, num_steps, layer_cache, ssm_pool, slot_indices) + return self._store_core(core, output) + + def _k3_verify(self, x, num_steps, layer_cache, ssm_pool, slot_indices) -> torch.Tensor: + """The core output ``[N T, H, 128]`` of N requests of T verify tokens: ``ssm/k3_kda_attn`` for one request of + 8 tokens, else ``ssm/k3_kda_verify`` on the projection's rows. Each slot's state after the golden token, its + drafts' states and its conv window are written in place.""" + num_requests = x.shape[0] // num_steps + slots = slot_indices[:num_requests] + w_q, w_k, w_v = self._get_mtp_conv_weights() + constants = (float(self.gate_lower_bound), self.head_k_dim**-0.5, float(self.o_norm.eps)) + if num_requests == 1 and num_steps == 8: + out = k3_kda_attn( + x.contiguous(), + self.k3_proj_weight, + self.f_b_proj.weight, + w_q, + w_k, + w_v, + self._A_log_f32, + self._dt_bias_f32, + self._onorm_w_f32, + layer_cache.kda_conv_q, + layer_cache.kda_conv_k, + layer_cache.kda_conv_v, + ssm_pool, + layer_cache.kda_state_tok, + slots, + layer_cache.prev_num_accepted_tokens, + self.k3_buffers, + num_steps - 1, + *constants, + ) + else: + out = k3_kda_verify( + torch.nn.functional.linear(x, self.k3_proj_weight), + self.f_b_proj.weight, + w_q, + w_k, + w_v, + self._A_log_f32, + self._dt_bias_f32, + self._onorm_w_f32, + layer_cache.kda_conv_q, + layer_cache.kda_conv_k, + layer_cache.kda_conv_v, + ssm_pool, + layer_cache.kda_state_tok, + slots, + layer_cache.prev_num_accepted_tokens, + num_steps - 1, + *constants, + ) + return out.view(-1, self.num_heads, self.head_dim) + + +class K3DecodeMLA(KimiK3MLAAttention): + """Kimi K3's MLA attention: the built-in module, with a decode step's attention on the decode kernels. + + A decode step runs ``x [W_a; W_g]^T`` as one GEMM with the gate columns through a sigmoid, then + ``attention/k3_mla_qkv`` (the q_a / kv_a RMSNorms, q_b and the k_b absorption into the fused query, the step's + latent rows stored into the paged cache) and ``attention/k3_mla_attn_vb_out`` (the attention over the paged cache, + v_b and the output gate in one launch), then the module's ``o_proj``. Every other step, and a decode step whose + cache the kernels do not read (``k3_mla_decode_view`` says why), runs the built-in module. + + ``[W_a; W_g]`` is built at load from the module's weights, which become views of it. The attention workspace is + the device's ``K3MlaAttnWorkspace``, which the target sets in ``post_load_weights``. + """ + + def __init__(self, **kwargs) -> None: + super().__init__(**kwargs) + # [W_a; W_g]: the fused q_a / kv_a projection's rows, then the output gate's. + self.k3_ag_weight: Optional[torch.Tensor] = None + # The decode attention's workspace: one per device, shared by every MLA layer. + self.k3_workspace: Optional[K3MlaAttnWorkspace] = None + + def post_load_weights(self) -> None: + """The built-in post-load, then ``[W_a; W_g]`` (once: CUDA graphs captured since read it).""" + super().post_load_weights() + if self.k3_ag_weight is not None: + return + gaps = self._k3_layout_gaps() + assert not gaps, ( + f"Kimi K3 MLA layer {self.layer_idx}: the decode kernels do not take {'; '.join(gaps)}" + ) + qkv_a, gate = self.kv_a_proj_with_mqa, self.g_proj + rows = qkv_a.weight.shape[0] + with torch.no_grad(): + fused = torch.cat([qkv_a.weight, gate.weight]) + qkv_a.weight = nn.Parameter(fused[:rows], requires_grad=False) + gate.weight = nn.Parameter(fused[rows:], requires_grad=False) + self.k3_ag_weight = fused + + def _k3_layout_gaps(self) -> list: + """What of this layer the decode kernels do not take (empty when they take all of it).""" + if not (self.use_output_gate and self.fuse_qkv_a_proj and not self.is_lite): + return ["a layer without the output gate or the fused q_a / kv_a projection"] + linears = (self.kv_a_proj_with_mqa, self.g_proj, self.q_b_proj, self.o_proj) + checks = ( + (not self.mapping.has_cp_helix(), "helix context parallelism"), + (not self.apply_rotary_emb and not self.llama_4_scaling, "RoPE or llama-4 scaling"), + (self.sparse_attn_hooks is None, "sparse attention"), + ( + self.kv_cache_dtype != "fp8_ds_mla" + and not getattr(self.mqa, "has_fp8_kv_cache", False) + and not getattr(self.mqa, "has_fp4_kv_cache", False), + "a quantized KV cache", + ), + ( + all(m.weight.dtype == torch.bfloat16 and m.bias is None for m in linears) + and self.k_b_proj_trans.dtype == self.v_b_proj.dtype == torch.bfloat16, + "projections other than bf16 and unbiased", + ), + ( + not getattr(self.q_a_layernorm, "is_nvfp4", False) + and not getattr(self.kv_a_layernorm, "use_gemma", False), + "an NVFP4 q_a norm or a Gemma kv_a norm", + ), + ( + self.num_heads_tp % 6 == 0 + and self.kv_lora_rank == 512 + and self.qk_rope_head_dim == 64 + and self.q_lora_rank == 1536, + f"{self.num_heads_tp} heads, latent {self.kv_lora_rank}, rope {self.qk_rope_head_dim}, " + f"q_lora {self.q_lora_rank}", + ), + ) + return [why for ok, why in checks if not ok] + + def forward( + self, + position_ids: Optional[torch.Tensor], + hidden_states: torch.Tensor, + attn_metadata: AttentionMetadata, + all_reduce_params=None, + latent_cache_gen: Optional[torch.Tensor] = None, + step: Optional[DecodeStep] = None, + ) -> torch.Tensor: + """The built-in forward, except on a decode step whose cache the decode kernels read.""" + view = None + if ( + step is not None + and step.decode + and latent_cache_gen is None + and self.k3_ag_weight is not None + and self.k3_workspace is not None + and not is_in_breakable_cuda_graph() + ): + view = self._k3_decode_view(attn_metadata, step.num_tokens) + if view is None: + return super().forward( + position_ids, hidden_states, attn_metadata, all_reduce_params, latent_cache_gen + ) + x = hidden_states[: step.num_tokens].contiguous() + rows = self.kv_a_proj_with_mqa.weight.shape[0] + ag = torch.nn.functional.linear(x, self.k3_ag_weight) + ag[:, rows:].sigmoid_() + fused_q = k3_mla_qkv( + ag, + self.q_a_layernorm.weight, + float(self.q_a_layernorm.variance_epsilon), + self.q_b_proj.weight, + self.k_b_proj_trans, + self.kv_a_layernorm.weight, + float(self.kv_a_layernorm.variance_epsilon), + view["pool"], + view["row_stride"], + view["page_table"], + view["page_offset"], + view["seq_len"], + ) + attn_output = self.create_output(x, 0) + k3_mla_attn_vb_out( + fused_q, + view["pool"], + view["row_stride"], + view["page_table"], + view["page_offset"], + view["seq_len"], + view["softmax_scale"], + self.v_b_proj, + attn_output, + self.k3_workspace, + gate=ag, + gate_col0=rows, + ) + return self._project_output([attn_output], position_ids, attn_metadata, all_reduce_params) + + def _k3_decode_view(self, attn_metadata: AttentionMetadata, num_tokens: int) -> Optional[dict]: + """The paged latent cache as the decode kernels read it this step, or None when they do not read it (the + reason is logged once).""" + view = k3_mla_decode_view(self.mqa, attn_metadata, num_tokens) + if isinstance(view, str): + logger.info_once( + f"Kimi K3 MLA: the built-in path for a decode step the decode kernels do not read ({view})", + key=f"k3_mla_decode_view_{view}", + ) + return None + return view + + # ---------------------------------------------------------------------------------------------------------------------- # The target: step classification, the construction checks and the registration shell. # ---------------------------------------------------------------------------------------------------------------------- @@ -1813,9 +2200,6 @@ def __init__(self, model_config: ModelConfig): vocab_size=cfg.vocab_size, ) self._step_checked = False - # The fused decode path and the state its kernels share, built in post_load_weights once the catalog entries - # it calls exist. None: every step takes the generic path. - self._fused_decode = None # The executor reads generation settings (eos_token_id, ...) off the model config the engine holds, which # must therefore be the text config, as the built-in wrapper leaves it. model_config._frozen = False @@ -1825,6 +2209,27 @@ def __init__(self, model_config: ModelConfig): def load_weights(self, weights, *args, **kwargs): _weights.load(self, weights) + def post_load_weights(self) -> None: + """The state the decode kernels share, built once per device before any CUDA-graph capture and handed to + every layer of its kind: the KDA projection's Lamport buffers and the MLA decode attention's workspace.""" + super().post_load_weights() + kda = [layer.linear_attn for layer in self.model.layers if layer.is_kda] + mla = [layer.self_attn.mixer for layer in self.model.layers if not layer.is_kda] + if kda[0].k3_buffers is not None: + return # built by an earlier call; CUDA graphs captured since hold it + device = self.model.embed_tokens.weight.device + buffers = K3KdaBuffers.create(device) + workspace = K3MlaAttnWorkspace.create(device, mla[0].num_heads_tp // 6) + for module in kda: + module.k3_buffers = buffers + for module in mla: + module.k3_workspace = workspace + logger.info( + "Kimi K3 decode kernels: KDA on k3_kda_decode_attn, k3_kda_attn and k3_kda_verify " + f"({sum(m.takes_k3_kernels for m in kda)} / {len(kda)} layers take them), MLA on k3_mla_qkv and " + f"k3_mla_attn_vb_out ({len(mla)} layers)" + ) + def _check_step_contract(self, attn_metadata: AttentionMetadata) -> None: """First-forward checks of the engine surface and the per-engine settings.""" objects = { @@ -1851,6 +2256,12 @@ def _check_step_contract(self, attn_metadata: AttentionMetadata) -> None: assert not manager.enable_block_reuse, ( "this target runs with kv_cache_config.enable_block_reuse false; the engine enabled it" ) + kda_layer = next(layer.layer_idx for layer in self.model.layers if layer.is_kda) + state_dtype = manager.mamba_layer_cache(kda_layer).temporal.dtype + assert state_dtype == torch.float32, ( + "this target's KDA kernels keep fp32 recurrent states (kv_cache_config.mamba_ssm_cache_dtype float32 or " + f"auto); the engine built a {state_dtype} state pool" + ) self._step_checked = True def forward( @@ -1870,19 +2281,6 @@ def forward( ) if not self._step_checked: self._check_step_contract(attn_metadata) - if self._fused_decode is not None: - rows = input_ids if input_ids is not None else inputs_embeds - step = None if rows is None else decode_step(attn_metadata, rows.shape[0]) - if step is not None: - return self._fused_decode( - step, - attn_metadata=attn_metadata, - input_ids=input_ids, - position_ids=position_ids, - spec_metadata=spec_metadata, - resource_manager=resource_manager, - **kwargs, - ) return super().forward( attn_metadata, input_ids, diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py index bd8376b2dcea..b62a1dad85fe 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py @@ -219,7 +219,8 @@ def test_targets_do_not_share_files(): #: tensorrt_llm imports that compute nothing, so a target's generic path need not declare them: types and configs -#: it reads, enums it passes, the metadata it is handed, its registration and its logger. +#: it reads, enums it passes, the metadata it is handed and the views it takes of it, the graph-capture context it +#: asks about, its registration and its logger. _NON_COMPUTE_IMPORTS = frozenset( { "AllReduceStrategy", @@ -232,6 +233,8 @@ def test_targets_do_not_share_files(): "QuantAlgo", "QuantConfig", "SiTuActivation", + "is_in_breakable_cuda_graph", + "k3_mla_decode_view", "logger", "register_auto_model", } From 0bf72e41bb1090109291eabcc6b1facab77c1a9d Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 19:54:23 -0700 Subject: [PATCH 039/161] [None][feat] modeling_v2 Kimi K3 target: decode GEMVs, LM head and embedding on the G2 entries decode_gemv.py puts the target's decode-size GEMVs on the catalog's single-GPU Kimi K3 entries, and the target calls it: - the MLA projections of at most 8 rows, per call site, on the kernel measured fastest at their weight shape: kv_a_proj_with_mqa on gemm/k3_decode_gemv, q_b_proj and the output gate on gemm/k3_ctm_gemv_wide (K3DecodeGemvs.project; the attention modules call it); - the LM head at most 8 rows on gemm/k3_head_gemv over a K3HeadGemvWorkspace, the vocabulary shards gathered with comm/allgather as the stock head does. K3LogitsProcessor wraps the shell's processor, so the speculative worker's target logits and a parallel drafter's logits take it too; - a decode step's embedding and layer 0's input RMSNorm in one norm/k3_embed_norm launch, the embedding written into the attention residual bank as layer 0's first snapshot (layer 0 then skips both). cache_derived_state builds the state once the weights are final: the head's workspace, and one eager call of every site's kernel and of the head. Under CUDA-graph capture a call whose kernel never ran eagerly takes the generic path. Every call returns None where its kernel does not take it, and the caller runs the generic module. The new test runs each site at 1..8 rows against its entry's bits and a float64 product, the LM head through the processor on a real LMHead, the embedding bit for bit against nn.Embedding + RMSNorm, the refusals, and graph replays. It is listed in l0_b200. Signed-off-by: Vasanth Sabavat --- .../decode_gemv.py | 282 ++++++++++++++++++ .../modeling.py | 92 +++++- .../test_lists/test-db/l0_b200.yml | 2 + .../test_modeling_v2_kimi_k3_decode_gemv.py | 214 +++++++++++++ 4 files changed, 578 insertions(+), 12 deletions(-) create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py create mode 100644 tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_gemv.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py new file mode 100644 index 000000000000..6dee9d21e49b --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py @@ -0,0 +1,282 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""The decode path's GEMVs, LM head and embedding on the catalog's single-GPU Kimi K3 entries. + +* **Per-site GEMVs** (`K3DecodeGemvs.project`): an MLA projection of at most `MAX_ROWS` rows runs on the kernel + measured fastest at its weight shape at every row count 1..8 (`SITES`): `gemm/k3_decode_gemv` for the fused + q_a / kv_a projection, `gemm/k3_ctm_gemv_wide` for q_b and the output gate. +* **LM head** (`K3LogitsProcessor`): at most `MAX_ROWS` rows of this rank's vocabulary shard on + `gemm/k3_head_gemv` over the target's `K3HeadGemvWorkspace`, then the shards gathered (`comm/allgather`) as the + stock head gathers them. It is the shell's logits processor, so the speculative worker's target logits and the + drafter's logits on the same head take it too. +* **Embedding** (`K3DecodeGemvs.embed_norm`): a decode step's embedding rows written into the attention-residual + bank's slot 0 (layer 0's first snapshot) and layer 0's input RMSNorm applied, in one `norm/k3_embed_norm` launch. + +Each returns None where its kernel does not take the call, and the caller then runs the generic path's module. + +A kernel compiles on its first call for a shape, which must not happen under CUDA-graph capture. +`K3DecodeGemvs.create`, run once the weights are final, runs every site's kernel and the head once, eagerly. The +embedding kernel compiles per token count, on the eager warm-up step before each capture. Under capture, a call +whose kernel has not run eagerly is refused. + +The head's workspace serves every `k3_head_gemv` call of its weight shape, so those calls must be ordered on one +stream: the logits are computed on the model's stream (see the `gemm/k3_head_gemv` contract's State section). +""" + +from __future__ import annotations + +from typing import Dict, Optional, Set, Tuple + +import torch +from torch import nn + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.allgather import allgather +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv_wide import ( + k3_ctm_gemv_wide, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_decode_gemv import k3_decode_gemv +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_head_gemv import ( + K3HeadGemvWorkspace, + k3_head_gemv, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.norm.k3_embed_norm import k3_embed_norm +from tensorrt_llm._torch._experimental.modeling_v2.catalog.torch.concat import concat +from tensorrt_llm._torch._experimental.modeling_v2.catalog.torch.split import split + +# The kernels' support predicates: metadata reads only. +from tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv import op as _ctm_op +from tensorrt_llm._torch.cute_dsl_kernels.k3_decode_gemv import op as _decode_op +from tensorrt_llm._torch.cute_dsl_kernels.k3_embed import op as _embed_op +from tensorrt_llm._torch.cute_dsl_kernels.k3_head_gemv import op as _head_op +from tensorrt_llm._torch.flashinfer_utils import IS_FLASHINFER_AVAILABLE +from tensorrt_llm._utils import mpi_disabled + +# The row limit of the decode GEMV and head kernels (one token tile). +MAX_ROWS = 8 + +#: Site -> (N, K) of its weight at this target's per-rank shapes, and the kernel that runs it at 1..MAX_ROWS rows. +SITES: Dict[str, Tuple[int, int, str]] = { + # MLA kv_a_proj_with_mqa: [q_a 1536 | kv_a 512 | k_pe 64] of the hidden size. + "kv_a": (2112, 7168, "decode"), + # MLA q_b_proj: 6 heads x 192 of q_lora_rank. + "q_b": (1152, 1536, "wide"), + # MLA output gate: 6 heads x 128 of the hidden size. + "g_proj": (768, 7168, "wide"), +} + + +def _capturing() -> bool: + return not torch.compiler.is_compiling() and torch.cuda.is_current_stream_capturing() + + +def _dense_rows(x2d: torch.Tensor) -> torch.Tensor: + """``x2d`` itself, or a dense copy: the kernels' TMA descriptors need dense, 16-byte-aligned rows.""" + if x2d.stride() != (x2d.shape[1], 1) or x2d.data_ptr() % 16: + return x2d.clone(memory_format=torch.contiguous_format) + return x2d + + +def _head_takes_module(lm_head: nn.Module) -> bool: + """Whether ``lm_head(rows)`` is a plain vocabulary-parallel GEMM of a bf16 weight whose shards it gathers along + the vocabulary in rank order, with nothing else applied: the stock head this path reproduces.""" + weight = getattr(lm_head, "weight", None) + mapping = getattr(lm_head, "mapping", None) + return ( + isinstance(weight, torch.Tensor) + and weight.dim() == 2 + and weight.dtype == torch.bfloat16 + and weight.is_contiguous() + and getattr(getattr(lm_head, "tp_mode", None), "name", None) == "COLUMN" + and getattr(lm_head, "gather_output", False) + and getattr(lm_head, "gather_output_sizes", None) is None + and getattr(lm_head, "padding_size", None) == 0 + and getattr(lm_head, "bias", None) is None + and not getattr(lm_head, "has_any_quant", True) + and mapping is not None + and not mapping.enable_attention_dp + ) + + +def _plain_rmsnorm(norm: nn.Module) -> bool: + """Whether ``norm`` is the stock RMSNorm that runs flashinfer's kernel, which ``k3_embed_norm`` reproduces bit for + bit. An unknown module fails closed.""" + weight = getattr(norm, "weight", None) + return ( + IS_FLASHINFER_AVAILABLE + and isinstance(weight, torch.Tensor) + and weight.dtype == torch.bfloat16 + and not getattr(norm, "use_gemma", True) + and not getattr(norm, "is_nvfp4", True) + and not getattr(norm, "use_cuda_tile", True) + and not getattr(norm, "return_hp_output", True) + and getattr(norm, "nvfp4_scale", None) is None + and hasattr(norm, "variance_epsilon") + ) + + +class K3DecodeGemvs: + """The decode GEMVs' state for one target: the LM head's `K3HeadGemvWorkspace`, and the calls whose kernels ran + eagerly (compiled), which are the only ones a CUDA-graph capture may take. Built by `create` once the weights are + final; owned by the target.""" + + def __init__(self, head_workspace: Optional[K3HeadGemvWorkspace] = None) -> None: + self.head_workspace = head_workspace + self._ran: Set[tuple] = set() + + @classmethod + def create( + cls, lm_head: Optional[nn.Module], site_weights: Dict[str, torch.Tensor] + ) -> "K3DecodeGemvs": + """The state for ``lm_head`` and the site weights (one weight per `SITES` key; every weight of a site has its + shape). Eager: it allocates the head's workspace and runs each kernel once on a zero row, so they compile here + rather than under a capture. A site or head its kernel does not take keeps the generic path.""" + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "K3DecodeGemvs.create allocates and compiles: run it before CUDA-graph capture" + ) + state = cls() + for site, weight in site_weights.items(): + state._project(site, weight.new_zeros(1, weight.shape[1]), weight, warm=True) + if lm_head is not None and _head_takes_module(lm_head): + weight = lm_head.weight + x = weight.new_zeros(1, weight.shape[1]) + if _head_op.supports(x, weight): + workspace = K3HeadGemvWorkspace.create( + weight.shape[0], weight.shape[1], weight.device + ) + k3_head_gemv(x, weight, workspace) + state.head_workspace = workspace + state._ran.add(("lm_head",)) + torch.cuda.synchronize() + return state + + def project(self, site: str, x: torch.Tensor, weight: torch.Tensor) -> Optional[torch.Tensor]: + """``x @ weight.T`` (bf16 ``[..., N]``) for ``site``'s weight on its decode kernel, or None where the kernel + does not take the call: more than `MAX_ROWS` rows, another shape or dtype, or, under capture, a kernel that + has not run eagerly.""" + return self._project(site, x, weight, warm=False) + + def _project( + self, site: str, x: torch.Tensor, weight: torch.Tensor, warm: bool + ) -> Optional[torch.Tensor]: + n, k, kernel = SITES[site] + if ( + weight.dtype != torch.bfloat16 + or tuple(weight.shape) != (n, k) + or not weight.is_contiguous() + or x.dtype != torch.bfloat16 + or x.dim() < 1 + or x.shape[-1] != k + ): + return None + rows = x.numel() // k + if not 0 < rows <= MAX_ROWS: + return None + key = (site,) + capturing = _capturing() + if capturing and not warm and key not in self._ran: + return None + x2d = _dense_rows(x.reshape(rows, k)) + if kernel == "decode": + if not _decode_op.supports(x2d, weight): + return None + y = k3_decode_gemv(x2d, weight) + else: + if not _ctm_op.supports_wide(x2d, weight, -1, False): + return None + y = k3_ctm_gemv_wide(x2d, weight) + if not capturing: + self._ran.add(key) + return y.view(*x.shape[:-1], n) + + def lm_head_logits(self, rows: torch.Tensor, lm_head: nn.Module) -> Optional[torch.Tensor]: + """``lm_head(rows)``, the gathered bf16 logits ``[M, vocab]``, with this rank's shard on + ``gemm/k3_head_gemv``; None where it does not take the call (more than `MAX_ROWS` rows, another head, ...).""" + workspace = self.head_workspace + if workspace is None or not _head_takes_module(lm_head): + return None + weight = lm_head.weight + if ( + tuple(weight.shape) != (workspace.n_out, workspace.k_in) + or weight.device != workspace.partials.device + or rows.dim() != 2 + or rows.dtype != torch.bfloat16 + or not 0 < rows.shape[0] <= MAX_ROWS + or rows.shape[1] != workspace.k_in + ): + return None + if _capturing() and ("lm_head",) not in self._ran: + return None + group = lm_head.mapping.tp_group + if len(group) > 1 and mpi_disabled(): + return None + x = _dense_rows(rows) + if not _head_op.supports(x, weight): + return None + local = k3_head_gemv(x, weight, workspace) + if len(group) == 1: + return local + gathered = allgather(local, None, group) + return concat(list(split(gathered, rows.shape[0], dim=0)), dim=-1) + + def embed_norm( + self, + input_ids: torch.Tensor, + table: torch.Tensor, + norm: nn.Module, + bank: torch.Tensor, + ) -> Optional[torch.Tensor]: + """Layer 0's normed input for the step's tokens, with their embedding rows written into ``bank[0]``, in one + ``norm/k3_embed_norm`` launch: bit-identical to the embedding followed by ``norm``. None where it does not + apply (another norm, a token count or table the kernel does not take, or, under capture, a token count whose + kernel has not run eagerly).""" + if not _plain_rmsnorm(norm): + return None + ids = input_ids.reshape(-1) + if not _embed_op.supports_norm(ids, table, norm.weight, bank[0]): + return None + key = ("embed_norm", ids.numel(), ids.dtype) + capturing = _capturing() + if capturing and key not in self._ran: + return None + normed = k3_embed_norm(ids, table, norm.weight, norm.variance_epsilon, bank[0]) + if not capturing: + self._ran.add(key) + return normed + + +class _K3Head: + """``lm_head(rows)``, with the rows on ``k3_head_gemv`` where it takes them.""" + + __slots__ = ("_gemvs", "_lm_head") + + def __init__(self, gemvs: K3DecodeGemvs, lm_head: nn.Module) -> None: + self._gemvs = gemvs + self._lm_head = lm_head + + def __call__(self, rows: torch.Tensor) -> torch.Tensor: + logits = self._gemvs.lm_head_logits(rows, self._lm_head) + return self._lm_head(rows) if logits is None else logits + + +class K3LogitsProcessor(nn.Module): + """The shell's logits processor, with its LM head call on ``gemm/k3_head_gemv`` where ``gemvs`` takes the rows. + + It wraps the stock processor instead of repeating it: the row selection and the fp32 conversion stay the stock + processor's. ``gemvs`` is set once the weights are final; until then every call is the stock one. + """ + + def __init__(self, stock: nn.Module) -> None: + super().__init__() + self.stock = stock + self.gemvs: Optional[K3DecodeGemvs] = None + + def forward( + self, + hidden_states: torch.Tensor, + lm_head: nn.Module, + attn_metadata, + return_context_logits: bool = False, + ) -> torch.Tensor: + head = lm_head if self.gemvs is None else _K3Head(self.gemvs, lm_head) + return self.stock.forward(hidden_states, head, attn_metadata, return_context_logits) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index 1b7a5cacbd5e..4af83170325a 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -82,6 +82,7 @@ from tensorrt_llm.mapping import Mapping from tensorrt_llm.models.modeling_utils import QuantAlgo, QuantConfig +from . import decode_gemv as _decode_gemv from . import weights as _weights if TYPE_CHECKING: @@ -105,6 +106,12 @@ "kda_mtp_decode", "dsv3_router_gemm_op", "dsv3_fused_a_gemm_op", + # The decode path's GEMVs, LM head and embedding (decode_gemv.py). + "k3_decode_gemv", + "k3_ctm_gemv_wide", + "k3_head_gemv", + "k3_embed_norm", + "allgather", ) #: The engine surface the first forward checks before this target relies on it: per object, the attributes read. @@ -1416,6 +1423,7 @@ def forward( num_snapshots: int, attn_metadata: AttentionMetadata, capture: Optional[Tuple[Any, int]] = None, + prenormed: bool = False, ) -> Tuple[torch.Tensor, int]: """Port of HF ``KimiDecoderLayer._forward_attn_residual`` (per token). @@ -1429,11 +1437,17 @@ def forward( below already is it. Reading it here beats recomputing it, and is only possible because K3 asserts pp_size == 1 -- layer j+1 is always local. PP support would need a recompute at the rank boundary. + + ``prenormed`` (layer 0 on a decode step): ``hidden_states`` already is + this layer's input norm, and the layer's input, the step's embedding, + already is in ``block_residual[0]`` (``K3DecodeGemvs.embed_norm``). """ prefix_sum = hidden_states valid_block_residual = block_residual[:num_snapshots] - if capture is not None: + if prenormed: + assert num_snapshots == 0 and self.layer_idx % self.attn_res_block_size == 0 + elif capture is not None: # The mixture tap needs the PRE-norm value, which the fused # attn-res + RMSNorm kernel does not expose. Keep the two steps # split on captured layers only and fuse everywhere else. @@ -1462,7 +1476,8 @@ def forward( hidden_states = self.input_layernorm(hidden_states) if self.layer_idx % self.attn_res_block_size == 0: - block_residual[num_snapshots].copy_(prefix_sum) + if not prenormed: + block_residual[num_snapshots].copy_(prefix_sum) num_snapshots += 1 valid_block_residual = block_residual[:num_snapshots] prefix_sum = None @@ -1550,6 +1565,9 @@ def __init__(self, model_config: ModelConfig): self.num_attn_res_snapshots = ( cfg.num_hidden_layers + cfg.attn_res_block_size - 1 ) // cfg.attn_res_block_size + # The decode path's GEMVs and embedding (decode_gemv.py), built by the target's cache_derived_state once the + # weights are final. None: every step embeds on the generic path. + self.decode_gemvs: Optional[_decode_gemv.K3DecodeGemvs] = None # Which convention the drafter tap is on is not recoverable from the # served output -- a mismatch only lowers acceptance -- so state it once @@ -1573,15 +1591,33 @@ def forward( if (input_ids is None) ^ (inputs_embeds is not None): raise ValueError("You must specify exactly one of input_ids or inputs_embeds") - if inputs_embeds is None: - inputs_embeds = self.embed_tokens(input_ids) - hidden_states = inputs_embeds - - block_residual = hidden_states.new_empty( - self.num_attn_res_snapshots, - hidden_states.shape[0], - hidden_states.shape[1], - ) + num_tokens = (input_ids if inputs_embeds is None else inputs_embeds).shape[0] + step = decode_step(attn_metadata, num_tokens) + # A decode step embeds and norms for layer 0 in one launch, the embedding written as layer 0's first snapshot. + prenormed = None + if ( + inputs_embeds is None + and step is not None + and self.decode_gemvs is not None + and len(self.layers) > 0 + and self.num_attn_res_snapshots > 0 + ): + table = self.embed_tokens.weight + block_residual = table.new_empty( + self.num_attn_res_snapshots, num_tokens, table.shape[1] + ) + prenormed = self.decode_gemvs.embed_norm( + input_ids, table, self.layers[0].input_layernorm, block_residual + ) + if prenormed is not None: + hidden_states = prenormed + else: + hidden_states = self.embed_tokens(input_ids) if inputs_embeds is None else inputs_embeds + block_residual = hidden_states.new_empty( + self.num_attn_res_snapshots, + hidden_states.shape[0], + hidden_states.shape[1], + ) num_snapshots = 0 capture_set = ( getattr(spec_metadata, "_capture_layer_set", None) @@ -1607,7 +1643,12 @@ def forward( ): capture = (spec_metadata, self.layers[i - 1].layer_idx) hidden_states, num_snapshots = layer( - hidden_states, block_residual, num_snapshots, attn_metadata, capture=capture + hidden_states, + block_residual, + num_snapshots, + attn_metadata, + capture=capture, + prenormed=i == 0 and prenormed is not None, ) # The last layer has no successor, so this one recompute is @@ -1807,6 +1848,13 @@ def __init__(self, model_config: ModelConfig): vocab_size=cfg.vocab_size, ) self._step_checked = False + # The LM head on gemm/k3_head_gemv at decode size, once cache_derived_state has built its state. It is the + # processor the speculative worker's target logits and a parallel drafter's logits call too. + stock_logits_processor = self.logits_processor + self.logits_processor = _decode_gemv.K3LogitsProcessor(stock_logits_processor) + draft_model = getattr(self, "draft_model", None) + if getattr(draft_model, "logits_processor", None) is stock_logits_processor: + draft_model.logits_processor = self.logits_processor # The fused decode path and the state its kernels share, built in post_load_weights once the catalog entries # it calls exist. None: every step takes the generic path. self._fused_decode = None @@ -1819,6 +1867,26 @@ def __init__(self, model_config: ModelConfig): def load_weights(self, weights, *args, **kwargs): _weights.load(self, weights) + def cache_derived_state(self) -> None: + """Build the decode GEMVs' state from the final weights: the LM head's workspace, and one eager call of each + decode GEMV kernel (the MLA projections' shapes, read off the first MLA layer), so none compiles under + capture.""" + super().cache_derived_state() + mla = next((layer.self_attn.mixer for layer in self.model.layers if not layer.is_kda), None) + sites = {} + if mla is not None: + for site, name in ( + ("kv_a", "kv_a_proj_with_mqa"), + ("q_b", "q_b_proj"), + ("g_proj", "g_proj"), + ): + module = getattr(mla, name, None) + if module is not None: + sites[site] = module.weight + gemvs = _decode_gemv.K3DecodeGemvs.create(self.lm_head, sites) + self.model.decode_gemvs = gemvs + self.logits_processor.gemvs = gemvs + def _check_step_contract(self, attn_metadata: AttentionMetadata) -> None: """First-forward checks of the engine surface and the per-engine settings.""" objects = { diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index 951d183ea564..8649f454a0ac 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -158,6 +158,8 @@ l0_b200: - unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_ctm_gemv_wide.py - unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_decode_gemv.py - unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_head_gemv.py + # The Kimi K3 target's decode GEMVs, LM head and embedding on those entries. + - unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_gemv.py # KDA runtime: host-derived prefill metadata and the bf16 state pool # round-trip. Both are single-device cases that use GPU 0 only. - unittest/_torch/modules/kimi_kda/test_kda_host_metadata.py diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_gemv.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_gemv.py new file mode 100644 index 000000000000..86366db7fad6 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_gemv.py @@ -0,0 +1,214 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""The Kimi K3 target's decode GEMVs, LM head and embedding on the catalog's single-GPU entries (``decode_gemv.py`` +of ``kimi_k3_mxfp4__sm_100__tp16_moetp4ep4``), on one GPU. + +* Each MLA projection site at 1..8 rows: the bits of its catalog entry's call, within 8e-3 of ``max |ref|`` of a + float64 product; declined above 8 rows, at another shape or dtype, and under capture before an eager call. +* The LM head through ``K3LogitsProcessor`` on a real ``LMHead`` (one rank): fp32 logits from ``k3_head_gemv`` + within the same bound, the stock processor's rows selected, the stock path above 8 rows and before the state is + built. +* The embedding: bit-identical to ``nn.Embedding`` followed by the stock RMSNorm, the rows in the bank's slot 0. +* CUDA-graph replays give the eager bits. +""" + +import types + +import pytest +import torch +from torch import nn + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv_wide import ( + k3_ctm_gemv_wide, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_decode_gemv import k3_decode_gemv +from tensorrt_llm._torch._experimental.modeling_v2.models.kimi_k3_vl.kimi_k3_mxfp4__sm_100__tp16_moetp4ep4 import ( # noqa: E501 + decode_gemv, +) +from tensorrt_llm._torch.modules.embedding import LMHead +from tensorrt_llm._torch.modules.linear import TensorParallelMode +from tensorrt_llm._torch.modules.logits_processor import LogitsProcessor +from tensorrt_llm._torch.modules.rms_norm import RMSNorm +from tensorrt_llm.mapping import Mapping + +pytestmark = pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.get_device_capability() != (10, 0), + reason="the K3 decode kernels run on sm_100 only", +) + +TOL = 8e-3 # max |y - ref| / max |ref|, ref a float64 product +VOCAB_SHARD, HIDDEN = 10240, 7168 +ROWS = range(1, 9) + + +def _weight(n, k, seed): + g = torch.Generator(device="cuda").manual_seed(seed) + return (torch.randn(n, k, generator=g, device="cuda") * 0.02).to(torch.bfloat16) + + +def _rows(m, k, seed): + g = torch.Generator(device="cuda").manual_seed(1000 + seed) + return torch.randn(m, k, generator=g, device="cuda").to(torch.bfloat16) + + +def _rel_err(y, x, w): + ref = x.double() @ w.double().t() + return ((y.double() - ref).abs().max() / ref.abs().max()).item() + + +def _bits(t): + return t.contiguous().view(torch.int16) + + +@pytest.mark.parametrize("site", list(decode_gemv.SITES)) +def test_site(site): + """Every row count 1..8 runs the site's catalog entry: its bits, within TOL of the float64 product.""" + n, k, kernel = decode_gemv.SITES[site] + w = _weight(n, k, seed=len(site)) + gemvs = decode_gemv.K3DecodeGemvs.create(None, {site: w}) + entry = k3_decode_gemv if kernel == "decode" else k3_ctm_gemv_wide + for m in ROWS: + x = _rows(m, k, seed=m) + y = gemvs.project(site, x, w) + assert y is not None and y.shape == (m, n), (site, m) + assert torch.equal(_bits(y), _bits(entry(x, w))), (site, m) + assert _rel_err(y, x, w) <= TOL, (site, m) + + +def test_site_declines(): + """Above 8 rows, another weight shape, a non-bf16 input, and under capture before any eager call: None, nothing + launched.""" + n, k, _ = decode_gemv.SITES["kv_a"] + w = _weight(n, k, seed=1) + gemvs = decode_gemv.K3DecodeGemvs.create(None, {"kv_a": w}) + assert gemvs.project("kv_a", _rows(9, k, seed=9), w) is None + assert gemvs.project("kv_a", _rows(4, k, seed=4), _weight(n + 128, k, seed=2)) is None + assert gemvs.project("kv_a", _rows(4, k, seed=4).float(), w) is None + fresh = decode_gemv.K3DecodeGemvs() + x = _rows(4, k, seed=4) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + y = fresh.project("kv_a", x, w) + assert y is None + + +def test_site_capture_replays_the_eager_bits(): + """A captured call, its input rewritten in place before each replay, gives the eager call's bits.""" + n, k, _ = decode_gemv.SITES["g_proj"] + w = _weight(n, k, seed=3) + gemvs = decode_gemv.K3DecodeGemvs.create(None, {"g_proj": w}) + x = _rows(4, k, seed=4) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + y = gemvs.project("g_proj", x, w) + assert y is not None + for seed in (5, 6): + x.copy_(_rows(4, k, seed=seed)) + graph.replay() + torch.cuda.synchronize() + assert torch.equal(_bits(y), _bits(gemvs.project("g_proj", x, w))) + + +def _lm_head(): + head = LMHead( + VOCAB_SHARD, + HIDDEN, + dtype=torch.bfloat16, + mapping=Mapping(), + tensor_parallel_mode=TensorParallelMode.COLUMN, + gather_output=True, + ).cuda() + with torch.no_grad(): + head.weight.copy_(_weight(VOCAB_SHARD, HIDDEN, seed=7)) + return head + + +def test_lm_head(): + """At most 8 rows: fp32 logits from k3_head_gemv within TOL; 9 rows and an unbuilt processor: the stock bits.""" + head = _lm_head() + stock = LogitsProcessor() + processor = decode_gemv.K3LogitsProcessor(stock) + for m in (1, 9): + x = _rows(m, HIDDEN, seed=m) + assert torch.equal(processor(x, head, None, True), stock(x, head, None, True)) + processor.gemvs = decode_gemv.K3DecodeGemvs.create(head, {}) + assert processor.gemvs.head_workspace is not None + for m in ROWS: + x = _rows(m, HIDDEN, seed=m) + logits = processor(x, head, None, True) + assert logits.dtype == torch.float32 and logits.shape == (m, VOCAB_SHARD) + assert _rel_err(logits, x, head.weight) <= TOL, m + x = _rows(9, HIDDEN, seed=9) + assert torch.equal(processor(x, head, None, True), stock(x, head, None, True)) + + +def test_lm_head_selects_the_stock_rows(): + """Without context logits, the stock processor's last-token selection feeds the head.""" + head = _lm_head() + processor = decode_gemv.K3LogitsProcessor(LogitsProcessor()) + processor.gemvs = decode_gemv.K3DecodeGemvs.create(head, {}) + x = _rows(7, HIDDEN, seed=3) + metadata = types.SimpleNamespace( + seq_lens_cuda=torch.tensor([3, 1, 3], dtype=torch.int32, device="cuda") + ) + logits = processor(x, head, metadata, False) + want = processor(x[[2, 3, 6]].contiguous(), head, None, True) + assert torch.equal(logits, want) + + +def test_lm_head_capture_replays_the_eager_bits(): + head = _lm_head() + gemvs = decode_gemv.K3DecodeGemvs.create(head, {}) + x = _rows(8, HIDDEN, seed=8) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + y = gemvs.lm_head_logits(x, head) + assert y is not None + x.copy_(_rows(8, HIDDEN, seed=11)) + graph.replay() + torch.cuda.synchronize() + assert torch.equal(_bits(y), _bits(gemvs.lm_head_logits(x, head))) + + +@pytest.mark.parametrize("n", [1, 8, 64]) +def test_embed_norm(n): + """Layer 0's normed input is bit-identical to nn.Embedding followed by the stock RMSNorm; the bank's slot 0 holds + the embedding rows.""" + vocab = 4096 + table = _weight(vocab, HIDDEN, seed=21) + norm = RMSNorm(hidden_size=HIDDEN, eps=1e-5, dtype=torch.bfloat16).cuda() + with torch.no_grad(): + norm.weight.copy_(1 + _weight(1, HIDDEN, seed=22)[0]) + embedding = nn.Embedding(vocab, HIDDEN, dtype=torch.bfloat16).cuda() + with torch.no_grad(): + embedding.weight.copy_(table) + g = torch.Generator(device="cuda").manual_seed(n) + ids = torch.randint(0, vocab, (n,), generator=g, device="cuda", dtype=torch.int32) + gemvs = decode_gemv.K3DecodeGemvs() + bank = table.new_empty(3, n, HIDDEN) + normed = gemvs.embed_norm(ids, table, norm, bank) + assert normed is not None + raw = embedding(ids) + assert torch.equal(_bits(bank[0]), _bits(raw)) + assert torch.equal(_bits(normed), _bits(norm(raw))) + + +def test_embed_norm_capture(): + """Under capture a token count that ran eagerly replays the eager bits; one that never ran is declined.""" + vocab = 4096 + table = _weight(vocab, HIDDEN, seed=31) + norm = RMSNorm(hidden_size=HIDDEN, eps=1e-5, dtype=torch.bfloat16).cuda() + gemvs = decode_gemv.K3DecodeGemvs() + ids = torch.arange(4, dtype=torch.int32, device="cuda") + bank = table.new_empty(2, 4, HIDDEN) + eager = gemvs.embed_norm(ids, table, norm, bank).clone() + other_ids = torch.arange(5, dtype=torch.int32, device="cuda") + other_bank = table.new_empty(2, 5, HIDDEN) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + captured = gemvs.embed_norm(ids, table, norm, bank) + declined = gemvs.embed_norm(other_ids, table, norm, other_bank) + assert captured is not None and declined is None + graph.replay() + torch.cuda.synchronize() + assert torch.equal(_bits(captured), _bits(eager)) From fabb6e65c8faec427c7c4cd35dcec53e9024ffd0 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 20:03:32 -0700 Subject: [PATCH 040/161] [None][feat] MNNVL all-reduce: Kimi K3 attention-residual epilogue (trtllm::mnnvl_allreduce_attn_res) A one-shot MNNVL all-reduce whose epilogue is Kimi K3's residual update: updated = prefix_sum + the sum over the ranks, and normed = RMSNorm of the attention-residual selection over the snapshot bank and updated. It rounds like the unfused all-reduce followed by attn_res_add_rmsnorm_fwd, and its reduction order is fixed, so every rank gets the same bits. - oneshotAllreduceAttnResKernel: one cluster of tokenDim / 1024 CTAs per token. The per-token statistics are summed in a fixed order, the cluster part through distributed shared memory. Each thread arrives on the cluster barrier early and waits before its first shared-memory write into a peer CTA. - reduceOneshotLamport is oneshotAllreduceFusionKernel's Lamport reduction as a function, for the new kernel. The existing kernels are unchanged. - The schema declares comm_buffer and buffer_flags mutable. A fake registration gives torch.compile the output shapes. Signed-off-by: Vasanth Sabavat --- .../mnnvlAllreduceKernels.cu | 410 ++++++++++++++++++ .../mnnvlAllreduceKernels.h | 28 ++ cpp/tensorrt_llm/thop/allreduceOp.cpp | 79 ++++ .../_torch/custom_ops/cpp_custom_ops.py | 6 + 4 files changed, 523 insertions(+) diff --git a/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllreduceKernels.cu b/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllreduceKernels.cu index e60a2f233c15..3b84bcfa759c 100644 --- a/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllreduceKernels.cu +++ b/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllreduceKernels.cu @@ -15,6 +15,7 @@ */ #include "mnnvlAllreduceKernels.h" #include "tensorrt_llm/common/config.h" +#include #include #include #include @@ -509,6 +510,58 @@ inline __device__ PackedVec reduceOneshotDeterministicFastPath( return reduceOneshotDeterministic(remoteValues, localValue); } +// The one-shot Lamport reduction of oneshotAllreduceFusionKernel as a function, for oneshotAllreduceAttnResKernel. +// Fully deterministic: every rank uses the exact same reduction order. For WorldSize <= 8, specialize the local +// slot so the fast path reuses `val` from registers without a dynamic `remoteValues[rank]` store. Larger world sizes +// use the compact fallback because the benefit is thin but specializing every rank significantly increases compile +// time. +template +inline __device__ PackedVec reduceOneshotLamport( + PackedVec const& val, T* stagePtrLocal, int token, int tokenDim, int packedIdx, int rank) +{ + PackedVec packedAccum; + if constexpr (WorldSize <= 8) + { + packedAccum = val; +#define RUN_ONESHOT_LOCAL_RANK(LOCAL_RANK) \ + case LOCAL_RANK: \ + if constexpr (WorldSize > LOCAL_RANK) \ + { \ + packedAccum = reduceOneshotDeterministicFastPath( \ + val, stagePtrLocal, token, tokenDim, packedIdx); \ + } \ + break + + switch (rank) + { + RUN_ONESHOT_LOCAL_RANK(0); + RUN_ONESHOT_LOCAL_RANK(1); + RUN_ONESHOT_LOCAL_RANK(2); + RUN_ONESHOT_LOCAL_RANK(3); + RUN_ONESHOT_LOCAL_RANK(4); + RUN_ONESHOT_LOCAL_RANK(5); + RUN_ONESHOT_LOCAL_RANK(6); + RUN_ONESHOT_LOCAL_RANK(7); + } +#undef RUN_ONESHOT_LOCAL_RANK + } + else + { + // Chunk Lamport polling so only a bounded rank set is live at once, avoiding register spills for large + // world sizes. + constexpr int kRankChunk = 8; + float accum[kELTS_PER_THREAD]; + accumulateLamportRanksChunked( + accum, stagePtrLocal, token, tokenDim, packedIdx); +#pragma unroll + for (int i = 0; i < kELTS_PER_THREAD; i++) + { + packedAccum.elements[i] = cuda_cast(accum[i]); + } + } + return packedAccum; +} + template inline __device__ void quantizeEpilogue(PackedVec const& value, MnnvlAllReduceKernelParams const& params, int packedAccessIdx, int accessIdInToken, int token) @@ -573,6 +626,7 @@ using detail::MnnvlAllReduceKernelParams; using detail::sanitizeLamportPayload; using detail::accumulateLamportRanksChunked; using detail::reduceOneshotDeterministicFastPath; +using detail::reduceOneshotLamport; using detail::writeEpilogueOutput; template +__global__ void __launch_bounds__(detail::kAttnResThreads) + oneshotAllreduceAttnResKernel(MnnvlAllReduceKernelParams<__nv_bfloat16> params, AttnResEpilogueParams epilogue) +{ +#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) + using T = __nv_bfloat16; + using PackedType = float4; + using Packed = PackedVec; + constexpr int kELTS = detail::kAttnResEltsPerThread; + constexpr int kWarps = detail::kAttnResThreads / detail::kWARP_SIZE; + constexpr int kStats = 2 * N; // Sum of squares and residual-projection dot product per candidate. + constexpr float kLog2E = 1.4426950408889634F; + static_assert(sizeof(PackedType) / sizeof(T) == kELTS); + static_assert(N >= 1 && N <= detail::kAttnResMaxCandidates && N <= 32, "one lane per candidate"); + + namespace cg = cooperative_groups; + cg::cluster_group cluster = cg::this_cluster(); + int const clusterRank = static_cast(cluster.block_rank()); + int const clusterSize = static_cast(cluster.num_blocks()); + int const packedIdx = static_cast(cluster.thread_rank()); + int const token = blockIdx.x; + int const threadOffset = token * params.tokenDim + packedIdx * kELTS; + int const lane = threadIdx.x & detail::kLANE_ID_MASK; + int const warp = threadIdx.x >> detail::kLOG2_WARP_SIZE; + + __shared__ float warpStats[kWarps][kStats]; + __shared__ float clusterStats[detail::kAttnResMaxClusterSize][kStats]; + __shared__ float clusterOutputSq[detail::kAttnResMaxClusterSize]; + __shared__ float candidateWeights[N]; + + // Peers write into this CTA's shared memory only after the matching wait, which also guarantees + // that every CTA of the cluster has started. + asm volatile("barrier.cluster.arrive.relaxed.aligned;\n" ::: "memory"); + + cudaGridDependencySynchronize(); + // Every consumer of the outputs waits for this whole grid, so triggering here only lets the next + // kernel launch and stream its weights while the ranks exchange their contributions. + cudaTriggerProgrammaticLaunchCompletion(); + + LamportFlags flag(params.bufferFlags, 1); + T* stagePtrMcast = reinterpret_cast(flag.getCurLamportBuf(params.mcastPtr, 0)); + T* stagePtrLocal = reinterpret_cast(flag.getCurLamportBuf(params.inputPtrs[params.rank], 0)); + + // ==================== Broadcast tokens to each rank ============================= + Packed val; + val.packed = loadPacked(¶ms.shardPtr[threadOffset]); + sanitizeLamportPayload(val); + reinterpret_cast( + &stagePtrMcast[token * params.tokenDim * WorldSize + params.rank * params.tokenDim])[packedIdx] + = val.packed; + flag.ctaArrive(); + flag.clearDirtyLamportBuf(params.inputPtrs[params.rank], -1); + + // The epilogue operands do not depend on the peers; load them while the peers' data is in flight. + auto const* blockResidual = static_cast(epilogue.blockResidual); + Packed snapshot[N > 1 ? N - 1 : 1]; +#pragma unroll + for (int n = 0; n < N - 1; n++) + { + snapshot[n].packed = loadPacked( + &blockResidual[static_cast(n) * params.numTokens * params.tokenDim + threadOffset]); + } + [[maybe_unused]] Packed prefix; + if constexpr (AddPrefix) + { + prefix.packed = loadPacked(¶ms.residualInPtr[threadOffset]); + } + Packed resWeight; + Packed rmsWeight; + Packed outputRmsWeight; + resWeight.packed = loadPacked(&static_cast(epilogue.resWeight)[packedIdx * kELTS]); + rmsWeight.packed = loadPacked(&static_cast(epilogue.rmsWeight)[packedIdx * kELTS]); + outputRmsWeight.packed + = loadPacked(&static_cast(epilogue.outputRmsWeight)[packedIdx * kELTS]); + + // ======================= Reduction ============================= + Packed const reduced = reduceOneshotLamport( + val, stagePtrLocal, token, params.tokenDim, packedIdx, params.rank); + + // ======================= Residual add: bf16(prefix + bf16(sum)) ============================= + Packed updated; +#pragma unroll + for (int i = 0; i < kELTS; i++) + { + if constexpr (AddPrefix) + { + updated.elements[i] + = __float2bfloat16_rn(__bfloat162float(prefix.elements[i]) + __bfloat162float(reduced.elements[i])); + } + else + { + updated.elements[i] = reduced.elements[i]; + } + } + reinterpret_cast(¶ms.residualOutPtr[threadOffset])[0] = updated.packed; + + // ======================= Attention-residual scores ============================= + // Candidates are the snapshots followed by the updated prefix sum. + float q[kELTS]; +#pragma unroll + for (int i = 0; i < kELTS; i++) + { + q[i] = __bfloat162float(resWeight.elements[i]) * __bfloat162float(rmsWeight.elements[i]); + } + float stats[kStats]; +#pragma unroll + for (int n = 0; n < N; n++) + { + Packed const& candidate = n < N - 1 ? snapshot[n] : updated; + float sumSq = 0.F; + float dot = 0.F; +#pragma unroll + for (int i = 0; i < kELTS; i++) + { + float const v = __bfloat162float(candidate.elements[i]); + sumSq = fmaf(v, v, sumSq); + dot = fmaf(v, q[i], dot); + } + stats[2 * n] = sumSq; + stats[2 * n + 1] = dot; + } +#pragma unroll + for (int offset = detail::kWARP_SIZE / 2; offset > 0; offset >>= 1) + { +#pragma unroll + for (int s = 0; s < kStats; s++) + { + stats[s] += __shfl_down_sync(0xffffffffU, stats[s], offset); + } + } + if (lane == 0) + { +#pragma unroll + for (int s = 0; s < kStats; s++) + { + warpStats[warp][s] = stats[s]; + } + } + __syncthreads(); + + // Each CTA's partial goes into slot clusterRank of every CTA, so all of them sum the same values in the same + // order. + asm volatile("barrier.cluster.wait.aligned;\n" ::: "memory"); + for (int i = threadIdx.x; i < clusterSize * kStats; i += blockDim.x) + { + int const peer = i / kStats; + int const s = i % kStats; + float partial = 0.F; +#pragma unroll + for (int w = 0; w < kWarps; w++) + { + partial += warpStats[w][s]; + } + *cluster.map_shared_rank(&clusterStats[clusterRank][s], peer) = partial; + } + cluster.sync(); + + // Warp 0 turns the cluster totals into the softmax weights; lane n owns candidate n. + if (warp == 0) + { + float logit = -FLT_MAX; + if (lane < N) + { + float sumSq = 0.F; + float dot = 0.F; + for (int r = 0; r < clusterSize; r++) + { + sumSq += clusterStats[r][2 * lane]; + dot += clusterStats[r][2 * lane + 1]; + } + logit = dot * rsqrtf(sumSq / params.tokenDim + epilogue.rmsEps); + } + float maxLogit = logit; +#pragma unroll + for (int offset = detail::kWARP_SIZE / 2; offset > 0; offset >>= 1) + { + maxLogit = fmaxf(maxLogit, __shfl_xor_sync(0xffffffffU, maxLogit, offset)); + } + float const weight = lane < N ? exp2f((logit - maxLogit) * kLog2E) : 0.F; + float denominator = weight; +#pragma unroll + for (int offset = detail::kWARP_SIZE / 2; offset > 0; offset >>= 1) + { + denominator += __shfl_xor_sync(0xffffffffU, denominator, offset); + } + if (lane < N) + { + candidateWeights[lane] = weight * (1.F / denominator); + } + } + __syncthreads(); + + // ======================= Selection and its RMSNorm ============================= + float weights[N]; +#pragma unroll + for (int n = 0; n < N; n++) + { + weights[n] = candidateWeights[n]; + } + Packed mixed; + float outputSq = 0.F; +#pragma unroll + for (int i = 0; i < kELTS; i++) + { + float value = 0.F; +#pragma unroll + for (int n = 0; n < N; n++) + { + Packed const& candidate = n < N - 1 ? snapshot[n] : updated; + value = fmaf(weights[n], __bfloat162float(candidate.elements[i]), value); + } + mixed.elements[i] = __float2bfloat16_rn(value); + float const rounded = __bfloat162float(mixed.elements[i]); + outputSq = fmaf(rounded, rounded, outputSq); + } +#pragma unroll + for (int offset = detail::kWARP_SIZE / 2; offset > 0; offset >>= 1) + { + outputSq += __shfl_down_sync(0xffffffffU, outputSq, offset); + } + // Every thread has read warpStats before the cluster barrier above, so it can be reused. + if (lane == 0) + { + warpStats[warp][0] = outputSq; + } + __syncthreads(); + if (threadIdx.x < clusterSize) + { + float partial = 0.F; +#pragma unroll + for (int w = 0; w < kWarps; w++) + { + partial += warpStats[w][0]; + } + *cluster.map_shared_rank(&clusterOutputSq[clusterRank], threadIdx.x) = partial; + } + cluster.sync(); + float totalSq = 0.F; + for (int r = 0; r < clusterSize; r++) + { + totalSq += clusterOutputSq[r]; + } + float const outputRsigma = rsqrtf(totalSq / params.tokenDim + epilogue.outputRmsEps); + Packed output; +#pragma unroll + for (int i = 0; i < kELTS; i++) + { + // KimiK3RMSNorm: normalize in fp32, round to bf16, then apply the bf16 weight. + T const normalized = __float2bfloat16_rn(__bfloat162float(mixed.elements[i]) * outputRsigma); + output.elements[i] + = __float2bfloat16_rn(__bfloat162float(normalized) * __bfloat162float(outputRmsWeight.elements[i])); + } + reinterpret_cast(¶ms.outputPtr[threadOffset])[0] = output.packed; + flag.waitAndUpdate({static_cast(params.numTokens * params.tokenDim * WorldSize * sizeof(T)), 0, 0, 0}); +#endif +} + +namespace +{ + +template +void launchOneshotAllreduceAttnRes(cudaLaunchConfig_t const& config, + MnnvlAllReduceKernelParams<__nv_bfloat16> const& kernelParams, AttnResEpilogueParams const& epilogue, + bool addPrefix) +{ + if (addPrefix) + { + TLLM_CUDA_CHECK( + cudaLaunchKernelEx(&config, &oneshotAllreduceAttnResKernel, kernelParams, epilogue)); + } + else + { + TLLM_CUDA_CHECK( + cudaLaunchKernelEx(&config, &oneshotAllreduceAttnResKernel, kernelParams, epilogue)); + } +} + +template +void dispatchOneshotAllreduceAttnRes(cudaLaunchConfig_t const& config, + MnnvlAllReduceKernelParams<__nv_bfloat16> const& kernelParams, AttnResEpilogueParams const& epilogue, + bool addPrefix, std::integer_sequence) +{ + bool const launched + = ((epilogue.numCandidates == CandidateIdx + 1 ? ( + launchOneshotAllreduceAttnRes(config, kernelParams, epilogue, addPrefix), + true) + : false) + || ...); + TLLM_CHECK_WITH_INFO(launched, "[MNNVL AllReduceAttnRes] unsupported number of candidates %d (1-%d).", + epilogue.numCandidates, detail::kAttnResMaxCandidates); +} + +} // namespace + +void oneshotAllreduceAttnResOp(AllReduceFusionParams const& params, AttnResEpilogueParams const& epilogue) +{ + static int const kSMVersion = tensorrt_llm::common::getSMVersion(); + TLLM_CHECK_WITH_INFO(kSMVersion >= 90, "[MNNVL AllReduceAttnRes] requires SM 90 or newer."); + TLLM_CHECK_WITH_INFO( + params.dType == tensorrt_llm::DataType::kBF16, "[MNNVL AllReduceAttnRes] supports BF16 tensors only."); + int const clusterSize = params.tokenDim / detail::kAttnResEltsPerCta; + TLLM_CHECK_WITH_INFO(params.tokenDim % detail::kAttnResEltsPerCta == 0 && clusterSize >= 1 + && clusterSize <= detail::kAttnResMaxClusterSize, + "[MNNVL AllReduceAttnRes] hidden dimension %d must be a multiple of %d and at most %d.", params.tokenDim, + detail::kAttnResEltsPerCta, detail::kAttnResEltsPerCta * detail::kAttnResMaxClusterSize); + + cudaLaunchAttribute attrs[2]; + attrs[0].id = cudaLaunchAttributeProgrammaticStreamSerialization; + attrs[0].val.programmaticStreamSerializationAllowed = tensorrt_llm::common::getEnvEnablePDL() ? 1 : 0; + attrs[1].id = cudaLaunchAttributeClusterDimension; + attrs[1].val.clusterDim.x = 1; + attrs[1].val.clusterDim.y = clusterSize; + attrs[1].val.clusterDim.z = 1; + cudaLaunchConfig_t config{ + .gridDim = dim3(params.numTokens, clusterSize, 1), + .blockDim = detail::kAttnResThreads, + .dynamicSmemBytes = 0, + .stream = params.stream, + .attrs = attrs, + .numAttrs = 2U, + }; + + using T = __nv_bfloat16; + MnnvlAllReduceKernelParams kernelParams{reinterpret_cast(params.output), + reinterpret_cast(params.residualOut), reinterpret_cast(params.input), + reinterpret_cast(params.residualIn), nullptr, reinterpret_cast(params.bufferPtrsDev), + reinterpret_cast(params.bufferPtrLocal), reinterpret_cast(params.multicastPtr), nullptr, nullptr, + nullptr, params.numTokens, params.tokenDim, params.nRanks, params.rank, 0.F, params.bufferFlags, false, + params.layout}; + bool const addPrefix = params.residualIn != nullptr; + auto constexpr kCandidates = std::make_integer_sequence{}; + + switch (params.nRanks) + { + case 2: dispatchOneshotAllreduceAttnRes<2>(config, kernelParams, epilogue, addPrefix, kCandidates); break; + case 4: dispatchOneshotAllreduceAttnRes<4>(config, kernelParams, epilogue, addPrefix, kCandidates); break; + case 8: dispatchOneshotAllreduceAttnRes<8>(config, kernelParams, epilogue, addPrefix, kCandidates); break; + case 16: dispatchOneshotAllreduceAttnRes<16>(config, kernelParams, epilogue, addPrefix, kCandidates); break; + default: + TLLM_CHECK_WITH_INFO( + false, "[MNNVL AllReduceAttnRes] unsupported world size %d (2, 4, 8 or 16).", params.nRanks); + } +} + enum MNNVLTwoShotStage : uint8_t { SCATTER = 0, diff --git a/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllreduceKernels.h b/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllreduceKernels.h index f2006f52240c..9ab0ef0a6581 100644 --- a/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllreduceKernels.h +++ b/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllreduceKernels.h @@ -73,8 +73,36 @@ struct AllReduceFusionParams //! @} }; +/** + * \brief Kimi K3 attention-residual epilogue of oneshotAllreduceAttnResOp (BF16 tensors). + */ +struct AttnResEpilogueParams +{ + void const* blockResidual; //!< [numCandidates - 1, numTokens, tokenDim] snapshots; unused when numCandidates == 1 + void const* resWeight; //!< [tokenDim] attention-residual projection + void const* rmsWeight; //!< [tokenDim] weight of the RMSNorm inside the attention-residual score + void const* outputRmsWeight; //!< [tokenDim] weight of the RMSNorm applied to the selected residual + float rmsEps; + float outputRmsEps; + int numCandidates; //!< Snapshots plus the running prefix sum, 1..12 +}; + void oneshotAllreduceFusionOp(AllReduceFusionParams const& params); void twoshotAllreduceFusionOp(AllReduceFusionParams const& params); + +/** + * \brief One-shot all-reduce with Kimi K3's residual add, attention-residual selection and RMSNorm as the epilogue. + * + * Per token, with r = bf16(allreduce(input)): + * residualOut = bf16(residualIn + r), or r when residualIn is null + * output = RMSNorm(attn_res(blockResidual[0 .. numCandidates-2], residualOut), outputRmsWeight) + * with the rounding of the unfused sequence all-reduce -> attn_res_add_rmsnorm_fwd. The reduction order is fixed, so + * the result is deterministic and identical on every rank. + * + * Requirements: BF16; tokenDim a multiple of 1024 and at most 8192; the one-shot footprint + * numTokens * tokenDim * nRanks elements fits in one Lamport buffer; nRanks in {2, 4, 8, 16}. + */ +void oneshotAllreduceAttnResOp(AllReduceFusionParams const& params, AttnResEpilogueParams const& epilogue); } // namespace kernels::mnnvl TRTLLM_NAMESPACE_END diff --git a/cpp/tensorrt_llm/thop/allreduceOp.cpp b/cpp/tensorrt_llm/thop/allreduceOp.cpp index acde3a3d3cd2..51e1be287964 100644 --- a/cpp/tensorrt_llm/thop/allreduceOp.cpp +++ b/cpp/tensorrt_llm/thop/allreduceOp.cpp @@ -2215,6 +2215,80 @@ std::vector mnnvlFusionAllReduce(torch::Tensor& input, torch::opt return {}; } +// Kimi K3 pre-MoE residual update in the MNNVL one-shot all-reduce epilogue; see +// tensorrt_llm::kernels::mnnvl::oneshotAllreduceAttnResOp. Returns {normed, updated_prefix_sum}. +std::vector mnnvlAllReduceAttnRes(torch::Tensor const& input, + torch::optional const& prefix_sum, torch::Tensor const& block_residual, + torch::Tensor const& res_weight, torch::Tensor const& rms_weight, torch::Tensor const& output_rms_weight, + double rms_eps, double output_rms_eps, torch::Tensor& comm_buffer, torch::Tensor& buffer_flags) +{ + auto* mcast_mem = tensorrt_llm::common::findMcastDevMemBuffer(comm_buffer.data_ptr()); + TORCH_CHECK( + mcast_mem != nullptr, "[mnnvlAllReduceAttnRes] comm_buffer must be obtained from a mcastBuffer instance."); + TORCH_CHECK(mcast_mem->isMapped(), "[mnnvlAllReduceAttnRes] MNNVL workspace handles are not attached."); + auto const checkBf16 = [](torch::Tensor const& tensor, char const* name) + { + TORCH_CHECK(tensor.is_cuda() && tensor.scalar_type() == torch::kBFloat16 && tensor.is_contiguous(), + "[mnnvlAllReduceAttnRes] ", name, " must be a contiguous CUDA bfloat16 tensor"); + }; + checkBf16(input, "input"); + checkBf16(block_residual, "block_residual"); + checkBf16(res_weight, "res_weight"); + checkBf16(rms_weight, "rms_weight"); + checkBf16(output_rms_weight, "output_rms_weight"); + TORCH_CHECK(input.dim() == 2, "[mnnvlAllReduceAttnRes] input must be [num_tokens, hidden]"); + int64_t const numTokens = input.size(0); + int64_t const hiddenDim = input.size(1); + if (prefix_sum.has_value()) + { + checkBf16(prefix_sum.value(), "prefix_sum"); + TORCH_CHECK(prefix_sum.value().sizes() == input.sizes(), + "[mnnvlAllReduceAttnRes] prefix_sum must have the shape of input"); + } + TORCH_CHECK(block_residual.dim() == 3 && block_residual.size(1) == numTokens && block_residual.size(2) == hiddenDim, + "[mnnvlAllReduceAttnRes] block_residual must be [num_snapshots, num_tokens, hidden]"); + for (auto const* weight : {&res_weight, &rms_weight, &output_rms_weight}) + { + TORCH_CHECK(weight->dim() == 1 && weight->size(0) == hiddenDim, + "[mnnvlAllReduceAttnRes] res_weight, rms_weight and output_rms_weight must be [hidden]"); + } + int64_t const nRanks = mcast_mem->getWorldSize(); + TORCH_CHECK(numTokens * hiddenDim * nRanks <= comm_buffer.size(-1), + "[mnnvlAllReduceAttnRes] the one-shot footprint of ", numTokens * hiddenDim * nRanks, + " elements exceeds one Lamport buffer of ", comm_buffer.size(-1), " elements"); + + torch::Tensor normOut = torch::empty_like(input); + torch::Tensor prefixOut = torch::empty_like(input); + + auto params = tensorrt_llm::kernels::mnnvl::AllReduceFusionParams(); + params.nRanks = static_cast(nRanks); + params.rank = mcast_mem->getRank(); + params.dType = tensorrt_llm::DataType::kBF16; + params.numTokens = static_cast(numTokens); + params.tokenDim = static_cast(hiddenDim); + params.bufferPtrsDev = reinterpret_cast(mcast_mem->getBufferPtrsDev()); + params.bufferPtrLocal = comm_buffer.mutable_data_ptr(); + params.multicastPtr = mcast_mem->getMulticastPtr(); + params.bufferFlags = reinterpret_cast(buffer_flags.mutable_data_ptr()); + params.input = input.const_data_ptr(); + params.residualIn = prefix_sum.has_value() ? prefix_sum.value().const_data_ptr() : nullptr; + params.residualOut = prefixOut.mutable_data_ptr(); + params.output = normOut.mutable_data_ptr(); + params.stream = at::cuda::getCurrentCUDAStream(input.get_device()); + + tensorrt_llm::kernels::mnnvl::AttnResEpilogueParams epilogue{}; + epilogue.blockResidual = block_residual.const_data_ptr(); + epilogue.resWeight = res_weight.const_data_ptr(); + epilogue.rmsWeight = rms_weight.const_data_ptr(); + epilogue.outputRmsWeight = output_rms_weight.const_data_ptr(); + epilogue.rmsEps = static_cast(rms_eps); + epilogue.outputRmsEps = static_cast(output_rms_eps); + epilogue.numCandidates = static_cast(block_residual.size(0)) + 1; + + tensorrt_llm::kernels::mnnvl::oneshotAllreduceAttnResOp(params, epilogue); + return {normOut, prefixOut}; +} + torch::Tensor minimax_allreduce_rms(torch::Tensor const& input, torch::Tensor const& norm_weight, torch::Tensor workspace, int64_t const rank, int64_t const nranks, double const eps, bool const trigger_completion_at_end_) @@ -2316,6 +2390,10 @@ TORCH_LIBRARY_FRAGMENT(trtllm, m) "float? epsilon, Tensor(a!) comm_buffer, Tensor buffer_flags, bool rmsnorm_fusion, " "Tensor? scale=None, int fusion_op=0) -> " "Tensor[]"); + m.def( + "mnnvl_allreduce_attn_res(Tensor input, Tensor? prefix_sum, Tensor block_residual, Tensor res_weight, " + "Tensor rms_weight, Tensor output_rms_weight, float rms_eps, float output_rms_eps, Tensor(a!) comm_buffer, " + "Tensor(b!) buffer_flags) -> Tensor[]"); m.def( "allreduce(" "Tensor input," @@ -2414,6 +2492,7 @@ TORCH_LIBRARY_FRAGMENT(trtllm, m) TORCH_LIBRARY_IMPL(trtllm, CUDA, m) { m.impl("mnnvl_fusion_allreduce", &tensorrt_llm::torch_ext::mnnvlFusionAllReduce); + m.impl("mnnvl_allreduce_attn_res", &tensorrt_llm::torch_ext::mnnvlAllReduceAttnRes); m.impl("allreduce", &tensorrt_llm::torch_ext::allreduce_raw); m.impl("autotuned_allreduce", &tensorrt_llm::torch_ext::autotunedAllreduce); m.impl("register_allreduce_tactic", &tensorrt_llm::torch_ext::registerAllReduceTactic); diff --git a/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py b/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py index d1b719f3a3fc..fb7560d1b7cc 100644 --- a/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py +++ b/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py @@ -121,6 +121,12 @@ def _( return allreduce(input, residual, norm_weight, scale, bias, workspace, group, strategy, op, eps, trigger_completion_at_end) + @torch.library.register_fake("trtllm::mnnvl_allreduce_attn_res") + def _(input, prefix_sum, block_residual, res_weight, rms_weight, + output_rms_weight, rms_eps, output_rms_eps, comm_buffer, + buffer_flags): + return [torch.empty_like(input), torch.empty_like(input)] + # MNNVL Allreduce @torch.library.register_fake("trtllm::mnnvl_fusion_allreduce") def _(input, From 494b1d9ff42b292c04858506a3a6efa875d9dac6 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 17:57:39 -0700 Subject: [PATCH 041/161] [None][feat] modeling_v2 catalog: stateful entries; comm/mnnvl_allreduce_attn_res over a caller-owned MnnvlWorkspace - Conventions (catalog/index.yaml) for entries whose result depends on state that outlives a call. The state is a typed object the caller owns and passes in, built by an explicit constructor (collective, eager, before any capture). The contract gains a "## State" section, and the test drives call sequences and a negative control, not only single calls. - comm/mnnvl_allreduce_attn_res: contract and wrapper, plus the state type comm/mnnvl_workspace.py (three Lamport buffers, the flag words, the multicast handle). Certified on sm_100 (GB200) at 4 ranks. - Its matrix, _mnnvl_allreduce_attn_res_op_matrix.py (shared helpers in _lockstep.py), runs these checks: - single calls over the shape grid; - 12 steps of 12 chained layers, with the token count dipping and growing back and a random rank late at every call; - two workspaces interleaved; - a captured step replayed between eager calls; - a call over one buffer, which must raise on every rank; - a negative control, in which one rank swaps two calls. - The matrix takes --world-size and --launcher: mpirun on one node, or srun for N ranks across nodes. _rank_job.run takes the world size. - CI: the matrix runs at 4 ranks in l0_gb200_multi_gpus.yml. Signed-off-by: Vasanth Sabavat --- .../_experimental/modeling_v2/README.md | 18 +- .../catalog/comm/mnnvl_allreduce_attn_res.md | 159 +++++++++++ .../catalog/comm/mnnvl_allreduce_attn_res.py | 41 +++ .../catalog/comm/mnnvl_workspace.py | 104 +++++++ .../modeling_v2/catalog/index.yaml | 28 +- .../test-db/l0_gb200_multi_gpus.yml | 1 + .../_torch/modeling_v2/comm/_lockstep.py | 167 +++++++++++ .../_mnnvl_allreduce_attn_res_op_matrix.py | 259 ++++++++++++++++++ .../_torch/modeling_v2/comm/_rank_job.py | 24 +- ...g_v2_mnnvl_allreduce_attn_res_op_matrix.py | 19 ++ 10 files changed, 803 insertions(+), 17 deletions(-) create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_workspace.py create mode 100644 tests/unittest/_torch/modeling_v2/comm/_lockstep.py create mode 100644 tests/unittest/_torch/modeling_v2/comm/_mnnvl_allreduce_attn_res_op_matrix.py create mode 100644 tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_allreduce_attn_res_op_matrix.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/README.md b/tensorrt_llm/_torch/_experimental/modeling_v2/README.md index 61be49f5be0d..cf0738128514 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/README.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/README.md @@ -115,8 +115,9 @@ tests/unittest/_torch/modeling_v2/ test_modeling_v2_claims.py routing tables vs the targets they name (no GPU) test_modeling_v2_routing.py what modeling_v2_resolve does (no GPU) /test_modeling_v2_.py - comm/__op_matrix.py the two collectives' 4-rank rank bodies - comm/_rank_job.py starts one of those and asserts on its exit code + comm/__op_matrix.py each collective's rank bodies + comm/_lockstep.py the stateful matrices' launcher and shared reference + comm/_rank_job.py starts one at a world size and asserts on its exit code ``` That split is not a preference; it is where this repo's CI collects from, and @@ -124,13 +125,17 @@ an in-package test is on no list. It costs one thing worth stating: a receipt is valid only if it post-dates the last write to *every* file of its entry, so that check now has to look in both trees. -The two collectives' rank bodies sit in the tests tree with everything else +The collectives' rank bodies sit in the tests tree with everything else that only tests run. Both halves are started by file path -- the launcher must not import `tensorrt_llm`, because that calls `MPI_Init` and an MPI-initialized process cannot start `mpirun`, and the ranks reach the catalog by absolute import -- so neither needs a package to live in. Neither those file names nor their `check_*` bodies match pytest's collection patterns: -each is one fixed 4-rank sequence that cannot run as independent cases. +each is one fixed sequence over one job's ranks that cannot run as +independent cases. CI runs them at 4 ranks on one node. A stateful entry's +launcher also takes `--world-size N --launcher srun`, which runs the same +body as one of N ranks started across nodes: that is how its multi-node +receipt is recorded. Identity is the directory name, and it carries all three segments: `gpt_oss_120b__sm_103__tp1`. They were three nested directories once, which @@ -181,8 +186,9 @@ Perf is measured, never gated. ## Status of every record in this tree -**The catalog is fully certified on sm_103. The targets construct but have -never executed.** +**The catalog is certified: 19 entries on sm_103, and +`comm/mnnvl_allreduce_attn_res` on sm_100 (GB200), where its first caller +runs. The targets construct but have never executed.** Two things voided every receipt in the move: each catalog test file was rewritten, and the targets moved from sm_100 (B200) to sm_103 (GB300), where diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.md new file mode 100644 index 000000000000..f6cc1af2f337 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.md @@ -0,0 +1,159 @@ +--- +receipts: + sm_100: {status: passed, world_size: 4} +--- + +# mnnvl_allreduce_attn_res + +**Wraps** `torch.ops.trtllm.mnnvl_allreduce_attn_res` (one call). + +A stateful entry: its correctness depends on a caller-owned state object, `MnnvlWorkspace`, passed as the last +argument and described under *State*. + +## Semantics + +A one-shot all-reduce over a TP group's MNNVL multicast workspace with Kimi K3's residual update as its epilogue. +Every rank of the group calls with its own rows `input` `[T, H]`; every rank gets back the same two tensors. Per +token, with `W` ranks: + +``` +r = bf16(sum over the W ranks of input) # every rank sums every rank's rows, fixed order +updated = bf16(prefix_sum + r) # r alone when prefix_sum is None +v = [block_residual[0], ..., block_residual[S-1], updated] # S + 1 candidates +score_c = sum_h rmsnorm(v_c)[h] * rms_weight[h] * res_weight[h] # rmsnorm: v_c / sqrt(mean(v_c^2) + rms_eps) +p = softmax over the candidates of score +normed = RMSNorm(sum_c p_c v_c; output_rms_weight, output_rms_eps) +``` + +and returns `(normed, updated)`. The selection is HF's `modeling_kimi._apply_attn_res`; the op rounds like the +unfused all-reduce followed by `trtllm::attn_res_add_rmsnorm_fwd` (its kernel's statement). The reduction order is +fixed: the result is deterministic and bitwise identical on every rank (certified, every call of the test). + +Fusion boundary. Inside: the all-reduce, the residual add, the attention-residual selection, the RMSNorm. Outside: +whatever produced `input` (a row-parallel projection), keeping the snapshot bank `block_residual` (which layers push +a snapshot, and when), and `prefix_sum`'s chaining from layer to layer. + +## Signature + +```python +def mnnvl_allreduce_attn_res( + input: torch.Tensor, + prefix_sum: Optional[torch.Tensor], + block_residual: torch.Tensor, + res_weight: torch.Tensor, + rms_weight: torch.Tensor, + output_rms_weight: torch.Tensor, + rms_eps: float, + output_rms_eps: float, + workspace: MnnvlWorkspace, +) -> Tuple[torch.Tensor, torch.Tensor] +``` + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `input` | `[T, H]`; `H` = 7168 certified (the op takes a multiple of 1024 up to 8192); `T` 1-8 and 16 certified, at most `workspace.max_one_shot_tokens(H)` | bf16 | contiguous | CUDA, this rank's device | +| `prefix_sum` | `None`, or `[T, H]` | bf16 | contiguous | CUDA | +| `block_residual` | `[S, T, H]`, `S` = 0..11 (certified at 0, 1, 2, 3, 5, 6, 7, 8, 9, 10, 11) | bf16 | contiguous | CUDA | +| `res_weight`, `rms_weight`, `output_rms_weight` | `[H]` | bf16 | contiguous | CUDA | +| `rms_eps`, `output_rms_eps` | scalar | Python float | — | — | +| `workspace` | an `MnnvlWorkspace` of this rank's TP group (see *State*) | — | — | — | +| returns | `(normed, updated)`, each `[T, H]` | bf16 | contiguous, newly allocated | = `input.device` | + +`input`, `prefix_sum` and `block_residual` are read only. The op writes the workspace's buffers and flag words; its +schema declares both mutable. + +## State + +**Object.** `MnnvlWorkspace` (`catalog/comm/mnnvl_workspace.py`), one per TP group, owned by the caller. A state +type, not an entry: it launches nothing per call. + +**Contents and size.** Three Lamport buffers of `buffer_bytes` each in one multicast allocation (this rank's +unicast view as `lamport`; every word `-0.0` when armed), the flag words `buffer_flags` (uint32 `[9]`: current +buffer, dirty buffer, bytes per buffer, dirty stages, bytes to clear x 4, access count), the `McastGPUBuffer` +handle that owns the memory, and the communicator the handles were exchanged over. A call of `T` tokens pushes +`T x H x W x 2` bytes into one buffer, so `buffer_bytes` must cover the largest call: Kimi K3 uses 4 MiB, i.e. +`T` <= 73 at `W` = 4 and `T` <= 18 at `W` = 16 for `H` = 7168 (`max_one_shot_tokens`). The workspace is never grown; +a call over one buffer raises (below). + +**Who creates it, and when.** The target, in `post_load_weights`, with +`MnnvlWorkspace.create(mapping, buffer_bytes, fabric_handle=None)`: + +- collective over `mapping`'s TP group: every rank calls it at the same point; it returns on every rank or raises + on every rank (the success of each rank's allocation is agreed before anyone proceeds); +- eager: it allocates and exchanges handles, so it refuses to run under CUDA-graph capture; +- it arms every buffer word and the flags, and returns only once every rank has armed its buffers, so no peer can + push into memory a rank has not armed; +- `fabric_handle`: share the memory by fabric handle (required across nodes) or POSIX file descriptor; default + `mapping.is_multi_node()`. No environment variable is read. + +**Which ops may share one object.** Every MNNVL one-shot op that takes the pair (`comm_buffer`, `buffer_flags`): +this entry and `trtllm::mnnvl_fusion_allreduce`. Ops that share one object share one rotation: their calls form one +sequence, interleaved in one order on every rank. Two objects are two independent rotations: calls alternating +between two workspaces in an irregular pattern, so that the two positions differ, are all correct (certified, 20 +calls). They are not independent *orders*, though: each call spins until its peers' rows of the same call arrive and +the calls of one stream run one after the other, so ranks that issue calls on two workspaces in different relative +orders on one stream deadlock (measured at `W` = 4: rank 0 issued a pair B-then-A, the others A-then-B; every GPU +spun at 100 % until the job was cancelled). The order of all of a stream's collectives must agree across ranks, +whatever object each belongs to. + +**Call-order invariant.** Every rank of the group makes the same sequence of calls on one workspace — the same +number of calls, the `k`-th call on every rank with the same `T` — across layers and decode steps, eager calls and +graph replays alike; and on one stream, the same order of calls across workspaces (above). Each call takes the next +buffer (the current one plus one, mod 3), pushes its rows into that buffer on every rank, and polls its own copy +until every rank's rows of *this* call are there. + +**What a later launch reads.** `buffer_flags`, which the previous call left: the current buffer, the dirty buffer +(the previous call's) and the bytes the previous call wrote into it; and the Lamport words of its own buffer, which +must all be `-0.0` except for this call's pushes. + +**How it is re-armed.** Each call clears the previous call's buffer by the bytes the previous call recorded, not by +its own size, and records its own (`cpp/tensorrt_llm/common/lamportUtils.cuh`, `LamportFlags`). So after a call with +fewer tokens than an earlier one, every word the earlier call pushed is cleared before that buffer comes round again. +Certified by the sequence below. + +**Why the test drives call sequences.** A re-arm sized by the current call instead passes every single-call test: +it fails only after a smaller call, when an older, larger call's words stay in the buffer and a later larger call +reads them as fresh rows whenever a peer has not pushed yet. In serving, that makes ranks disagree, then hang. This +entry's test therefore runs that shape of sequence: 12 decode steps of 12 chained layers at `T` = 8, 8, 8, 2, 7, 8, +1, 1, 8, 16, 3, 8, a random rank 5 ms late at every call, each call against the reference. + +**What a wrong order does.** Measured at `W` = 4 (the test's negative control): rank 0 issues two same-shaped calls on +one workspace in swapped order. Nothing raises and nothing hangs — the rotation positions still agree — but every +rank's two results are wrong (each call paired with the peers' call at the same position; more than half the +`updated` elements differ on every rank). A plain call right after is correct again: a swapped pair realigns the +positions. A rank making one call more or fewer than its peers was not exercised; its positions never realign. + +## Metadata consumed + +None besides `workspace`, which is an explicit argument. The op keeps no process cache and compiles nothing. + +## Preconditions + +- bf16, contiguous, `input` 2-D; `H` a multiple of 1024 and at most 8192; `W` in {2, 4, 8, 16}; `S` <= 11. +- `T x H x W x 2 <= workspace.buffer_bytes`. A call over one buffer raises `RuntimeError` ("the one-shot footprint + ... exceeds one Lamport buffer") on every rank before it touches the workspace: certified, and the next call is + correct. +- Every rank calls with the same `T`, `S` and `prefix_sum` presence; the call order is the *State* invariant. +- `workspace` was created before any capture. Calls may be captured: certified with a captured step of 12 chained + calls at `T` = 8 replayed 8 times with rewritten inputs, an eager call of another `T` on the same workspace between + replays, every replayed and eager call against the reference. + +## Notes + +- Certified path: 4 ranks of one GB200 tray (sm_100), one rank per GPU, POSIX-fd handles, `H` 7168. Test: + `tests/unittest/_torch/modeling_v2/comm/_mnnvl_allreduce_attn_res_op_matrix.py` (the reference is native torch: + the sum's inputs are multiples of 1/16, so `updated` is exact whatever the summation order and is compared bit for + bit; `normed` against an fp32 reference within 2e-2 of its largest magnitude). +- World size: the matrix takes `--world-size` and `--launcher` (`mpirun` on one node, `srun` across nodes). Kimi K3 + runs the op over 16 ranks on four GB200 trays, where the kernel sums ranks in chunks of 8, which a 4-rank run never + reaches. At 16 ranks (four trays, fabric handles) the same kernel passed a recorded check of 110 cases (`T` 1-8, + 16, 32, 64; 0, 1, 3, 8 and 11 snapshots; with and without `prefix_sum`) against the unfused all-reduce followed by + `attn_res_add_rmsnorm_fwd`; this matrix itself has not run at 16 ranks. +- The op finds its multicast mapping by looking `comm_buffer`'s data pointer up in a process registry of multicast + buffers, which the workspace's handle keeps registered, rather than taking the handle as an argument (as + `mnnvl_fusion_allreduce` does). +- `MNNVLAllReduce` keeps its own workspaces in a dict keyed by `Mapping`, grown on demand by the first eager call + that needs more. This entry takes the explicit object instead; its size is the caller's decision, made at + construction. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.py new file mode 100644 index 000000000000..fc02c49b4b53 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.py @@ -0,0 +1,41 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""One-shot MNNVL all-reduce with Kimi K3's residual update (attention-residual selection + RMSNorm) as its +epilogue, over a caller-owned :class:`MnnvlWorkspace`.""" + +from typing import Optional, Tuple + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + +from .mnnvl_workspace import MnnvlWorkspace + + +def mnnvl_allreduce_attn_res( + input: torch.Tensor, + prefix_sum: Optional[torch.Tensor], + block_residual: torch.Tensor, + res_weight: torch.Tensor, + rms_weight: torch.Tensor, + output_rms_weight: torch.Tensor, + rms_eps: float, + output_rms_eps: float, + workspace: MnnvlWorkspace, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Return ``(normed, updated)``: ``updated = prefix_sum + allreduce(input)`` over ``workspace``'s TP group (the + sum alone without ``prefix_sum``), ``normed`` = RMSNorm(attn_res(block_residual..., updated)). Advances + ``workspace`` by one call: every rank of the group makes the same calls on it in the same order.""" + normed, updated = torch.ops.trtllm.mnnvl_allreduce_attn_res( + input, + prefix_sum, + block_residual, + res_weight, + rms_weight, + output_rms_weight, + rms_eps, + output_rms_eps, + workspace.comm_buffer(input.dtype), + workspace.buffer_flags, + ) + return normed, updated diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_workspace.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_workspace.py new file mode 100644 index 000000000000..d8d7eec9bbe5 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_workspace.py @@ -0,0 +1,104 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Caller-owned state of the MNNVL one-shot collectives: one multicast workspace of a TP group. + +A state type, not an entry: it launches nothing per call. Its constructor is collective and eager, the +target builds one in ``post_load_weights`` (before any CUDA-graph capture) and passes it to every MNNVL +one-shot op of that group (``comm/mnnvl_allreduce_attn_res``, and ``trtllm::mnnvl_fusion_allreduce`` given +the same ``comm_buffer`` / ``buffer_flags``). The contract is the ``## State`` section of +``mnnvl_allreduce_attn_res.md``. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, Optional + +import torch + +NUM_LAMPORT_BUFFERS = 3 +FLAG_WORDS = 9 +"""``buffer_flags``: [current buffer, dirty buffer, bytes per buffer, dirty stages, bytes to clear x 4, +access count].""" + + +@dataclass(eq=False) +class MnnvlWorkspace: + """One TP group's MNNVL Lamport workspace: three buffers of ``buffer_bytes`` behind one multicast mapping, and + the flag words that rotate them. Every MNNVL call on it takes the next buffer, so all of a group's ranks must + make the same calls on it in the same order (see the contract's ``## State``).""" + + lamport: torch.Tensor + """fp32 [3 * buffer_bytes / 4]: this rank's unicast view of the three buffers (every word -0.0 when armed).""" + buffer_flags: torch.Tensor + """uint32 [9], see ``FLAG_WORDS``; read and advanced by every call.""" + buffer_bytes: int + rank: int + world_size: int + handle: Any + """The ``McastGPUBuffer`` that owns the memory; the workspace is valid while this object lives.""" + comm: Any + """The TP-group communicator the handles were exchanged over.""" + + def comm_buffer(self, dtype: torch.dtype) -> torch.Tensor: + """The three buffers as the ops take them: ``dtype`` [3, buffer_bytes / itemsize] (a view, no copy).""" + return self.lamport.view(dtype).view(NUM_LAMPORT_BUFFERS, -1) + + def max_one_shot_tokens(self, hidden: int, dtype: torch.dtype = torch.bfloat16) -> int: + """The most tokens of ``hidden`` columns a one-shot call fits in one buffer.""" + itemsize = torch.empty((), dtype=dtype).element_size() + return self.buffer_bytes // (hidden * self.world_size * itemsize) + + @classmethod + def create( + cls, mapping, buffer_bytes: int, fabric_handle: Optional[bool] = None + ) -> "MnnvlWorkspace": + """Allocate and arm a workspace for ``mapping``'s TP group. Collective: every rank of the group calls it at + the same point, eagerly (not under CUDA-graph capture); it returns on every rank or raises on every rank. + ``fabric_handle``: share the memory by fabric handle (required across nodes) rather than POSIX file + descriptor; default ``mapping.is_multi_node()``.""" + from tensorrt_llm._torch.distributed.ops import ( + _get_mnnvl_workspace_comm, + _initialize_allreduce_mnnvl_protocol, + _make_mnnvl_mcast_buffer, + _mnnvl_workspace_all_succeeded, + ) + + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "MnnvlWorkspace.create is collective and allocates: call it before capture" + ) + if buffer_bytes <= 0 or buffer_bytes % 16: + raise ValueError(f"buffer_bytes must be a positive multiple of 16, got {buffer_bytes}") + use_fabric = mapping.is_multi_node() if fabric_handle is None else bool(fabric_handle) + comm = _get_mnnvl_workspace_comm(mapping) + error: Optional[Exception] = None + workspace = None + try: + total = NUM_LAMPORT_BUFFERS * buffer_bytes + handle = _make_mnnvl_mcast_buffer(comm, total, mapping, use_fabric) + lamport = handle.get_uc_buffer(mapping.tp_rank, (total // 4,), torch.float32, 0) + flags = torch.zeros(FLAG_WORDS, dtype=torch.uint32, device=lamport.device) + workspace = cls( + lamport=lamport, + buffer_flags=flags, + buffer_bytes=buffer_bytes, + rank=mapping.tp_rank, + world_size=mapping.tp_size, + handle=handle, + comm=comm, + ) + except Exception as exc: # noqa: BLE001 -- reported to every rank below, then re-raised + error = exc + if not _mnnvl_workspace_all_succeeded(comm, error is None): + raise RuntimeError("MnnvlWorkspace: allocation failed on at least one rank") from error + # Arms every buffer word and the flags; also the barrier after which a peer may push into this rank. + _initialize_allreduce_mnnvl_protocol( + dict( + uc_buffer=workspace.lamport, + buffer_flags=workspace.buffer_flags, + buffer_size_bytes=buffer_bytes, + comm=comm, + ) + ) + return workspace diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml index 7382f31d99a1..456e86438283 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml @@ -28,7 +28,7 @@ # limits, e.g. no float8 arithmetic, are noted in the entry docstring). # Entries wrapping trtllm ops carry all three. # -# ── RECEIPT STATUS: all 19 entries certified on sm_103 ──────────────────── +# ── RECEIPT STATUS: 19 entries certified on sm_103, 1 on sm_100 ─────────── # # A receipt says this entry's test passed on a stated GPU architecture. The # architecture is the whole key: it is a real axis -- see the @@ -66,8 +66,28 @@ # exactly as written. They are true records of what was observed on another # device, and rewriting them would manufacture GB300 evidence that does not # exist. Read them as provenance; the frontmatter is the certification. +# +# An entry first called by a GB200 (sm_100) target is certified there: its +# receipt key is sm_100 (comm/mnnvl_allreduce_attn_res). An entry already +# certified elsewhere gains its sm_100 key in the change that adds its first +# sm_100 caller. A collective's receipt records the world size its matrix ran +# at. # ────────────────────────────────────────────────────────────────────────── # +# STATEFUL ENTRIES. When an op's result depends on state that outlives the +# call (a workspace, Lamport buffers, counters, a record one launch writes and +# a later one reads), that state is a typed object the caller owns and passes +# as an argument: created once by an explicit constructor (collective where +# the op is, eager, before any CUDA-graph capture), never a module-level dict +# or an environment switch. The object's type sits beside the entries that +# take it (comm/mnnvl_workspace.py) and is not an entry: it launches nothing +# per call. The contract gains a `## State` section: the object's contents and +# size; who creates it and when; which ops may share one object; the call +# order every rank must keep across layers and steps; what a later launch +# reads; how it is re-armed. Its test drives call sequences (chained calls +# across steps, capture and replay, two objects interleaved) and a negative +# control that breaks the call order, not only single calls. +# # 30 of the source catalog's 43 entries are here — the union of what the two # migrated targets call. The 13 left behind (silu_and_mul, attn_custom_op_inplace, # create_attn_outputs, create_mla_outputs, load_chunked_kv_cache_for_mla, @@ -206,3 +226,9 @@ entries: - path: comm/reducescatter.py impl: torch.ops.trtllm.reducescatter summary: "Summing reduce-scatter over a group of MPI-session ranks into a fresh tensor: every rank's tensor summed elementwise, then split along dim 0 in ascending rank order so each rank keeps its own rows, even (sizes=None) or uneven (per-rank row counts in `sizes`, as attention data parallelism produces); the sum is accumulated in the input dtype in a destination-dependent ring order — deterministic and replay-reproducible, but not the correctly rounded fp32 sum — and float8_e4m3fn is summed as raw bytes rather than as floats; calls pair by position on the communicator whichever `sizes` form they take, and because this one computes, ranks that disagree on call order make every rank wrong nearly everywhere rather than a quarter of the way — silently, and bitwise reproducibly, whenever the mispaired calls move the same number of bytes, hanging only when they do not" + + # Stateful (a `## State` section and a caller-owned state object; the state type comm/mnnvl_workspace.py is not + # an entry: it launches nothing per call). + - path: comm/mnnvl_allreduce_attn_res.py + impl: torch.ops.trtllm.mnnvl_allreduce_attn_res + summary: "One-shot MNNVL all-reduce with Kimi K3's residual update as its epilogue (updated = prefix + sum over ranks, normed = RMSNorm of the attention-residual selection over the snapshot bank and updated), over a caller-owned MnnvlWorkspace: three Lamport buffers rotated by every call of every MNNVL op that shares the object, so every rank makes the same calls on it in the same order; a swapped pair of same-shaped calls is silently wrong on every rank, two ranks ordering calls on two workspaces differently on one stream deadlock" diff --git a/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml b/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml index 5102abe33232..57ed3a53b716 100644 --- a/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml +++ b/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml @@ -51,6 +51,7 @@ l0_gb200_multi_gpus: - unittest/_torch/multi_gpu/test_mnnvl_allreduce.py::test_mnnvl_checkpoint_preserves_cuda_graph_addresses - unittest/_torch/multi_gpu/test_mnnvl_allreduce.py::test_mnnvl_checkpoint_rejects_wrong_group_membership - unittest/_torch/multi_gpu/test_mnnvl_allreduce.py::test_mnnvl_workspace_growth_keeps_captured_graphs + - unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_allreduce_attn_res_op_matrix.py - unittest/_torch/moe/test_moe_comm.py::TestMoEComm::test_mnnvl_checkpoint_preserves_moe_graph_addresses - unittest/_torch/moe/test_moe_comm.py::TestMoEComm::test_mnnvl_engine_checkpoint_coordination - unittest/_torch/moe/test_moe_comm.py::TestMoEComm::test_mnnvl_checkpoint_failure_is_collective_and_bounded diff --git a/tests/unittest/_torch/modeling_v2/comm/_lockstep.py b/tests/unittest/_torch/modeling_v2/comm/_lockstep.py new file mode 100644 index 000000000000..012498ab36f5 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/comm/_lockstep.py @@ -0,0 +1,167 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Shared pieces of the stateful collective matrices (``_mnnvl_allreduce_attn_res_op_matrix.py``): the launcher +with the world size and launcher as parameters, the rank setup, the native-torch reference of Kimi K3's residual +update, and exact-arithmetic payloads. + +Started by file path from the launchers (this tree is not a package). Importing it pulls in torch only; the +rank body imports tensorrt_llm. + +Launchers (``--launcher``): + mpirun this process re-executes the entry under ``mpirun -n ``, one rank per device named in + CUDA_VISIBLE_DEVICES, under a deadline it enforces by killing the process group (CI on one tray); + srun this process is already one of ```` ranks started by an external launcher (e.g. + ``srun -N 4 --ntasks-per-node 4 --mpi=pmix python --launcher srun --world-size 16`` across trays); + the launcher owns the deadline. Rank r drives device r % (devices per node). +""" + +from __future__ import annotations + +import argparse +import faulthandler +import os +import signal +import subprocess +import sys +import time +from pathlib import Path +from typing import Callable, List, Optional, Sequence + +import torch + +WORKER_FLAG = "--rank-worker" +EMPTY_ROWS_EPS = 1e-6 +STACK_DUMP_S = 600 +"""A rank still running after this long prints every thread's Python stack (and again every period): a wedged +collective shows where it waits instead of only timing out.""" + + +def parse_args(argv: Sequence[str]) -> argparse.Namespace: + parser = argparse.ArgumentParser(add_help=False) + parser.add_argument("--world-size", type=int, default=None) + parser.add_argument("--launcher", choices=("mpirun", "srun"), default="mpirun") + parser.add_argument("--fabric-handle", choices=("auto", "on", "off"), default="auto") + parser.add_argument(WORKER_FLAG, action="store_true") + args, _ = parser.parse_known_args(argv) + return args + + +def spawn(entry_file: str, args: argparse.Namespace, deadline_s: int) -> None: + """Re-exec ``entry_file`` under mpirun with ``args.world_size`` ranks (default: one per visible device).""" + visible = [d for d in os.environ.get("CUDA_VISIBLE_DEVICES", "").split(",") if d.strip()] + assert visible, "set CUDA_VISIBLE_DEVICES to the devices this run owns, e.g. 0,1,2,3" + world = args.world_size or len(visible) + assert 2 <= world <= len(visible), ( + f"world size {world} needs 2..{len(visible)} visible devices (one rank per device)" + ) + command = ["mpirun", "-n", str(world), sys.executable, str(Path(entry_file).resolve()), WORKER_FLAG, + "--world-size", str(world), "--fabric-handle", args.fabric_handle] # fmt: skip + print(f"[launcher] {' '.join(command)}", flush=True) + process = subprocess.Popen(command, start_new_session=True) + try: + code = process.wait(timeout=deadline_s) + except subprocess.TimeoutExpired: + os.killpg(process.pid, signal.SIGKILL) + process.wait() + raise AssertionError( + f"the {world}-rank run did not finish in {deadline_s}s (wedged)" + ) from None + assert code == 0, f"the {world}-rank run exited {code}" + + +class Rank: + """This process's rank: the MPI world, its device, the TP Mapping over every rank of the job.""" + + def __init__(self, args: argparse.Namespace): + from mpi4py import MPI + + from tensorrt_llm.mapping import Mapping + + faulthandler.dump_traceback_later(STACK_DUMP_S, repeat=True) + self.MPI = MPI + self.comm = MPI.COMM_WORLD + self.rank, self.world = self.comm.Get_rank(), self.comm.Get_size() + if args.world_size is not None: + assert self.world == args.world_size, ( + f"the launcher started {self.world} ranks, --world-size says {args.world_size}" + ) + assert self.world >= 2, f"a collective needs at least 2 ranks, got {self.world}" + per_node = torch.cuda.device_count() + torch.cuda.set_device(self.rank % per_node) + self.mapping = Mapping( + world_size=self.world, rank=self.rank, gpus_per_node=per_node, tp_size=self.world + ) + self.fabric = {"auto": None, "on": True, "off": False}[args.fabric_handle] + + def barrier(self) -> None: + torch.cuda.synchronize() + self.comm.Barrier() + + def all_true(self, flag: bool) -> bool: + return all(self.comm.allgather(bool(flag))) + + def same_on_ranks(self, *tensors: torch.Tensor) -> bool: + """Bitwise equality of the tensors across ranks (their int16 words summed per row, compared exactly).""" + mine = [t.contiguous().view(torch.int16).long().sum(dim=-1).cpu() for t in tensors] + every = self.comm.allgather(mine) + return all(torch.equal(a, b) for other in every[1:] for a, b in zip(every[0], other)) + + def late(self, which: Optional[int], seconds: float = 0.005) -> None: + """Rank ``which`` (None: nobody) starts its next launch ``seconds`` late.""" + if which == self.rank: + time.sleep(seconds) + + +def run_checks(rank: Rank, checks: List[Callable[[], None]]) -> int: + for check in checks: + try: + check() + except BaseException: + import traceback + + print(f"[rank {rank.rank}] FAILED {check.__name__}", flush=True) + traceback.print_exc() + sys.stdout.flush() + sys.stderr.flush() + # A rank that leaves a collective early wedges every other rank in it. + rank.comm.Abort(1) + if rank.rank == 0: + print(f"[rank 0] passed {check.__name__}", flush=True) + rank.barrier() + print(f"[rank {rank.rank}] {len(checks)} checks passed", flush=True) + return 0 + + +def exact_bf16(gen: torch.Generator, shape, lo: int, hi: int, scale: float) -> torch.Tensor: + """bf16 integers in [lo, hi) times ``scale`` (a power of two): sums of a few of them are exact in fp32 and + bf16, so a reference sum does not depend on the summation order.""" + ints = torch.randint(lo, hi, tuple(shape), generator=gen, device="cuda") + return (ints.float() * scale).bfloat16() + + +def residual_update_ref( + updated: torch.Tensor, + block_residual: torch.Tensor, + res_weight: torch.Tensor, + rms_weight: torch.Tensor, + rms_eps: float, + output_rms_weight: torch.Tensor, + output_rms_eps: float, +) -> torch.Tensor: + """Kimi K3's residual update after the all-reduce, in fp32 (HF ``_apply_attn_res`` + RMSNorm): over the + candidates ``v = [block_residual..., updated]`` score ``sum(rmsnorm(v) * rms_weight * res_weight)``, softmax + over the candidates, mix, RMSNorm with ``output_rms_weight``; bf16 out.""" + v = torch.cat([block_residual, updated.unsqueeze(0)], dim=0).float() + rs = (v.square().mean(dim=-1) + rms_eps).rsqrt() + logits = (v * rs[..., None] * (rms_weight.float() * res_weight.float())).sum(dim=-1) + probs = torch.softmax(logits, dim=0) + mixed = (probs[..., None] * v).sum(dim=0).bfloat16().float() + out = mixed * (mixed.square().mean(dim=-1, keepdim=True) + output_rms_eps).rsqrt() + return (out * output_rms_weight.float()).bfloat16() + + +def rel_err(got: torch.Tensor, want: torch.Tensor) -> float: + return ( + (got.float() - want.float()).abs().max() + / want.float().abs().max().clamp_min(EMPTY_ROWS_EPS) + ).item() diff --git a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allreduce_attn_res_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allreduce_attn_res_op_matrix.py new file mode 100644 index 000000000000..ee5d18600687 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allreduce_attn_res_op_matrix.py @@ -0,0 +1,259 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU certification matrix for the ``comm/mnnvl_allreduce_attn_res`` catalog entry and its ``MnnvlWorkspace``. + +The op's correctness depends on state that outlives a call (the workspace's Lamport rotation), so beyond single +calls this drives call *sequences*: layers x steps with the token count dipping and growing back and a random rank +late, two workspaces interleaved, CUDA-graph capture and replay mixed with eager calls, and a negative control in +which one rank swaps two calls and every rank gets a wrong answer without an error -- the failure the sequence tests +exist to catch. + + CUDA_VISIBLE_DEVICES=0,1,2,3 python _mnnvl_allreduce_attn_res_op_matrix.py [--world-size 4] + srun -n 16 --mpi=pmix python _mnnvl_allreduce_attn_res_op_matrix.py --launcher srun --world-size 16 + +Not a pytest module: one fixed sequence of checks inside one W-rank job (they share the workspaces and their +rotation). The collected entry point is ``test_modeling_v2_mnnvl_allreduce_attn_res_op_matrix.py``. + +Every rank draws every rank's inputs from one seed, so each rank holds the whole reference. The inputs of the sum +are small multiples of 1/16, so ``updated`` is exact in fp32 and bf16 whatever the summation order and is compared +bit for bit; ``normed`` (softmax, rsqrt) against the fp32 reference within ``TOL``. Every output is also compared +bitwise across the ranks. +""" + +import random +import sys +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parent)) +import _lockstep as ls # noqa: E402 + +assert torch.cuda.is_available(), "mnnvl_allreduce_attn_res requires CUDA devices" + +DEADLINE_S = 900 +H = 7168 +BUFFER_BYTES = 4 << 20 # the Lamport buffer size Kimi K3 decodes with +RMS_EPS = 1e-6 +OUT_EPS = 1e-6 +TOL = 2e-2 # normed: max |err| / max |ref| +TOKENS = (1, 2, 3, 4, 5, 6, 7, 8, 16) +SNAPSHOTS = (0, 1, 5, 8, 11) # candidates 1..12 +LAYERS = 12 +DIP_STEPS = (8, 8, 8, 2, 7, 8, 1, 1, 8, 16, 3, 8) + +R = None +entry = None +WS_A = None +WS_B = None +STATS = {"normed_err": 0.0} + + +class Call: + """One call's arguments on every rank and its reference. ``prefix``: a tensor (chained from the previous call's + ``updated``), True (drawn), or None.""" + + def __init__(self, seed, tokens, snapshots, prefix=True): + g = torch.Generator(device="cuda").manual_seed(seed) + self.inputs = [ls.exact_bf16(g, (tokens, H), -4, 5, 1 / 16) for _ in range(R.world)] + if prefix is True: + prefix = ls.exact_bf16(g, (tokens, H), -32, 33, 1 / 16) + self.prefix = prefix + self.block = torch.randn(snapshots, tokens, H, generator=g, device="cuda").bfloat16() + self.res_w = (torch.randn(H, generator=g, device="cuda") * 0.05).bfloat16() + self.rms_w = (1.0 + 0.1 * torch.randn(H, generator=g, device="cuda")).bfloat16() + self.out_w = (1.0 + 0.1 * torch.randn(H, generator=g, device="cuda")).bfloat16() + + def ref(self): + total = sum(x.float() for x in self.inputs) + if self.prefix is not None: + total = total + self.prefix.float() + updated = total.bfloat16() + normed = ls.residual_update_ref( + updated, self.block, self.res_w, self.rms_w, RMS_EPS, self.out_w, OUT_EPS + ) + return normed, updated + + def run(self, ws): + return entry(self.inputs[R.rank], self.prefix, self.block, self.res_w, self.rms_w, self.out_w, RMS_EPS, + OUT_EPS, ws) # fmt: skip + + +def verify(call: Call, got, where: str) -> None: + normed, updated = got + want_normed, want_updated = call.ref() + assert torch.equal(updated, want_updated), f"{where}: updated differs from the exact sum" + err = ls.rel_err(normed, want_normed) + STATS["normed_err"] = max(STATS["normed_err"], err) + assert err <= TOL, f"{where}: normed rel err {err:.3e} > {TOL}" + assert R.same_on_ranks(normed, updated), f"{where}: ranks disagree" + + +def check_workspace_is_armed_and_sized() -> None: + assert WS_A.world_size == R.world and WS_A.rank == R.rank + assert WS_A.comm_buffer(torch.bfloat16).shape == (3, BUFFER_BYTES // 2) + assert WS_A.max_one_shot_tokens(H) == BUFFER_BYTES // (H * R.world * 2) + armed = WS_A.lamport.view(torch.int32) + assert bool((armed == torch.tensor(-(2**31), dtype=torch.int32, device="cuda")).all()), ( + "every word -0.0" + ) + assert WS_A.buffer_flags.view(torch.int32).tolist()[:3] == [0, 2, BUFFER_BYTES] + + +def check_single_calls() -> None: + for t in TOKENS: + for s in SNAPSHOTS: + for with_prefix in (True, False): + call = Call(1000 + 37 * t + s, t, s, prefix=with_prefix or None) + verify(call, call.run(WS_A), f"T {t} snapshots {s} prefix {with_prefix}") + + +def check_a_call_over_one_buffer_raises_on_every_rank() -> None: + t = WS_A.max_one_shot_tokens(H) + 1 + call = Call(2000, t, 1) + try: + call.run(WS_A) + raised = False + except RuntimeError as exc: + raised = "exceeds one Lamport buffer" in str(exc) + assert R.all_true(raised), f"T {t} over one Lamport buffer did not raise on every rank" + # The rotation did not move: the next call is still correct. + call = Call(2001, 8, 2) + verify(call, call.run(WS_A), "after the rejected call") + + +def run_step(ws, seed, tokens, layers=LAYERS, late_rng=None): + """One decode step: ``layers`` calls, each layer's prefix the previous layer's ``updated``; verified.""" + prefix = ls.exact_bf16( + torch.Generator(device="cuda").manual_seed(seed), (tokens, H), -32, 33, 1 / 16 + ) + for layer in range(layers): + call = Call(seed + 1 + layer, tokens, (layer * 5) % 12, prefix=prefix) + R.barrier() + R.late(late_rng.randrange(R.world) if late_rng is not None else None) + got = call.run(ws) + verify(call, got, f"step seed {seed} T {tokens} layer {layer}") + prefix = got[1] + + +def check_dip_and_regrow_sequence() -> None: + """A call after a smaller one must not read what an older, larger call left in the buffer (the failure of a + re-arm sized by the current call). Steps of 12 layers at T 8, 8, 8, 2, 7, 8, 1, 1, 8, 16, 3, 8, a random rank + late at every call.""" + late = random.Random(7) + for i, t in enumerate(DIP_STEPS): + run_step(WS_A, 3000 + 100 * i, t, late_rng=late) + + +def check_two_workspaces_interleaved() -> None: + """Two workspaces are two rotations: calls alternate between them in an irregular pattern (A A B A B B ...), so + the two objects' positions differ, and every call is correct. The pattern is the same on every rank: calls on one + stream are serialized and each waits for its peers, so two ranks issuing calls on two workspaces in different + orders deadlock (measured; not exercised here).""" + pattern = "AABABBAAAB" * 2 + for i, which in enumerate(pattern): + ws = WS_A if which == "A" else WS_B + call = Call(4000 + i, (3, 8, 1, 8, 5)[i % 5], i % 12) + verify(call, call.run(ws), f"interleaved {which} {i}") + + +def check_graph_capture_and_replay() -> None: + """A captured step of 12 chained calls (T 8) replayed with rewritten inputs, eager calls of other shapes on the + same workspace between replays: replays and eager calls share one rotation, in the same order on every rank.""" + t = 8 + calls = [Call(5000 + layer, t, (layer * 5) % 12, prefix=True) for layer in range(LAYERS)] + bufs = [(c.inputs[R.rank].clone(), c.block.clone()) for c in calls] + prefix0 = calls[0].prefix.clone() + + def step(): + outs, prefix = [], prefix0 + for c, (x, blk) in zip(calls, bufs): + outs.append(entry(x, prefix, blk, c.res_w, c.rms_w, c.out_w, RMS_EPS, OUT_EPS, WS_B)) + prefix = outs[-1][1] + return outs + + step() # the first call of every shape eagerly + R.barrier() + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + outs = step() + R.barrier() + for rep in range(8): + fresh = [ + Call(6000 + 100 * rep + layer, t, (layer * 5) % 12, prefix=True) + for layer in range(LAYERS) + ] + prefix0.copy_(fresh[0].prefix) + for (x, blk), f, c in zip(bufs, fresh, calls): + x.copy_(f.inputs[R.rank]) + blk.copy_(f.block) + c.inputs, c.block = f.inputs, f.block + R.barrier() + graph.replay() + prefix = prefix0 + for layer, (c, got) in enumerate(zip(calls, outs)): + c.prefix = prefix + verify(c, got, f"replay {rep} layer {layer}") + prefix = got[1] + eager = Call(7000 + rep, (3, 1, 16, 5)[rep % 4], rep % 12) + verify(eager, eager.run(WS_B), f"eager after replay {rep}") + del graph + + +def check_wrong_call_order_is_detected() -> None: + """Negative control: rank 0 swaps two same-shaped calls on one workspace. Every call still returns and nothing + raises or hangs (the rotation positions agree), but every rank's two results are wrong: a call pairs with the + peers' call at the same position. Then a plain call is correct again: a swapped pair realigns the positions.""" + c1, c2 = Call(8000, 8, 3), Call(8001, 8, 3) + R.barrier() + if R.rank == 0: + got2, got1 = c2.run(WS_A), c1.run(WS_A) + else: + got1, got2 = c1.run(WS_A), c2.run(WS_A) + torch.cuda.synchronize() + wrong = [(got[1] != c.ref()[1]).float().mean().item() for c, got in ((c1, got1), (c2, got2))] + assert R.all_true(min(wrong) > 0.5), f"the swap went unnoticed: wrong fractions {wrong}" + c3 = Call(8002, 8, 3) + verify(c3, c3.run(WS_A), "after the swapped pair") + + +CHECKS = [ + check_workspace_is_armed_and_sized, + check_single_calls, + check_a_call_over_one_buffer_raises_on_every_rank, + check_dip_and_regrow_sequence, + check_two_workspaces_interleaved, + check_graph_capture_and_replay, + # Stays last: it deliberately disagrees on call order. + check_wrong_call_order_is_detected, +] + + +def _run_one_rank(args) -> int: + global R, entry, WS_A, WS_B + R = ls.Rank(args) + from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm import ( + mnnvl_allreduce_attn_res as module, + ) + from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.mnnvl_workspace import ( + MnnvlWorkspace, + ) + + entry = module.mnnvl_allreduce_attn_res + with torch.inference_mode(): + WS_A = MnnvlWorkspace.create(R.mapping, BUFFER_BYTES, fabric_handle=R.fabric) + WS_B = MnnvlWorkspace.create(R.mapping, BUFFER_BYTES, fabric_handle=R.fabric) + code = ls.run_checks(R, CHECKS) + if R.rank == 0: + print(f"[rank 0] world {R.world}; max normed rel err {STATS['normed_err']:.3e}", flush=True) + return code + + +if __name__ == "__main__": + ARGS = ls.parse_args(sys.argv[1:]) + if ARGS.rank_worker or ARGS.launcher == "srun": + sys.exit(_run_one_rank(ARGS)) + ls.spawn(__file__, ARGS, DEADLINE_S) + print("OK") diff --git a/tests/unittest/_torch/modeling_v2/comm/_rank_job.py b/tests/unittest/_torch/modeling_v2/comm/_rank_job.py index a24b842d6f45..197f237d2c2f 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_rank_job.py +++ b/tests/unittest/_torch/modeling_v2/comm/_rank_job.py @@ -32,13 +32,13 @@ from pathlib import Path WORLD_SIZE = 4 -"""dep4's world size -- the topology these entries are certified for.""" +"""dep4's world size -- the default topology these entries are certified for (one node).""" _LAUNCHER_GRACE_S = 300 """Headroom over the entry's own deadline, so its message wins the race.""" -def _devices() -> str: +def _devices(world_size: int) -> str: visible = os.environ.get("CUDA_VISIBLE_DEVICES") if visible: devices = [d for d in visible.split(",") if d.strip()] @@ -46,11 +46,10 @@ def _devices() -> str: import torch devices = [str(i) for i in range(torch.cuda.device_count())] - assert len(devices) >= WORLD_SIZE, ( - f"this entry is certified at world size {WORLD_SIZE}; only " - f"{len(devices)} device(s) are visible" + assert len(devices) >= world_size, ( + f"this run is at world size {world_size}; only {len(devices)} device(s) are visible" ) - return ",".join(devices[:WORLD_SIZE]) + return ",".join(devices[:world_size]) def _load_launcher_constants(launcher: Path): @@ -67,10 +66,15 @@ def _load_launcher_constants(launcher: Path): return module -def run(entry: str) -> None: - """Run ``__op_matrix``'s launcher over ``WORLD_SIZE`` devices.""" +def run(entry: str, world_size: int = WORLD_SIZE) -> None: + """Run ``__op_matrix``'s launcher over ``world_size`` devices of this node (its local ``mpirun``). + + A launcher that takes ``--world-size`` (the stateful entries' ``_lockstep`` launchers) also runs as one of N + ranks started by an external launcher, ``--launcher srun --world-size N``: that is how a multi-node receipt is + recorded, outside pytest; the others ignore the argument and size themselves from CUDA_VISIBLE_DEVICES. + """ launcher = Path(__file__).resolve().with_name(f"_{entry}_op_matrix.py") - env = dict(os.environ, CUDA_VISIBLE_DEVICES=_devices()) + env = dict(os.environ, CUDA_VISIBLE_DEVICES=_devices(world_size)) # The ranks import tensorrt_llm absolutely, and a source checkout is not # necessarily installed. tests/unittest/_torch/modeling_v2/comm -> repo root. @@ -89,7 +93,7 @@ def run(entry: str) -> None: timeout = constants.DEADLINE_S + getattr(constants, "WEDGE_CAP_S", 0) + _LAUNCHER_GRACE_S completed = subprocess.run( - [sys.executable, str(launcher)], + [sys.executable, str(launcher), "--world-size", str(world_size)], env=env, capture_output=True, text=True, diff --git a/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_allreduce_attn_res_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_allreduce_attn_res_op_matrix.py new file mode 100644 index 000000000000..c93308cda874 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_allreduce_attn_res_op_matrix.py @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Collected entry point for the mnnvl_allreduce_attn_res op's certification matrix. + +The matrix is ``_mnnvl_allreduce_attn_res_op_matrix.py`` beside this file, its own W-rank launcher (call sequences +over one caller-owned workspace, so one job, not independent cases); see ``_rank_job`` for why that is left intact. +""" + +import _rank_job +import pytest +import torch + +assert torch.cuda.is_available(), "mnnvl_allreduce_attn_res requires CUDA devices" + + +# Each case starts its own W-rank mpirun over the visible devices; under xdist several workers would fight for them. +@pytest.mark.no_xdist +def test_mnnvl_allreduce_attn_res_op_matrix() -> None: + _rank_job.run("mnnvl_allreduce_attn_res") From b97811aaf78060c278a077fff9a14fe7073bb775 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 18:31:13 -0700 Subject: [PATCH 042/161] [None][feat] MNNVLAllReduce.allreduce_attn_res_rmsnorm, and the op's 4-rank test - MNNVLAllReduce.allreduce_attn_res_rmsnorm calls trtllm::mnnvl_allreduce_attn_res on the module's MNNVL workspace. It grows the workspace to the call's one-shot footprint when needed, so the first call of a shape runs outside capture. - test_k3_mnnvl_comm.py (4 ranks, l0_gb200_multi_gpus.yml) compares the op, at M 1-8, 16, 32 and 64, with 0, 1, 3, 8 and 11 snapshots, with and without the prefix, against: - the unfused path it replaces (the MNNVL all-reduce, then attn_res_add_rmsnorm_fwd): the updated prefix sum bit for bit, normed within 1e-2; - an fp32 reference: within 2e-2. - It also checks: - run-to-run identical bits and the same bits on every rank; - every M's rows equal to the same rows of the 64-row call; - one rank's changed input changes every rank's result. Signed-off-by: Vasanth Sabavat --- tensorrt_llm/_torch/distributed/ops.py | 41 +++ .../test-db/l0_gb200_multi_gpus.yml | 1 + .../kimi_k3/test_k3_mnnvl_comm.py | 240 ++++++++++++++++++ 3 files changed, 282 insertions(+) create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mnnvl_comm.py diff --git a/tensorrt_llm/_torch/distributed/ops.py b/tensorrt_llm/_torch/distributed/ops.py index 31d06468876d..4a4ec2fc5a81 100644 --- a/tensorrt_llm/_torch/distributed/ops.py +++ b/tensorrt_llm/_torch/distributed/ops.py @@ -983,6 +983,47 @@ def forward( ) return tuple(outputs) if is_fusion else outputs[0] + def allreduce_attn_res_rmsnorm( + self, + input: torch.Tensor, + prefix_sum: Optional[torch.Tensor], + block_residual: torch.Tensor, + res_weight: torch.Tensor, + rms_weight: torch.Tensor, + output_rms_weight: torch.Tensor, + rms_eps: float, + output_rms_eps: float, + ) -> Tuple[torch.Tensor, torch.Tensor]: + """One-shot all-reduce of ``input`` with Kimi K3's residual update as its epilogue. + + Computes ``updated = prefix_sum + allreduce(input)`` (``allreduce(input)`` when + ``prefix_sum`` is None), then the attention-residual selection over + ``block_residual`` ``[num_snapshots, num_tokens, hidden]`` and ``updated``, then an + RMSNorm with ``output_rms_weight``. Returns ``(normed, updated)``, rounded like the + unfused all-reduce followed by ``trtllm::attn_res_add_rmsnorm_fwd``. + + The workspace is grown to the one-shot footprint when needed, so the first call for a + shape must happen outside CUDA graph capture (warmup does this). + """ + num_tokens, hidden_dim = input.shape + one_shot_bytes = (num_tokens * hidden_dim * self.mapping.tp_size * + input.element_size()) + workspace = get_or_scale_allreduce_mnnvl_workspace( + self.mapping, self.dtype, buffer_size_bytes=one_shot_bytes) + normed, updated = torch.ops.trtllm.mnnvl_allreduce_attn_res( + input, + prefix_sum, + block_residual, + res_weight, + rms_weight, + output_rms_weight, + rms_eps, + output_rms_eps, + workspace["uc_buffer"].view(self.dtype).view(3, -1), + workspace["buffer_flags"], + ) + return normed, updated + class AllReduce(nn.Module): diff --git a/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml b/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml index 57ed3a53b716..c53cb87f897d 100644 --- a/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml +++ b/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml @@ -52,6 +52,7 @@ l0_gb200_multi_gpus: - unittest/_torch/multi_gpu/test_mnnvl_allreduce.py::test_mnnvl_checkpoint_rejects_wrong_group_membership - unittest/_torch/multi_gpu/test_mnnvl_allreduce.py::test_mnnvl_workspace_growth_keeps_captured_graphs - unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_allreduce_attn_res_op_matrix.py + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mnnvl_comm.py - unittest/_torch/moe/test_moe_comm.py::TestMoEComm::test_mnnvl_checkpoint_preserves_moe_graph_addresses - unittest/_torch/moe/test_moe_comm.py::TestMoEComm::test_mnnvl_engine_checkpoint_coordination - unittest/_torch/moe/test_moe_comm.py::TestMoEComm::test_mnnvl_checkpoint_failure_is_collective_and_bounded diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mnnvl_comm.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mnnvl_comm.py new file mode 100644 index 000000000000..0004cff4fe64 --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mnnvl_comm.py @@ -0,0 +1,240 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""The Kimi K3 decode collectives on the MNNVL all-reduce workspace, one process per GPU over the TP group of this +run (4 on one GB200 tray), at every M in 1..8 and at 16, 32, 64 tokens: + trtllm::mnnvl_allreduce_attn_res (MNNVLAllReduce.allreduce_attn_res_rmsnorm): against the unfused path it replaces, + the MNNVL all-reduce then trtllm::attn_res_add_rmsnorm_fwd (attn_res_rmsnorm_fwd without a prefix sum): the + updated prefix sum bit for bit, the normed rows within 1e-2 (max |d| / max |ref|), both against an fp32 port of the + attention-residual selection within 2e-2; 0, 1, 3, 8 and 11 snapshots, with and without the prefix; +each with run-to-run identical bits, the same bits on every rank, each M's rows bit-identical to the same rows of the +64-row call, and one rank's perturbed input changing every rank's result. + +Run under pytest (a pool of 4 MPI workers) or directly, one process per GPU: + srun -N1 -n4 --mpi=pmix python3 test_k3_mnnvl_comm.py [attn_res] +""" + +import hashlib +import os +import pickle +import sys +import traceback +from types import SimpleNamespace + +import pytest +import torch + +try: + import cloudpickle + from mpi4py import MPI +except ImportError: # the test is skipped below + cloudpickle = MPI = None + +if cloudpickle is not None: + cloudpickle.register_pickle_by_value(sys.modules[__name__]) + MPI.pickle.__init__(cloudpickle.dumps, cloudpickle.loads, pickle.HIGHEST_PROTOCOL) + +WORLD = 4 +H = 7168 +EPS, OUT_EPS = 1e-6, 1e-5 +M_CASES = list(range(1, 9)) + [16, 32, 64] +M_MAX = 64 +SNAPSHOTS = (0, 1, 3, 8, 11) + + +def _supported() -> bool: + if MPI is None or not torch.cuda.is_available() or torch.cuda.device_count() < WORLD: + return False + return torch.cuda.get_device_capability()[0] == 10 + + +pytestmark = [ + pytest.mark.threadleak(enabled=False), + pytest.mark.skipif(not _supported(), reason=f"needs {WORLD} SM100 GPUs with MNNVL and mpi4py"), +] + + +def _bits(t: torch.Tensor) -> torch.Tensor: + t = t.contiguous() + return t.view(torch.int16) if t.element_size() == 2 else t.view(torch.int32) + + +def _same(a: torch.Tensor, b: torch.Tensor) -> bool: + return a.shape == b.shape and a.dtype == b.dtype and torch.equal(_bits(a), _bits(b)) + + +def _rel(a, b) -> float: + """max |a - b| / max |b|: normalized like the GEMV checks (a 2-ulp rounding difference at the tail of a large M + is within it; a wrong row is not).""" + return ((a.float() - b.float()).abs().max() / b.float().abs().max().clamp_min(1e-30)).item() + + +def _context(): + """One process per GPU, ranks filling the nodes in order (gpus_per_node = the node's GPU count, so a rank's + local_rank is its device on every node); the multicast buffers use fabric handles within a tray as across trays.""" + os.environ.setdefault("TRTLLM_FORCE_MNNVL_AR", "1") + comm = MPI.COMM_WORLD + rank, world = comm.Get_rank(), comm.Get_size() + gpus = torch.cuda.device_count() + torch.cuda.set_device(rank % gpus) + import tensorrt_llm # noqa: F401 + from tensorrt_llm._torch.distributed.ops import MNNVLAllReduce + from tensorrt_llm.mapping import Mapping + + mapping = Mapping(world_size=world, rank=rank, gpus_per_node=gpus, tp_size=world) + return SimpleNamespace( + comm=comm, rank=rank, world=world, mnnvl=MNNVLAllReduce(mapping, torch.bfloat16) + ) + + +def _allreduce(ctx, x): + """The plain MNNVL all-reduce. Up to 8 ranks its one-shot and two-shot kernels both sum the ranks in rank order in + fp32, the fused op's order, so the reference is exact whichever one the size picks.""" + from tensorrt_llm._torch.distributed import AllReduceParams + + assert ctx.world <= 8, "above 8 ranks the two-shot kernel sums the ranks in another order" + return ctx.mnnvl(x, AllReduceParams()) + + +def _all_ranks(ctx, good) -> bool: + return all(ctx.comm.allgather(bool(good))) + + +def _digest(t: torch.Tensor) -> str: + return hashlib.sha256(_bits(t).cpu().numpy().tobytes()).hexdigest() + + +def _attn_res_inputs(ctx, snapshots, with_prefix): + shared = torch.Generator(device="cuda").manual_seed(1000 + 17 * snapshots + int(with_prefix)) + prefix = ( + torch.randn(M_MAX, H, generator=shared, device="cuda").bfloat16() if with_prefix else None + ) + scale = (1 + torch.arange(snapshots, device="cuda")).view(-1, 1, 1) + block = (torch.randn(snapshots, M_MAX, H, generator=shared, device="cuda") * scale).bfloat16() + res_w = (torch.randn(H, generator=shared, device="cuda") * 0.05).bfloat16() + rms_w = (1 + 0.1 * torch.randn(H, generator=shared, device="cuda")).bfloat16() + out_w = (1 + 0.1 * torch.randn(H, generator=shared, device="cuda")).bfloat16() + own = torch.Generator(device="cuda").manual_seed(7 * (1000 + snapshots) + ctx.rank + 1) + partial = (torch.randn(M_MAX, H, generator=own, device="cuda") * 0.5).bfloat16() + return partial, prefix, block, res_w, rms_w, out_w + + +def _fp32_reference(updated, block, res_w, rms_w, out_w): + """HF _apply_attn_res + RMSNorm in fp32 over [snapshots..., updated].""" + v = torch.cat([block.float(), updated.float().unsqueeze(0)], 0) + k = v * torch.rsqrt(v.pow(2).mean(-1, keepdim=True) + EPS) + probs = (k * (rms_w.float() * res_w.float())).sum(-1).softmax(0) + mixed = (probs.unsqueeze(-1) * v).sum(0).bfloat16().float() + normalized = (mixed * torch.rsqrt(mixed.pow(2).mean(-1, keepdim=True) + OUT_EPS)).bfloat16() + return (normalized.float() * out_w.float()).bfloat16() + + +def _unfused(ctx, partial, prefix, block, res_w, rms_w, out_w): + """MNNVL all-reduce, then attn_res_add_rmsnorm_fwd (attn_res_rmsnorm_fwd at a block start); (normed, updated). + With no snapshot (one candidate) the selection is the identity: normed is None (compared with fp32 only).""" + m, s = partial.shape[0], block.shape[0] + reduced = _allreduce(ctx, partial) + if s == 0: + return None, (reduced if prefix is None else (prefix.float() + reduced.float()).bfloat16()) + if prefix is None: + out = torch.ops.trtllm.attn_res_rmsnorm_fwd(reduced.reshape(m, 1, H), block.reshape(s, m, 1, H), res_w, rms_w, + out_w, EPS, OUT_EPS) # fmt: skip + return out.reshape(m, H), reduced + updated, out = torch.ops.trtllm.attn_res_add_rmsnorm_fwd(prefix.reshape(m, 1, H), reduced.reshape(m, 1, H), + block.reshape(s, m, 1, H), res_w, rms_w, out_w, EPS, + OUT_EPS) # fmt: skip + return out.reshape(m, H), updated.reshape(m, H) + + +def check_attn_res(ctx): + results = [] + for snapshots in SNAPSHOTS: + for with_prefix in (True, False): + partial64, prefix64, block64, res_w, rms_w, out_w = _attn_res_inputs( + ctx, snapshots, with_prefix + ) + + def fused(rows, part=None): + pre = prefix64[:rows].contiguous() if with_prefix else None + return ctx.mnnvl.allreduce_attn_res_rmsnorm( + (part if part is not None else partial64[:rows]).contiguous(), pre, + block64[:, :rows].contiguous(), res_w, rms_w, out_w, EPS, OUT_EPS) # fmt: skip + + n64, u64 = fused(M_MAX) + for m in M_CASES: + partial = partial64[:m].contiguous() + prefix = prefix64[:m].contiguous() if with_prefix else None + block = block64[:, :m].contiguous() + n, u = fused(m) + ref_n, ref_u = _unfused(ctx, partial, prefix, block, res_w, rms_w, out_w) + fp32 = _fp32_reference(ref_u, block, res_w, rms_w, out_w) + again = [fused(m) for _ in range(2)] + bad = partial.clone() + if ctx.rank == ctx.world - 1: + bad[0, 0] += 1.0 + bad_n, bad_u = fused(m, bad) + row = dict( + op="mnnvl_allreduce_attn_res", case=f"S{snapshots}_{'prefix' if with_prefix else 'noprefix'}", + M=m, updated_exact=_same(u, ref_u), + normed_vs_unfused=_rel(n, ref_n) if ref_n is not None else 0.0, normed_vs_fp32=_rel(n, fp32), + det=all(_same(a, n) and _same(b, u) for a, b in again), + rows_as_m64=_same(n, n64[:m]) and _same(u, u64[:m]), control=not _same(bad_u, u), + ranks_agree=len(set(ctx.comm.allgather(_digest(n)))) == 1, + ) # fmt: skip + good = (row["updated_exact"] and row["normed_vs_unfused"] <= 1e-2 and row["normed_vs_fp32"] <= 2e-2 + and row["det"] and row["rows_as_m64"] and row["control"] and row["ranks_agree"]) # fmt: skip + row["ok"] = _all_ranks(ctx, good) + results.append(row) + return results + + +CHECKS = {"attn_res": check_attn_res} + + +def _run_checks(names): + try: + ctx = _context() + with torch.inference_mode(): + return [row for name in names for row in CHECKS[name](ctx)] + except Exception: + traceback.print_exc() + raise + + +def _report(rows): + for row in rows: + fields = " ".join(f"{k}={(f'{v:.3e}' if isinstance(v, float) else v)}" for k, v in row.items() + if k not in ("op", "case", "M")) # fmt: skip + print(f"OPCHECK op={row['op']} case={row['case']} M={row['M']} {fields}", flush=True) + + +@pytest.mark.parametrize("mpi_pool_executor", [WORLD], indirect=True) +@pytest.mark.parametrize("check", list(CHECKS)) +def test_k3_mnnvl_comm(mpi_pool_executor, check): + per_rank = list(mpi_pool_executor.map(_run_checks, [[check]] * WORLD)) + _report(per_rank[0]) + assert all(row["ok"] for rows in per_rank for row in rows) + + +def main() -> int: + names = sys.argv[1:] or list(CHECKS) + rows = _run_checks(names) + if MPI.COMM_WORLD.Get_rank() == 0: + _report(rows) + print("PASS" if all(r["ok"] for r in rows) else "FAIL", flush=True) + return 0 if all(r["ok"] for r in rows) else 1 + + +if __name__ == "__main__": + sys.exit(main()) From c922c7a29205a39b15c423ade7e759a2202b9772 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 18:31:35 -0700 Subject: [PATCH 043/161] [None][doc] comm/mnnvl_allreduce_attn_res: every snapshot count 0-11 is certified The matrix's chained steps use (layer * 5) % 12 snapshots, which covers 4 as well. Signed-off-by: Vasanth Sabavat --- .../modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.md index f6cc1af2f337..8275156a4f7f 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.md @@ -55,7 +55,7 @@ def mnnvl_allreduce_attn_res( |---|---|---|---|---| | `input` | `[T, H]`; `H` = 7168 certified (the op takes a multiple of 1024 up to 8192); `T` 1-8 and 16 certified, at most `workspace.max_one_shot_tokens(H)` | bf16 | contiguous | CUDA, this rank's device | | `prefix_sum` | `None`, or `[T, H]` | bf16 | contiguous | CUDA | -| `block_residual` | `[S, T, H]`, `S` = 0..11 (certified at 0, 1, 2, 3, 5, 6, 7, 8, 9, 10, 11) | bf16 | contiguous | CUDA | +| `block_residual` | `[S, T, H]`, `S` = 0..11 (all certified) | bf16 | contiguous | CUDA | | `res_weight`, `rms_weight`, `output_rms_weight` | `[H]` | bf16 | contiguous | CUDA | | `rms_eps`, `output_rms_eps` | scalar | Python float | — | — | | `workspace` | an `MnnvlWorkspace` of this rank's TP group (see *State*) | — | — | — | From bf1f60ae0838fdb22572574ea58eab0199d169e4 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 18:36:41 -0700 Subject: [PATCH 044/161] [None][test] mnnvl_allreduce_attn_res matrix: run on sm_100 only The entry is certified on sm_100 (GB200). l0_gb300_multi_gpus.yml collects the whole modeling_v2/comm directory, so the collected entry point skips on other architectures. Signed-off-by: Vasanth Sabavat --- .../test_modeling_v2_mnnvl_allreduce_attn_res_op_matrix.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_allreduce_attn_res_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_allreduce_attn_res_op_matrix.py index c93308cda874..077390400645 100644 --- a/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_allreduce_attn_res_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_allreduce_attn_res_op_matrix.py @@ -12,6 +12,10 @@ assert torch.cuda.is_available(), "mnnvl_allreduce_attn_res requires CUDA devices" +if torch.cuda.get_device_capability() != (10, 0): + # The entry is certified on sm_100 (GB200) only; see its contract's receipts. + pytest.skip("mnnvl_allreduce_attn_res is certified on sm_100 only", allow_module_level=True) + # Each case starts its own W-rank mpirun over the visible devices; under xdist several workers would fight for them. @pytest.mark.no_xdist From 67f82d42c664b207cf95f2f7dea30865b6866e6a Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 20:05:49 -0700 Subject: [PATCH 045/161] [None][feat] modeling_v2 Kimi K3 target: decode GEMV sites for the fused projections K3DecodeGemvs.project covers the decode kernels' own projections, as the K3 stack runs them, besides the built-in MLA path's three: - mla_ag, MLA's [W_a; W_g] (2880 x 7168) with the gate rows through a sigmoid: gemm/k3_ctm_gemv_long (split 6, ring 6, pushed partials) up to 8 rows, gemm/k3_ctm_gemv_wide up to 64; - kda_proj, KDA's [q | k | v | g | f_a | b | pad] (3208 x 7168): the long kernel (split 5, ring 6) up to 8 rows where it fits one wave, the wide one up to 64; - o_proj, the attention output projection (7168 x 768): gemm/k3_decode_gemv up to 8 rows, the wide kernel up to 64. create() no longer needs the weights: it runs each site's kernels once on zero rows of a zero weight of the site's shape, each wide row class (16, 32, 64 tokens) once, so they compile at load and the shell's state build does not depend on when the attention modules fuse their weights. Under capture a kernel and row class that never ran eagerly is refused. The test runs every site at every row count it takes against its entry's bits and a float64 product (the sigmoid columns against its sigmoid), and replays the fused MLA site's long and wide kernels from a graph. Signed-off-by: Vasanth Sabavat --- .../decode_gemv.py | 168 +++++++++++++----- .../modeling.py | 18 +- .../test_modeling_v2_kimi_k3_decode_gemv.py | 109 ++++++++---- 3 files changed, 198 insertions(+), 97 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py index 6dee9d21e49b..b5c378e30e3c 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py @@ -2,9 +2,12 @@ # SPDX-License-Identifier: Apache-2.0 """The decode path's GEMVs, LM head and embedding on the catalog's single-GPU Kimi K3 entries. -* **Per-site GEMVs** (`K3DecodeGemvs.project`): an MLA projection of at most `MAX_ROWS` rows runs on the kernel - measured fastest at its weight shape at every row count 1..8 (`SITES`): `gemm/k3_decode_gemv` for the fused - q_a / kv_a projection, `gemm/k3_ctm_gemv_wide` for q_b and the output gate. +* **Per-site GEMVs** (`K3DecodeGemvs.project`): a projection of a decode step runs on the kernel measured fastest + at its call site's weight shape (`SITES`): at most `MAX_ROWS` rows on `gemm/k3_decode_gemv`, + `gemm/k3_ctm_gemv_wide` or `gemm/k3_ctm_gemv_long`, and, where the site lists it, up to `WIDE_ROWS` rows on + `gemm/k3_ctm_gemv_wide`. The sites are the decode kernels' fused projections (MLA's [W_a; W_g] with the gate rows + through a sigmoid, KDA's [q | k | v | g | f_a | b]), the attention output projection, and the built-in MLA path's + q_a / kv_a, q_b and gate projections. * **LM head** (`K3LogitsProcessor`): at most `MAX_ROWS` rows of this rank's vocabulary shard on `gemm/k3_head_gemv` over the target's `K3HeadGemvWorkspace`, then the shards gathered (`comm/allgather`) as the stock head gathers them. It is the shell's logits processor, so the speculative worker's target logits and the @@ -25,12 +28,17 @@ from __future__ import annotations -from typing import Dict, Optional, Set, Tuple +import math +from dataclasses import dataclass +from typing import Dict, Iterable, Optional, Set import torch from torch import nn from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.allgather import allgather +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv_long import ( + k3_ctm_gemv_long, +) from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv_wide import ( k3_ctm_gemv_wide, ) @@ -51,20 +59,78 @@ from tensorrt_llm._torch.flashinfer_utils import IS_FLASHINFER_AVAILABLE from tensorrt_llm._utils import mpi_disabled -# The row limit of the decode GEMV and head kernels (one token tile). +# The row limit of the decode GEMV and head kernels (one token tile), and of k3_ctm_gemv_wide (a decode step of 8 +# requests of 8 tokens). MAX_ROWS = 8 - -#: Site -> (N, K) of its weight at this target's per-rank shapes, and the kernel that runs it at 1..MAX_ROWS rows. -SITES: Dict[str, Tuple[int, int, str]] = { - # MLA kv_a_proj_with_mqa: [q_a 1536 | kv_a 512 | k_pe 64] of the hidden size. - "kv_a": (2112, 7168, "decode"), - # MLA q_b_proj: 6 heads x 192 of q_lora_rank. - "q_b": (1152, 1536, "wide"), - # MLA output gate: 6 heads x 128 of the hidden size. - "g_proj": (768, 7168, "wide"), +WIDE_ROWS = 64 + + +@dataclass(frozen=True) +class Site: + """A call site's weight shape (this target's per-rank shapes) and its kernels: ``small`` at 1..MAX_ROWS rows + ("decode", "wide" or "long"), and k3_ctm_gemv_wide at MAX_ROWS+1..WIDE_ROWS rows where ``wide``. Output columns + from ``sig_col0`` on are stored through a sigmoid. ``split`` / ``ring``: k3_ctm_gemv_long's CTAs per 128-row + weight tile and weight-ring stages.""" + + n: int + k: int + small: str + wide: bool = False + sig_col0: int = -1 + split: int = 0 + ring: int = 0 + + +SITES: Dict[str, Site] = { + # MLA's [W_a; W_g] on a decode step: [q_a 1536 | kv_a 512 | k_pe 64] then the output gate (6 heads x 128), + # the gate rows through a sigmoid. + "mla_ag": Site(2880, 7168, "long", wide=True, sig_col0=2112, split=6, ring=6), + # KDA's [q | k | v | g | f_a | b] on a decode step (6 heads x 128 each, then 128 and 6, padded to 3208 rows). + "kda_proj": Site(3208, 7168, "long", wide=True, split=5, ring=6), + # The attention output projection (row parallel; 6 heads x 128 in). + "o_proj": Site(7168, 768, "decode", wide=True), + # The built-in MLA path's projections: kv_a_proj_with_mqa, q_b_proj and the output gate. + "kv_a": Site(2112, 7168, "decode"), + "q_b": Site(1152, 1536, "wide"), + "g_proj": Site(768, 7168, "wide"), } +def _wide_tile(rows: int) -> int: + from tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv import k3_ctm_gemv_kernel + + return k3_ctm_gemv_kernel.wide_tile(rows) + + +def _run( + spec: Site, kernel: str, x2d: torch.Tensor, weight: torch.Tensor +) -> Optional[torch.Tensor]: + """``spec``'s ``kernel`` on dense rows ``x2d``, or None where it does not take them.""" + if kernel == "decode": + if spec.sig_col0 >= 0 or not _decode_op.supports(x2d, weight): + return None + return k3_decode_gemv(x2d, weight) + if kernel == "wide": + if not _ctm_op.supports_wide(x2d, weight, spec.sig_col0, False): + return None + return k3_ctm_gemv_wide(x2d, weight, sig_col0=spec.sig_col0) + # One wave of the GPU's SMs: beyond it the long GEMV loses to the others. + sms = torch.cuda.get_device_properties(x2d.device).multi_processor_count + if math.ceil(spec.n / 128) * spec.split > sms or not _ctm_op.supports_long( + x2d, weight, spec.split, spec.ring + ): + return None + return k3_ctm_gemv_long( + x2d, + weight, + sig_col0=spec.sig_col0, + split=spec.split, + ring=spec.ring, + trigger_early=True, + push=True, + ) + + def _capturing() -> bool: return not torch.compiler.is_compiling() and torch.cuda.is_current_stream_capturing() @@ -125,69 +191,81 @@ def __init__(self, head_workspace: Optional[K3HeadGemvWorkspace] = None) -> None @classmethod def create( - cls, lm_head: Optional[nn.Module], site_weights: Dict[str, torch.Tensor] + cls, + lm_head: Optional[nn.Module] = None, + sites: Iterable[str] = tuple(SITES), + device: Optional[torch.device] = None, ) -> "K3DecodeGemvs": - """The state for ``lm_head`` and the site weights (one weight per `SITES` key; every weight of a site has its - shape). Eager: it allocates the head's workspace and runs each kernel once on a zero row, so they compile here - rather than under a capture. A site or head its kernel does not take keeps the generic path.""" + """The state for ``lm_head`` and ``sites`` on ``device`` (default: the head's, else the current one). Eager: + it allocates the head's workspace and runs every kernel of every site once (each wide row class once) on + zero rows of a zero weight of the site's shape, so they compile here and not under a capture. A site or + head whose kernel does not take its shape keeps the generic path.""" if torch.cuda.is_current_stream_capturing(): raise RuntimeError( "K3DecodeGemvs.create allocates and compiles: run it before CUDA-graph capture" ) + head_weight = getattr(lm_head, "weight", None) if lm_head is not None else None + if device is None: + device = ( + head_weight.device + if isinstance(head_weight, torch.Tensor) + else torch.device("cuda", torch.cuda.current_device()) + ) state = cls() - for site, weight in site_weights.items(): - state._project(site, weight.new_zeros(1, weight.shape[1]), weight, warm=True) + for site in sites: + spec = SITES[site] + weight = torch.zeros(spec.n, spec.k, dtype=torch.bfloat16, device=device) + for rows in (1, 16, 32, 64) if spec.wide else (1,): + state._project(site, weight.new_zeros(rows, spec.k), weight, warm=True) + del weight if lm_head is not None and _head_takes_module(lm_head): - weight = lm_head.weight - x = weight.new_zeros(1, weight.shape[1]) - if _head_op.supports(x, weight): + x = head_weight.new_zeros(1, head_weight.shape[1]) + if _head_op.supports(x, head_weight): workspace = K3HeadGemvWorkspace.create( - weight.shape[0], weight.shape[1], weight.device + head_weight.shape[0], head_weight.shape[1], head_weight.device ) - k3_head_gemv(x, weight, workspace) + k3_head_gemv(x, head_weight, workspace) state.head_workspace = workspace state._ran.add(("lm_head",)) - torch.cuda.synchronize() + torch.cuda.synchronize(device) return state def project(self, site: str, x: torch.Tensor, weight: torch.Tensor) -> Optional[torch.Tensor]: - """``x @ weight.T`` (bf16 ``[..., N]``) for ``site``'s weight on its decode kernel, or None where the kernel - does not take the call: more than `MAX_ROWS` rows, another shape or dtype, or, under capture, a kernel that - has not run eagerly.""" + """``x @ weight.T`` (bf16 ``[..., N]``, the site's sigmoid columns through the sigmoid) for ``site``'s weight + on its decode kernel, or None where none takes the call: more rows than the site's kernels take, another + shape or dtype, or, under capture, a kernel that has not run eagerly. The caller then runs its GEMM.""" return self._project(site, x, weight, warm=False) def _project( self, site: str, x: torch.Tensor, weight: torch.Tensor, warm: bool ) -> Optional[torch.Tensor]: - n, k, kernel = SITES[site] + spec = SITES[site] if ( weight.dtype != torch.bfloat16 - or tuple(weight.shape) != (n, k) + or tuple(weight.shape) != (spec.n, spec.k) or not weight.is_contiguous() or x.dtype != torch.bfloat16 or x.dim() < 1 - or x.shape[-1] != k + or x.shape[-1] != spec.k ): return None - rows = x.numel() // k - if not 0 < rows <= MAX_ROWS: + rows = x.numel() // spec.k + if 0 < rows <= MAX_ROWS: + kernel = spec.small + elif spec.wide and MAX_ROWS < rows <= WIDE_ROWS: + kernel = "wide" + else: return None - key = (site,) + key = (site, kernel, _wide_tile(rows) if kernel == "wide" else 0) capturing = _capturing() if capturing and not warm and key not in self._ran: return None - x2d = _dense_rows(x.reshape(rows, k)) - if kernel == "decode": - if not _decode_op.supports(x2d, weight): - return None - y = k3_decode_gemv(x2d, weight) - else: - if not _ctm_op.supports_wide(x2d, weight, -1, False): - return None - y = k3_ctm_gemv_wide(x2d, weight) + y = _run(spec, kernel, _dense_rows(x.reshape(rows, spec.k)), weight) + if y is None: + return None if not capturing: self._ran.add(key) - return y.view(*x.shape[:-1], n) + return y.view(*x.shape[:-1], spec.n) def lm_head_logits(self, rows: torch.Tensor, lm_head: nn.Module) -> Optional[torch.Tensor]: """``lm_head(rows)``, the gathered bf16 logits ``[M, vocab]``, with this rank's shard on diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index 4af83170325a..d2d6203dfb2c 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -1868,22 +1868,10 @@ def load_weights(self, weights, *args, **kwargs): _weights.load(self, weights) def cache_derived_state(self) -> None: - """Build the decode GEMVs' state from the final weights: the LM head's workspace, and one eager call of each - decode GEMV kernel (the MLA projections' shapes, read off the first MLA layer), so none compiles under - capture.""" + """Build the decode GEMVs' state once the weights are final: the LM head's workspace, and one eager call of + every decode GEMV kernel at its site's shape (decode_gemv.SITES), so none compiles under capture.""" super().cache_derived_state() - mla = next((layer.self_attn.mixer for layer in self.model.layers if not layer.is_kda), None) - sites = {} - if mla is not None: - for site, name in ( - ("kv_a", "kv_a_proj_with_mqa"), - ("q_b", "q_b_proj"), - ("g_proj", "g_proj"), - ): - module = getattr(mla, name, None) - if module is not None: - sites[site] = module.weight - gemvs = _decode_gemv.K3DecodeGemvs.create(self.lm_head, sites) + gemvs = _decode_gemv.K3DecodeGemvs.create(self.lm_head) self.model.decode_gemvs = gemvs self.logits_processor.gemvs = gemvs diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_gemv.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_gemv.py index 86366db7fad6..c511657259b1 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_gemv.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_gemv.py @@ -3,8 +3,9 @@ """The Kimi K3 target's decode GEMVs, LM head and embedding on the catalog's single-GPU entries (``decode_gemv.py`` of ``kimi_k3_mxfp4__sm_100__tp16_moetp4ep4``), on one GPU. -* Each MLA projection site at 1..8 rows: the bits of its catalog entry's call, within 8e-3 of ``max |ref|`` of a - float64 product; declined above 8 rows, at another shape or dtype, and under capture before an eager call. +* Each GEMV site at every row count it takes (1..8, and 9..64 where k3_ctm_gemv_wide takes it): the bits of its + catalog entry's call, within 8e-3 of ``max |ref|`` of a float64 product (the sigmoid columns within 1e-2 of the + sigmoid of it); declined above its rows, at another shape or dtype, and under capture before an eager call. * The LM head through ``K3LogitsProcessor`` on a real ``LMHead`` (one rank): fp32 logits from ``k3_head_gemv`` within the same bound, the stock processor's rows selected, the stock path above 8 rows and before the state is built. @@ -18,6 +19,9 @@ import torch from torch import nn +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv_long import ( + k3_ctm_gemv_long, +) from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv_wide import ( k3_ctm_gemv_wide, ) @@ -51,6 +55,17 @@ def _rows(m, k, seed): return torch.randn(m, k, generator=g, device="cuda").to(torch.bfloat16) +def _check_product(y, x, w, sig_col0=-1): + """Within TOL of the float64 product (columns from ``sig_col0`` on: within 1e-2 of its sigmoid).""" + ref = x.double() @ w.double().t() + sig = ref.shape[1] if sig_col0 < 0 else sig_col0 + lin = ((y[:, :sig].double() - ref[:, :sig]).abs().max() / ref[:, :sig].abs().max()).item() + assert lin <= TOL, lin + if sig < ref.shape[1]: + err = (y[:, sig:].double() - ref[:, sig:].sigmoid()).abs().max().item() + assert err <= 1e-2, err + + def _rel_err(y, x, w): ref = x.double() @ w.double().t() return ((y.double() - ref).abs().max() / ref.abs().max()).item() @@ -60,53 +75,73 @@ def _bits(t): return t.contiguous().view(torch.int16) +def _entry(spec, rows, x, w): + """The site's catalog entry called directly: what ``project`` must reproduce bit for bit.""" + if rows > decode_gemv.MAX_ROWS or spec.small == "wide": + return k3_ctm_gemv_wide(x, w, sig_col0=spec.sig_col0) + if spec.small == "decode": + return k3_decode_gemv(x, w) + return k3_ctm_gemv_long( + x, w, sig_col0=spec.sig_col0, split=spec.split, ring=spec.ring, push=True + ) + + +def _site_rows(spec): + return list(ROWS) + ([9, 16, 24, 40, 64] if spec.wide else []) + + @pytest.mark.parametrize("site", list(decode_gemv.SITES)) def test_site(site): - """Every row count 1..8 runs the site's catalog entry: its bits, within TOL of the float64 product.""" - n, k, kernel = decode_gemv.SITES[site] - w = _weight(n, k, seed=len(site)) - gemvs = decode_gemv.K3DecodeGemvs.create(None, {site: w}) - entry = k3_decode_gemv if kernel == "decode" else k3_ctm_gemv_wide - for m in ROWS: - x = _rows(m, k, seed=m) + """Every row count the site takes runs its catalog entry: its bits, within TOL of the float64 product.""" + spec = decode_gemv.SITES[site] + w = _weight(spec.n, spec.k, seed=len(site)) + gemvs = decode_gemv.K3DecodeGemvs.create(None, sites=[site]) + for m in _site_rows(spec): + x = _rows(m, spec.k, seed=m) y = gemvs.project(site, x, w) - assert y is not None and y.shape == (m, n), (site, m) - assert torch.equal(_bits(y), _bits(entry(x, w))), (site, m) - assert _rel_err(y, x, w) <= TOL, (site, m) - - -def test_site_declines(): - """Above 8 rows, another weight shape, a non-bf16 input, and under capture before any eager call: None, nothing - launched.""" - n, k, _ = decode_gemv.SITES["kv_a"] - w = _weight(n, k, seed=1) - gemvs = decode_gemv.K3DecodeGemvs.create(None, {"kv_a": w}) - assert gemvs.project("kv_a", _rows(9, k, seed=9), w) is None - assert gemvs.project("kv_a", _rows(4, k, seed=4), _weight(n + 128, k, seed=2)) is None - assert gemvs.project("kv_a", _rows(4, k, seed=4).float(), w) is None + assert y is not None and y.shape == (m, spec.n), (site, m) + assert torch.equal(_bits(y), _bits(_entry(spec, m, x, w))), (site, m) + _check_product(y, x, w, spec.sig_col0) + + +@pytest.mark.parametrize("site", ["kv_a", "mla_ag"]) +def test_site_declines(site): + """More rows than the site takes, another weight shape, a non-bf16 input, and under capture before any eager + call: None, nothing launched.""" + spec = decode_gemv.SITES[site] + w = _weight(spec.n, spec.k, seed=1) + gemvs = decode_gemv.K3DecodeGemvs.create(None, sites=[site]) + too_many = decode_gemv.WIDE_ROWS + 1 if spec.wide else decode_gemv.MAX_ROWS + 1 + assert gemvs.project(site, _rows(too_many, spec.k, seed=9), w) is None + assert ( + gemvs.project(site, _rows(4, spec.k, seed=4), _weight(spec.n + 128, spec.k, seed=2)) is None + ) + assert gemvs.project(site, _rows(4, spec.k, seed=4).float(), w) is None fresh = decode_gemv.K3DecodeGemvs() - x = _rows(4, k, seed=4) + x = _rows(4, spec.k, seed=4) graph = torch.cuda.CUDAGraph() with torch.cuda.graph(graph): - y = fresh.project("kv_a", x, w) + y = fresh.project(site, x, w) assert y is None -def test_site_capture_replays_the_eager_bits(): - """A captured call, its input rewritten in place before each replay, gives the eager call's bits.""" - n, k, _ = decode_gemv.SITES["g_proj"] - w = _weight(n, k, seed=3) - gemvs = decode_gemv.K3DecodeGemvs.create(None, {"g_proj": w}) - x = _rows(4, k, seed=4) +@pytest.mark.parametrize("rows", [4, 16]) +def test_site_capture_replays_the_eager_bits(rows): + """A captured call (the long kernel at 4 rows, the wide one at 16, the gate columns through the sigmoid), its + input rewritten in place before each replay, gives the eager call's bits.""" + spec = decode_gemv.SITES["mla_ag"] + w = _weight(spec.n, spec.k, seed=3) + gemvs = decode_gemv.K3DecodeGemvs.create(None, sites=["mla_ag"]) + x = _rows(rows, spec.k, seed=4) graph = torch.cuda.CUDAGraph() with torch.cuda.graph(graph): - y = gemvs.project("g_proj", x, w) + y = gemvs.project("mla_ag", x, w) assert y is not None for seed in (5, 6): - x.copy_(_rows(4, k, seed=seed)) + x.copy_(_rows(rows, spec.k, seed=seed)) graph.replay() torch.cuda.synchronize() - assert torch.equal(_bits(y), _bits(gemvs.project("g_proj", x, w))) + assert torch.equal(_bits(y), _bits(gemvs.project("mla_ag", x, w))) def _lm_head(): @@ -131,7 +166,7 @@ def test_lm_head(): for m in (1, 9): x = _rows(m, HIDDEN, seed=m) assert torch.equal(processor(x, head, None, True), stock(x, head, None, True)) - processor.gemvs = decode_gemv.K3DecodeGemvs.create(head, {}) + processor.gemvs = decode_gemv.K3DecodeGemvs.create(head, sites=()) assert processor.gemvs.head_workspace is not None for m in ROWS: x = _rows(m, HIDDEN, seed=m) @@ -146,7 +181,7 @@ def test_lm_head_selects_the_stock_rows(): """Without context logits, the stock processor's last-token selection feeds the head.""" head = _lm_head() processor = decode_gemv.K3LogitsProcessor(LogitsProcessor()) - processor.gemvs = decode_gemv.K3DecodeGemvs.create(head, {}) + processor.gemvs = decode_gemv.K3DecodeGemvs.create(head, sites=()) x = _rows(7, HIDDEN, seed=3) metadata = types.SimpleNamespace( seq_lens_cuda=torch.tensor([3, 1, 3], dtype=torch.int32, device="cuda") @@ -158,7 +193,7 @@ def test_lm_head_selects_the_stock_rows(): def test_lm_head_capture_replays_the_eager_bits(): head = _lm_head() - gemvs = decode_gemv.K3DecodeGemvs.create(head, {}) + gemvs = decode_gemv.K3DecodeGemvs.create(head, sites=()) x = _rows(8, HIDDEN, seed=8) graph = torch.cuda.CUDAGraph() with torch.cuda.graph(graph): From 754d7060af17cfabafdec4aac40614c3df0adf6a Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 20:07:27 -0700 Subject: [PATCH 046/161] [None][fix] modeling_v2 Kimi K3 target: build the decode GEMVs' state once cache_derived_state may run again after the weights are loaded (a later load or a weight update). It built a new K3DecodeGemvs each time, and the old one's head workspace was freed while CUDA graphs captured in between still launched k3_head_gemv on it. A later call now keeps the state, as post_load_weights keeps the KDA buffers and the MLA workspace. Signed-off-by: Vasanth Sabavat --- .../kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index d9d2a653d3eb..b0d54b1fba5f 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -2254,8 +2254,11 @@ def load_weights(self, weights, *args, **kwargs): def cache_derived_state(self) -> None: """Build the decode GEMVs' state once the weights are final: the LM head's workspace, and one eager call of - every decode GEMV kernel at its site's shape (decode_gemv.SITES), so none compiles under capture.""" + every decode GEMV kernel at its site's shape (decode_gemv.SITES), so none compiles under capture. Built once: + a later call keeps it, since CUDA graphs captured in between hold its workspace.""" super().cache_derived_state() + if self.model.decode_gemvs is not None: + return gemvs = _decode_gemv.K3DecodeGemvs.create(self.lm_head) self.model.decode_gemvs = gemvs self.logits_processor.gemvs = gemvs From 8c6682190b55ec2f182e07669c78e7b7faab1d62 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 20:08:35 -0700 Subject: [PATCH 047/161] [None][fix] MnnvlWorkspace.create: the ranks agree before allocating; the contract states the failure model Before allocating, the ranks now agree that each of them can: not capturing, a valid buffer_bytes, and the three buffers within that rank's free device memory. If one rank cannot, every rank raises and none enters the collective allocation, where a rank that failed alone would leave its peers waiting in the handle exchange. The contract and the docstring state what remains: a rank that fails inside the handle exchange itself can still leave its peers waiting. The matrix gains check_create_refuses_on_every_rank: one rank asks for more than its free memory, every rank raises, and the next call is correct. Signed-off-by: Vasanth Sabavat --- .../catalog/comm/mnnvl_allreduce_attn_res.md | 12 +++++-- .../catalog/comm/mnnvl_workspace.py | 31 ++++++++++++++----- .../_mnnvl_allreduce_attn_res_op_matrix.py | 22 +++++++++++++ 3 files changed, 55 insertions(+), 10 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.md index 8275156a4f7f..d661a1fd0272 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.md @@ -80,9 +80,15 @@ a call over one buffer raises (below). **Who creates it, and when.** The target, in `post_load_weights`, with `MnnvlWorkspace.create(mapping, buffer_bytes, fabric_handle=None)`: -- collective over `mapping`'s TP group: every rank calls it at the same point; it returns on every rank or raises - on every rank (the success of each rank's allocation is agreed before anyone proceeds); -- eager: it allocates and exchanges handles, so it refuses to run under CUDA-graph capture; +- collective over `mapping`'s TP group: every rank calls it at the same point; +- failure model: + - before allocating, the ranks agree that each of them can (not capturing, a valid `buffer_bytes`, the three + buffers within that rank's free device memory). If one cannot, every rank raises `RuntimeError` and none + allocates (certified: one rank short of memory, every rank raises, the next call is correct); + - a failure that returns from the allocation is agreed the same way; + - a rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange; this + is not turned into an error on the other ranks; +- eager: it allocates and exchanges handles, so it refuses to run under CUDA-graph capture (every rank raises); - it arms every buffer word and the flags, and returns only once every rank has armed its buffers, so no peer can push into memory a rank has not armed; - `fabric_handle`: share the memory by fabric handle (required across nodes) or POSIX file descriptor; default diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_workspace.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_workspace.py index d8d7eec9bbe5..0d3c39ed9bc1 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_workspace.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_workspace.py @@ -54,28 +54,45 @@ def create( cls, mapping, buffer_bytes: int, fabric_handle: Optional[bool] = None ) -> "MnnvlWorkspace": """Allocate and arm a workspace for ``mapping``'s TP group. Collective: every rank of the group calls it at - the same point, eagerly (not under CUDA-graph capture); it returns on every rank or raises on every rank. + the same point, eagerly (not under CUDA-graph capture). + + Failure model: before allocating, the ranks agree that each of them can (not capturing, a valid + ``buffer_bytes``, the three buffers within its device's free memory); if one cannot, every rank raises + ``RuntimeError`` and none allocates. A failure that returns from the allocation is agreed the same way. A + rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange: that + failure is not turned into an error on the other ranks. + ``fabric_handle``: share the memory by fabric handle (required across nodes) rather than POSIX file descriptor; default ``mapping.is_multi_node()``.""" from tensorrt_llm._torch.distributed.ops import ( _get_mnnvl_workspace_comm, _initialize_allreduce_mnnvl_protocol, _make_mnnvl_mcast_buffer, + _mnnvl_device_index, _mnnvl_workspace_all_succeeded, ) + use_fabric = mapping.is_multi_node() if fabric_handle is None else bool(fabric_handle) + total = NUM_LAMPORT_BUFFERS * buffer_bytes + comm = _get_mnnvl_workspace_comm(mapping) + # Every condition one rank alone can fail is checked before the allocation, and the ranks agree on it: a + # rank failing inside the allocation would leave its peers in the handle exchange. + problem: Optional[str] = None if torch.cuda.is_current_stream_capturing(): + problem = "it is collective and allocates: call it before capture" + elif buffer_bytes <= 0 or buffer_bytes % 16: + problem = f"buffer_bytes must be a positive multiple of 16, got {buffer_bytes}" + else: + free_bytes, _ = torch.cuda.mem_get_info(_mnnvl_device_index(mapping)) + if free_bytes < total: + problem = f"its {total} bytes exceed the {free_bytes} free on this rank's device" + if not _mnnvl_workspace_all_succeeded(comm, problem is None): raise RuntimeError( - "MnnvlWorkspace.create is collective and allocates: call it before capture" + f"MnnvlWorkspace.create: not every rank can allocate ({problem or 'another rank cannot'})" ) - if buffer_bytes <= 0 or buffer_bytes % 16: - raise ValueError(f"buffer_bytes must be a positive multiple of 16, got {buffer_bytes}") - use_fabric = mapping.is_multi_node() if fabric_handle is None else bool(fabric_handle) - comm = _get_mnnvl_workspace_comm(mapping) error: Optional[Exception] = None workspace = None try: - total = NUM_LAMPORT_BUFFERS * buffer_bytes handle = _make_mnnvl_mcast_buffer(comm, total, mapping, use_fabric) lamport = handle.get_uc_buffer(mapping.tp_rank, (total // 4,), torch.float32, 0) flags = torch.zeros(FLAG_WORDS, dtype=torch.uint32, device=lamport.device) diff --git a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allreduce_attn_res_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allreduce_attn_res_op_matrix.py index ee5d18600687..536aecb2d7e9 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allreduce_attn_res_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allreduce_attn_res_op_matrix.py @@ -100,6 +100,27 @@ def check_workspace_is_armed_and_sized() -> None: assert WS_A.buffer_flags.view(torch.int32).tolist()[:3] == [0, 2, BUFFER_BYTES] +def check_create_refuses_on_every_rank() -> None: + """One rank asks for three buffers past its device's free memory: every rank raises before any allocates, and the + workspaces in use stay correct.""" + from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.mnnvl_workspace import ( + MnnvlWorkspace, + ) + + free_bytes, _ = torch.cuda.mem_get_info() + too_big = (free_bytes // 3 // 16 + (64 << 20) // 16) * 16 + try: + MnnvlWorkspace.create( + R.mapping, too_big if R.rank == 0 else BUFFER_BYTES, fabric_handle=R.fabric + ) + raised = False + except RuntimeError as exc: + raised = "not every rank can allocate" in str(exc) + assert R.all_true(raised), "a rank short of memory did not make every rank raise" + call = Call(2100, 8, 3) + verify(call, call.run(WS_A), "after the refused create") + + def check_single_calls() -> None: for t in TOKENS: for s in SNAPSHOTS: @@ -221,6 +242,7 @@ def check_wrong_call_order_is_detected() -> None: CHECKS = [ check_workspace_is_armed_and_sized, + check_create_refuses_on_every_rank, check_single_calls, check_a_call_over_one_buffer_raises_on_every_rank, check_dip_and_regrow_sequence, From a96450058d3679d09b23d84406e0eb86ad8abd8e Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 20:16:51 -0700 Subject: [PATCH 048/161] [None][feat] modeling_v2 Kimi K3 target: layer 0's dense MLP on the decode kernels On a step of at most 8 tokens, layer 0's SiTU MLP runs as the K3 stack runs it: gate_up on gemm/k3_ctm_gemv_long (split 4, ring 5), the SiTU gate on activation/k3_situ_mul, down on gemm/k3_ctm_gemv_long (split 2, ring 6), then the down projection's all-reduce (K3DecodeGemvs.dense_mlp). Any other step, and a call the kernels do not take, runs the module. The kernels take the MLP split over the whole TP group (per rank gate_up [4224, 7168], down [7168, 2112]). As in the stack, the layer keeps that split when its all-reduce runs over MNNVL (one NVLink domain across the nodes, where a cross-node all-reduce costs what a node's does) instead of capping it at the GPUs of one node. This applies on every step, prefill included, so layer 0's partial sums now reduce over 16 ranks instead of 4. create() warms both GEMVs and the activation (with and without its linear beta). The test checks the chained entries' bits and the float64 SituAndMul MLP at 1..8 rows, the refusals, and a graph replay. Signed-off-by: Vasanth Sabavat --- .../decode_gemv.py | 61 +++++++++++++++-- .../modeling.py | 37 ++++++++-- .../test_modeling_v2_kimi_k3_decode_gemv.py | 68 ++++++++++++++++++- 3 files changed, 156 insertions(+), 10 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py index b5c378e30e3c..b384010b65dd 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py @@ -14,6 +14,9 @@ drafter's logits on the same head take it too. * **Embedding** (`K3DecodeGemvs.embed_norm`): a decode step's embedding rows written into the attention-residual bank's slot 0 (layer 0's first snapshot) and layer 0's input RMSNorm applied, in one `norm/k3_embed_norm` launch. +* **Dense MLP** (`K3DecodeGemvs.dense_mlp`): layer 0's MLP at most `MAX_ROWS` rows, split over the whole TP group: + gate_up on `gemm/k3_ctm_gemv_long`, `activation/k3_situ_mul`, down on `gemm/k3_ctm_gemv_long`; the caller then + runs the down projection's all-reduce. Each returns None where its kernel does not take the call, and the caller then runs the generic path's module. @@ -35,6 +38,7 @@ import torch from torch import nn +from tensorrt_llm._torch._experimental.modeling_v2.catalog.activation.k3_situ_mul import k3_situ_mul from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.allgather import allgather from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv_long import ( k3_ctm_gemv_long, @@ -69,8 +73,8 @@ class Site: """A call site's weight shape (this target's per-rank shapes) and its kernels: ``small`` at 1..MAX_ROWS rows ("decode", "wide" or "long"), and k3_ctm_gemv_wide at MAX_ROWS+1..WIDE_ROWS rows where ``wide``. Output columns - from ``sig_col0`` on are stored through a sigmoid. ``split`` / ``ring``: k3_ctm_gemv_long's CTAs per 128-row - weight tile and weight-ring stages.""" + from ``sig_col0`` on are stored through a sigmoid. ``split`` / ``ring`` / ``push``: k3_ctm_gemv_long's CTAs per + 128-row weight tile, weight-ring stages, and whether the partial sums are pushed to each row's owner.""" n: int k: int @@ -79,20 +83,24 @@ class Site: sig_col0: int = -1 split: int = 0 ring: int = 0 + push: bool = False SITES: Dict[str, Site] = { # MLA's [W_a; W_g] on a decode step: [q_a 1536 | kv_a 512 | k_pe 64] then the output gate (6 heads x 128), # the gate rows through a sigmoid. - "mla_ag": Site(2880, 7168, "long", wide=True, sig_col0=2112, split=6, ring=6), + "mla_ag": Site(2880, 7168, "long", wide=True, sig_col0=2112, split=6, ring=6, push=True), # KDA's [q | k | v | g | f_a | b] on a decode step (6 heads x 128 each, then 128 and 6, padded to 3208 rows). - "kda_proj": Site(3208, 7168, "long", wide=True, split=5, ring=6), + "kda_proj": Site(3208, 7168, "long", wide=True, split=5, ring=6, push=True), # The attention output projection (row parallel; 6 heads x 128 in). "o_proj": Site(7168, 768, "decode", wide=True), # The built-in MLA path's projections: kv_a_proj_with_mqa, q_b_proj and the output gate. "kv_a": Site(2112, 7168, "decode"), "q_b": Site(1152, 1536, "wide"), "g_proj": Site(768, 7168, "wide"), + # Layer 0's dense MLP split over the 16-way TP group: gate_up [gate 2112 | up 2112] and down. + "dense_gate_up": Site(4224, 7168, "long", split=4, ring=5), + "dense_down": Site(7168, 2112, "long", split=2, ring=6), } @@ -127,7 +135,7 @@ def _run( split=spec.split, ring=spec.ring, trigger_early=True, - push=True, + push=spec.push, ) @@ -218,6 +226,12 @@ def create( for rows in (1, 16, 32, 64) if spec.wide else (1,): state._project(site, weight.new_zeros(rows, spec.k), weight, warm=True) del weight + if "dense_gate_up" in sites: + gu = torch.zeros(1, SITES["dense_gate_up"].n, dtype=torch.bfloat16, device=device) + if _ctm_op.supports_situ_mul(gu): + for linear_beta in (None, 1.0): + k3_situ_mul(gu, 1.0, linear_beta) + state._ran.add(("situ_mul", linear_beta is not None)) if lm_head is not None and _head_takes_module(lm_head): x = head_weight.new_zeros(1, head_weight.shape[1]) if _head_op.supports(x, head_weight): @@ -297,6 +311,43 @@ def lm_head_logits(self, rows: torch.Tensor, lm_head: nn.Module) -> Optional[tor gathered = allgather(local, None, group) return concat(list(split(gathered, rows.shape[0], dim=0)), dim=-1) + def dense_mlp( + self, + x: torch.Tensor, + gate_up_weight: torch.Tensor, + down_weight: torch.Tensor, + beta: float, + linear_beta: Optional[float], + ) -> Optional[torch.Tensor]: + """Layer 0's dense MLP of at most `MAX_ROWS` rows, before its down projection's all-reduce: gate_up on + ``gemm/k3_ctm_gemv_long``, ``SituAndMul(beta, linear_beta)`` on ``activation/k3_situ_mul``, down on + ``gemm/k3_ctm_gemv_long``. None, with nothing launched, where a kernel does not take the call; the caller then + runs the module.""" + gate_up, down = SITES["dense_gate_up"], SITES["dense_down"] + rows = x.shape[0] if x.dim() == 2 else 0 + if ( + not 0 < rows <= MAX_ROWS + or linear_beta == 0.0 + or tuple(gate_up_weight.shape) != (gate_up.n, gate_up.k) + or tuple(down_weight.shape) != (down.n, down.k) + or gate_up.n != 2 * down.k + ): + return None + if ( + _capturing() + and not { + ("dense_gate_up", "long", 0), + ("dense_down", "long", 0), + ("situ_mul", linear_beta is not None), + } + <= self._ran + ): + return None + gu = self.project("dense_gate_up", x, gate_up_weight) + if gu is None or not _ctm_op.supports_situ_mul(gu): + return None + return self.project("dense_down", k3_situ_mul(gu, beta, linear_beta), down_weight) + def embed_norm( self, input_ids: torch.Tensor, diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index b0d54b1fba5f..2ad5654095ec 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -1408,7 +1408,11 @@ def __init__( self.mlp_tp_size = 1 else: self.mlp_tp_size = math.gcd(cfg.intermediate_size, model_config.mapping.tp_size) - if self.mlp_tp_size > model_config.mapping.gpus_per_node: + # Over MNNVL (one NVLink domain across the nodes, where a cross-node all-reduce costs what a node's + # does) the MLP stays split over the whole TP group, the per-rank shapes its decode GEMVs take + # (decode_gemv.SITES); otherwise it stays within one node. + spans_nodes = self._mnnvl_allreduce() is not None + if self.mlp_tp_size > model_config.mapping.gpus_per_node and not spans_nodes: self.mlp_tp_size = math.gcd( self.mlp_tp_size, model_config.mapping.gpus_per_node ) @@ -1416,8 +1420,7 @@ def __init__( mlp_model_config.quant_config = QuantConfig() # K3's dense layer is BF16, so a unit block size gives the same # subgroup selection as DeepSeek-V3. Attention DP replicates the - # MLP because ranks own different tokens; otherwise the subgroup - # is block-aligned and stays within one node. + # MLP because ranks own different tokens. self.mlp = GatedMLP( hidden_size=cfg.hidden_size, intermediate_size=cfg.intermediate_size, @@ -1433,6 +1436,9 @@ def __init__( reduce_output=self.mlp_tp_size > 1, layer_idx=layer_idx, ) + self._situ = (situ_beta, situ_linear_beta) + # The decode GEMVs' state (decode_gemv.py), set by the target's cache_derived_state. + self.decode_gemvs: Optional[_decode_gemv.K3DecodeGemvs] = None # Stock fused RMSNorm for the plain (whole-tensor) norms; numerics # are drop-in for KimiK3RMSNorm (fp32 variance, weight applied @@ -1556,11 +1562,31 @@ def forward( hidden_states, getattr(attn_metadata, "all_rank_num_tokens", None) ) else: - hidden_states = self.mlp(hidden_states) + hidden_states = self._dense_mlp(hidden_states, step) prefix_sum = prefix_sum + hidden_states return prefix_sum, num_snapshots + def _dense_mlp(self, hidden_states: torch.Tensor, step: Optional[DecodeStep]) -> torch.Tensor: + """The dense MLP: on a step of at most DECODE_MAX_TOKENS tokens, its GEMVs and activation on the decode + kernels (``K3DecodeGemvs.dense_mlp``), then the down projection's all-reduce; else the module.""" + gemvs = self.decode_gemvs + if gemvs is not None and step is not None and step.small: + out = gemvs.dense_mlp( + hidden_states, + self.mlp.gate_up_proj.weight, + self.mlp.down_proj.weight, + *self._situ, + ) + if out is not None: + return self.mlp.down_proj.all_reduce(out) if self.mlp_tp_size > 1 else out + return self.mlp(hidden_states) + + def _mnnvl_allreduce(self): + """The MNNVL all-reduce of this layer's attention output, or None.""" + attention = self.linear_attn if self.is_kda else self.self_attn + return getattr(getattr(attention, "_o_allreduce", None), "mnnvl_allreduce", None) + def skip_forward( self, hidden_states: torch.Tensor, @@ -2262,6 +2288,9 @@ def cache_derived_state(self) -> None: gemvs = _decode_gemv.K3DecodeGemvs.create(self.lm_head) self.model.decode_gemvs = gemvs self.logits_processor.gemvs = gemvs + for layer in self.model.layers: + if not layer.is_moe: + layer.decode_gemvs = gemvs def post_load_weights(self) -> None: """The state the decode kernels share, built once per device before any CUDA-graph capture and handed to diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_gemv.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_gemv.py index c511657259b1..f87c8c6a6832 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_gemv.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_gemv.py @@ -10,6 +10,8 @@ within the same bound, the stock processor's rows selected, the stock path above 8 rows and before the state is built. * The embedding: bit-identical to ``nn.Embedding`` followed by the stock RMSNorm, the rows in the bank's slot 0. +* Layer 0's dense MLP at 1..8 rows: the bits of its three entries chained, within 3e-2 of ``max |ref|`` of the + float64 SituAndMul MLP; declined above 8 rows, at linear_beta 0 and under capture before an eager call. * CUDA-graph replays give the eager bits. """ @@ -19,6 +21,7 @@ import torch from torch import nn +from tensorrt_llm._torch._experimental.modeling_v2.catalog.activation.k3_situ_mul import k3_situ_mul from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv_long import ( k3_ctm_gemv_long, ) @@ -82,7 +85,7 @@ def _entry(spec, rows, x, w): if spec.small == "decode": return k3_decode_gemv(x, w) return k3_ctm_gemv_long( - x, w, sig_col0=spec.sig_col0, split=spec.split, ring=spec.ring, push=True + x, w, sig_col0=spec.sig_col0, split=spec.split, ring=spec.ring, push=spec.push ) @@ -247,3 +250,66 @@ def test_embed_norm_capture(): graph.replay() torch.cuda.synchronize() assert torch.equal(_bits(captured), _bits(eager)) + + +SITU = (4.0, 25.0) # the Kimi K3 checkpoint's (beta, linear_beta) + + +def _dense_weights(): + gate_up, down = decode_gemv.SITES["dense_gate_up"], decode_gemv.SITES["dense_down"] + return _weight(gate_up.n, gate_up.k, seed=41), _weight(down.n, down.k, seed=42) + + +def _dense_reference(x, w_gu, w_down, beta, linear_beta): + gu = x.double() @ w_gu.double().t() + g, u = gu.chunk(2, dim=1) + a = beta * torch.tanh(g / beta) * torch.sigmoid(g) + v = u if linear_beta is None else linear_beta * torch.tanh(u / linear_beta) + return (a * v) @ w_down.double().t() + + +@pytest.mark.parametrize("situ", [SITU, (1.0, None)], ids=["k3", "default"]) +def test_dense_mlp(situ): + """At 1..8 rows: the bits of the gate_up GEMV, k3_situ_mul and the down GEMV called directly, within 3e-2 of the + float64 MLP.""" + w_gu, w_down = _dense_weights() + gemvs = decode_gemv.K3DecodeGemvs.create(None, sites=["dense_gate_up", "dense_down"]) + gate_up, down = decode_gemv.SITES["dense_gate_up"], decode_gemv.SITES["dense_down"] + for m in ROWS: + x = _rows(m, HIDDEN, seed=m) + y = gemvs.dense_mlp(x, w_gu, w_down, *situ) + assert y is not None and y.shape == (m, down.n), m + gu = k3_ctm_gemv_long(x, w_gu, split=gate_up.split, ring=gate_up.ring) + want = k3_ctm_gemv_long(k3_situ_mul(gu, *situ), w_down, split=down.split, ring=down.ring) + assert torch.equal(_bits(y), _bits(want)), m + ref = _dense_reference(x, w_gu, w_down, *situ) + err = ((y.double() - ref).abs().max() / ref.abs().max()).item() + assert err <= 3e-2, (m, err) + + +def test_dense_mlp_declines(): + """9 rows, linear_beta 0 (the activation would run it as 1), and under capture before any eager call: None.""" + w_gu, w_down = _dense_weights() + gemvs = decode_gemv.K3DecodeGemvs.create(None, sites=["dense_gate_up", "dense_down"]) + assert gemvs.dense_mlp(_rows(9, HIDDEN, seed=9), w_gu, w_down, *SITU) is None + assert gemvs.dense_mlp(_rows(4, HIDDEN, seed=4), w_gu, w_down, 4.0, 0.0) is None + fresh = decode_gemv.K3DecodeGemvs() + x = _rows(4, HIDDEN, seed=4) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + y = fresh.dense_mlp(x, w_gu, w_down, *SITU) + assert y is None + + +def test_dense_mlp_capture_replays_the_eager_bits(): + w_gu, w_down = _dense_weights() + gemvs = decode_gemv.K3DecodeGemvs.create(None, sites=["dense_gate_up", "dense_down"]) + x = _rows(8, HIDDEN, seed=8) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + y = gemvs.dense_mlp(x, w_gu, w_down, *SITU) + assert y is not None + x.copy_(_rows(8, HIDDEN, seed=12)) + graph.replay() + torch.cuda.synchronize() + assert torch.equal(_bits(y), _bits(gemvs.dense_mlp(x, w_gu, w_down, *SITU))) From f224355f930b15f9c707fd0ecb1654974043858a Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 20:50:40 -0700 Subject: [PATCH 049/161] [None][feat] modeling_v2 Kimi K3 target: the attention's projections on the decode GEMV sites The decode branches' projections run on decode_gemv.py's sites: - K3DecodeMLA: [W_a; W_g] on mla_ag (the gate columns through the sigmoid in the kernel) and o_proj on o_proj; - K3DecodeKDA: the verify rows on kda_proj, and o_proj on o_proj on every step decode_step classifies, then the module's all-reduce. Those steps now take the core from the built-in dispatch when the plain-decode kernel does not run them. Each falls back to its GEMM, or to the module, where the site does not take the call. The target's post_load_weights hands the decode GEMVs' state, which cache_derived_state builds, to every attention module. Signed-off-by: Vasanth Sabavat --- .../modeling.py | 99 +++++++++++++------ 1 file changed, 70 insertions(+), 29 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index 2ad5654095ec..d513e2fb7e24 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -39,11 +39,12 @@ * MLA (`K3DecodeMLA`): `attention/k3_mla_qkv` (the query path and the step's latent KV rows into the paged cache), then `attention/k3_mla_attn_vb_out` (the attention, v_b and the output gate in one launch). -The state those kernels share (the KDA projection's Lamport buffers, the MLA attention workspace) lives in typed -objects this target creates in `post_load_weights`, before any graph capture. The token-count kernels (the decode -GEMVs, the MoE front and routed experts, the sandwiches, the embedding and residual epilogues) come with their own -entries; until then every module but the attention runs the generic path on every step, as do the projections around -the attention kernels (the [W_a; W_g] and verify-row GEMMs, `o_proj` and its all-reduce). +The projections around them run on the decode GEMV sites of `decode_gemv.py` (the [W_a; W_g] and KDA verify-row +projections, `o_proj` on every classified step), as do the LM head, the embedding and layer 0's dense MLP. The state +those kernels share (the KDA projection's Lamport buffers, the MLA attention workspace, the decode GEMVs' state) +lives in typed objects this target creates in `post_load_weights`, before any graph capture. The MoE front and routed +experts, the sandwiches and the residual epilogues come with their own entries; until then they run the generic path +on every step. **What this target asserts rather than adapts**: SM 10.0; the topology above; bf16 weights and a bf16 KV pool; tokens_per_block 64 (the MLA generation kernels K3's 96 heads reach exist only at 64); the V2 hybrid KV / state @@ -1777,15 +1778,18 @@ class K3DecodeKDA(KimiKDALinearAttention): """Kimi K3's KDA attention: the built-in module, with the decode kernels on the steps they take. * A decode step of one token per request runs the fused input projection and the plain decode in one - ``ssm/k3_kda_decode_attn`` launch, then the module's ``o_proj`` and all-reduce. + ``ssm/k3_kda_decode_attn`` launch. * With the cache manager's per-token states (``KimiLinearModel.kda_token_states``), every verify of the layer, on any step, runs the kernels that keep them: ``ssm/k3_kda_attn`` for one request of 8 tokens (the projection fused - in), else ``ssm/k3_kda_verify`` on the projection's rows. The built-in verify replays drafts from caches these - kernels do not fill, so the two never run on one manager. + in), else ``ssm/k3_kda_verify`` on the projection's rows (the ``kda_proj`` decode GEMV site where it takes + them). The built-in verify replays drafts from caches these kernels do not fill, so the two never run on one + manager. + * On every step ``decode_step`` classifies, ``o_proj`` runs on the ``o_proj`` decode GEMV site where it takes the + rows, then the module's all-reduce. Every other step runs the built-in module. The kernels read one ``[q | k | v | g | f_a | b]`` weight, built at the - checkpoint load from the module's own, and the device's ``K3KdaBuffers``, which the target sets in - ``post_load_weights``. + checkpoint load from the module's own, and the device's ``K3KdaBuffers`` and decode GEMVs' state, which the + target sets in ``post_load_weights``. """ def __init__(self, *args, **kwargs) -> None: @@ -1795,6 +1799,8 @@ def __init__(self, *args, **kwargs) -> None: self.k3_proj_weight: Optional[torch.Tensor] = None # The fused projection's Lamport buffers: one set per device, shared by every KDA layer. self.k3_buffers: Optional[K3KdaBuffers] = None + # The decode GEMVs' state (decode_gemv.py), shared by the target's layers. + self.decode_gemvs: Optional[_decode_gemv.K3DecodeGemvs] = None @property def takes_k3_kernels(self) -> bool: @@ -1829,19 +1835,33 @@ def forward( attn_metadata: AttentionMetadata, step: Optional[DecodeStep] = None, ) -> torch.Tensor: - """The built-in forward, except on a decode step of one token per request.""" + """The built-in forward on a step ``decode_step`` does not classify, and under a breakable CUDA graph. On the + others: the plain decode on ``ssm/k3_kda_decode_attn`` on a decode step of one token per request, else the + built-in dispatch; then ``o_proj`` on its decode GEMV site.""" + if step is None or is_in_breakable_cuda_graph(): + return super().forward(hidden_states, attn_metadata) if ( - step is not None - and step.decode + step.decode and step.tokens_per_request == 1 and self.takes_k3_kernels and self.k3_buffers is not None - and not is_in_breakable_cuda_graph() ): - return self._project_output( - self._k3_decode(hidden_states[: step.num_tokens], attn_metadata) + core = self._k3_decode(hidden_states[: step.num_tokens], attn_metadata) + else: + core = self._forward_impl(hidden_states, attn_metadata) + return self._k3_project_output(core) + + def _k3_project_output(self, core: torch.Tensor) -> torch.Tensor: + """``o_proj`` on the ``o_proj`` decode GEMV site where it takes the rows (else the module), then the TP + all-reduce.""" + out = None + if self.decode_gemvs is not None: + out = self.decode_gemvs.project( + "o_proj", core.reshape(-1, self.proj_size), self.o_proj.weight ) - return super().forward(hidden_states, attn_metadata) + if out is None: + return self._project_output(core) + return out if self._o_allreduce is None else self._o_allreduce(out) def _k3_decode(self, x: torch.Tensor, attn_metadata: AttentionMetadata) -> torch.Tensor: """``ssm/k3_kda_decode_attn``: the core output ``[R, H, 128]`` of one token of each of the step's R requests; @@ -1898,8 +1918,9 @@ def forward_verify( def _k3_verify(self, x, num_steps, layer_cache, ssm_pool, slot_indices) -> torch.Tensor: """The core output ``[N T, H, 128]`` of N requests of T verify tokens: ``ssm/k3_kda_attn`` for one request of - 8 tokens, else ``ssm/k3_kda_verify`` on the projection's rows. Each slot's state after the golden token, its - drafts' states and its conv window are written in place.""" + 8 tokens, else ``ssm/k3_kda_verify`` on the projection's rows (the ``kda_proj`` decode GEMV site where it + takes them). Each slot's state after the golden token, its drafts' states and its conv window are written in + place.""" num_requests = x.shape[0] // num_steps slots = slot_indices[:num_requests] w_q, w_k, w_v = self._get_mtp_conv_weights() @@ -1927,8 +1948,13 @@ def _k3_verify(self, x, num_steps, layer_cache, ssm_pool, slot_indices) -> torch *constants, ) else: + rows = None + if self.decode_gemvs is not None: + rows = self.decode_gemvs.project("kda_proj", x, self.k3_proj_weight) + if rows is None: + rows = torch.nn.functional.linear(x, self.k3_proj_weight) out = k3_kda_verify( - torch.nn.functional.linear(x, self.k3_proj_weight), + rows, self.f_b_proj.weight, w_q, w_k, @@ -1952,14 +1978,16 @@ def _k3_verify(self, x, num_steps, layer_cache, ssm_pool, slot_indices) -> torch class K3DecodeMLA(KimiK3MLAAttention): """Kimi K3's MLA attention: the built-in module, with a decode step's attention on the decode kernels. - A decode step runs ``x [W_a; W_g]^T`` as one GEMM with the gate columns through a sigmoid, then - ``attention/k3_mla_qkv`` (the q_a / kv_a RMSNorms, q_b and the k_b absorption into the fused query, the step's - latent rows stored into the paged cache) and ``attention/k3_mla_attn_vb_out`` (the attention over the paged cache, - v_b and the output gate in one launch), then the module's ``o_proj``. Every other step, and a decode step whose + A decode step runs ``x [W_a; W_g]^T`` with the gate columns through a sigmoid (the ``mla_ag`` decode GEMV site + where it takes the rows, else one GEMM), then ``attention/k3_mla_qkv`` (the q_a / kv_a RMSNorms, q_b and the k_b + absorption into the fused query, the step's latent rows stored into the paged cache) and + ``attention/k3_mla_attn_vb_out`` (the attention over the paged cache, v_b and the output gate in one launch), then + ``o_proj`` (the ``o_proj`` site where it takes the rows, else the module). Every other step, and a decode step whose cache the kernels do not read (``k3_mla_decode_view`` says why), runs the built-in module. ``[W_a; W_g]`` is built at load from the module's weights, which become views of it. The attention workspace is - the device's ``K3MlaAttnWorkspace``, which the target sets in ``post_load_weights``. + the device's ``K3MlaAttnWorkspace``; it and the decode GEMVs' state are set by the target in + ``post_load_weights``. """ def __init__(self, **kwargs) -> None: @@ -1968,6 +1996,8 @@ def __init__(self, **kwargs) -> None: self.k3_ag_weight: Optional[torch.Tensor] = None # The decode attention's workspace: one per device, shared by every MLA layer. self.k3_workspace: Optional[K3MlaAttnWorkspace] = None + # The decode GEMVs' state (decode_gemv.py), shared by the target's layers. + self.decode_gemvs: Optional[_decode_gemv.K3DecodeGemvs] = None def post_load_weights(self) -> None: """The built-in post-load, then ``[W_a; W_g]`` (once: CUDA graphs captured since read it).""" @@ -2048,8 +2078,11 @@ def forward( ) x = hidden_states[: step.num_tokens].contiguous() rows = self.kv_a_proj_with_mqa.weight.shape[0] - ag = torch.nn.functional.linear(x, self.k3_ag_weight) - ag[:, rows:].sigmoid_() + gemvs = self.decode_gemvs + ag = None if gemvs is None else gemvs.project("mla_ag", x, self.k3_ag_weight) + if ag is None: + ag = torch.nn.functional.linear(x, self.k3_ag_weight) + ag[:, rows:].sigmoid_() fused_q = k3_mla_qkv( ag, self.q_a_layernorm.weight, @@ -2079,7 +2112,12 @@ def forward( gate=ag, gate_col0=rows, ) - return self._project_output([attn_output], position_ids, attn_metadata, all_reduce_params) + out = None if gemvs is None else gemvs.project("o_proj", attn_output, self.o_proj.weight) + if out is None: + out = self._project_output( + [attn_output], position_ids, attn_metadata, all_reduce_params + ) + return out def _k3_decode_view(self, attn_metadata: AttentionMetadata, num_tokens: int) -> Optional[dict]: """The paged latent cache as the decode kernels read it this step, or None when they do not read it (the @@ -2294,7 +2332,8 @@ def cache_derived_state(self) -> None: def post_load_weights(self) -> None: """The state the decode kernels share, built once per device before any CUDA-graph capture and handed to - every layer of its kind: the KDA projection's Lamport buffers and the MLA decode attention's workspace.""" + every layer of its kind: the KDA projection's Lamport buffers and the MLA decode attention's workspace; and + the decode GEMVs' state (built by ``cache_derived_state``) handed to every attention module.""" super().post_load_weights() kda = [layer.linear_attn for layer in self.model.layers if layer.is_kda] mla = [layer.self_attn.mixer for layer in self.model.layers if not layer.is_kda] @@ -2307,6 +2346,8 @@ def post_load_weights(self) -> None: module.k3_buffers = buffers for module in mla: module.k3_workspace = workspace + for module in kda + mla: + module.decode_gemvs = self.model.decode_gemvs logger.info( "Kimi K3 decode kernels: KDA on k3_kda_decode_attn, k3_kda_attn and k3_kda_verify " f"({sum(m.takes_k3_kernels for m in kda)} / {len(kda)} layers take them), MLA on k3_mla_qkv and " From 65cfd479a605fae5ed1d07205d41aec293fc1c9c Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 20:54:02 -0700 Subject: [PATCH 050/161] [None][fix] MnnvlWorkspace.create frees its communicator split when it raises create() splits an MPI communicator for each call. On the two failures every rank agrees on (a rank cannot allocate; the allocation failed on a rank) it raised and kept the split. Every rank now frees it on both paths. Under Ray the communicator is the TP ProcessGroup, which belongs to c10d and is only dropped. An allocated handle only borrows the communicator and makes no MPI call when destroyed. check_create_refuses_on_every_rank also asserts that every rank freed the communicator it split. Signed-off-by: Vasanth Sabavat --- .../catalog/comm/mnnvl_allreduce_attn_res.md | 7 ++++--- .../modeling_v2/catalog/comm/mnnvl_workspace.py | 15 ++++++++++++--- .../comm/_mnnvl_allreduce_attn_res_op_matrix.py | 17 +++++++++++++++-- 3 files changed, 31 insertions(+), 8 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.md index d661a1fd0272..61a3faeffdd6 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.md @@ -83,9 +83,10 @@ a call over one buffer raises (below). - collective over `mapping`'s TP group: every rank calls it at the same point; - failure model: - before allocating, the ranks agree that each of them can (not capturing, a valid `buffer_bytes`, the three - buffers within that rank's free device memory). If one cannot, every rank raises `RuntimeError` and none - allocates (certified: one rank short of memory, every rank raises, the next call is correct); - - a failure that returns from the allocation is agreed the same way; + buffers within that rank's free device memory). If one cannot, every rank raises `RuntimeError`, none + allocates, and under MPI each frees the communicator it split for the call (certified: one rank short of + memory, every rank raises and frees its split, the next call is correct); + - a failure that returns from the allocation is agreed and handled the same way; - a rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange; this is not turned into an error on the other ranks; - eager: it allocates and exchanges handles, so it refuses to run under CUDA-graph capture (every rank raises); diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_workspace.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_workspace.py index 0d3c39ed9bc1..e9de6d592f7a 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_workspace.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_workspace.py @@ -58,9 +58,10 @@ def create( Failure model: before allocating, the ranks agree that each of them can (not capturing, a valid ``buffer_bytes``, the three buffers within its device's free memory); if one cannot, every rank raises - ``RuntimeError`` and none allocates. A failure that returns from the allocation is agreed the same way. A - rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange: that - failure is not turned into an error on the other ranks. + ``RuntimeError``, none allocates, and under MPI each frees the communicator it split for the call. A failure + that returns from the allocation is agreed and handled the same way. A rank that fails inside the + allocation's handle exchange can leave its peers waiting in that exchange: that failure is not turned into + an error on the other ranks. ``fabric_handle``: share the memory by fabric handle (required across nodes) rather than POSIX file descriptor; default ``mapping.is_multi_node()``.""" @@ -71,6 +72,7 @@ def create( _mnnvl_device_index, _mnnvl_workspace_all_succeeded, ) + from tensorrt_llm._utils import mpi_disabled use_fabric = mapping.is_multi_node() if fabric_handle is None else bool(fabric_handle) total = NUM_LAMPORT_BUFFERS * buffer_bytes @@ -87,6 +89,9 @@ def create( if free_bytes < total: problem = f"its {total} bytes exceed the {free_bytes} free on this rank's device" if not _mnnvl_workspace_all_succeeded(comm, problem is None): + # Every rank takes this path: free the MPI communicator split above (a ProcessGroup is c10d's). + if not mpi_disabled(): + comm.Free() raise RuntimeError( f"MnnvlWorkspace.create: not every rank can allocate ({problem or 'another rank cannot'})" ) @@ -108,6 +113,10 @@ def create( except Exception as exc: # noqa: BLE001 -- reported to every rank below, then re-raised error = exc if not _mnnvl_workspace_all_succeeded(comm, error is None): + # Every rank takes this path too. A handle only borrows the communicator and makes no MPI call when + # destroyed, so it may outlive the communicator. + if not mpi_disabled(): + comm.Free() raise RuntimeError("MnnvlWorkspace: allocation failed on at least one rank") from error # Arms every buffer word and the flags; also the barrier after which a peer may push into this rank. _initialize_allreduce_mnnvl_protocol( diff --git a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allreduce_attn_res_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allreduce_attn_res_op_matrix.py index 536aecb2d7e9..fd58b8d10d0c 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allreduce_attn_res_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allreduce_attn_res_op_matrix.py @@ -101,14 +101,23 @@ def check_workspace_is_armed_and_sized() -> None: def check_create_refuses_on_every_rank() -> None: - """One rank asks for three buffers past its device's free memory: every rank raises before any allocates, and the - workspaces in use stay correct.""" + """One rank asks for three buffers past its device's free memory: every rank raises before any allocates and frees + the communicator it split, and the workspaces in use stay correct.""" from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.mnnvl_workspace import ( MnnvlWorkspace, ) + from tensorrt_llm._torch.distributed import ops + + split = ops._get_mnnvl_workspace_comm + comms = [] + + def recording_split(mapping): + comms.append(split(mapping)) + return comms[-1] free_bytes, _ = torch.cuda.mem_get_info() too_big = (free_bytes // 3 // 16 + (64 << 20) // 16) * 16 + ops._get_mnnvl_workspace_comm = recording_split try: MnnvlWorkspace.create( R.mapping, too_big if R.rank == 0 else BUFFER_BYTES, fabric_handle=R.fabric @@ -116,7 +125,11 @@ def check_create_refuses_on_every_rank() -> None: raised = False except RuntimeError as exc: raised = "not every rank can allocate" in str(exc) + finally: + ops._get_mnnvl_workspace_comm = split assert R.all_true(raised), "a rank short of memory did not make every rank raise" + freed = len(comms) == 1 and comms[0] == R.MPI.COMM_NULL + assert R.all_true(freed), "a refused create kept the communicator it split" call = Call(2100, 8, 3) verify(call, call.run(WS_A), "after the refused create") From 2f444c96d34be915b97adf30e1dca1247b9dbce0 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 20:55:03 -0700 Subject: [PATCH 051/161] [None][fix] modeling_v2 Kimi K3 target: no measurement date in a copied comment The MoE runtime copied from modeling_kimi_linear.py dated a measurement in a comment ("Measured 2026-09-08"), which test_modeling_v2_no_stale_claims rejects anywhere under modeling_v2. The comment now says it was measured, not when. Comments are not part of the AST, so the copy stays AST-identical to the built-in definition. Signed-off-by: Vasanth Sabavat --- .../kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index d513e2fb7e24..0b2d741361bf 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -877,7 +877,7 @@ def __init__( # as if it were the backend that was asked for, and the decline is # easy to trigger: MegaMoE has its own token / top-k limits and is # EP-only, and CuteDSL declines on activation shape, SM version and - # the CuTe DSL dependency. Measured 2026-09-08: a CUTEDSL request + # the CuTe DSL dependency. As measured once: a CUTEDSL request # was turned down on every one of the 92 MoE layers, on all 16 # ranks, and still produced correct text and a zero exit -- the # only trace was a warning line per layer. Fail in the resolver From 549ffd4198c2dc2ec15dfaa9b4677742365f0a42 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 20:55:26 -0700 Subject: [PATCH 052/161] [None][fix] modeling_v2 Kimi K3 target: register trtllm::kda_mtp_decode on import REQUIRED_TRTLLM_OPS names kda_mtp_decode, the built-in KDA module's fused verify, but that op is a CuTe DSL op the module registers only on its first verify (_kda_kernels loads custom_ops/cute_dsl_kimi_k3_kda_mtp_ops lazily). On a build of main, test_modeling_v2_target_contract therefore reported it as missing. The target now imports the registering module, and UNCERTIFIED_GENERIC_CALLS declares it as generic-path code. Signed-off-by: Vasanth Sabavat --- .../kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index 0b2d741361bf..3cee67b5f481 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -85,6 +85,9 @@ from tensorrt_llm._torch._experimental.modeling_v2.catalog.ssm.k3_kda_verify import k3_kda_verify from tensorrt_llm._torch.attention.backends import AttentionMetadata from tensorrt_llm._torch.attention.backends.fmha.cute_dsl_mla import k3_mla_decode_view + +# Registers trtllm::kda_mtp_decode (REQUIRED_TRTLLM_OPS), which the built-in KDA module loads on its first verify. +from tensorrt_llm._torch.custom_ops import cute_dsl_kimi_k3_kda_mtp_ops # noqa: F401 from tensorrt_llm._torch.distributed import AllReduce from tensorrt_llm._torch.model_config import ModelConfig from tensorrt_llm._torch.models.modeling_kimi_linear import KimiLinearForCausalLM @@ -175,6 +178,7 @@ "tensorrt_llm._torch.models.modeling_utils.DecoderModel", # The text model's stock modules. "tensorrt_llm._torch.modules.kimi_kda.KimiKDALinearAttention", + "tensorrt_llm._torch.custom_ops.cute_dsl_kimi_k3_kda_mtp_ops", "tensorrt_llm._torch.modules.kimi_k3_mla.KimiK3MLAAttention", "tensorrt_llm._torch.moe.fused_moe.create_moe", "tensorrt_llm._torch.moe.fused_moe.ConfigurableMoE", From f01a46b67fe8e11c785fbdae0325334b6f6de169 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 21:13:20 -0700 Subject: [PATCH 053/161] [None][chore] modeling_v2 Kimi K3 target: trim the text model to the target's settings The target's copy of the built-in Kimi K3 text model kept branches that no deployment of this target reaches. They are removed: - attention data parallelism: the router's upcast default, the replicated shared expert and dense MLP, and the MoE-TP refusal under attention DP. The topology assert sets enable_attention_dp false. - helix context parallelism: KimiMLARuntime's mapping_with_cp. The shell never sets up helix mappings, and world 16 = tp 16 x pp 1. - the automatic MoE split and its checks: _select_moe_tp_ep's EP-only default, the split-product and DWDP refusals. The topology assert fixes moe_tp 4 x moe_ep 4, which Mapping produces only from an explicit split. - the NVFP4 routed-expert layout: its loader, its spec entry and the per-layer quantization lookup. A new construction assert requires the MXFP4 checkpoint's quantization, with no quant algo or per-layer declaration in the model config. Routing already requires no quant algo. Each changed definition equals the old one with those settings substituted and constant-folded, so the dead branches and their locals drop out. The removed definitions are referenced nowhere. Signed-off-by: Vasanth Sabavat --- .../modeling.py | 172 ++++-------------- 1 file changed, 40 insertions(+), 132 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index 3cee67b5f481..dc8a27202f56 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -46,7 +46,8 @@ experts, the sandwiches and the residual epilogues come with their own entries; until then they run the generic path on every step. -**What this target asserts rather than adapts**: SM 10.0; the topology above; bf16 weights and a bf16 KV pool; +**What this target asserts rather than adapts**: SM 10.0; the topology above; the MXFP4 checkpoint's quantization (none +the model config reads, so the routed experts keep the MXFP4 default); bf16 weights and a bf16 KV pool; tokens_per_block 64 (the MLA generation kernels K3's 96 heads reach exist only at 64); the V2 hybrid KV / state manager, which holds the KDA states, with block reuse off and fp32 recurrent states; an all-reduce strategy of AUTO or MNNVL. The construction-time ones fail in `__init__`, the per-engine ones on the first forward, each naming the @@ -711,19 +712,12 @@ def _apply_attn_res_add_and_rmsnorm( ) -# Routed-expert key spellings that ModelOpt emits for Kimi K3. The NVFP4 -# checkpoint (``nvidia/Kimi-K3-NVFP4``) lists every prefix x module-name -# combination in ``quantized_layers``, so a lookup over this product finds it -# without needing the MiniMax-M3-style prefix normalization in ``ModelConfig``. -_K3_ROUTED_EXPERT_KEY_PREFIXES = ("language_model.model.", "model.", "") - - _K3_ROUTED_EXPERT_KEY_SUFFIXES = ("block_sparse_moe.experts", "mlp.experts") -# The subset of the above that can be a real module path. ``exclude_modules`` -# matches with wildcards and walks ancestor prefixes, so an empty prefix would -# widen what matches instead of just missing, as it does in the dict lookup. +# The routed experts' module path prefixes. ``exclude_modules`` matches with +# wildcards and walks ancestor prefixes, so an empty prefix would widen what +# matches instead of just missing. _K3_ROUTED_EXPERT_MODULE_PREFIXES = ("language_model.model.", "model.") @@ -748,19 +742,6 @@ def _load_packed_mxfp4_expert(backend, base, expert_idx, local_slot_id, get_tens ) -def _load_nvfp4_expert(backend, base, expert_idx, local_slot_id, get_tensor) -> None: - backend.quant_method.load_streaming_nvfp4_expert( - backend, - global_expert_id=expert_idx, - local_slot_id=local_slot_id, - **{ - f"{w}_{kind}": get_tensor(f"{base}.{expert_idx}.{w}.{kind}") - for w in ("w1", "w2", "w3") - for kind in ("weight", "weight_scale", "weight_scale_2", "input_scale") - }, - ) - - class _K3ExpertCkptSpec(NamedTuple): """How one routed-expert quantization is spelled and loaded.""" @@ -769,8 +750,7 @@ class _K3ExpertCkptSpec(NamedTuple): loader: Callable[..., None] # Set of filled slots the loader maintains, checked after the load. loaded_slots_attr: str - # NVFP4 defers cat/pad/interleave and the alpha computation to - # ``process_weights_after_loading``; the MXFP4 loaders write through. + # Whether the layer is finalized after its experts load (the MXFP4 loaders write through). needs_layer_finalize: bool @@ -781,12 +761,6 @@ class _K3ExpertCkptSpec(NamedTuple): loaded_slots_attr="_packed_mxfp4_loaded_slots", needs_layer_finalize=False, ), - QuantAlgo.NVFP4: _K3ExpertCkptSpec( - kinds=("weight", "weight_scale", "weight_scale_2", "input_scale"), - loader=_load_nvfp4_expert, - loaded_slots_attr="_streamed_expert_slots", - needs_layer_finalize=True, - ), } @@ -837,20 +811,12 @@ def __init__( situ_beta, situ_linear_beta = _resolve_kimi_situ_betas(cfg) dtype = torch.bfloat16 - # Routing scores stay fp32; with attention-DP off the gate GEMM runs - # bf16xbf16 with fp32 accumulate/output (checkpoint stores the gate - # weight in bf16; saves a per-layer input cast + fp32 splitK pair on - # the bs1 decode path). Under attention-DP the legacy upcast-to-fp32 - # GEMM is kept: the bf16-input min-latency GEMM's different reduction - # order flips borderline top-16 picks (GSM8K 96.7 -> 96.1/96.4, - # 3-run bisect on 62b20dd868), and the bs1-latency win is irrelevant - # at DEP batch sizes. KIMI_K3_ROUTER_BF16=1/0 forces either path. + # Routing scores stay fp32; the gate GEMM runs bf16xbf16 with fp32 + # accumulate/output (checkpoint stores the gate weight in bf16; saves a + # per-layer input cast + fp32 splitK pair on the bs1 decode path). + # KIMI_K3_ROUTER_BF16=0 forces the upcast-to-fp32 GEMM. _router_bf16_env = os.environ.get("KIMI_K3_ROUTER_BF16") - _router_bf16 = ( - _router_bf16_env == "1" - if _router_bf16_env is not None - else not model_config.mapping.enable_attention_dp - ) + _router_bf16 = _router_bf16_env == "1" if _router_bf16_env is not None else True self.gate = KimiK3MoEGate(cfg, logits_gemm_dtype=torch.bfloat16 if _router_bf16 else None) routed_moe_model_config = self._routed_moe_model_config(model_config) @@ -920,14 +886,11 @@ def __init__( self.expert_hi = self.expert_lo + self.experts_per_rank shared_intermediate = cfg.moe_intermediate_size * cfg.num_shared_experts - attention_dp = model_config.mapping.enable_attention_dp shared_model_config = copy.copy(model_config) shared_model_config.quant_config = QuantConfig() - # Under attention DP each rank owns different tokens, so the shared - # expert is replicated (TP size 1) and must not reduce across ranks. # Direct MoE-TP leaves both branches as partials for one concatenated # all-reduce. - use_shared_tp = not attention_dp and model_config.mapping.tp_size > 1 + use_shared_tp = model_config.mapping.tp_size > 1 self._reduce_routed_output = ( use_shared_tp and self.routed_experts.backend.scheduler_kind != MoESchedulerKind.FUSED_COMM @@ -948,7 +911,6 @@ def __init__( ), dtype=dtype, config=shared_model_config, - overridden_tp_size=1 if attention_dp else None, reduce_output=use_shared_tp, layer_idx=layer_idx, is_shared_expert=True, @@ -982,38 +944,21 @@ def _routed_projection(hidden_states: torch.Tensor, projection: nn.Module) -> to @staticmethod def _select_moe_tp_ep(mapping: Mapping) -> Tuple[int, int]: - """Resolve the routed-expert ``(moe_tp, moe_ep)`` split. - - Precedence: - - 1. Explicit ``moe_tensor_parallel_size`` / ``moe_expert_parallel_size`` - from the user config. Detected via - ``mapping.moe_tp_ep_user_specified`` so the auto-resolved mapping - default (``moe_tp=tp_size, moe_ep=1``) is NOT mistaken for a TP - request. - 2. Default: EP-only (``moe_tp=1, moe_ep=tp_size``), the historical - K3 layout. - """ - tp_size = mapping.tp_size - if getattr(mapping, "moe_tp_ep_user_specified", False): - return mapping.moe_tp_size, mapping.moe_ep_size - return 1, tp_size + """The routed-expert ``(moe_tp, moe_ep)`` split: the user config's explicit + ``moe_tensor_parallel_size`` / ``moe_expert_parallel_size`` (4 x 4, asserted at + construction).""" + return mapping.moe_tp_size, mapping.moe_ep_size @staticmethod def _resolve_routed_quant_config(model_config: ModelConfig, layer_idx: int) -> QuantConfig: - """Routed-expert quantization for ``layer_idx``, taken from the checkpoint. - - ``nvidia/Kimi-K3-NVFP4`` declares the routed experts per layer as - ``NVFP4`` with ``group_size=16``; the original ``moonshotai/Kimi-K3`` - declares nothing per layer and keeps the historical - ``W4A8_MXFP4_MXFP8`` default. Reading the checkpoint instead of - hardcoding is what lets one code path serve both. - - An exclusion outranks the per-layer entry and the default below: - ``create_weights`` treats an override as authoritative over anything - ``__post_init__`` wrote, so this return value stands in for both - quantization passes and exclusion is the one that runs second. It is - matched as a pattern, so it is asked only about real module names. + """Routed-expert quantization for ``layer_idx``: the MXFP4 checkpoint declares nothing per layer (asserted at + construction), so the experts keep the ``W4A8_MXFP4_MXFP8`` default. + + An exclusion outranks the default: ``create_weights`` treats an override + as authoritative over anything ``__post_init__`` wrote, so this return + value stands in for both quantization passes and exclusion is the one + that runs second. It is matched as a pattern, so it is asked only about + real module names. """ quant_config = model_config.quant_config if quant_config is not None and any( @@ -1030,22 +975,6 @@ def _resolve_routed_quant_config(model_config: ModelConfig, layer_idx: int) -> Q ) return QuantConfig(kv_cache_quant_algo=quant_config.kv_cache_quant_algo) - per_layer = getattr(model_config, "quant_config_dict", None) - if per_layer: - for prefix in _K3_ROUTED_EXPERT_KEY_PREFIXES: - for suffix in _K3_ROUTED_EXPERT_KEY_SUFFIXES: - cfg = per_layer.get(f"{prefix}layers.{layer_idx}.{suffix}") - if cfg is not None and cfg.quant_algo is not None: - # Logged once per layer: the routed-expert format decides - # which MoE backends can serve this checkpoint at all. - logger.debug( - "Kimi K3 layer %d routed experts: %s (group_size=%s) " - "from the checkpoint", - layer_idx, - cfg.quant_algo, - cfg.group_size, - ) - return cfg logger.debug( "Kimi K3 layer %d routed experts: no per-layer quant config in the " "checkpoint, defaulting to %s", @@ -1095,7 +1024,7 @@ def _check_trtllm_situ_quant(moe_backend: str, quant_algo: Optional[QuantAlgo]) @staticmethod def _routed_moe_model_config(model_config: ModelConfig) -> ModelConfig: """Build a private routed-expert mapping without mutating the shared - config. Default split is EP-only; see ``_select_moe_tp_ep``.""" + config, with the split of ``_select_moe_tp_ep``.""" # Every backend here declares ``ActivationType.SiTu`` in its # ``activation_support``; the list is not a preference order. CUTEDSL # joined once its act-fusion kernel grew the SiTU epilogue. @@ -1118,21 +1047,8 @@ def _routed_moe_model_config(model_config: ModelConfig) -> ModelConfig: "EPLB or replicated expert slots." ) mapping = model_config.mapping - if getattr(mapping, "_dwdp_size", 0) > 1: - raise NotImplementedError("Kimi K3 packed-checkpoint streaming does not support DWDP.") moe_tp, moe_ep = KimiK3MoERuntime._select_moe_tp_ep(mapping) - if moe_tp < 1 or moe_ep < 1 or moe_tp * moe_ep != mapping.tp_size: - raise ValueError( - f"Kimi K3 routed MoE split moe_tp={moe_tp} x moe_ep={moe_ep} " - f"must multiply to tp_size={mapping.tp_size}." - ) - if moe_tp > 1 and mapping.enable_attention_dp: - raise NotImplementedError( - "Kimi K3 MoE tensor parallelism requires " - "enable_attention_dp=false (the attention-DP dispatch/combine " - "path is validated for EP-only splits)." - ) logger.info_once( f"Kimi K3 routed MoE parallelism: moe_tp={moe_tp}, " f"moe_ep={moe_ep} (tp_size={mapping.tp_size})", @@ -1273,7 +1189,6 @@ def __init__( layer_idx: int, model_config: ModelConfig, aux_stream_dict: Dict[AuxStreamType, torch.cuda.Stream], - mapping_with_cp: Optional[Mapping] = None, ) -> None: super().__init__() @@ -1287,11 +1202,8 @@ def __init__( # KimiK3MLAAttention owns MLA projection/head sharding. Keep only the # final output reduction in this wrapper so the output gate remains # between attention and the row-parallel o_proj. - # Helix: mapping_with_cp (the CP original) activates the base MLA's - # helix machinery; this wrapper's allreduce over the repurposed - # mapping sums the base o_proj's tp*cp partials. mapping = model_config.mapping - reduce_output = not mapping.enable_attention_dp and mapping.tp_size > 1 + reduce_output = mapping.tp_size > 1 self._o_allreduce = ( AllReduce( mapping=mapping, @@ -1335,7 +1247,6 @@ def __init__( max_position_embeddings=max_positions, model_config=attention_config, aux_stream_dict=aux_stream_dict, - mapping_with_cp=mapping_with_cp, ) def forward( @@ -1394,8 +1305,6 @@ def __init__( layer_idx, model_config=model_config, aux_stream_dict=aux_stream_dict, - # CP original stashed by _setup_helix_mappings; None outside helix. - mapping_with_cp=getattr(model_config, "_helix_mapping_with_cp", None), ) self.is_moe = ( @@ -1408,24 +1317,17 @@ def __init__( else: situ_beta = getattr(cfg, "activation_situ_beta", None) or 1.0 situ_linear_beta = getattr(cfg, "activation_situ_linear_beta", None) - attention_dp = model_config.mapping.enable_attention_dp - if attention_dp: - self.mlp_tp_size = 1 - else: - self.mlp_tp_size = math.gcd(cfg.intermediate_size, model_config.mapping.tp_size) - # Over MNNVL (one NVLink domain across the nodes, where a cross-node all-reduce costs what a node's - # does) the MLP stays split over the whole TP group, the per-rank shapes its decode GEMVs take - # (decode_gemv.SITES); otherwise it stays within one node. - spans_nodes = self._mnnvl_allreduce() is not None - if self.mlp_tp_size > model_config.mapping.gpus_per_node and not spans_nodes: - self.mlp_tp_size = math.gcd( - self.mlp_tp_size, model_config.mapping.gpus_per_node - ) + self.mlp_tp_size = math.gcd(cfg.intermediate_size, model_config.mapping.tp_size) + # Over MNNVL (one NVLink domain across the nodes, where a cross-node all-reduce costs what a node's + # does) the MLP stays split over the whole TP group, the per-rank shapes its decode GEMVs take + # (decode_gemv.SITES); otherwise it stays within one node. + spans_nodes = self._mnnvl_allreduce() is not None + if self.mlp_tp_size > model_config.mapping.gpus_per_node and not spans_nodes: + self.mlp_tp_size = math.gcd(self.mlp_tp_size, model_config.mapping.gpus_per_node) mlp_model_config = copy.copy(model_config) mlp_model_config.quant_config = QuantConfig() # K3's dense layer is BF16, so a unit block size gives the same - # subgroup selection as DeepSeek-V3. Attention DP replicates the - # MLP because ranks own different tokens. + # subgroup selection as DeepSeek-V3. self.mlp = GatedMLP( hidden_size=cfg.hidden_size, intermediate_size=cfg.intermediate_size, @@ -2255,6 +2157,12 @@ def _check_construction(model_config: ModelConfig) -> None: assert kv_algo is None, ( f"this target's MLA kernels read a bf16 KV pool; kv_cache_config.dtype resolved to {kv_algo}" ) + quant_algo = model_config.quant_config.quant_algo + assert quant_algo is None and not model_config.quant_config_dict, ( + "this target loads the MXFP4 checkpoint, which declares no quantization the model config reads (its routed " + f"experts keep the W4A8_MXFP4_MXFP8 default); the engine read {quant_algo} with " + f"{len(model_config.quant_config_dict or {})} per-layer declarations" + ) strategy = model_config.allreduce_strategy assert strategy in (AllReduceStrategy.AUTO, AllReduceStrategy.MNNVL), ( f"this target runs its all-reduces over MNNVL; allreduce_strategy is {strategy.name}" From 992679301283986cd82065738c97691fd75aaad2 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 21:18:42 -0700 Subject: [PATCH 054/161] [None][fix] MnnvlWorkspace.create involves only the TP group's ranks create() made its communicator with a split of the MPI session, which every rank of the session has to enter together. With PP or CP > 1 the ranks of another TP group then had to call create() at the same point. _get_mnnvl_tp_group_comm (distributed/ops.py) makes a communicator of exactly mapping.tp_group with MPI_Comm_create_group, so only the group's ranks take part. Its rank i is TP rank i. Under Ray it is the TP ProcessGroup, as before. The matrix gains check_create_involves_only_the_tp_group. Under TP W/2 x PP 2, the first group creates a workspace while the other group's ranks make no MNNVL call; the workspace is armed and sized for the group. The refusal check records the new helper's communicator. Signed-off-by: Vasanth Sabavat --- .../catalog/comm/mnnvl_allreduce_attn_res.md | 8 +-- .../catalog/comm/mnnvl_workspace.py | 13 ++--- tensorrt_llm/_torch/distributed/ops.py | 18 +++++++ .../_mnnvl_allreduce_attn_res_op_matrix.py | 52 ++++++++++++++++--- 4 files changed, 75 insertions(+), 16 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.md index 61a3faeffdd6..3ac74f6100aa 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allreduce_attn_res.md @@ -80,12 +80,14 @@ a call over one buffer raises (below). **Who creates it, and when.** The target, in `post_load_weights`, with `MnnvlWorkspace.create(mapping, buffer_bytes, fabric_handle=None)`: -- collective over `mapping`'s TP group: every rank calls it at the same point; +- collective over `mapping`'s TP group only: every rank of the group calls it at the same point, and ranks outside + the group take no part (certified: under TP W/2 x PP 2, one group creates a workspace while the other makes no + MNNVL call); - failure model: - before allocating, the ranks agree that each of them can (not capturing, a valid `buffer_bytes`, the three buffers within that rank's free device memory). If one cannot, every rank raises `RuntimeError`, none - allocates, and under MPI each frees the communicator it split for the call (certified: one rank short of - memory, every rank raises and frees its split, the next call is correct); + allocates, and under MPI each frees the communicator it made for the call (certified: one rank short of + memory, every rank raises and frees that communicator, the next call is correct); - a failure that returns from the allocation is agreed and handled the same way; - a rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange; this is not turned into an error on the other ranks; diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_workspace.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_workspace.py index e9de6d592f7a..b3124bdd0bf7 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_workspace.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_workspace.py @@ -53,12 +53,12 @@ def max_one_shot_tokens(self, hidden: int, dtype: torch.dtype = torch.bfloat16) def create( cls, mapping, buffer_bytes: int, fabric_handle: Optional[bool] = None ) -> "MnnvlWorkspace": - """Allocate and arm a workspace for ``mapping``'s TP group. Collective: every rank of the group calls it at - the same point, eagerly (not under CUDA-graph capture). + """Allocate and arm a workspace for ``mapping``'s TP group. Collective over that group only: every rank of + the group calls it at the same point, eagerly (not under CUDA-graph capture); ranks outside it take no part. Failure model: before allocating, the ranks agree that each of them can (not capturing, a valid ``buffer_bytes``, the three buffers within its device's free memory); if one cannot, every rank raises - ``RuntimeError``, none allocates, and under MPI each frees the communicator it split for the call. A failure + ``RuntimeError``, none allocates, and under MPI each frees the communicator it made for the call. A failure that returns from the allocation is agreed and handled the same way. A rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange: that failure is not turned into an error on the other ranks. @@ -66,7 +66,7 @@ def create( ``fabric_handle``: share the memory by fabric handle (required across nodes) rather than POSIX file descriptor; default ``mapping.is_multi_node()``.""" from tensorrt_llm._torch.distributed.ops import ( - _get_mnnvl_workspace_comm, + _get_mnnvl_tp_group_comm, _initialize_allreduce_mnnvl_protocol, _make_mnnvl_mcast_buffer, _mnnvl_device_index, @@ -76,7 +76,7 @@ def create( use_fabric = mapping.is_multi_node() if fabric_handle is None else bool(fabric_handle) total = NUM_LAMPORT_BUFFERS * buffer_bytes - comm = _get_mnnvl_workspace_comm(mapping) + comm = _get_mnnvl_tp_group_comm(mapping) # Every condition one rank alone can fail is checked before the allocation, and the ranks agree on it: a # rank failing inside the allocation would leave its peers in the handle exchange. problem: Optional[str] = None @@ -89,7 +89,8 @@ def create( if free_bytes < total: problem = f"its {total} bytes exceed the {free_bytes} free on this rank's device" if not _mnnvl_workspace_all_succeeded(comm, problem is None): - # Every rank takes this path: free the MPI communicator split above (a ProcessGroup is c10d's). + # Every rank of the group takes this path: free the MPI communicator made above (a ProcessGroup is + # c10d's). if not mpi_disabled(): comm.Free() raise RuntimeError( diff --git a/tensorrt_llm/_torch/distributed/ops.py b/tensorrt_llm/_torch/distributed/ops.py index 4a4ec2fc5a81..c7a61101d920 100644 --- a/tensorrt_llm/_torch/distributed/ops.py +++ b/tensorrt_llm/_torch/distributed/ops.py @@ -183,6 +183,24 @@ def _get_mnnvl_workspace_comm(mapping: Mapping): mapping.tp_rank) +def _get_mnnvl_tp_group_comm(mapping: Mapping): + """A new communicator of exactly mapping.tp_group (its rank i = TP rank i) for a caller-owned + MNNVL state's create(); only the group's ranks take part (MPI_Comm_create_group). The caller + frees it. Under Ray the TP ProcessGroup (c10d's). + """ + if mpi_disabled(): + pg = mapping.tp_group_pg + assert pg is not None, "TP ProcessGroup not initialised" + return pg + session = mpi_comm() + session_group = session.Get_group() + group = session_group.Incl(mapping.tp_group) + session_group.Free() + comm = session.Create_group(group) + group.Free() + return comm + + def _mnnvl_device_index(mapping: Mapping) -> int: """CUDA device index backing this rank's MNNVL buffers. diff --git a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allreduce_attn_res_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allreduce_attn_res_op_matrix.py index fd58b8d10d0c..7b20ecad466e 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allreduce_attn_res_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allreduce_attn_res_op_matrix.py @@ -102,22 +102,22 @@ def check_workspace_is_armed_and_sized() -> None: def check_create_refuses_on_every_rank() -> None: """One rank asks for three buffers past its device's free memory: every rank raises before any allocates and frees - the communicator it split, and the workspaces in use stay correct.""" + the communicator it made, and the workspaces in use stay correct.""" from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.mnnvl_workspace import ( MnnvlWorkspace, ) from tensorrt_llm._torch.distributed import ops - split = ops._get_mnnvl_workspace_comm + make = ops._get_mnnvl_tp_group_comm comms = [] - def recording_split(mapping): - comms.append(split(mapping)) + def recording_make(mapping): + comms.append(make(mapping)) return comms[-1] free_bytes, _ = torch.cuda.mem_get_info() too_big = (free_bytes // 3 // 16 + (64 << 20) // 16) * 16 - ops._get_mnnvl_workspace_comm = recording_split + ops._get_mnnvl_tp_group_comm = recording_make try: MnnvlWorkspace.create( R.mapping, too_big if R.rank == 0 else BUFFER_BYTES, fabric_handle=R.fabric @@ -126,14 +126,51 @@ def recording_split(mapping): except RuntimeError as exc: raised = "not every rank can allocate" in str(exc) finally: - ops._get_mnnvl_workspace_comm = split + ops._get_mnnvl_tp_group_comm = make assert R.all_true(raised), "a rank short of memory did not make every rank raise" freed = len(comms) == 1 and comms[0] == R.MPI.COMM_NULL - assert R.all_true(freed), "a refused create kept the communicator it split" + assert R.all_true(freed), "a refused create kept the communicator it made" call = Call(2100, 8, 3) verify(call, call.run(WS_A), "after the refused create") +def check_create_involves_only_the_tp_group() -> None: + """Under TP W/2 x PP 2 only the first TP group's ranks create a workspace, while the other group's ranks make no + MNNVL call: create() completes on the group without them, armed and sized for the group.""" + from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.mnnvl_workspace import ( + MnnvlWorkspace, + ) + from tensorrt_llm.mapping import Mapping + + if R.world < 4 or R.world % 2: + if R.rank == 0: + print( + f"[rank 0] check_create_involves_only_the_tp_group: not run at world {R.world}", + flush=True, + ) + return + half = R.world // 2 + mapping = Mapping( + world_size=R.world, + rank=R.rank, + gpus_per_node=torch.cuda.device_count(), + tp_size=half, + pp_size=2, + ) + ok = True + if mapping.pp_rank == 0: + ws = MnnvlWorkspace.create(mapping, BUFFER_BYTES, fabric_handle=R.fabric) + armed = ws.lamport.view(torch.int32) + ok = ( + ws.world_size == half + and ws.rank == mapping.tp_rank + and ws.comm.Get_size() == half + and bool((armed == torch.tensor(-(2**31), dtype=torch.int32, device="cuda")).all()) + ) + ws.comm.Free() + assert R.all_true(ok), "a TP group's create() did not complete without the other group's ranks" + + def check_single_calls() -> None: for t in TOKENS: for s in SNAPSHOTS: @@ -256,6 +293,7 @@ def check_wrong_call_order_is_detected() -> None: CHECKS = [ check_workspace_is_armed_and_sized, check_create_refuses_on_every_rank, + check_create_involves_only_the_tp_group, check_single_calls, check_a_call_over_one_buffer_raises_on_every_rank, check_dip_and_regrow_sequence, From 3691863da290a60ff8d4e04ffe74c7bfee60be01 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 20:21:05 -0700 Subject: [PATCH 055/161] [None][feat] MNNVL: one-shot size per call, early PDL trigger; split all-gather - trtllm::mnnvl_fusion_allreduce takes one_shot_max_bytes: messages up to that many bytes go one-shot, larger ones two-shot. The default is 1 MiB, main's threshold. MNNVLAllReduce keeps the value as an attribute, sizes its workspace with it, and forward() takes a per-call override. Kimi K3 sends its decode-size all-reduces one-shot up to 4 MiB. - The one-shot fusion kernel triggers its programmatic dependents right after its grid-dependency wait instead of after the reduction. A dependent GEMV can then stream its weights while this kernel waits on the other ranks. Dependents still read the output and the Lamport flags only after their own grid wait. This applies to every caller of the one-shot kernel. - The one-shot kernel's Lamport reduction is now the shared reduceOneshotLamport, which the attention-residual one-shot already uses. The code is moved, not changed: the reduction order and every result are unchanged, but the kernel's SASS changes. - trtllm::mnnvl_allgather_split (mnnvlAllGatherKernels.{h,cu}, exposed as MNNVLAllReduce.allgather_split) is a one-shot all-gather of fp32 rows. The leading columns travel and arrive as bf16, the rest as fp32. It runs on the all-reduce's workspace and takes one turn of its Lamport rotation. Kimi K3 gathers its sharded MoE head with it: the latent columns in bf16, the router logits in fp32. Its schema declares comm_buffer and buffer_flags as mutable. Signed-off-by: Vasanth Sabavat --- .../mnnvlAllGatherKernels.cu | 165 ++++++++++++++++++ .../mnnvlAllGatherKernels.h | 66 +++++++ .../mnnvlAllreduceKernels.cu | 62 ++----- cpp/tensorrt_llm/thop/allreduceOp.cpp | 79 ++++++++- .../_torch/custom_ops/cpp_custom_ops.py | 3 +- tensorrt_llm/_torch/distributed/ops.py | 46 ++++- 6 files changed, 363 insertions(+), 58 deletions(-) create mode 100644 cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllGatherKernels.cu create mode 100644 cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllGatherKernels.h diff --git a/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllGatherKernels.cu b/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllGatherKernels.cu new file mode 100644 index 000000000000..e43fe8c4378b --- /dev/null +++ b/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllGatherKernels.cu @@ -0,0 +1,165 @@ +/* + * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * 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. + */ + +#include "tensorrt_llm/kernels/communicationKernels/mnnvlAllGatherKernels.h" + +#include "tensorrt_llm/common/cudaUtils.h" +#include "tensorrt_llm/common/envUtils.h" +#include "tensorrt_llm/common/lamportUtils.cuh" + +TRTLLM_NAMESPACE_BEGIN + +namespace kernels::mnnvl +{ + +using tensorrt_llm::common::isLamportDirty; +using tensorrt_llm::common::LamportFlags; +using tensorrt_llm::common::loadPackedVolatile; +using tensorrt_llm::common::VolatilePackedLoad; + +namespace +{ + +constexpr int kThreads = 256; + +// Payload vectors (16 bytes) of one rank's slot for one token: the bf16 part, then the fp32 part. +__host__ __device__ inline int slotVectors(int bf16Columns, int fp32Columns) +{ + return bf16Columns / 8 + fp32Columns / 4; +} + +// No payload word may equal the Lamport sentinel (the fp32 -0.0 word): -0.0 halves of a bf16 +// pair and -0.0 fp32 words become +0.0. +__device__ inline uint32_t sanitizeBf16Pair(uint32_t word) +{ + if ((word & 0xffffu) == 0x8000u) + { + word &= 0xffff0000u; + } + if ((word >> 16) == 0x8000u) + { + word &= 0x0000ffffu; + } + return word; +} + +__device__ inline uint32_t sanitizeFp32(uint32_t word) +{ + return word == 0x80000000u ? 0u : word; +} + +__device__ inline uint32_t packBf16Pair(float lo, float hi) +{ + __nv_bfloat162 const pair = __floats2bfloat162_rn(lo, hi); + return sanitizeBf16Pair(*reinterpret_cast(&pair)); +} + +// One CTA per token: broadcast this rank's slot of the token, then poll every rank's slot and +// scatter it to the outputs. +__global__ void __launch_bounds__(kThreads) mnnvlAllGatherSplitKernel(AllGatherSplitParams const p) +{ +#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900) + cudaGridDependencySynchronize(); + // Dependents read the outputs (and the Lamport flags) only after their own wait for this grid. + cudaTriggerProgrammaticLaunchCompletion(); +#endif + int const token = blockIdx.x; + int const bf16Vectors = p.bf16Columns / 8; + int const vectors = slotVectors(p.bf16Columns, p.fp32Columns); + int const inputColumns = p.bf16Columns + p.fp32Columns; + + LamportFlags flag(p.bufferFlags, 1); + auto* lamportMcast = reinterpret_cast(flag.getCurLamportBuf(p.multicastPtr, 0)); + auto* lamportLocal = reinterpret_cast(flag.getCurLamportBuf(p.bufferPtrsDev[p.rank], 0)); + + float const* row = p.input + static_cast(token) * inputColumns; + for (int v = threadIdx.x; v < vectors; v += blockDim.x) + { + uint4 packed; + if (v < bf16Vectors) + { + float4 const a = reinterpret_cast(row)[2 * v]; + float4 const b = reinterpret_cast(row)[2 * v + 1]; + packed = make_uint4( + packBf16Pair(a.x, a.y), packBf16Pair(a.z, a.w), packBf16Pair(b.x, b.y), packBf16Pair(b.z, b.w)); + } + else + { + uint4 const words = reinterpret_cast(row + p.bf16Columns)[v - bf16Vectors]; + packed = make_uint4( + sanitizeFp32(words.x), sanitizeFp32(words.y), sanitizeFp32(words.z), sanitizeFp32(words.w)); + } + lamportMcast[(static_cast(token) * p.nRanks + p.rank) * vectors + v] = packed; + } + + flag.ctaArrive(); + flag.clearDirtyLamportBuf(p.bufferPtrsDev[p.rank], -1); + + for (int i = threadIdx.x; i < p.nRanks * vectors; i += blockDim.x) + { + int const r = i / vectors; + int const v = i % vectors; + VolatilePackedLoad value; + do + { + value + = loadPackedVolatile(&lamportLocal[(static_cast(token) * p.nRanks + r) * vectors + v]); + } while (isLamportDirty(value)); + uint4 const words = make_uint4(value.words[0], value.words[1], value.words[2], value.words[3]); + if (v < bf16Vectors) + { + reinterpret_cast( + p.bf16Output + static_cast(token) * p.nRanks * p.bf16Columns + r * p.bf16Columns)[v] + = words; + } + else + { + reinterpret_cast(p.fp32Output + static_cast(token) * p.nRanks * p.fp32Columns + + r * p.fp32Columns)[v - bf16Vectors] + = words; + } + } + + flag.waitAndUpdate({static_cast(p.numTokens * p.nRanks * vectors * sizeof(uint4)), 0, 0, 0}); +} + +} // namespace + +int64_t mnnvlAllGatherSplitFootprint(int numTokens, int bf16Columns, int fp32Columns, int nRanks) +{ + return static_cast(numTokens) * nRanks * slotVectors(bf16Columns, fp32Columns) * sizeof(uint4); +} + +void mnnvlAllGatherSplitOp(AllGatherSplitParams const& params) +{ + TLLM_CHECK_WITH_INFO(params.bf16Columns % 8 == 0 && params.fp32Columns % 4 == 0, + "[mnnvlAllGatherSplit] needs bf16 columns in multiples of 8 and fp32 columns in multiples of 4"); + TLLM_CHECK_WITH_INFO(params.numTokens > 0, "[mnnvlAllGatherSplit] needs at least one token"); + cudaLaunchAttribute attrs[1]; + attrs[0].id = cudaLaunchAttributeProgrammaticStreamSerialization; + attrs[0].val.programmaticStreamSerializationAllowed = tensorrt_llm::common::getEnvEnablePDL() ? 1 : 0; + cudaLaunchConfig_t config{}; + config.gridDim = dim3(params.numTokens); + config.blockDim = dim3(kThreads); + config.stream = params.stream; + config.attrs = attrs; + config.numAttrs = 1; + TLLM_CUDA_CHECK(cudaLaunchKernelEx(&config, mnnvlAllGatherSplitKernel, params)); +} + +} // namespace kernels::mnnvl + +TRTLLM_NAMESPACE_END diff --git a/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllGatherKernels.h b/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllGatherKernels.h new file mode 100644 index 000000000000..b646e91897c7 --- /dev/null +++ b/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllGatherKernels.h @@ -0,0 +1,66 @@ +/* + * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * 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. + */ + +#pragma once + +#include "tensorrt_llm/common/config.h" + +#include +#include +#include + +TRTLLM_NAMESPACE_BEGIN + +namespace kernels::mnnvl +{ + +/** + * \brief Parameters of mnnvlAllGatherSplitOp: a one-shot all-gather over the MNNVL workspace of + * fp32 rows whose first bf16Columns columns travel, and are gathered, as bf16. + * + * Every rank contributes input [numTokens, bf16Columns + fp32Columns] (fp32). On every rank, + * bf16Output[t, r * bf16Columns + j] = bf16(input_r[t, j]) and + * fp32Output[t, r * fp32Columns + j] = input_r[t, bf16Columns + j], for every rank r. The bf16 + * rounding is round-to-nearest, as a GEMV storing bf16 would round. -0.0 arrives as +0.0. + * + * The kernel follows the one-shot all-reduce's Lamport protocol on the shared workspace (it takes + * one turn of the buffer rotation) and triggers its dependents as soon as it starts. + */ +struct AllGatherSplitParams +{ + float const* input; //!< [numTokens, bf16Columns + fp32Columns], this rank's slice + __nv_bfloat16* bf16Output; //!< [numTokens, nRanks * bf16Columns] + float* fp32Output; //!< [numTokens, nRanks * fp32Columns]; unused if fp32Columns == 0 + int numTokens; + int bf16Columns; //!< Multiple of 8 + int fp32Columns; //!< Multiple of 4 + + int nRanks; + int rank; + void** bufferPtrsDev; + void* multicastPtr; + uint32_t* bufferFlags; + cudaStream_t stream; +}; + +//! Bytes of one Lamport buffer the all-gather occupies. +int64_t mnnvlAllGatherSplitFootprint(int numTokens, int bf16Columns, int fp32Columns, int nRanks); + +void mnnvlAllGatherSplitOp(AllGatherSplitParams const& params); + +} // namespace kernels::mnnvl + +TRTLLM_NAMESPACE_END diff --git a/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllreduceKernels.cu b/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllreduceKernels.cu index 3b84bcfa759c..3d972b773908 100644 --- a/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllreduceKernels.cu +++ b/cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllreduceKernels.cu @@ -337,6 +337,9 @@ struct MnnvlAllReduceKernelParams uint32_t* bufferFlags; bool waitForResults; QuantizationSFLayout layout; + // One-shot kernel: trigger the dependent launch at the start instead of after the reduction, + // so a dependent GEMV streams its weights while this kernel waits on the other ranks. + bool earlyTrigger; }; template @@ -510,7 +513,6 @@ inline __device__ PackedVec reduceOneshotDeterministicFastPath( return reduceOneshotDeterministic(remoteValues, localValue); } -// The one-shot Lamport reduction of oneshotAllreduceFusionKernel as a function, for oneshotAllreduceAttnResKernel. // Fully deterministic: every rank uses the exact same reduction order. For WorldSize <= 8, specialize the local // slot so the fast path reuses `val` from registers without a dynamic `remoteValues[rank]` store. Larger world sizes // use the compact fallback because the benefit is thin but specializing every rank significantly increases compile @@ -625,7 +627,6 @@ using detail::copyF4; using detail::MnnvlAllReduceKernelParams; using detail::sanitizeLamportPayload; using detail::accumulateLamportRanksChunked; -using detail::reduceOneshotDeterministicFastPath; using detail::reduceOneshotLamport; using detail::writeEpilogueOutput; @@ -644,6 +645,13 @@ __global__ void __launch_bounds__(1024) oneshotAllreduceFusionKernel(MnnvlAllRed int threadOffset = token * params.tokenDim + packedIdx * kELTS_PER_THREAD; cudaGridDependencySynchronize(); + if (params.earlyTrigger) + { + // Dependents read our output and the Lamport flags only after their own + // griddepcontrol.wait, i.e. after this grid completes; launching them now only lets + // them start streaming weights. + cudaTriggerProgrammaticLaunchCompletion(); + } #else int packedIdx = blockIdx.y * blockDim.x + threadIdx.x; int token = blockIdx.x; @@ -681,50 +689,8 @@ __global__ void __launch_bounds__(1024) oneshotAllreduceFusionKernel(MnnvlAllRed } // ======================= Reduction ============================= - // Fully deterministic: every rank uses the exact same reduction order. For WorldSize <= 8, specialize the local - // slot so the fast path reuses `val` from registers without a dynamic `remoteValues[params.rank]` store. Larger - // world sizes use the compact fallback because the benefit is thin but specializing every rank significantly - // increases compile time. - PackedVec packedAccum; - if constexpr (WorldSize <= 8) - { - packedAccum = val; -#define RUN_ONESHOT_LOCAL_RANK(LOCAL_RANK) \ - case LOCAL_RANK: \ - if constexpr (WorldSize > LOCAL_RANK) \ - { \ - packedAccum = reduceOneshotDeterministicFastPath( \ - val, stagePtrLocal, token, params.tokenDim, packedIdx); \ - } \ - break - - switch (params.rank) - { - RUN_ONESHOT_LOCAL_RANK(0); - RUN_ONESHOT_LOCAL_RANK(1); - RUN_ONESHOT_LOCAL_RANK(2); - RUN_ONESHOT_LOCAL_RANK(3); - RUN_ONESHOT_LOCAL_RANK(4); - RUN_ONESHOT_LOCAL_RANK(5); - RUN_ONESHOT_LOCAL_RANK(6); - RUN_ONESHOT_LOCAL_RANK(7); - } -#undef RUN_ONESHOT_LOCAL_RANK - } - else - { - // Chunk Lamport polling so only a bounded rank set is live at once, avoiding register spills for large - // world sizes. - constexpr int kRankChunk = 8; - float accum[kELTS_PER_THREAD]; - accumulateLamportRanksChunked( - accum, stagePtrLocal, token, params.tokenDim, packedIdx); -#pragma unroll - for (int i = 0; i < kELTS_PER_THREAD; i++) - { - packedAccum.elements[i] = cuda_cast(accum[i]); - } - } + PackedVec packedAccum = reduceOneshotLamport( + val, stagePtrLocal, token, params.tokenDim, packedIdx, params.rank); #if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)) cudaTriggerProgrammaticLaunchCompletion(); #endif @@ -834,6 +800,8 @@ void oneshotAllreduceFusionOp(AllReduceFusionParams const& params) .attrs = attrs, .numAttrs = 2U, }; + // See MnnvlAllReduceKernelParams::earlyTrigger. + constexpr bool kEarlyTrigger = true; #define LAUNCH_ALLREDUCE_KERNEL(WORLD_SIZE, T, PATTERN) \ TLLM_CUDA_CHECK(cudaLaunchKernelEx(&config, &oneshotAllreduceFusionKernel, kernelParams)); @@ -898,7 +866,7 @@ void oneshotAllreduceFusionOp(AllReduceFusionParams const& params) MnnvlAllReduceKernelParams kernelParams{output, residualOut, input, residualIn, gamma, ucPtrs, reinterpret_cast(params.bufferPtrLocal), mcPtr, params.quantOut, params.scaleOut, params.scaleFactor, numTokens, tokenDim, params.nRanks, params.rank, static_cast(params.epsilon), params.bufferFlags, - false, params.layout}; + false, params.layout, kEarlyTrigger}; switch (params.nRanks) { diff --git a/cpp/tensorrt_llm/thop/allreduceOp.cpp b/cpp/tensorrt_llm/thop/allreduceOp.cpp index 51e1be287964..9ef1ed39b46e 100644 --- a/cpp/tensorrt_llm/thop/allreduceOp.cpp +++ b/cpp/tensorrt_llm/thop/allreduceOp.cpp @@ -27,6 +27,7 @@ #include "tensorrt_llm/kernels/communicationKernels/MiniMaxReduceRMSKernel.h" #include "tensorrt_llm/kernels/communicationKernels/allReduceFusionKernels.h" #include "tensorrt_llm/kernels/communicationKernels/customLowPrecisionAllReduceKernels.h" +#include "tensorrt_llm/kernels/communicationKernels/mnnvlAllGatherKernels.h" #include "tensorrt_llm/kernels/communicationKernels/mnnvlAllreduceKernels.h" #include "tensorrt_llm/kernels/customAllReduceKernels.h" #include "tensorrt_llm/kernels/moe/communication/moeAllReduceFusionKernels.h" @@ -2061,7 +2062,8 @@ bool hasMnnvlNormOutput(AllReduceFusionOp fusionOp) std::vector mnnvlFusionAllReduce(torch::Tensor& input, torch::optional const& gamma, torch::optional const& residual_in, torch::optional epsilon, torch::Tensor& comm_buffer, - torch::Tensor& buffer_flags, bool rmsnorm_fusion, torch::optional const& scale, int64_t fusion_op_) + torch::Tensor& buffer_flags, bool rmsnorm_fusion, torch::optional const& scale, int64_t fusion_op_, + int64_t one_shot_max_bytes) { auto* mcast_mem = tensorrt_llm::common::findMcastDevMemBuffer(comm_buffer.data_ptr()); TORCH_CHECK( @@ -2187,11 +2189,11 @@ std::vector mnnvlFusionAllReduce(torch::Tensor& input, torch::opt allreduce_params.rmsNormFusion = hasRmsNormFusion; allreduce_params.stream = at::cuda::getCurrentCUDAStream(input.get_device()); - // Threshold to switch between one-shot and two-shot allreduce kernel. - // Empirical value from the MNNVL sweep, matching FlashInfer's byte threshold. - constexpr size_t kOneShotSizeThreshold = 64 * 1024 * 8 * 2; + // Largest message sent one-shot; larger ones go two-shot. MNNVLAllReduce sizes the workspace with the same + // value (default: the empirical value from the MNNVL sweep, matching FlashInfer's byte threshold). + TORCH_CHECK(one_shot_max_bytes >= 0, "[mnnvlFusionAllReduce] one_shot_max_bytes must be non-negative"); - if (numTokens * hiddenDim * allreduce_params.nRanks * input.itemsize() <= kOneShotSizeThreshold) + if (numTokens * hiddenDim * allreduce_params.nRanks * input.itemsize() <= static_cast(one_shot_max_bytes)) { tensorrt_llm::kernels::mnnvl::oneshotAllreduceFusionOp(allreduce_params); } @@ -2289,6 +2291,67 @@ std::vector mnnvlAllReduceAttnRes(torch::Tensor const& input, return {normOut, prefixOut}; } +// One-shot all-gather over the MNNVL workspace of this rank's fp32 rows [num_tokens, columns]: +// the first bf16_columns columns of every rank are gathered as bf16 into [num_tokens, nRanks * +// bf16_columns], the rest as fp32 into [num_tokens, nRanks * (columns - bf16_columns)]; see +// tensorrt_llm::kernels::mnnvl::mnnvlAllGatherSplitOp. +namespace +{ + +// The all-gather's checks, outputs and params. +tensorrt_llm::kernels::mnnvl::AllGatherSplitParams makeAllGatherSplitParams(torch::Tensor const& input, + int64_t bf16_columns, torch::Tensor& comm_buffer, torch::Tensor& buffer_flags, torch::Tensor& bf16Out, + torch::Tensor& fp32Out) +{ + namespace mnnvl = tensorrt_llm::kernels::mnnvl; + auto* mcast_mem = tensorrt_llm::common::findMcastDevMemBuffer(comm_buffer.data_ptr()); + TORCH_CHECK( + mcast_mem != nullptr, "[mnnvlAllGatherSplit] comm_buffer must be obtained from a mcastBuffer instance."); + TORCH_CHECK(mcast_mem->isMapped(), "[mnnvlAllGatherSplit] MNNVL workspace handles are not attached."); + TORCH_CHECK(input.is_cuda() && input.scalar_type() == torch::kFloat32 && input.is_contiguous() && input.dim() == 2, + "[mnnvlAllGatherSplit] input must be a contiguous [num_tokens, columns] fp32 CUDA tensor"); + TORCH_CHECK(reinterpret_cast(input.const_data_ptr()) % 16 == 0, + "[mnnvlAllGatherSplit] input must be 16-byte aligned"); + int64_t const numTokens = input.size(0); + int64_t const fp32Columns = input.size(1) - bf16_columns; + TORCH_CHECK(bf16_columns >= 0 && fp32Columns >= 0 && bf16_columns % 8 == 0 && fp32Columns % 4 == 0, + "[mnnvlAllGatherSplit] needs bf16_columns a multiple of 8 and the remaining columns a multiple of 4"); + int64_t const nRanks = mcast_mem->getWorldSize(); + TORCH_CHECK(mnnvl::mnnvlAllGatherSplitFootprint(numTokens, bf16_columns, fp32Columns, nRanks) + <= comm_buffer.size(-1) * comm_buffer.element_size(), + "[mnnvlAllGatherSplit] the exchange does not fit in one Lamport buffer"); + + auto const options = input.options(); + bf16Out = torch::empty({numTokens, nRanks * bf16_columns}, options.dtype(torch::kBFloat16)); + fp32Out = torch::empty({numTokens, nRanks * fp32Columns}, options); + mnnvl::AllGatherSplitParams params{}; + params.input = input.const_data_ptr(); + params.bf16Output = reinterpret_cast<__nv_bfloat16*>(bf16Out.mutable_data_ptr()); + params.fp32Output = fp32Out.mutable_data_ptr(); + params.numTokens = static_cast(numTokens); + params.bf16Columns = static_cast(bf16_columns); + params.fp32Columns = static_cast(fp32Columns); + params.nRanks = static_cast(nRanks); + params.rank = mcast_mem->getRank(); + params.bufferPtrsDev = reinterpret_cast(mcast_mem->getBufferPtrsDev()); + params.multicastPtr = mcast_mem->getMulticastPtr(); + params.bufferFlags = reinterpret_cast(buffer_flags.mutable_data_ptr()); + params.stream = at::cuda::getCurrentCUDAStream(input.get_device()); + return params; +} + +} // namespace + +std::vector mnnvlAllGatherSplit( + torch::Tensor const& input, int64_t bf16_columns, torch::Tensor& comm_buffer, torch::Tensor& buffer_flags) +{ + torch::Tensor bf16Out; + torch::Tensor fp32Out; + auto const params = makeAllGatherSplitParams(input, bf16_columns, comm_buffer, buffer_flags, bf16Out, fp32Out); + tensorrt_llm::kernels::mnnvl::mnnvlAllGatherSplitOp(params); + return {bf16Out, fp32Out}; +} + torch::Tensor minimax_allreduce_rms(torch::Tensor const& input, torch::Tensor const& norm_weight, torch::Tensor workspace, int64_t const rank, int64_t const nranks, double const eps, bool const trigger_completion_at_end_) @@ -2388,12 +2451,15 @@ TORCH_LIBRARY_FRAGMENT(trtllm, m) m.def( "mnnvl_fusion_allreduce(Tensor input, Tensor? gamma, Tensor? residual, " "float? epsilon, Tensor(a!) comm_buffer, Tensor buffer_flags, bool rmsnorm_fusion, " - "Tensor? scale=None, int fusion_op=0) -> " + "Tensor? scale=None, int fusion_op=0, int one_shot_max_bytes=1048576) -> " "Tensor[]"); m.def( "mnnvl_allreduce_attn_res(Tensor input, Tensor? prefix_sum, Tensor block_residual, Tensor res_weight, " "Tensor rms_weight, Tensor output_rms_weight, float rms_eps, float output_rms_eps, Tensor(a!) comm_buffer, " "Tensor(b!) buffer_flags) -> Tensor[]"); + m.def( + "mnnvl_allgather_split(Tensor input, int bf16_columns, Tensor(a!) comm_buffer, Tensor(b!) buffer_flags) " + "-> Tensor[]"); m.def( "allreduce(" "Tensor input," @@ -2493,6 +2559,7 @@ TORCH_LIBRARY_IMPL(trtllm, CUDA, m) { m.impl("mnnvl_fusion_allreduce", &tensorrt_llm::torch_ext::mnnvlFusionAllReduce); m.impl("mnnvl_allreduce_attn_res", &tensorrt_llm::torch_ext::mnnvlAllReduceAttnRes); + m.impl("mnnvl_allgather_split", &tensorrt_llm::torch_ext::mnnvlAllGatherSplit); m.impl("allreduce", &tensorrt_llm::torch_ext::allreduce_raw); m.impl("autotuned_allreduce", &tensorrt_llm::torch_ext::autotunedAllreduce); m.impl("register_allreduce_tactic", &tensorrt_llm::torch_ext::registerAllReduceTactic); diff --git a/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py b/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py index fb7560d1b7cc..020cf6658bcb 100644 --- a/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py +++ b/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py @@ -137,7 +137,8 @@ def _(input, buffer_flags, rmsnorm_fusion, scale=None, - fusion_op: int = 0): + fusion_op: int = 0, + one_shot_max_bytes: int = 1048576): from tensorrt_llm.functional import AllReduceFusionOp op = AllReduceFusionOp(fusion_op) if op == AllReduceFusionOp.NONE and rmsnorm_fusion: diff --git a/tensorrt_llm/_torch/distributed/ops.py b/tensorrt_llm/_torch/distributed/ops.py index c7a61101d920..1f7a6a3a6065 100644 --- a/tensorrt_llm/_torch/distributed/ops.py +++ b/tensorrt_llm/_torch/distributed/ops.py @@ -818,6 +818,9 @@ def __init__(self, mapping: Mapping, dtype: torch.dtype): super().__init__() self.mapping = mapping self.dtype = dtype + # Largest num_tokens * hidden * ranks * element size this all-reduce sends one-shot when a call does not + # say; a model may raise it on its own all-reduces. + self.one_shot_max_bytes = _MNNVL_ONE_SHOT_THRESHOLD_BYTES if dtype not in MNNVLAllReduce.get_supported_dtypes() or ( mapping.has_cp()): # This is safe as we always capture the exception when create this object @@ -859,12 +862,16 @@ def is_mnnvl(mapping: Mapping, return supported and (explicitly_requested or mapping.is_multi_node()) @staticmethod - def get_required_workspace_size(num_tokens: int, hidden_dim: int, - group_size: int, dtype: torch.dtype) -> int: + def get_required_workspace_size( + num_tokens: int, + hidden_dim: int, + group_size: int, + dtype: torch.dtype, + one_shot_max_bytes: int = _MNNVL_ONE_SHOT_THRESHOLD_BYTES) -> int: elem_size = torch.tensor([], dtype=dtype).element_size() # This should match the heuristic in allreduceOp.cpp. is_one_shot = (num_tokens * hidden_dim * group_size * elem_size - <= _MNNVL_ONE_SHOT_THRESHOLD_BYTES) + <= one_shot_max_bytes) if is_one_shot: # For one-shot, each rank needs to store num_tokens * group_size tokens workspace_size = num_tokens * hidden_dim * group_size * elem_size @@ -945,12 +952,15 @@ def forward( self, input: torch.Tensor, all_reduce_params: AllReduceParams, + one_shot_max_bytes: Optional[int] = None, ) -> Union[torch.Tensor, Tuple[torch.Tensor, ...]]: """Forward pass for MNNVL AllReduce. Args: input (torch.Tensor): Input tensor to be reduced; last dim is the hidden dimension. all_reduce_params (Optional[AllReduceParams]): Parameters for fused operations. + one_shot_max_bytes (Optional[int]): Largest num_tokens * hidden * ranks * element size + sent one-shot (larger messages go two-shot); default ``self.one_shot_max_bytes``. Returns: Union[torch.Tensor, Tuple[torch.Tensor, ...]]: Reduced tensor(s). Output tensors @@ -958,12 +968,15 @@ def forward( NVFP4 scale-factor output is 1-D). """ + if one_shot_max_bytes is None: + one_shot_max_bytes = self.one_shot_max_bytes fusion_op = all_reduce_params.fusion_op hidden_dim = input.shape[-1] num_tokens = input.numel() // hidden_dim workspace_size_bytes = self.get_required_workspace_size( - num_tokens, hidden_dim, self.mapping.tp_size, self.dtype) + num_tokens, hidden_dim, self.mapping.tp_size, self.dtype, + one_shot_max_bytes) # We use uint32_t to store workspace size related info. Safeguard against overflow. if workspace_size_bytes >= 2**32 - 1: @@ -998,6 +1011,7 @@ def forward( is_fusion, # rmsnorm_fusion all_reduce_params.scale, # scale int(fusion_op), + one_shot_max_bytes, ) return tuple(outputs) if is_fusion else outputs[0] @@ -1042,6 +1056,30 @@ def allreduce_attn_res_rmsnorm( ) return normed, updated + def allgather_split(self, input: torch.Tensor, + bf16_columns: int) -> Tuple[torch.Tensor, torch.Tensor]: + """One-shot all-gather of this rank's fp32 rows over the MNNVL workspace. + + ``input`` is ``[num_tokens, columns]`` fp32. Returns ``(bf16_out, fp32_out)``: the + first ``bf16_columns`` columns of every rank, rounded to bf16, as ``[num_tokens, tp * + bf16_columns]`` in rank order, and the remaining columns as fp32 ``[num_tokens, tp * + (columns - bf16_columns)]``. It takes a turn of the one-shot all-reduce's Lamport + rotation, so it must run in the same stream order as the other MNNVL collectives of + this workspace. The first call for a shape must happen outside CUDA graph capture. + """ + num_tokens, columns = input.shape + footprint = num_tokens * self.mapping.tp_size * ( + bf16_columns * 2 + (columns - bf16_columns) * 4) + workspace = get_or_scale_allreduce_mnnvl_workspace( + self.mapping, self.dtype, buffer_size_bytes=footprint) + bf16_out, fp32_out = torch.ops.trtllm.mnnvl_allgather_split( + input, + bf16_columns, + workspace["uc_buffer"].view(self.dtype).view(3, -1), + workspace["buffer_flags"], + ) + return bf16_out, fp32_out + class AllReduce(nn.Module): From 65def521c268861efe1b67da3f645483f5c0233c Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 18:03:27 -0700 Subject: [PATCH 056/161] [None][feat] Kimi K3 decode: collective CuTe DSL kernels with their op tests The Kimi K3 decode kernels whose calls exchange data across the TP group, as of the K3 stack 82a110a92a, verbatim apart from k3_fused_moe/op.py: - k3_sandwich/: trtllm::k3_sandwich_oproj, _tail and _plain (a row-parallel projection, its TP all-reduce and the residual update in one kernel; the post-attention step, the pre-attention step after the MoE tail, and the drafter's plain residual add + RMSNorm); - k3_fused_moe/: trtllm::k3_moe_front (the MoE front: head GEMV, head all-gather, top-16 routing, MXFP8 latent, shared gate_up + SiTU), trtllm::k3_fused_moe and k3_fused_moe_front (the routed experts of up to 8 tokens), K3MoeWideState (up to 64 tokens) and trtllm::k3_latent_reduce (the latent all-reduce as the consumer of pushed partials); - k3_route_quant/: trtllm::k3_route_quant (top-16 routing and the MXFP8 input quantization). op.py leaves out the route-B engines (k3_moe_m1 / k3_moe_m2 and the push wrappers trtllm::k3_fused_moe_push / k3_fused_moe_front_push), which follow in their own PR. The op tests come along; they carry the regression tests of the fixes made to these kernels: the FC2 partial rows past M never read unwritten memory, the collective buffers refuse a first use under CUDA-graph capture, the sandwich call counters across the int32 wrap (also with the latent all-reduce folded into the tail), the front's ready words acquired by k3_moe before the front's grid ends and published in order, and mixed token counts back to back on the wide build. Signed-off-by: Vasanth Sabavat --- .../cute_dsl_kernels/k3_fused_moe/__init__.py | 19 + .../cute_dsl_kernels/k3_fused_moe/front_op.py | 231 ++ .../k3_fused_moe/k3_latent_reduce.py | 191 + .../k3_fused_moe/k3_moe_front.py | 1297 +++++++ .../k3_fused_moe/k3_moe_kernel.py | 3144 +++++++++++++++++ .../k3_fused_moe/k3_route_quant_ag.py | 360 ++ .../k3_fused_moe/latent_op.py | 158 + .../cute_dsl_kernels/k3_fused_moe/op.py | 973 +++++ .../k3_route_quant/__init__.py | 19 + .../k3_route_quant/k3_route_quant_kernel.py | 431 +++ .../cute_dsl_kernels/k3_route_quant/op.py | 146 + .../cute_dsl_kernels/k3_sandwich/__init__.py | 19 + .../k3_sandwich/k3_sandwich_kernel.py | 2105 +++++++++++ .../_torch/cute_dsl_kernels/k3_sandwich/op.py | 616 ++++ .../kimi_k3/test_k3_fused_moe.py | 341 ++ .../kimi_k3/test_k3_latent_reduce.py | 259 ++ .../kimi_k3/test_k3_moe_front.py | 555 +++ .../kimi_k3/test_k3_moe_wide.py | 459 +++ .../kimi_k3/test_k3_route_quant.py | 161 + .../kimi_k3/test_k3_sandwich.py | 524 +++ 20 files changed, 12008 insertions(+) create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/__init__.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/front_op.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_latent_reduce.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_moe_front.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_moe_kernel.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_route_quant_ag.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_route_quant/__init__.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_route_quant/k3_route_quant_kernel.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_route_quant/op.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/__init__.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/k3_sandwich_kernel.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_latent_reduce.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_wide.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_route_quant.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/__init__.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/__init__.py new file mode 100644 index 000000000000..c20fcf7c5dfa --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/__init__.py @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Kimi K3 fused routed-expert decode path (``trtllm::k3_fused_moe``). + +Importing :mod:`.op` registers the torch op; nothing is imported eagerly here so that +the CuTe DSL dependency stays optional for every other model. +""" diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/front_op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/front_op.py new file mode 100644 index 000000000000..994c9715b7b4 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/front_op.py @@ -0,0 +1,231 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``trtllm::k3_moe_front``: the Kimi K3 MoE front at decode size in one kernel (``k3_moe_front.py``). + +Replaces, per MoE layer: the sharded head GEMV, ``trtllm::k3_route_quant_ag`` (head all-gather, top-16 routing, MXFP8 +latent) and, on the shared-expert stream, the shared gate_up GEMV and SiTU-and-mul. The head all-gather uses +``op.head_workspace``'s buffers and protocol, so the front and ``k3_route_quant_ag`` must not both serve one layer. +The kernel compiles on the first call for each configuration, which must happen outside CUDA-graph capture. +""" + +from __future__ import annotations + +import functools +import os +import threading +from typing import Dict, Optional, Tuple + +import torch + +TOP_K = 16 +HIDDEN_SIZE = 3584 +NUM_EXPERTS = 896 +SF_VEC = 32 +MAX_TOKENS = 8 +DEFAULT_RING = 4 + +_lock = threading.Lock() +_compiled: Dict[tuple, object] = {} + + +def _kernel(): + from . import k3_moe_front as kernel + + return kernel + + +def front_weight( + head_weight: torch.Tensor, gate_up_weight: Optional[torch.Tensor] = None +) -> torch.Tensor: + """The front's one weight: ``head_weight`` (``[3584/W + 896/W, K]``: this rank's latent-down rows, then its router + rows) zero-padded to whole 128-row tiles, then ``gate_up_weight`` (``[2 I, K]``: gate rows, then up rows) with + every 32 rows holding 16 gate rows and the 16 up rows of the same columns. Without ``gate_up_weight`` the front + computes the head alone (``shared_cols`` 0).""" + rows, k_in = head_weight.shape + if gate_up_weight is None: + gate_up_weight = head_weight[:0] + inter = gate_up_weight.shape[0] // 2 + if ( + gate_up_weight.shape[1] != k_in + or inter % 64 != 0 + or head_weight.dtype != gate_up_weight.dtype + ): + raise ValueError( + f"k3_moe_front weight: head {tuple(head_weight.shape)}, gate_up {tuple(gate_up_weight.shape)} " + f"(needs equal K and dtype and a multiple of 64 shared columns)" + ) + padded = -(-rows // 128) * 128 + head = torch.zeros(padded, k_in, dtype=head_weight.dtype, device=head_weight.device) + head[:rows].copy_(head_weight) + gate = gate_up_weight[:inter].reshape(inter // 16, 16, k_in) + up = gate_up_weight[inter:].reshape(inter // 16, 16, k_in) + shared = torch.stack([gate, up], dim=1).reshape(2 * inter, k_in) + return torch.cat([head, shared]).contiguous() + + +@functools.lru_cache(maxsize=None) +def _max_clusters_of(device_index: int) -> int: + from cutlass.utils.hardware_info import HardwareInfo + + return HardwareInfo(device_index).get_max_active_clusters(_kernel().SPLIT) + + +def max_clusters(device: torch.device) -> int: + """Clusters of the kernel's SPLIT CTAs, one CTA per SM, that ``device`` holds at once.""" + index = device.index if device.index is not None else torch.cuda.current_device() + return _max_clusters_of(index) + + +def weight_supported( + world: int, shared_cols: int, k_in: int, device: torch.device, ring: int = DEFAULT_RING +) -> bool: + """Whether the front runs this TP world, shared activation width and hidden size (any M <= 8).""" + return _kernel().supports(world, shared_cols, max_clusters(device), k_in, ring) + + +def supports( + x: torch.Tensor, w_front: torch.Tensor, world: int, shared_cols: int, ring: int = DEFAULT_RING +) -> bool: + kernel = _kernel() + return ( + x.dim() == 2 + and 0 < x.shape[0] <= MAX_TOKENS + and x.dtype == torch.bfloat16 + and w_front.dtype == torch.bfloat16 + and x.shape[1] == w_front.shape[1] + and w_front.shape[0] == (kernel.head_tiles(world) * 128 + 2 * shared_cols) + and kernel.supports(world, shared_cols, max_clusters(x.device), x.shape[1], ring) + ) + + +@torch.library.custom_op("trtllm::k3_moe_front", mutates_args=()) +def k3_moe_front( + x: torch.Tensor, + w_front: torch.Tensor, + bias: torch.Tensor, + routed_scaling_factor: float, + shared_cols: int, + gate_cap: float, + linear_cap: float, + ag_uc: torch.Tensor, + ag_mc: torch.Tensor, + ag_flags: torch.Tensor, + ag_rank: int, + ag_world: int, + ring: int = DEFAULT_RING, + ag_ready: Optional[torch.Tensor] = None, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """``x`` bf16 ``[M <= 8, 7168]`` (the MoE input, the same on every rank), ``w_front`` from ``front_weight``, + ``bias`` the routing bias fp32 ``[896]``. Returns ``(topk_ids, topk_weights, quantized, scales, shared)``: what + ``trtllm::k3_route_quant_ag`` returns for the gathered head, and the shared experts' activation bf16 + ``[M, shared_cols]``. With ``ag_ready`` it also releases the per-token ready words as route_quant_ag does.""" + import cuda.bindings.driver as cuda_driver + import cutlass.cute as cute + from cutlass.cute.runtime import from_dlpack + + kernel = _kernel() + if not supports(x, w_front, ag_world, shared_cols, ring): + raise ValueError( + f"k3_moe_front: unsupported call x {tuple(x.shape)} {x.dtype}, w_front {tuple(w_front.shape)} " + f"{w_front.dtype}, world {ag_world}, shared_cols {shared_cols}, ring {ring}" + ) + from . import k3_route_quant_ag as layout + + ag_words = layout.workspace_words(ag_world) + if ag_uc.numel() < ag_words or ag_mc.numel() < ag_words: + raise ValueError( + f"k3_moe_front: the head workspace holds {ag_uc.numel()} words, the front needs {ag_words} " + "(op.head_workspace: the all-gather's buffers, then the router partials)" + ) + num_tokens, k_in = x.shape + device = x.device + topk_ids = torch.empty(num_tokens, TOP_K, dtype=torch.int32, device=device) + topk_weights = torch.empty(num_tokens, TOP_K, dtype=torch.bfloat16, device=device) + quantized = torch.empty(num_tokens, HIDDEN_SIZE, dtype=torch.float8_e4m3fn, device=device) + scales = torch.empty(num_tokens, HIDDEN_SIZE // SF_VEC, dtype=torch.uint8, device=device) + shared = torch.empty(num_tokens, shared_cols, dtype=torch.bfloat16, device=device) + # Head-only fronts (shared_cols 0) still hand the kernel an aligned shared-output argument it never writes. + shared_arg = shared if shared_cols else ag_flags.view(torch.int16).view(-1) + + def arg2(t): + return from_dlpack(t.detach(), assumed_align=16).mark_layout_dynamic(leading_dim=1) + + def arg(t): + return from_dlpack(t.detach(), assumed_align=16).mark_layout_dynamic(leading_dim=0) + + args = ( + arg2(w_front), + arg2(x.contiguous()), + arg(bias.contiguous().view(-1)), + arg(ag_uc.view(-1)), + arg(ag_mc.view(-1)), + arg(ag_flags.view(-1)), + arg(ag_ready.view(-1) if ag_ready is not None else ag_flags.view(-1)), + arg(topk_ids.view(-1)), + arg(topk_weights.view(-1).view(torch.int16)), + arg(quantized.view(-1).view(torch.int32)), + arg(scales.view(-1)), + arg(shared_arg.view(-1).view(torch.int16)), + ) + capacity = max_clusters(device) + # One round with the head in 64-row half-tiles when it fits (TP16): every head k-tile on chip before the wait. + half = kernel.half_geometry(ag_world, shared_cols, capacity, k_in, ring) + head_half = half is not None + ht, st, clusters = half if head_half else kernel.geometry(ag_world, shared_cols, capacity) + publish = ag_ready is not None + use_pdl = os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" + stream = cuda_driver.CUstream(torch.cuda.current_stream(device).cuda_stream) + key = ( + ag_world, + shared_cols, + ht + st, + clusters, + k_in, + ring, + float(gate_cap), + float(linear_cap), + publish, + use_pdl, + head_half, + ) + fn = _compiled.get(key) + if fn is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "trtllm::k3_moe_front must run once per configuration outside CUDA-graph capture first" + ) + with _lock: + fn = _compiled.get(key) + if fn is None: + fn = _compiled[key] = cute.compile( + kernel.k3_moe_front, *args, num_tokens, ag_rank, float(routed_scaling_factor), ag_world, + shared_cols, ht, ht + st, clusters, k_in, ring, float(gate_cap), float(linear_cap), publish, + use_pdl, head_half, stream, + ) # fmt: skip + fn(*args, num_tokens, ag_rank, float(routed_scaling_factor), stream) + return topk_ids, topk_weights, quantized, scales, shared + + +@k3_moe_front.register_fake +def _(x, w_front, bias, routed_scaling_factor, shared_cols, gate_cap, linear_cap, ag_uc, ag_mc, ag_flags, ag_rank, + ag_world, ring=DEFAULT_RING, ag_ready=None): # fmt: skip + num_tokens = x.shape[0] + return ( + x.new_empty((num_tokens, TOP_K), dtype=torch.int32), + x.new_empty((num_tokens, TOP_K), dtype=torch.bfloat16), + x.new_empty((num_tokens, HIDDEN_SIZE), dtype=torch.float8_e4m3fn), + x.new_empty((num_tokens, HIDDEN_SIZE // SF_VEC), dtype=torch.uint8), + x.new_empty((num_tokens, shared_cols), dtype=torch.bfloat16), + ) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_latent_reduce.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_latent_reduce.py new file mode 100644 index 000000000000..a83f870cefe8 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_latent_reduce.py @@ -0,0 +1,191 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Kimi K3 latent all-reduce for decode (M <= 8 tokens) as the consumer of k3_moe's pushed partials. + +The push-only k3_moe (``trtllm::k3_fused_moe_push`` / ``trtllm::k3_fused_moe_front_push``) stores every rank's routed +partial through the multicast mapping into slot [rank] of half ``flags[0] & 1`` of every rank's buffer, int32 +``[2][8][world][1792]`` (bf16 pairs, ``0x80000000`` not written, -0.0 stored as +0.0), and writes no flag. This kernel +sums the slots of the live rows in the MNNVL one-shot's order (``reduceOneshotLamport``: fp32 over chunks of 8 ranks +in rank order, each from 0, the chunks added in order, then bf16), so its row equals the one-shot all-reduce of the +partials bit for bit. It owns the protocol's state: ``flags[0]`` counts its calls (the half), ``flags[2]`` counts the +CTAs of the running call in. + +Grid M x CTAS, 448 / CTAS threads, one 16-byte vector of a token's 3584-wide row per thread. Before the grid-dependency +wait thread 0 reads the count and counts its CTA in: the predecessor (the push) read the same count after its own +wait, and this kernel's previous call ended before that. After the wait every thread loads its vector from all slots +in one pass, repeated until no word is empty, sums, stores the row and empties the words it read (the next push into +this half comes two calls later, after this grid). CTA 0's thread 0 advances the count once every CTA is in. +""" + +from __future__ import annotations + +import cuda.bindings.driver as cuda_driver +import cutlass +import cutlass.cute as cute +from cutlass._mlir.dialects import llvm as _llvm +from cutlass._mlir.extras import types as _T +from cutlass.cutlass_dsl import dsl_user_op +from cutlass.experimental import primitives as prims + +HIDDEN_SIZE = 3584 +ROW_WORDS = HIDDEN_SIZE // 2 # int32 words (bf16 pairs) of one row +ROW_VECS = HIDDEN_SIZE // 8 # 16-byte vectors of one row +MAX_TOKENS = 8 +BUFFERS = 2 +RANK_CHUNK = 8 +EMPTY_WORD = -(2**31) # 0x80000000 +FLAG_COUNT = 0 +FLAG_ARRIVED = 2 + + +def buffer_words(world: int) -> int: + """Int32 words of one rank's buffer (both halves).""" + return BUFFERS * MAX_TOKENS * world * ROW_WORDS + + +@dsl_user_op +def _pack_bf16x2(hi, lo, *, loc=None, ip=None): + """(bf16(hi) << 16) | bf16(lo), round to nearest even.""" + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [hi.ir_value(loc=loc, ip=ip), lo.ir_value(loc=loc, ip=ip)], + "cvt.rn.bf16x2.f32 $0, $1, $2;", "=r,f,f", has_side_effects=False, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _red_add_release(addr_i64, val, *, loc=None, ip=None): + """red.release.gpu.global.add.u32: this thread's earlier reads are performed before the add.""" + _llvm.inline_asm( + None, [addr_i64.ir_value(loc=loc, ip=ip), cutlass.Int32(val).ir_value(loc=loc, ip=ip)], + "red.release.gpu.global.add.u32 [$0], $1;", "l,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _ld_acquire_gpu(addr_i64, *, loc=None, ip=None): + """ld.acquire.gpu.global.u32.""" + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [addr_i64.ir_value(loc=loc, ip=ip)], "ld.acquire.gpu.global.u32 $0, [$1];", "=r,l", + has_side_effects=True, is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +def _lo(word): + return (word << cutlass.Int32(16)).bitcast(cutlass.Float32) + + +def _hi(word): + return (word & cutlass.Int32(-65536)).bitcast(cutlass.Float32) + + +@cute.kernel +def k3_latent_reduce_kernel( + buf: cutlass.Array, # int32 words of this rank's buffer, [2][8][world][1792] + flags: cutlass.Array, # int32 [4]: [0] the call count, [2] the CTAs of this call counted in + out: cutlass.Array, # int32 view of the bf16 output [M, 3584]: [M * 1792] + world: cutlass.Constexpr[int], + threads: cutlass.Constexpr[int], +): + tidx, _, _ = cute.arch.thread_idx() + tok, cta, _ = cute.arch.block_idx() + n_tok, n_cta, _ = cute.arch.grid_dim() + s_count = cutlass.Array(cutlass.Int32, 4, space=cutlass.AddressSpace.smem, alignment=16) + if tidx == 0: + count = cutlass.Int32(flags.load(idx=FLAG_COUNT, is_volatile=True)) + s_count.store(count, idx=0) + # Released after the read: CTA 0 advances the count only once every CTA has read it. + _red_add_release(flags.data_ptr(FLAG_ARRIVED).toint(), cutlass.Int32(1)) + cute.arch.barrier() + count = s_count.load(idx=0) + + prims.griddepcontrol(prims.GridDepAction.WAIT) + # The dependents read this grid's output only after their own grid-dependency wait. + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + + vec = cta * cutlass.Int32(threads) + tidx + base = ( + ((count & cutlass.Int32(1)) * cutlass.Int32(MAX_TOKENS) + tok) * cutlass.Int32(world) + ) * cutlass.Int32(ROW_WORDS) + vec * cutlass.Int32(4) + a0 = cutlass.Float32(0.0) + a1 = cutlass.Float32(0.0) + a2 = cutlass.Float32(0.0) + a3 = cutlass.Float32(0.0) + a4 = cutlass.Float32(0.0) + a5 = cutlass.Float32(0.0) + a6 = cutlass.Float32(0.0) + a7 = cutlass.Float32(0.0) + pending = cutlass.Boolean(True) + while pending: + dirty = cutlass.Boolean(False) + total = [cutlass.Float32(0.0)] * 8 + for rb in cutlass.range_constexpr(0, world, RANK_CHUNK): + chunk = [cutlass.Float32(0.0)] * 8 + for rr in cutlass.range_constexpr(min(RANK_CHUNK, world - rb)): + v = buf.load(idx=base + cutlass.Int32((rb + rr) * ROW_WORDS), vector_size=4, alignment=16, + is_volatile=True) # fmt: skip + for q in cutlass.range_constexpr(4): + word = cutlass.Int32(v[q]) + dirty = dirty | (word == cutlass.Int32(EMPTY_WORD)) + chunk[2 * q] = chunk[2 * q] + _lo(word) + chunk[2 * q + 1] = chunk[2 * q + 1] + _hi(word) + for e in cutlass.range_constexpr(8): + total[e] = total[e] + chunk[e] + a0, a1, a2, a3, a4, a5, a6, a7 = total + pending = dirty + + out.store( + (_pack_bf16x2(a1, a0), _pack_bf16x2(a3, a2), _pack_bf16x2(a5, a4), _pack_bf16x2(a7, a6)), + idx=tok * cutlass.Int32(ROW_WORDS) + vec * cutlass.Int32(4), + alignment=16, + ) + empty = cutlass.Int32(EMPTY_WORD) + for r in cutlass.range_constexpr(world): + buf.store( + (empty, empty, empty, empty), idx=base + cutlass.Int32(r * ROW_WORDS), alignment=16 + ) + + if (tok == 0) & (cta == 0) & (tidx == 0): + everyone = n_tok * n_cta + arrived = _ld_acquire_gpu(flags.data_ptr(FLAG_ARRIVED).toint()) + while arrived < cutlass.Int32(everyone): + arrived = _ld_acquire_gpu(flags.data_ptr(FLAG_ARRIVED).toint()) + flags.store(cutlass.Int32(0), idx=FLAG_ARRIVED) + flags.store(count + cutlass.Int32(1), idx=FLAG_COUNT) + + +@cute.jit +def k3_latent_reduce( + buf: cute.Tensor, + flags: cute.Tensor, + out: cute.Tensor, + num_tokens: cutlass.Int32, + world: cutlass.Constexpr[int], + ctas: cutlass.Constexpr[int], + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + threads = ROW_VECS // ctas + k3_latent_reduce_kernel(buf, flags, out, world, threads).launch( + grid=[num_tokens, ctas, 1], + block=[threads, 1, 1], + stream=stream, + use_pdl=use_pdl, + ) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_moe_front.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_moe_front.py new file mode 100644 index 000000000000..ce3e2c39fdc8 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_moe_front.py @@ -0,0 +1,1297 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +# ============================================================================= +# Kimi K3 MoE front for decode, one CTM (prims/cute) kernel per MoE layer, M <= 8 tokens: +# +# head = x @ [latent-down slice; router slice]^T (fp32; this rank's 3584/W latent columns, 896/W logits) +# gather = every rank's head slice, through the multicast mapping of the head workspace (Lamport buffers) +# route = top-16 of sigmoid(logits) + bias per token; MXFP8 latent + UE8M0 scales per token +# shared = SiTU(bf16(x @ gate^T)) * SiTU_lin(bf16(x @ up^T)) (the shared experts' gate_up and activation) +# +# One weight W [N, K = 7168] bf16 in 128-row tiles: the head rows (latent slice, then router slice, then zero rows +# up to a whole tile), then the shared gate_up rows re-ordered so that every 32-row slice holds 16 gate rows and +# the 16 up rows of the same columns (`front_weight`). +# +# Geometry (the long-K CTM decode GEMV): clusters of SPLIT = 8 CTAs; a GEMV cluster owns one 128-row tile per round +# (at most three rounds), rank r streaming k-tiles r, r + 8, ... through a RING-stage ring filled, with every later +# k-tile of every round prefetched into L2, before any wait; x's k-tiles of the rank stay resident. tcgen05 MMA +# M 128, N 8, K 16 into a TMEM accumulator per round. Split-K: rank w < 4 owns rows [32 w, 32 w + 32) of the tile; +# the other ranks write their fp32 partials of those rows into slot [rank] of the owner's round mailbox by st.async, +# which completes the mailbox barrier by bytes; the owner, which expected the bytes at setup, spins on test_wait (a +# warp suspended in try_wait is woken late by remote completions) and adds the 8 partials in rank order. Then, by +# tile kind: +# head latent rows -> bf16 pairs (cvt.rn.bf16x2.f32, -0.0 halves -> +0.0), 16-byte vectors pushed into slot +# [buffer][token][rank] of every rank's head buffer (the head all-gather's packing: see k3_route_quant_ag.py); +# head router rows skip the owner: every rank of the cluster pushes its own fp32 partial (-0.0 -> +0.0) into +# [buffer][token][rank][q = cluster rank] of the partials region after the buffers, and the route CTAs add the +# 8 partials in rank order from +0.0 (the owner's sum, bit for bit); no fence after the pushes, the readers poll +# the words themselves; +# shared rows -> gate and up of one column in lanes l and l + 16: bf16(SiTU(bf16 gate) * SiTU_lin(bf16 up)). +# Role clusters, the first two of the grid (resident before any GEMV cluster): CTA t < M routes token t, CTA M + t +# quantizes it, k3_route_quant_ag's device code minus its push (the GEMV epilogues push), the route CTA selecting +# with all its warps (_top16_cta) instead of top16_warp's one warp. Every CTA that reads the +# buffer index (the 2 M role CTAs and every GEMV CTA) signs in on flags[3] right after the read; thread 0 of token 0's +# quantization CTA waits for all of them and flips the index for the next call, long before the data it polls arrives, +# so no count sits on the role CTAs' output path. With `publish`, per-token ready words as route_quant_ag's. +# +# Warps of a GEMV CTA: 0 weight TMA (+ early dependent trigger), 1 grid-dependency wait, head workspace flags, +# activation TMA, 2 TMEM allocation + MMA, 3 idle, 4-7 epilogue, 8-13 exit after setup. Role CTAs: all 14 warps +# (the quantization CTA holds one 8-element vector per thread). +# ============================================================================= +"""Kimi K3 MoE front (head GEMV + all-gather + routing + MXFP8, shared gate_up + SiTU) as one kernel.""" + +from __future__ import annotations + +import cuda.bindings.driver as cuda_driver +import cutlass +import cutlass.cute as cute +import cutlass.experimental.cuda as cuda +from cutlass import dsl_user_op +from cutlass.experimental import primitives as prims + +from tensorrt_llm._torch.cute_dsl_kernels.k3_route_quant import k3_route_quant_kernel as _rq + +from . import k3_route_quant_ag as _ag + +NUM_EXPERTS = 896 +TOP_K = 16 +HIDDEN_SIZE = 3584 # the routed latent +MODEL_HIDDEN = 7168 # K +SF_VEC_SIZE = 32 +M_MAX = 8 +EMPTY_WORD = _ag.EMPTY_WORD + +CTA_M = 128 +HALF_M = 64 # a head half-tile's rows (one round, e.g. TP16: every head k-tile of the CTA on chip before the wait) +MAX_HALF_TILES = 8 # k-tiles a head half-tile CTA can hold (its full-barrier count) +MMA_N = 8 +CTA_K = 128 +MMA_K = 16 +TMA_K_BOX = 64 +TMA_COPY_ITERS = CTA_K // TMA_K_BOX +K_BLOCKS_PER_HALF = TMA_K_BOX // MMA_K +SPLIT = 8 +assert SPLIT == _ag.FRONT_SPLIT # the router partials region is sized by the all-gather's layout +OWNERS = 4 # ranks owning 32-row slices of a tile (epilogue warps 4-7 hold TMEM lanes 32 w ..) +ROWS_PER_WARP = 32 +MAX_ROUNDS = 3 +ROLE_CLUSTERS = 2 +THREADS = HIDDEN_SIZE // 8 # 448: one 8-element latent vector per thread in a quantization CTA +TMEM_COLS = 32 # one 8-column accumulator per round (3 x 8 columns used) +# Head half-tile CTAs (head_half) copy every A k-tile into TMEM before the grid wait and run their MMAs from there +# (A from shared memory costs ~4 KB of smem reads per M 128 K 16 MMA): k-tile i at columns A_TMEM_BASE + 64 i, +# 8 columns per K = 16 step, all of TMEM. +A_TMEM_BASE = 64 +A_TMEM_COLS = CTA_K // 2 +A_KSTEP_COLS = MMA_K // 2 +HALF_TMEM_COLS = 512 +ELEM_BYTES = 2 +EVICT_FIRST = 0x12F0000000000000 +# The route CTA's top-16 (_top16_cta): candidate slots warp 0 ranks (more candidates, from heavy ties, take +# _rq.top16_warp), an empty slot (below every candidate), the id complement that makes lower ids rank first. +CAND_SLOTS = 32 +CAND_EMPTY = -(2**63) +NO_ID = 0x7FFFFFFF +KEY_MAX = 2**31 - 1 + +LEADING = 16 +STRIDE = 8 * TMA_K_BOX * ELEM_BYTES +A_HALF_ELEMS = CTA_M * TMA_K_BOX +B_HALF_ELEMS = MMA_N * TMA_K_BOX +STEP = (MMA_K * ELEM_BYTES) >> 4 +A_BOX = A_HALF_ELEMS >> 3 +B_BOX = B_HALF_ELEMS >> 3 +STAGE_A = (CTA_M * CTA_K * ELEM_BYTES) >> 4 +STAGE_AH = (HALF_M * CTA_K * ELEM_BYTES) >> 4 # a head half-tile's k-tile (16 KB) +A_BOX_H = (HALF_M * TMA_K_BOX) >> 3 # its 64-element K box (8 KB) +STAGE_B = (MMA_N * CTA_K * ELEM_BYTES) >> 4 + +io_dtype = cutlass.BFloat16 + + +def head_rows(world: int) -> int: + return HIDDEN_SIZE // world + NUM_EXPERTS // world + + +def head_tiles(world: int) -> int: + return (head_rows(world) + CTA_M - 1) // CTA_M + + +def geometry(world: int, shared_cols: int, max_clusters: int) -> tuple[int, int, int]: + """(head tiles, shared tiles, GEMV clusters) for this TP world and shared activation width. + + ``max_clusters``: how many clusters of SPLIT CTAs (one CTA per SM) the GPU holds at once. The role clusters and + every GEMV cluster must fit together: the role clusters stay until the routing is done, so a GEMV cluster that + does not fit waits for them (a 152-SM GB200 holds 15, not the 16 its 8 GPCs suggest). + """ + ht = head_tiles(world) + st = shared_cols // 64 + tiles = ht + st + clusters = min(tiles, max(1, max_clusters - ROLE_CLUSTERS)) + return ht, st, clusters + + +def half_geometry(world: int, shared_cols: int, max_clusters: int, k_in: int, ring: int): + """(head half-tiles, shared tiles, GEMV clusters) when the head fits as 64-row half-tiles in one round next to the + shared tiles, each head CTA holding all its k-tiles in shared memory (the ring's space); else None.""" + hh = (head_rows(world) + HALF_M - 1) // HALF_M + st = shared_cols // 64 + my_tiles = k_in // CTA_K // SPLIT + fits = ( + hh + st <= max_clusters - ROLE_CLUSTERS + and my_tiles * HALF_M * CTA_K + HALF_M * TMA_K_BOX <= ring * CTA_M * CTA_K + and my_tiles <= MAX_HALF_TILES + and A_TMEM_BASE + my_tiles * A_TMEM_COLS <= HALF_TMEM_COLS + ) + return (hh, st, hh + st) if fits else None + + +def _weight_rows(world: int, n_head_tiles: int, n_tiles: int, head_half: bool) -> int: + """Rows of the front weight: the head padded to 128-row tiles, then the shared tiles.""" + if not head_half: + return n_tiles * CTA_M + return (head_tiles(world) + n_tiles - n_head_tiles) * CTA_M + + +def _half_plan( + head_half: bool, world: int, n_head_tiles: int, ring: int, my_tiles: int +) -> tuple[int, int]: + """(weight-row shift of the shared tiles, weight full/empty barriers) for the kernel's tile mode.""" + if not head_half: + return 0, ring + return (head_tiles(world) - n_head_tiles) * CTA_M, max(ring, my_tiles) + + +def supports(world: int, shared_cols: int, max_clusters: int, k_in: int, ring: int) -> bool: + ht, st, clusters = geometry(world, shared_cols, max_clusters) + k_tiles = k_in // CTA_K + return ( + world in (4, 8, 16) + and k_in % (CTA_K * SPLIT) == 0 + and shared_cols % 64 == 0 + and ht + st <= MAX_ROUNDS * clusters + and 1 <= ring <= k_tiles // SPLIT + and ring * CTA_M * CTA_K * ELEM_BYTES + (k_tiles // SPLIT) * MMA_N * CTA_K * ELEM_BYTES + <= 176 * 1024 + ) + + +def _situ(gate, up, gate_cap: float, linear_cap: float): + """SiTU(gate) * SiTU_lin(up) in fp32 (the fused MoE's device functions).""" + log2e = 1.4426950408889634 + + def tanh_f32(v): + e = cute.math.exp2(cute.math.abs(v) * cutlass.Float32(-2.0 * log2e), fastmath=True) + t = (cutlass.Float32(1.0) - e) * cute.arch.rcp_approx(cutlass.Float32(1.0) + e) + return cutlass.select_(v < cutlass.Float32(0.0), -t, t) + + sig = cute.arch.rcp_approx( + cutlass.Float32(1.0) + cute.math.exp2(gate * cutlass.Float32(-log2e), fastmath=True) + ) + g = cutlass.Float32(gate_cap) * tanh_f32(gate * cutlass.Float32(1.0 / gate_cap)) * sig + u = cutlass.Float32(linear_cap) * tanh_f32(up * cutlass.Float32(1.0 / linear_cap)) + return g * u + + +@dsl_user_op +def _mapa_u32(smem_ptr, peer, *, loc=None, ip=None): + """The shared::cluster address of this CTA's shared-memory location in cluster CTA ``peer``.""" + from cutlass._mlir.dialects import llvm as _llvm + from cutlass._mlir.extras import types as _T + + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [smem_ptr.toint(loc=loc, ip=ip).ir_value(), cutlass.Int32(peer).ir_value(loc=loc, ip=ip)], + "mapa.shared::cluster.u32 $0, $1, $2;", "=r,r,r", has_side_effects=False, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _st_async_v4_f32(dst, a, b, c, d, mbar, *, loc=None, ip=None): + """st.async of four fp32 to a shared::cluster address, completing ``mbar`` (shared::cluster) by 16 bytes.""" + from cutlass._mlir.dialects import llvm as _llvm + + _llvm.inline_asm( + None, + [cutlass.Int32(dst).ir_value(loc=loc, ip=ip), cutlass.Float32(a).ir_value(loc=loc, ip=ip), + cutlass.Float32(b).ir_value(loc=loc, ip=ip), cutlass.Float32(c).ir_value(loc=loc, ip=ip), + cutlass.Float32(d).ir_value(loc=loc, ip=ip), cutlass.Int32(mbar).ir_value(loc=loc, ip=ip)], + "st.async.shared::cluster.mbarrier::complete_tx::bytes.v4.f32 [$0], {$1, $2, $3, $4}, [$5];", "r,f,f,f,f,r", + has_side_effects=True, is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _test_wait_cluster(mbar, parity, *, loc=None, ip=None): + """mbarrier.test_wait.parity.acquire.cluster on a barrier of this CTA (``mbar``, its shared-memory pointer): whether + phase ``parity`` has completed, acquiring at cluster scope. The barrier is completed by other CTAs' st.async, whose + complete_tx releases at cluster scope.""" + from cutlass._mlir.dialects import llvm as _llvm + from cutlass._mlir.extras import types as _T + + done = cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [mbar.toint(loc=loc, ip=ip).ir_value(), cutlass.Int32(parity).ir_value(loc=loc, ip=ip)], + "{\n\t.reg .pred p;\n\tmbarrier.test_wait.parity.acquire.cluster.shared::cta.b64 p, [$1], $2;\n\t" + "selp.u32 $0, 1, 0, p;\n\t}", "=r,r,r,~{memory}", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + return done != cutlass.Int32(0) + + +@cute.jit +def _poll_logit_parts(arr, idx, stride: cutlass.Constexpr[int]): + """Spin (no back-off) until none of the SPLIT 4-word vectors at idx + q * stride is empty, every sweep loading all + of them at once. Returns the 4 logits, each the fp32 sum of its partials in rank order from +0.0 (the owner's + order; a -0.0 partial pushed as +0.0 cannot change a sum that starts at +0.0), -0.0 as +0.0 (the owner's push).""" + empty = cutlass.Int32(EMPTY_WORD) + s0 = empty + s1 = empty + s2 = empty + s3 = empty + pending = cutlass.Boolean(True) + while pending: + vecs = [] + for q in cutlass.range_constexpr(SPLIT): + vecs.append( + arr.load( + idx=idx + cutlass.Int32(q * stride), + vector_size=4, + alignment=16, + is_volatile=True, + ) + ) + missing = cutlass.Boolean(False) + totals = [cutlass.Float32(0.0)] * 4 + for q in cutlass.range_constexpr(SPLIT): + for j in cutlass.range_constexpr(4): + word = cutlass.Int32(vecs[q][j]) + missing = missing | (word == empty) + totals[j] = totals[j] + word.bitcast(cutlass.Float32) + s0 = _ag._sanitize_f32(totals[0].bitcast(cutlass.Int32)) + s1 = _ag._sanitize_f32(totals[1].bitcast(cutlass.Int32)) + s2 = _ag._sanitize_f32(totals[2].bitcast(cutlass.Int32)) + s3 = _ag._sanitize_f32(totals[3].bitcast(cutlass.Int32)) + pending = missing + return s0, s1, s2, s3 + + +def _bf16_rn(v): + """fp32 -> bf16 precision (round to nearest even), kept as fp32, in the integer domain.""" + u = v.bitcast(cutlass.Int32) + u = u + (((u >> 16) & 1) + 0x7FFF) + u = (u >> 16) << 16 + return u.bitcast(cutlass.Float32) + + +@dsl_user_op +def _atom_add_shared(addr_u32, val, *, loc=None, ip=None): + """atom.shared.add.u32; returns the old value.""" + from cutlass._mlir.dialects import llvm as _llvm + from cutlass._mlir.extras import types as _T + + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [cutlass.Int32(addr_u32).ir_value(loc=loc, ip=ip), cutlass.Int32(val).ir_value(loc=loc, ip=ip)], + "atom.shared.add.u32 $0, [$1], $2;", "=r,r,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _add_opaque(a, b, *, loc=None, ip=None): + """add.s32 in inline PTX: the compiler cannot re-associate it, so a tree of them stays a tree (LLVM turns a + tree of plain adds of 0/1 selects into one dependent chain of conditional increments).""" + from cutlass._mlir.dialects import llvm as _llvm + from cutlass._mlir.extras import types as _T + + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [cutlass.Int32(a).ir_value(loc=loc, ip=ip), cutlass.Int32(b).ir_value(loc=loc, ip=ip)], + "add.s32 $0, $1, $2;", "=r,r,r", has_side_effects=False, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +def _tree_sum(vals): + """Sum of a list of Int32 as a balanced tree of independent adds.""" + while len(vals) > 1: + vals = [_add_opaque(vals[i], vals[i + 1]) for i in range(0, len(vals) - 1, 2)] + ( + [vals[-1]] if len(vals) % 2 else [] + ) + return vals[0] + + +@cute.jit +def _top16_cta( + s_key, s_sigmoid, s_top2, s_cnt, s_cand, s_csig, s_win, s_wsig, tx, routed_scaling_factor +): + """Top-16 of one token by all THREADS threads of the route CTA: the experts, their order and the weight bits of + ``_rq.top16_warp`` (key descending, ties to the lower expert id; the weights from the same fp32 sum and fp64 + division), in two CTA barriers instead of one warp's 16 dependent redux rounds. + + 1. Thread t holds the keys of experts t and t + THREADS; each warp takes its two largest keys (two experts). + 2. Every warp computes B, the 16th largest of those 28 keys: at least 16 experts have a key >= B, so the 16th + largest key of all is >= B and every winner has a key >= B. + 3. The C experts with a key >= B (C >= 16; ~16-30 for router keys) go to s_cand as (key << 32) | (NO_ID - id), + whose signed order is the selection order, and their sigmoids to s_csig; one shared atomic per warp hands out + the slots. + 4. Warp 0 ranks the candidates (lane i: how many beat candidate i). Every expert that beats a winner is a + candidate, so a winner's rank is exact; the rank-r winner goes to slot r. With C > CAND_SLOTS (heavy ties), + warp 0 runs ``_rq.top16_warp`` instead. + Returns (expert id, weight bf16 bits) in lanes 0-15 of warp 0 (the rank-lane expert); other lanes hold garbage. + Shared memory: s_top2 [32] (warp w's keys at w and 16 + w), s_cnt [1], s_cand [CAND_SLOTS] (Int64), s_csig + [CAND_SLOTS], s_win and s_wsig [16]; every word is written and read between this call's barriers and the caller's + barriers around it. Counts are balanced trees of adds: a chain of 28 or 32 dependent adds cost ~0.25 us each.""" + lane = tx % cutlass.Int32(32) + warp = tx // cutlass.Int32(32) + key_a = s_key.load(idx=tx) + key_b = s_key.load(idx=tx + cutlass.Int32(THREADS)) + sig_a = s_sigmoid.load(idx=tx) + sig_b = s_sigmoid.load(idx=tx + cutlass.Int32(THREADS)) + a_high = key_a > key_b + hi = cutlass.Int32(cutlass.select_(a_high, key_a, key_b)) + lo = cutlass.Int32(cutlass.select_(a_high, key_b, key_a)) + top1 = prims.redux_sync(hi, prims.ReductionKind.MAX, _rq.FULL_MASK) + holders = prims.vote_sync(_rq.FULL_MASK, hi == top1, prims.VoteSync.BALLOT) + # The lowest lane holding the largest key gives it up for its other key. + gives = (cutlass.Int32(1) << lane) == (holders & (cutlass.Int32(0) - holders)) + top2 = prims.redux_sync( + cutlass.Int32(cutlass.select_(gives, lo, hi)), prims.ReductionKind.MAX, _rq.FULL_MASK + ) + if lane == cutlass.Int32(0): + s_top2.store(top1, idx=warp) + s_top2.store(top2, idx=warp + cutlass.Int32(16)) + if tx == cutlass.Int32(0): + s_cnt.store(cutlass.Int32(0), idx=0) + if tx < cutlass.Int32(CAND_SLOTS): + s_cand.store(cutlass.Int64(CAND_EMPTY), idx=tx) + cute.arch.barrier() + + # B: the smallest of the 28 keys with at most 15 of them above it. + valid = (lane % cutlass.Int32(16)) < cutlass.Int32(THREADS // 32) + mine = cutlass.Int32(cutlass.select_(valid, s_top2.load(idx=lane), cutlass.Int32(KEY_MAX))) + greater = [] + for i in cutlass.range_constexpr(8): + quad = s_top2.load(idx=cutlass.Int32(4 * i), vector_size=4, alignment=16) + for q in cutlass.range_constexpr(4): + if cutlass.const_expr((4 * i + q) % 16 < THREADS // 32): + greater.append(cutlass.Int32(cutlass.select_(cutlass.Int32(quad[q]) > mine, 1, 0))) + above = _tree_sum(greater) + bound = prims.redux_sync( + cutlass.Int32( + cutlass.select_(above <= cutlass.Int32(TOP_K - 1), mine, cutlass.Int32(KEY_MAX)) + ), + prims.ReductionKind.MIN, + _rq.FULL_MASK, + ) + in_a = key_a >= bound + in_b = key_b >= bound + ballot_a = prims.vote_sync(_rq.FULL_MASK, in_a, prims.VoteSync.BALLOT) + ballot_b = prims.vote_sync(_rq.FULL_MASK, in_b, prims.VoteSync.BALLOT) + count_a = cute.arch.popc(ballot_a) + count_w = count_a + cute.arch.popc(ballot_b) + first_slot = cutlass.Int32(0) + if lane == cutlass.Int32(0): + if count_w > cutlass.Int32(0): + first_slot = _atom_add_shared(s_cnt.data_ptr().toint(), count_w) + first_slot = cute.arch.shuffle_sync(first_slot, 0) + below = cutlass.Int32(cute.arch.lanemask_lt()) + if in_a: + slot_a = first_slot + cute.arch.popc(ballot_a & below) + if slot_a < cutlass.Int32(CAND_SLOTS): + s_cand.store( + (cutlass.Int64(key_a) << cutlass.Int64(32)) + | cutlass.Int64(cutlass.Int32(NO_ID) - tx), + idx=slot_a, + ) + s_csig.store(sig_a, idx=slot_a) + if in_b: + slot_b = first_slot + count_a + cute.arch.popc(ballot_b & below) + if slot_b < cutlass.Int32(CAND_SLOTS): + s_cand.store( + (cutlass.Int64(key_b) << cutlass.Int64(32)) + | cutlass.Int64(cutlass.Int32(NO_ID) - tx - cutlass.Int32(THREADS)), + idx=slot_b, + ) + s_csig.store(sig_b, idx=slot_b) + cute.arch.barrier() + + n_cand = s_cnt.load(idx=0) + expert = cutlass.Int32(0) + weight_bits = cutlass.Int16(0) + if n_cand <= cutlass.Int32(CAND_SLOTS): + if warp == cutlass.Int32(0): + cand = s_cand.load(idx=lane) + cand_id = cutlass.Int32(NO_ID) - cutlass.Int32(cand & cutlass.Int64(NO_ID)) + # An empty slot (id out of range, sigmoid stale) ranks >= 16, so neither is ever used. + cand_sig = s_csig.load(idx=lane) + beaten = [] + for i in cutlass.range_constexpr(CAND_SLOTS // 2): + pair = s_cand.load(idx=cutlass.Int32(2 * i), vector_size=2, alignment=16) + for q in cutlass.range_constexpr(2): + beaten.append( + cutlass.Int32(cutlass.select_(cutlass.Int64(pair[q]) > cand, 1, 0)) + ) + rank = _tree_sum(beaten) + if rank < cutlass.Int32(TOP_K): + s_win.store(cand_id, idx=rank) + s_wsig.store(cand_sig, idx=rank) + cute.arch.sync_warp() + r = lane % cutlass.Int32(TOP_K) + expert = s_win.load(idx=r) + # top16_warp's sum: an xor butterfly over the warp (offsets 16, 8, 4, 2, 1) with lanes 16-31 at 0, which + # leaves every lane with this tree (the partners' adds are the same fp32 adds). + sig = [] + for i in cutlass.range_constexpr(TOP_K // 4): + quad_s = s_wsig.load(idx=cutlass.Int32(4 * i), vector_size=4, alignment=16) + for q in cutlass.range_constexpr(4): + sig.append(cutlass.Float32(quad_s[q]) + cutlass.Float32(0.0)) + sum8 = [sig[j] + sig[j + 8] for j in range(8)] + sum4 = [sum8[j] + sum8[j + 4] for j in range(4)] + sum2 = [sum4[j] + sum4[j + 2] for j in range(2)] + weight = (cutlass.Float64(s_wsig.load(idx=r)) * routed_scaling_factor) / ( + cutlass.Float64(sum2[0] + sum2[1]) + cutlass.Float64(1e-20) + ) + weight_bits = _rq.cvt_rn_bf16_f64(weight) + else: + if warp == cutlass.Int32(0): + expert, weight_bits = _rq.top16_warp(s_key, s_sigmoid, lane, routed_scaling_factor) + return expert, weight_bits + + +@cute.kernel +def k3_moe_front_kernel( + tma_desc_w: cutlass.GridConstant[cuda.TensorMap], # W [N, K] bf16, 5-D, one call per k-tile + tma_desc_w64: cutlass.GridConstant[ + cuda.TensorMap + ], # the same W with a 64-row box (head half-tiles) + tma_desc_x: cutlass.GridConstant[cuda.TensorMap], # x [M, K] bf16, box 64 x 8 + bias: cutlass.Array, # fp32 [896] + buf_uc: cutlass.Array, # int32 words of this rank's head buffers + buf_mc: cutlass.Array, # int32 words of their multicast mapping + flags: cutlass.Array, # int32 [0] buffer of this call, [1] unused (0), [2] epoch, [3] CTAs that read [0] + ready: cutlass.Array, # int32 [16] per-token ready words (publish) + topk_ids: cutlass.Array, # int32 [M * 16] + topk_weight_bits: cutlass.Array, # int16 view of bf16 [M * 16] + quant_words: cutlass.Array, # int32 view of e4m3 [M, 3584] + scales: cutlass.Array, # uint8 [M * 112] + shared_bits: cutlass.Array, # int16 view of bf16 [M, shared_cols] + num_tokens: cutlass.Int32, + tp_rank: cutlass.Int32, + routed_scaling_factor: cutlass.Float64, + world: cutlass.Constexpr[int], + shared_cols: cutlass.Constexpr[int], + n_head_tiles: cutlass.Constexpr[int], + n_tiles: cutlass.Constexpr[int], + gemv_clusters: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + ring: cutlass.Constexpr[int], + gate_cap: cutlass.Constexpr[float], + linear_cap: cutlass.Constexpr[float], + publish: cutlass.Constexpr[bool], + head_half: cutlass.Constexpr[bool], +): + WL = HIDDEN_SIZE // world + WE = NUM_EXPERTS // world + LV = WL // 8 + EV = WE // 4 + SLOT = (LV + EV) * 4 + BUF = M_MAX * world * SLOT + PBASE = ( + _ag.BUFFERS * BUF + ) # the router logits' split-K partials, after the buffers: [buffer][token][rank][q][WE] + k_tiles = k_in // CTA_K + my_tiles = k_tiles // SPLIT + rounds_max = (n_tiles + gemv_clusters - 1) // gemv_clusters + # head_half (one round): clusters < n_head_tiles each own a 64-row head half-tile and hold all its k-tiles; the + # shared tiles' weight rows start after the 128-row-padded head, hence the row shift for tiles >= n_head_tiles. + row_shift, nbar = _half_plan(head_half, world, n_head_tiles, ring, my_tiles) + + tx, _, _ = cute.arch.thread_idx() + bx, _, _ = cute.arch.block_idx() + warp_id = cute.arch.warp_idx() + lane = tx % 32 + crank = cute.arch.block_idx_in_cluster() + # The role clusters come first, so they are resident before any GEMV cluster (they spin; the GEMV CTAs never + # wait on them). + role_or_gemv = bx // cutlass.Int32(SPLIT) + is_gemv = role_or_gemv >= cutlass.Int32(ROLE_CLUSTERS) + cluster = role_or_gemv - cutlass.Int32(ROLE_CLUSTERS) + tma_ptr_w = tma_desc_w.get_ptr() + tma_ptr_w64 = tma_desc_w64.get_ptr() + # A GEMV CTA of a head half-tile (head_half: one round, the head in 64-row tiles on the first clusters). + head_cta = cutlass.Boolean(False) + stream_cta = cutlass.Boolean(True) # its k-tiles stream through the ring + if cutlass.const_expr(head_half): + head_cta = cluster < cutlass.Int32(n_head_tiles) + stream_cta = cluster >= cutlass.Int32(n_head_tiles) + tma_ptr_x = tma_desc_x.get_ptr() + + smem_a = cutlass.Array( + io_dtype, ring * CTA_M * CTA_K, space=cutlass.AddressSpace.smem, alignment=1024 + ) + smem_b = cutlass.Array( + io_dtype, my_tiles * MMA_N * CTA_K, space=cutlass.AddressSpace.smem, alignment=1024 + ) + tma_full = cutlass.Array(cutlass.Int64, nbar, space=cutlass.AddressSpace.smem, alignment=8) + mma_done = cutlass.Array(cutlass.Int64, nbar, space=cutlass.AddressSpace.smem, alignment=8) + act_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + flags_ready = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + acc_done = cutlass.Array( + cutlass.Int64, MAX_ROUNDS, space=cutlass.AddressSpace.smem, alignment=8 + ) + mail_full = cutlass.Array( + cutlass.Int64, MAX_ROUNDS, space=cutlass.AddressSpace.smem, alignment=8 + ) + tmem_ptr_i32 = cutlass.Array(cutlass.Int32, 1, space=cutlass.AddressSpace.smem) + # [round][source rank][row of the owned slice][token] fp32 partials. + mailbox = cutlass.Array( + cutlass.Float32, + MAX_ROUNDS * SPLIT * ROWS_PER_WARP * MMA_N, + space=cutlass.AddressSpace.smem, + alignment=16, + ) + s_flags = cutlass.Array(cutlass.Int32, 4, space=cutlass.AddressSpace.smem, alignment=16) + # The owner warp's reduced rows [row][token], for the vector pushes. + s_rows = cutlass.Array( + cutlass.Float32, + OWNERS * ROWS_PER_WARP * MMA_N, + space=cutlass.AddressSpace.smem, + alignment=16, + ) + s_key = cutlass.Array(cutlass.Int32, NUM_EXPERTS, space=cutlass.AddressSpace.smem, alignment=16) + s_sigmoid = cutlass.Array( + cutlass.Float32, NUM_EXPERTS, space=cutlass.AddressSpace.smem, alignment=16 + ) + # _top16_cta's scratch (route CTAs). + s_top2 = cutlass.Array(cutlass.Int32, 32, space=cutlass.AddressSpace.smem, alignment=16) + s_cnt = cutlass.Array(cutlass.Int32, 4, space=cutlass.AddressSpace.smem, alignment=16) + s_cand = cutlass.Array(cutlass.Int64, CAND_SLOTS, space=cutlass.AddressSpace.smem, alignment=16) + s_csig = cutlass.Array( + cutlass.Float32, CAND_SLOTS, space=cutlass.AddressSpace.smem, alignment=16 + ) + s_win = cutlass.Array(cutlass.Int32, TOP_K, space=cutlass.AddressSpace.smem, alignment=16) + s_wsig = cutlass.Array(cutlass.Float32, TOP_K, space=cutlass.AddressSpace.smem, alignment=16) + + if is_gemv: + if warp_id == 0: + prims.prefetch_tensormap(tma_ptr_w) + if cutlass.const_expr(head_half): + prims.prefetch_tensormap(tma_ptr_w64) + prims.prefetch_tensormap(tma_ptr_x) + if prims.elect_sync(): + for s in cutlass.range_constexpr(nbar): + prims.mbarrier_init(tma_full.subview(s), 1) + prims.mbarrier_init(mma_done.subview(s), 1) + prims.mbarrier_init(act_full, 1) + prims.mbarrier_init(flags_ready, 1) + for r in cutlass.range_constexpr(MAX_ROUNDS): + prims.mbarrier_init(acc_done.subview(r), 1) + # Owner ranks (< 4): the other 7 ranks' partials of the owned 32 rows x 8 tokens arrive by + # st.async; the phase completes on their bytes (expected here, before the cluster forms). + prims.mbarrier_init(mail_full.subview(r), 1) + prims.mbarrier_arrive_expect_tx( + mail_full.subview(r), (SPLIT - 1) * ROWS_PER_WARP * MMA_N * 4 + ) + if warp_id == 2: + prims.tcgen05_alloc(tmem_ptr_i32, HALF_TMEM_COLS if head_half else TMEM_COLS) + prims.tcgen05_relinquish_alloc_permit() + prims.fence_mbarrier_init() + # Cluster formation: the peers' shared memory and barriers are addressable from here on. + prims.barrier_cluster_arrive_relaxed() + prims.barrier_cluster_wait() + prims.barrier_cta_sync(0) + if is_gemv: + tmem_ptr = cutlass.inttoptr(tmem_ptr_i32.load(), 6, cutlass.Int32) + if warp_id == 0: + # ================================================================= + # Weight TMA: fill the ring, prefetch every later k-tile of both + # rounds into L2, all before any wait; then refill as MMAs finish. + # ================================================================= + if prims.elect_sync(): + if head_cta: + # A head half-tile: every k-tile of the rank, 64-row boxes 16 KB apart; nothing loads later. + mh = cluster * cutlass.Int32(HALF_M) + for hi in cutlass.range_constexpr(my_tiles): + kh = crank + cutlass.Int32(hi * SPLIT) + prims.mbarrier_arrive_expect_tx( + tma_full.subview(hi), HALF_M * CTA_K * ELEM_BYTES + ) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a.subview(hi * HALF_M * CTA_K), + tma_ptr_w64, + (cutlass.Int32(0), mh, kh * cutlass.Int32(TMA_COPY_ITERS), cutlass.Int32(0), + cutlass.Int32(0)), + tma_full.subview(hi), + l2_cache_hint=EVICT_FIRST, + ) # fmt: skip + else: + m0 = cluster * cutlass.Int32(CTA_M) + cutlass.Int32(row_shift) + for i in cutlass.range_constexpr(ring): + k = crank + cutlass.Int32(i * SPLIT) + prims.mbarrier_arrive_expect_tx( + tma_full.subview(i), CTA_M * CTA_K * ELEM_BYTES + ) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a.subview(i * CTA_M * CTA_K), + tma_ptr_w, + ( + cutlass.Int32(0), + m0, + k * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), + cutlass.Int32(0), + ), + tma_full.subview(i), + l2_cache_hint=EVICT_FIRST, + ) + for rd in cutlass.range_constexpr(rounds_max): + tile = cluster + cutlass.Int32(rd * gemv_clusters) + if tile < cutlass.Int32(n_tiles): + first = ring if rd == 0 else 0 + for i in range(first, my_tiles): + k = crank + i * cutlass.Int32(SPLIT) + prims.cp_async_bulk_tensor_prefetch( + tma_ptr_w, + [cutlass.Int32(0), tile * cutlass.Int32(CTA_M) + cutlass.Int32(row_shift), + k * cutlass.Int32(TMA_COPY_ITERS), cutlass.Int32(0), cutlass.Int32(0)], + [], + ) # fmt: skip + # Dependents may launch now; they wait for this whole grid before reading its outputs. + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + if prims.elect_sync(): + if stream_cta: + stage = cutlass.Int32(0) + phase = cutlass.Int32(0) + seq = cutlass.Int32(0) # k-tiles issued so far, over the rounds + for rd in cutlass.range_constexpr(rounds_max): + tile = cluster + cutlass.Int32(rd * gemv_clusters) + if tile < cutlass.Int32(n_tiles): + for i in range(my_tiles): + if seq >= cutlass.Int32(ring): + while not cute.arch.mbarrier_try_wait( + mma_done.subview(stage).data_ptr(), phase + ): + pass + k = crank + i * cutlass.Int32(SPLIT) + prims.mbarrier_arrive_expect_tx( + tma_full.subview(stage), CTA_M * CTA_K * ELEM_BYTES + ) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a.subview(stage * cutlass.Int32(CTA_M * CTA_K)), + tma_ptr_w, + (cutlass.Int32(0), tile * cutlass.Int32(CTA_M) + cutlass.Int32(row_shift), + k * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), cutlass.Int32(0)), + tma_full.subview(stage), + l2_cache_hint=EVICT_FIRST, + ) # fmt: skip + seq = seq + cutlass.Int32(1) + stage = stage + cutlass.Int32(1) + if stage == cutlass.Int32(ring): + stage = cutlass.Int32(0) + if seq > cutlass.Int32(ring): + phase = phase ^ cutlass.Int32(1) + elif warp_id == 1: + # ================================================================= + # Grid dependency, the head workspace's flags, then x's k-tiles. + # ================================================================= + prims.griddepcontrol(prims.GridDepAction.WAIT) + if prims.elect_sync(): + prims.mbarrier_arrive_expect_tx(act_full, my_tiles * MMA_N * CTA_K * ELEM_BYTES) + for i in range(my_tiles): + k = crank + i * cutlass.Int32(SPLIT) + for half in cutlass.range_constexpr(TMA_COPY_ITERS): + prims.cp_async_bulk_tensor_shared_cta_global( + smem_b.subview( + i * cutlass.Int32(MMA_N * CTA_K) + + cutlass.Int32(half * B_HALF_ELEMS) + ), + tma_ptr_x, + ( + k * cutlass.Int32(CTA_K) + cutlass.Int32(half * TMA_K_BOX), + cutlass.Int32(0), + ), + act_full, + ) + # The buffer index is for the epilogue's pushes, microseconds later: read it once x's loads are + # issued, so they do not wait for its round trip. + s_flags.store(flags.load(idx=0, is_volatile=True), idx=0) + _red_add_release(flags.data_ptr(3).toint(), cutlass.Int32(1)) + prims.mbarrier_arrive(flags_ready) + elif warp_id == 2: + # ================================================================= + # MMA: per round, the rank's k-tiles into that round's accumulator. + # ================================================================= + idesc = prims.Tcgen05InstrDesc.build( + c_dtype=cutlass.Float32, + a_dtype=io_dtype, + b_dtype=io_dtype, + n_dim=MMA_N, + m_dim=CTA_M, + ) + desc_a_base = prims.Tcgen05SmemDesc.build( + start_address=smem_a, leading_byte_offset=LEADING, stride_byte_offset=STRIDE, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + desc_b_base = prims.Tcgen05SmemDesc.build( + start_address=smem_b, leading_byte_offset=LEADING, stride_byte_offset=STRIDE, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + if head_cta: + # Before the grid wait: every A k-tile of the head half-tile into TMEM, one 128 x 16 copy per K = 16 + # step with the MMA's own A descriptor as the source (rows 64-127 alias the next box, as in smem). + for hc in cutlass.range(my_tiles, unroll=1): + while not cute.arch.mbarrier_try_wait(tma_full.subview(hc).data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + for kc in cutlass.range_constexpr(TMA_COPY_ITERS * K_BLOCKS_PER_HALF): + desc_c = desc_a_base + ( + hc * cutlass.Int32(STAGE_AH) + + cutlass.Int32( + (kc // K_BLOCKS_PER_HALF) * A_BOX_H + + (kc % K_BLOCKS_PER_HALF) * STEP + ) + ) + a_dst = cutlass.inttoptr( + tmem_ptr_i32.load() + + cutlass.Int32(A_TMEM_BASE + kc * A_KSTEP_COLS) + + hc * cutlass.Int32(A_TMEM_COLS), + 6, + cutlass.Int32, + ) + if prims.elect_sync(): + prims.tcgen05_cp(prims.Tcgen05CpShape.SHAPE_128X256B, a_dst, desc_c) + while not cute.arch.mbarrier_try_wait(act_full.data_ptr(), 0): + pass + if head_cta: + # A head half-tile: A from TMEM (staged before the wait), B (x) from shared memory. The MMA is M 128: + # rows 64-127 are the aliased garbage and land in accumulator rows nobody reads. + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc_h = cutlass.inttoptr(tmem_ptr_i32.load(), 6, cutlass.Int32) + for hi in cutlass.range(my_tiles, unroll=1): + for kh in cutlass.range_constexpr(TMA_COPY_ITERS * K_BLOCKS_PER_HALF): + a_src = cutlass.inttoptr( + tmem_ptr_i32.load() + + cutlass.Int32(A_TMEM_BASE + kh * A_KSTEP_COLS) + + hi * cutlass.Int32(A_TMEM_COLS), + 6, + cutlass.Int32, + ) + desc_bh = desc_b_base + ( + hi * cutlass.Int32(STAGE_B) + + cutlass.Int32( + (kh // K_BLOCKS_PER_HALF) * B_BOX + (kh % K_BLOCKS_PER_HALF) * STEP + ) + ) + accumulate_h = cutlass.Boolean(True) + if cutlass.const_expr(kh == 0): + accumulate_h = hi > cutlass.Int32(0) + if prims.elect_sync(): + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, prims.CTAGroup.CTA_1, acc_h, a_src, desc_bh, idesc, + accumulate_h, + ) # fmt: skip + if prims.elect_sync(): + prims.tcgen05_commit(acc_done.subview(0)) + else: + stage = cutlass.Int32(0) + phase = cutlass.Int32(0) + for rd in cutlass.range_constexpr(rounds_max): + tile = cluster + cutlass.Int32(rd * gemv_clusters) + if tile < cutlass.Int32(n_tiles): + acc = cutlass.inttoptr( + tmem_ptr_i32.load() + cutlass.Int32(rd * MMA_N), 6, cutlass.Int32 + ) + for i in range(my_tiles): + while not cute.arch.mbarrier_try_wait( + tma_full.subview(stage).data_ptr(), phase + ): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + for kb in cutlass.range_constexpr(TMA_COPY_ITERS * K_BLOCKS_PER_HALF): + box = kb // K_BLOCKS_PER_HALF + within = kb % K_BLOCKS_PER_HALF + desc_a = desc_a_base + ( + stage * cutlass.Int32(STAGE_A) + + cutlass.Int32(box * A_BOX + within * STEP) + ) + desc_b = desc_b_base + ( + i * cutlass.Int32(STAGE_B) + + cutlass.Int32(box * B_BOX + within * STEP) + ) + accumulate = cutlass.Boolean(True) + if cutlass.const_expr(kb == 0): + accumulate = i > cutlass.Int32(0) + if prims.elect_sync(): + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, prims.CTAGroup.CTA_1, acc, desc_a, desc_b, idesc, + accumulate, + ) # fmt: skip + if prims.elect_sync(): + prims.tcgen05_commit(mma_done.subview(stage)) + stage = stage + cutlass.Int32(1) + if stage == cutlass.Int32(ring): + stage = cutlass.Int32(0) + phase = phase ^ cutlass.Int32(1) + if prims.elect_sync(): + prims.tcgen05_commit(acc_done.subview(rd)) + elif warp_id >= 4 and warp_id < 8: + # ================================================================= + # Epilogue: per round, TMEM -> registers, push to / reduce at the + # row owner; the owner pushes the head slice or writes the shared + # activation. + # ================================================================= + w = warp_id - 4 # TMEM lanes 32 w ..: tile rows owned by rank w + while not cute.arch.mbarrier_try_wait(act_full.data_ptr(), 0): + pass + # flags_ready orders warp 1's store of the buffer index before this read. + while not cute.arch.mbarrier_try_wait(flags_ready.data_ptr(), 0): + pass + b = s_flags.load(idx=0) + slot_row = lane * cutlass.Int32(MMA_N) + # A head half-tile's rows are TMEM lanes 0-63 (warps 4-5); warps 6-7 hold garbage rows and sit out. + live = cutlass.Boolean(True) + if cutlass.const_expr(head_half): + live = stream_cta | (w < cutlass.Int32(2)) + for rd in cutlass.range_constexpr(rounds_max): + tile = cluster + cutlass.Int32(rd * gemv_clusters) + if (tile < cutlass.Int32(n_tiles)) & live: + while not cute.arch.mbarrier_try_wait(acc_done.subview(rd).data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc = prims.tcgen05_ld( + "32x32b", + cutlass.inttoptr( + tmem_ptr_i32.load() + cutlass.Int32(rd * MMA_N), 6, cutlass.Float32 + ), + num=MMA_N, + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + box_base = cutlass.Int32(rd * SPLIT * ROWS_PER_WARP * MMA_N) + tile_row = cutlass.select_( + head_cta, tile * cutlass.Int32(HALF_M), tile * cutlass.Int32(CTA_M) + ) + row0 = cutlass.Int32(tile_row) + w * cutlass.Int32(ROWS_PER_WARP) + srow0 = w * cutlass.Int32(ROWS_PER_WARP) + if (tile < cutlass.Int32(n_head_tiles)) & (row0 >= cutlass.Int32(WL)): + # Router rows: every rank of the cluster pushes its own split-K partial (fp32, -0.0 as +0.0) + # into slot [crank] of the partials region; the route CTAs add the 8 in rank order. + pe0 = row0 - cutlass.Int32(WL) + if pe0 < cutlass.Int32(WE): + for tt in cutlass.range_constexpr(MMA_N): + s_rows.store( + cutlass.Float32(acc[tt]), + idx=(srow0 + lane) * cutlass.Int32(MMA_N) + cutlass.Int32(tt), + ) + cute.arch.sync_warp() + # 8 router vectors (4 logits each) per token: two (token, vector) pairs per lane. + for ph in cutlass.range_constexpr(2): + ppair = lane + cutlass.Int32(ph * 32) + ptok = ppair // cutlass.Int32(8) + pvec = ppair % cutlass.Int32(8) + pe = pe0 + pvec * cutlass.Int32(4) + if ptok < num_tokens: + if pe < cutlass.Int32(WE): + pwords = [] + for pp in cutlass.range_constexpr(4): + pval = cutlass.Float32( + s_rows.load( + idx=( + srow0 + + pvec * cutlass.Int32(4) + + cutlass.Int32(pp) + ) + * cutlass.Int32(MMA_N) + + ptok + ) + ) + pwords.append( + _ag._sanitize_f32(pval.bitcast(cutlass.Int32)) + ) + pdst = ( + cutlass.Int32(PBASE) + + ( + (b * cutlass.Int32(M_MAX) + ptok) + * cutlass.Int32(world) + + tp_rank + ) + * cutlass.Int32(SPLIT * WE) + + crank * cutlass.Int32(WE) + + pe + ) + buf_mc.store( + (pwords[0], pwords[1], pwords[2], pwords[3]), + idx=pdst, + alignment=16, + ) + elif w != crank: + base = box_base + crank * cutlass.Int32(ROWS_PER_WARP * MMA_N) + slot_row + bar = _mapa_u32(mail_full.subview(rd).data_ptr(), w) + for h in cutlass.range_constexpr(MMA_N // 4): + _st_async_v4_f32( + _mapa_u32(mailbox.subview(base + cutlass.Int32(4 * h)).data_ptr(), w), + cutlass.Float32(acc[4 * h]), cutlass.Float32(acc[4 * h + 1]), + cutlass.Float32(acc[4 * h + 2]), cutlass.Float32(acc[4 * h + 3]), bar, + ) # fmt: skip + else: + # Completed by the peers' st.async: acquire at cluster scope. + while not _test_wait_cluster(mail_full.subview(rd).data_ptr(), 0): + pass + tot = [cutlass.Float32(0.0)] * MMA_N + for t in cutlass.range_constexpr(MMA_N): + total = cutlass.Float32(0.0) + for q in cutlass.range_constexpr(SPLIT): + part = mailbox.load( + idx=box_base + + cutlass.Int32(q * ROWS_PER_WARP * MMA_N) + + slot_row + + cutlass.Int32(t) + ) + total = total + cutlass.Float32( + cutlass.select_( + crank == cutlass.Int32(q), cutlass.Float32(acc[t]), part + ) + ) + tot[t] = total + if tile < cutlass.Int32(n_head_tiles): + # Head: stage the warp's 32 rows x 8 tokens, then 16-byte vectors to every rank. + for t in cutlass.range_constexpr(MMA_N): + s_rows.store( + tot[t], + idx=(w * cutlass.Int32(ROWS_PER_WARP) + lane) + * cutlass.Int32(MMA_N) + + cutlass.Int32(t), + ) + cute.arch.sync_warp() + # Latent rows (router rows push their partials above). + # 4 latent vectors (8 columns each) per token: lane = token * 4 + vector. + t = lane // cutlass.Int32(4) + v = lane % cutlass.Int32(4) + if t < num_tokens: + words = [] + for p in cutlass.range_constexpr(4): + c0 = srow0 + v * cutlass.Int32(8) + cutlass.Int32(2 * p) + lo = cutlass.Float32( + s_rows.load(idx=c0 * cutlass.Int32(MMA_N) + t) + ) + hi = cutlass.Float32( + s_rows.load( + idx=(c0 + cutlass.Int32(1)) * cutlass.Int32(MMA_N) + t + ) + ) + words.append(_ag._sanitize_bf16x2(_ag._pack_bf16x2(hi, lo))) + dst = b * cutlass.Int32(BUF) + ( + t * cutlass.Int32(world) + tp_rank + ) * cutlass.Int32(SLOT) + dst = dst + ((row0 // cutlass.Int32(8)) + v) * cutlass.Int32(4) + buf_mc.store( + (words[0], words[1], words[2], words[3]), + idx=dst, + alignment=16, + ) + else: + # Shared: lane l < 16 holds gate column j, lane l + 16 its up row. + col = ( + (tile - cutlass.Int32(n_head_tiles)) * cutlass.Int32(64) + + w * cutlass.Int32(16) + + lane + ) + for t in cutlass.range_constexpr(MMA_N): + up = cutlass.Float32(cute.arch.shuffle_sync_bfly(tot[t], offset=16)) + if lane < cutlass.Int32(16): + if cutlass.Int32(t) < num_tokens: + out = _bf16_rn( + _situ( + _bf16_rn(tot[t]), _bf16_rn(up), gate_cap, linear_cap + ) + ) + shared_bits.store( + cutlass.Int16(out.bitcast(cutlass.Int32) >> 16), + idx=cutlass.Int32(t * shared_cols) + col, + ) + prims.barrier_cta_sync(1, thread_count=128) + if warp_id == 4: + prims.tcgen05_dealloc(tmem_ptr, HALF_TMEM_COLS if head_half else TMEM_COLS) + else: + # ===================================================================== + # Role CTAs: route (CTA t < M) or quantize (CTA M + t) token t from + # every rank's pushed slot, as k3_route_quant_ag after its push. + # ===================================================================== + role = role_or_gemv * cutlass.Int32(SPLIT) + crank + if role < num_tokens * cutlass.Int32(2): + routes = role < num_tokens + tok = cutlass.select_(routes, role, role - num_tokens) + polls_logits = tx < cutlass.Int32(NUM_EXPERTS // 4) + e0 = cutlass.select_(polls_logits, tx * cutlass.Int32(4), cutlass.Int32(0)) + bias4 = bias.load(idx=e0, vector_size=4, alignment=16) + prims.griddepcontrol(prims.GridDepAction.WAIT) + if cutlass.const_expr(not publish): + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + if tx == 0: + s_flags.store(flags.load(idx=0, is_volatile=True), idx=0) + s_flags.store(flags.load(idx=2, is_volatile=True), idx=1) + cute.arch.barrier() + if cutlass.const_expr(publish): + # The dependent k3_moe (head_flags) advances the epoch at its last claim, so the epoch is read before + # this CTA's trigger, as in k3_route_quant_ag. + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + b = s_flags.load(idx=0) + epoch = s_flags.load(idx=1) + # A role CTA signs in as soon as it has read the buffer index (it reads nothing of flags after this). + # Token 0's quantization CTA hands the buffer over once every reader is in; its polls wait for the pushes, + # which come microseconds later. + if tx == 0: + _red_add_release(flags.data_ptr(3).toint(), cutlass.Int32(1)) + if role == num_tokens: + _flip( + flags, + b, + num_tokens * cutlass.Int32(2) + cutlass.Int32(gemv_clusters * SPLIT), + ) + tok_base = b * cutlass.Int32(BUF) + tok * cutlass.Int32(world * SLOT) + empty = cutlass.Int32(EMPTY_WORD) + if routes: + # The keys and the selection run once per layer, so their code is cold in the SM's instruction cache. + # Pass 0 runs the same code (one copy: a dynamic loop) on synthetic logits while the pushes are still + # on their way, its outputs discarded; pass 1 polls the pushes and runs it for real, from warm code. + for p in cutlass.range(2, unroll=1): + real = p == cutlass.Int32(1) + if polls_logits: + # Pass 0's logits: distinct values in [-4.5, 4.5), a permutation of the experts (e * 37 mod + # 896) that spreads the 16 largest over 8 warps as router logits spread, so pass 0 takes + # _top16_cta's common path (18 candidates; the bias is left out of pass 0's keys). + dry = [] + for q in cutlass.range_constexpr(4): + perm = ((e0 + cutlass.Int32(q)) * cutlass.Int32(37)) % cutlass.Int32( + NUM_EXPERTS + ) + dry.append( + ( + cutlass.Float32(perm) * cutlass.Float32(0.01) + - cutlass.Float32(4.5) + ).bitcast(cutlass.Int32) + ) + w0, w1, w2, w3 = dry + if real: + # The 8 split-K partials of these 4 logits, pushed by the pushing rank's GEMV CTAs. + addr = ( + cutlass.Int32(PBASE) + + ( + (b * cutlass.Int32(M_MAX) + tok) * cutlass.Int32(world) + + tx // cutlass.Int32(EV) + ) + * cutlass.Int32(SPLIT * WE) + + (tx % cutlass.Int32(EV)) * cutlass.Int32(4) + ) + w0, w1, w2, w3 = _poll_logit_parts(buf_uc, addr, WE) + for q in cutlass.range_constexpr(SPLIT): + buf_uc.store( + (empty, empty, empty, empty), + idx=addr + cutlass.Int32(q * WE), + alignment=16, + ) + words = [w0, w1, w2, w3] + for q in cutlass.range_constexpr(4): + sig = _rq.sigmoid_accurate(words[q].bitcast(cutlass.Float32)) + s_sigmoid.store(sig, idx=e0 + q) + key_bias = cutlass.Float32( + cutlass.select_(real, bias4[q], cutlass.Float32(0.0)) + ) + s_key.store(_rq.selection_key(sig + key_bias), idx=e0 + q) + cute.arch.barrier() + expert, weight_bits = _top16_cta( + s_key, + s_sigmoid, + s_top2, + s_cnt, + s_cand, + s_csig, + s_win, + s_wsig, + tx, + routed_scaling_factor, + ) + if tx < cutlass.Int32(32): + if real: + if tx < cutlass.Int32(TOP_K): + out = tok * cutlass.Int32(TOP_K) + tx + topk_ids.store(expert, idx=out) + topk_weight_bits.store(weight_bits, idx=out) + if cutlass.const_expr(publish): + cute.arch.fence_acq_rel_gpu() + if cutlass.const_expr(publish): + cute.arch.sync_warp() + if tx == cutlass.Int32(0): + _ag._store_release( + ready.data_ptr(tok).toint(), epoch + cutlass.Int32(1) + ) + # Pass 0's readers of the keys and of the selection's scratch are done before pass 1 rewrites them. + cute.arch.barrier() + else: + addr = ( + tok_base + + (tx // cutlass.Int32(LV)) * cutlass.Int32(SLOT) + + (tx % cutlass.Int32(LV)) * cutlass.Int32(4) + ) + w0, w1, w2, w3 = _ag._poll4(buf_uc, addr) + buf_uc.store((empty, empty, empty, empty), idx=addr, alignment=16) + q_lo, q_hi, sf_byte = _rq.mxfp8_quant_vec8([w0, w1, w2, w3]) + quant_words.store( + (q_lo, q_hi), + idx=tok * cutlass.Int32(HIDDEN_SIZE // 4) + tx * cutlass.Int32(2), + alignment=8, + ) + if tx % cutlass.Int32(SF_VEC_SIZE // 8) == cutlass.Int32(0): + scales.store( + cutlass.Uint8(sf_byte), + idx=tok * cutlass.Int32(HIDDEN_SIZE // SF_VEC_SIZE) + + tx // cutlass.Int32(SF_VEC_SIZE // 8), + ) + if cutlass.const_expr(publish): + prims.fence_proxy("async_global") + cute.arch.fence_acq_rel_gpu() + cute.arch.barrier() + if tx == cutlass.Int32(0): + _ag._store_release( + ready.data_ptr(tok + cutlass.Int32(M_MAX)).toint(), + epoch + cutlass.Int32(1), + ) + cute.arch.barrier() + else: + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + + +@dsl_user_op +def _red_add_release(addr_i64, val, *, loc=None, ip=None): + """red.release.gpu.global.add.u32 (no round trip).""" + from cutlass._mlir.dialects import llvm as _llvm + + _llvm.inline_asm( + None, [addr_i64.ir_value(loc=loc, ip=ip), val.ir_value(loc=loc, ip=ip)], + "red.release.gpu.global.add.u32 [$0], $1;", "l,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _load_acquire(addr_i64, *, loc=None, ip=None): + """ld.acquire.gpu.global.u32.""" + from cutlass._mlir.dialects import llvm as _llvm + from cutlass._mlir.extras import types as _T + + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [addr_i64.ir_value(loc=loc, ip=ip)], "ld.acquire.gpu.global.u32 $0, [$1];", "=r,l", + has_side_effects=True, is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@cute.jit +def _flip(flags, b, readers): + """Hands the other buffer to the next call once all ``readers`` CTAs of this call that read the buffer index have + signed in (flags[3], each with a release right after its read), so none reads it any more. The next call reads the + index only after this grid has completed.""" + while _load_acquire(flags.data_ptr(3).toint()) < readers: + pass + flags.store(cutlass.Int32(0), idx=3) + flags.store(b ^ cutlass.Int32(1), idx=0) + + +def _weight_tensor_map(w, n_out, k_in, box_rows=CTA_M): + """W as five TMA dimensions (64-element column chunk, row, 64-element chunk index, 1, 1), box_rows rows a call.""" + return cuda.create_tensor_map_tiled( + global_address=w.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[TMA_K_BOX, n_out, k_in // TMA_K_BOX, 1, 1], + global_strides=[ + (k_in * ELEM_BYTES) // 16, + (TMA_K_BOX * ELEM_BYTES) // 16, + (n_out * k_in * ELEM_BYTES) // 16, + (n_out * k_in * ELEM_BYTES) // 16, + ], + box_dims=[TMA_K_BOX, box_rows, TMA_COPY_ITERS, 1, 1], + swizzle=cuda.TensorMapSwizzle.s128b, + ) + + +def _activation_tensor_map(x, cols, num_tokens): + """x [M, cols] as (cols, M) with an 8-row box: rows past num_tokens arrive as zeros.""" + return cuda.create_tensor_map_tiled( + global_address=x.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[cols, num_tokens], + global_strides=[(cols * ELEM_BYTES) // 16], + box_dims=[TMA_K_BOX, MMA_N], + swizzle=cuda.TensorMapSwizzle.s128b, + ) + + +@cute.jit +def k3_moe_front( + w: cute.Tensor, # [N, K] bf16 + x: cute.Tensor, # [M, K] bf16 + bias: cute.Tensor, + buf_uc: cute.Tensor, + buf_mc: cute.Tensor, + flags: cute.Tensor, + ready: cute.Tensor, + topk_ids: cute.Tensor, + topk_weight_bits: cute.Tensor, + quant_words: cute.Tensor, + scales: cute.Tensor, + shared_bits: cute.Tensor, + num_tokens: cutlass.Int32, + tp_rank: cutlass.Int32, + routed_scaling_factor: cutlass.Float64, + world: cutlass.Constexpr[int], + shared_cols: cutlass.Constexpr[int], + n_head_tiles: cutlass.Constexpr[int], + n_tiles: cutlass.Constexpr[int], + gemv_clusters: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + ring: cutlass.Constexpr[int], + gate_cap: cutlass.Constexpr[float], + linear_cap: cutlass.Constexpr[float], + publish: cutlass.Constexpr[bool], + use_pdl: cutlass.Constexpr[bool], + head_half: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + rows = _weight_rows(world, n_head_tiles, n_tiles, head_half) + tma_desc_w = _weight_tensor_map(w, rows, k_in) + tma_desc_w64 = _weight_tensor_map(w, rows, k_in, box_rows=HALF_M) + tma_desc_x = _activation_tensor_map(x, k_in, num_tokens) + k3_moe_front_kernel( + tma_desc_w, tma_desc_w64, tma_desc_x, bias, buf_uc, buf_mc, flags, ready, topk_ids, topk_weight_bits, + quant_words, scales, shared_bits, num_tokens, tp_rank, routed_scaling_factor, world, shared_cols, + n_head_tiles, n_tiles, gemv_clusters, k_in, ring, gate_cap, linear_cap, publish, head_half, + ).launch( + grid=((gemv_clusters + ROLE_CLUSTERS) * SPLIT, 1, 1), + block=(THREADS, 1, 1), + cluster=(SPLIT, 1, 1), + stream=stream, + use_pdl=use_pdl, + # One CTA per SM (shared memory): without it ptxas may cap registers for two and spill (bias4 did). + min_blocks_per_mp=1, + ) # fmt: skip diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_moe_kernel.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_moe_kernel.py new file mode 100644 index 000000000000..64f3d8a279fa --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_moe_kernel.py @@ -0,0 +1,3144 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: LicenseRef-NvidiaProprietary +# +# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual +# property and proprietary rights in and to this material, related +# documentation and any modifications thereto. Any use, reproduction, +# disclosure or distribution of this material and related documentation +# without an express license agreement from NVIDIA CORPORATION or +# its affiliates is strictly prohibited. +"""Kimi K3 routed-expert core for decode (M <= 8 tokens; up to 64 with m_max 64): one persistent CuTe DSL kernel. + +k3_moe reads the top-k of trtllm::kimi_k3_noaux_tc_mxfp8_quant (global expert ids and +weights) and its MXFP8 activations, and computes this rank's routed partial [M, 3584]: +- With the route + quant folded in (fold): the kernel reads the router logits and the bf16 + latent instead. Every CTA computes the top-16 of every token itself (the sigmoid keys in + shared memory aliased on the weight ring, one warp per token selecting, with the device + code of trtllm::k3_route_quant, so the ids and weights are the same bits), and quantizes + the latent to MXFP8 into its own scratch rows, which its activation producer gathers. +- Prologue: every CTA derives the same (local expert, <= 8 token slots) groups in shared + memory: a token bitmask per local expert, a ballot prefix over the experts present, and + each token's slots in top-k order for the combine. +- FC1 (gate_up, MXFP4 x MXFP8) + SiTU + MXFP8 requantization -> FC2 (down) through an + FC1->FC2 Lamport handoff (the intermediate slab holds FP8 -0.0 / E8M0 NaN sentinels until + FC1 writes it) or, with FC2_SYNC counter, per-group counts of the FC1 tiles done (hint: + both, the count only deciding when to load), then a deterministic routing-weighted combine. +- The last FC2 tile that reads a group's intermediate re-arms it, so the slab is armed + again when the kernel ends; the tile counters reset themselves the same way. +- The m_max 64 build (wide decode steps, e.g. R x 8 speculative verify tokens) gives an + expert with t tokens ceil(t / 8) groups of consecutive tokens (64-bit token masks; the + group's slots name their (token, top-k) pairs), keeps each FC2 task's per-token sums in + TMEM columns, and combines each m-tile in 8-token chunk tasks. No fold, fused all-reduce, + head flags or latent slab there. +- With the fused all-reduce (ar_world > 0), the partial is not stored: each m-tile's + reducer pushes its bf16 rows into every rank's Lamport buffer through a multicast + mapping; once a CTA has no tile left, its epilogue takes reduction tasks (one m-tile + each, from their own cursor, one at a time), waits for all ranks' rows, sums them in rank + order (as the MNNVL one-shot all-reduce does) into the output and empties the buffer. + Two buffers alternate between calls; the grid's last task claim flips a flag. + +Weights are read in place in the trtllm-gen W4A8_MXFP4_MXFP8 layout +(MXFP4WeightTRTLLMGenFusedMoEMethod): + w3_w1_weight [E, 2I, H/2], rows [up ; gate], interleaved (2i = up_i, 2i+1 = + gate_i), then shuffled in 32-row blocks: physical row p holds source row + 4*(p%8) + p//8 of its block + w2_weight [E, H, I/2], same 32-row block shuffle + *_weight_scale: same row order, then block_scale_interleave (128x4), which is + exactly the SF-atom layout the tcgen05 block-scaled MMA consumes. +So FC1 epilogue lane L of warp w in m-tile t holds intermediate column +64t + 16w + 2(L%8) + L//16 (up rows: bit 3 of L clear; the gate row is lane L^8), +and FC2 lane L holds output channel 128t + 32w + 4(L%8) + L//8. + +Tile order: FC1 tiles of every group, then FC2 tiles. With the dynamic queue (default) +tiles are claimed from one global cursor, so a CTA only ever waits for FC1 tiles already +claimed by a running CTA: the kernel makes progress with any number of resident CTAs, +e.g. while other streams hold SMs. The static schedule (tile = CTA + k * grid) is faster +when every CTA is resident. + +Launched with programmatic dependent launch: barrier setup and the TMEM allocation run +before griddepcontrol.wait, which precedes every read of the producer's outputs and every +global write. + +Configuration is per module instance (the shapes are trace-time constants): the +loader injects K3_CONFIG = {"i_tp": ..., "num_ctas": ..., "num_local": ..., ...} before +executing the module; options it leaves out take their defaults. +""" + +from __future__ import annotations + +import os + +import cuda.bindings.driver as cuda_driver +import cutlass +import cutlass.cute as cute +import cutlass.experimental.cuda as cuda +from cutlass import dsl_user_op +from cutlass.experimental import primitives as prims +from cutlass.experimental.cuda.tensor_map import TensorMapDataType + +_CFG = globals().get("K3_CONFIG") or {} + + +def _cfg(key: str, default): + """A kernel option from the op's configuration (K3_CONFIG), else its default.""" + return type(default)(_CFG.get(key, default)) + + +# ============================================================================= +# Problem shape and tunables (trace-time constants). +# ============================================================================= +N = 8 # token slots per group (MMA N) +H = 3584 # latent hidden = FC1 K = FC2 M +TOP_K = 16 +# Tokens per call: 8 (decode), or up to 64 in the wide build (16, 32 or 64), which groups an +# expert's tokens 8 at a time and keeps the FC2 per-token sums in TMEM. +M_MAX = int(_cfg("m_max", 8)) +WIDE = M_MAX > N +assert M_MAX == N or (WIDE and M_MAX % N == 0 and M_MAX <= 64), M_MAX +I_TP = int(_cfg("i_tp", 3072)) +assert I_TP % 128 == 0, "I_TP must be a multiple of 128" +NUM_LOCAL = int(_cfg("num_local", 256)) # this rank's experts +SITU_GATE_CAP = float(_cfg("gate_cap", 4.0)) +SITU_LINEAR_CAP = float(_cfg("linear_cap", 25.0)) +SF_RECIPE = str(_cfg("sf_recipe", "ceil")) +assert SF_RECIPE in ("ceil", "ocp") +# Tile schedule. dynamic: every tile is claimed from one global cursor, in queue order. +# static: tile = CTA + k * grid. hybrid: each CTA's first tile is its own index, the rest are +# claimed. Only the dynamic queue makes progress with any number of resident CTAs. +SCHED = str(_cfg("sched", "dynamic" if int(_cfg("dyn_all", 1)) else "static")) +assert SCHED in ("dynamic", "hybrid", "static"), SCHED +DYN_ALL = SCHED != "static" # tiles (all, or all but the first) come through the claim ring +HYBRID = SCHED == "hybrid" +PRECLAIM_T = 1 if HYBRID else 0 # the tile whose claim is issued in the prologue +USE_PDL = bool(int(_cfg("pdl", 0 if os.environ.get("TRTLLM_ENABLE_PDL") == "0" else 1))) +# PDL trigger: right after griddepcontrol.wait (0), so the next kernel's CTAs take SMs as this +# kernel's CTAs exit and may stream their weights while its last tasks run; or at CTA exit (1). +LATE_TRIGGER = bool(int(_cfg("late_trigger", 0))) +# All-reduce of the routed partial in the kernel (0: off, else the group size): each FC2 +# m-tile's reducer pushes its bf16 rows into every rank's Lamport buffer through the +# multicast mapping, and 28 reduction tasks (one per m-tile), taken by CTAs with no tile +# left, reduce the ranks' rows in rank order into the output. +AR_WORLD = int(_cfg("ar_world", 0)) +FUSED_AR = AR_WORLD > 0 +# Push only (edge E7): the m-tiles' rows still go to every rank's buffer (half flags[0] & 1), but no reduction tasks +# run, the output is not written and nothing here writes the flags word: the consumer (sandwich (b)) sums the +# ranks in the reducers' order, empties the half it read and owns the word as its call count. +AR_PUSH_ONLY = bool(int(_cfg("ar_push_only", 0))) +assert not AR_PUSH_ONLY or FUSED_AR, "push-only is a mode of the fused all-reduce" +# The reduction tasks spin on other GPUs; with the static schedule, CTAs spinning there could +# hold SMs that this kernel's unscheduled tiles (whose pushes the peers wait for) need while +# another collective holds the rest. +assert SCHED == "dynamic" or not FUSED_AR, "the fused all-reduce needs the dynamic tile queue" +AR_POLL_NS = int(_cfg("ar_poll_ns", 64)) +# Poll loops back off with nanosleep between polls (0) or spin tightly (1): a sleeping warp can +# wake well after the value it waits for has landed. 2 also spins in the FC2 start delay. +SPIN = int(_cfg("spin", 0)) +assert SPIN in (0, 1, 2), SPIN +# Issue the prologue claim before griddepcontrol.wait: the queue is this layer's, reset by +# its previous call, which ended before the producer kernel (launched after it and waiting +# on it) triggered this launch. (The all-reduce buffer flag is read after the wait; its flip +# follows every CTA's last reduction task, not the tile queue.) +PREWAIT_CLAIM = bool(int(_cfg("prewait_claim", 1))) and DYN_ALL +# A CTA whose first task is an FC2 task claims its next task only once that task's MMAs are done +# (its last stage released). Claiming as soon as its loads are issued, such a CTA takes a second +# FC2 task ahead of the CTAs still issuing their FC1 tile and runs both back to back while those +# get none. Dynamic queue (the hybrid queue claims a CTA's second task up front), M <= 8 build. +FC2_FIRST_HOLD = bool(int(_cfg("fc2_first_hold", 1))) and SCHED == "dynamic" and not WIDE +# FC2 epilogue hand-off. last: each tile's arrival is an atomic round trip and the last +# arrival of an m-tile combines it (the last reader of a group re-arms it). designated: the +# tile of the last group combines each m-tile and the tile of the last m-tile re-arms each +# group, after waiting for the others' fire-and-forget arrivals; a tile only ever waits for +# tiles earlier in the queue, which every CTA processes in order. +EPI_SYNC = str(_cfg("epi_sync", "last")) +assert EPI_SYNC in ("last", "designated"), EPI_SYNC +# FC2 work split. tile: one task per (group, m-tile); its partial rows are stored and the +# m-tile's last (or designated) arrival combines them. slice: one task per (m-tile, slice of +# consecutive groups); the CTA accumulates its groups' rows per token in registers, stores one +# partial per task, and the m-tile's last slice sums the slices in slice order. The model runs +# slice, with or without the fused all-reduce; tile is a config choice. +FC2_MODE = str(_CFG.get("fc2", "slice")) +assert FC2_MODE in ("tile", "slice"), FC2_MODE +SLICE = FC2_MODE == "slice" +FC2_SLICES = int(_cfg("fc2_slices", 5)) +SLICE_MAX = 8 +# Wide build: at least one slice per FC2_GROUP_TARGET groups (up to WIDE_SLICE_MAX slices), so that the FC2 tasks +# left when the FC1 tiles run out stay short as the group count grows. +FC2_GROUP_TARGET = int(_cfg("fc2_group_target", 6)) +WIDE_SLICE_MAX = 16 +# Where the slice FC2 combines an m-tile and re-arms a slice's groups. inline: the m-tile's last +# slice task and the slice's m-tile-27 task do it after their own groups (and wait for the +# others). task: every FC2 task only publishes; 28 combine tasks and one re-arm task per slice +# follow the FC2 tasks in the queue, taken by CTAs whose FC2 work is done. +# last: every FC2 task's arrival is an atomic round trip on its m-tile's counter and the last +# arrival combines the m-tile at once; the re-arm tasks follow the FC2 tasks. +# chunks (the wide build's only mode): every FC2 task only publishes; 28 x ceil(M / 8) combine +# tasks, one per (m-tile, 8 tokens), then the re-arm tasks follow the FC2 tasks. +FC2_COMBINE = str(_cfg("fc2_combine", "chunks" if WIDE else "last")) +assert FC2_COMBINE in (("chunks",) if WIDE else ("inline", "task", "last")), FC2_COMBINE +SLICE_TASKS = SLICE and FC2_COMBINE in ("task", "last", "chunks") # re-arm tasks in the queue +COMBINE_TASKS = SLICE and FC2_COMBINE == "task" +COMBINE_LAST = SLICE and FC2_COMBINE == "last" +COMBINE_CHUNKS = SLICE and FC2_COMBINE == "chunks" +assert 1 <= FC2_SLICES <= SLICE_MAX, FC2_SLICES +# FC1 -> FC2 handoff. scan: FC2's activation stages are loaded as soon as the ring allows and the +# Lamport warps re-load them until their sentinels are gone. counter (slice FC2): the FC1 epilogues +# count their tiles per group with a release; the activation producer acquires a group's count +# before loading its intermediate once, and the re-arm tasks reset the counts instead of re-arming. +# hint (slice FC2): the counts are relaxed and only decide when a group is loaded; the scan still +# validates what was loaded (no fences in the FC1 epilogue). The model runs hint (the slice FC2's re-arm +# tasks and the dynamic queue); builds without them default to scan. +FC2_SYNC = str(_CFG.get("fc2_sync", "hint" if SLICE_TASKS and SCHED == "dynamic" else "scan")) +assert FC2_SYNC in ("scan", "hint", "counter"), FC2_SYNC +FC1_COUNTS = FC2_SYNC != "scan" # the FC1 epilogues count their tiles per group +SCAN_FC2 = FC2_SYNC != "counter" # the Lamport warps validate FC2's activation stages +assert not FC1_COUNTS or (SLICE_TASKS and SCHED == "dynamic"), ( + "the FC1 counts need the slice FC2's re-arm tasks and the dynamic queue" +) +# CTAs that drew fewer FC1 tiles start FC2 later, so FC2's loads do not compete with the FC1 +# tiles still streaming (worth 3-6 us at 11-32 groups with the slice FC2's scan). +FC2_DELAY_NS = int(_cfg("fc2_delay_ns", 0 if FC1_COUNTS else 1000)) +FC2_DELAY_LAG_NS = int(_cfg("fc2_delay_lag_ns", 0 if FC1_COUNTS else 4000)) +NUM_LAMPORT_WARPS = 4 +# Route + quant in the prologue: the kernel takes the router logits and the bf16 latent instead +# of trtllm::k3_route_quant's outputs (that kernel's device code, imported, so the same bits). +FOLD = bool(int(_cfg("fold", 0))) +# The weights producer starts once grouping phase 3 is done (groups_ready) instead of at the prologue's CTA-wide +# sync; not with FOLD, whose routing scratch lives in the A stages until that sync. +EARLY_RELEASE = bool(int(_cfg("early_release", 1))) and not FOLD +# The routing comes from trtllm::k3_route_quant_ag of the same call without a grid boundary: the kernel +# does not wait for that grid before its prologue; it acquires route_quant_ag's per-token ready words +# (ready[t] for the ids and weights, ready[8 + t] for the MXFP8 row) against this call's epoch (the head +# workspace's hflags[2], advanced by this kernel's last claim, which also re-arms the words of the tokens +# past the call's), and waits for the grid only before writing memory it does not own (the output). +HEAD_FLAGS = bool(int(_cfg("head_flags", 0))) +# With FUSED_AR (Track 5 edge E5): the reduced latent rows also go into the consumer's slab (int32 [3][8][H/2], the +# all-ones word empty; a computed all-ones word is published as 0x7FC07FC0) in buffer ``lat_buf`` (the call's +# ordinal mod 3), and this call re-arms buffer (lat_buf + 1) % 3, plus buffer 0 when ``lat_rearm0`` (the step's last +# call), once its grid wait has returned: the buffer's previous reader completed before this grid could start. +LAT_SLAB = bool(int(_cfg("lat_slab", 0))) +assert not LAT_SLAB or (FUSED_AR and not AR_PUSH_ONLY), ( + "the latent slab is written by the reduction tasks" +) +LAT_SLAB_BUFS = 3 +LAT_SLAB_EMPTY = -1 +assert not (HEAD_FLAGS and FOLD), ( + "the ready words come from route_quant_ag, which the fold replaces" +) +assert not HEAD_FLAGS or DYN_ALL, ( + "the head epoch advances, and the unused ready words are re-armed, at the queue's last claim" +) +# With HEAD_FLAGS, the PDL trigger: at launch (launch) or once the routing is acquired (ready), where +# the grid wait would have returned. +HEAD_TRIGGER = str(_cfg("head_trigger", "launch")) +assert HEAD_TRIGGER in ("launch", "ready"), HEAD_TRIGGER +HEAD_TRIGGER_READY = HEAD_FLAGS and HEAD_TRIGGER == "ready" +assert not WIDE or (SLICE and SCHED == "dynamic" and FC2_SYNC == "hint"), ( + "the wide build runs the slice FC2 on the dynamic queue with the hint handoff" +) +assert not WIDE or not (FOLD or FUSED_AR or HEAD_FLAGS), ( + "the wide build takes route_quant's outputs and returns the partial" +) +if FOLD: + from tensorrt_llm._torch.cute_dsl_kernels.k3_route_quant import k3_route_quant_kernel as _rq + + +def _device_sm_count() -> int: + (err,) = cuda_driver.cuInit(0) + assert err == cuda_driver.CUresult.CUDA_SUCCESS, err + err, ctx_dev = cuda_driver.cuCtxGetDevice() + if err != cuda_driver.CUresult.CUDA_SUCCESS: + err, ctx_dev = cuda_driver.cuDeviceGet(0) + err, n = cuda_driver.cuDeviceGetAttribute( + cuda_driver.CUdevice_attribute.CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT, ctx_dev + ) + assert err == cuda_driver.CUresult.CUDA_SUCCESS, err + return int(n) + + +NUM_CTAS = int(_CFG["num_ctas"]) if "num_ctas" in _CFG else _device_sm_count() +CLAIM_BASE = NUM_CTAS if HYBRID else 0 # queue index of the first claim + +a_dtype = cutlass.Float4E2M1FN +b_dtype = cutlass.Float8E4M3FN +sf_dtype = cutlass.Float8E8M0FNU +c_dtype = cutlass.BFloat16 + +a_smem_width = 8 +sf_vec_size = 32 +num_m0_per_sf_atom = 32 +num_m1_per_sf_atom = 4 +num_k_per_sf_atom = 4 +num_elts_atom_sf_e8 = num_m0_per_sf_atom * num_m1_per_sf_atom * num_k_per_sf_atom +num_elts_atom_sf_fp16 = num_elts_atom_sf_e8 // 2 +num_tmem_cols_per_sf_atom = 4 +smem_capacity = cutlass.memory.get_smem_capacity_in_bytes("sm_100") +num_mbar_bytes = 1024 +# Accumulator (8 columns) + SFA/SFB for every stage (2 x 12 x 4) fit in 128 columns, which +# leaves TMEM for other kernels' CTAs on the same SM. The wide build adds an FC2 task's per-token +# sums from column TOK_COL (one per token, then one scratch column for empty slots). +num_tmem_alloc_cols = 256 if WIDE else 128 +TOK_COL = 128 + +epilog_warp_id = (0, 1, 2, 3) +producer_weights_warp_id = 4 +producer_acts_warp_id = 5 +scales_tmem_warp_id = 6 +consumer_warp_id = 7 +lamport_acts_warp_id = 8 +EPI_THREADS = len(epilog_warp_id) * 32 +EPI_BAR_ID = 1 +GROUPING_BAR_ID = 2 +threads_per_cta = (8 + NUM_LAMPORT_WARPS) * 32 +NUM_WARPS = threads_per_cta // 32 +# The prologue grouping runs on every warp but the claimer, whose first atomic claim is in +# flight meanwhile. +GROUPING_WARPS = NUM_WARPS - 1 +GROUPING_THREADS = GROUPING_WARPS * 32 + +_LOG2E = 1.4426950408889634 +MX_BLOCK = 32 +SF_TMEM_COL = 4 +SF_TMEM_DP = 4 +E4M3_MAX = 448.0 +FP8_SENTINEL_I8 = -128 # 0x80 = FP8 -0.0 +SF_SENTINEL_I8 = -1 # 0xFF = E8M0 NaN (never produced) +C_ARM_WORD = -2139062144 # 0x80808080 +DYN_RING = 8 + +MMA_M, MMA_TILER_N, MMA_TILE_K, MMA_INST_K = 128, 8, 128, 32 +_mma_tiler_mnk = [MMA_M, MMA_TILER_N, MMA_TILE_K] +_mma_inst_mnk = [MMA_M, MMA_TILER_N, MMA_INST_K] +# SF atoms per tile / TMEM columns per MMA k-block for this tiling +# (compute_sf_rest of the CuTeDSL qmma example): one atom per tile in every dim. +REST_K_SF = REST_M_SF = REST_N_SF = 1 +NUM_TMEM_COLS_PER_KBLOCK_SFA = NUM_TMEM_COLS_PER_KBLOCK_SFB = 4 + +NUM_BYTES_A = MMA_M * MMA_TILE_K * a_smem_width // 8 +NUM_BYTES_A_GMEM = MMA_M * MMA_TILE_K * 4 // 8 +NUM_BYTES_B = MMA_TILER_N * MMA_TILE_K * b_dtype.width // 8 +NUM_BYTES_SFA = num_elts_atom_sf_e8 * sf_dtype.width // 8 +NUM_BYTES_SFB = num_elts_atom_sf_e8 * sf_dtype.width // 8 +SFB_GROUP_BYTES = num_tmem_cols_per_sf_atom * num_m1_per_sf_atom # 16 + +K1_TILES = H // MMA_TILE_K # 28 +K2_TILES = I_TP // MMA_TILE_K +NUM_KBLOCKS = MMA_TILE_K // MMA_INST_K +NUM_TMA_LOAD_BYTES_WEIGHTS = NUM_BYTES_A_GMEM + NUM_BYTES_SFA +# The expert weights and their scales are read once per call: their TMA loads are L2 evict-first +# (createpolicy.fractional.L2::evict_first, fraction 1.0), so the layer's stream does not push +# out the code, states and small weights the other kernels of the step reuse. evict_first 0 +# (A/B builds) loads them at normal priority. +EVICT_FIRST = 0x12F0000000000000 +W_L2_HINT = EVICT_FIRST if int(_cfg("evict_first", 1)) else None +NUM_TMA_LOAD_BYTES_ACTS_FC2 = NUM_BYTES_B + N * SFB_GROUP_BYTES +SFB_SRC_STRIDE_FC1 = H // sf_vec_size # 112 = the linear activation-scale row +SF_STRIDE0 = (I_TP // MX_BLOCK) * SF_TMEM_DP # I/8: per-(group, slot) intermediate-scale bytes +assert SF_STRIDE0 == K2_TILES * SFB_GROUP_BYTES + +M_TILES_FC1 = 2 * I_TP // MMA_M +M_TILES_FC2 = H // MMA_M # 28 +M_TILES_TOTAL = M_TILES_FC1 + M_TILES_FC2 +# FC2_SYNC counter / hint: FC2 loads a group once this many of its FC1 tiles have counted (a negative +# control of counter sets fewer, so FC2 can read an incomplete intermediate). +FC2_SYNC_NEED = int(_cfg("fc2_sync_need", M_TILES_FC1)) + + +def group_capacity(num_local_experts: int) -> int: + """Groups this rank can see in one decode step (M <= 8, top-16). Wide build: an expert with + t tokens has ceil(t / 8) <= 1 + (t - 1) / 8 groups, so P = M_MAX * 16 pairs on E present + experts make at most E + (P - E) / 8 groups, largest with E = min(experts, P).""" + if not WIDE: + return min(num_local_experts, M_MAX * TOP_K) + pairs = M_MAX * TOP_K + present = min(num_local_experts, pairs) + return present + (pairs - present) // N + + +G_CAP = group_capacity(NUM_LOCAL) +# Slice partials [m-tile][slice][token][128] in the partial buffer (G_CAP * N * H floats; the +# wide build's buffer holds exactly PART_ROWS rows of H). +S_CAP = WIDE_SLICE_MAX if WIDE else SLICE_MAX +PART_ROWS = S_CAP * M_MAX +assert WIDE or M_TILES_FC2 * S_CAP * M_MAX * MMA_M <= G_CAP * N * H +assert M_TILES_FC2 * MMA_M == H +NUM_CHUNKS = (NUM_LOCAL + 31) // 32 # 32 local experts per ballot +CHUNKS_PER_WARP = (NUM_CHUNKS + GROUPING_WARPS - 1) // GROUPING_WARPS +ROUTE_PAIRS = M_MAX * TOP_K # (token, top-k slot) pairs, one per prologue thread +# Wide build: the pairs per grouping thread, and the second word of each expert's token mask. +PAIR_ITERS = (ROUTE_PAIRS + GROUPING_THREADS - 1) // GROUPING_THREADS +MASK_HI = NUM_CHUNKS * 32 +PAIR_CHUNK = 16 # combine: pairs whose partial rows are loaded together +# A/B switches for the m-tile combine (batched: the local pairs compacted in the prologue, +# PAIR_CHUNK loads in flight; loop: per token, its slots one guarded load at a time) and for +# where the reducer resets its counters (before or after the epilogue barrier). +COMBINE = str(_cfg("combine", "batched")) +assert COMBINE in ("batched", "loop"), COMBINE +EPI_RESET = str(_cfg("epi_reset", "before")) +assert EPI_RESET in ("before", "after"), EPI_RESET +assert WIDE or ROUTE_PAIRS <= producer_weights_warp_id * 32 + +# Per-layer state words (int32). Every counter is back at zero when the kernel ends: +# [0] tile-queue cursor (reset by the last claim: every CTA claims until the queue is empty, +# so exactly one claim per CTA finds it empty), +# [4 + m] FC2 m-tile arrivals (reset by that m-tile's reducer), +# [32 + g] FC2 tiles done reading group g's intermediate (reset by the last one, which +# also re-arms that intermediate), +# [32 + G_CAP + g] FC1 tiles of group g done (FC2_SYNC counter or hint; reset by the slice's re-arm +# task). +ST_CURSOR = 0 +ST_AR_CURSOR = 1 +ST_MTILE = 4 +ST_GROUP = 32 +assert ST_MTILE + M_TILES_FC2 <= ST_GROUP +ST_FC1 = ST_GROUP + G_CAP +NUM_STATE = ST_FC1 + G_CAP + + +NUM_AB_STAGE = (smem_capacity - num_mbar_bytes - 2048 - 8192) // ( + NUM_BYTES_A + NUM_BYTES_B + NUM_BYTES_SFA + NUM_BYTES_SFB +) +NUM_AB_STAGE = (NUM_AB_STAGE // NUM_LAMPORT_WARPS) * NUM_LAMPORT_WARPS # 12 +assert not WIDE or 3 * MASK_HI * 4 <= NUM_BYTES_B * NUM_AB_STAGE # the wide prologue's scratch + +B_SCAN_BYTES_PER_LANE = NUM_BYTES_B // 32 +B_SCAN_VEC = 16 +B_SCAN_ITERS = B_SCAN_BYTES_PER_LANE // B_SCAN_VEC +SFB_SCAN_VEC = N * SFB_GROUP_BYTES // 32 +REARM_VEC4 = N * I_TP // 16 # 16-byte stores per group intermediate +REARM_SF_WORDS = N * K2_TILES # armed scale words per group + +# Fused all-reduce buffers: 2 x [M_MAX][AR_WORLD][H] bf16 per rank, as int32 words. A word +# of -0.0 (fp32 bits 0x80000000) means "not written"; pushes turn bf16 -0.0 into +0.0, so a +# written pair of bf16 values never has that pattern. The reducer puts the pattern back. +AR_BUFS = 2 +AR_BUF_WORDS = M_MAX * max(AR_WORLD, 1) * H // 2 +AR_EMPTY_WORD = -2147483648 # 0x80000000 +AR_RANK_CHUNK = 8 # ranks summed per chunk, as the MNNVL one-shot does for > 8 ranks +AR_TASKS = M_TILES_FC2 if FUSED_AR and not AR_PUSH_ONLY else 0 + +# Fold: routing scratch in the weight ring's A stages, which no TMA writes before the prologue's +# CTA-wide sync: keys (int32) and sigmoids (fp32) of [M_MAX][896], then the selected ids and +# weights (bf16 bits in int32) of [M_MAX][16]. +NUM_EXPERTS = 896 +KEY_ITERS = (NUM_EXPERTS + GROUPING_THREADS - 1) // GROUPING_THREADS # experts per grouping thread +RS_KEY = 0 +RS_SIG = RS_KEY + M_MAX * NUM_EXPERTS * 4 +RS_ID = RS_SIG + M_MAX * NUM_EXPERTS * 4 +RS_W = RS_ID + M_MAX * TOP_K * 4 +assert not FOLD or RS_W + M_MAX * TOP_K * 4 <= NUM_BYTES_A * NUM_AB_STAGE +# Fold: the quantization runs on the epilogue and Lamport warps, idle until their first tile; the +# activation producer waits for their arrival on QUANT_BAR before its first gather. +QUANT_BAR_ID = 3 +QUANT_THREADS = (len(epilog_warp_id) + NUM_LAMPORT_WARPS) * 32 +VEC8_PER_ROW = H // 8 +QUANT_ITERS = M_MAX * VEC8_PER_ROW // QUANT_THREADS +assert QUANT_ITERS * QUANT_THREADS == M_MAX * VEC8_PER_ROW and VEC8_PER_ROW % 64 == 0 + + +# ============================================================================= +# DSL helpers +# ============================================================================= +@dsl_user_op +def _read_globaltimer(*, loc=None, ip=None): + from cutlass._mlir.dialects import llvm as _llvm + from cutlass._mlir.extras import types as _T + + return cutlass.Int64( + _llvm.inline_asm( + _T.i64(), [], "mov.u64 $0, %globaltimer;", "=l", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _atomic_fetch_add(addr_i64, val, *, loc=None, ip=None): + """atom.acq_rel.gpu.add.u32 returning the old value.""" + from cutlass._mlir.dialects import llvm as _llvm + from cutlass._mlir.extras import types as _T + + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [addr_i64.ir_value(loc=loc, ip=ip), val.ir_value(loc=loc, ip=ip)], + "atom.acq_rel.gpu.global.add.u32 $0, [$1], $2;", "=r,l,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _red_release_add(addr_i64, val, *, loc=None, ip=None): + """red.release.gpu.global.add.u32 (no return value, no round trip).""" + from cutlass._mlir.dialects import llvm as _llvm + + _llvm.inline_asm( + None, [addr_i64.ir_value(loc=loc, ip=ip), val.ir_value(loc=loc, ip=ip)], + "red.release.gpu.global.add.u32 [$0], $1;", "l,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _red_relaxed_add(addr_i64, val, *, loc=None, ip=None): + """red.relaxed.gpu.global.add.u32 (no ordering, no round trip).""" + from cutlass._mlir.dialects import llvm as _llvm + + _llvm.inline_asm( + None, [addr_i64.ir_value(loc=loc, ip=ip), val.ir_value(loc=loc, ip=ip)], + "red.relaxed.gpu.global.add.u32 [$0], $1;", "l,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _load_relaxed(addr_i64, *, loc=None, ip=None): + """ld.relaxed.gpu.global.u32.""" + from cutlass._mlir.dialects import llvm as _llvm + from cutlass._mlir.extras import types as _T + + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [addr_i64.ir_value(loc=loc, ip=ip)], + "ld.relaxed.gpu.global.u32 $0, [$1];", "=r,l", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _load_acquire(addr_i64, *, loc=None, ip=None): + """ld.acquire.gpu.global.u32.""" + from cutlass._mlir.dialects import llvm as _llvm + from cutlass._mlir.extras import types as _T + + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [addr_i64.ir_value(loc=loc, ip=ip)], + "ld.acquire.gpu.global.u32 $0, [$1];", "=r,l", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _store_release(addr_i64, val, *, loc=None, ip=None): + """st.release.gpu.global.u32.""" + from cutlass._mlir.dialects import llvm as _llvm + + _llvm.inline_asm( + None, [addr_i64.ir_value(loc=loc, ip=ip), val.ir_value(loc=loc, ip=ip)], + "st.release.gpu.global.u32 [$0], $1;", "l,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _pack_bf16x2(hi, lo, *, loc=None, ip=None): + """(bf16(hi) << 16) | bf16(lo), round to nearest even.""" + from cutlass._mlir.dialects import llvm as _llvm + from cutlass._mlir.extras import types as _T + + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [hi.ir_value(loc=loc, ip=ip), lo.ir_value(loc=loc, ip=ip)], + "cvt.rn.bf16x2.f32 $0, $1, $2;", "=r,f,f", has_side_effects=False, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +def _bf16_lo(word): + return (word << cutlass.Int32(16)).bitcast(cutlass.Float32) + + +def _bf16_hi(word): + return (word & cutlass.Int32(-65536)).bitcast(cutlass.Float32) + + +@cute.jit +def _ar_word(cur, tok, rank, ch): + """Int32 word index of (buffer cur, token, rank, channel) in a fused all-reduce buffer.""" + return cur * AR_BUF_WORDS + ((tok * AR_WORLD + rank) * H + ch) // 2 + + +@cute.jit +def _ar_reduce_tile(ar_uc, y_words, cur, m_tile, tidx, num_tokens, ar_rank, lat_slab, lat_buf): + """One m-tile of the cross-rank reduction (epilogue warps): wait until every rank's + rows are in, sum them in rank order (chunks of AR_RANK_CHUNK ranks, like the MNNVL + one-shot), write the output and put the empty pattern back.""" + tok = tidx // 16 + ch0 = m_tile * MMA_M + (tidx % 16) * 8 + if tok < num_tokens: + base = _ar_word(cur, tok, cutlass.Int32(0), ch0) + # A task can be claimed long before its m-tile is done here: until this rank's own + # rows arrive, poll only them (one load per thread) and back off longer. + own = cutlass.Boolean(True) + while own: + v = ar_uc.load( + idx=base + ar_rank * (H // 2), vector_size=4, alignment=16, is_volatile=True + ) + dirty = cutlass.Boolean(False) + for q in cutlass.range_constexpr(4): + dirty = dirty | (v[q] == cutlass.Int32(AR_EMPTY_WORD)) + own = dirty + if own: + _backoff(4 * AR_POLL_NS) + # Poll every rank's row until none holds an empty word; sum the rows of that last + # poll: fp32, rank order, chunks of AR_RANK_CHUNK ranks summed from 0 and then added. + a0 = cutlass.Float32(0.0) + a1 = cutlass.Float32(0.0) + a2 = cutlass.Float32(0.0) + a3 = cutlass.Float32(0.0) + a4 = cutlass.Float32(0.0) + a5 = cutlass.Float32(0.0) + a6 = cutlass.Float32(0.0) + a7 = cutlass.Float32(0.0) + pending = cutlass.Boolean(True) + while pending: + dirty = cutlass.Boolean(False) + acc = [cutlass.Float32(0.0)] * 8 + for rb in cutlass.range_constexpr(0, AR_WORLD, AR_RANK_CHUNK): + chunk = [cutlass.Float32(0.0)] * 8 + for rr in cutlass.range_constexpr(min(AR_RANK_CHUNK, AR_WORLD - rb)): + v = ar_uc.load( + idx=base + (rb + rr) * (H // 2), + vector_size=4, + alignment=16, + is_volatile=True, + ) + for q in cutlass.range_constexpr(4): + w = cutlass.Int32(v[q]) + dirty = dirty | (w == cutlass.Int32(AR_EMPTY_WORD)) + chunk[2 * q] = chunk[2 * q] + _bf16_lo(w) + chunk[2 * q + 1] = chunk[2 * q + 1] + _bf16_hi(w) + for e in cutlass.range_constexpr(8): + acc[e] = acc[e] + chunk[e] + a0, a1, a2, a3, a4, a5, a6, a7 = acc + pending = dirty + if pending: + _backoff(AR_POLL_NS) + acc = [a0, a1, a2, a3, a4, a5, a6, a7] + words = ( + _pack_bf16x2(acc[1], acc[0]), + _pack_bf16x2(acc[3], acc[2]), + _pack_bf16x2(acc[5], acc[4]), + _pack_bf16x2(acc[7], acc[6]), + ) + y_words.store(words, idx=(tok * H + ch0) // 2, alignment=16) + if cutlass.const_expr(LAT_SLAB): + ones = cutlass.Int32(LAT_SLAB_EMPTY) + nan = cutlass.Int32(0x7FC07FC0) + lat_slab.store( + tuple(cutlass.select_(w == ones, nan, w) for w in words), + idx=lat_buf * cutlass.Int32(M_MAX * (H // 2)) + (tok * H + ch0) // 2, + alignment=16, + ) + empty = cutlass.Int32(AR_EMPTY_WORD) + for r in cutlass.range_constexpr(AR_WORLD): + ar_uc.store((empty, empty, empty, empty), idx=base + r * (H // 2), alignment=16) + + +@cute.jit +def _emit_row(acc, tok, ch, lane, warp_idx, m_tile, y_tensor, ar_mc, ar_cur, ar_rank): + """Epilogue: token tok's combined value of channel ch. Without the fused all-reduce it + is stored; with it, lanes p, p+8, p+16, p+24, p+1, ... hold channels 4p .. 4p+7 of the + warp's 32, so lanes 0, 2, 4, 6 gather them and push 16 bytes each (-0.0 as +0.0).""" + if cutlass.const_expr(FUSED_AR): + bits = _pack_bf16x2(cutlass.Float32(0.0), acc) & cutlass.Int32(0xFFFF) + bits = cutlass.select_(bits == cutlass.Int32(0x8000), cutlass.Int32(0), bits) + b1 = cute.arch.shuffle_sync_down(bits, 8) + b2 = cute.arch.shuffle_sync_down(bits, 16) + b3 = cute.arch.shuffle_sync_down(bits, 24) + b4 = cute.arch.shuffle_sync_down(bits, 1) + b5 = cute.arch.shuffle_sync_down(bits, 9) + b6 = cute.arch.shuffle_sync_down(bits, 17) + b7 = cute.arch.shuffle_sync_down(bits, 25) + if (lane < 8) & (lane % 2 == 0): + ar_mc.store( + ( + bits | (b1 << cutlass.Int32(16)), + b2 | (b3 << cutlass.Int32(16)), + b4 | (b5 << cutlass.Int32(16)), + b6 | (b7 << cutlass.Int32(16)), + ), + idx=_ar_word(ar_cur, tok, ar_rank, m_tile * MMA_M + warp_idx * 32 + 4 * lane), + alignment=16, + ) + else: + y_tensor.store(acc.to(c_dtype), idx=tok * H + ch) + + +@cute.jit +def _wide_slices(num_groups): + """The wide build's FC2 slice count for num_groups groups: at least min(groups, FC2_SLICES) and one per + FC2_GROUP_TARGET groups, at most S_CAP, at least 1.""" + s = cutlass.select_( + num_groups < cutlass.Int32(FC2_SLICES), num_groups, cutlass.Int32(FC2_SLICES) + ) + per = (num_groups + cutlass.Int32(FC2_GROUP_TARGET - 1)) // cutlass.Int32(FC2_GROUP_TARGET) + s = cutlass.select_(s < per, per, s) + s = cutlass.select_(s > cutlass.Int32(S_CAP), cutlass.Int32(S_CAP), s) + return cutlass.select_(s < cutlass.Int32(1), cutlass.Int32(1), s) + + +def _backoff(ns, delay=False): + """Between two polls (delay: two timer reads of the FC2 start delay): nanosleep, or + nothing when the spin option spins there.""" + if SPIN < (2 if delay else 1): + prims.nanosleep(ns) + + +@cute.jit +def _smem_or(arr, idx, val): + prims.inline_ptx_hl( + "red.shared.or.b32 [{$r0}], {$r1};", read_only_args=[arr.subview(idx).data_ptr(), val] + ) + + +@cute.jit +def _tma_gather4_cta(smem_dst, tma_ptr, k_coord, row0, row1, row2, row3, barrier): + prims.inline_ptx_hl( + "cp.async.bulk.tensor.2d.shared::cta.global.tile::gather4.mbarrier::complete_tx::bytes" + " [{$r0}], [{$r1}, {{$r2}, {$r3}, {$r4}, {$r5}, {$r6}}], [{$r7}];", + read_only_args=[ + smem_dst.data_ptr(), + tma_ptr, + k_coord, + row0, + row1, + row2, + row3, + barrier.data_ptr(), + ], + ) + + +@cute.jit +def _claim(dyn_slot, dyn_ready, dyn_consumed, ctr_ptr, t): + """Warp 4: take the next tile of the global queue and publish the raw index CTA-wide. + Before reusing a ring slot, wait for every reader lane's release.""" + slot = t % cutlass.Int32(DYN_RING) + if t >= cutlass.Int32(DYN_RING): + phase = (t // cutlass.Int32(DYN_RING) - 1) % 2 + while not cute.arch.mbarrier_try_wait(dyn_consumed.subview(slot).data_ptr(), phase): + pass + if prims.elect_sync(): + v = _atomic_fetch_add(ctr_ptr, cutlass.Int32(1)) + cutlass.Int32(CLAIM_BASE) + dyn_slot.store(v, idx=slot) + prims.mbarrier_arrive(dyn_ready.subview(slot)) + cute.arch.sync_warp() + value = dyn_slot.load(idx=slot, is_volatile=True) + cute.arch.sync_warp() + return value + + +@cute.jit +def _recv(dyn_slot, dyn_ready, dyn_consumed, t, phase): + slot = t % cutlass.Int32(DYN_RING) + while not cute.arch.mbarrier_try_wait(dyn_ready.subview(slot).data_ptr(), phase): + pass + value = dyn_slot.load(idx=slot, is_volatile=True) + prims.mbarrier_arrive(dyn_consumed.subview(slot)) + return value + + +@cute.jit +def _next_tile( + t, + claimer: cutlass.Constexpr[bool], + ph, + bidx, + total_tiles, + dyn_slot, + dyn_ready, + dyn_consumed, + ctr_ptr, +): + """Merged-queue index of this CTA's t-th tile (FC1 tiles < tiles_fc1 <= FC2 + tiles), or -1. Dynamic: claimed from the global cursor. Static: bidx + t*grid. + The claimer's own claims go through _claimer_next.""" + if cutlass.const_expr(DYN_ALL): + # Hybrid too: its first tile (the CTA's index) goes through ring slot 0, so every + # slot's publish and release counts match the dynamic queue's. + raw = _recv(dyn_slot, dyn_ready, dyn_consumed, t, ph) + return cutlass.select_(raw < total_tiles, raw, cutlass.Int32(-1)) + lin = bidx + t * NUM_CTAS + return cutlass.select_(lin < total_tiles, lin, cutlass.Int32(-1)) + + +@cute.jit +def _retire_claim(raw, total_tasks, state_ptr, ar_flags, ar_cur, hflags, head_e, ready, num_tokens): + """Claimer warp, after each claim: every CTA claims until the queue is empty, so the + grid's last claim is the one that finds it empty for the NUM_CTAS-th time. It resets + the queue for this layer's next call. With HEAD_FLAGS it also advances the epoch, which + every CTA read before its first claim, and stores the advanced epoch into the ready words + of the tokens past this call's, which no CTA of this call polls; the routing kernel stores + it into the others. Every word then holds the next call's epoch, never the epoch + 1 that + call waits for, also across the int32 wrap (a word left at its initial 0 would match the + epoch -1).""" + succ = cutlass.select_( + total_tasks > cutlass.Int32(CLAIM_BASE), + total_tasks - cutlass.Int32(CLAIM_BASE), + cutlass.Int32(0), + ) + if raw == cutlass.Int32(CLAIM_BASE + NUM_CTAS - 1) + succ: + if prims.elect_sync(): + _store_release(state_ptr + cutlass.Int64(ST_CURSOR * 4), cutlass.Int32(0)) + if cutlass.const_expr(HEAD_FLAGS): + for tok in cutlass.range_constexpr(M_MAX): + if cutlass.Int32(tok) >= num_tokens: + ready.store(head_e + cutlass.Int32(1), idx=tok, is_volatile=True) + ready.store(head_e + cutlass.Int32(1), idx=tok + M_MAX, is_volatile=True) + hflags.store(head_e + cutlass.Int32(1), idx=2, is_volatile=True) + + +@cute.jit +def _claim_ar_task(state_ptr, epi_flag, tidx, ar_flags, ar_cur): + """Epilogue warps: the next cross-rank reduction task (an m-tile), or >= AR_TASKS. + Every CTA claims until none is left, so the grid's last claim comes after every + reduction and push of this call: it resets the cursor and hands the other buffer to + the next call.""" + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + if tidx == 0: + task = _atomic_fetch_add(state_ptr + cutlass.Int64(ST_AR_CURSOR * 4), cutlass.Int32(1)) + if task == cutlass.Int32(AR_TASKS + NUM_CTAS - 1): + _store_release(state_ptr + cutlass.Int64(ST_AR_CURSOR * 4), cutlass.Int32(0)) + ar_flags.store(ar_cur ^ cutlass.Int32(1), idx=0, is_volatile=True) + epi_flag.store(task, idx=0) + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + return epi_flag.load(idx=0) + + +@cute.jit +def _claimer_next(t, first_raw, dyn_slot, dyn_ready, dyn_consumed, ctr_ptr, bidx, total_tasks, + state_ptr, ar_flags, ar_cur, hflags, head_e, ready, num_tokens): # fmt: skip + """The claimer's t-th tile (t >= 1): the claim issued in the prologue, or a new one.""" + if cutlass.const_expr(not DYN_ALL): + lin = bidx + t * NUM_CTAS + return cutlass.select_(lin < total_tasks, lin, cutlass.Int32(-1)) + raw = first_raw + if t != cutlass.Int32(PRECLAIM_T): + raw = _claim(dyn_slot, dyn_ready, dyn_consumed, ctr_ptr, t) + _retire_claim( + raw, total_tasks, state_ptr, ar_flags, ar_cur, hflags, head_e, ready, num_tokens + ) + return cutlass.select_(raw < total_tasks, raw, cutlass.Int32(-1)) + + +def _tanh_f32(x): + e = cute.math.exp2(cute.math.abs(x) * cutlass.Float32(-2.0 * _LOG2E), fastmath=True) + t = (cutlass.Float32(1.0) - e) * cute.arch.rcp_approx(cutlass.Float32(1.0) + e) + return cutlass.select_(x < cutlass.Float32(0.0), -t, t) + + +def _sigmoid_f32(x): + return cute.arch.rcp_approx( + cutlass.Float32(1.0) + cute.math.exp2(x * cutlass.Float32(-_LOG2E), fastmath=True) + ) + + +def _situ(gate, up): + g = ( + cutlass.Float32(SITU_GATE_CAP) + * _tanh_f32(gate * cutlass.Float32(1.0 / SITU_GATE_CAP)) + * _sigmoid_f32(gate) + ) + u = cutlass.Float32(SITU_LINEAR_CAP) * _tanh_f32(up * cutlass.Float32(1.0 / SITU_LINEAR_CAP)) + return g * u + + +def _block_e8m0(amax): + """E8M0 byte of an MX block and 2^(127-byte) as f32 (ceil: trtllm-gen's SiTU cubin + on sm_100; ocp: floor(log2(amax)) - 8, saturated by the caller).""" + bits = cutlass.Int32(amax.bitcast(cutlass.Int32)) + expf = (bits >> cutlass.Int32(23)) & cutlass.Int32(0xFF) + if cutlass.const_expr(SF_RECIPE == "ocp"): + byte = expf - cutlass.Int32(8) + byte = cutlass.select_(byte < cutlass.Int32(0), cutlass.Int32(0), byte) + else: + sf = amax * cutlass.Float32(1.0 / 448.0) + sbits = cutlass.Int32(sf.bitcast(cutlass.Int32)) + sexp = (sbits >> cutlass.Int32(23)) & cutlass.Int32(0xFF) + mant = sbits & cutlass.Int32(0x7FFFFF) + byte = sexp + cutlass.select_(mant != cutlass.Int32(0), cutlass.Int32(1), cutlass.Int32(0)) + byte = cutlass.select_(byte > cutlass.Int32(0xFE), cutlass.Int32(0xFE), byte) + byte = cutlass.select_(amax > cutlass.Float32(0.0), byte, cutlass.Int32(0)) + inv = cutlass.Int32((cutlass.Int32(254) - byte) << cutlass.Int32(23)).bitcast(cutlass.Float32) + return byte, inv + + +def _scan_b_sentinel(sB, stage, lane): + base = stage * cutlass.Int32(NUM_BYTES_B) + lane * cutlass.Int32(B_SCAN_BYTES_PER_LANE) + found = cutlass.Boolean(False) + for i in range(B_SCAN_ITERS): + v = sB.subview(base + cutlass.Int32(i * B_SCAN_VEC)).load(vector_size=B_SCAN_VEC) + for j in range(B_SCAN_VEC): + found = found | (v[j] == cutlass.Int8(FP8_SENTINEL_I8)) + return prims.vote_sync(0xFFFFFFFF, found, "any") + + +def _scan_sfb_sentinel(sSFB, stage, lane): + base = stage * cutlass.Int32(NUM_BYTES_SFB) + lane * cutlass.Int32(SFB_SCAN_VEC) + v = sSFB.subview(base).load(vector_size=SFB_SCAN_VEC) + found = cutlass.Boolean(False) + for j in range(SFB_SCAN_VEC): + found = found | (v[j] == cutlass.Int8(SF_SENTINEL_I8)) + return prims.vote_sync(0xFFFFFFFF, found, "any") + + +@cute.jit +def _rearm_group(c_words, cs_words, group, tid): + """Epilogue warps: put the sentinels back into group's intermediate values and into + bytes 0..3 of each 16-byte scale group (bytes 4..15 are never written and stay 0).""" + c_base = group * (REARM_VEC4 * 4) + for j in cutlass.range_constexpr((REARM_VEC4 + EPI_THREADS - 1) // EPI_THREADS): + q = j * EPI_THREADS + tid + if q < REARM_VEC4: + arm = cutlass.Int32(C_ARM_WORD) + c_words.store((arm, arm, arm, arm), idx=c_base + q * 4, alignment=16) + s_base = group * (N * SF_STRIDE0 // 4) + for j in cutlass.range_constexpr((REARM_SF_WORDS + EPI_THREADS - 1) // EPI_THREADS): + q = j * EPI_THREADS + tid + if q < REARM_SF_WORDS: + cs_words.store(cutlass.Int32(-1), idx=s_base + q * 4) + + +# ============================================================================= +# k3_moe: host function (TMA descriptors + launch) +# ============================================================================= +@cute.jit +def k3_moe( + a1_tensor: cute.Tensor, # w3_w1_weight viewed (H/2, 2I, E) FP4 bytes, K-major + b1_tensor: cute.Tensor, # MXFP8 activations viewed (H, M) FP8 (fold: the scratch, (H, NUM_CTAS*8)) + sfa1_tensor: cute.Tensor, # w3_w1_weight_scale viewed (512, H/128, 2I/128, E) + sfb1_tensor: cute.Tensor, # activation scales (M, H/32) E8M0, linear (fold: (NUM_CTAS*8, H/32)) + c_tensor: cute.Tensor, # intermediate values (G_cap, N, I) FP8, armed + c_scale_tensor: cute.Tensor, # intermediate scales (G_cap, N, I/8), armed + c_words_tensor: cute.Tensor, # int32 view of the intermediate values + cs_words_tensor: cute.Tensor, # int32 view of the intermediate scales + a2_tensor: cute.Tensor, # w2_weight viewed (I/2, H, E) + b2_tensor: cute.Tensor, # c_tensor viewed (I, N, G_cap) + sfa2_tensor: cute.Tensor, # w2_weight_scale viewed (512, I/128, H/128, E) + sfb2_tensor: cute.Tensor, # c_scale viewed (16, I/128, N, G_cap) + y_tensor: cute.Tensor, # out (M, H) bf16: the partial, or with FUSED_AR the reduced sum + y_words_tensor: cute.Tensor, # the same, as int32 words + part_tensor: cute.Tensor, # scratch fp32 (G_cap*N, H) + topk_ids_tensor: cute.Tensor, # int32 (M, 16) global expert ids + topk_w_tensor: cute.Tensor, # bf16 (M, 16) routing weights + state_tensor: cute.Tensor, # int32 (NUM_STATE,) per-layer counters, zero between calls + ar_uc_tensor: cute.Tensor, # int32 words: this rank's all-reduce buffers (FUSED_AR) + ar_mc_tensor: cute.Tensor, # int32 words: their multicast mapping (FUSED_AR) + ar_flags_tensor: cute.Tensor, # int32 [0] = buffer of the next call (FUSED_AR) + logits_tensor: cute.Tensor, # fp32 (M * 896,) router logits (FOLD) + bias_tensor: cute.Tensor, # fp32 (896,) routing bias (FOLD) + xin_words_tensor: cute.Tensor, # int32 view of the bf16 latent (M * H/2,) (FOLD) + xq_words_tensor: cute.Tensor, # int32 view of the e4m3 scratch (NUM_CTAS*8 * H/4,) (FOLD) + xsf_tensor: cute.Tensor, # uint8 scale scratch (NUM_CTAS*8 * H/32,) (FOLD) + ready_tensor: cute.Tensor, # int32 route_quant_ag ready words [16] (HEAD_FLAGS) + hflags_tensor: cute.Tensor, # int32 head workspace flags, [2] = epoch (HEAD_FLAGS) + lat_slab_tensor: cute.Tensor, # int32 words of the consumer's latent slab [3][8][H/2] (LAT_SLAB) + num_tokens: cutlass.Int32, + local_offset: cutlass.Int32, + num_local: cutlass.Int32, + ar_rank: cutlass.Int32, + routed_scaling_factor: cutlass.Float64, + lat_buf: cutlass.Int32, + lat_rearm0: cutlass.Int32, + stream: cuda_driver.CUstream, +) -> None: + _kpp = a1_tensor.shape[0] + _mw = a1_tensor.shape[1] + _ew = a1_tensor.shape[2] + tma_a1_desc = cuda.create_tensor_map_tiled( + global_address=a1_tensor.iterator.toint(), dtype=a_dtype, global_dims=[_kpp * 2, _mw, _ew], + global_strides=[_kpp // 16, (_mw * _kpp) // 16], box_dims=(MMA_TILE_K, MMA_M, 1), + swizzle=cuda.TensorMapSwizzle.s128b, tma_format=TensorMapDataType.f416u4_align16b, + ) # fmt: skip + tma_b1_desc = cuda.create_tensor_map_tiled_from_view( + b1_tensor, + box_dims=(MMA_TILE_K, 1), + stride_order=(0, 1), + swizzle=cuda.TensorMapSwizzle.s128b, + ) + sfa1_fp16 = cute.recast_tensor(sfa1_tensor, cutlass.Uint16) + tma_sfa1_desc = cuda.create_tensor_map_tiled_from_view( + sfa1_fp16, + box_dims=(num_elts_atom_sf_fp16, 1, 1, 1), + stride_order=(0, 1, 2, 3), + swizzle=cuda.TensorMapSwizzle.none, + ) + sfb1_ptr = sfb1_tensor.iterator.toint() + _kpp2 = a2_tensor.shape[0] + _mw2 = a2_tensor.shape[1] + _ew2 = a2_tensor.shape[2] + tma_a2_desc = cuda.create_tensor_map_tiled( + global_address=a2_tensor.iterator.toint(), dtype=a_dtype, global_dims=[_kpp2 * 2, _mw2, _ew2], + global_strides=[_kpp2 // 16, (_mw2 * _kpp2) // 16], box_dims=(MMA_TILE_K, MMA_M, 1), + swizzle=cuda.TensorMapSwizzle.s128b, tma_format=TensorMapDataType.f416u4_align16b, + ) # fmt: skip + tma_b2_desc = cuda.create_tensor_map_tiled_from_view( + b2_tensor, + box_dims=(MMA_TILE_K, MMA_TILER_N, 1), + stride_order=(0, 1, 2), + swizzle=cuda.TensorMapSwizzle.s128b, + ) + sfa2_fp16 = cute.recast_tensor(sfa2_tensor, cutlass.Uint16) + tma_sfa2_desc = cuda.create_tensor_map_tiled_from_view( + sfa2_fp16, + box_dims=(num_elts_atom_sf_fp16, 1, 1, 1), + stride_order=(0, 1, 2, 3), + swizzle=cuda.TensorMapSwizzle.none, + ) + sfb2_fp16 = cute.recast_tensor(sfb2_tensor, cutlass.Uint16) + tma_sfb2_desc = cuda.create_tensor_map_tiled_from_view( + sfb2_fp16, + box_dims=(SFB_GROUP_BYTES // 2, 1, N, 1), + stride_order=(0, 1, 2, 3), + swizzle=cuda.TensorMapSwizzle.none, + ) + state_ptr = state_tensor.iterator.toint() + k3_moe_kernel( + tma_a1_desc, tma_b1_desc, tma_sfa1_desc, tma_a2_desc, tma_b2_desc, tma_sfa2_desc, tma_sfb2_desc, + sfb1_ptr, state_ptr, c_tensor, c_scale_tensor, c_words_tensor, cs_words_tensor, + y_tensor, y_words_tensor, part_tensor, topk_ids_tensor, topk_w_tensor, ar_uc_tensor, ar_mc_tensor, + ar_flags_tensor, logits_tensor, bias_tensor, xin_words_tensor, xq_words_tensor, xsf_tensor, ready_tensor, + hflags_tensor, lat_slab_tensor, num_tokens, local_offset, num_local, ar_rank, routed_scaling_factor, + lat_buf, lat_rearm0, + ).launch( + grid=[NUM_CTAS, 1, 1], block=[threads_per_cta, 1, 1], cluster=(1, 1, 1), stream=stream, + use_pdl=USE_PDL, + ) # fmt: skip + return + + +# ============================================================================= +# k3_moe kernel +# ============================================================================= +@cute.kernel +def k3_moe_kernel( + tma_a1_desc: cutlass.GridConstant[cuda.TensorMap], + tma_b1_desc: cutlass.GridConstant[cuda.TensorMap], + tma_sfa1_desc: cutlass.GridConstant[cuda.TensorMap], + tma_a2_desc: cutlass.GridConstant[cuda.TensorMap], + tma_b2_desc: cutlass.GridConstant[cuda.TensorMap], + tma_sfa2_desc: cutlass.GridConstant[cuda.TensorMap], + tma_sfb2_desc: cutlass.GridConstant[cuda.TensorMap], + sfb1_ptr: cutlass.Int64, + state_ptr: cutlass.Int64, + c_tensor: cutlass.Array, + c_scale_tensor: cutlass.Array, + c_words: cutlass.Array, + cs_words: cutlass.Array, + y_tensor: cutlass.Array, + y_words: cutlass.Array, + part_tensor: cutlass.Array, + topk_ids: cutlass.Array, + topk_w: cutlass.Array, + ar_uc: cutlass.Array, + ar_mc: cutlass.Array, + ar_flags: cutlass.Array, + logits: cutlass.Array, + bias: cutlass.Array, + xin_words: cutlass.Array, + xq_words: cutlass.Array, + xsf: cutlass.Array, + ready: cutlass.Array, + hflags: cutlass.Array, + lat_slab: cutlass.Array, + num_tokens: cutlass.Int32, + local_offset: cutlass.Int32, + num_local: cutlass.Int32, + ar_rank: cutlass.Int32, + routed_scaling_factor: cutlass.Float64, + lat_buf: cutlass.Int32, + lat_rearm0: cutlass.Int32, +) -> None: + mma_tiler_mnk = _mma_tiler_mnk + mma_inst_mnk = _mma_inst_mnk + num_ab_stage = NUM_AB_STAGE + rest_k_sf = REST_K_SF + rest_m_sf = REST_M_SF + rest_n_sf = REST_N_SF + + warp_idx = cute.arch.make_warp_uniform(cute.arch.warp_idx()) + tidx, _, _ = cute.arch.thread_idx() + bidx, _, _ = cute.arch.block_idx() + lane_id = tidx % 32 + dyn_ctr_ptr = state_ptr + cutlass.Int64(ST_CURSOR * 4) + + ab_full_fc1 = cutlass.Array( + cutlass.Int64, num_ab_stage, space=cutlass.AddressSpace.smem, alignment=8 + ) + ab_full_fc2 = cutlass.Array( + cutlass.Int64, num_ab_stage, space=cutlass.AddressSpace.smem, alignment=8 + ) + ab_empty = cutlass.Array( + cutlass.Int64, num_ab_stage, space=cutlass.AddressSpace.smem, alignment=8 + ) + scales_in_tmem = cutlass.Array( + cutlass.Int64, num_ab_stage, space=cutlass.AddressSpace.smem, alignment=8 + ) + lamport_arrived = cutlass.Array( + cutlass.Int64, num_ab_stage, space=cutlass.AddressSpace.smem, alignment=8 + ) + lamport_retry = cutlass.Array( + cutlass.Int64, NUM_LAMPORT_WARPS, space=cutlass.AddressSpace.smem, alignment=8 + ) + dyn_slot = cutlass.Array(cutlass.Int32, DYN_RING, space=cutlass.AddressSpace.smem, alignment=4) + dyn_ready = cutlass.Array(cutlass.Int64, DYN_RING, space=cutlass.AddressSpace.smem, alignment=8) + dyn_consumed = cutlass.Array( + cutlass.Int64, DYN_RING, space=cutlass.AddressSpace.smem, alignment=8 + ) + acc_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + acc_empty = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + tmem_ready = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + # The group -> expert map, the group count and the task count are in shared memory (grouping phase 3): + # the weights producer starts on this while the other warps finish the pair slots (phases 4-5). + groups_ready = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + if cutlass.const_expr(WIDE): + # The epilogue finished an epilogue-only task (combine chunk, re-arm): the claimer may take the next. + epi_free = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + localmax_smem = cutlass.Array( + cutlass.Float32, 2 * len(epilog_warp_id), space=cutlass.AddressSpace.smem + ) + epi_flag = cutlass.Array(cutlass.Int32, 2, space=cutlass.AddressSpace.smem) + tmem_ptr_i32 = cutlass.Array(cutlass.Int32, 1, space=cutlass.AddressSpace.smem) + if cutlass.const_expr(WIDE): + # Grouping (same in every CTA): per group its expert and token count; per (group, slot) + # the pair index token * 16 + k of the slot's token (the routing weight is read from the + # top-k weights). The token masks and each expert's first group live in the B stages + # during the prologue (below). + chunk_cnt = cutlass.Array(cutlass.Int32, 32, space=cutlass.AddressSpace.smem, alignment=16) + g_expert = cutlass.Array( + cutlass.Int32, G_CAP, space=cutlass.AddressSpace.smem, alignment=16 + ) + g_cnt = cutlass.Array(cutlass.Int32, G_CAP, space=cutlass.AddressSpace.smem, alignment=16) + g_pair = cutlass.Array( + cutlass.Int16, G_CAP * N, space=cutlass.AddressSpace.smem, alignment=16 + ) + # Per FC2 slice, the tokens its groups hold (tokens 0..31 at [2 s], 32..63 at [2 s + 1]): the slice tasks + # store, and the combine reads, only those tokens' partial rows. + slice_mask = cutlass.Array( + cutlass.Int32, 2 * S_CAP, space=cutlass.AddressSpace.smem, alignment=16 + ) + else: + # Grouping (same in every CTA): emap[e] = this step's token mask of local expert e, then + # (group << 8) | mask; per group: expert, token count, the 8 slots' token rows (4 bits + # each, padded with the first token) and routing weights (0 on pads); each token's + # slots in top-k order; [0] = number of groups. + emap = cutlass.Array( + cutlass.Int32, NUM_CHUNKS * 32, space=cutlass.AddressSpace.smem, alignment=16 + ) + chunk_cnt = cutlass.Array(cutlass.Int32, 32, space=cutlass.AddressSpace.smem, alignment=16) + g_expert = cutlass.Array( + cutlass.Int32, G_CAP, space=cutlass.AddressSpace.smem, alignment=16 + ) + g_cnt = cutlass.Array(cutlass.Int32, G_CAP, space=cutlass.AddressSpace.smem, alignment=16) + g_rows = cutlass.Array(cutlass.Int32, G_CAP, space=cutlass.AddressSpace.smem, alignment=16) + g_rw = cutlass.Array( + cutlass.Float32, G_CAP * N, space=cutlass.AddressSpace.smem, alignment=16 + ) + s_tok_slots = cutlass.Array( + cutlass.Int32, ROUTE_PAIRS, space=cutlass.AddressSpace.smem, alignment=16 + ) + # The local pairs compacted in (token, top-k) order: slot | token << 16, for the combine. + s_pairs = cutlass.Array( + cutlass.Int32, ROUTE_PAIRS, space=cutlass.AddressSpace.smem, alignment=16 + ) + s_pcnt = cutlass.Array(cutlass.Int32, 4, space=cutlass.AddressSpace.smem, alignment=16) + # [0] groups, [1] all-reduce buffer of this call, [2] local pairs + s_meta = cutlass.Array(cutlass.Int32, 4, space=cutlass.AddressSpace.smem, alignment=16) + sA = cutlass.Array( + cutlass.Int8, NUM_BYTES_A * num_ab_stage, space=cutlass.AddressSpace.smem, alignment=1024 + ) + sB = cutlass.Array( + cutlass.Int8, NUM_BYTES_B * num_ab_stage, space=cutlass.AddressSpace.smem, alignment=1024 + ) + sSFA = cutlass.Array( + cutlass.Int8, NUM_BYTES_SFA * num_ab_stage, space=cutlass.AddressSpace.smem, alignment=1024 + ) + sSFB = cutlass.Array( + cutlass.Int8, NUM_BYTES_SFB * num_ab_stage, space=cutlass.AddressSpace.smem, alignment=1024 + ) + if cutlass.const_expr(WIDE): + # Prologue only, in the B stages (no TMA writes them before the grouping's last barrier): + # each local expert's token mask, tokens 0..31 at [e] and 32..63 at [MASK_HI + e], and + # its first group. + emask = cutlass.Array( + sB.data_ptr(0), shape=(2 * MASK_HI,), dtype=cutlass.Int32, alignment=16 + ) + e_first = cutlass.Array( + sB.data_ptr(2 * MASK_HI * 4), shape=(MASK_HI,), dtype=cutlass.Int32, alignment=16 + ) + # This thread's index among the grouping threads (every warp but the claimer). + gtid = tidx - cutlass.select_( + warp_idx > producer_weights_warp_id, cutlass.Int32(32), cutlass.Int32(0) + ) + bias_vals = [] + if cutlass.const_expr(FOLD): + rs_key = cutlass.Array( + sA.data_ptr(RS_KEY), shape=(M_MAX * NUM_EXPERTS,), dtype=cutlass.Int32, alignment=16 + ) + rs_sig = cutlass.Array( + sA.data_ptr(RS_SIG), shape=(M_MAX * NUM_EXPERTS,), dtype=cutlass.Float32, alignment=16 + ) + rs_id = cutlass.Array( + sA.data_ptr(RS_ID), shape=(M_MAX * TOP_K,), dtype=cutlass.Int32, alignment=16 + ) + rs_w = cutlass.Array( + sA.data_ptr(RS_W), shape=(M_MAX * TOP_K,), dtype=cutlass.Int32, alignment=16 + ) + # The routing bias is a weight: this thread's experts' values are read before the grid + # dependency (the claimer warp's reads are in bounds and unused). + for i in cutlass.range_constexpr(KEY_ITERS): + e = gtid + cutlass.Int32(i * GROUPING_THREADS) + bias_vals.append( + bias.load(idx=cutlass.select_(e < NUM_EXPERTS, e, cutlass.Int32(NUM_EXPERTS - 1))) + ) + + # ------------------------------------------------ prologue, before the grid dependency + if warp_idx == 0: + if tidx < num_ab_stage: + prims.mbarrier_init(ab_full_fc1.subview(tidx), 2 + N) + prims.mbarrier_init(ab_full_fc2.subview(tidx), 2) + prims.mbarrier_init(ab_empty.subview(tidx), 1) + prims.mbarrier_init(scales_in_tmem.subview(tidx), 1) + prims.mbarrier_init(lamport_arrived.subview(tidx), 1) + if tidx < NUM_LAMPORT_WARPS: + prims.mbarrier_init(lamport_retry.subview(tidx), 1) + if cutlass.const_expr(DYN_ALL): + if tidx < DYN_RING: + prims.mbarrier_init(dyn_ready.subview(tidx), 1) + prims.mbarrier_init(dyn_consumed.subview(tidx), threads_per_cta - 32) + if tidx == 0: + prims.mbarrier_init(acc_full.subview(0), 1) + prims.mbarrier_init(acc_empty.subview(0), EPI_THREADS) + prims.mbarrier_init(tmem_ready.subview(0), 32) + prims.mbarrier_init(groups_ready.subview(0), 1) + if cutlass.const_expr(WIDE): + prims.mbarrier_init(epi_free.subview(0), 1) + if cutlass.const_expr(WIDE): + for j in cutlass.range_constexpr((2 * MASK_HI + threads_per_cta - 1) // threads_per_cta): + if tidx + j * threads_per_cta < 2 * MASK_HI: + emask.store(cutlass.Int32(0), idx=tidx + j * threads_per_cta) + if tidx < 2 * S_CAP: + slice_mask.store(cutlass.Int32(0), idx=tidx) + else: + for j in cutlass.range_constexpr( + (NUM_CHUNKS * 32 + threads_per_cta - 1) // threads_per_cta + ): + if tidx + j * threads_per_cta < NUM_CHUNKS * 32: + emap.store(cutlass.Int32(0), idx=tidx + j * threads_per_cta) + if cutlass.const_expr(HEAD_FLAGS): + # This call's epoch: k3_moe of the previous call advanced it and completed before + # route_quant_ag (which never writes it) let this grid launch. + if tidx == 0: + s_meta.store(hflags.load(idx=2, is_volatile=True), idx=3) + prims.fence_mbarrier_init() + prims.barrier_cta_sync(0) + head_e = cutlass.Int32(0) + if cutlass.const_expr(HEAD_FLAGS): + head_e = s_meta.load(idx=3) + if warp_idx == consumer_warp_id: + prims.tcgen05_alloc(tmem_ptr_i32, num_tmem_alloc_cols) + prims.mbarrier_arrive(tmem_ready) + prims.tcgen05_relinquish_alloc_permit() + + first_raw = cutlass.Int32(0) + if cutlass.const_expr(HYBRID): + if warp_idx == producer_weights_warp_id: + if prims.elect_sync(): + dyn_slot.store(bidx, idx=0) + prims.mbarrier_arrive(dyn_ready.subview(0)) + cute.arch.sync_warp() + if cutlass.const_expr(PREWAIT_CLAIM): + if warp_idx == producer_weights_warp_id: + first_raw = _claim( + dyn_slot, dyn_ready, dyn_consumed, dyn_ctr_ptr, cutlass.Int32(PRECLAIM_T) + ) + if cutlass.const_expr(USE_PDL): + if cutlass.const_expr(not HEAD_FLAGS): + cute.arch.griddepcontrol_wait() + if cutlass.const_expr(not LATE_TRIGGER and not HEAD_TRIGGER_READY): + # Every consumer of this kernel's outputs waits for the whole grid, so the next + # kernel may launch (and set itself up) now. + cute.arch.griddepcontrol_launch_dependents() + # ------------------------------------------------ prologue: first claim + grouping + if warp_idx == producer_weights_warp_id: + if cutlass.const_expr(DYN_ALL and not PREWAIT_CLAIM): + first_raw = _claim( + dyn_slot, dyn_ready, dyn_consumed, dyn_ctr_ptr, cutlass.Int32(PRECLAIM_T) + ) + else: + gw = warp_idx - cutlass.select_( + warp_idx > producer_weights_warp_id, cutlass.Int32(1), cutlass.Int32(0) + ) + if cutlass.const_expr(FOLD): + # (0) every token's selection keys and sigmoids (all logit loads in flight first, rows + # clamped to M-1), then warp gw < M selects token gw's top-16 and their weights. + logit = [] + for t in cutlass.range_constexpr(M_MAX): + row = cutlass.select_( + cutlass.Int32(t) < num_tokens, cutlass.Int32(t), num_tokens - cutlass.Int32(1) + ) + for i in cutlass.range_constexpr(KEY_ITERS): + e = gtid + cutlass.Int32(i * GROUPING_THREADS) + ec = cutlass.select_(e < NUM_EXPERTS, e, cutlass.Int32(NUM_EXPERTS - 1)) + logit.append(logits.load(idx=row * NUM_EXPERTS + ec)) + for t in cutlass.range_constexpr(M_MAX): + if cutlass.Int32(t) < num_tokens: + for i in cutlass.range_constexpr(KEY_ITERS): + e = gtid + cutlass.Int32(i * GROUPING_THREADS) + if e < NUM_EXPERTS: + sig = _rq.sigmoid_accurate(logit[t * KEY_ITERS + i]) + rs_sig.store(sig, idx=t * NUM_EXPERTS + e) + rs_key.store( + _rq.selection_key(sig + bias_vals[i]), idx=t * NUM_EXPERTS + e + ) + cute.arch.barrier(barrier_id=GROUPING_BAR_ID, number_of_threads=GROUPING_THREADS) + if gw < num_tokens: + expert, weight_bits = _rq.top16_warp( + rs_key.subview(gw * NUM_EXPERTS), + rs_sig.subview(gw * NUM_EXPERTS), + lane_id, + routed_scaling_factor, + ) + if lane_id < TOP_K: + rs_id.store(expert, idx=gw * TOP_K + lane_id) + rs_w.store( + cutlass.Int32(weight_bits) & cutlass.Int32(0xFFFF), idx=gw * TOP_K + lane_id + ) + cute.arch.barrier(barrier_id=GROUPING_BAR_ID, number_of_threads=GROUPING_THREADS) + if cutlass.const_expr(HEAD_FLAGS): + # One lane per token acquires route_quant_ag's ready word (every CTA polling with all its + # pair threads would hammer 8 L2 words with ~20K threads); the barrier hands the + # acquired writes to the other grouping threads. + if warp_idx == 0: + if lane_id < num_tokens: + ready_addr = ready.data_ptr(lane_id).toint() + while _load_acquire(ready_addr) != head_e + cutlass.Int32(1): + pass + cute.arch.barrier(barrier_id=GROUPING_BAR_ID, number_of_threads=GROUPING_THREADS) + if cutlass.const_expr(USE_PDL and HEAD_TRIGGER_READY and not LATE_TRIGGER): + # The CTA's first launch_dependents by any thread is its trigger. + cute.arch.griddepcontrol_launch_dependents() + if cutlass.const_expr(FUSED_AR): + # This call's all-reduce buffer: written by the previous call's last task claim, or with + # AR_PUSH_ONLY the consumer's call count, whose low bit alternates the halves. + if tidx == 0: + s_meta.store(ar_flags.load(idx=0, is_volatile=True) & cutlass.Int32(1), idx=1) + if cutlass.const_expr(WIDE): + # (1) each local expert's token mask; this thread's pairs p = gtid + i * GROUPING_THREADS + # (token p // 16, top-k slot p % 16). + pair_experts = [] + for i in cutlass.range_constexpr(PAIR_ITERS): + p = gtid + cutlass.Int32(i * GROUPING_THREADS) + pt = p // TOP_K + pe = cutlass.Int32(-1) + if pt < num_tokens: + el = topk_ids.load(idx=p) - local_offset + if el >= cutlass.Int32(0): + if el < num_local: + pe = el + if pe >= cutlass.Int32(0): + _smem_or( + emask, + pe + (pt >> cutlass.Int32(5)) * cutlass.Int32(MASK_HI), + cutlass.Int32(1) << (pt & cutlass.Int32(31)), + ) + pair_experts.append(pe) + cute.arch.barrier(barrier_id=GROUPING_BAR_ID, number_of_threads=GROUPING_THREADS) + # (2) groups per 32-expert chunk: ceil(tokens / 8) per expert (at most 8: four ballots). + for j in cutlass.range_constexpr(CHUNKS_PER_WARP): + c = gw + j * GROUPING_WARPS + if c < NUM_CHUNKS: + x = c * 32 + lane_id + cnt = cute.arch.popc(emask.load(idx=x)) + cute.arch.popc( + emask.load(idx=x + MASK_HI) + ) + ng = (cnt + cutlass.Int32(N - 1)) // cutlass.Int32(N) + chunk_groups = cutlass.Int32(0) + for b in cutlass.range_constexpr(4): + bal = cute.arch.vote_ballot_sync( + ((ng >> cutlass.Int32(b)) & cutlass.Int32(1)) != cutlass.Int32(0) + ) + chunk_groups = chunk_groups + (cute.arch.popc(bal) << cutlass.Int32(b)) + if lane_id == 0: + chunk_cnt.store(chunk_groups, idx=c) + cute.arch.barrier(barrier_id=GROUPING_BAR_ID, number_of_threads=GROUPING_THREADS) + # (3) an expert's first group = the groups of the experts before it (ascending id); its + # groups take its tokens 8 at a time in ascending order. + for j in cutlass.range_constexpr(CHUNKS_PER_WARP): + c = gw + j * GROUPING_WARPS + if c < NUM_CHUNKS: + x = c * 32 + lane_id + cnt = cute.arch.popc(emask.load(idx=x)) + cute.arch.popc( + emask.load(idx=x + MASK_HI) + ) + ng = (cnt + cutlass.Int32(N - 1)) // cutlass.Int32(N) + g0 = cutlass.Int32(0) + for c2 in cutlass.range_constexpr(NUM_CHUNKS): + if cutlass.Int32(c2) < c: + g0 = g0 + chunk_cnt.load(idx=c2) + for b in cutlass.range_constexpr(4): + bal = cute.arch.vote_ballot_sync( + ((ng >> cutlass.Int32(b)) & cutlass.Int32(1)) != cutlass.Int32(0) + ) + g0 = g0 + ( + cute.arch.popc(bal & cutlass.Int32(cute.arch.lanemask_lt())) + << cutlass.Int32(b) + ) + e_first.store(g0, idx=x) + for q in cutlass.range(ng, unroll=1): + rest = cnt - q * cutlass.Int32(N) + g_expert.store(x, idx=g0 + q) + g_cnt.store( + cutlass.select_(rest > cutlass.Int32(N), cutlass.Int32(N), rest), + idx=g0 + q, + ) + if tidx == 0: + total = cutlass.Int32(0) + for c2 in cutlass.range_constexpr(NUM_CHUNKS): + total = total + chunk_cnt.load(idx=c2) + s_meta.store(total, idx=0) + s_meta.store(_wide_slices(total), idx=2) + cute.arch.barrier(barrier_id=GROUPING_BAR_ID, number_of_threads=GROUPING_THREADS) + if cutlass.const_expr(EARLY_RELEASE): + # Phases 1-3 and thread 0's totals are behind the barrier above; the weights + # producer needs nothing from phase 4. + if tidx == 0: + prims.mbarrier_arrive(groups_ready) + # (4) slot of each pair: its expert's first group * N + the rank of its token among + # the expert's tokens; the slot keeps the pair index. + for i in cutlass.range_constexpr(PAIR_ITERS): + pe = pair_experts[i] + if pe >= cutlass.Int32(0): + p = gtid + cutlass.Int32(i * GROUPING_THREADS) + pt = p // TOP_K + below = (cutlass.Int32(1) << (pt & cutlass.Int32(31))) - cutlass.Int32(1) + below_lo = cutlass.select_(pt >= cutlass.Int32(32), cutlass.Int32(-1), below) + below_hi = cutlass.select_(pt > cutlass.Int32(32), below, cutlass.Int32(0)) + rank = cute.arch.popc(emask.load(idx=pe) & below_lo) + cute.arch.popc( + emask.load(idx=pe + MASK_HI) & below_hi + ) + first = e_first.load(idx=pe) + g_pair.store(cutlass.Int16(p), idx=first * N + rank) + # The slice of the pair's group: g in [s G / S, (s + 1) G / S). + grp = first + rank // N + n_sl = s_meta.load(idx=2) + sl_p = ((grp + cutlass.Int32(1)) * n_sl - cutlass.Int32(1)) // s_meta.load( + idx=0 + ) + _smem_or( + slice_mask, + sl_p * 2 + (pt >> cutlass.Int32(5)), + cutlass.Int32(1) << (pt & cutlass.Int32(31)), + ) + # The masks and first groups were written and read through the generic proxy in the B + # stages; the activation producer's TMA writes (async proxy) follow the barrier below. + prims.fence_proxy("async_shared", space=prims.SharedSpace.shared_cta) + else: + # (1) token bitmask per local expert; this thread's pair (token pt, top-k slot). + pt = tidx // TOP_K + pe = cutlass.Int32(-1) + pw = cutlass.Float32(0.0) + if tidx < ROUTE_PAIRS: + if pt < num_tokens: + if cutlass.const_expr(FOLD): + el = rs_id.load(idx=tidx) - local_offset + else: + el = topk_ids.load(idx=tidx) - local_offset + if el >= cutlass.Int32(0): + if el < num_local: + pe = el + if cutlass.const_expr(FOLD): + pw = (rs_w.load(idx=tidx) << cutlass.Int32(16)).bitcast( + cutlass.Float32 + ) + else: + pw = cutlass.Float32(topk_w.load(idx=tidx)) + if pe >= cutlass.Int32(0): + _smem_or(emap, pe, cutlass.Int32(1) << pt) + cute.arch.barrier(barrier_id=GROUPING_BAR_ID, number_of_threads=GROUPING_THREADS) + # (2) experts present per 32-expert chunk. + for j in cutlass.range_constexpr(CHUNKS_PER_WARP): + c = gw + j * GROUPING_WARPS + if c < NUM_CHUNKS: + m = emap.load(idx=c * 32 + lane_id) + bal = cute.arch.vote_ballot_sync(m != cutlass.Int32(0)) + if lane_id == 0: + chunk_cnt.store(cute.arch.popc(bal), idx=c) + cute.arch.barrier(barrier_id=GROUPING_BAR_ID, number_of_threads=GROUPING_THREADS) + # (3) group index = rank of the expert among those present (ascending id). + for j in cutlass.range_constexpr(CHUNKS_PER_WARP): + c = gw + j * GROUPING_WARPS + if c < NUM_CHUNKS: + x = c * 32 + lane_id + m = emap.load(idx=x) + bal = cute.arch.vote_ballot_sync(m != cutlass.Int32(0)) + base = cutlass.Int32(0) + for c2 in cutlass.range_constexpr(NUM_CHUNKS): + if cutlass.Int32(c2) < c: + base = base + chunk_cnt.load(idx=c2) + g = base + cute.arch.popc(bal & cutlass.Int32(cute.arch.lanemask_lt())) + if m != cutlass.Int32(0): + emap.store((g << cutlass.Int32(8)) | m, idx=x) + g_expert.store(x, idx=g) + cnt = cutlass.Int32(0) + rows = cutlass.Int32(0) + first = cutlass.Int32(0) + for tb in cutlass.range_constexpr(M_MAX): + if ((m >> cutlass.Int32(tb)) & cutlass.Int32(1)) != cutlass.Int32(0): + rows = rows | (cutlass.Int32(tb) << (cnt * cutlass.Int32(4))) + if cnt == cutlass.Int32(0): + first = cutlass.Int32(tb) + cnt = cnt + cutlass.Int32(1) + for n2 in cutlass.range_constexpr(N): + if cutlass.Int32(n2) >= cnt: + rows = rows | (first << cutlass.Int32(4 * n2)) + g_rw.store(cutlass.Float32(0.0), idx=g * N + n2) + g_cnt.store(cnt, idx=g) + g_rows.store(rows, idx=g) + if tidx == 0: + total = cutlass.Int32(0) + for c2 in cutlass.range_constexpr(NUM_CHUNKS): + total = total + chunk_cnt.load(idx=c2) + s_meta.store(total, idx=0) + cute.arch.barrier(barrier_id=GROUPING_BAR_ID, number_of_threads=GROUPING_THREADS) + if cutlass.const_expr(EARLY_RELEASE): + # Phases 1-3 and thread 0's totals are behind the barrier above; the arrive releases them to the + # weights producer, which needs nothing from phases 4-5. + if tidx == 0: + prims.mbarrier_arrive(groups_ready) + # (4) slot of each pair: group * N + rank of the token among the expert's tokens. + slot = cutlass.Int32(-1) + pbal = cutlass.Int32(0) + if tidx < ROUTE_PAIRS: + if pe >= cutlass.Int32(0): + word = emap.load(idx=pe) + mask_lt = word & ((cutlass.Int32(1) << pt) - cutlass.Int32(1)) + slot = (word >> cutlass.Int32(8)) * N + cute.arch.popc( + mask_lt & cutlass.Int32(0xFF) + ) + g_rw.store(pw, idx=slot) + s_tok_slots.store(slot, idx=tidx) + if cutlass.const_expr(COMBINE == "batched"): + pbal = cute.arch.vote_ballot_sync(slot >= cutlass.Int32(0)) + if lane_id == 0: + s_pcnt.store(cute.arch.popc(pbal), idx=warp_idx) + if cutlass.const_expr(COMBINE == "batched"): + cute.arch.barrier(barrier_id=GROUPING_BAR_ID, number_of_threads=GROUPING_THREADS) + # (5) the local pairs, compacted in (token, top-k) order. + if cutlass.const_expr(COMBINE == "batched"): + if tidx < ROUTE_PAIRS: + pbase = cutlass.Int32(0) + ptotal = cutlass.Int32(0) + for w2 in cutlass.range_constexpr(ROUTE_PAIRS // 32): + c2 = s_pcnt.load(idx=w2) + pbase = pbase + cutlass.select_( + cutlass.Int32(w2) < warp_idx, c2, cutlass.Int32(0) + ) + ptotal = ptotal + c2 + if slot >= cutlass.Int32(0): + s_pairs.store( + slot | (pt << cutlass.Int32(16)), + idx=pbase + + cute.arch.popc(pbal & cutlass.Int32(cute.arch.lanemask_lt())), + ) + if tidx == 0: + s_meta.store(ptotal, idx=2) + if cutlass.const_expr(not EARLY_RELEASE): + if cutlass.const_expr(FOLD): + # The routing scratch in the A stages was written and read through the generic proxy; + # the ring's TMA writes (async proxy) start after the CTA-wide sync below. + prims.fence_proxy("async_shared", space=prims.SharedSpace.shared_cta) + prims.barrier_cta_sync(0) + else: + if warp_idx == producer_weights_warp_id: + while not cute.arch.mbarrier_test_wait(groups_ready.data_ptr(), 0): + pass + else: + cute.arch.barrier(barrier_id=GROUPING_BAR_ID, number_of_threads=GROUPING_THREADS) + num_groups = cute.arch.make_warp_uniform(s_meta.load(idx=0)) + ar_cur = cutlass.Int32(0) + if cutlass.const_expr(FUSED_AR): + ar_cur = cute.arch.make_warp_uniform(s_meta.load(idx=1)) + tiles_fc1 = M_TILES_FC1 * num_groups + total_tiles = M_TILES_TOTAL * num_groups + num_slices = cutlass.Int32(1) + if cutlass.const_expr(SLICE): + # FC2 tasks are (slice, m-tile), slice-major, so the first FC2 tasks read the first + # groups, whose FC1 tiles come first in the queue; at most 32 groups per slice (one + # lane each when the tasks count their groups' readers; the wide build loops). + if cutlass.const_expr(WIDE): + num_slices = _wide_slices(num_groups) + else: + num_slices = cutlass.select_( + num_groups < cutlass.Int32(FC2_SLICES), num_groups, cutlass.Int32(FC2_SLICES) + ) + min_slices = (num_groups + cutlass.Int32(31)) // cutlass.Int32(32) + num_slices = cutlass.select_(num_slices < min_slices, min_slices, num_slices) + num_slices = cutlass.select_( + num_slices < cutlass.Int32(1), cutlass.Int32(1), num_slices + ) + total_tiles = cutlass.select_( + num_groups > cutlass.Int32(0), + tiles_fc1 + cutlass.Int32(M_TILES_FC2) * num_slices, + cutlass.Int32(0), + ) + # Queue order: FC1 tiles, then FC2 tiles (then the slice FC2's combine and re-arm tasks). The + # all-reduce tasks have their own cursor. + total_tasks = total_tiles + n_chunks = 0 # combine tasks per m-tile (chunks of 8 tokens) + if cutlass.const_expr(SLICE_TASKS): + if cutlass.const_expr(COMBINE_CHUNKS): + n_chunks = (num_tokens + cutlass.Int32(N - 1)) // cutlass.Int32(N) + total_tasks = total_tiles + cutlass.select_( + num_groups > cutlass.Int32(0), + cutlass.Int32(M_TILES_FC2) * n_chunks + num_slices, + cutlass.Int32(0), + ) + else: + total_tasks = total_tiles + cutlass.select_( + num_groups > cutlass.Int32(0), + cutlass.Int32(M_TILES_FC2 if COMBINE_TASKS else 0) + num_slices, + cutlass.Int32(0), + ) + active = total_tasks > 0 + if cutlass.const_expr(DYN_ALL): + if warp_idx == producer_weights_warp_id: + _retire_claim( + first_raw, total_tasks, state_ptr, ar_flags, ar_cur, hflags, head_e, ready, + num_tokens, + ) # fmt: skip + coord_n = 0 + + # Fold: the MXFP8 latent, quantized by this CTA into its own scratch rows [8 * bidx, 8 * bidx + M) + # by the epilogue and Lamport warps (idle until their first tile), all loads in flight first + # (vectors past M clamped to this thread's first one). The activation producer gathers FC1's B + # rows and scales from there once QUANT_BAR completes. + if cutlass.const_expr(FOLD): + if active and ((warp_idx < len(epilog_warp_id)) | (warp_idx >= lamport_acts_warp_id)): + qtid = tidx - cutlass.select_( + warp_idx >= lamport_acts_warp_id, + cutlass.Int32((lamport_acts_warp_id - len(epilog_warp_id)) * 32), + cutlass.Int32(0), + ) + n_vec = num_tokens * VEC8_PER_ROW + xrow0 = bidx * M_MAX + vecs = [] + for k in cutlass.range_constexpr(QUANT_ITERS): + c = qtid + cutlass.Int32(k * QUANT_THREADS) + cc = cutlass.select_(c < n_vec, c, qtid) + tok = cc // VEC8_PER_ROW + vecs.append( + xin_words.load( + idx=tok * (H // 2) + (cc - tok * VEC8_PER_ROW) * 4, + vector_size=4, + alignment=16, + ) + ) + for k in cutlass.range_constexpr(QUANT_ITERS): + c = qtid + cutlass.Int32(k * QUANT_THREADS) + # n_vec and k * QUANT_THREADS are multiples of 64: the condition is warp-uniform, as + # the scale's 4-lane shuffles need. + if c < n_vec: + tok = c // VEC8_PER_ROW + j = c - tok * VEC8_PER_ROW + v = vecs[k] + q_lo, q_hi, sf_byte = _rq.mxfp8_quant_vec8([v[0], v[1], v[2], v[3]]) + xq_words.store((q_lo, q_hi), idx=(xrow0 + tok) * (H // 4) + j * 2, alignment=8) + if lane_id % 4 == 0: + xsf.store( + cutlass.Uint8(sf_byte), idx=(xrow0 + tok) * (H // MX_BLOCK) + j // 4 + ) + # Generic-proxy stores read by the producer's TMA gathers (async proxy). + prims.fence_proxy("async_global") + prims.barrier_cta_arrive(QUANT_BAR_ID, QUANT_THREADS + 32) + # Nothing routed to this rank: its partial is zero. + if warp_idx < len(epilog_warp_id) and num_groups == 0 and bidx < M_TILES_FC2: + if cutlass.const_expr(HEAD_FLAGS and USE_PDL): + cute.arch.griddepcontrol_wait() + if cutlass.const_expr(FUSED_AR): + ztok = tidx // 16 + if ztok < num_tokens: + zero = cutlass.Int32(0) + ar_mc.store( + (zero, zero, zero, zero), + idx=_ar_word(ar_cur, ztok, ar_rank, bidx * MMA_M + (tidx % 16) * 8), + alignment=16, + ) + else: + zch = bidx * MMA_M + tidx + for t in cutlass.range(num_tokens, unroll=1): + y_tensor.store(cutlass.Float32(0.0).to(c_dtype), idx=t * H + zch) + + # ------------------------------------------------ producerWeights (4), claimer + if warp_idx == producer_weights_warp_id and active: + g = 0 + ab_empty_phase = 1 + t = cutlass.Int32(0) + dyn_ph = 0 + v = cutlass.select_(bidx < total_tasks, bidx, cutlass.Int32(-1)) + if cutlass.const_expr(SCHED == "dynamic"): + v = cutlass.select_(first_raw < total_tasks, first_raw, cutlass.Int32(-1)) + in_fc1 = (v >= cutlass.Int32(0)) & (v < tiles_fc1) + while in_fc1: + m_tile = v % M_TILES_FC1 + group = v // M_TILES_FC1 + coord_m = m_tile * mma_tiler_mnk[0] + coord_m_sf = coord_m // (num_m0_per_sf_atom * num_m1_per_sf_atom) + coord_expert = g_expert.load(idx=group) + for k_tile in cutlass.range(K1_TILES, unroll=1): + stage = g % num_ab_stage + if stage == 0 and g != 0: + ab_empty_phase = ab_empty_phase ^ 1 + while not cute.arch.mbarrier_try_wait( + ab_empty.subview(stage).data_ptr(), ab_empty_phase + ): + pass + coord_k = k_tile * mma_tiler_mnk[2] + if prims.elect_sync(): + prims.mbarrier_arrive_expect_tx( + ab_full_fc1.subview(stage), NUM_TMA_LOAD_BYTES_WEIGHTS + ) + prims.cp_async_bulk_tensor_shared_cta_global( + sA.subview(NUM_BYTES_A * stage), + tma_a1_desc.get_ptr(), + (coord_k, coord_m, coord_expert), + ab_full_fc1.subview(stage), + l2_cache_hint=W_L2_HINT, + ) + prims.cp_async_bulk_tensor_shared_cta_global( + sSFA.subview(NUM_BYTES_SFA * stage), + tma_sfa1_desc.get_ptr(), + (cutlass.Int32(0), k_tile * rest_k_sf, coord_m_sf, coord_expert), + ab_full_fc1.subview(stage), + l2_cache_hint=W_L2_HINT, + ) + g = g + 1 + t = t + cutlass.Int32(1) + if cutlass.const_expr(DYN_ALL): + if t % cutlass.Int32(DYN_RING) == cutlass.Int32(0): + dyn_ph = dyn_ph ^ 1 + v = _claimer_next( + t, first_raw, dyn_slot, dyn_ready, dyn_consumed, dyn_ctr_ptr, bidx, total_tasks, + state_ptr, ar_flags, ar_cur, hflags, head_e, ready, num_tokens, + ) # fmt: skip + in_fc1 = (v >= cutlass.Int32(0)) & (v < tiles_fc1) + in_fc2 = (v >= cutlass.Int32(0)) & (v < total_tiles) + while in_fc2: + lin2 = v - tiles_fc1 + m_tile = lin2 % M_TILES_FC2 + g_lo = lin2 // M_TILES_FC2 + g_hi = g_lo + cutlass.Int32(1) + if cutlass.const_expr(SLICE): + sl = lin2 // M_TILES_FC2 + if cutlass.const_expr(WIDE): + sl = lin2 % num_slices + m_tile = lin2 // num_slices + g_lo = sl * num_groups // num_slices + g_hi = (sl + cutlass.Int32(1)) * num_groups // num_slices + coord_m = m_tile * mma_tiler_mnk[0] + coord_m_sf = coord_m // (num_m0_per_sf_atom * num_m1_per_sf_atom) + for group in cutlass.range(g_lo, g_hi, unroll=1): + coord_expert = g_expert.load(idx=group) + for k_tile in cutlass.range(K2_TILES, unroll=1): + stage = g % num_ab_stage + if stage == 0 and g != 0: + ab_empty_phase = ab_empty_phase ^ 1 + while not cute.arch.mbarrier_try_wait( + ab_empty.subview(stage).data_ptr(), ab_empty_phase + ): + pass + coord_k = k_tile * mma_tiler_mnk[2] + if prims.elect_sync(): + prims.mbarrier_arrive_expect_tx( + ab_full_fc2.subview(stage), NUM_TMA_LOAD_BYTES_WEIGHTS + ) + prims.cp_async_bulk_tensor_shared_cta_global( + sA.subview(NUM_BYTES_A * stage), + tma_a2_desc.get_ptr(), + (coord_k, coord_m, coord_expert), + ab_full_fc2.subview(stage), + l2_cache_hint=W_L2_HINT, + ) + prims.cp_async_bulk_tensor_shared_cta_global( + sSFA.subview(NUM_BYTES_SFA * stage), + tma_sfa2_desc.get_ptr(), + (cutlass.Int32(0), k_tile * rest_k_sf, coord_m_sf, coord_expert), + ab_full_fc2.subview(stage), + l2_cache_hint=W_L2_HINT, + ) + g = g + 1 + t = t + cutlass.Int32(1) + if cutlass.const_expr(DYN_ALL): + if t % cutlass.Int32(DYN_RING) == cutlass.Int32(0): + dyn_ph = dyn_ph ^ 1 + if cutlass.const_expr(FC2_FIRST_HOLD): + if t == cutlass.Int32(1): # this FC2 task was the CTA's first + while not cute.arch.mbarrier_try_wait( + ab_empty.subview((g - 1) % num_ab_stage).data_ptr(), ab_empty_phase ^ 1 + ): + pass + v = _claimer_next( + t, first_raw, dyn_slot, dyn_ready, dyn_consumed, dyn_ctr_ptr, bidx, total_tasks, + state_ptr, ar_flags, ar_cur, hflags, head_e, ready, num_tokens, + ) # fmt: skip + in_fc2 = (v >= cutlass.Int32(0)) & (v < total_tiles) + t_epi0 = t # the claim of this CTA's first epilogue-only task + while v >= cutlass.Int32(0): # the all-reduce tasks are the epilogue's + if cutlass.const_expr(WIDE): + # The combine and re-arm tasks are the epilogue's alone: claim the next one only once + # the epilogue has finished this one, so that they spread over the CTAs that are free. + while not cute.arch.mbarrier_try_wait( + epi_free.data_ptr(), (t - t_epi0) % cutlass.Int32(2) + ): + pass + t = t + cutlass.Int32(1) + if cutlass.const_expr(DYN_ALL): + if t % cutlass.Int32(DYN_RING) == cutlass.Int32(0): + dyn_ph = dyn_ph ^ 1 + v = _claimer_next( + t, first_raw, dyn_slot, dyn_ready, dyn_consumed, dyn_ctr_ptr, bidx, total_tasks, + state_ptr, ar_flags, ar_cur, hflags, head_e, ready, num_tokens, + ) # fmt: skip + + # ------------------------------------------------ producerActivations (5) + if warp_idx == producer_acts_warp_id and active: + lane = tidx % 32 + sfb_lane = cutlass.select_(lane < cutlass.Int32(N), lane, cutlass.Int32(0)) + # Fold: B rows and scales come from this CTA's scratch rows, complete once every quantizing + # thread has arrived. + xrow0 = cutlass.Int32(0) + if cutlass.const_expr(HEAD_FLAGS): + # route_quant_ag's MXFP8 rows (its writers fenced the async proxy before releasing), + # acquired by one lane per token. + if lane < num_tokens: + qready_addr = ready.data_ptr(lane + cutlass.Int32(M_MAX)).toint() + while _load_acquire(qready_addr) != head_e + cutlass.Int32(1): + pass + cute.arch.sync_warp() + if cutlass.const_expr(FOLD): + xrow0 = bidx * M_MAX + prims.barrier_cta_sync(QUANT_BAR_ID, thread_count=QUANT_THREADS + 32) + g = 0 + ab_empty_phase = 1 + t = cutlass.Int32(0) + dyn_ph = 0 + n_fc1_done = cutlass.Int32(0) + v = _next_tile( + t, False, dyn_ph, bidx, total_tasks, dyn_slot, dyn_ready, dyn_consumed, dyn_ctr_ptr + ) + in_fc1 = (v >= cutlass.Int32(0)) & (v < tiles_fc1) + while in_fc1: + group = v // M_TILES_FC1 + if cutlass.const_expr(WIDE): + # The slots' token rows (empty slots: row 0). + grp_cnt = g_cnt.load(idx=group) + rw = [] + for n in cutlass.range_constexpr(N): + pair_n = cutlass.Int32(g_pair.load(idx=group * N + n)) + rw.append( + cutlass.select_( + cutlass.Int32(n) < grp_cnt, pair_n // TOP_K, cutlass.Int32(0) + ) + ) + r0, r1, r2, r3, r4, r5, r6, r7 = rw + sfb_pair = cutlass.Int32(g_pair.load(idx=group * N + sfb_lane)) + sfb_token = cutlass.select_(sfb_lane < grp_cnt, sfb_pair // TOP_K, cutlass.Int32(0)) + else: + rows = g_rows.load(idx=group) + r0 = (rows & cutlass.Int32(0xF)) + xrow0 + r1 = ((rows >> cutlass.Int32(4)) & cutlass.Int32(0xF)) + xrow0 + r2 = ((rows >> cutlass.Int32(8)) & cutlass.Int32(0xF)) + xrow0 + r3 = ((rows >> cutlass.Int32(12)) & cutlass.Int32(0xF)) + xrow0 + r4 = ((rows >> cutlass.Int32(16)) & cutlass.Int32(0xF)) + xrow0 + r5 = ((rows >> cutlass.Int32(20)) & cutlass.Int32(0xF)) + xrow0 + r6 = ((rows >> cutlass.Int32(24)) & cutlass.Int32(0xF)) + xrow0 + r7 = ((rows >> cutlass.Int32(28)) & cutlass.Int32(0xF)) + xrow0 + grp_cnt = g_cnt.load(idx=group) + sfb_token = ((rows >> (sfb_lane * cutlass.Int32(4))) & cutlass.Int32(0xF)) + xrow0 + sfb_cp_sz = cutlass.select_(lane < grp_cnt, cutlass.Int32(4), cutlass.Int32(0)) + for k_tile in cutlass.range(K1_TILES, unroll=1): + stage = g % num_ab_stage + if stage == 0 and g != 0: + ab_empty_phase = ab_empty_phase ^ 1 + while not cute.arch.mbarrier_try_wait( + ab_empty.subview(stage).data_ptr(), ab_empty_phase + ): + pass + coord_k = k_tile * mma_tiler_mnk[2] + if prims.elect_sync(): + prims.mbarrier_arrive_expect_tx(ab_full_fc1.subview(stage), NUM_BYTES_B) + _tma_gather4_cta( + sB.subview(NUM_BYTES_B * stage), + tma_b1_desc.get_ptr(), + coord_k, + r0, + r1, + r2, + r3, + ab_full_fc1.subview(stage), + ) + _tma_gather4_cta( + sB.subview(NUM_BYTES_B * stage + NUM_BYTES_B // 2), + tma_b1_desc.get_ptr(), + coord_k, + r4, + r5, + r6, + r7, + ab_full_fc1.subview(stage), + ) + if lane < cutlass.Int32(N): + sfb_gmem = ( + sfb1_ptr + + cutlass.Int64(sfb_token) * cutlass.Int64(SFB_SRC_STRIDE_FC1) + + cutlass.Int64(k_tile * NUM_KBLOCKS) + ) + sfb_gmem_ir = cutlass.inttoptr(sfb_gmem, mem_space=1, dtype=sf_dtype) + sfb_smem_ir = sSFB.subview( + NUM_BYTES_SFB * stage + lane * SFB_GROUP_BYTES + ).data_ptr() + prims.cp_async_shared_global( + sfb_smem_ir, sfb_gmem_ir, size=4, modifier="ca", cp_size=sfb_cp_sz + ) + prims.cp_async_mbarrier_arrive(ab_full_fc1.subview(stage), noinc=True) + g = g + 1 + n_fc1_done = n_fc1_done + cutlass.Int32(1) + t = t + cutlass.Int32(1) + if cutlass.const_expr(DYN_ALL): + if t % cutlass.Int32(DYN_RING) == cutlass.Int32(0): + dyn_ph = dyn_ph ^ 1 + v = _next_tile( + t, False, dyn_ph, bidx, total_tasks, dyn_slot, dyn_ready, dyn_consumed, dyn_ctr_ptr + ) + in_fc1 = (v >= cutlass.Int32(0)) & (v < tiles_fc1) + + # FC1 -> FC2 activation delay: CTAs that drew fewer FC1 tiles start later. + if cutlass.const_expr(FC2_DELAY_NS > 0 or FC2_DELAY_LAG_NS > 0): + if (v >= cutlass.Int32(0)) & (v < total_tiles): + delay_ns = cutlass.Int64(FC2_DELAY_NS) + if cutlass.const_expr(FC2_DELAY_LAG_NS > 0): + n1_max = (tiles_fc1 + NUM_CTAS - 1) // NUM_CTAS + lag = cutlass.select_( + n1_max > n_fc1_done, n1_max - n_fc1_done, cutlass.Int32(0) + ) + delay_ns = delay_ns + cutlass.Int64(FC2_DELAY_LAG_NS) * cutlass.Int64(lag) + t_end = _read_globaltimer() + delay_ns + while _read_globaltimer() < t_end: + _backoff(32, delay=True) + + # counter: the loads complete on the MMA's barrier directly (the scan's arrival is this one). + acts_fc2_full = lamport_arrived if SCAN_FC2 else ab_full_fc2 + in_fc2 = (v >= cutlass.Int32(0)) & (v < total_tiles) + while in_fc2: + lin2 = v - tiles_fc1 + g_lo = lin2 // M_TILES_FC2 + g_hi = g_lo + cutlass.Int32(1) + if cutlass.const_expr(SLICE): + sl = lin2 // M_TILES_FC2 + if cutlass.const_expr(WIDE): + sl = lin2 % num_slices + g_lo = sl * num_groups // num_slices + g_hi = (sl + cutlass.Int32(1)) * num_groups // num_slices + for group in cutlass.range(g_lo, g_hi, unroll=1): + if cutlass.const_expr(FC1_COUNTS): + # The group's intermediate is complete once its FC1 tiles have all counted + # (every lane polls the same word: one request, no divergence). + fc1_ctr = state_ptr + cutlass.Int64((ST_FC1 + group) * 4) + if cutlass.const_expr(SCAN_FC2): + while _load_relaxed(fc1_ctr) < cutlass.Int32(FC2_SYNC_NEED): + pass + else: + while _load_acquire(fc1_ctr) < cutlass.Int32(FC2_SYNC_NEED): + pass + prims.fence_proxy("async_global") + for k_tile in cutlass.range(K2_TILES, unroll=1): + stage = g % num_ab_stage + if stage == 0 and g != 0: + ab_empty_phase = ab_empty_phase ^ 1 + while not cute.arch.mbarrier_try_wait( + ab_empty.subview(stage).data_ptr(), ab_empty_phase + ): + pass + coord_k = k_tile * mma_tiler_mnk[2] + if prims.elect_sync(): + prims.mbarrier_arrive_expect_tx( + acts_fc2_full.subview(stage), NUM_TMA_LOAD_BYTES_ACTS_FC2 + ) + prims.cp_async_bulk_tensor_shared_cta_global( + sB.subview(NUM_BYTES_B * stage), + tma_b2_desc.get_ptr(), + (coord_k, coord_n, group), + acts_fc2_full.subview(stage), + ) + prims.cp_async_bulk_tensor_shared_cta_global( + sSFB.subview(NUM_BYTES_SFB * stage), + tma_sfb2_desc.get_ptr(), + (cutlass.Int32(0), k_tile, cutlass.Int32(0), group), + acts_fc2_full.subview(stage), + ) + g = g + 1 + t = t + cutlass.Int32(1) + if cutlass.const_expr(DYN_ALL): + if t % cutlass.Int32(DYN_RING) == cutlass.Int32(0): + dyn_ph = dyn_ph ^ 1 + v = _next_tile( + t, False, dyn_ph, bidx, total_tasks, dyn_slot, dyn_ready, dyn_consumed, dyn_ctr_ptr + ) + in_fc2 = (v >= cutlass.Int32(0)) & (v < total_tiles) + while v >= cutlass.Int32(0): # the all-reduce tasks are the epilogue's + t = t + cutlass.Int32(1) + if cutlass.const_expr(DYN_ALL): + if t % cutlass.Int32(DYN_RING) == cutlass.Int32(0): + dyn_ph = dyn_ph ^ 1 + v = _next_tile( + t, False, dyn_ph, bidx, total_tasks, dyn_slot, dyn_ready, dyn_consumed, dyn_ctr_ptr + ) + + # ------------------------------------------------ producerScalesTmem (6) + if warp_idx == scales_tmem_warp_id and active: + while not cute.arch.mbarrier_try_wait(tmem_ready.data_ptr(), 0): + pass + tmem_raw_addr = tmem_ptr_i32.load() + base_col_id = tmem_raw_addr & 0xFFFF + base_row_id = tmem_raw_addr >> 16 + sfa_cols_per_stage = rest_k_sf * rest_m_sf * num_tmem_cols_per_sf_atom + sfb_cols_per_stage = rest_k_sf * rest_n_sf * num_tmem_cols_per_sf_atom + sfa_col_id0 = base_col_id + mma_tiler_mnk[1] + sfb_col_id0 = sfa_col_id0 + num_ab_stage * sfa_cols_per_stage + s2t_shape, s2t_multicast = prims.S2TCopyMode.S2T_32x128b_WARPX4 + g = 0 + ab_full_fc1_phase = 0 + t = cutlass.Int32(0) + dyn_ph = 0 + v = _next_tile( + t, False, dyn_ph, bidx, total_tasks, dyn_slot, dyn_ready, dyn_consumed, dyn_ctr_ptr + ) + in_fc1 = (v >= cutlass.Int32(0)) & (v < tiles_fc1) + while in_fc1: + for k_tile in cutlass.range(K1_TILES, unroll=1): + stage = g % num_ab_stage + if stage == 0 and g != 0: + ab_full_fc1_phase = ab_full_fc1_phase ^ 1 + while not cute.arch.mbarrier_try_wait( + ab_full_fc1.subview(stage).data_ptr(), ab_full_fc1_phase + ): + pass + # The activation scales came by cp.async (generic proxy); tcgen05.cp reads them (async proxy). + prims.fence_proxy("async_shared", space=prims.SharedSpace.shared_cta) + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + sfa_tmem_ptr = cutlass.inttoptr( + (base_row_id << 16) | (sfa_col_id0 + stage * sfa_cols_per_stage), + 6, + cutlass.Int32, + ) + sfb_tmem_ptr = cutlass.inttoptr( + (base_row_id << 16) | (sfb_col_id0 + stage * sfb_cols_per_stage), + 6, + cutlass.Int32, + ) + desc_a = prims.Tcgen05SmemDesc.build( + sSFA.subview(stage * NUM_BYTES_SFA), + leading_byte_offset=16, + stride_byte_offset=128, + base_offset=0, + layout=0, + ) + desc_b = prims.Tcgen05SmemDesc.build( + sSFB.subview(stage * NUM_BYTES_SFB), + leading_byte_offset=16, + stride_byte_offset=128, + base_offset=0, + layout=0, + ) + if prims.elect_sync(): + prims.tcgen05_cp(s2t_shape, sfa_tmem_ptr, desc_a, multicast=s2t_multicast) + prims.tcgen05_cp(s2t_shape, sfb_tmem_ptr, desc_b, multicast=s2t_multicast) + prims.tcgen05_commit(scales_in_tmem.subview(stage)) + g = g + 1 + t = t + cutlass.Int32(1) + if cutlass.const_expr(DYN_ALL): + if t % cutlass.Int32(DYN_RING) == cutlass.Int32(0): + dyn_ph = dyn_ph ^ 1 + v = _next_tile( + t, False, dyn_ph, bidx, total_tasks, dyn_slot, dyn_ready, dyn_consumed, dyn_ctr_ptr + ) + in_fc1 = (v >= cutlass.Int32(0)) & (v < tiles_fc1) + j = 0 + ab_full_fc2_phase = 0 + in_fc2 = (v >= cutlass.Int32(0)) & (v < total_tiles) + while in_fc2: + lin2 = v - tiles_fc1 + g_lo = lin2 // M_TILES_FC2 + g_hi = g_lo + cutlass.Int32(1) + if cutlass.const_expr(SLICE): + sl = lin2 // M_TILES_FC2 + if cutlass.const_expr(WIDE): + sl = lin2 % num_slices + g_lo = sl * num_groups // num_slices + g_hi = (sl + cutlass.Int32(1)) * num_groups // num_slices + for _group in cutlass.range(g_lo, g_hi, unroll=1): + for k_tile in cutlass.range(K2_TILES, unroll=1): + stage = g % num_ab_stage + if j % num_ab_stage == 0 and j != 0: + ab_full_fc2_phase = ab_full_fc2_phase ^ 1 + while not cute.arch.mbarrier_try_wait( + ab_full_fc2.subview(stage).data_ptr(), ab_full_fc2_phase + ): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + sfa_tmem_ptr = cutlass.inttoptr( + (base_row_id << 16) | (sfa_col_id0 + stage * sfa_cols_per_stage), + 6, + cutlass.Int32, + ) + sfb_tmem_ptr = cutlass.inttoptr( + (base_row_id << 16) | (sfb_col_id0 + stage * sfb_cols_per_stage), + 6, + cutlass.Int32, + ) + desc_a = prims.Tcgen05SmemDesc.build( + sSFA.subview(stage * NUM_BYTES_SFA), + leading_byte_offset=16, + stride_byte_offset=128, + base_offset=0, + layout=0, + ) + desc_b = prims.Tcgen05SmemDesc.build( + sSFB.subview(stage * NUM_BYTES_SFB), + leading_byte_offset=16, + stride_byte_offset=128, + base_offset=0, + layout=0, + ) + if prims.elect_sync(): + prims.tcgen05_cp(s2t_shape, sfa_tmem_ptr, desc_a, multicast=s2t_multicast) + prims.tcgen05_cp(s2t_shape, sfb_tmem_ptr, desc_b, multicast=s2t_multicast) + prims.tcgen05_commit(scales_in_tmem.subview(stage)) + g = g + 1 + j = j + 1 + t = t + cutlass.Int32(1) + if cutlass.const_expr(DYN_ALL): + if t % cutlass.Int32(DYN_RING) == cutlass.Int32(0): + dyn_ph = dyn_ph ^ 1 + v = _next_tile( + t, False, dyn_ph, bidx, total_tasks, dyn_slot, dyn_ready, dyn_consumed, dyn_ctr_ptr + ) + in_fc2 = (v >= cutlass.Int32(0)) & (v < total_tiles) + while v >= cutlass.Int32(0): # the all-reduce tasks are the epilogue's + t = t + cutlass.Int32(1) + if cutlass.const_expr(DYN_ALL): + if t % cutlass.Int32(DYN_RING) == cutlass.Int32(0): + dyn_ph = dyn_ph ^ 1 + v = _next_tile( + t, False, dyn_ph, bidx, total_tasks, dyn_slot, dyn_ready, dyn_consumed, dyn_ctr_ptr + ) + + # ------------------------------------------------ MMA (7) + if warp_idx == consumer_warp_id and active: + tmem_raw_addr = tmem_ptr_i32.load() + acc_tmem_ptr = cutlass.inttoptr(tmem_raw_addr, 6, cutlass.Float32) + idesc = prims.Tcgen05MxInstrDesc.build( + a_dtype=a_dtype, + b_dtype=b_dtype, + scale_format=1, + n_dim=mma_tiler_mnk[1], + m_dim=mma_tiler_mnk[0], + ) + base_col_id = tmem_raw_addr & 0xFFFF + base_row_id = tmem_raw_addr >> 16 + sfa_cols_per_stage = rest_k_sf * rest_m_sf * num_tmem_cols_per_sf_atom + sfb_cols_per_stage = rest_k_sf * rest_n_sf * num_tmem_cols_per_sf_atom + sfa_col_id0 = base_col_id + mma_tiler_mnk[1] + sfb_col_id0 = sfa_col_id0 + num_ab_stage * sfa_cols_per_stage + num_kblocks = mma_tiler_mnk[2] // mma_inst_mnk[2] + num_sf_ids = num_k_per_sf_atom * sf_vec_size // mma_inst_mnk[2] + g = 0 + scales_in_tmem_phase = 0 + acc_empty_phase = 1 + t = cutlass.Int32(0) + dyn_ph = 0 + v = _next_tile( + t, False, dyn_ph, bidx, total_tasks, dyn_slot, dyn_ready, dyn_consumed, dyn_ctr_ptr + ) + in_tiles = (v >= cutlass.Int32(0)) & (v < total_tiles) + while in_tiles: + # FC1 and FC2 tiles share the MMA body; only the k-tile count differs. + k_tiles = cutlass.select_( + v < tiles_fc1, cutlass.Int32(K1_TILES), cutlass.Int32(K2_TILES) + ) + # One accumulation per FC1 tile and per group of an FC2 task (one in tile mode). + n_acc = cutlass.Int32(1) + if cutlass.const_expr(SLICE): + lin2 = v - tiles_fc1 + sl = lin2 // M_TILES_FC2 + if cutlass.const_expr(WIDE): + sl = lin2 % num_slices + n_acc = cutlass.select_( + v < tiles_fc1, + cutlass.Int32(1), + (sl + cutlass.Int32(1)) * num_groups // num_slices + - sl * num_groups // num_slices, + ) + for _acc_i in cutlass.range(n_acc, unroll=1): + while not cute.arch.mbarrier_try_wait(acc_empty.data_ptr(), acc_empty_phase): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc_empty_phase = acc_empty_phase ^ 1 + scale_d = False + for k_tile in cutlass.range(k_tiles, unroll=1): + stage = g % num_ab_stage + if stage == 0 and g != 0: + scales_in_tmem_phase = scales_in_tmem_phase ^ 1 + while not cute.arch.mbarrier_try_wait( + scales_in_tmem.subview(stage).data_ptr(), scales_in_tmem_phase + ): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + sfa_tmem_addr_base = (base_row_id << 16) | ( + sfa_col_id0 + stage * sfa_cols_per_stage + ) + sfb_tmem_addr_base = (base_row_id << 16) | ( + sfb_col_id0 + stage * sfb_cols_per_stage + ) + desc_a_mma_base = prims.Tcgen05SmemDesc.build( + sA.subview(stage * NUM_BYTES_A), + leading_byte_offset=16, + stride_byte_offset=1024, + base_offset=0, + layout=2, + ) + desc_b_mma_base = prims.Tcgen05SmemDesc.build( + sB.subview(stage * NUM_BYTES_B), + leading_byte_offset=16, + stride_byte_offset=1024, + base_offset=0, + layout=2, + ) + for kblock_idx in cutlass.range(num_kblocks, unroll_full=True): + sf_inside = kblock_idx % num_sf_ids + sf_col = kblock_idx // num_sf_ids + sfa_tmem_ptr = cutlass.inttoptr( + sfa_tmem_addr_base + sf_col * NUM_TMEM_COLS_PER_KBLOCK_SFA, + 6, + cutlass.Int32, + ) + sfb_tmem_ptr = cutlass.inttoptr( + sfb_tmem_addr_base + sf_col * NUM_TMEM_COLS_PER_KBLOCK_SFB, + 6, + cutlass.Int32, + ) + idesc_u = idesc.set_sf_ids(a_sf_id=sf_inside, b_sf_id=sf_inside) + inc = ((mma_inst_mnk[2] * a_smem_width // 8) >> 4) * kblock_idx + if prims.elect_sync(): + prims.tcgen05_mma_block_scale( + prims.MMABlockScaleKind.MXF8F6F4, prims.CTAGroup.CTA_1, acc_tmem_ptr, + desc_a_mma_base + inc, desc_b_mma_base + inc, idesc_u, scale_d, sfa_tmem_ptr, + sfb_tmem_ptr, + ) # fmt: skip + scale_d = True + if prims.elect_sync(): + prims.tcgen05_commit(ab_empty.subview(stage)) + g = g + 1 + if prims.elect_sync(): + prims.tcgen05_commit(acc_full) + t = t + cutlass.Int32(1) + if cutlass.const_expr(DYN_ALL): + if t % cutlass.Int32(DYN_RING) == cutlass.Int32(0): + dyn_ph = dyn_ph ^ 1 + v = _next_tile( + t, False, dyn_ph, bidx, total_tasks, dyn_slot, dyn_ready, dyn_consumed, dyn_ctr_ptr + ) + in_tiles = (v >= cutlass.Int32(0)) & (v < total_tiles) + while v >= cutlass.Int32(0): # the all-reduce tasks are the epilogue's + t = t + cutlass.Int32(1) + if cutlass.const_expr(DYN_ALL): + if t % cutlass.Int32(DYN_RING) == cutlass.Int32(0): + dyn_ph = dyn_ph ^ 1 + v = _next_tile( + t, False, dyn_ph, bidx, total_tasks, dyn_slot, dyn_ready, dyn_consumed, dyn_ctr_ptr + ) + + # ------------------------------------------------ lamport (8..11): FC2 only + if warp_idx >= lamport_acts_warp_id and active: + lane = tidx % 32 + my_group = warp_idx - cutlass.Int32(lamport_acts_warp_id) + stages_per_warp = num_ab_stage // NUM_LAMPORT_WARPS + g = 0 + t = cutlass.Int32(0) + dyn_ph = 0 + v = _next_tile( + t, False, dyn_ph, bidx, total_tasks, dyn_slot, dyn_ready, dyn_consumed, dyn_ctr_ptr + ) + in_fc1 = (v >= cutlass.Int32(0)) & (v < tiles_fc1) + while in_fc1: # FC1 stages need no validation; keep ring position and claim ring in step + g = g + K1_TILES + t = t + cutlass.Int32(1) + if cutlass.const_expr(DYN_ALL): + if t % cutlass.Int32(DYN_RING) == cutlass.Int32(0): + dyn_ph = dyn_ph ^ 1 + v = _next_tile( + t, False, dyn_ph, bidx, total_tasks, dyn_slot, dyn_ready, dyn_consumed, dyn_ctr_ptr + ) + in_fc1 = (v >= cutlass.Int32(0)) & (v < tiles_fc1) + j = 0 + lam_phase = 0 + retry_phase = 0 + in_fc2 = (v >= cutlass.Int32(0)) & (v < total_tiles) + while in_fc2: + lin2 = v - tiles_fc1 + g_lo = lin2 // M_TILES_FC2 + g_hi = g_lo + cutlass.Int32(1) + if cutlass.const_expr(SLICE): + sl = lin2 // M_TILES_FC2 + if cutlass.const_expr(WIDE): + sl = lin2 % num_slices + g_lo = sl * num_groups // num_slices + g_hi = (sl + cutlass.Int32(1)) * num_groups // num_slices + if cutlass.const_expr(SCAN_FC2): # counter: the producer acquired the group + for group in cutlass.range(g_lo, g_hi, unroll=1): + for k_tile in cutlass.range(K2_TILES, unroll=1): + stage = g % num_ab_stage + if j % num_ab_stage == 0 and j != 0: + lam_phase = lam_phase ^ 1 + if (stage // stages_per_warp) == my_group: + while not cute.arch.mbarrier_try_wait( + lamport_arrived.subview(stage).data_ptr(), lam_phase + ): + pass + coord_k = k_tile * mma_tiler_mnk[2] + need_retry = _scan_b_sentinel(sB, stage, lane) | _scan_sfb_sentinel( + sSFB, stage, lane + ) + while need_retry: + if prims.elect_sync(): + prims.mbarrier_arrive_expect_tx( + lamport_retry.subview(my_group), NUM_TMA_LOAD_BYTES_ACTS_FC2 + ) + prims.cp_async_bulk_tensor_shared_cta_global( + sB.subview(NUM_BYTES_B * stage), + tma_b2_desc.get_ptr(), + (coord_k, coord_n, group), + lamport_retry.subview(my_group), + ) + prims.cp_async_bulk_tensor_shared_cta_global( + sSFB.subview(NUM_BYTES_SFB * stage), + tma_sfb2_desc.get_ptr(), + (cutlass.Int32(0), k_tile, cutlass.Int32(0), group), + lamport_retry.subview(my_group), + ) + while not cute.arch.mbarrier_try_wait( + lamport_retry.subview(my_group).data_ptr(), retry_phase + ): + pass + retry_phase = retry_phase ^ 1 + need_retry = _scan_b_sentinel(sB, stage, lane) | _scan_sfb_sentinel( + sSFB, stage, lane + ) + if prims.elect_sync(): + prims.mbarrier_arrive(ab_full_fc2.subview(stage)) + g = g + 1 + j = j + 1 + t = t + cutlass.Int32(1) + if cutlass.const_expr(DYN_ALL): + if t % cutlass.Int32(DYN_RING) == cutlass.Int32(0): + dyn_ph = dyn_ph ^ 1 + v = _next_tile( + t, False, dyn_ph, bidx, total_tasks, dyn_slot, dyn_ready, dyn_consumed, dyn_ctr_ptr + ) + in_fc2 = (v >= cutlass.Int32(0)) & (v < total_tiles) + while v >= cutlass.Int32(0): # the all-reduce tasks are the epilogue's + t = t + cutlass.Int32(1) + if cutlass.const_expr(DYN_ALL): + if t % cutlass.Int32(DYN_RING) == cutlass.Int32(0): + dyn_ph = dyn_ph ^ 1 + v = _next_tile( + t, False, dyn_ph, bidx, total_tasks, dyn_slot, dyn_ready, dyn_consumed, dyn_ctr_ptr + ) + + # ------------------------------------------------ epilogue (0-3) + if warp_idx < len(epilog_warp_id) and (active or (FUSED_AR and not AR_PUSH_ONLY)): + # With HEAD_FLAGS, FC1 writes only the op's own buffers (c, its scales), so it runs while the + # producer grid (route_quant_ag, or the MoE front's shared tiles) finishes; the grid wait comes + # before the FC2 phase, whose combine and all-reduce write the output (the allocator's memory). + while not cute.arch.mbarrier_try_wait(tmem_ready.data_ptr(), 0): + pass + tmem_raw_addr = tmem_ptr_i32.load() + base_col_id = tmem_raw_addr & 0xFFFF + base_row_id = tmem_raw_addr >> 16 + row_id_with_warp_offset = base_row_id + warp_idx * 32 + t2r_repx = min(32, mma_tiler_mnk[1]) + lane = tidx % 32 + is_up = ((lane // 8) % 2) == 0 + up_mask = cutlass.select_(is_up, cutlass.Float32(1.0), cutlass.Float32(0.0)) + fc1_col_in_tile = warp_idx * 16 + 2 * (lane % 8) + lane // 16 + fc2_ch_in_tile = warp_idx * 32 + 4 * (lane % 8) + lane // 8 + acc_full_phase = 0 + t = cutlass.Int32(0) + dyn_ph = 0 + v = _next_tile( + t, False, dyn_ph, bidx, total_tasks, dyn_slot, dyn_ready, dyn_consumed, dyn_ctr_ptr + ) + in_fc1 = (v >= cutlass.Int32(0)) & (v < tiles_fc1) + while in_fc1: + m_tile = v % M_TILES_FC1 + group = v // M_TILES_FC1 + while not cute.arch.mbarrier_try_wait(acc_full.data_ptr(), acc_full_phase): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc_full_phase = acc_full_phase ^ 1 + tmem_ld = cutlass.inttoptr( + (row_id_with_warp_offset << 16) | base_col_id, 6, cutlass.Float32 + ) + t2r_rmem = prims.tcgen05_ld("32x32b", tmem_ld, num=t2r_repx) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + prims.mbarrier_arrive(acc_empty) + out_col = m_tile * (MMA_M // 2) + fc1_col_in_tile + block_idx = m_tile * (MMA_M // 2 // MX_BLOCK) + (warp_idx // 2) + feat_off = (block_idx % SF_TMEM_COL) + (block_idx // SF_TMEM_COL) * ( + SF_TMEM_COL * SF_TMEM_DP + ) + gbase = group * N + for n in cutlass.range_constexpr(N): + xn = cutlass.Float32(t2r_rmem[n]) + partner = cute.arch.shuffle_sync_bfly(xn, 8) + res = _situ(partner, xn) + absv = cute.math.abs(res) * up_mask + warp_amax = prims.redux_sync(absv, prims.ReductionKind.FMAX, 0xFFFFFFFF, abs=True) + if lane == 0: + localmax_smem.store(warp_amax, idx=(n % 2) * 4 + warp_idx) + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + partner_amax = localmax_smem.load(idx=(n % 2) * 4 + (warp_idx ^ 1)) + block_amax = cute.arch.fmax(warp_amax, partner_amax) + byte, inv_scale = _block_e8m0(block_amax) + qv = cute.arch.fmax( + cute.arch.fmin(res * inv_scale, cutlass.Float32(E4M3_MAX)), + cutlass.Float32(-E4M3_MAX), + ) + fp8_i8 = cutlass.Float8E4M3FN(qv).bitcast(cutlass.Int8) + if fp8_i8 == cutlass.Int8(FP8_SENTINEL_I8): + fp8_i8 = cutlass.Int8(0) + if is_up: + c_tensor.store(fp8_i8, idx=(gbase + n) * I_TP + out_col, alignment=1) + if (tidx % 64) == 0: + sc8 = cutlass.Int8(byte & cutlass.Int32(0xFF)) + if sc8 == cutlass.Int8(SF_SENTINEL_I8): + sc8 = cutlass.Int8(0) + c_scale_tensor.store(sc8, idx=(gbase + n) * SF_STRIDE0 + feat_off, alignment=1) + if cutlass.const_expr(FC1_COUNTS): + fc1_ctr = state_ptr + cutlass.Int64((ST_FC1 + group) * 4) + if cutlass.const_expr(SCAN_FC2): + # hint: counted once every thread has issued its stores; the scan validates. + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + if tidx == 0: + _red_relaxed_add(fc1_ctr, cutlass.Int32(1)) + else: + # Other CTAs' TMA loads read this tile's columns once the group's count is + # complete: every writer fences the async proxy, thread 0 releases after the barrier. + prims.fence_proxy("async_global") + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + if tidx == 0: + _red_release_add(fc1_ctr, cutlass.Int32(1)) + t = t + cutlass.Int32(1) + if cutlass.const_expr(DYN_ALL): + if t % cutlass.Int32(DYN_RING) == cutlass.Int32(0): + dyn_ph = dyn_ph ^ 1 + v = _next_tile( + t, False, dyn_ph, bidx, total_tasks, dyn_slot, dyn_ready, dyn_consumed, dyn_ctr_ptr + ) + in_fc1 = (v >= cutlass.Int32(0)) & (v < tiles_fc1) + + if cutlass.const_expr(HEAD_FLAGS and USE_PDL): + cute.arch.griddepcontrol_wait() + in_fc2 = (v >= cutlass.Int32(0)) & (v < total_tiles) + if cutlass.const_expr(WIDE): + zeros = cutlass.Array(cutlass.Float32, M_MAX, space=cutlass.AddressSpace.rmem) + for i in cutlass.range_constexpr(M_MAX): + zeros[i] = cutlass.Float32(0.0) + if cutlass.const_expr(SLICE): + while in_fc2: + lin2 = v - tiles_fc1 + m_tile = lin2 % M_TILES_FC2 + sl = lin2 // M_TILES_FC2 + if cutlass.const_expr(WIDE): + # FC2 tasks are m-tile-major here: an m-tile's slices finish together, so its combine + # chunks run while later m-tiles stream. + sl = lin2 % num_slices + m_tile = lin2 // num_slices + g_lo = sl * num_groups // num_slices + g_hi = (sl + cutlass.Int32(1)) * num_groups // num_slices + ch = m_tile * MMA_M + fc2_ch_in_tile + zero = cutlass.Float32(0.0) + if cutlass.const_expr(WIDE): + # Per-token sums of this thread's channel over the slice's groups, in group + # order, in TMEM column TOK_COL + token (TOK_COL + M_MAX takes empty slots). + tok_col0 = (row_id_with_warp_offset << 16) | (base_col_id + TOK_COL) + prims.tcgen05_st( + prims.Tcgen05LdStShape.SHAPE_32X32B, + cutlass.inttoptr(tok_col0, 6, cutlass.Float32), + zeros.load(0, M_MAX, alignment=32), + ) + prims.tcgen05_wait(kind=prims.Tcgen05Wait.STORE) + for group in cutlass.range(g_lo, g_hi, unroll=1): + grp_cnt = g_cnt.load(idx=group) + valids = [] + cols = [] + rws = [] + for n in cutlass.range_constexpr(N): + valid = cutlass.Int32(n) < grp_cnt + pair_n = cutlass.select_( + valid, + cutlass.Int32(g_pair.load(idx=group * N + n)), + cutlass.Int32(0), + ) + valids.append(valid) + cols.append( + tok_col0 + + cutlass.select_(valid, pair_n // TOP_K, cutlass.Int32(M_MAX)) + ) + rws.append(cutlass.Float32(topk_w.load(idx=pair_n))) + while not cute.arch.mbarrier_try_wait(acc_full.data_ptr(), acc_full_phase): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc_full_phase = acc_full_phase ^ 1 + tmem_ld = cutlass.inttoptr( + (row_id_with_warp_offset << 16) | base_col_id, 6, cutlass.Float32 + ) + t2r_rmem = prims.tcgen05_ld("32x32b", tmem_ld, num=t2r_repx) + sums = [ + prims.tcgen05_ld( + "32x32b", cutlass.inttoptr(cols[n], 6, cutlass.Float32), num=1 + ) + for n in range(N) + ] + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + prims.mbarrier_arrive(acc_empty) + for n in cutlass.range_constexpr(N): + val = cutlass.select_( + valids[n], cutlass.Float32(t2r_rmem[n]) * rws[n], zero + ) + prims.tcgen05_st( + prims.Tcgen05LdStShape.SHAPE_32X32B, + cutlass.inttoptr(cols[n], 6, cutlass.Float32), + cutlass.Float32(sums[n][0]) + val, + ) + prims.tcgen05_wait(kind=prims.Tcgen05Wait.STORE) + # This task's partial rows, [m-tile][slice][token][128]. + accs = prims.tcgen05_ld( + "32x32b", cutlass.inttoptr(tok_col0, 6, cutlass.Float32), num=M_MAX + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + pbase = (m_tile * S_CAP + sl) * (M_MAX * MMA_M) + fc2_ch_in_tile + mask_lo = slice_mask.load(idx=sl * 2) + mask_hi = slice_mask.load(idx=sl * 2 + 1) + for tok in cutlass.range_constexpr(M_MAX): + word = mask_lo if tok < 32 else mask_hi + if ((word >> cutlass.Int32(tok % 32)) & cutlass.Int32(1)) != cutlass.Int32( + 0 + ): + part_tensor.store(cutlass.Float32(accs[tok]), idx=pbase + tok * MMA_M) + else: + # Per-token sums of this thread's channel over the slice's groups, in group order. + a0 = zero + a1 = zero + a2 = zero + a3 = zero + a4 = zero + a5 = zero + a6 = zero + a7 = zero + for group in cutlass.range(g_lo, g_hi, unroll=1): + rows = g_rows.load(idx=group) + grp_cnt = g_cnt.load(idx=group) + gbase = group * N + while not cute.arch.mbarrier_try_wait(acc_full.data_ptr(), acc_full_phase): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc_full_phase = acc_full_phase ^ 1 + tmem_ld = cutlass.inttoptr( + (row_id_with_warp_offset << 16) | base_col_id, 6, cutlass.Float32 + ) + t2r_rmem = prims.tcgen05_ld("32x32b", tmem_ld, num=t2r_repx) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + prims.mbarrier_arrive(acc_empty) + for n in cutlass.range_constexpr(N): + valid = cutlass.Int32(n) < grp_cnt + val = cutlass.select_( + valid, cutlass.Float32(t2r_rmem[n]) * g_rw.load(idx=gbase + n), zero + ) + tok = cutlass.select_( + valid, + (rows >> cutlass.Int32(4 * n)) & cutlass.Int32(0xF), + cutlass.Int32(-1), + ) + a0 = a0 + cutlass.select_(tok == cutlass.Int32(0), val, zero) + a1 = a1 + cutlass.select_(tok == cutlass.Int32(1), val, zero) + a2 = a2 + cutlass.select_(tok == cutlass.Int32(2), val, zero) + a3 = a3 + cutlass.select_(tok == cutlass.Int32(3), val, zero) + a4 = a4 + cutlass.select_(tok == cutlass.Int32(4), val, zero) + a5 = a5 + cutlass.select_(tok == cutlass.Int32(5), val, zero) + a6 = a6 + cutlass.select_(tok == cutlass.Int32(6), val, zero) + a7 = a7 + cutlass.select_(tok == cutlass.Int32(7), val, zero) + # This task's partial rows, [m-tile][slice][token][128]. + accs = [a0, a1, a2, a3, a4, a5, a6, a7] + pbase = (m_tile * S_CAP + sl) * (M_MAX * MMA_M) + fc2_ch_in_tile + for tok in cutlass.range_constexpr(M_MAX): + if cutlass.Int32(tok) < num_tokens: + part_tensor.store(accs[tok], idx=pbase + tok * MMA_M) + mtile_ctr = state_ptr + cutlass.Int64((ST_MTILE + m_tile) * 4) + is_combiner = sl == num_slices - cutlass.Int32(1) + is_rearmer = m_tile == cutlass.Int32(M_TILES_FC2 - 1) + if cutlass.const_expr(SLICE_TASKS): + is_combiner = cutlass.Boolean(False) + is_rearmer = cutlass.Boolean(False) + # The partial rows are CTA-visible after the barrier; tid 0's release publishes + # them GPU-wide. The m-tile's last slice combines, after the others' arrivals. + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + if cutlass.const_expr(COMBINE_LAST): + if tidx == 0: + arrived = _atomic_fetch_add(mtile_ctr, cutlass.Int32(1)) + last_arrival = arrived == num_slices - cutlass.Int32(1) + epi_flag.store( + cutlass.select_(last_arrival, cutlass.Int32(1), cutlass.Int32(0)), idx=0 + ) + if last_arrival: + _store_release(mtile_ctr, cutlass.Int32(0)) + else: + if tidx == 0: + if is_combiner: + while _load_acquire(mtile_ctr) < num_slices - cutlass.Int32(1): + _backoff(32) + _store_release(mtile_ctr, cutlass.Int32(0)) + else: + _red_release_add(mtile_ctr, cutlass.Int32(1)) + # Every group this task read counts one reader (its MMAs are done); the slice's + # m-tile-27 task re-arms the slice's groups once the other 27 have counted. + if warp_idx == 1: + if cutlass.const_expr(WIDE): + for gb in cutlass.range(g_lo, g_hi, 32, unroll=1): + if (gb + lane) < g_hi: + _red_release_add( + state_ptr + cutlass.Int64((ST_GROUP + gb + lane) * 4), + cutlass.Int32(1), + ) + else: + if (g_lo + lane) < g_hi: + gctr = state_ptr + cutlass.Int64((ST_GROUP + g_lo + lane) * 4) + if is_rearmer: + while _load_acquire(gctr) < cutlass.Int32(M_TILES_FC2 - 1): + _backoff(32) + _store_release(gctr, cutlass.Int32(0)) + else: + _red_release_add(gctr, cutlass.Int32(1)) + # Only a combiner or a re-armer waits for its spinning thread(s); every other task + # moves on to its next accumulator. + if cutlass.const_expr(COMBINE_LAST): + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + is_combiner = epi_flag.load(idx=0) != cutlass.Int32(0) + else: + if is_combiner | is_rearmer: + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + if cutlass.const_expr(not WIDE): # the wide build combines in chunk tasks + if is_combiner: + # The slices' partials summed in slice order. + cute.arch.fence_acq_rel_gpu() + cbase = m_tile * S_CAP * (M_MAX * MMA_M) + fc2_ch_in_tile + c0 = zero + c1 = zero + c2 = zero + c3 = zero + c4 = zero + c5 = zero + c6 = zero + c7 = zero + # 8 slices x 8 tokens of loads in flight per round. + for sb in cutlass.range(0, num_slices, 8, unroll=1): + for q in cutlass.range_constexpr(8): + sq = sb + cutlass.Int32(q) + ok = sq < num_slices + qbase = cbase + cutlass.select_(ok, sq, cutlass.Int32(0)) * ( + M_MAX * MMA_M + ) + vq = [ + part_tensor.load(idx=qbase + tk * MMA_M, is_volatile=True) + for tk in range(M_MAX) + ] + c0 = c0 + cutlass.select_(ok, vq[0], zero) + c1 = c1 + cutlass.select_(ok, vq[1], zero) + c2 = c2 + cutlass.select_(ok, vq[2], zero) + c3 = c3 + cutlass.select_(ok, vq[3], zero) + c4 = c4 + cutlass.select_(ok, vq[4], zero) + c5 = c5 + cutlass.select_(ok, vq[5], zero) + c6 = c6 + cutlass.select_(ok, vq[6], zero) + c7 = c7 + cutlass.select_(ok, vq[7], zero) + cs = [c0, c1, c2, c3, c4, c5, c6, c7] + for tok in cutlass.range_constexpr(M_MAX): + if cutlass.Int32(tok) < num_tokens: + _emit_row( + cs[tok], cutlass.Int32(tok), ch, lane, warp_idx, m_tile, y_tensor, ar_mc, + ar_cur, ar_rank, + ) # fmt: skip + if is_rearmer: + for b in cutlass.range(g_hi - g_lo, unroll=1): + _rearm_group(c_words, cs_words, g_lo + b, tidx) + t = t + cutlass.Int32(1) + if cutlass.const_expr(DYN_ALL): + if t % cutlass.Int32(DYN_RING) == cutlass.Int32(0): + dyn_ph = dyn_ph ^ 1 + v = _next_tile( + t, + False, + dyn_ph, + bidx, + total_tasks, + dyn_slot, + dyn_ready, + dyn_consumed, + dyn_ctr_ptr, + ) + in_fc2 = (v >= cutlass.Int32(0)) & (v < total_tiles) + if cutlass.const_expr(SLICE_TASKS): + zero = cutlass.Float32(0.0) + if cutlass.const_expr(not WIDE): + # Combine tasks: an m-tile's slices summed in slice order once all have published. + in_comb = (v >= total_tiles) & (v < total_tiles + cutlass.Int32(M_TILES_FC2)) + if cutlass.const_expr(not COMBINE_TASKS): + in_comb = cutlass.Boolean(False) + while in_comb: + m_tile = v - total_tiles + mtile_ctr = state_ptr + cutlass.Int64((ST_MTILE + m_tile) * 4) + if tidx == 0: + while _load_acquire(mtile_ctr) < num_slices: + _backoff(32) + _store_release(mtile_ctr, cutlass.Int32(0)) + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + cute.arch.fence_acq_rel_gpu() + ch = m_tile * MMA_M + fc2_ch_in_tile + cbase = m_tile * S_CAP * (M_MAX * MMA_M) + fc2_ch_in_tile + c0 = zero + c1 = zero + c2 = zero + c3 = zero + c4 = zero + c5 = zero + c6 = zero + c7 = zero + for sb in cutlass.range(0, num_slices, 8, unroll=1): + for q in cutlass.range_constexpr(8): + sq = sb + cutlass.Int32(q) + ok = sq < num_slices + qbase = cbase + cutlass.select_(ok, sq, cutlass.Int32(0)) * ( + M_MAX * MMA_M + ) + vq = [ + part_tensor.load(idx=qbase + tk * MMA_M, is_volatile=True) + for tk in range(M_MAX) + ] + c0 = c0 + cutlass.select_(ok, vq[0], zero) + c1 = c1 + cutlass.select_(ok, vq[1], zero) + c2 = c2 + cutlass.select_(ok, vq[2], zero) + c3 = c3 + cutlass.select_(ok, vq[3], zero) + c4 = c4 + cutlass.select_(ok, vq[4], zero) + c5 = c5 + cutlass.select_(ok, vq[5], zero) + c6 = c6 + cutlass.select_(ok, vq[6], zero) + c7 = c7 + cutlass.select_(ok, vq[7], zero) + cs = [c0, c1, c2, c3, c4, c5, c6, c7] + for tok in cutlass.range_constexpr(M_MAX): + if cutlass.Int32(tok) < num_tokens: + _emit_row( + cs[tok], cutlass.Int32(tok), ch, lane, warp_idx, m_tile, y_tensor, ar_mc, + ar_cur, ar_rank, + ) # fmt: skip + t = t + cutlass.Int32(1) + if cutlass.const_expr(DYN_ALL): + if t % cutlass.Int32(DYN_RING) == cutlass.Int32(0): + dyn_ph = dyn_ph ^ 1 + v = _next_tile( + t, + False, + dyn_ph, + bidx, + total_tasks, + dyn_slot, + dyn_ready, + dyn_consumed, + dyn_ctr_ptr, + ) + in_comb = (v >= total_tiles) & ( + v < total_tiles + cutlass.Int32(M_TILES_FC2) + ) + if cutlass.const_expr(COMBINE_CHUNKS): + # Combine tasks, one per (m-tile, 8 tokens): the m-tile's slices summed in slice + # order once all have published; the m-tile's counter then counts the chunks that + # have read it, and the last one resets it. + n_comb = cutlass.Int32(M_TILES_FC2) * n_chunks + in_chunk = (v >= total_tiles) & (v < total_tiles + n_comb) + while in_chunk: + cidx = v - total_tiles + m_tile = cidx // n_chunks + tok0 = (cidx - m_tile * n_chunks) * cutlass.Int32(N) + mtile_ctr = state_ptr + cutlass.Int64((ST_MTILE + m_tile) * 4) + if tidx == 0: + while _load_acquire(mtile_ctr) < num_slices: + _backoff(32) + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + cute.arch.fence_acq_rel_gpu() + ch = m_tile * MMA_M + fc2_ch_in_tile + cbase = m_tile * S_CAP * (M_MAX * MMA_M) + tok0 * MMA_M + fc2_ch_in_tile + c0 = zero + c1 = zero + c2 = zero + c3 = zero + c4 = zero + c5 = zero + c6 = zero + c7 = zero + mword = tok0 >> cutlass.Int32(5) + mshift = tok0 & cutlass.Int32(31) + # 8 slices x 8 tokens of loads in flight per round; a slice without one of these + # tokens stored nothing for it, and it adds zero (as the stored zero would). + for sb in cutlass.range(0, num_slices, 8, unroll=1): + vals = [] + for q in cutlass.range_constexpr(8): + sq = sb + cutlass.Int32(q) + ok = sq < num_slices + sq_c = cutlass.select_(ok, sq, cutlass.Int32(0)) + bits = cutlass.select_( + ok, + (slice_mask.load(idx=sq_c * 2 + mword) >> mshift) + & cutlass.Int32(0xFF), + cutlass.Int32(0), + ) + qbase = cbase + sq_c * (M_MAX * MMA_M) + for tk in cutlass.range_constexpr(N): + pv = zero + if ( + (bits >> cutlass.Int32(tk)) & cutlass.Int32(1) + ) != cutlass.Int32(0): + pv = part_tensor.load( + idx=qbase + tk * MMA_M, is_volatile=True + ) + vals.append(pv) + for q in cutlass.range_constexpr(8): + c0 = c0 + vals[q * N + 0] + c1 = c1 + vals[q * N + 1] + c2 = c2 + vals[q * N + 2] + c3 = c3 + vals[q * N + 3] + c4 = c4 + vals[q * N + 4] + c5 = c5 + vals[q * N + 5] + c6 = c6 + vals[q * N + 6] + c7 = c7 + vals[q * N + 7] + cs = [c0, c1, c2, c3, c4, c5, c6, c7] + for k in cutlass.range_constexpr(N): + if tok0 + cutlass.Int32(k) < num_tokens: + _emit_row( + cs[k], tok0 + cutlass.Int32(k), ch, lane, warp_idx, m_tile, + y_tensor, ar_mc, ar_cur, ar_rank, + ) # fmt: skip + if tidx == 0: + seen = _atomic_fetch_add(mtile_ctr, cutlass.Int32(1)) + if seen == num_slices + n_chunks - cutlass.Int32(1): + _store_release(mtile_ctr, cutlass.Int32(0)) + prims.mbarrier_arrive(epi_free) + t = t + cutlass.Int32(1) + if cutlass.const_expr(DYN_ALL): + if t % cutlass.Int32(DYN_RING) == cutlass.Int32(0): + dyn_ph = dyn_ph ^ 1 + v = _next_tile( + t, + False, + dyn_ph, + bidx, + total_tasks, + dyn_slot, + dyn_ready, + dyn_consumed, + dyn_ctr_ptr, + ) + in_chunk = (v >= total_tiles) & (v < total_tiles + n_comb) + # Re-arm tasks: a slice's groups re-armed once all 28 of their readers have counted + # (and their FC1 counts reset; counter: only the counts, nothing scans the intermediate). + if cutlass.const_expr(COMBINE_CHUNKS): + rearm0 = total_tiles + cutlass.Int32(M_TILES_FC2) * n_chunks + else: + rearm0 = total_tiles + cutlass.Int32(M_TILES_FC2 if COMBINE_TASKS else 0) + in_rearm = v >= rearm0 + while in_rearm: + sl = v - rearm0 + g_lo = sl * num_groups // num_slices + g_hi = (sl + cutlass.Int32(1)) * num_groups // num_slices + if warp_idx == 1: + if cutlass.const_expr(WIDE): + for gb in cutlass.range(g_lo, g_hi, 32, unroll=1): + if (gb + lane) < g_hi: + gctr = state_ptr + cutlass.Int64((ST_GROUP + gb + lane) * 4) + while _load_acquire(gctr) < cutlass.Int32(M_TILES_FC2): + _backoff(32) + _store_release(gctr, cutlass.Int32(0)) + _store_release( + state_ptr + cutlass.Int64((ST_FC1 + gb + lane) * 4), + cutlass.Int32(0), + ) + else: + if (g_lo + lane) < g_hi: + gctr = state_ptr + cutlass.Int64((ST_GROUP + g_lo + lane) * 4) + while _load_acquire(gctr) < cutlass.Int32(M_TILES_FC2): + _backoff(32) + _store_release(gctr, cutlass.Int32(0)) + if cutlass.const_expr(FC1_COUNTS): + _store_release( + state_ptr + cutlass.Int64((ST_FC1 + g_lo + lane) * 4), + cutlass.Int32(0), + ) + if cutlass.const_expr(SCAN_FC2): + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + for b in cutlass.range(g_hi - g_lo, unroll=1): + _rearm_group(c_words, cs_words, g_lo + b, tidx) + if cutlass.const_expr(WIDE): + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + if tidx == 0: + prims.mbarrier_arrive(epi_free) + t = t + cutlass.Int32(1) + if cutlass.const_expr(DYN_ALL): + if t % cutlass.Int32(DYN_RING) == cutlass.Int32(0): + dyn_ph = dyn_ph ^ 1 + v = _next_tile( + t, + False, + dyn_ph, + bidx, + total_tasks, + dyn_slot, + dyn_ready, + dyn_consumed, + dyn_ctr_ptr, + ) + in_rearm = v >= rearm0 + else: + while in_fc2: + lin2 = v - tiles_fc1 + m_tile = lin2 % M_TILES_FC2 + group = lin2 // M_TILES_FC2 + grp_cnt = g_cnt.load(idx=group) + gbase = group * N + ch = m_tile * MMA_M + fc2_ch_in_tile + while not cute.arch.mbarrier_try_wait(acc_full.data_ptr(), acc_full_phase): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc_full_phase = acc_full_phase ^ 1 + tmem_ld = cutlass.inttoptr( + (row_id_with_warp_offset << 16) | base_col_id, 6, cutlass.Float32 + ) + t2r_rmem = prims.tcgen05_ld("32x32b", tmem_ld, num=t2r_repx) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + prims.mbarrier_arrive(acc_empty) + for n in cutlass.range_constexpr(N): + if n < grp_cnt: + rw = g_rw.load(idx=gbase + n) + part_tensor.store( + cutlass.Float32(t2r_rmem[n]) * rw, idx=(gbase + n) * H + ch + ) + mtile_ctr = state_ptr + cutlass.Int64((ST_MTILE + m_tile) * 4) + group_ctr = state_ptr + cutlass.Int64((ST_GROUP + group) * 4) + if cutlass.const_expr(EPI_SYNC == "last"): + cute.arch.fence_acq_rel_gpu() + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + if tidx == 0: + # Both round trips in flight at once. + arrived = _atomic_fetch_add(mtile_ctr, cutlass.Int32(1)) + readers = _atomic_fetch_add(group_ctr, cutlass.Int32(1)) + epi_flag.store( + cutlass.select_( + arrived == num_groups - 1, cutlass.Int32(1), cutlass.Int32(0) + ), + idx=0, + ) + epi_flag.store( + cutlass.select_( + readers == M_TILES_FC2 - 1, cutlass.Int32(1), cutlass.Int32(0) + ), + idx=1, + ) + if cutlass.const_expr(EPI_RESET == "before"): + if arrived == num_groups - 1: + _store_release(mtile_ctr, cutlass.Int32(0)) + if readers == M_TILES_FC2 - 1: + _store_release(group_ctr, cutlass.Int32(0)) + else: + # This tile's partial rows are CTA-visible after the barrier; tid 0's release + # publishes them GPU-wide. + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + if tidx == 0: + if group != num_groups - 1: + _red_release_add(mtile_ctr, cutlass.Int32(1)) + if m_tile != M_TILES_FC2 - 1: + _red_release_add(group_ctr, cutlass.Int32(1)) + if group == num_groups - 1: + while _load_acquire(mtile_ctr) < num_groups - 1: + _backoff(32) + _store_release(mtile_ctr, cutlass.Int32(0)) + if m_tile == M_TILES_FC2 - 1: + while _load_acquire(group_ctr) < cutlass.Int32(M_TILES_FC2 - 1): + _backoff(32) + _store_release(group_ctr, cutlass.Int32(0)) + epi_flag.store( + cutlass.select_( + group == num_groups - 1, cutlass.Int32(1), cutlass.Int32(0) + ), + idx=0, + ) + epi_flag.store( + cutlass.select_( + m_tile == M_TILES_FC2 - 1, cutlass.Int32(1), cutlass.Int32(0) + ), + idx=1, + ) + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + if epi_flag.load(idx=1) != cutlass.Int32(0): + # Every FC2 tile of this group has read its intermediate: re-arm it. + _rearm_group(c_words, cs_words, group, tidx) + if cutlass.const_expr(EPI_SYNC == "last" and EPI_RESET == "after"): + if tidx == 0: + _store_release(group_ctr, cutlass.Int32(0)) + if epi_flag.load(idx=0) != cutlass.Int32(0): + if cutlass.const_expr(EPI_SYNC == "last" and EPI_RESET == "after"): + if tidx == 0: + _store_release(mtile_ctr, cutlass.Int32(0)) + cute.arch.fence_acq_rel_gpu() + if cutlass.const_expr(COMBINE == "loop"): + for tok in cutlass.range(num_tokens, unroll=1): + acc = cutlass.Float32(0.0) + for s in cutlass.range_constexpr(TOP_K): + slot = s_tok_slots.load(idx=tok * TOP_K + s) + if slot >= cutlass.Int32(0): + acc = acc + part_tensor.load( + idx=slot * H + ch, is_volatile=True + ) + _emit_row( + acc, + tok, + ch, + lane, + warp_idx, + m_tile, + y_tensor, + ar_mc, + ar_cur, + ar_rank, + ) + else: + # Every token's slots in top-k order; the loads of PAIR_CHUNK pairs (any + # tokens) are in flight together. + n_pairs = s_meta.load(idx=2) + acc0 = cutlass.Float32(0.0) + acc1 = cutlass.Float32(0.0) + acc2 = cutlass.Float32(0.0) + acc3 = cutlass.Float32(0.0) + acc4 = cutlass.Float32(0.0) + acc5 = cutlass.Float32(0.0) + acc6 = cutlass.Float32(0.0) + acc7 = cutlass.Float32(0.0) + for pb in cutlass.range(0, n_pairs, PAIR_CHUNK, unroll=1): + vals = [] + toks = [] + for q in cutlass.range_constexpr(PAIR_CHUNK): + pq = pb + q + valid = pq < n_pairs + packed = s_pairs.load( + idx=cutlass.select_(valid, pq, cutlass.Int32(0)) + ) + part_v = part_tensor.load( + idx=(packed & cutlass.Int32(0xFFFF)) * H + ch, is_volatile=True + ) + vals.append(cutlass.select_(valid, part_v, cutlass.Float32(0.0))) + toks.append( + cutlass.select_( + valid, packed >> cutlass.Int32(16), cutlass.Int32(-1) + ) + ) + for q in cutlass.range_constexpr(PAIR_CHUNK): + zero = cutlass.Float32(0.0) + acc0 = acc0 + cutlass.select_( + toks[q] == cutlass.Int32(0), vals[q], zero + ) + acc1 = acc1 + cutlass.select_( + toks[q] == cutlass.Int32(1), vals[q], zero + ) + acc2 = acc2 + cutlass.select_( + toks[q] == cutlass.Int32(2), vals[q], zero + ) + acc3 = acc3 + cutlass.select_( + toks[q] == cutlass.Int32(3), vals[q], zero + ) + acc4 = acc4 + cutlass.select_( + toks[q] == cutlass.Int32(4), vals[q], zero + ) + acc5 = acc5 + cutlass.select_( + toks[q] == cutlass.Int32(5), vals[q], zero + ) + acc6 = acc6 + cutlass.select_( + toks[q] == cutlass.Int32(6), vals[q], zero + ) + acc7 = acc7 + cutlass.select_( + toks[q] == cutlass.Int32(7), vals[q], zero + ) + accs = [acc0, acc1, acc2, acc3, acc4, acc5, acc6, acc7] + for tok in cutlass.range_constexpr(M_MAX): + if cutlass.Int32(tok) < num_tokens: + _emit_row( + accs[tok], cutlass.Int32(tok), ch, lane, warp_idx, m_tile, y_tensor, ar_mc, + ar_cur, ar_rank, + ) # fmt: skip + t = t + cutlass.Int32(1) + if cutlass.const_expr(DYN_ALL): + if t % cutlass.Int32(DYN_RING) == cutlass.Int32(0): + dyn_ph = dyn_ph ^ 1 + v = _next_tile( + t, + False, + dyn_ph, + bidx, + total_tasks, + dyn_slot, + dyn_ready, + dyn_consumed, + dyn_ctr_ptr, + ) + in_fc2 = (v >= cutlass.Int32(0)) & (v < total_tiles) + if cutlass.const_expr(FUSED_AR and not AR_PUSH_ONLY): + # No tile is left for this CTA and the queue is empty, so every push of this rank + # is issued or in a running CTA's epilogue. Reduction tasks, one at a time: an + # idle CTA takes the next only when it is free again. + if cutlass.const_expr(LAT_SLAB): + # E5: re-arm the next buffer (and buffer 0 at the step's last call); this grid's wait returned. + slab_words = M_MAX * (H // 2) + ones = cutlass.Int32(LAT_SLAB_EMPTY) + nxt = (lat_buf + cutlass.Int32(1)) % cutlass.Int32(LAT_SLAB_BUFS) + for w in range(bidx * EPI_THREADS + tidx, slab_words, NUM_CTAS * EPI_THREADS): + lat_slab.store(ones, idx=nxt * cutlass.Int32(slab_words) + w) + if lat_rearm0 != cutlass.Int32(0): + lat_slab.store(ones, idx=w) + task = _claim_ar_task(state_ptr, epi_flag, tidx, ar_flags, ar_cur) + while task < cutlass.Int32(AR_TASKS): + _ar_reduce_tile( + ar_uc, y_words, ar_cur, task, tidx, num_tokens, ar_rank, lat_slab, lat_buf + ) + task = _claim_ar_task(state_ptr, epi_flag, tidx, ar_flags, ar_cur) + # ------------------------------------------------ teardown + # Every role is done: MMAs committed and waited on, TMEM loads waited on. + if cutlass.const_expr(HEAD_FLAGS and USE_PDL): + cute.arch.griddepcontrol_wait() + prims.barrier_cta_sync(0) + if warp_idx == consumer_warp_id: + tmem_raw_addr = tmem_ptr_i32.load() + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + prims.tcgen05_dealloc( + cutlass.inttoptr(tmem_raw_addr, 6, cutlass.Float32), num_tmem_alloc_cols + ) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_route_quant_ag.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_route_quant_ag.py new file mode 100644 index 000000000000..9688b95e238c --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_route_quant_ag.py @@ -0,0 +1,360 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Kimi K3 MoE head all-gather + routing + MXFP8 input quantization for decode, in CuTe DSL. + +The MoE head is row-sharded over the TP group of W ranks: rank r's GEMV gives fp32 +``[M, H/W + E/W]``, the latent down's columns ``[r*H/W, (r+1)*H/W)`` followed by the router +logits of experts ``[r*E/W, (r+1)*E/W)``. This kernel gathers every rank's slice and returns +what ``trtllm::k3_route_quant`` returns for the gathered logits and latent (top-16 ids, bf16 +weights, the MXFP8 latent, its UE8M0 scales), bit for bit what ``mnnvl_allgather_split`` +followed by ``k3_route_quant`` give: the pushed values are rounded and sanitized as that +all-gather does (latent to bf16, -0.0 to +0.0), and the routing and quantization are +``k3_route_quant``'s device code. + +Grid 2M CTAs of 448 threads: +- CTA t < M pushes this rank's 14 logit vectors of token t, through the multicast mapping, + into every rank's Lamport buffer, polls token t's logit vectors of all W ranks, and routes + token t (keys in shared memory, warp 0 selects). +- CTA M + t pushes this rank's 28 latent vectors of token t (8 bf16 each), polls token t's + latent vectors of all ranks (one per thread: the thread's 8-element quantization vector) + and quantizes them. +Every reader writes the empty word (0x80000000) back over what it read. Two buffers alternate +between calls: flags[0] holds the buffer of the next call, flipped by the grid's last CTA to +count itself in flags[1]. + +With ``publish``, each CTA also releases a per-token ready word once its outputs are written +(ready[t] = epoch + 1 for the ids and weights, ready[8 + t] for the MXFP8 row, the epoch being +flags[2]), so that k3_moe (built with head_flags, which advances flags[2]) can acquire them +instead of waiting for this grid. + +Buffer words: ``[buffer][token < 8][rank < W][latent 4*H/(8W) words | logits 4*E/(4W) words]``. +""" + +from __future__ import annotations + +import cuda.bindings.driver as cuda_driver +import cutlass +import cutlass.cute as cute +from cutlass import dsl_user_op +from cutlass.experimental import primitives as prims + +from tensorrt_llm._torch.cute_dsl_kernels.k3_route_quant import k3_route_quant_kernel as _rq + +NUM_EXPERTS = 896 +TOP_K = 16 +HIDDEN_SIZE = 3584 +SF_VEC_SIZE = 32 +THREADS = HIDDEN_SIZE // 8 # one 8-element vector per thread in a quantization CTA +M_MAX = 8 +BUFFERS = 2 +EMPTY_WORD = -(2**31) # 0x80000000: fp32 -0.0, never a pushed word + + +def slot_words(world: int) -> int: + """Int32 words of one (token, rank) slot: the latent's 8-element vectors, then the logits'.""" + return (HIDDEN_SIZE // world // 8 + NUM_EXPERTS // world // 4) * 4 + + +def buffer_words(world: int) -> int: + return BUFFERS * M_MAX * world * slot_words(world) + + +FRONT_SPLIT = 8 # k3_moe_front's cluster size: its router rows reach the route CTAs as this many split-K partials + + +def partial_words() -> int: + """Int32 words after the buffers: k3_moe_front's router logits as split-K partials, + [buffer][token][rank][cluster rank][expert of the rank] (the same size for every world).""" + return BUFFERS * M_MAX * FRONT_SPLIT * NUM_EXPERTS + + +def workspace_words(world: int) -> int: + """The head workspace: the all-gather's buffers, then the front's router partials.""" + return buffer_words(world) + partial_words() + + +@dsl_user_op +def _atomic_add_acq_rel(addr_i64, val, *, loc=None, ip=None): + """atom.acq_rel.gpu.global.add.u32, returning the old value.""" + from cutlass._mlir.dialects import llvm as _llvm + from cutlass._mlir.extras import types as _T + + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [addr_i64.ir_value(loc=loc, ip=ip), val.ir_value(loc=loc, ip=ip)], + "atom.acq_rel.gpu.global.add.u32 $0, [$1], $2;", "=r,l,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _store_release(addr_i64, val, *, loc=None, ip=None): + """st.release.gpu.global.u32.""" + from cutlass._mlir.dialects import llvm as _llvm + + _llvm.inline_asm( + None, [addr_i64.ir_value(loc=loc, ip=ip), val.ir_value(loc=loc, ip=ip)], + "st.release.gpu.global.u32 [$0], $1;", "l,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _pack_bf16x2(hi, lo, *, loc=None, ip=None): + """(bf16(hi) << 16) | bf16(lo), round to nearest even (``__floats2bfloat162_rn``).""" + from cutlass._mlir.dialects import llvm as _llvm + from cutlass._mlir.extras import types as _T + + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [hi.ir_value(loc=loc, ip=ip), lo.ir_value(loc=loc, ip=ip)], + "cvt.rn.bf16x2.f32 $0, $1, $2;", "=r,f,f", has_side_effects=False, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +def _sanitize_bf16x2(word): + """-0.0 halves become +0.0 (the all-gather's sanitizeBf16Pair), so no word is 0x80000000.""" + word = cutlass.Int32( + cutlass.select_( + (word & cutlass.Int32(0xFFFF)) == cutlass.Int32(0x8000), + word & cutlass.Int32(-65536), + word, + ) + ) + return cutlass.Int32( + cutlass.select_( + (word & cutlass.Int32(-65536)) == cutlass.Int32(-2147483648), + word & cutlass.Int32(0xFFFF), + word, + ) + ) + + +def _sanitize_f32(word): + return cutlass.Int32(cutlass.select_(word == cutlass.Int32(EMPTY_WORD), cutlass.Int32(0), word)) + + +@cute.jit +def _poll4(arr, idx): + """Spin (no back-off) until none of the 4 words at idx is empty; returns them.""" + w0 = cutlass.Int32(EMPTY_WORD) + w1 = cutlass.Int32(EMPTY_WORD) + w2 = cutlass.Int32(EMPTY_WORD) + w3 = cutlass.Int32(EMPTY_WORD) + pending = cutlass.Boolean(True) + while pending: + v = arr.load(idx=idx, vector_size=4, alignment=16, is_volatile=True) + w0 = cutlass.Int32(v[0]) + w1 = cutlass.Int32(v[1]) + w2 = cutlass.Int32(v[2]) + w3 = cutlass.Int32(v[3]) + empty = cutlass.Int32(EMPTY_WORD) + pending = (w0 == empty) | (w1 == empty) | (w2 == empty) | (w3 == empty) + return w0, w1, w2, w3 + + +@cute.kernel +def k3_route_quant_ag_kernel( + head_words: cutlass.Array, # int32 view of this rank's fp32 head [M, WL + WE] + bias: cutlass.Array, # fp32 [896] + buf_uc: cutlass.Array, # int32 words of this rank's Lamport buffers + buf_mc: cutlass.Array, # int32 words of their multicast mapping + flags: cutlass.Array, # int32 [0] buffer of this call, [1] CTAs of this call counted so far + topk_ids: cutlass.Array, # int32 [M * 16] + topk_weight_bits: cutlass.Array, # int16 view of bf16 [M * 16] + quant_words: cutlass.Array, # int32 view of e4m3 [M, 3584]: [M * 896] + scales: cutlass.Array, # uint8 [M * 112] + ready: cutlass.Array, # int32 [16] per-token ready words (publish) + num_tokens: cutlass.Int32, + rank: cutlass.Int32, + routed_scaling_factor: cutlass.Float64, + world: cutlass.Constexpr[int], + early_trigger: cutlass.Constexpr[bool], + publish: cutlass.Constexpr[bool], +): + WL = HIDDEN_SIZE // world # latent columns per rank + WE = NUM_EXPERTS // world # logits per rank + LV = WL // 8 # latent vectors per slot + EV = WE // 4 # logit vectors per slot + SLOT = (LV + EV) * 4 + HEAD_ROW = WL + WE # fp32 words per head row + BUF = M_MAX * world * SLOT + tidx, _, _ = cute.arch.thread_idx() + bidx, _, _ = cute.arch.block_idx() + s_key = cutlass.Array(cutlass.Int32, NUM_EXPERTS, space=cutlass.AddressSpace.smem, alignment=16) + s_sigmoid = cutlass.Array( + cutlass.Float32, NUM_EXPERTS, space=cutlass.AddressSpace.smem, alignment=16 + ) + s_buf = cutlass.Array(cutlass.Int32, 4, space=cutlass.AddressSpace.smem, alignment=16) + routes = bidx < num_tokens + tok = cutlass.select_(routes, bidx, bidx - num_tokens) + + # The bias is a weight: this thread's 4 experts (the gathered logits are rank-major, so + # route thread i polls experts 4i .. 4i+3), read before the grid dependency. + polls_logits = tidx < cutlass.Int32(NUM_EXPERTS // 4) + e0 = cutlass.select_(polls_logits, tidx * cutlass.Int32(4), cutlass.Int32(0)) + bias4 = bias.load(idx=e0, vector_size=4, alignment=16) + + prims.griddepcontrol(prims.GridDepAction.WAIT) + if tidx == 0: + s_buf.store(flags.load(idx=0, is_volatile=True), idx=0) + if cutlass.const_expr(publish): + s_buf.store(flags.load(idx=2, is_volatile=True), idx=1) + cute.arch.barrier() + b = s_buf.load(idx=0) + epoch = s_buf.load(idx=1) + if cutlass.const_expr(early_trigger): + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + # Push this rank's slot of token tok: logits (route CTAs) or latent (quantization CTAs). + slot_base = b * BUF + (tok * world + rank) * SLOT + head_row = tok * HEAD_ROW + if routes: + if tidx < EV: + w = head_words.load(idx=head_row + WL + tidx * 4, vector_size=4, alignment=16) + buf_mc.store( + ( + _sanitize_f32(cutlass.Int32(w[0])), + _sanitize_f32(cutlass.Int32(w[1])), + _sanitize_f32(cutlass.Int32(w[2])), + _sanitize_f32(cutlass.Int32(w[3])), + ), + idx=slot_base + LV * 4 + tidx * 4, + alignment=16, + ) + # Drain the posted multicast store now: without a fence it can sit in the SM's write path + # until this thread's next release, and the polls below have none (a cluster-scope fence + # does not wait for the remote ranks' acknowledgement). + prims.fence_acq_rel(prims.MemScope.CLUSTER) + else: + if tidx < LV: + lo = head_words.load(idx=head_row + tidx * 8, vector_size=4, alignment=16) + hi = head_words.load(idx=head_row + tidx * 8 + 4, vector_size=4, alignment=16) + f = [cutlass.Int32(lo[q]).bitcast(cutlass.Float32) for q in range(4)] + f += [cutlass.Int32(hi[q]).bitcast(cutlass.Float32) for q in range(4)] + buf_mc.store( + ( + _sanitize_bf16x2(_pack_bf16x2(f[1], f[0])), + _sanitize_bf16x2(_pack_bf16x2(f[3], f[2])), + _sanitize_bf16x2(_pack_bf16x2(f[5], f[4])), + _sanitize_bf16x2(_pack_bf16x2(f[7], f[6])), + ), + idx=slot_base + tidx * 4, + alignment=16, + ) + prims.fence_acq_rel(prims.MemScope.CLUSTER) + + # Every CTA has read the buffer index before counting itself (a thread of the last warp, which + # never pushes): the grid's last one hands the other buffer to the next call, which reads it + # after its own grid-dependency wait. + if tidx == cutlass.Int32(THREADS - 32): + arrived = _atomic_add_acq_rel(flags.data_ptr(1).toint(), cutlass.Int32(1)) + if arrived == num_tokens * cutlass.Int32(2) - cutlass.Int32(1): + flags.store(cutlass.Int32(0), idx=1) + flags.store(b ^ cutlass.Int32(1), idx=0) + + tok_base = b * BUF + tok * world * SLOT + empty = cutlass.Int32(EMPTY_WORD) + if routes: + # Poll every rank's logit vectors of this token (one per thread), keys and sigmoids into + # shared memory, empty what was read. + if polls_logits: + addr = tok_base + (tidx // EV) * SLOT + LV * 4 + (tidx % EV) * 4 + w0, w1, w2, w3 = _poll4(buf_uc, addr) + buf_uc.store((empty, empty, empty, empty), idx=addr, alignment=16) + words = [w0, w1, w2, w3] + for q in cutlass.range_constexpr(4): + sig = _rq.sigmoid_accurate(words[q].bitcast(cutlass.Float32)) + s_sigmoid.store(sig, idx=e0 + q) + s_key.store(_rq.selection_key(sig + cutlass.Float32(bias4[q])), idx=e0 + q) + cute.arch.barrier() + if tidx < cutlass.Int32(32): + expert, weight_bits = _rq.top16_warp(s_key, s_sigmoid, tidx, routed_scaling_factor) + if tidx < cutlass.Int32(TOP_K): + out = tok * cutlass.Int32(TOP_K) + tidx + topk_ids.store(expert, idx=out) + topk_weight_bits.store(weight_bits, idx=out) + if cutlass.const_expr(publish): + cute.arch.fence_acq_rel_gpu() + if cutlass.const_expr(publish): + cute.arch.sync_warp() + if tidx == cutlass.Int32(0): + _store_release(ready.data_ptr(tok).toint(), epoch + cutlass.Int32(1)) + else: + # Poll this thread's latent vector (rank tidx // LV, columns 8 * tidx ..), quantize it with + # its 4-lane scale group, empty what was read. + addr = tok_base + (tidx // LV) * SLOT + (tidx % LV) * 4 + w0, w1, w2, w3 = _poll4(buf_uc, addr) + buf_uc.store((empty, empty, empty, empty), idx=addr, alignment=16) + q_lo, q_hi, sf_byte = _rq.mxfp8_quant_vec8([w0, w1, w2, w3]) + quant_words.store( + (q_lo, q_hi), + idx=tok * cutlass.Int32(HIDDEN_SIZE // 4) + tidx * cutlass.Int32(2), + alignment=8, + ) + if tidx % cutlass.Int32(SF_VEC_SIZE // 8) == cutlass.Int32(0): + scales.store( + cutlass.Uint8(sf_byte), + idx=tok * cutlass.Int32(HIDDEN_SIZE // SF_VEC_SIZE) + + tidx // cutlass.Int32(SF_VEC_SIZE // 8), + ) + if cutlass.const_expr(publish): + # The row is read by k3_moe's TMA gathers (async proxy). + prims.fence_proxy("async_global") + cute.arch.fence_acq_rel_gpu() + cute.arch.barrier() + if tidx == cutlass.Int32(0): + _store_release( + ready.data_ptr(tok + cutlass.Int32(M_MAX)).toint(), epoch + cutlass.Int32(1) + ) + + if cutlass.const_expr(not early_trigger): + cute.arch.fence_acq_rel_gpu() + cute.arch.barrier() + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + + +@cute.jit +def k3_route_quant_ag( + head_words: cute.Tensor, + bias: cute.Tensor, + buf_uc: cute.Tensor, + buf_mc: cute.Tensor, + flags: cute.Tensor, + topk_ids: cute.Tensor, + topk_weight_bits: cute.Tensor, + quant_words: cute.Tensor, + scales: cute.Tensor, + ready: cute.Tensor, + num_tokens: cutlass.Int32, + rank: cutlass.Int32, + routed_scaling_factor: cutlass.Float64, + world: cutlass.Constexpr[int], + early_trigger: cutlass.Constexpr[bool], + publish: cutlass.Constexpr[bool], + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + k3_route_quant_ag_kernel( + head_words, bias, buf_uc, buf_mc, flags, topk_ids, topk_weight_bits, quant_words, scales, ready, + num_tokens, rank, routed_scaling_factor, world, early_trigger, publish, + ).launch( + grid=[num_tokens * 2, 1, 1], + block=[THREADS, 1, 1], + stream=stream, + use_pdl=use_pdl, + ) # fmt: skip diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py new file mode 100644 index 000000000000..220a7b8a0d8a --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py @@ -0,0 +1,158 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``trtllm::k3_latent_reduce``: the Kimi K3 latent all-reduce at decode size as the consumer of the push-only k3_moe +(``k3_latent_reduce.py``). + +``LatentExchange`` owns a TP group's buffers: the push-only ops (``trtllm::k3_fused_moe_push`` / +``trtllm::k3_fused_moe_front_push``) store every rank's routed partial into them, and ``trtllm::k3_latent_reduce`` +returns the sum, bit-identical to ``MNNVLAllReduce``'s one-shot of the partials. Each push must be followed by exactly +one reduce of the same token count on the same exchange before the next push, on every rank in the same order. A push +reads the half from the call count after its grid-dependency wait and triggers its dependents only after that wait, +and the reduce reads the count before its own wait, so every kernel from a reduce to the next push must end only after +its predecessor has ended (it calls ``griddepcontrol.wait``, or launches without PDL). The kernel compiles on the first +call for each configuration, which must happen outside CUDA-graph capture. +""" + +from __future__ import annotations + +import os +import threading +from typing import Dict, Optional + +import torch + +HIDDEN_SIZE = 3584 +MAX_TOKENS = 8 +EMPTY_WORD = -(2**31) + +_lock = threading.Lock() +_compiled: Dict[tuple, object] = {} + + +def _kernel(): + from . import k3_latent_reduce as kernel + + return kernel + + +def default_ctas(world: int) -> int: + """CTAs per token row: 4 (112 threads) up to 8 ranks; 14 (32 threads) at 16, where a poll pass loads 16 slots.""" + return 4 if world <= 8 else 14 + + +class LatentExchange: + """The latent all-reduce buffers of ``mapping``'s TP group: int32 ``[2][8][world][1792]`` per rank behind one + multicast mapping, every word ``0x80000000``, and ``flags`` (int32 ``[4]``: the consumer's call count and its + CTA arrivals). Collective: every rank of the group constructs it at the same point (outside graph capture). + Separate from the MNNVL all-reduce workspace.""" + + def __init__(self, mapping): + from tensorrt_llm._torch.distributed.ops import ( + _get_mnnvl_workspace_comm, + _make_mnnvl_mcast_buffer, + _mnnvl_workspace_all_succeeded, + ) + + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError("Kimi K3 latent exchange buffers must be built outside CUDA-graph capture") + self.world = mapping.tp_size + self.rank = mapping.tp_rank + words = _kernel().buffer_words(self.world) + comm = _get_mnnvl_workspace_comm(mapping) + use_fabric_handle = ( + os.environ.get("TRTLLM_FORCE_MNNVL_AR", "0") == "1" or mapping.is_multi_node() + ) + error: Optional[Exception] = None + try: + self.handle = _make_mnnvl_mcast_buffer(comm, words * 4, mapping, use_fabric_handle) + self.uc = self.handle.get_uc_buffer(self.rank, (words,), torch.int32, 0) + self.mc = self.handle.get_mc_buffer((words,), torch.int32, 0) + with torch.inference_mode(): + self.uc.fill_(EMPTY_WORD) + self.flags = torch.zeros(4, dtype=torch.int32, device=self.uc.device) + torch.cuda.synchronize() + except Exception as exc: # noqa: BLE001 -- reported to every rank below, then re-raised + error = exc + # Also the barrier that keeps any rank from pushing into a peer's buffer before the peer has filled it. + if not _mnnvl_workspace_all_succeeded(comm, error is None): + raise RuntimeError( + "Kimi K3 latent exchange buffers failed on at least one rank" + ) from error + self.comm = comm + + def push_args(self): + """``(ar_uc, ar_mc, ar_flags, ar_rank)`` of the push-only ops.""" + return self.uc, self.mc, self.flags, self.rank + + +def supports(world: int, num_tokens: int) -> bool: + return world in (4, 8, 16) and 0 < num_tokens <= MAX_TOKENS + + +@torch.library.custom_op("trtllm::k3_latent_reduce", mutates_args=("lat_uc", "lat_flags")) +def k3_latent_reduce( + lat_uc: torch.Tensor, lat_flags: torch.Tensor, num_tokens: int, ctas_per_token: int = 0 +) -> torch.Tensor: + """The latent rows ``[num_tokens, 3584]`` bf16: the sum over the ranks of the routed partials the push-only k3_moe + stored into ``lat_uc`` (``LatentExchange.uc``) since the last call, in the MNNVL one-shot's order. Empties the + words it read and advances ``lat_flags``' call count. ``ctas_per_token``: 4, 14 or 28 (0: ``default_ctas``).""" + import cuda.bindings.driver as cuda_driver + import cutlass.cute as cute + from cutlass.cute.runtime import from_dlpack + + kernel = _kernel() + world = lat_uc.numel() // kernel.buffer_words(1) + if ( + lat_uc.dtype != torch.int32 + or lat_uc.numel() != kernel.buffer_words(world) + or lat_flags.dtype != torch.int32 + or lat_flags.numel() < 4 + or not supports(world, num_tokens) + ): + raise ValueError( + f"k3_latent_reduce: int32 buffer of [2][8][world][1792] words with world 4 / 8 / 16, int32 flags[4], " + f"1..{MAX_TOKENS} tokens; got {lat_uc.numel()} words, flags {tuple(lat_flags.shape)}, {num_tokens} tokens" + ) + ctas = ctas_per_token or default_ctas(world) + if ctas not in (4, 14, 28): + raise ValueError(f"k3_latent_reduce: ctas_per_token must be 4, 14 or 28, got {ctas}") + out = torch.empty(num_tokens, HIDDEN_SIZE, dtype=torch.bfloat16, device=lat_uc.device) + + def arg(t): + return from_dlpack(t.detach(), assumed_align=16).mark_layout_dynamic(leading_dim=0) + + args = (arg(lat_uc.view(-1)), arg(lat_flags.view(-1)), arg(out.view(-1).view(torch.int32))) + stream = cuda_driver.CUstream(torch.cuda.current_stream(lat_uc.device).cuda_stream) + use_pdl = os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" + key = (world, ctas, use_pdl) + fn = _compiled.get(key) + if fn is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "trtllm::k3_latent_reduce must run once outside CUDA-graph capture first" + ) + with _lock: + fn = _compiled.get(key) + if fn is None: + fn = _compiled[key] = cute.compile( + kernel.k3_latent_reduce, *args, num_tokens, world, ctas, use_pdl, stream + ) + fn(*args, num_tokens, stream) + return out + + +@k3_latent_reduce.register_fake +def _(lat_uc, lat_flags, num_tokens, ctas_per_token=0): + return lat_uc.new_empty((num_tokens, HIDDEN_SIZE), dtype=torch.bfloat16) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py new file mode 100644 index 000000000000..2f97a353a35d --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py @@ -0,0 +1,973 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``trtllm::k3_fused_moe``: Kimi K3 routed experts for decode (M <= 8 tokens). + +Two kernels on the current stream, no host synchronization: + +1. ``trtllm::k3_route_quant`` -- the routing and MXFP8 input quantization the TRTLLM-Gen + path uses under separated routing (sigmoid, top-16 of sigmoid + bias, unbiased scores + renormalized times the routed scaling factor; ties to the lower id), the CuTe DSL form of + ``trtllm::kimi_k3_noaux_tc_mxfp8_quant`` with the same outputs bit for bit. +2. ``k3_moe`` -- one persistent CuTe DSL kernel, launched as a programmatic dependent of + the first: this rank's (expert, token) groups in its prologue, then FC1 + SiTU + FC2 + with the routing-weighted, deterministic combine. + +With the ``fold`` option there is one kernel: ``k3_moe`` computes the routing and the quantization in +its prologue (``trtllm::k3_route_quant``'s device code, so the same ids, weights and MXFP8 bits) +from the router logits and the latent, and its PDL predecessor is whatever produced them. + +The result is this rank's routed partial ``[M, 3584]`` bf16, the tensor the TRTLLM-Gen +W4A8_MXFP4_MXFP8 op returns, so the routed-latent all-reduce and the latent-up tail are +unchanged. Weights are the TRTLLM-Gen buffers, read in place. + +``trtllm::k3_fused_moe_ar`` also performs the all-reduce of that partial over the TP group +inside ``k3_moe`` (buffers from ``ar_workspace``) and returns the reduced latent. + +The CuTe DSL kernel is compiled on the first call for each intermediate size (for the +model: the warmup that precedes CUDA-graph capture) and the scratch buffers are allocated +then too, so captured calls only launch. Each layer (keyed by its weight buffer) owns a +few counters that the kernel returns to zero; the intermediate slab, shared by all layers, +is left armed by every call. + +Steps of up to 64 tokens use :class:`K3MoeWideState` (the m_max 64 build of ``k3_moe``, +launched after ``trtllm::k3_route_quant``), whose scratch and per-layer counters the caller +owns. +""" + +from __future__ import annotations + +import importlib.util +import os +import sys +import threading +from typing import Dict, Optional, Tuple + +import torch + +HIDDEN_SIZE = 3584 +NUM_EXPERTS = 896 +TOP_K = 16 +MAX_TOKENS = 8 +_TOKEN_SLOTS = 8 +_SF_VEC = 32 + +_KERNEL_PATH = os.path.join(os.path.dirname(os.path.abspath(__file__)), "k3_moe_kernel.py") +_lock = threading.Lock() +_modules: Dict[tuple, object] = {} +_states: Dict[Tuple[int, int, int], "_K3FusedMoE"] = {} + + +def _kernel_module(config: dict): + """One kernel module per configuration: shapes are trace-time constants.""" + key = tuple(sorted(config.items())) + mod = _modules.get(key) + if mod is None: + tag = "_".join(f"{k}{v}" for k, v in key) + name = f"{__name__}_kernel_{tag}" + spec = importlib.util.spec_from_file_location(name, _KERNEL_PATH) + mod = importlib.util.module_from_spec(spec) + mod.K3_CONFIG = dict(config) + sys.modules[name] = mod + spec.loader.exec_module(mod) + _modules[key] = mod + return mod + + +def _view(t: torch.Tensor, align: int, leading_dim: int, element_type=None): + from cutlass.cute.runtime import from_dlpack + + v = from_dlpack(t, assumed_align=align).mark_layout_dynamic(leading_dim=leading_dim) + if element_type is not None: + v.element_type = element_type + return v + + +def is_supported(w3_w1_weight: torch.Tensor, w3_w1_weight_scale: torch.Tensor, w2_weight: torch.Tensor, + w2_weight_scale: torch.Tensor, local_num_experts: int) -> Tuple[bool, str]: # fmt: skip + """Whether these TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers fit the fused kernels. + + Only metadata is read, so the buffers may still be on the meta device.""" + if not torch.cuda.is_available() or torch.cuda.get_device_capability() != (10, 0): + return False, "needs sm_100" + try: + import cutlass # noqa: F401 + except ImportError: + return False, "CuTe DSL (nvidia-cutlass-dsl) is not installed" + e, two_i, k_half = w3_w1_weight.shape + i_tp = two_i // 2 + expected = { + "w3_w1_weight": (w3_w1_weight, (local_num_experts, two_i, HIDDEN_SIZE // 2)), + "w3_w1_weight_scale": ( + w3_w1_weight_scale, + (local_num_experts, two_i, HIDDEN_SIZE // _SF_VEC), + ), + "w2_weight": (w2_weight, (local_num_experts, HIDDEN_SIZE, i_tp // 2)), + "w2_weight_scale": (w2_weight_scale, (local_num_experts, HIDDEN_SIZE, i_tp // _SF_VEC)), + } + for name, (t, shape) in expected.items(): + if t.dtype != torch.uint8 or tuple(t.shape) != shape or not t.is_contiguous(): + return False, f"{name} is {t.dtype} {tuple(t.shape)}, expected contiguous uint8 {shape}" + if i_tp % 128 != 0: + return False, f"intermediate size {i_tp} is not a multiple of 128" + if local_num_experts > NUM_EXPERTS: + return False, f"{local_num_experts} local experts" + return True, "" + + +class _K3FusedMoE: + """Compiled kernel and scratch for one (device, intermediate size, local experts). + + ``config`` overrides kernel options (tests and A/B runs: ``pdl``, + ``num_ctas``, ...); anything it leaves out takes the kernel's default.""" + + def __init__( + self, device: torch.device, i_tp: int, num_local: int, config: Optional[dict] = None + ): + # One persistent CTA per SM (config "num_ctas" caps it, e.g. for a grid-size A/B). + num_ctas = torch.cuda.get_device_properties(device).multi_processor_count + cfg = {"i_tp": i_tp, "num_ctas": num_ctas, "num_local": num_local} + cfg.update(config or {}) + self.mod = mod = _kernel_module(cfg) + self.device = device + g_cap = mod.G_CAP + kw = dict(device=device) + # Lamport slab, armed: FP8 -0.0 values; E8M0 NaN in bytes 0..3 of each 16-byte scale + # group. Every call leaves the groups it used armed again. + self.c = torch.full((g_cap, _TOKEN_SLOTS, i_tp), -128, dtype=torch.int8, **kw) + self.cs = torch.zeros(g_cap, _TOKEN_SLOTS, mod.SF_STRIDE0, dtype=torch.int8, **kw) + self.cs.view(g_cap, _TOKEN_SLOTS, mod.K2_TILES, mod.SFB_GROUP_BYTES)[..., :4] = -1 + # FC2 partial rows. A call writes the rows of its M tokens; the combine also loads the rows past M (their sums + # are dropped), so the buffer starts zeroed and those loads never read unwritten memory. + self.part = torch.zeros(g_cap * _TOKEN_SLOTS, HIDDEN_SIZE, dtype=torch.float32, **kw) + self.c_t = _view(self.c, 16, 2) + self.cs_t = _view(self.cs, 4, 2) + self.c_words_t = _view(self.c.view(-1).view(torch.int32), 16, 0) + self.cs_words_t = _view(self.cs.view(-1).view(torch.int32), 16, 0) + self.b2_t = _view(self.c.permute(2, 1, 0), 16, 0, mod.b_dtype) + sfb2 = ( + self.cs.view(torch.uint8) + .reshape(g_cap, _TOKEN_SLOTS, mod.K2_TILES, mod.SFB_GROUP_BYTES) + .permute(3, 2, 1, 0) + ) + self.sfb2_t = _view(sfb2, 16, 0, mod.sf_dtype) + self.part_t = _view(self.part, 16, 1) + # Stand-ins for the all-reduce buffers of builds without the fused all-reduce, and for the + # inputs a build does not read (fold: the top-k; otherwise the logits and the latent). + self.no_ar = torch.zeros(4, dtype=torch.int32, **kw) + self.no_ar_t = _view(self.no_ar, 16, 0) + self.no_w = torch.zeros(8, dtype=torch.bfloat16, **kw) + self.no_w_t = _view(self.no_w, 16, 0) + self.ar_views: Dict[Tuple[int, int, int], tuple] = {} + self.fold = mod.FOLD + if self.fold: + # Each CTA's MXFP8 latent rows [8 * cta, 8 * cta + M) and their linear scales. + rows = mod.NUM_CTAS * mod.M_MAX + self.xq = torch.empty(rows, HIDDEN_SIZE, dtype=torch.uint8, **kw) + self.xsf = torch.empty(rows, HIDDEN_SIZE // _SF_VEC, dtype=torch.uint8, **kw) + self.b1_t = _view(self.xq.permute(1, 0), 16, 0, mod.b_dtype) + self.sfb1_t = _view(self.xsf, 16, 1, mod.sf_dtype) + self.xq_words_t = _view(self.xq.view(-1).view(torch.int32), 16, 0) + self.xsf_t = _view(self.xsf.view(-1), 16, 0) + # The route+quant kernel triggers k3_moe's launch right after its own grid dependency: + # k3_moe waits for the whole route+quant grid before reading its outputs. + self.route_kwargs = {"early_trigger": True} if mod.USE_PDL else {} + from ..k3_route_quant import op as _k3_route_quant_op # noqa: F401 + + self.route_quant = torch.ops.trtllm.k3_route_quant + # Per layer (keyed by its w3_w1 buffer): weight views and the kernel's counters. + self.layers: Dict[Tuple[int, int, int, int], tuple] = {} + self.moe = None + + def _layer(self, w31, w31s, w2, w2s): + key = (w31.data_ptr(), w31s.data_ptr(), w2.data_ptr(), w2s.data_ptr()) + layer = self.layers.get(key) + if layer is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "trtllm::k3_fused_moe must run once per layer outside CUDA-graph capture " + "first (it allocates the layer's counters on the first call)." + ) + mod = self.mod + e, two_i, _ = w31.shape + i_tp = two_i // 2 + sfa1 = w31s.view(e, two_i // 128, HIDDEN_SIZE // 128, 512).permute(3, 2, 1, 0) + sfa2 = w2s.view(e, HIDDEN_SIZE // 128, i_tp // 128, 512).permute(3, 2, 1, 0) + state = torch.zeros(mod.NUM_STATE, dtype=torch.int32, device=self.device) + layer = ( + _view(w31.view(torch.int8).permute(2, 1, 0), 16, 0), + _view(sfa1, 16, 0, mod.sf_dtype), + _view(w2.view(torch.int8).permute(2, 1, 0), 16, 0), + _view(sfa2, 16, 0, mod.sf_dtype), + state, + _view(state, 4, 0), + ) + self.layers[key] = layer + return layer + + def _flat_view(self, t: torch.Tensor): + key = ("flat", t.data_ptr(), t.numel()) + v = self.ar_views.get(key) + if v is None: + v = self.ar_views[key] = _view(t.view(-1), 16, 0) + return v + + def _ar_views(self, ar): + if ar is None: + return self.no_ar_t, self.no_ar_t, self.no_ar_t, 0 + uc, mc, flags, rank = ar + key = (uc.data_ptr(), mc.data_ptr(), flags.data_ptr()) + views = self.ar_views.get(key) + if views is None: + views = self.ar_views[key] = (_view(uc, 16, 0), _view(mc, 16, 0), _view(flags, 4, 0)) + return (*views, rank) + + def __call__( + self, + x, + router_logits, + bias, + w31, + w31s, + w2, + w2s, + local_offset: int, + num_local: int, + scale: float, + ar: Optional[tuple] = None, + head_ag: Optional[tuple] = None, + head_ready: Optional[torch.Tensor] = None, + front: Optional[tuple] = None, + lat_slab: Optional[tuple] = None, + ): + """``ar``: (uc words, mc words, flags, rank) of the fused all-reduce, for a build + with ``ar_world`` set; the result is then the all-reduced sum. With ``ar_push_only`` the + rows only go to every rank's buffer (half ``flags[0] & 1``) and the result has no rows. + + ``head_ag``: (uc words, mc words, flags, rank) of the head all-gather's buffers; ``x`` + is then this rank's slice of the sharded MoE head (fp32 ``[M, (H + E) / world]``), + ``router_logits`` is None, and ``trtllm::k3_route_quant_ag`` gathers, routes and + quantizes before ``k3_moe``. ``head_ready``: route_quant_ag's ready words, for a build with + ``head_flags`` (k3_moe then acquires them instead of waiting for that grid). + + ``front``: (front weight, shared columns, gate cap, linear cap, world) of + ``trtllm::k3_moe_front``, with ``head_ag``; ``x`` is then the MoE input (bf16 ``[M, 7168]``, + the same on every rank), the front stands in for the head GEMV and route_quant_ag, and + the call returns ``(y, shared activation)``. + + ``lat_slab``: (slab, buffer, re-arm buffer 0) for a build with ``lat_slab`` (and the fused + all-reduce): the reduced latent rows also go into ``slab`` (int32 ``[3, 8, 1792]``, all-ones + empty) in ``buffer`` (the call's ordinal mod 3); the call re-arms the next buffer, and buffer + 0 too when the flag is set (the step's last call, whose own buffer must not be 0). + """ + import cuda.bindings.driver as cuda_driver + import cutlass.cute as cute + + mod = self.mod + if (ar is not None) != mod.FUSED_AR: + raise ValueError("the all-reduce buffers go with an ar_world build, and only with one") + if head_ag is not None and self.fold: + raise ValueError("the head all-gather feeds the unfolded route + quant") + if (head_ready is not None) != mod.HEAD_FLAGS or (mod.HEAD_FLAGS and head_ag is None): + raise ValueError( + "the ready words go with a head_flags build and the head all-gather, and only there" + ) + if front is not None and head_ag is None: + raise ValueError("the front pushes its head into the head all-gather's buffers") + if (lat_slab is not None) != mod.LAT_SLAB: + raise ValueError("the latent slab goes with a lat_slab build, and only with one") + lat_buf, lat_rearm0 = 0, 0 + if lat_slab is not None: + slab, lat_buf, rearm0 = lat_slab + if ( + slab.dtype != torch.int32 + or not slab.is_contiguous() + or slab.numel() != mod.LAT_SLAB_BUFS * MAX_TOKENS * HIDDEN_SIZE // 2 + or not 0 <= lat_buf < mod.LAT_SLAB_BUFS + or (rearm0 and lat_buf == 0) + ): + raise ValueError( + f"k3_moe latent slab: int32 [3, 8, 1792] contiguous, buffer in [0, 3), and not buffer 0 " + f"with the buffer-0 re-arm; got {tuple(slab.shape)} {slab.dtype}, buffer {lat_buf}, " + f"re-arm 0 {rearm0}" + ) + lat_rearm0 = int(bool(rearm0)) + num_tokens = x.shape[0] + shared = None + a1, sfa1, a2, sfa2, _, state_t = self._layer(w31, w31s, w2, w2s) + ar_uc_t, ar_mc_t, ar_flags_t, ar_rank = self._ar_views(ar) + y = torch.empty(num_tokens, HIDDEN_SIZE, dtype=torch.bfloat16, device=x.device) + stream = cuda_driver.CUstream(torch.cuda.current_stream().cuda_stream) + if self.fold: + ids_t, weights_t, b1, sfb1 = self.no_ar_t, self.no_w_t, self.b1_t, self.sfb1_t + fold_in = [ + _view(router_logits.contiguous().view(-1), 16, 0), + _view(bias.detach().contiguous().view(-1), 16, 0), + _view(x.contiguous().view(-1).view(torch.int32), 16, 0), + self.xq_words_t, + self.xsf_t, + ] + else: + if front is not None: + # Registers trtllm::k3_moe_front. + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import front_op # noqa: F401 + + w_front, shared_cols, gate_cap, linear_cap, world = front + ids, weights, x_fp8, x_sf, shared = torch.ops.trtllm.k3_moe_front( + x, w_front, bias, scale, shared_cols, gate_cap, linear_cap, *head_ag, world, + ag_ready=head_ready, + ) # fmt: skip + elif head_ag is not None: + ids, weights, x_fp8, x_sf = torch.ops.trtllm.k3_route_quant_ag( + x, bias, scale, *head_ag, early_trigger=mod.USE_PDL, ag_ready=head_ready + ) + else: + ids, weights, x_fp8, x_sf = self.route_quant( + router_logits, bias, x, scale, **self.route_kwargs + ) + ids_t, weights_t = _view(ids, 4, 1), _view(weights, 4, 1) + b1 = _view(x_fp8.view(torch.uint8).permute(1, 0), 16, 0, mod.b_dtype) + sfb1 = _view( + x_sf.view(torch.uint8).view(num_tokens, HIDDEN_SIZE // _SF_VEC), 16, 1, mod.sf_dtype + ) + fold_in = [self.no_ar_t] * 5 + if mod.HEAD_FLAGS: + flag_in = [self._flat_view(head_ready), self._flat_view(head_ag[2])] + else: + flag_in = [self.no_ar_t, self.no_ar_t] + args = [ + a1, b1, sfa1, sfb1, self.c_t, self.cs_t, self.c_words_t, self.cs_words_t, a2, + self.b2_t, sfa2, self.sfb2_t, _view(y, 16, 1), _view(y.view(torch.int32), 16, 1), + self.part_t, ids_t, weights_t, state_t, ar_uc_t, ar_mc_t, ar_flags_t, + *fold_in, *flag_in, + self._flat_view(lat_slab[0]) if lat_slab is not None else self.no_ar_t, + ] # fmt: skip + scalars = (num_tokens, local_offset, num_local, ar_rank, float(scale), lat_buf, lat_rearm0) + if self.moe is None: + self.moe = cute.compile(mod.k3_moe, *args, *scalars, stream) + self.moe(*args, *scalars, stream) + if mod.AR_PUSH_ONLY: + y = y[:0] # the rows went to every rank's buffer; the consumer reduces them + return y if shared is None else (y, shared) + + +def _state( + device: torch.device, + i_tp: int, + num_local: int, + ar_world: int = 0, + head_flags: bool = False, + lat_slab: bool = False, + ar_push_only: bool = False, +) -> _K3FusedMoE: + key = ( + device.index if device.index is not None else torch.cuda.current_device(), + i_tp, + num_local, + ar_world, + head_flags, + lat_slab, + ar_push_only, + ) + st = _states.get(key) + if st is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "trtllm::k3_fused_moe must run once outside CUDA-graph capture first " + "(it compiles its kernels and allocates its scratch on the first call)." + ) + with _lock: + st = _states.get(key) + if st is None: + # head_flags always explicit: its environment fallback must not reach the plain op. + config = {"head_flags": int(head_flags), "lat_slab": int(lat_slab)} + if ar_world: + config["ar_world"] = ar_world + config["ar_push_only"] = int(ar_push_only) + st = _states[key] = _K3FusedMoE(device, i_tp, num_local, config) + return st + + +# --------------------------------------------------------------------------- +# Fused all-reduce: this TP group's Lamport buffers, 2 x [MAX_TOKENS][group][3584] bf16 per +# rank behind one multicast mapping (the MNNVL all-reduce's allocator), and a local flag +# holding the buffer the next call uses. The kernel leaves both buffers empty (every word +# 0x80000000) after each call. Separate from the model's all-reduce workspace, which the +# shared expert's all-reduce uses concurrently on its auxiliary stream. +# --------------------------------------------------------------------------- +AR_BUFFERS = 2 +AR_EMPTY_WORD = -(2**31) +_ar_workspaces: Dict[object, dict] = {} +_head_workspaces: Dict[object, dict] = {} + + +def ar_buffer_words(world: int) -> int: + return AR_BUFFERS * MAX_TOKENS * world * HIDDEN_SIZE // 2 + + +def ar_workspace(mapping) -> dict: + """The fused all-reduce buffers of ``mapping``'s TP group; collective on first use, so + every rank of the group must make the first call at the same point, outside CUDA-graph + capture.""" + return _lamport_workspace(mapping, ar_buffer_words(mapping.tp_size), _ar_workspaces) + + +def _rqag_module(): + """k3_route_quant_ag.py next to this file (also when op.py is loaded outside the package).""" + mod = _modules.get("rqag") + if mod is None: + path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "k3_route_quant_ag.py") + spec = importlib.util.spec_from_file_location(f"{__name__}_k3_route_quant_ag", path) + mod = importlib.util.module_from_spec(spec) + sys.modules[spec.name] = mod + spec.loader.exec_module(mod) + _modules["rqag"] = mod + return mod + + +def head_workspace(mapping) -> dict: + """The head all-gather's buffers of ``mapping``'s TP group (``trtllm::k3_route_quant_ag``), then + ``trtllm::k3_moe_front``'s router partials; collective on first use like ``ar_workspace``. + ``ready``: route_quant_ag's per-token ready words for k3_moe's flag handoff; flags[2] is their + epoch.""" + ws = _lamport_workspace( + mapping, _rqag_module().workspace_words(mapping.tp_size), _head_workspaces + ) + if "ready" not in ws: + ws["ready"] = torch.zeros(32, dtype=torch.int32, device=ws["uc"].device) + return ws + + +def _lamport_workspace(mapping, words: int, cache: Dict[object, dict]) -> dict: + ws = cache.get(mapping) + if ws is not None: + return ws + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "k3_fused_moe: the head all-gather and fused all-reduce buffers are collective on first " + "use and must be allocated outside CUDA-graph capture" + ) + from tensorrt_llm._torch.distributed.ops import ( + _get_mnnvl_workspace_comm, + _make_mnnvl_mcast_buffer, + _mnnvl_workspace_all_succeeded, + ) + + world = mapping.tp_size + comm = _get_mnnvl_workspace_comm(mapping) + use_fabric_handle = ( + os.environ.get("TRTLLM_FORCE_MNNVL_AR", "0") == "1" or mapping.is_multi_node() + ) + error: Optional[Exception] = None + ws = None + try: + handle = _make_mnnvl_mcast_buffer(comm, words * 4, mapping, use_fabric_handle) + uc = handle.get_uc_buffer(mapping.tp_rank, (words,), torch.int32, 0) + mc = handle.get_mc_buffer((words,), torch.int32, 0) + with torch.inference_mode(): + uc.fill_(AR_EMPTY_WORD) + flags = torch.zeros(4, dtype=torch.int32, device=uc.device) + torch.cuda.synchronize() + ws = dict( + handle=handle, comm=comm, uc=uc, mc=mc, flags=flags, rank=mapping.tp_rank, world=world + ) + except Exception as exc: # noqa: BLE001 -- reported to every rank below, then re-raised + error = exc + # Also the barrier that keeps any rank from pushing into a peer's buffer before the + # peer has emptied it. + if not _mnnvl_workspace_all_succeeded(comm, error is None): + raise RuntimeError("k3_fused_moe Lamport buffers failed on at least one rank") from error + cache[mapping] = ws + return ws + + +# --------------------------------------------------------------------------- +# The head all-gather + route + quant (k3_route_quant_ag.py): one kernel instead of the MNNVL +# all-gather of the sharded head followed by trtllm::k3_route_quant. +# --------------------------------------------------------------------------- +_rqag_compiled: Dict[Tuple[int, bool, bool, bool], object] = {} + + +@torch.library.custom_op("trtllm::k3_route_quant_ag", mutates_args=()) +def k3_route_quant_ag( + head: torch.Tensor, + bias: torch.Tensor, + routed_scaling_factor: float, + ag_uc: torch.Tensor, + ag_mc: torch.Tensor, + ag_flags: torch.Tensor, + ag_rank: int, + early_trigger: bool = False, + ag_ready: Optional[torch.Tensor] = None, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Gathers every rank's slice of the sharded MoE head (``head``: this rank's fp32 + ``[M, (3584 + 896) / world]``, latent columns then router logits) through ``head_workspace``'s + buffers and returns ``trtllm::k3_route_quant``'s outputs for the gathered logits and latent: + ``(topk_ids, topk_weights, quantized, scales)``. With ``ag_ready`` it also releases the + per-token ready words for a k3_moe built with head_flags.""" + import cuda.bindings.driver as cuda_driver + import cutlass.cute as cute + from cutlass.cute.runtime import from_dlpack + + rqag = _rqag_module() + num_tokens, width = head.shape + world = (HIDDEN_SIZE + NUM_EXPERTS) // width + if ( + head.dtype != torch.float32 + or width * world != HIDDEN_SIZE + NUM_EXPERTS + or not 0 < num_tokens <= MAX_TOKENS + ): + raise ValueError( + f"k3_route_quant_ag: head must be fp32 [M <= {MAX_TOKENS}, (3584 + 896) / world]" + ) + device = head.device + topk_ids = torch.empty(num_tokens, TOP_K, dtype=torch.int32, device=device) + topk_weights = torch.empty(num_tokens, TOP_K, dtype=torch.bfloat16, device=device) + quantized = torch.empty(num_tokens, HIDDEN_SIZE, dtype=torch.float8_e4m3fn, device=device) + scales = torch.empty(num_tokens, HIDDEN_SIZE // _SF_VEC, dtype=torch.uint8, device=device) + + def arg(t): + return from_dlpack(t.detach(), assumed_align=16).mark_layout_dynamic(leading_dim=0) + + args = ( + arg(head.contiguous().view(-1).view(torch.int32)), + arg(bias.contiguous().view(-1)), + arg(ag_uc.view(-1)), + arg(ag_mc.view(-1)), + arg(ag_flags.view(-1)), + arg(topk_ids.view(-1)), + arg(topk_weights.view(-1).view(torch.int16)), + arg(quantized.view(-1).view(torch.int32)), + arg(scales.view(-1)), + arg(ag_ready.view(-1) if ag_ready is not None else ag_flags.view(-1)), + ) + publish = ag_ready is not None + stream = cuda_driver.CUstream(torch.cuda.current_stream(device).cuda_stream) + use_pdl = os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" + key = (world, bool(early_trigger), publish, use_pdl) + fn = _rqag_compiled.get(key) + if fn is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "trtllm::k3_route_quant_ag must run once outside CUDA-graph capture first" + ) + with _lock: + fn = _rqag_compiled.get(key) + if fn is None: + fn = _rqag_compiled[key] = cute.compile( + rqag.k3_route_quant_ag, *args, num_tokens, ag_rank, float(routed_scaling_factor), world, + bool(early_trigger), publish, use_pdl, stream, + ) # fmt: skip + fn(*args, num_tokens, ag_rank, float(routed_scaling_factor), stream) + return topk_ids, topk_weights, quantized, scales + + +@k3_route_quant_ag.register_fake +def _( + head, + bias, + routed_scaling_factor, + ag_uc, + ag_mc, + ag_flags, + ag_rank, + early_trigger=False, + ag_ready=None, +): + num_tokens = head.shape[0] + return ( + head.new_empty((num_tokens, TOP_K), dtype=torch.int32), + head.new_empty((num_tokens, TOP_K), dtype=torch.bfloat16), + head.new_empty((num_tokens, HIDDEN_SIZE), dtype=torch.float8_e4m3fn), + head.new_empty((num_tokens, HIDDEN_SIZE // _SF_VEC), dtype=torch.uint8), + ) + + +@torch.library.custom_op("trtllm::k3_fused_moe_head", mutates_args=()) +def k3_fused_moe_head( + head: torch.Tensor, + e_score_correction_bias: torch.Tensor, + w3_w1_weight: torch.Tensor, + w3_w1_weight_scale: torch.Tensor, + w2_weight: torch.Tensor, + w2_weight_scale: torch.Tensor, + local_expert_offset: int, + local_num_experts: int, + routed_scaling_factor: float, + ag_uc: torch.Tensor, + ag_mc: torch.Tensor, + ag_flags: torch.Tensor, + ag_rank: int, + ar_uc: Optional[torch.Tensor] = None, + ar_mc: Optional[torch.Tensor] = None, + ar_flags: Optional[torch.Tensor] = None, + ar_rank: int = -1, + ag_ready: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """``trtllm::k3_fused_moe`` (or ``_ar`` with the ``ar_*`` buffers) from this rank's slice of + the sharded MoE head: ``trtllm::k3_route_quant_ag`` gathers it, routes and quantizes, then + ``k3_moe``. Returns ``[M, 3584]`` bf16 (the partial, or the reduced latent). With ``ag_ready`` + (``head_workspace``'s ready words) k3_moe acquires route_quant_ag's outputs through them + instead of waiting for its grid.""" + if head.shape[0] > MAX_TOKENS: + raise ValueError(f"k3_fused_moe handles at most {MAX_TOKENS} tokens, got {head.shape[0]}") + ar = None + world = 0 + if ar_uc is not None: + ar = (ar_uc, ar_mc, ar_flags, ar_rank) + world = ar_uc.numel() // ar_buffer_words(1) + st = _state( + head.device, w3_w1_weight.shape[1] // 2, local_num_experts, world, ag_ready is not None + ) + return st( + head, None, e_score_correction_bias, w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale, + local_expert_offset, local_num_experts, routed_scaling_factor, ar=ar, + head_ag=(ag_uc, ag_mc, ag_flags, ag_rank), head_ready=ag_ready, + ) # fmt: skip + + +@k3_fused_moe_head.register_fake +def _(head, e_score_correction_bias, w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale, + local_expert_offset, local_num_experts, routed_scaling_factor, ag_uc, ag_mc, ag_flags, ag_rank, + ar_uc=None, ar_mc=None, ar_flags=None, ar_rank=-1, ag_ready=None): # fmt: skip + return head.new_empty((head.shape[0], HIDDEN_SIZE), dtype=torch.bfloat16) + + +@torch.library.custom_op("trtllm::k3_fused_moe_front", mutates_args=()) +def k3_fused_moe_front( + x: torch.Tensor, + w_front: torch.Tensor, + e_score_correction_bias: torch.Tensor, + w3_w1_weight: torch.Tensor, + w3_w1_weight_scale: torch.Tensor, + w2_weight: torch.Tensor, + w2_weight_scale: torch.Tensor, + local_expert_offset: int, + local_num_experts: int, + routed_scaling_factor: float, + shared_cols: int, + gate_cap: float, + linear_cap: float, + ag_uc: torch.Tensor, + ag_mc: torch.Tensor, + ag_flags: torch.Tensor, + ag_rank: int, + ag_world: int, + ar_uc: Optional[torch.Tensor] = None, + ar_mc: Optional[torch.Tensor] = None, + ar_flags: Optional[torch.Tensor] = None, + ar_rank: int = -1, + ag_ready: Optional[torch.Tensor] = None, +) -> Tuple[torch.Tensor, torch.Tensor]: + """``trtllm::k3_moe_front`` (head GEMV, head all-gather, routing, MXFP8 latent, shared gate_up + SiTU) then + ``k3_moe`` (or its fused all-reduce build with the ``ar_*`` buffers) for the MoE input ``x`` bf16 ``[M, 7168]``. + Returns ``(y [M, 3584] bf16, shared activation [M, shared_cols] bf16)``. With ``ag_ready`` + (``head_workspace``'s ready words) k3_moe acquires the front's routing and MXFP8 rows through them instead of + waiting for the front's grid, whose shared tiles may still be running.""" + if x.shape[0] > MAX_TOKENS: + raise ValueError(f"k3_fused_moe handles at most {MAX_TOKENS} tokens, got {x.shape[0]}") + ar = None + world = 0 + if ar_uc is not None: + ar = (ar_uc, ar_mc, ar_flags, ar_rank) + world = ar_uc.numel() // ar_buffer_words(1) + st = _state( + x.device, w3_w1_weight.shape[1] // 2, local_num_experts, world, ag_ready is not None + ) + return st( + x, None, e_score_correction_bias, w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale, + local_expert_offset, local_num_experts, routed_scaling_factor, ar=ar, + head_ag=(ag_uc, ag_mc, ag_flags, ag_rank), head_ready=ag_ready, + front=(w_front, shared_cols, gate_cap, linear_cap, ag_world), + ) # fmt: skip + + +@k3_fused_moe_front.register_fake +def _(x, w_front, e_score_correction_bias, w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale, + local_expert_offset, local_num_experts, routed_scaling_factor, shared_cols, gate_cap, linear_cap, ag_uc, ag_mc, + ag_flags, ag_rank, ag_world, ar_uc=None, ar_mc=None, ar_flags=None, ar_rank=-1, ag_ready=None): # fmt: skip + return ( + x.new_empty((x.shape[0], HIDDEN_SIZE), dtype=torch.bfloat16), + x.new_empty((x.shape[0], shared_cols), dtype=torch.bfloat16), + ) + + +@torch.library.custom_op("trtllm::k3_fused_moe", mutates_args=()) +def k3_fused_moe( + hidden_states: torch.Tensor, + router_logits: torch.Tensor, + e_score_correction_bias: torch.Tensor, + w3_w1_weight: torch.Tensor, + w3_w1_weight_scale: torch.Tensor, + w2_weight: torch.Tensor, + w2_weight_scale: torch.Tensor, + local_expert_offset: int, + local_num_experts: int, + routed_scaling_factor: float, +) -> torch.Tensor: + """Kimi K3 routed experts for M <= 8 decode tokens; returns ``[M, 3584]`` bf16. + + ``hidden_states``: bf16 ``[M, 3584]`` latent; ``router_logits``: fp32 ``[M, 896]``; + ``e_score_correction_bias``: fp32 ``[896]``; weights: the TRTLLM-Gen + W4A8_MXFP4_MXFP8 buffers of this rank's ``local_num_experts`` experts, which hold + global ids ``[local_expert_offset, local_expert_offset + local_num_experts)``. + """ + if hidden_states.shape[0] > MAX_TOKENS: + raise ValueError( + f"k3_fused_moe handles at most {MAX_TOKENS} tokens, got {hidden_states.shape[0]}" + ) + st = _state(hidden_states.device, w3_w1_weight.shape[1] // 2, local_num_experts) + return st( + hidden_states, router_logits, e_score_correction_bias, w3_w1_weight, w3_w1_weight_scale, + w2_weight, w2_weight_scale, local_expert_offset, local_num_experts, routed_scaling_factor, + ) # fmt: skip + + +@k3_fused_moe.register_fake +def _(hidden_states, router_logits, e_score_correction_bias, w3_w1_weight, w3_w1_weight_scale, w2_weight, + w2_weight_scale, local_expert_offset, local_num_experts, routed_scaling_factor): # fmt: skip + return hidden_states.new_empty((hidden_states.shape[0], HIDDEN_SIZE), dtype=torch.bfloat16) + + +@torch.library.custom_op("trtllm::k3_fused_moe_ar", mutates_args=("lat_slab",)) +def k3_fused_moe_ar( + hidden_states: torch.Tensor, + router_logits: torch.Tensor, + e_score_correction_bias: torch.Tensor, + w3_w1_weight: torch.Tensor, + w3_w1_weight_scale: torch.Tensor, + w2_weight: torch.Tensor, + w2_weight_scale: torch.Tensor, + local_expert_offset: int, + local_num_experts: int, + routed_scaling_factor: float, + ar_uc: torch.Tensor, + ar_mc: torch.Tensor, + ar_flags: torch.Tensor, + ar_rank: int, + lat_slab: Optional[torch.Tensor] = None, + lat_buf: int = 0, + lat_rearm0: bool = False, +) -> torch.Tensor: + """``trtllm::k3_fused_moe`` followed by the all-reduce of the routed partial over the + group of ``ar_workspace``'s buffers (``ar_uc``, ``ar_mc``, ``ar_flags``), in one kernel. + Returns the reduced ``[M, 3584]`` bf16 latent, identical on every rank. With ``lat_slab`` + (int32 ``[3, 8, 1792]``, all-ones empty) the reduced rows are also published into buffer + ``lat_buf`` for a consumer that polls them, and the call re-arms buffer ``(lat_buf + 1) % 3`` + (and buffer 0 with ``lat_rearm0``, the step's last call).""" + if hidden_states.shape[0] > MAX_TOKENS: + raise ValueError( + f"k3_fused_moe handles at most {MAX_TOKENS} tokens, got {hidden_states.shape[0]}" + ) + world = ar_uc.numel() // ar_buffer_words(1) + st = _state( + hidden_states.device, w3_w1_weight.shape[1] // 2, local_num_experts, world, + lat_slab=lat_slab is not None, + ) # fmt: skip + return st( + hidden_states, router_logits, e_score_correction_bias, w3_w1_weight, w3_w1_weight_scale, + w2_weight, w2_weight_scale, local_expert_offset, local_num_experts, routed_scaling_factor, + ar=(ar_uc, ar_mc, ar_flags, ar_rank), + lat_slab=(lat_slab, lat_buf, lat_rearm0) if lat_slab is not None else None, + ) # fmt: skip + + +@k3_fused_moe_ar.register_fake +def _(hidden_states, router_logits, e_score_correction_bias, w3_w1_weight, w3_w1_weight_scale, w2_weight, + w2_weight_scale, local_expert_offset, local_num_experts, routed_scaling_factor, ar_uc, ar_mc, ar_flags, + ar_rank, lat_slab=None, lat_buf=0, lat_rearm0=False): # fmt: skip + return hidden_states.new_empty((hidden_states.shape[0], HIDDEN_SIZE), dtype=torch.bfloat16) + + +# --------------------------------------------------------------------------- +# Steps of up to 64 tokens (e.g. R x 8 speculative verify tokens): the m_max 64 build of +# ``k3_moe``, one launch per call, after ``trtllm::k3_route_quant``. Its state belongs to the +# caller: nothing here is cached per process beyond the kernel module of the configuration. + +WIDE_MAX_TOKENS = 64 + + +class K3MoeWideState: + """``k3_moe`` for 1..64 tokens on one device: the compiled kernel and the scratch its layers + share, i.e. the FC1 -> FC2 intermediate slab (armed between calls: FP8 -0.0 values, E8M0 NaN + scale words) and the FC2 slice partials, all sized for 64 tokens and kept at fixed addresses. + Build it eagerly before CUDA-graph capture and keep it with the model; every layer takes its + own counters from :meth:`layer`. The layers of one state run in one stream order (they share + the scratch). The kernel is compiled (TVM-FFI, explicit stream) by the first call, which must + therefore come before capture. + + ``use_pdl``: launch ``k3_moe`` as a programmatic dependent; its producer must then be + ``trtllm::k3_route_quant`` with ``early_trigger=True`` (or any kernel whose outputs ``k3_moe`` + may read once that grid has completed).""" + + def __init__(self, device: torch.device, i_tp: int, num_local: int, use_pdl: bool = True): + num_ctas = torch.cuda.get_device_properties(device).multi_processor_count + config = { + "i_tp": i_tp, + "num_ctas": num_ctas, + "num_local": num_local, + "m_max": WIDE_MAX_TOKENS, + "pdl": int(use_pdl), + } + self.mod = mod = _kernel_module(config) + self.device = device + self.i_tp = i_tp + self.num_local = num_local + g_cap = mod.G_CAP + kw = dict(device=device) + self.c = torch.full((g_cap, _TOKEN_SLOTS, i_tp), -128, dtype=torch.int8, **kw) + self.cs = torch.zeros(g_cap, _TOKEN_SLOTS, mod.SF_STRIDE0, dtype=torch.int8, **kw) + self.cs.view(g_cap, _TOKEN_SLOTS, mod.K2_TILES, mod.SFB_GROUP_BYTES)[..., :4] = -1 + self.part = torch.empty(mod.PART_ROWS, HIDDEN_SIZE, dtype=torch.float32, **kw) + # Stand-in for the buffers of the options this build does not have (fused all-reduce, + # fold, head flags, latent slab). + self.unused = torch.zeros(4, dtype=torch.int32, **kw) + # The kernel's views of the scratch, in its argument order (FP8 / E8M0 data as bytes). + sfb2 = ( + self.cs.view(torch.uint8) + .reshape(g_cap, _TOKEN_SLOTS, mod.K2_TILES, mod.SFB_GROUP_BYTES) + .permute(3, 2, 1, 0) + ) + self.scratch = ( + self.c, + self.cs, + self.c.view(-1).view(torch.int32), + self.cs.view(-1).view(torch.int32), + self.c.view(torch.uint8).permute(2, 1, 0), + sfb2, + self.part, + ) + self.compiled = None + + def layer( + self, + w3_w1_weight: torch.Tensor, + w3_w1_weight_scale: torch.Tensor, + w2_weight: torch.Tensor, + w2_weight_scale: torch.Tensor, + ) -> "K3MoeWideLayer": + """A layer's handle: its experts' TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers (read in place) and + its counters.""" + return K3MoeWideLayer(self, w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale) + + def _compile(self, args, scalars): + """TVM-FFI build for these torch arguments' types and layouts (any M up to 64).""" + import cutlass.cute as cute + + # As the M <= 8 op's views: (alignment, leading dim) per tensor argument; the 11 stand-ins last. + aligns = [16, 16, 16, 16, 16, 4, 16, 16, 16, 16, 16, 16, 16, 16, 16, 4, 4, 4] + [16] * 11 + leading = [0, 0, 0, 1, 2, 2, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 0] + [0] * 11 + assert len(args) == len(aligns) + signature = [_view(t, a, d) for t, a, d in zip(args, aligns, leading)] + return cute.compile( + self.mod.k3_moe, + *signature, + *scalars, + cute.runtime.make_fake_stream(), + options="--enable-tvm-ffi", + ) + + +class K3MoeWideLayer: + """One MoE layer on a :class:`K3MoeWideState`: its weights as the kernel reads them and its + counters (int32, zero between calls; every call leaves them zero).""" + + def __init__( + self, + state: K3MoeWideState, + w3_w1_weight: torch.Tensor, + w3_w1_weight_scale: torch.Tensor, + w2_weight: torch.Tensor, + w2_weight_scale: torch.Tensor, + ): + ok, why = is_supported( + w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale, state.num_local + ) + if not ok or w3_w1_weight.shape[1] != 2 * state.i_tp: + raise ValueError( + f"k3_moe wide layer: {why or 'intermediate size differs from the state'}" + ) + e, two_i, _ = w3_w1_weight.shape + i_tp = two_i // 2 + self.state = state + self.counters = torch.zeros(state.mod.NUM_STATE, dtype=torch.int32, device=state.device) + self.weights = ( + w3_w1_weight.view(torch.int8).permute(2, 1, 0), + w3_w1_weight_scale.view(e, two_i // 128, HIDDEN_SIZE // 128, 512).permute(3, 2, 1, 0), + w2_weight.view(torch.int8).permute(2, 1, 0), + w2_weight_scale.view(e, HIDDEN_SIZE // 128, i_tp // 128, 512).permute(3, 2, 1, 0), + ) + + def __call__( + self, + x_fp8: torch.Tensor, + x_sf: torch.Tensor, + topk_ids: torch.Tensor, + topk_weights: torch.Tensor, + local_expert_offset: int, + out: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + """This rank's routed partial ``[M, 3584]`` bf16 for the outputs of + ``trtllm::k3_route_quant``: ``x_fp8`` float8_e4m3fn ``[M, 3584]``, ``x_sf`` its E8M0 + scales (``M * 112`` bytes), ``topk_ids`` int32 ``[M, 16]`` global expert ids, + ``topk_weights`` bf16 ``[M, 16]``; 1 <= M <= 64. The layer's experts hold global ids + ``[local_expert_offset, local_expert_offset + num_local)``. ``out``: bf16, contiguous, at + least ``[M, 3584]``; its first M rows are the result (a fresh tensor without it). Writes + the state's slab (left armed) and partials and this layer's counters (left zero).""" + st = self.state + num_tokens = topk_ids.shape[0] + if not 0 < num_tokens <= WIDE_MAX_TOKENS: + raise ValueError(f"k3_moe wide: M must be in [1, {WIDE_MAX_TOKENS}], got {num_tokens}") + if ( + topk_ids.dtype != torch.int32 + or tuple(topk_ids.shape) != (num_tokens, TOP_K) + or topk_weights.dtype != torch.bfloat16 + or tuple(topk_weights.shape) != (num_tokens, TOP_K) + or x_fp8.dtype != torch.float8_e4m3fn + or tuple(x_fp8.shape) != (num_tokens, HIDDEN_SIZE) + or x_sf.numel() != num_tokens * (HIDDEN_SIZE // _SF_VEC) + or not (topk_ids.is_contiguous() and topk_weights.is_contiguous()) + or not (x_fp8.is_contiguous() and x_sf.is_contiguous()) + ): + raise ValueError("k3_moe wide: expects trtllm::k3_route_quant's outputs for M tokens") + if out is None: + y = torch.empty(num_tokens, HIDDEN_SIZE, dtype=torch.bfloat16, device=x_fp8.device) + else: + if ( + out.dtype != torch.bfloat16 + or out.dim() != 2 + or out.shape[0] < num_tokens + or out.shape[1] != HIDDEN_SIZE + or not out.is_contiguous() + ): + raise ValueError("k3_moe wide: out must be contiguous bf16 [>= M, 3584]") + y = out[:num_tokens] + a1, sfa1, a2, sfa2 = self.weights + c, cs, c_words, cs_words, b2, sfb2, part = st.scratch + u = st.unused + args = ( + a1, x_fp8.view(torch.uint8).permute(1, 0), sfa1, + x_sf.view(torch.uint8).view(num_tokens, HIDDEN_SIZE // _SF_VEC), c, cs, c_words, cs_words, a2, b2, + sfa2, sfb2, y, y.view(torch.int32), part, topk_ids, topk_weights, self.counters, + u, u, u, u, u, u, u, u, u, u, u, + ) # fmt: skip + scalars = (num_tokens, local_expert_offset, st.num_local, 0, 1.0, 0, 0) + if st.compiled is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "k3_moe wide compiles on its first call: call it once before CUDA-graph capture" + ) + st.compiled = st._compile(args, scalars) + st.compiled(*args, *scalars, torch.cuda.current_stream().cuda_stream) + return y diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_route_quant/__init__.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_route_quant/__init__.py new file mode 100644 index 000000000000..a6869da3dbfb --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_route_quant/__init__.py @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Kimi K3 MoE routing + MXFP8 input quantization in CuTe DSL (``trtllm::k3_route_quant``). + +Importing :mod:`.op` registers the torch op; nothing is imported eagerly here so that +the CuTe DSL dependency stays optional for every other model. +""" diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_route_quant/k3_route_quant_kernel.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_route_quant/k3_route_quant_kernel.py new file mode 100644 index 000000000000..ede52bb522cc --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_route_quant/k3_route_quant_kernel.py @@ -0,0 +1,431 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Kimi K3 MoE routing and MXFP8 input quantization for decode, in CuTe DSL. + +Computes exactly what ``trtllm::kimi_k3_noaux_tc_mxfp8_quant`` computes, bit for bit: + +* routing: ``sigmoid(logit) = 0.5 * tanhf(0.5 * logit) + 0.5``; the 16 experts with the largest + ``sigmoid + bias``, in descending order, ties to the lower expert id; weights + ``bf16(sigmoid * scale / (sum of the 16 sigmoids + 1e-20))`` evaluated in fp64, the sum being the + xor-butterfly warp reduction ``cg::reduce`` performs over lanes 0..15; +* quantization: MXFP8 (e4m3) with one UE8M0 scale per 32 elements in the linear [M, 112] layout, + the recipe of ``cvt_warp_fp16_to_mxfp8`` (scale rounded up from ``amax / 448``). + +The selection differs from the C++ kernel's: one warp holds 28 keys per lane, each lane sorts its +six largest, and the 16 winners come out of 16 rounds of ``redux.sync.max`` over the lanes' heads +plus ``redux.sync.min`` over the candidate ids of the lanes holding that maximum (the tie-break); +the winner's lane shifts its list. The C++ path sorts 32 packed 64-bit keys per lane and runs 16 +rounds of 64-bit shuffle arg-max instead. + +``top16_warp`` (a ``cute.jit`` function) and ``mxfp8_quant_vec8`` (a trace-time helper) can be +inlined into other kernels, e.g. a fused MoE kernel routing in its prologue. +""" + +from __future__ import annotations + +import cuda.bindings.driver as cuda_driver +import cutlass +import cutlass.cute as cute +from cutlass import dsl_user_op +from cutlass.experimental import primitives as prims + +NUM_EXPERTS = 896 +TOP_K = 16 +HIDDEN_SIZE = 3584 +SF_VEC_SIZE = 32 +THREADS = HIDDEN_SIZE // 8 # one 8-element vector per thread in a quantization CTA +KEYS_PER_LANE = NUM_EXPERTS // 32 +FULL_MASK = 0xFFFFFFFF +REMOVED_KEY = -(2**31) # below the key of every float +NO_CANDIDATE = 0x7FFFFFFF +LANE_LIST = 6 # sorted keys each lane keeps for the fast selection rounds + +assert NUM_EXPERTS % 32 == 0 and NUM_EXPERTS == 2 * THREADS and TOP_K <= 32 + + +# ============================================================================= +# PTX the C++ reference uses, emitted verbatim so the results match bit for bit. +# ============================================================================= +def _asm(result_type, operands, text, constraints, *, loc, ip): + from cutlass._mlir.dialects import llvm as _llvm + + return _llvm.inline_asm( + result_type, + [op.ir_value(loc=loc, ip=ip) for op in operands], + text, + constraints, + has_side_effects=True, + is_align_stack=False, + asm_dialect=_llvm.AsmDialect.AD_ATT, + loc=loc, + ip=ip, + ) + + +@dsl_user_op +def rcp_approx_ftz(a, *, loc=None, ip=None): + """``rcp.approx.ftz.f32`` (``reciprocal_approximate_ftz``); never constant-folded.""" + from cutlass._mlir.extras import types as _T + + return cutlass.Float32( + _asm(_T.f32(), [cutlass.Float32(a)], "rcp.approx.ftz.f32 $0, $1;", "=f,f", loc=loc, ip=ip) + ) + + +@dsl_user_op +def cvt_rp_satfinite_ue8m0(a, *, loc=None, ip=None): + """UE8M0 byte of ``a`` rounded toward +inf, saturated to finite (``__nv_cvt_float_to_e8m0``).""" + from cutlass._mlir.extras import types as _T + + pair = cutlass.Int16( + _asm( + _T.i16(), + [cutlass.Float32(a), cutlass.Float32(a)], + "cvt.rp.satfinite.ue8m0x2.f32 $0, $1, $2;", + "=h,f,f", + loc=loc, + ip=ip, + ) + ) + return cutlass.Int32(pair) & cutlass.Int32(0xFF) + + +@dsl_user_op +def cvt_rn_satfinite_e4m3x2(hi, lo, *, loc=None, ip=None): + """Two e4m3 bytes, ``lo`` in bits 0-7 and ``hi`` in bits 8-15 (``__nv_fp8x2_e4m3(float2)``).""" + from cutlass._mlir.extras import types as _T + + pair = cutlass.Int16( + _asm( + _T.i16(), + [cutlass.Float32(hi), cutlass.Float32(lo)], + "cvt.rn.satfinite.e4m3x2.f32 $0, $1, $2;", + "=h,f,f", + loc=loc, + ip=ip, + ) + ) + return cutlass.Int32(pair) & cutlass.Int32(0xFFFF) + + +@dsl_user_op +def cvt_rn_bf16_f64(a, *, loc=None, ip=None): + """bf16 bits of an fp64 value, one rounding (``__double2bfloat16`` on sm_90+).""" + from cutlass._mlir.extras import types as _T + + return cutlass.Int16( + _asm(_T.i16(), [cutlass.Float64(a)], "cvt.rn.bf16.f64 $0, $1;", "=h,d", loc=loc, ip=ip) + ) + + +# ============================================================================= +# Routing +# ============================================================================= +def sigmoid_accurate(x): + """``0.5f * tanhf(0.5f * x) + 0.5f`` with the accurate libdevice ``tanhf``. Both scalings by + 0.5 are exact, so contracting the outer multiply-add into an FMA cannot change the result.""" + half = cutlass.Float32(0.5) + return half * cute.math.tanh(half * x, fastmath=False) + half + + +def selection_key(score): + """Int32 whose signed order is the order ``cub::Traits::TwiddleIn`` gives the bits of + ``score`` as unsigned integers (the C++ top-k's comparison key).""" + bits = score.bitcast(cutlass.Int32) + return bits ^ ((bits >> cutlass.Int32(31)) & cutlass.Int32(0x7FFFFFFF)) + + +def _argmax_lowest_index(values): + """(max, index of its first occurrence) over a list of Int32, as a balanced tree.""" + items = [(v, cutlass.Int32(j)) for j, v in enumerate(values)] + while len(items) > 1: + paired = [] + for i in range(0, len(items) - 1, 2): + (va, ja), (vb, jb) = items[i], items[i + 1] + take_b = vb > va # the left item has the lower indices, so it keeps ties + paired.append( + ( + cutlass.Int32(cutlass.select_(take_b, vb, va)), + cutlass.Int32(cutlass.select_(take_b, jb, ja)), + ) + ) + if len(items) % 2 == 1: + paired.append(items[-1]) + items = paired + return items[0] + + +def _select_i32(pred, a, b): + return cutlass.Int32(cutlass.select_(pred, a, b)) + + +def _round_winner(head, head_slot, lane): + """One selection round: the largest head over the warp, the lowest expert id among equal heads. + Returns (winner id, this lane's candidate id).""" + top = prims.redux_sync(head, prims.ReductionKind.MAX, FULL_MASK) + candidate = _select_i32( + head == top, head_slot * cutlass.Int32(32) + lane, cutlass.Int32(NO_CANDIDATE) + ) + return prims.redux_sync(candidate, prims.ReductionKind.MIN, FULL_MASK), candidate + + +def _load_lane_keys(s_key, lane): + return [s_key.load(idx=lane + cutlass.Int32(32 * slot)) for slot in range(KEYS_PER_LANE)] + + +def _rounds_exact(lane_keys, lane): + """The 16 rounds over all 28 keys of every lane, the lane's argmax recomputed after each pop. + Returns the round-r winner in lane r.""" + keys = list(lane_keys) + best, best_slot = _argmax_lowest_index(keys) + expert = cutlass.Int32(0) + for rank in range(TOP_K): + winner, candidate = _round_winner(best, best_slot, lane) + expert = _select_i32(lane == cutlass.Int32(rank), winner, expert) + if rank < TOP_K - 1: + popped = candidate == winner + for slot in range(KEYS_PER_LANE): + keys[slot] = _select_i32( + popped & (best_slot == cutlass.Int32(slot)), + cutlass.Int32(REMOVED_KEY), + keys[slot], + ) + best, best_slot = _argmax_lowest_index(keys) + return expert + + +def _lane_list(lane_keys): + """The lane's LANE_LIST largest keys and their slots, descending, the lower slot first among equal + keys (insertion with a strict comparison).""" + values = [cutlass.Int32(REMOVED_KEY)] * LANE_LIST + slots = [cutlass.Int32(0)] * LANE_LIST + for slot, key in enumerate(lane_keys): + above = [key > v for v in values] + new_values = [_select_i32(above[0], key, values[0])] + new_slots = [_select_i32(above[0], cutlass.Int32(slot), slots[0])] + for i in range(1, LANE_LIST): + new_values.append( + _select_i32(above[i - 1], values[i - 1], _select_i32(above[i], key, values[i])) + ) + new_slots.append( + _select_i32( + above[i - 1], slots[i - 1], _select_i32(above[i], cutlass.Int32(slot), slots[i]) + ) + ) + values, slots = new_values, new_slots + return values, slots + + +def _rounds_from_lists(lane_keys, lane): + """The 16 rounds over each lane's sorted list of its LANE_LIST largest keys, a pop being a shift. + Returns (the round-r winner in lane r, nonzero in a lane whose list ran empty before a round).""" + values, slots = _lane_list(lane_keys) + expert = cutlass.Int32(0) + emptied = cutlass.Int32(0) + for rank in range(TOP_K): + if rank > 0: + emptied = emptied | _select_i32( + values[0] == cutlass.Int32(REMOVED_KEY), cutlass.Int32(1), cutlass.Int32(0) + ) + winner, candidate = _round_winner(values[0], slots[0], lane) + expert = _select_i32(lane == cutlass.Int32(rank), winner, expert) + if rank < TOP_K - 1: + popped = candidate == winner + for i in range(LANE_LIST - 1): + values[i] = _select_i32(popped, values[i + 1], values[i]) + slots[i] = _select_i32(popped, slots[i + 1], slots[i]) + values[-1] = _select_i32(popped, cutlass.Int32(REMOVED_KEY), values[-1]) + return expert, emptied + + +def _routing_weight(s_sigmoid, expert, lane, routed_scaling_factor): + """bf16 bits of ``sigmoid * scale / (sum of the 16 sigmoids + 1e-20)`` in fp64, the sum being + ``cg::reduce``'s xor butterfly over the warp with lanes 16-31 contributing 0.""" + selected = lane < cutlass.Int32(TOP_K) + sig = cutlass.Float32( + cutlass.select_(selected, s_sigmoid.load(idx=expert), cutlass.Float32(0.0)) + ) + total = sig + for offset in (16, 8, 4, 2, 1): + total = total + cute.arch.shuffle_sync_bfly(total, offset=offset) + weight = (cutlass.Float64(sig) * routed_scaling_factor) / ( + cutlass.Float64(total) + cutlass.Float64(1e-20) + ) + return cvt_rn_bf16_f64(weight) + + +@cute.jit +def top16_warp(s_key, s_sigmoid, lane, routed_scaling_factor): + """Top-16 experts of one token and their routing weights, computed by one whole warp. + + ``s_key`` (Int32 [896], from ``selection_key(sigmoid + bias)``) and ``s_sigmoid`` (Float32 + [896]) hold the token in shared memory. Returns ``(expert_id, weight_bf16_bits)``; lane r < 16 + holds the rank-r expert, the other lanes hold garbage. All 32 lanes must call it together. + + Each lane first sorts its LANE_LIST largest keys, so a round costs two ``redux.sync`` and a + shift. A lane can hold more of the winners than that (P ~ 1.4e-5 per lane for random keys); the + warp then redoes the rounds over all 28 keys per lane, so the result is always exact. + """ + lane_keys = _load_lane_keys(s_key, lane) + expert, emptied = _rounds_from_lists(lane_keys, lane) + if prims.vote_sync(FULL_MASK, emptied != cutlass.Int32(0), prims.VoteSync.ANY): + expert = _rounds_exact(lane_keys, lane) + return expert, _routing_weight(s_sigmoid, expert, lane, routed_scaling_factor) + + +# ============================================================================= +# MXFP8 quantization +# ============================================================================= +def mxfp8_quant_vec8(words): + """MXFP8 of eight bf16 values given as four Int32 words (element 2i in the low half of word i). + + The four consecutive lanes that share one 32-element scale must call it together. Returns + ``(q_lo, q_hi, sf_byte)``: the e4m3 bytes of elements 0-3 and 4-7 (element order, little + endian) and the UE8M0 scale byte of the lane's group. + """ + values = [] + for w in words: + values.append((w << cutlass.Int32(16)).bitcast(cutlass.Float32)) + values.append((w & cutlass.Int32(-65536)).bitcast(cutlass.Float32)) + amax = cute.arch.fmax(cute.math.abs(values[0]), cute.math.abs(values[1])) + for v in values[2:]: + amax = cute.arch.fmax(amax, cute.math.abs(v)) + for offset in (1, 2): + amax = cute.arch.fmax(cute.arch.shuffle_sync_bfly(amax, offset=offset), amax) + + sf_byte = cvt_rp_satfinite_ue8m0(amax * rcp_approx_ftz(cutlass.Float32(448.0))) + # static_cast(__nv_fp8_e8m0): 2^(byte - 127), with byte 0 the fp32 denormal 2^-127. + sf_bits = cutlass.Int32( + cutlass.select_( + sf_byte == cutlass.Int32(0), cutlass.Int32(0x00400000), sf_byte << cutlass.Int32(23) + ) + ) + sf_bits = cutlass.Int32( + cutlass.select_(sf_byte == cutlass.Int32(0xFF), cutlass.Int32(0x7FFFFFFF), sf_bits) + ) + out_scale = cutlass.Float32( + cutlass.select_( + amax != cutlass.Float32(0.0), + rcp_approx_ftz(sf_bits.bitcast(cutlass.Float32)), + cutlass.Float32(0.0), + ) + ) + pairs = [ + cvt_rn_satfinite_e4m3x2(values[2 * i + 1] * out_scale, values[2 * i] * out_scale) + for i in range(4) + ] + q_lo = pairs[0] | (pairs[1] << cutlass.Int32(16)) + q_hi = pairs[2] | (pairs[3] << cutlass.Int32(16)) + return q_lo, q_hi, sf_byte + + +# ============================================================================= +# Kernel: CTAs [0, M) route token b, CTAs [M, 2M) quantize row b - M (the C++ kernel's split). +# ============================================================================= +@cute.kernel +def k3_route_quant_kernel( + scores: cutlass.Array, # fp32 [M * 896] router logits + bias: cutlass.Array, # fp32 [896] + hidden_words: cutlass.Array, # int32 view of bf16 [M, 3584]: [M * 1792] + topk_ids: cutlass.Array, # int32 [M * 16] + topk_weight_bits: cutlass.Array, # int16 view of bf16 [M * 16] + quant_words: cutlass.Array, # int32 view of e4m3 [M, 3584]: [M * 896] + scales: cutlass.Array, # uint8 [M * 112] + num_tokens: cutlass.Int32, + routed_scaling_factor: cutlass.Float64, + early_trigger: cutlass.Constexpr[bool], +): + tidx, _, _ = cute.arch.thread_idx() + bidx, _, _ = cute.arch.block_idx() + s_key = cutlass.Array(cutlass.Int32, NUM_EXPERTS, space=cutlass.AddressSpace.smem, alignment=16) + s_sigmoid = cutlass.Array( + cutlass.Float32, NUM_EXPERTS, space=cutlass.AddressSpace.smem, alignment=16 + ) + routes = bidx < num_tokens + + # The bias is a weight: read it before waiting for the producer of the logits (every CTA + # reads it, the index is valid for all). + bias_lo = bias.load(idx=tidx) + bias_hi = bias.load(idx=tidx + cutlass.Int32(THREADS)) + + prims.griddepcontrol(prims.GridDepAction.WAIT) + if cutlass.const_expr(early_trigger): + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + + if routes: + row_base = bidx * cutlass.Int32(NUM_EXPERTS) + sig_lo = sigmoid_accurate(scores.load(idx=row_base + tidx)) + sig_hi = sigmoid_accurate(scores.load(idx=row_base + tidx + cutlass.Int32(THREADS))) + s_sigmoid.store(sig_lo, idx=tidx) + s_sigmoid.store(sig_hi, idx=tidx + cutlass.Int32(THREADS)) + s_key.store(selection_key(sig_lo + bias_lo), idx=tidx) + s_key.store(selection_key(sig_hi + bias_hi), idx=tidx + cutlass.Int32(THREADS)) + cute.arch.barrier() + if tidx < cutlass.Int32(32): + expert, weight_bits = top16_warp(s_key, s_sigmoid, tidx, routed_scaling_factor) + if tidx < cutlass.Int32(TOP_K): + out = bidx * cutlass.Int32(TOP_K) + tidx + topk_ids.store(expert, idx=out) + topk_weight_bits.store(weight_bits, idx=out) + else: + row = bidx - num_tokens + words = hidden_words.load( + idx=row * cutlass.Int32(HIDDEN_SIZE // 2) + tidx * cutlass.Int32(4), + vector_size=4, + alignment=16, + ) + q_lo, q_hi, sf_byte = mxfp8_quant_vec8([words[0], words[1], words[2], words[3]]) + quant_words.store( + (q_lo, q_hi), + idx=row * cutlass.Int32(HIDDEN_SIZE // 4) + tidx * cutlass.Int32(2), + alignment=8, + ) + if tidx % cutlass.Int32(SF_VEC_SIZE // 8) == cutlass.Int32(0): + scales.store( + cutlass.Uint8(sf_byte), + idx=row * cutlass.Int32(HIDDEN_SIZE // SF_VEC_SIZE) + + tidx // cutlass.Int32(SF_VEC_SIZE // 8), + ) + + if cutlass.const_expr(not early_trigger): + cute.arch.fence_acq_rel_gpu() + cute.arch.barrier() + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + + +@cute.jit +def k3_route_quant( + scores: cute.Tensor, + bias: cute.Tensor, + hidden_words: cute.Tensor, + topk_ids: cute.Tensor, + topk_weight_bits: cute.Tensor, + quant_words: cute.Tensor, + scales: cute.Tensor, + num_tokens: cutlass.Int32, + routed_scaling_factor: cutlass.Float64, + early_trigger: cutlass.Constexpr[bool], + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + k3_route_quant_kernel( + scores, bias, hidden_words, topk_ids, topk_weight_bits, quant_words, scales, num_tokens, + routed_scaling_factor, early_trigger, + ).launch( + grid=[num_tokens * 2, 1, 1], + block=[THREADS, 1, 1], + stream=stream, + use_pdl=use_pdl, + ) # fmt: skip diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_route_quant/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_route_quant/op.py new file mode 100644 index 000000000000..d7c63fc8394a --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_route_quant/op.py @@ -0,0 +1,146 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``trtllm::k3_route_quant``: the CuTe DSL form of ``trtllm::kimi_k3_noaux_tc_mxfp8_quant``. + +Same arguments, same outputs bit for bit: top-16 expert ids (int32 [M, 16]), routing weights +(bf16 [M, 16]), the MXFP8 latent (e4m3 [M, 3584]) and its UE8M0 scales (uint8 [M, 112], linear). +The kernel is compiled on the first call for each (early trigger, PDL) pair, which must happen +outside CUDA-graph capture (the model's warmup does it). +""" + +from __future__ import annotations + +import os +import threading +from typing import Dict, Tuple + +import torch + +NUM_EXPERTS = 896 +TOP_K = 16 +HIDDEN_SIZE = 3584 +SF_VEC_SIZE = 32 +MAX_TOKENS = 64 + +_lock = threading.Lock() +_compiled: Dict[Tuple[bool, bool], object] = {} + + +def _arg(t: torch.Tensor): + from cutlass.cute.runtime import from_dlpack + + # detach(): DLPack refuses tensors that require grad, e.g. the routing bias parameter. + return from_dlpack(t.detach(), assumed_align=16).mark_layout_dynamic(leading_dim=0) + + +def _use_pdl() -> bool: + # Same switch as the C++ launches (tensorrt_llm::common::getEnvEnablePDL). + return os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" + + +def _check(scores: torch.Tensor, bias: torch.Tensor, hidden_states: torch.Tensor) -> None: + if not (scores.is_cuda and bias.is_cuda and hidden_states.is_cuda): + raise ValueError("k3_route_quant: all inputs must be CUDA tensors") + if scores.dtype != torch.float32 or bias.dtype != torch.float32: + raise ValueError("k3_route_quant: scores and bias must be float32") + if hidden_states.dtype != torch.bfloat16: + raise ValueError("k3_route_quant: hidden_states must be bfloat16") + if not (scores.is_contiguous() and bias.is_contiguous() and hidden_states.is_contiguous()): + raise ValueError("k3_route_quant: all inputs must be contiguous") + if scores.dim() != 2 or scores.shape[1] != NUM_EXPERTS or bias.numel() != NUM_EXPERTS: + raise ValueError( + f"k3_route_quant: scores must be [M, {NUM_EXPERTS}] and bias [{NUM_EXPERTS}]" + ) + if hidden_states.shape != (scores.shape[0], HIDDEN_SIZE): + raise ValueError( + f"k3_route_quant: hidden_states must be [M, {HIDDEN_SIZE}] with the M of scores" + ) + if not 0 < scores.shape[0] <= MAX_TOKENS: + raise ValueError(f"k3_route_quant: M must be in [1, {MAX_TOKENS}]") + + +def _kernel(early_trigger: bool, use_pdl: bool, args, num_tokens: int, scale: float, stream): + key = (early_trigger, use_pdl) + fn = _compiled.get(key) + if fn is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "trtllm::k3_route_quant must run once outside CUDA-graph capture first " + "(it compiles its kernel on the first call)." + ) + import cutlass.cute as cute + + from . import k3_route_quant_kernel as kernel + + with _lock: + fn = _compiled.get(key) + if fn is None: + fn = _compiled[key] = cute.compile( + kernel.k3_route_quant, *args, num_tokens, scale, early_trigger, use_pdl, stream + ) + return fn + + +@torch.library.custom_op("trtllm::k3_route_quant", mutates_args=()) +def k3_route_quant( + scores: torch.Tensor, + bias: torch.Tensor, + hidden_states: torch.Tensor, + routed_scaling_factor: float, + early_trigger: bool = False, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Kimi K3 top-16 routing of ``scores`` (fp32 [M, 896]) with ``bias`` (fp32 [896]) and MXFP8 + quantization of ``hidden_states`` (bf16 [M, 3584]); M <= 64. + + Returns ``(topk_ids, topk_weights, quantized, scales)`` as ``kimi_k3_noaux_tc_mxfp8_quant``. + ``early_trigger`` lets the dependent grid launch once every CTA has passed its own dependency + wait (for dependents that wait for this whole grid before reading its outputs). + """ + import cuda.bindings.driver as cuda_driver + + _check(scores, bias, hidden_states) + num_tokens = scores.shape[0] + device = scores.device + topk_ids = torch.empty(num_tokens, TOP_K, dtype=torch.int32, device=device) + topk_weights = torch.empty(num_tokens, TOP_K, dtype=torch.bfloat16, device=device) + quantized = torch.empty(num_tokens, HIDDEN_SIZE, dtype=torch.float8_e4m3fn, device=device) + scales = torch.empty(num_tokens, HIDDEN_SIZE // SF_VEC_SIZE, dtype=torch.uint8, device=device) + args = ( + _arg(scores.view(-1)), + _arg(bias.view(-1)), + _arg(hidden_states.view(-1).view(torch.int32)), + _arg(topk_ids.view(-1)), + _arg(topk_weights.view(-1).view(torch.int16)), + _arg(quantized.view(-1).view(torch.int32)), + _arg(scales.view(-1)), + ) + stream = cuda_driver.CUstream(torch.cuda.current_stream(device).cuda_stream) + use_pdl = _use_pdl() + scale = float(routed_scaling_factor) + fn = _kernel(early_trigger, use_pdl, args, num_tokens, scale, stream) + # The compiled function takes the runtime arguments only (the Constexpr ones are baked in). + fn(*args, num_tokens, scale, stream) + return topk_ids, topk_weights, quantized, scales + + +@k3_route_quant.register_fake +def _(scores, bias, hidden_states, routed_scaling_factor, early_trigger=False): + num_tokens = scores.shape[0] + return ( + scores.new_empty((num_tokens, TOP_K), dtype=torch.int32), + scores.new_empty((num_tokens, TOP_K), dtype=torch.bfloat16), + hidden_states.new_empty((num_tokens, HIDDEN_SIZE), dtype=torch.float8_e4m3fn), + hidden_states.new_empty((num_tokens, HIDDEN_SIZE // SF_VEC_SIZE), dtype=torch.uint8), + ) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/__init__.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/__init__.py new file mode 100644 index 000000000000..1a30a934d4cd --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/__init__.py @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Kimi K3 collective sandwiches in CuTe DSL (``trtllm::k3_sandwich_oproj``). + +Importing :mod:`.op` registers the torch op; nothing is imported eagerly here so that +the CuTe DSL dependency stays optional for every other model. +""" diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/k3_sandwich_kernel.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/k3_sandwich_kernel.py new file mode 100644 index 000000000000..fb4488da2b36 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/k3_sandwich_kernel.py @@ -0,0 +1,2105 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Kimi K3 post-attention sandwich: the attention output projection, the TP all-reduce and the residual update +(attention-residual selection + RMSNorm) in one kernel, for M <= 8 decode tokens. + + partial_r = core_r @ W_o,r^T (this rank's [M, 7168] share, rounded to bf16) + updated = bf16(prefix + bf16(sum_r partial_r)) (or the sum alone without a prefix) + normed = KimiK3RMSNorm(attn_res(snapshots..., updated)) + +as ``o_proj`` followed by ``trtllm::mnnvl_allreduce_attn_res``; the reduction and the epilogue follow +``oneshotAllreduceAttnResKernel``'s arithmetic (rank chunks of 8, the same statistics, summation orders and +roundings), so the outputs are bit-identical to that pair. + +56 CTAs = 8 clusters of 7, 256 threads. Phase 1, every CTA: the o_proj rows [128 c, 128 c + 128) with the whole +weight slice (192 KB) TMA-loaded and copied on into TMEM (``tcgen05.cp``, 8 columns per K16 granule, the MMAs' own +descriptors) as each k-tile lands, all before ``griddepcontrol.wait`` under PDL; then the core [M, 768] and 48 tcgen05 +MMAs (M 128, N 8, fp32 in TMEM) with A read from TMEM (the tensor core reads A out of shared memory at ~110 B/clk). +The bf16 rows are pushed as 16-byte stores through the multicast mapping of the all-reduce buffer into slot +[token][rank] of every rank, and the dependents are launched. Phase 2, cluster t for token t < M: every thread owns 8 +columns (7 x 128 x 8 = 7168). The snapshots' statistics do not depend on the peers, so they are reduced over the +cluster and turned into the snapshots' logits and their maximum while the rows are in flight; then the thread polls +the ranks' slots until no word is empty, sums them, empties the slots it read, and runs the rest of the epilogue. Its +reductions over the cluster are st.async stores into every peer's shared memory that complete one-shot mailbox +mbarriers there (no cluster barrier): the snapshots' statistics (before the poll), the updated sum's statistics and +the selection's sum of squares. The mailbox waits spin on +``mbarrier.test_wait`` (a warp suspended in ``try_wait`` on a mailbox completed by remote st.async wakes late), +acquiring at cluster scope: the st.async complete_tx of the other CTAs releases at cluster scope. + +The buffer: two alternating halves of [8 tokens][world][7168] bf16 per rank, as int32 words; a word of +0x80000000 is empty (pushes turn bf16 -0.0 into +0.0, so a pushed pair never has that pattern). ``flags[b]`` +counts the calls of CTA b (its parity selects the half; every CTA runs every call, so all CTAs and ranks agree). +A reducing thread empties the words it read right after reading them: the next push into that half comes from a +peer's call after next, which starts only after this call has ended here. + +Published output (``x_slab``, when given): the normed rows also go to a Lamport slab [3][8][7168] bf16 (int32 +words, sentinel 0xFFFFFFFF, a computed all-ones word stored as 0x7FC07FC0) that the next kernel polls tile by tile +instead of waiting for this grid. Call ``slab_buf`` writes buffer ``slab_buf`` and, after its grid wait, re-arms +buffer ``(slab_buf + 1) % 3`` (every row, every CTA), whose last readers ran two grid completions earlier. + +The pre-attention sandwich (``k3_sandwich_tail``) runs the same phases with the row-parallel MoE tail as phase 1: +``[rmsnorm(latent)[:, lo:lo+224] | act] @ [W_lat (zero-padded to 256) | W_act]^T``, K = 256 + 384, the latent and +the activation k-tiles in two TMEM accumulators and the latent RMS (over the whole reduced latent row) applied to +the first; its phase 2 adds the MoE partial sum to the running prefix sum. The latent RMS of token t is computed once +per cluster, by epilogue warp t // 7 of CTA t % 7, from a bulk copy of the row into shared memory (the sums of a +per-lane chain in row order, then a butterfly), and st.async'd into every cluster CTA's row scales. + +Warps: 0 weight TMA, 1 the activation TMA (tail: first the bulk copies of the latent rows whose RMS the CTA owns), 2 +TMEM allocation, the weight copies into TMEM and the MMA, 3 idle, 4-7 the phase 1 epilogue and phase 2. After cluster +formation only warps 4-7 synchronize (named barrier 1 and the mailboxes), so warps 0-3 finish once their work is +issued. + +The CTA's work is ``sandwich_role`` (a ``cute.jit`` role function on shared memory carved from its ``smem`` +pointer, ``smem_bytes`` of it, e.g. one ``SmemAllocator`` block), so a kernel with further roles in other clusters +can run it in 56 of its CTAs; its ``act_hook`` replaces the activation TMA by the caller's own delivery into +``smem_b`` (see ``sandwich_role``). +``k3_sandwich_kernel`` is the sandwich alone. +""" + +from __future__ import annotations + +import cuda.bindings.driver as cuda_driver +import cutlass +import cutlass.cute as cute +import cutlass.experimental.cuda as cuda +from cutlass._mlir.dialects import llvm as _llvm +from cutlass._mlir.extras import types as _T +from cutlass.cutlass_dsl import dsl_user_op +from cutlass.experimental import primitives as prims + +try: + from cutlass.memory.smem import SmemAllocator +except ImportError: # DSL builds that keep it under cutlass.utils + from cutlass.utils import SmemAllocator + +H = 7168 +K_IN = 768 # this rank's o_proj input: 96 heads x 128 / TP16 +TAIL_LAT = 256 # the MoE tail's latent slice (224 columns at TP16) zero-padded to whole k-tiles +TAIL_ACT = 384 # the shared-expert activation: 2 x 3072 / TP16 / 2 +LATENT = 3584 +CTA_M = 128 +MMA_N = 8 +CTA_K = 128 +MMA_K = 16 +TMA_K_BOX = 64 +TMA_COPY_ITERS = CTA_K // TMA_K_BOX +K_BLOCKS_PER_HALF = TMA_K_BOX // MMA_K +MAX_K_TILES = K_IN // CTA_K +# The plain form takes up to 7 k-tiles (the drafter's down projection, K 896): all of them staged in TMEM (64 + 7 x 64 +# of 512 columns), but through a ring of W_RING shared-memory stages once they no longer all fit there. +W_RESIDENT = 6 +W_RING = 5 +PLAIN_MAX_K = 896 +THREADS = 256 +EPI_THREADS = 128 +EPI_WARPS = EPI_THREADS // 32 +ELTS = 8 # bf16 columns per phase 2 thread +ELEM_BYTES = 2 +TMEM_COLS = ( + 512 # the whole TMEM: accumulators at columns 0 (and 8), the weight slice from TMEM_A_COL +) +TMEM_A_COL = 64 # A (the weight slice) staged in TMEM: 8 columns (128 lanes x K16 bf16) per MMA +TMEM_A_COLS_PER_MMA = MMA_K * ELEM_BYTES // 4 +EVICT_FIRST = 0x12F0000000000000 + +CLUSTER = H // (EPI_THREADS * ELTS) # 7 CTAs own one token in phase 2 +NUM_CTAS = H // CTA_M # 56 +MAX_TOKENS = NUM_CTAS // CLUSTER # 8 +LATENT_ROWS_PER_CTA = ( + MAX_TOKENS + CLUSTER - 1 +) // CLUSTER # tail: cluster CTA c owns the RMS of tokens c and c + 7 +WORDS_PER_ROW = H // 2 +ROW_WORDS_PER_CTA = CTA_M // 2 +MAX_CANDIDATES = ( + 9 # snapshots + the updated sum; Kimi K3 (93 layers, a snapshot every 12) needs at most 9 +) +MAX_SNAPSHOTS = MAX_CANDIDATES - 1 +SNAP_STATS = 2 * MAX_SNAPSHOTS # [2 n] sum of squares, [2 n + 1] residual projection of snapshot n +FLAG_WORDS = 64 # per-CTA call counters (NUM_CTAS of them) +RANK_CHUNK = 8 # ranks summed per chunk, as the MNNVL one-shot does for more than 8 ranks +EMPTY_WORD = -2147483648 # 0x80000000 +LOG2E = 1.4426950408889634 +# The published slab (x_slab): [SLAB_BUFS][MAX_TOKENS][WORDS_PER_ROW] int32 words. +SLAB_BUFS = 3 +SLAB_WORDS = MAX_TOKENS * WORDS_PER_ROW +SLAB_SENTINEL = -1 # 0xFFFFFFFF +SLAB_CANON = 0x7FC07FC0 # what a computed all-ones word is stored as +POLL_BACKOFF = ( + 256 # cycles between two polling rounds of phase 1's input slab that found a sentinel +) +# The latent exchange (x_src 2): k3_moe pushes every rank's routed latent partial into two alternating halves of +# [8 tokens][world][3584] bf16 per rank (int32 words, EMPTY_WORD = not written); the tail sums them itself. +LAT_SLICE = 224 # latent columns of a rank's slice +LAT_SLICE_VECS = LAT_SLICE // 8 # its 16-byte vectors per token row +LAT_WORDS = 3584 // 2 # int32 words of one latent row +LAT_VECS = 3584 // 8 # 16-byte vectors of one latent row +LAT_FLAG_WORDS = ( + 64 # int32: [0] the tail's call count mod 6, [LAT_SCALES + 8 b + t] buffer b of the latent scale slab +) +LAT_SCALES = 32 +LAT_SCALE_BUFS = 3 +SCALE_SENTINEL = -1 # 0xFFFFFFFF: an unwritten scale (a computed rsqrt is never this NaN pattern) + +LEADING = 16 +STRIDE = 8 * TMA_K_BOX * ELEM_BYTES +A_HALF_ELEMS = CTA_M * TMA_K_BOX +B_HALF_ELEMS = MMA_N * TMA_K_BOX +STEP = (MMA_K * ELEM_BYTES) >> 4 +A_BOX = A_HALF_ELEMS >> 3 +B_BOX = B_HALF_ELEMS >> 3 +STAGE_A = (CTA_M * CTA_K * ELEM_BYTES) >> 4 +STAGE_B = (MMA_N * CTA_K * ELEM_BYTES) >> 4 + +io_dtype = cutlass.BFloat16 + +assert NUM_CTAS == MAX_TOKENS * CLUSTER +assert NUM_CTAS <= FLAG_WORDS +assert SNAP_STATS * CLUSTER <= EPI_THREADS +assert TMEM_A_COL + MAX_K_TILES * (CTA_K // MMA_K) * TMEM_A_COLS_PER_MMA <= TMEM_COLS +assert TMEM_A_COL + (PLAIN_MAX_K // CTA_K) * (CTA_K // MMA_K) * TMEM_A_COLS_PER_MMA <= TMEM_COLS + + +def buffer_words(world: int) -> int: + """Int32 words of one rank's all-reduce buffer: both halves of rows.""" + return 2 * MAX_TOKENS * world * WORDS_PER_ROW + + +def lat_buffer_words(world: int) -> int: + """Int32 words of one rank's latent exchange buffer (both halves).""" + return 2 * MAX_TOKENS * world * LAT_WORDS + + +@dsl_user_op +def _clock64(*, loc=None, ip=None): + return cutlass.Int64( + _llvm.inline_asm( + _T.i64(), [], "mov.u64 $0, %clock64;", "=l", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _pack_bf16x2(hi, lo, *, loc=None, ip=None): + """(bf16(hi) << 16) | bf16(lo), round to nearest even.""" + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [hi.ir_value(loc=loc, ip=ip), lo.ir_value(loc=loc, ip=ip)], + "cvt.rn.bf16x2.f32 $0, $1, $2;", "=r,f,f", has_side_effects=False, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _mapa_u32(smem_ptr, peer, *, loc=None, ip=None): + """The shared::cluster address of this CTA's shared-memory location in cluster CTA ``peer``.""" + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [smem_ptr.toint(loc=loc, ip=ip).ir_value(), cutlass.Int32(peer).ir_value(loc=loc, ip=ip)], + "mapa.shared::cluster.u32 $0, $1, $2;", "=r,r,r", has_side_effects=False, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _st_async_f32(dst, value, mbar, *, loc=None, ip=None): + """st.async of one fp32 to a shared::cluster address, completing ``mbar`` (a shared::cluster address) by 4 bytes.""" + _llvm.inline_asm( + None, + [cutlass.Int32(dst).ir_value(loc=loc, ip=ip), cutlass.Float32(value).ir_value(loc=loc, ip=ip), + cutlass.Int32(mbar).ir_value(loc=loc, ip=ip)], + "st.async.shared::cluster.mbarrier::complete_tx::bytes.f32 [$0], $1, [$2];", "r,f,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _st_async_v4(dst, w0, w1, w2, w3, mbar, *, loc=None, ip=None): + """st.async of four int32 words (16 bytes) to a shared::cluster address, completing ``mbar`` (a shared::cluster + address) by 16 bytes.""" + _llvm.inline_asm( + None, + [cutlass.Int32(dst).ir_value(loc=loc, ip=ip), cutlass.Int32(w0).ir_value(loc=loc, ip=ip), + cutlass.Int32(w1).ir_value(loc=loc, ip=ip), cutlass.Int32(w2).ir_value(loc=loc, ip=ip), + cutlass.Int32(w3).ir_value(loc=loc, ip=ip), cutlass.Int32(mbar).ir_value(loc=loc, ip=ip)], + "st.async.shared::cluster.mbarrier::complete_tx::bytes.v4.b32 [$0], {$1, $2, $3, $4}, [$5];", "r,r,r,r,r,r", + has_side_effects=True, is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _test_wait_cluster(mbar, parity, *, loc=None, ip=None): + """mbarrier.test_wait.parity.acquire.cluster on a barrier of this CTA (``mbar``, its shared-memory pointer): whether + phase ``parity`` has completed, acquiring at cluster scope. The barriers it is used on are completed by other CTAs + (st.async complete_tx, remote arrives), which release at cluster scope.""" + done = cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [mbar.toint(loc=loc, ip=ip).ir_value(), cutlass.Int32(parity).ir_value(loc=loc, ip=ip)], + "{\n\t.reg .pred p;\n\tmbarrier.test_wait.parity.acquire.cluster.shared::cta.b64 p, [$1], $2;\n\t" + "selp.u32 $0, 1, 0, p;\n\t}", "=r,r,r,~{memory}", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + return done != cutlass.Int32(0) + + +@dsl_user_op +def _rcp_rn(x, *, loc=None, ip=None): + """rcp.rn.f32: the correctly rounded 1 / x, the same value as the IEEE division div.rn.f32 1.0, x.""" + return cutlass.Float32( + _llvm.inline_asm( + _T.f32(), [cutlass.Float32(x).ir_value(loc=loc, ip=ip)], "rcp.rn.f32 $0, $1;", "=f,f", + has_side_effects=False, is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _ld_acquire_sys(addr_i64, *, loc=None, ip=None): + """ld.acquire.sys.global.u32: a flag word whose writer released its earlier writes at system scope.""" + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [addr_i64.ir_value(loc=loc, ip=ip)], "ld.acquire.sys.global.u32 $0, [$1];", "=r,l", + has_side_effects=True, is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _fence_sys(*, loc=None, ip=None): + """fence.acq_rel.sys: this thread's earlier writes (and those it has observed) before its later ones, for every + observer.""" + _llvm.inline_asm( + None, [], "fence.acq_rel.sys;", "", has_side_effects=True, is_align_stack=False, + asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +def _lo(word): + return (word << cutlass.Int32(16)).bitcast(cutlass.Float32) + + +def _hi(word): + return (word & cutlass.Int32(-65536)).bitcast(cutlass.Float32) + + +def _sanitize(word): + """A pushed pair must never read as empty: bf16 -0.0 halves become +0.0.""" + w = cutlass.Int32( + cutlass.select_( + (word & cutlass.Int32(0xFFFF)) == cutlass.Int32(0x8000), + word & cutlass.Int32(-65536), + word, + ) + ) + return cutlass.Int32( + cutlass.select_( + (w & cutlass.Int32(-65536)) == cutlass.Int32(EMPTY_WORD), w & cutlass.Int32(0xFFFF), w + ) + ) + + +def _slab_word(word): + """A published word must never read as the slab sentinel: an all-ones word (two all-ones bf16 NaNs) is stored + as two quiet NaNs.""" + return cutlass.Int32( + cutlass.select_(word == cutlass.Int32(SLAB_SENTINEL), cutlass.Int32(SLAB_CANON), word) + ) + + +def _fmax(a, b): + return cutlass.Float32(cutlass.select_(a > b, a, b)) + + +def _warp_allsum(value): + """Sum over the warp, in every lane. The butterfly adds the same pairs as a shfl.down tree does into lane 0 + (fp32 addition is commutative), so every lane holds lane 0's shfl.down result bit for bit.""" + for offset in (16, 8, 4, 2, 1): + value = value + cute.arch.shuffle_sync_bfly(value, offset=offset) + return value + + +def _warp_sum_scatter(values, lane): + """Sum each of len(values) (a power of two, 2..32) values over the warp; value i ends in the lanes l with + (l >> (5 - log2(len))) == i, reduced over the same pairs as a shfl.down tree (so bit-identical to it). + Each level keeps half of the values and sends the other half: len - 1 shuffles in all, then a plain + butterfly over the remaining lane bits.""" + vals = list(values) + offset = 16 + while len(vals) > 1: + half = len(vals) // 2 + upper = (lane & cutlass.Int32(offset)) != cutlass.Int32(0) + nxt = [] + for j in range(half): + keep = cutlass.Float32(cutlass.select_(upper, vals[j + half], vals[j])) + send = cutlass.Float32(cutlass.select_(upper, vals[j], vals[j + half])) + nxt.append(keep + cute.arch.shuffle_sync_bfly(send, offset=offset)) + vals = nxt + offset //= 2 + value = vals[0] + while offset >= 1: + value = value + cute.arch.shuffle_sync_bfly(value, offset=offset) + offset //= 2 + return value + + +def weight_stages(k_tiles: int) -> int: + """Shared-memory stages of the weight: every k-tile, or a ring of W_RING past W_RESIDENT k-tiles.""" + return k_tiles if k_tiles <= W_RESIDENT else W_RING + + +def _role_layout(k_tiles: int, rms_cols: int, swiglu: int = 0): + """(dtype, elements, alignment) of the role's shared arrays, in carve order.""" + stages = weight_stages(k_tiles) + return ( + ( + io_dtype, + stages * CTA_M * CTA_K, + 1024, + ), # smem_a: the weight k-tiles or the ring's stages (TMA, SW128) + ( + io_dtype, + k_tiles * MMA_N * CTA_K, + 1024, + ), # smem_b: the activation k-tiles (SW128, see act_word_index) + (cutlass.Int32, MMA_N * ROW_WORDS_PER_CTA, 16), # stage: the bf16 rows before the push + ( + cutlass.Int32, + LATENT_ROWS_PER_CTA * (rms_cols // 2) if rms_cols > 0 else 4, + 128, + ), # rms_row + (cutlass.Int64, stages, 8), # weight_full, per stage + (cutlass.Int64, 1, 8), # act_full + (cutlass.Int64, 1, 8), # acc_done + (cutlass.Int64, 1, 8), # mb_snap + (cutlass.Int64, 1, 8), # mb_upd + (cutlass.Int64, 1, 8), # mb_sq + (cutlass.Int64, 1, 8), # mb_rms + (cutlass.Int64, LATENT_ROWS_PER_CTA, 8), # rms_full + (cutlass.Int32, 1, 16), # the TMEM address + (cutlass.Float32, EPI_WARPS * SNAP_STATS, 16), # warp_snap + (cutlass.Float32, CLUSTER * SNAP_STATS, 16), # box_snap + (cutlass.Float32, CLUSTER * EPI_WARPS * 2, 16), # box_upd + (cutlass.Float32, CLUSTER * EPI_WARPS, 16), # box_sq + (cutlass.Float32, MMA_N, 16), # row_scale + (cutlass.Int64, 1, 8), # mb_ts (plain) + ( + cutlass.Float32, + CLUSTER * EPI_THREADS, + 16, + ), # box_ts (plain): every cluster thread's sum of squares + (cutlass.Int64, 1, 8), # lat_full (x_src 2): the latent slice by st.async + ( + io_dtype, + k_tiles * MMA_N * CTA_K if swiglu else 8, + 1024 if swiglu else 16, + ), # smem_g (swiglu): the gate tiles + (cutlass.Int64, stages, 8), # stage_free (ring): a stage's copies into TMEM have completed + (cutlass.Int64, 1, 8), # b_ready (swiglu): the activation rewritten in smem_b + ) + + +def smem_bytes(k_in: int, rms_cols: int, swiglu: int = 0) -> int: + """Bytes of shared memory ``sandwich_role`` carves from its ``smem`` pointer (1024-byte aligned): + ``smem_bytes(K_IN, 0)`` for the post-attention form, ``smem_bytes(TAIL_LAT + TAIL_ACT, LATENT)`` for the tail.""" + off = 0 + for dtype, n, align in _role_layout(k_in // CTA_K, rms_cols, swiglu): + off = -(-off // align) * align + n * dtype.width // 8 + return off + + +def _view(raw, off: int, dtype, n: int, align: int): + return cutlass.Array( + cute.recast_ptr(raw if off == 0 else raw + off, dtype=dtype), shape=(n,), dtype=dtype, bounds_check=False, + addrspace=cutlass.AddressSpace.smem.value, alignment=align, + ) # fmt: skip + + +def _carve(raw, k_tiles: int, rms_cols: int, swiglu: int = 0): + """The role's arrays as views into the shared-memory bytes at ``raw`` (1024-byte aligned), in ``_role_layout`` + order, then an int32 view of ``smem_b``. Every CTA of a cluster gets the same offsets, so a peer's address of a + mailbox is its mapa.""" + views = [] + off = 0 + b_off = 0 + for i, (dtype, n, align) in enumerate(_role_layout(k_tiles, rms_cols, swiglu)): + off = -(-off // align) * align + if i == 1: + b_off = off + views.append(_view(raw, off, dtype, n, align)) + off += n * dtype.width // 8 + views.append(_view(raw, b_off, cutlass.Int32, k_tiles * MMA_N * CTA_K // 2, 1024)) + return views + + +def act_word_index(t, col): + """Int32-word index in ``smem_b`` of the bf16 activation columns [col, col + 8) of token row t (col a multiple of + 8): the SW128 layout the activation TMA box lands, 16-byte unit j = (col % 64) // 8 of row t at unit j ^ t of + the row's 128 bytes, 64-column halves 1 KB apart and k-tiles of 128 columns 2 KB apart.""" + return (col // 128) * 512 + ((col % 128) // 64) * 256 + t * 32 + (((col % 64) // 8) ^ t) * 4 + + +@cute.jit +def swiglu_share(smem_b: cutlass.Array, smem_g: cutlass.Array, act_full, b_ready, tid: cutlass.Int32, + crank: cutlass.Int32): # fmt: skip + """The SwiGLU of k-tile ``crank``, this CTA's share of the cluster's 7 k-tiles: epilogue thread ``tid`` takes + 16-byte chunk ``tid`` of the tile, B = bf16((g * sigmoid(g)) * a) with silu_and_mul's fp32 order and sigmoid(g) + = 1 / (1 + exp(-g)) correctly rounded (rcp.rn, the value of the IEEE division), and st.async's the 16 bytes into + the same offsets of tile ``crank`` in every cluster CTA's ``smem_b`` (a and g share one swizzled layout), + completing their ``b_ready`` by 16 bytes each.""" + while not cute.arch.mbarrier_try_wait(act_full.data_ptr(), 0): + pass + c_el = crank * cutlass.Int32(MMA_N * CTA_K) + tid * cutlass.Int32(ELTS) + av = smem_b.load(idx=c_el, vector_size=ELTS, alignment=16) + gv = smem_g.load(idx=c_el, vector_size=ELTS, alignment=16) + b_words = [] + for q in cutlass.range_constexpr(ELTS // 2): + pair = [] + for e in (2 * q, 2 * q + 1): + g_e = cutlass.Float32(gv[e]) + sig = _rcp_rn(cutlass.Float32(1.0) + cute.math.exp(-g_e, fastmath=False)) + pair.append((g_e * sig) * cutlass.Float32(av[e])) + b_words.append(_pack_bf16x2(pair[1], pair[0])) + for peer in cutlass.range_constexpr(CLUSTER): + _st_async_v4( + _mapa_u32(smem_b.subview(c_el).data_ptr(), peer), b_words[0], b_words[1], b_words[2], b_words[3], + _mapa_u32(b_ready.data_ptr(), peer), + ) # fmt: skip + + +@cute.jit +def sandwich_role( + tma_desc_w, # W [7168, k_in] bf16, 5-D, one call per k-tile + tma_desc_x, # core [M, 768] (tail: latent [M, 3584]), box 8 x 64 (unused with act_hook) + tma_desc_x2, # tail: act [M, 384], box 8 x 64 + rms_src: cutlass.Array, # tail: int32 words of the latent rows whose RMS scales accumulator 0 + ws_uc: cutlass.Array, # int32 words: this rank's all-reduce buffer + ws_mc: cutlass.Array, # int32 words: its multicast mapping + ws_flags: cutlass.Array, # int32 [FLAG_WORDS]: call count of CTA b at [b] + prefix: cutlass.Array, # int32 words of bf16 [M, 7168] (read only with add_prefix) + snapshots: cutlass.Array, # int32 words of bf16 [num_cand - 1, M, 7168] (at least one [M, 7168] row) + res_w: cutlass.Array, # int32 words of bf16 [7168]: the residual projection + rms_w: cutlass.Array, # int32 words of bf16 [7168]: its RMSNorm weight + out_w: cutlass.Array, # int32 words of bf16 [7168]: the output RMSNorm weight + updated: cutlass.Array, # int32 words of bf16 [M, 7168], out + normed: cutlass.Array, # int32 words of bf16 [M, 7168], out + x_slab: cutlass.Array, # int32 words [SLAB_BUFS, MAX_TOKENS, 3584]: the published normed rows (publish) + lat_flags: cutlass.Array, # int32 [LAT_FLAG_WORDS] (x_src 2): the tail's call count and the latent scale slab + tap: cutlass.Array, # int32 words of bf16 rows (tap_out), tap_stride words apart: the capture layer's tap, out + smem, # shared-memory pointer, 1024-byte aligned, at least smem_bytes(k_in, rms_cols) bytes (SmemAllocator) + num_tokens: cutlass.Int32, + rank: cutlass.Int32, + num_cand: cutlass.Int32, # snapshots + 1, at most MAX_CANDIDATES + add_prefix: cutlass.Int32, # nonzero: updated = prefix + the sum + rms_eps: cutlass.Float32, + out_eps: cutlass.Float32, + x_col0: cutlass.Int32, # tail: first latent column of this rank's slice + lat_eps: cutlass.Float32, # tail: the latent RMSNorm's epsilon + slab_buf: cutlass.Int32, # publish: the slab buffer this call writes (0-2) + x_buf: cutlass.Int32, # x_src 1: the slab buffer of phase 1's input (0-2) + tap_stride: cutlass.Int32, # tap_out: int32 words from one tap row to the next (3584 when contiguous) + world: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], # K of phase 1: 768 (o_proj) or 640 (tail) + x_tiles: cutlass.Constexpr[ + int + ], # k-tiles from tma_desc_x; the rest from tma_desc_x2 into accumulator 1 + rms_cols: cutlass.Constexpr[ + int + ], # > 0 (tail): y = rsqrt(mean(rms_src row^2) + lat_eps) * acc0 + acc1 + publish: cutlass.Constexpr[ + int + ], # 1: also publish normed into x_slab (and re-arm the next buffer) + x_src: cutlass.Constexpr[ + int + ], # 0: phase 1's input by TMA after the grid wait; 1: polled from a slab; 2 (tail): + # the latent summed from every rank's pushed partial (below) + plain: cutlass.Constexpr[ + int + ], # 1 (MNNVL order) / 2 (IPC order): residual + RMSNorm, not attn_res (below) + tap_out: cutlass.Constexpr[ + int + ], # attn_res forms: also store into ``tap`` the pre-norm mixture rows (1) or updated (2) + swiglu: cutlass.Constexpr[ + int + ], # plain: phase 1's B is silu_and_mul of x [M, 2 k_in] (gate columns first), below + k_split: cutlass.Constexpr[ + int + ], # 2: accumulate the even and the odd k-tiles apart, then (0 + even) + odd + cta_base: cutlass.Constexpr[int], # grid index of the role's first CTA (a multiple of CLUSTER) + act_hook: cutlass.Constexpr, # None: the activation by TMA after the grid wait; else see below +): + """One CTA of the sandwich: grid indices cta_base .. cta_base + 55 (8 clusters of 7, 256 threads). A kernel may + run other roles in other CTAs (their own clusters) as long as they never touch ``ws_flags`` or the all-reduce + buffer. + + ``act_hook`` (post-attention form only): warp 1 calls ``act_hook(smem_b_words, act_full, grid_bx, lane)`` (all 32 + lanes; grid_bx the CTA's grid index) instead of waiting on the grid and loading the activation by TMA. The hook + must store every token row's 768 columns into ``smem_b_words`` (int32 view of the bf16 tile) at + ``act_word_index`` (rows >= M zero), then every storing lane ``fence.proxy.async.shared::cta``, a warp sync, and + one plain arrive of lane 0 on ``act_full`` (count 1, no transaction bytes). Warps 4-7 still wait on the grid + before the flags, the prefix and the snapshots. + + ``x_src`` 1 (the producer publishes phase 1's input as a Lamport slab, int32 [3][8][x_cols / 2] in ``rms_src``, + sentinel 0xFFFFFFFF, buffer ``x_buf``): warp 1 polls this rank's k-tiles of the live rows into ``smem_b`` (tail: + and the whole rows of the CTA's RMS tokens into ``rms_row``; the shared-expert activation by TMA at once), and no + warp waits on the grid before its reads. That needs the producer to launch its dependents only after its own + grid wait: then this launch implies that every kernel before the producer has completed, which covers the call + count, the prefix, the snapshots and the slab buffer this call re-arms (read L1-bypassing). Warps 4-7 wait on the + grid at their end, so the grid completes only after its producer. + + ``x_src`` 2 (tail; the latent all-reduce folded in): k3_moe, the producer, pushes its routed latent partial rows + through the multicast mapping into slot [rank] of half ``n & 1`` of every rank's exchange buffer (``rms_src``: int32 + [2][8][world][1792], EMPTY_WORD = not written, pushes never write that pattern) and exits; ``n = lat_flags[0]`` + counts this op's calls modulo 6. No warp waits on the grid before its reads (as for ``x_src`` 1, the producer + launches its dependents right after its own grid wait). The epilogue warps first empty the other half (every row: + its readers, the previous call, have completed, and the fence before this call's push orders the empties before any + rank's next push into it, which follows that rank's phase 2 of this call), then sum the partials in the MNNVL + one-shot's order (fp32 per chunk of 8 ranks in rank order, the chunks added from 0, one bf16 rounding), so the + latent is the all-reduce's bit for bit: + - every cluster: the rank's 224 slice columns of the live rows, one (token, vector, rank chunk) per thread, st.async + into all 7 CTAs' ``smem_b`` (``lat_full``; the latent k-tiles are zeroed by warp 3 before the cluster forms); + - cluster t: token t's whole row, st.async into CTA 0's ``rms_row``; CTA 0's warp 4 takes its RMS as the bulk-copy + path does and publishes the scale into buffer ``n % 3`` of the scale slab in ``lat_flags``, which every + epilogue polls after the MMA (each call re-arms buffer ``(n + 1) % 3``). + The MMA runs the shared-expert k-tiles first and the latent ones once ``lat_full`` completes. CTA 0 stores + ``n + 1`` after its final grid wait. Every k3_moe push call must be followed by exactly one such call. + + ``plain`` 1 (post-attention form, no snapshots): the epilogue is the MNNVL one-shot's residual + RMSNorm + (``kARResidualRMSNorm``) with its arithmetic: updated = bf16(float(sum) + float(prefix)); per thread the sum of + the bf16 squares of its 8 values in order; that kernel's reduction tree for a 7168-wide row (8 CTAs of 112 + threads: xor-butterfly warps, the partial fourth warp, warp 0's butterfly over the warp sums, then the block sums + in order), replayed on every thread's sum gathered into each cluster CTA; normed = bf16(float(x) * r * float(w)) + with r = rsqrt(S / 7168 + out_eps). ``plain`` 2: the same in the order of the IPC one-shot + (``allreduce_fusion_kernel_oneshot_lamport``, kARResidualRMSNorm, fp32 accumulation), which TRT-LLM runs within one + node: the ranks' sum in rank order (world <= 8: one chunk), per thread the fp32 squares of its 8 values in order, + and that kernel's tree for a 7168-wide row (4 CTAs of 224 threads: xor-butterfly warps, blockReduceSumV2's + butterfly over the 7 warp sums, then the CTA sums in order); thread g's 8 columns are [8 g, 8 g + 8) in both. + + ``swiglu`` (plain, the drafter MLP's down projection after gate_up, K 896 = one k-tile per cluster CTA): warp 1 + lands the up half of k-tile ``crank`` of x [M, 2 k_in] in ``smem_b`` and its gate half in ``smem_g`` (identical + SW128 layouts); after the grid wait the epilogue warps compute that tile's B = bf16((g * sigmoid(g)) * a) with + k3_ctm_gemv_swiglu's arithmetic (``swiglu_share``) and st.async it into every cluster CTA's ``smem_b``, whose + ``b_ready`` expects the 7 tiles' bytes; warp 2 then fences the generic-proxy stores for the MMA's async-proxy + reads. ``k_split`` 2 accumulates the even and + the odd k-tiles in two TMEM accumulators and adds them as (0 + even) + odd before the one bf16 rounding: the sums + of k3_ctm_gemv split 2, whose rank r takes the k-tiles r, r + 2, .... Past W_RESIDENT k-tiles the weight goes + through a ring of W_RING shared-memory stages: warp 2 commits each stage's copies into TMEM to ``stage_free``, + which warp 0 waits on before it refills the stage.""" + tail = rms_cols > 0 + hooked = act_hook is not None + k_tiles = k_in // CTA_K + split_acc = x_tiles < k_tiles + stages = weight_stages(k_tiles) + ring = stages < k_tiles + acc_pair = split_acc or k_split == 2 + assert not (ring and (x_src != 0 or hooked)), "the weight ring serves the plain form only" + assert not swiglu or (plain and k_split == 2 and k_tiles == CLUSTER), ( + "swiglu: the plain form, split 2's order, K 896" + ) + tx, _, _ = cute.arch.thread_idx() + grid_bx, _, _ = cute.arch.block_idx() + bx = grid_bx - cutlass.Int32(cta_base) # the role's CTA index (0-55) + warp_id = cute.arch.warp_idx() + crank = cute.arch.block_idx_in_cluster() + token = bx // cutlass.Int32(CLUSTER) + tma_ptr_w = tma_desc_w.get_ptr() + tma_ptr_x = tma_desc_x.get_ptr() + m_offset = bx * cutlass.Int32(CTA_M) + + tma_ptr_x2 = tma_desc_x2.get_ptr() + # Views into the role's shared memory (_role_layout order). One-shot mailboxes: the cluster's statistics + # (snapshots CTA-summed, the updated sum per warp), the per-warp output sums of squares and (tail) the tokens' + # latent RMS land by st.async in box_snap [source CTA rank][statistic], box_upd [source CTA rank][source warp][sum + # of squares, residual projection], box_sq [source CTA rank][source warp] and row_scale, completing mb_snap, + # mb_upd, mb_sq and mb_rms. warp_snap [epilogue warp][snapshot statistic] holds this CTA's warp sums before the + # CTA sum is sent. Tail: rms_row holds the latent rows of the tokens whose RMS this CTA computes (bulk copies). + (smem_a, smem_b, stage, rms_row, weight_full, act_full, acc_done, mb_snap, mb_upd, mb_sq, mb_rms, rms_full, + tmem_ptr_i32, warp_snap, box_snap, box_upd, box_sq, row_scale, mb_ts, box_ts, lat_full, smem_g, + stage_free, b_ready, smem_b_words) = _carve(smem, k_tiles, rms_cols, swiglu) # fmt: skip + + if warp_id == 0: + prims.prefetch_tensormap(tma_ptr_w) + if cutlass.const_expr(not hooked and x_src == 0): + prims.prefetch_tensormap(tma_ptr_x) + if cutlass.const_expr(split_acc): + prims.prefetch_tensormap(tma_ptr_x2) + if prims.elect_sync(): + for k in cutlass.range_constexpr(stages): + prims.mbarrier_init(weight_full.subview(k), 1) + if cutlass.const_expr(ring): + prims.mbarrier_init(stage_free.subview(k), 1) + if cutlass.const_expr(swiglu): + # Armed before cluster formation: every cluster CTA's share of B, by st.async. + prims.mbarrier_init(b_ready, 1) + prims.mbarrier_arrive_expect_tx(b_ready, CLUSTER * EPI_THREADS * 16) + # x_src 1 with a split accumulator: the activation TMA's arrive and warp 1's arrive after the slab rows. + prims.mbarrier_init(act_full, 2 if (x_src == 1 and split_acc) else 1) + prims.mbarrier_init(acc_done, 1) + prims.mbarrier_init(mb_snap, 1) + prims.mbarrier_init(mb_upd, 1) + prims.mbarrier_init(mb_sq, 1) + # Armed before cluster formation, so no peer's bytes can arrive before the expectation. Each peer sends + # its CTA sums of the num_cand - 1 snapshots' two statistics (none: the phase completes here) and, per + # warp, the updated sum's two and the selection's sum of squares. + prims.mbarrier_arrive_expect_tx( + mb_snap, cutlass.Int32(CLUSTER * 4 * 2) * (num_cand - cutlass.Int32(1)) + ) + prims.mbarrier_arrive_expect_tx(mb_upd, CLUSTER * EPI_WARPS * 2 * 4) + prims.mbarrier_arrive_expect_tx(mb_sq, CLUSTER * EPI_WARPS * 4) + if cutlass.const_expr(plain): + # Every cluster thread's sum of squares, from every cluster CTA. + prims.mbarrier_init(mb_ts, 1) + prims.mbarrier_arrive_expect_tx(mb_ts, CLUSTER * EPI_THREADS * 4) + if cutlass.const_expr(tail): + # The latent RMS of each live token, from the one cluster CTA that computes it (x_src 2: from the + # scale slab instead). + prims.mbarrier_init(mb_rms, 1) + if cutlass.const_expr(x_src != 2): + prims.mbarrier_arrive_expect_tx(mb_rms, num_tokens * cutlass.Int32(4)) + for rj in cutlass.range_constexpr(LATENT_ROWS_PER_CTA): + prims.mbarrier_init(rms_full.subview(rj), 1) + if cutlass.const_expr(x_src == 2): + # The live rows' slice vectors from the cluster's reducers, and (CTA 0 of a live token's cluster) the + # token's whole summed row. + prims.mbarrier_init(lat_full, 1) + prims.mbarrier_arrive_expect_tx( + lat_full, num_tokens * cutlass.Int32(LAT_SLICE_VECS * 16) + ) + if (crank == cutlass.Int32(0)) & (token < num_tokens): + prims.mbarrier_arrive_expect_tx( + rms_full.subview(0), cutlass.Int32(LAT_VECS * 16) + ) + + if warp_id == 2: + prims.tcgen05_alloc(tmem_ptr_i32, TMEM_COLS) + prims.tcgen05_relinquish_alloc_permit() + if cutlass.const_expr(x_src == 2): + if warp_id == 3: + # Zero the latent k-tiles of smem_b (rows >= M and the 32 padding columns stay zero; the reducers' st.async + # land after cluster formation), for the MMA's async-proxy reads and before any peer's store. + lane3 = tx % cutlass.Int32(32) + zero = cutlass.Int32(0) + for zi in cutlass.range_constexpr(x_tiles * MMA_N * CTA_K // (2 * 4 * 32)): + smem_b_words.store((zero, zero, zero, zero), idx=cutlass.Int32(zi * 128) + lane3 * cutlass.Int32(4), + alignment=16) # fmt: skip + prims.fence_proxy("async_shared", space=prims.SharedSpace.shared_cta) + prims.fence_acq_rel(prims.MemScope.CLUSTER) + prims.fence_mbarrier_init() + # Cluster formation: the peers' shared memory and armed mailboxes are addressable from here on. + prims.barrier_cluster_arrive_relaxed() + prims.barrier_cluster_wait() + prims.barrier_cta_sync(0) + tmem_ptr = cutlass.inttoptr(tmem_ptr_i32.load(), 6, cutlass.Int32) + + if warp_id == 0: + if prims.elect_sync(): + for k in cutlass.range_constexpr(k_tiles): + ws_ = k % stages + if cutlass.const_expr(k >= stages): + # The ring: this stage's previous k-tile has gone on into TMEM (warp 2's commit). + while not cute.arch.mbarrier_test_wait( + stage_free.subview(ws_).data_ptr(), (k // stages - 1) % 2 + ): + pass + prims.mbarrier_arrive_expect_tx( + weight_full.subview(ws_), CTA_M * CTA_K * ELEM_BYTES + ) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a.subview(ws_ * CTA_M * CTA_K), + tma_ptr_w, + ( + cutlass.Int32(0), + m_offset, + cutlass.Int32(k * TMA_COPY_ITERS), + cutlass.Int32(0), + cutlass.Int32(0), + ), + weight_full.subview(ws_), + l2_cache_hint=EVICT_FIRST, + ) + elif warp_id == 1: + if cutlass.const_expr(x_src == 1): + # Phase 1's input from the producer's Lamport slab: a 16-byte vector is ready when none of its words is the + # sentinel (relaxed gpu-scope loads, no sleep). + lane1 = tx % cutlass.Int32(32) + sent1 = cutlass.Int32(SLAB_SENTINEL) + x_words = (rms_cols if tail else k_in) // 2 + x_row0 = x_buf * cutlass.Int32(MAX_TOKENS * x_words) + if cutlass.const_expr(split_acc): + # The shared-expert activation (accumulator 1's k-tiles) is complete at launch. + if prims.elect_sync(): + prims.mbarrier_arrive_expect_tx( + act_full, (k_tiles - x_tiles) * MMA_N * CTA_K * ELEM_BYTES + ) + for k in cutlass.range_constexpr(x_tiles, k_tiles): + for half in cutlass.range_constexpr(TMA_COPY_ITERS): + prims.cp_async_bulk_tensor_shared_cta_global( + smem_b.subview(k * MMA_N * CTA_K + half * B_HALF_ELEMS), + tma_ptr_x2, + ( + cutlass.Int32((k - x_tiles) * CTA_K + half * TMA_K_BOX), + cutlass.Int32(0), + ), + act_full, + ) + # Poll once this CTA's weights are in (the MMA needs both; earlier polls only add L2 traffic beside the + # producer). Each round loads every vector of the lane (all loads in flight at once), stores them, and + # repeats after a short back-off while any live one still holds the sentinel. + for k in cutlass.range_constexpr(k_tiles): + while not cute.arch.mbarrier_try_wait(weight_full.subview(k).data_ptr(), 0): + pass + # This rank's x_tiles k-tiles of the live rows (columns x_col0 + [0, 128 x_tiles)); rows >= M zero. + x_vecs = x_tiles * CTA_K // 8 + x_col_word = x_col0 // cutlass.Int32(2) + x_pending = cutlass.Boolean(True) + while x_pending: + x_pending = cutlass.Boolean(False) + xv_in = [] + for i in cutlass.range_constexpr(MAX_TOKENS * x_vecs // 32): + xi = cutlass.Int32(i * 32) + lane1 + x_idx = ( + x_row0 + + (xi // cutlass.Int32(x_vecs)) * cutlass.Int32(x_words) + + x_col_word + + (xi % cutlass.Int32(x_vecs)) * cutlass.Int32(4) + ) + xv_in.append( + prims.load_ext( + rms_src.subview(x_idx), + dtype=cutlass.Int32, + count=4, + order="relaxed", + scope="gpu", + ) + ) + for i in cutlass.range_constexpr(MAX_TOKENS * x_vecs // 32): + xi = cutlass.Int32(i * 32) + lane1 + xt = xi // cutlass.Int32(x_vecs) + live = xt < num_tokens + xw = [ + cutlass.Int32( + cutlass.select_(live, cutlass.Int32(xv_in[i][q]), cutlass.Int32(0)) + ) + for q in range(4) + ] + smem_b_words.store( + (xw[0], xw[1], xw[2], xw[3]), + idx=act_word_index(xt, (xi % cutlass.Int32(x_vecs)) * cutlass.Int32(8)), + alignment=16, + ) + x_pending = x_pending | ( + (xw[0] == sent1) | (xw[1] == sent1) | (xw[2] == sent1) | (xw[3] == sent1) + ) + if x_pending: + t_back = _clock64() + while _clock64() - t_back < cutlass.Int64(POLL_BACKOFF): + pass + prims.fence_proxy("async_shared", space=prims.SharedSpace.shared_cta) + cute.arch.sync_warp() + if lane1 == cutlass.Int32(0): + prims.mbarrier_arrive(act_full) + if cutlass.const_expr(tail): + # The whole latent row of each token this CTA takes the RMS of (the last of the producer's tiles), in + # the same rounds. + for rj in cutlass.range_constexpr(LATENT_ROWS_PER_CTA): + rt_s = crank + cutlass.Int32(rj * CLUSTER) + if rt_s < num_tokens: + r_pending = cutlass.Boolean(True) + while r_pending: + r_pending = cutlass.Boolean(False) + rv_in = [] + for i in cutlass.range_constexpr(x_words // (4 * 32)): + r_idx = ( + x_row0 + rt_s * cutlass.Int32(x_words) + cutlass.Int32(i * 32 * 4) + + lane1 * cutlass.Int32(4) + ) # fmt: skip + rv_in.append( + prims.load_ext( + rms_src.subview(r_idx), + dtype=cutlass.Int32, + count=4, + order="relaxed", + scope="gpu", + ) # fmt: skip + ) + for i in cutlass.range_constexpr(x_words // (4 * 32)): + rw = [cutlass.Int32(rv_in[i][q]) for q in range(4)] + rms_row.store( + (rw[0], rw[1], rw[2], rw[3]), + idx=cutlass.Int32(rj * x_words + i * 32 * 4) + + lane1 * cutlass.Int32(4), + alignment=16, + ) + r_pending = r_pending | ( + (rw[0] == sent1) + | (rw[1] == sent1) + | (rw[2] == sent1) + | (rw[3] == sent1) + ) + if r_pending: + t_rback = _clock64() + while _clock64() - t_rback < cutlass.Int64(POLL_BACKOFF): + pass + cute.arch.sync_warp() + if lane1 == cutlass.Int32(0): + prims.mbarrier_arrive(rms_full.subview(rj)) + elif cutlass.const_expr(x_src == 2): + # The shared-expert activation (accumulator 1's k-tiles) is complete at launch; the latent k-tiles come + # from the epilogue warps' reducers (lat_full). + if prims.elect_sync(): + prims.mbarrier_arrive_expect_tx( + act_full, (k_tiles - x_tiles) * MMA_N * CTA_K * ELEM_BYTES + ) + for k in cutlass.range_constexpr(x_tiles, k_tiles): + for half in cutlass.range_constexpr(TMA_COPY_ITERS): + prims.cp_async_bulk_tensor_shared_cta_global( + smem_b.subview(k * MMA_N * CTA_K + half * B_HALF_ELEMS), + tma_ptr_x2, + ( + cutlass.Int32((k - x_tiles) * CTA_K + half * TMA_K_BOX), + cutlass.Int32(0), + ), + act_full, + ) + elif cutlass.const_expr(hooked): + act_hook(smem_b_words, act_full, grid_bx, tx % cutlass.Int32(32)) + else: + prims.griddepcontrol(prims.GridDepAction.WAIT) + if cutlass.const_expr(tail and x_src == 0): + # The whole latent row of each token this CTA takes the RMS of, one bulk copy each (the TMA engine, not + # the epilogue warps' load path, which their phase 2 operand loads fill at the same time). + if prims.elect_sync(): + for rj in cutlass.range_constexpr(LATENT_ROWS_PER_CTA): + rt_w1 = crank + cutlass.Int32(rj * CLUSTER) + if rt_w1 < num_tokens: + prims.mbarrier_arrive_expect_tx(rms_full.subview(rj), rms_cols * ELEM_BYTES) + prims.cp_async_bulk_shared_cluster_global( + rms_row.subview(rj * (rms_cols // 2)), + rms_src.subview(rt_w1 * cutlass.Int32(rms_cols // 2)), + rms_full.subview(rj), + rms_cols * ELEM_BYTES, + ) + if cutlass.const_expr(not hooked and x_src == 0): + if prims.elect_sync(): + if cutlass.const_expr(swiglu): + # This CTA's k-tile only: the up half (from column x_col0 = k_in) into smem_b, the gate half into + # smem_g (the cluster CTAs share the SwiGLU tile by tile). + prims.mbarrier_arrive_expect_tx(act_full, 2 * MMA_N * CTA_K * ELEM_BYTES) + for half in cutlass.range_constexpr(TMA_COPY_ITERS): + prims.cp_async_bulk_tensor_shared_cta_global( + smem_b.subview( + crank * cutlass.Int32(MMA_N * CTA_K) + + cutlass.Int32(half * B_HALF_ELEMS) + ), + tma_ptr_x, + ( + x_col0 + + crank * cutlass.Int32(CTA_K) + + cutlass.Int32(half * TMA_K_BOX), + cutlass.Int32(0), + ), + act_full, + ) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_g.subview( + crank * cutlass.Int32(MMA_N * CTA_K) + + cutlass.Int32(half * B_HALF_ELEMS) + ), + tma_ptr_x, + ( + crank * cutlass.Int32(CTA_K) + cutlass.Int32(half * TMA_K_BOX), + cutlass.Int32(0), + ), + act_full, + ) + else: + prims.mbarrier_arrive_expect_tx(act_full, k_tiles * MMA_N * CTA_K * ELEM_BYTES) + for k in cutlass.range_constexpr(0 if swiglu else k_tiles): + for half in cutlass.range_constexpr(TMA_COPY_ITERS): + if cutlass.const_expr(k < x_tiles): + prims.cp_async_bulk_tensor_shared_cta_global( + smem_b.subview(k * MMA_N * CTA_K + half * B_HALF_ELEMS), + tma_ptr_x, + ( + x_col0 + cutlass.Int32(k * CTA_K + half * TMA_K_BOX), + cutlass.Int32(0), + ), + act_full, + ) + else: + prims.cp_async_bulk_tensor_shared_cta_global( + smem_b.subview(k * MMA_N * CTA_K + half * B_HALF_ELEMS), + tma_ptr_x2, + ( + cutlass.Int32((k - x_tiles) * CTA_K + half * TMA_K_BOX), + cutlass.Int32(0), + ), + act_full, + ) + elif warp_id == 2: + idesc = prims.Tcgen05InstrDesc.build( + c_dtype=cutlass.Float32, a_dtype=io_dtype, b_dtype=io_dtype, n_dim=MMA_N, m_dim=CTA_M + ) + desc_a_base = prims.Tcgen05SmemDesc.build( + start_address=smem_a, leading_byte_offset=LEADING, stride_byte_offset=STRIDE, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + desc_b_base = prims.Tcgen05SmemDesc.build( + start_address=smem_b, leading_byte_offset=LEADING, stride_byte_offset=STRIDE, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + tmem_raw = tmem_ptr_i32.load() + tmem_acc1 = tmem_ptr + if cutlass.const_expr(acc_pair): + tmem_acc1 = cutlass.inttoptr(tmem_raw + cutlass.Int32(MMA_N), 6, cutlass.Int32) + # Before the grid wait (under PDL): each landed weight k-tile goes on into TMEM, one 128 x K16 granule per copy + # with the MMA's own descriptor; this thread issues the MMAs that read those columns later, in order. (The ring: + # stage k % stages, its fill k // stages; the copies of a stage that is refilled later are committed to + # stage_free, which warp 0 waits on.) + for k in cutlass.range_constexpr(k_tiles): + ws2 = k % stages + while not cute.arch.mbarrier_try_wait( + weight_full.subview(ws2).data_ptr(), (k // stages) % 2 + ): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + if prims.elect_sync(): + for kb in cutlass.range_constexpr(TMA_COPY_ITERS * K_BLOCKS_PER_HALF): + box = kb // K_BLOCKS_PER_HALF + within = kb % K_BLOCKS_PER_HALF + prims.tcgen05_cp( + prims.Tcgen05CpShape.SHAPE_128X256B, + cutlass.inttoptr( + tmem_raw + + cutlass.Int32( + TMEM_A_COL + + (k * TMA_COPY_ITERS * K_BLOCKS_PER_HALF + kb) + * TMEM_A_COLS_PER_MMA + ), + 6, + cutlass.Int32, + ), + desc_a_base + (ws2 * STAGE_A + box * A_BOX + within * STEP), + group=prims.CTAGroup.CTA_1, + ) + if cutlass.const_expr(ring and k + stages < k_tiles): + prims.tcgen05_commit(stage_free.subview(ws2)) + if cutlass.const_expr(swiglu): + # Completed by the cluster CTAs' st.async of B (test_wait: a suspended try_wait wakes late on those); their + # generic-proxy stores are then made visible to the MMA's async-proxy reads. + while not _test_wait_cluster(b_ready.data_ptr(), 0): + pass + prims.fence_proxy("async_shared", space=prims.SharedSpace.shared_cta) + elif cutlass.const_expr(hooked or x_src == 1): + # Completed by a thread's arrive (not TMA bytes): a suspended try_wait wakes up to µs late on those. + while not cute.arch.mbarrier_test_wait(act_full.data_ptr(), 0): + pass + else: + while not cute.arch.mbarrier_try_wait(act_full.data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + # x_src 2: the shared-expert k-tiles first; the latent ones once the reducers' slice vectors are in. + mma_order = ( + list(range(x_tiles, k_tiles)) + list(range(x_tiles)) + if x_src == 2 + else list(range(k_tiles)) + ) + for ki in cutlass.range_constexpr(k_tiles): + k = mma_order[ki] + second = (split_acc and k >= x_tiles) or (k_split == 2 and k % 2 == 1) + first_tile = k == 0 or (split_acc and k == x_tiles) or (k_split == 2 and k == 1) + if cutlass.const_expr(x_src == 2 and k == 0): + # Completed by the peers' st.async (test_wait: a suspended try_wait wakes late on those); their + # generic-proxy stores are then made visible to the MMA's async-proxy reads. + while not _test_wait_cluster(lat_full.data_ptr(), 0): + pass + prims.fence_proxy("async_shared", space=prims.SharedSpace.shared_cta) + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + for kb in cutlass.range_constexpr(TMA_COPY_ITERS * K_BLOCKS_PER_HALF): + box = kb // K_BLOCKS_PER_HALF + within = kb % K_BLOCKS_PER_HALF + a_staged = cutlass.inttoptr( + tmem_raw + + cutlass.Int32( + TMEM_A_COL + + (k * TMA_COPY_ITERS * K_BLOCKS_PER_HALF + kb) * TMEM_A_COLS_PER_MMA + ), + 6, + cutlass.Int32, + ) + desc_b = desc_b_base + (k * STAGE_B + box * B_BOX + within * STEP) + if prims.elect_sync(): + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, prims.CTAGroup.CTA_1, tmem_acc1 if second else tmem_ptr, + a_staged, desc_b, idesc, not (first_tile and kb == 0), + ) # fmt: skip + if prims.elect_sync(): + prims.tcgen05_commit(acc_done) + + if warp_id >= 4: + tid = tx - cutlass.Int32(EPI_THREADS) + lane = tx % 32 + w = warp_id - 4 + # The flag word, the prefix and the snapshots are written upstream (x_src 1: see the docstring). + if cutlass.const_expr(x_src == 0): + prims.griddepcontrol(prims.GridDepAction.WAIT) + if cutlass.const_expr(swiglu): + swiglu_share(smem_b, smem_g, act_full, b_ready, tid, crank) + # This CTA's call count. The previous call of this CTA has completed (every kernel between two calls waits + # for its predecessor), and the next one reads it only after this call has completed. + calls = ws_flags.load(idx=bx, is_volatile=True) + cur = calls & cutlass.Int32(1) + + col_word = (crank * cutlass.Int32(EPI_THREADS) + tid) * cutlass.Int32(ELTS // 2) + if cutlass.const_expr(publish): + # Re-arm the slab buffer after next: its last readers (the consumer of the call two back) completed + # before this grid's wait returned (x_src 1: before this grid launched). Every thread empties its own 16 + # bytes of its cluster's token row, live token or not, so every row of the buffer is re-armed whatever M + # the next calls run. + rearm = slab_buf + cutlass.Int32(1) + rearm = cutlass.Int32( + cutlass.select_(rearm == cutlass.Int32(SLAB_BUFS), cutlass.Int32(0), rearm) + ) + sent = cutlass.Int32(SLAB_SENTINEL) + x_slab.store( + (sent, sent, sent, sent), + idx=rearm * cutlass.Int32(SLAB_WORDS) + + token * cutlass.Int32(WORDS_PER_ROW) + + col_word, + alignment=16, + ) + + # ---- Phase 2 operands, loaded while the MMA runs (cluster `token` reduces token `token`). Clusters past + # the last token load a valid row and discard it. + reducer = token < num_tokens + row_tok = cutlass.Int32(cutlass.select_(reducer, token, num_tokens - cutlass.Int32(1))) + row_word = row_tok * cutlass.Int32(WORDS_PER_ROW) + col_word + if cutlass.const_expr(tap_out): + tap_word = row_tok * tap_stride + col_word + # (x_src 1 and 2: L1-bypassing, as no grid wait orders them.) + pre = prefix.load(idx=row_word, vector_size=4, alignment=16, is_volatile=x_src >= 1) + # Snapshots past the last one reload snapshot 0 (their candidates are masked below). + snap = [] + rw = pre + mw = pre + if cutlass.const_expr(not plain): + for n in cutlass.range_constexpr(MAX_SNAPSHOTS): + n_row = cutlass.Int32( + cutlass.select_( + cutlass.Int32(n) < num_cand - cutlass.Int32(1), + cutlass.Int32(n), + cutlass.Int32(0), + ) + ) + snap.append( + snapshots.load( + idx=(n_row * num_tokens + row_tok) * cutlass.Int32(WORDS_PER_ROW) + + col_word, + vector_size=4, + alignment=16, + is_volatile=x_src >= 1, + ) + ) + rw = res_w.load(idx=col_word, vector_size=4, alignment=16) + mw = rms_w.load(idx=col_word, vector_size=4, alignment=16) + ow = out_w.load(idx=col_word, vector_size=4, alignment=16) + + n_calls = cutlass.Int32(0) + if cutlass.const_expr(x_src == 2): + # ---- The latent all-reduce, folded in (see the docstring). The previous call's CTA 0 stored this call's + # count after its grid wait, and that call completed before this launch. + n_calls = cutlass.Int32(lat_flags.load(idx=0, is_volatile=True)) + half_words = MAX_TOKENS * world * LAT_WORDS + cur_h = n_calls & cutlass.Int32(1) + h_base = cur_h * cutlass.Int32(half_words) + # The other half (the next call's) emptied first: this CTA's 1/56 of it, 16-byte stores. + e_base = (cur_h ^ cutlass.Int32(1)) * cutlass.Int32(half_words) + bx * cutlass.Int32( + half_words // NUM_CTAS + ) + empty_l = cutlass.Int32(EMPTY_WORD) + for q in cutlass.range_constexpr(half_words // NUM_CTAS // (EPI_THREADS * 4)): + rms_src.store((empty_l, empty_l, empty_l, empty_l), + idx=e_base + cutlass.Int32(q * EPI_THREADS * 4) + tid * cutlass.Int32(4), + alignment=16) # fmt: skip + if bx == cutlass.Int32(0): + if tid < cutlass.Int32(MAX_TOKENS): + next_sb = (n_calls + cutlass.Int32(1)) % cutlass.Int32(LAT_SCALE_BUFS) + lat_flags.store(cutlass.Int32(SCALE_SENTINEL), + idx=cutlass.Int32(LAT_SCALES) + next_sb * cutlass.Int32(MAX_TOKENS) + tid, + is_volatile=True) # fmt: skip + # Rank chunks of 8 (world <= 8: one chunk of world ranks); a thread sums one chunk of one 16-byte vector, + # the pair's even thread adds the two chunks from 0 as the one-shot does, rounds once and sends. + chunks = (world + RANK_CHUNK - 1) // RANK_CHUNK + per_chunk = min(RANK_CHUNK, world) + g = crank * cutlass.Int32(EPI_THREADS) + tid + g_chunk = g % cutlass.Int32(chunks) + g_even = g_chunk == cutlass.Int32(0) + + # (1) The rank's slice of every live row, into all 7 CTAs' smem_b (the latent k-tiles). + s_live = g < num_tokens * cutlass.Int32(LAT_SLICE_VECS * chunks) + s_pair = g // cutlass.Int32(chunks) + s_tok = cutlass.Int32( + cutlass.select_(s_live, s_pair // cutlass.Int32(LAT_SLICE_VECS), cutlass.Int32(0)) + ) + s_vec = s_pair % cutlass.Int32(LAT_SLICE_VECS) + s_base = ( + h_base + (s_tok * cutlass.Int32(world) + g_chunk * cutlass.Int32(RANK_CHUNK)) * cutlass.Int32(LAT_WORDS) + + x_col0 // cutlass.Int32(2) + s_vec * cutlass.Int32(4) + ) # fmt: skip + s0 = cutlass.Float32(0.0) + s1 = cutlass.Float32(0.0) + s2 = cutlass.Float32(0.0) + s3 = cutlass.Float32(0.0) + s4 = cutlass.Float32(0.0) + s5 = cutlass.Float32(0.0) + s6 = cutlass.Float32(0.0) + s7 = cutlass.Float32(0.0) + s_pending = s_live + while s_pending: + s_dirty = cutlass.Boolean(False) + sc = [cutlass.Float32(0.0)] * ELTS + for rr in cutlass.range_constexpr(per_chunk): + spv = rms_src.load(idx=s_base + cutlass.Int32(rr * LAT_WORDS), vector_size=4, alignment=16, + is_volatile=True) # fmt: skip + for q in cutlass.range_constexpr(4): + sword = cutlass.Int32(spv[q]) + s_dirty = s_dirty | (sword == cutlass.Int32(EMPTY_WORD)) + sc[2 * q] = sc[2 * q] + _lo(sword) + sc[2 * q + 1] = sc[2 * q + 1] + _hi(sword) + s0, s1, s2, s3, s4, s5, s6, s7 = sc + s_pending = s_dirty + cute.arch.sync_warp() + s_mine = [s0, s1, s2, s3, s4, s5, s6, s7] + s_tot = [] + for e in cutlass.range_constexpr(ELTS): + if cutlass.const_expr(chunks == 2): + s_other = cute.arch.shuffle_sync_bfly(s_mine[e], offset=1) + s_c0 = cutlass.Float32(cutlass.select_(g_even, s_mine[e], s_other)) + s_c1 = cutlass.Float32(cutlass.select_(g_even, s_other, s_mine[e])) + s_tot.append((cutlass.Float32(0.0) + s_c0) + s_c1) + else: + s_tot.append(cutlass.Float32(0.0) + s_mine[e]) + if s_live & g_even: + s_dst = act_word_index(s_tok, s_vec * cutlass.Int32(8)) + for pc in cutlass.range_constexpr(CLUSTER): + _st_async_v4( + _mapa_u32(smem_b_words.subview(s_dst).data_ptr(), cutlass.Int32(pc)), + _pack_bf16x2(s_tot[1], s_tot[0]), _pack_bf16x2(s_tot[3], s_tot[2]), + _pack_bf16x2(s_tot[5], s_tot[4]), _pack_bf16x2(s_tot[7], s_tot[6]), + _mapa_u32(lat_full.data_ptr(), cutlass.Int32(pc)), + ) # fmt: skip + + # (2) Cluster t: token t's whole row, into CTA 0's rms_row. + r_live = (token < num_tokens) & (g < cutlass.Int32(LAT_VECS * chunks)) + r_tok = cutlass.Int32(cutlass.select_(token < num_tokens, token, cutlass.Int32(0))) + r_vec = g // cutlass.Int32(chunks) + r_base = ( + h_base + (r_tok * cutlass.Int32(world) + g_chunk * cutlass.Int32(RANK_CHUNK)) * cutlass.Int32(LAT_WORDS) + + r_vec * cutlass.Int32(4) + ) # fmt: skip + r0 = cutlass.Float32(0.0) + r1 = cutlass.Float32(0.0) + r2 = cutlass.Float32(0.0) + r3 = cutlass.Float32(0.0) + r4 = cutlass.Float32(0.0) + r5 = cutlass.Float32(0.0) + r6 = cutlass.Float32(0.0) + r7 = cutlass.Float32(0.0) + r_pending = r_live + while r_pending: + r_dirty = cutlass.Boolean(False) + rc = [cutlass.Float32(0.0)] * ELTS + for rr in cutlass.range_constexpr(per_chunk): + rpv = rms_src.load(idx=r_base + cutlass.Int32(rr * LAT_WORDS), vector_size=4, alignment=16, + is_volatile=True) # fmt: skip + for q in cutlass.range_constexpr(4): + rword = cutlass.Int32(rpv[q]) + r_dirty = r_dirty | (rword == cutlass.Int32(EMPTY_WORD)) + rc[2 * q] = rc[2 * q] + _lo(rword) + rc[2 * q + 1] = rc[2 * q + 1] + _hi(rword) + r0, r1, r2, r3, r4, r5, r6, r7 = rc + r_pending = r_dirty + cute.arch.sync_warp() + r_mine = [r0, r1, r2, r3, r4, r5, r6, r7] + r_tot = [] + for e in cutlass.range_constexpr(ELTS): + if cutlass.const_expr(chunks == 2): + r_other = cute.arch.shuffle_sync_bfly(r_mine[e], offset=1) + r_c0 = cutlass.Float32(cutlass.select_(g_even, r_mine[e], r_other)) + r_c1 = cutlass.Float32(cutlass.select_(g_even, r_other, r_mine[e])) + r_tot.append((cutlass.Float32(0.0) + r_c0) + r_c1) + else: + r_tot.append(cutlass.Float32(0.0) + r_mine[e]) + if r_live & g_even: + _st_async_v4( + _mapa_u32(rms_row.subview(r_vec * cutlass.Int32(4)).data_ptr(), cutlass.Int32(0)), + _pack_bf16x2(r_tot[1], r_tot[0]), _pack_bf16x2(r_tot[3], r_tot[2]), + _pack_bf16x2(r_tot[5], r_tot[4]), _pack_bf16x2(r_tot[7], r_tot[6]), + _mapa_u32(rms_full.subview(0).data_ptr(), cutlass.Int32(0)), + ) # fmt: skip + + # (3) CTA 0's warp 4 of cluster t: token t's latent RMS, as the bulk-copy path computes it, published into + # buffer n % 3 of the scale slab. + if (crank == cutlass.Int32(0)) & (w == cutlass.Int32(0)) & (token < num_tokens): + while not _test_wait_cluster(rms_full.subview(0).data_ptr(), 0): + pass + cute.arch.sync_warp() + lat_rows_o = [] + for rv in cutlass.range_constexpr(rms_cols // (8 * 32)): + lat_rows_o.append( + rms_row.load( + idx=cutlass.Int32(rv * 32 * 4) + lane * cutlass.Int32(4), + vector_size=4, + alignment=16, + ) + ) + lat_sq_o = cutlass.Float32(0.0) + for rv in cutlass.range_constexpr(rms_cols // (8 * 32)): + for ri in cutlass.range_constexpr(4): + lat_lo_o = _lo(cutlass.Int32(lat_rows_o[rv][ri])) + lat_hi_o = _hi(cutlass.Int32(lat_rows_o[rv][ri])) + lat_sq_o = lat_sq_o + lat_lo_o * lat_lo_o + lat_hi_o * lat_hi_o + for offset in (16, 8, 4, 2, 1): + lat_sq_o = lat_sq_o + cute.arch.shuffle_sync_bfly(lat_sq_o, offset=offset) + scale_o = cute.math.rsqrt(lat_sq_o * cutlass.Float32(1.0 / rms_cols) + lat_eps) + if lane == cutlass.Int32(0): + lat_flags.store( + scale_o.bitcast(cutlass.Int32), + idx=cutlass.Int32(LAT_SCALES) + + (n_calls % cutlass.Int32(LAT_SCALE_BUFS)) * cutlass.Int32(MAX_TOKENS) + + token, + is_volatile=True, + ) # fmt: skip + + if cutlass.const_expr(tail and x_src != 2): + # Per-token latent RMS while the MMA runs: cluster CTA c computes tokens c and c + 7 (epilogue warps 0 and + # 1) from its bulk-copied row: each lane sums the squares of its 14 16-byte vectors in order, then a + # butterfly; the scale goes by st.async into row_scale of every cluster CTA, completing mb_rms there. + for rj in cutlass.range_constexpr(LATENT_ROWS_PER_CTA): + rt_c = crank + cutlass.Int32(rj * CLUSTER) + if w == cutlass.Int32(rj): + if rt_c < num_tokens: + if cutlass.const_expr(x_src == 1): + # Completed by warp 1's arrive after its polled stores (test_wait: see act_full). + while not cute.arch.mbarrier_test_wait( + rms_full.subview(rj).data_ptr(), 0 + ): + pass + else: + while not cute.arch.mbarrier_try_wait( + rms_full.subview(rj).data_ptr(), 0 + ): + pass + cute.arch.sync_warp() + lat_rows_c = [] + for rv in cutlass.range_constexpr(rms_cols // (8 * 32)): + lat_rows_c.append( + rms_row.load( + idx=cutlass.Int32(rj * (rms_cols // 2) + rv * 32 * 4) + + lane * cutlass.Int32(4), + vector_size=4, + alignment=16, + ) + ) + lat_sq_c = cutlass.Float32(0.0) + for rv in cutlass.range_constexpr(rms_cols // (8 * 32)): + for ri in cutlass.range_constexpr(4): + lat_lo_c = _lo(cutlass.Int32(lat_rows_c[rv][ri])) + lat_hi_c = _hi(cutlass.Int32(lat_rows_c[rv][ri])) + lat_sq_c = lat_sq_c + lat_lo_c * lat_lo_c + lat_hi_c * lat_hi_c + for offset in (16, 8, 4, 2, 1): + lat_sq_c = lat_sq_c + cute.arch.shuffle_sync_bfly( + lat_sq_c, offset=offset + ) + scale_c = cute.math.rsqrt( + lat_sq_c * cutlass.Float32(1.0 / rms_cols) + lat_eps + ) + if lane < cutlass.Int32(CLUSTER): + _st_async_f32( + _mapa_u32(row_scale.subview(rt_c).data_ptr(), lane), + scale_c, + _mapa_u32(mb_rms.data_ptr(), lane), + ) + + # ---- Phase 1 epilogue: TMEM -> bf16 pairs (even lane: its row and the next) -> 16-byte pushes. + # (Every wait loop and lane-dependent branch before a warp shuffle ends in a warp sync: lanes may leave a + # wait loop in different iterations, and a shuffle on a diverged warp takes the slow collective path.) + while not cute.arch.mbarrier_try_wait(acc_done.data_ptr(), 0): + pass + cute.arch.sync_warp() + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc = prims.tcgen05_ld( + "32x32b", cutlass.inttoptr(tmem_ptr_i32.load(), 6, cutlass.Float32), num=MMA_N + ) + acc1 = acc + if cutlass.const_expr(acc_pair): + acc1 = prims.tcgen05_ld( + "32x32b", + cutlass.inttoptr(tmem_ptr_i32.load() + cutlass.Int32(MMA_N), 6, cutlass.Float32), + num=MMA_N, + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + scales = [] + if cutlass.const_expr(tail and x_src == 2): + # The live tokens' scales from buffer n % 3 of the scale slab (published by each token's cluster). + sb_idx = cutlass.Int32(LAT_SCALES) + ( + n_calls % cutlass.Int32(LAT_SCALE_BUFS) + ) * cutlass.Int32(MAX_TOKENS) + sw0 = cutlass.Int32(SCALE_SENTINEL) + sw1 = cutlass.Int32(SCALE_SENTINEL) + sw2 = cutlass.Int32(SCALE_SENTINEL) + sw3 = cutlass.Int32(SCALE_SENTINEL) + sw4 = cutlass.Int32(SCALE_SENTINEL) + sw5 = cutlass.Int32(SCALE_SENTINEL) + sw6 = cutlass.Int32(SCALE_SENTINEL) + sw7 = cutlass.Int32(SCALE_SENTINEL) + sc_pending = cutlass.Boolean(True) + while sc_pending: + sv_a = lat_flags.load(idx=sb_idx, vector_size=4, alignment=16, is_volatile=True) + sv_b = lat_flags.load( + idx=sb_idx + cutlass.Int32(4), vector_size=4, alignment=16, is_volatile=True + ) + sws = [cutlass.Int32(sv_a[q]) for q in range(4)] + [ + cutlass.Int32(sv_b[q]) for q in range(4) + ] + sc_dirty = cutlass.Boolean(False) + for t in cutlass.range_constexpr(MMA_N): + sc_dirty = sc_dirty | ( + (cutlass.Int32(t) < num_tokens) & (sws[t] == cutlass.Int32(SCALE_SENTINEL)) + ) + sw0, sw1, sw2, sw3, sw4, sw5, sw6, sw7 = sws + sc_pending = sc_dirty + cute.arch.sync_warp() + for sword_t in (sw0, sw1, sw2, sw3, sw4, sw5, sw6, sw7): + scales.append(cutlass.Int32(sword_t).bitcast(cutlass.Float32)) + elif cutlass.const_expr(tail): + while not _test_wait_cluster(mb_rms.data_ptr(), 0): + pass + cute.arch.sync_warp() + # Every token's scale at once (rows past the last token are scaled but never pushed). + for sv in cutlass.range_constexpr(MMA_N // 4): + scale_vec = row_scale.load(idx=sv * 4, vector_size=4, alignment=16) + for si in cutlass.range_constexpr(4): + scales.append(cutlass.Float32(scale_vec[si])) + row = w * cutlass.Int32(32) + lane + for t in cutlass.range_constexpr(MMA_N): + mine = cutlass.Float32(acc[t]) + if cutlass.const_expr(tail): + mine = mine * scales[t] + if cutlass.const_expr(split_acc): + mine = mine + cutlass.Float32(acc1[t]) + if cutlass.const_expr(k_split == 2): + # Split 2's sums: the rank partials added from zero in rank order. + mine = (cutlass.Float32(0.0) + mine) + cutlass.Float32(acc1[t]) + odd = cute.arch.shuffle_sync_bfly(mine, offset=1) + if lane % 2 == 0: + stage.store( + _pack_bf16x2(odd, mine), idx=cutlass.Int32(t * ROW_WORDS_PER_CTA) + row // 2 + ) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + prims.barrier_cta_sync(1, thread_count=EPI_THREADS) + push_tok = tid // cutlass.Int32(16) + push_vec = tid % cutlass.Int32(16) + if push_tok < num_tokens: + v = stage.load( + idx=push_tok * cutlass.Int32(ROW_WORDS_PER_CTA) + push_vec * cutlass.Int32(4), + vector_size=4, + alignment=16, + ) + if cutlass.const_expr(x_src == 2): + # The epilogue warps' empties of the latent half (ordered before this thread by the stage barrier) + # before the push: a rank's next push into that half follows its phase 2 of this call. + _fence_sys() + slot = ( + (cur * cutlass.Int32(MAX_TOKENS) + push_tok) * cutlass.Int32(world) + rank + ) * cutlass.Int32(WORDS_PER_ROW) + ws_mc.store( + (_sanitize(v[0]), _sanitize(v[1]), _sanitize(v[2]), _sanitize(v[3])), + idx=slot + (m_offset // cutlass.Int32(2)) + push_vec * cutlass.Int32(4), + alignment=16, + ) + prims.fence_acq_rel(prims.MemScope.CLUSTER) + # The dependents launch once every CTA has pushed: they stream their weights during the exchange without + # competing with phase 1, and they wait for (or poll) this grid's outputs. + if tid == 0: + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + # After the stage barrier: every TMEM read of this CTA is done and every epilogue thread has read the count. + if warp_id == 4: + prims.tcgen05_dealloc(tmem_ptr, TMEM_COLS) + if tid == 0: + ws_flags.store(calls + cutlass.Int32(1), idx=bx, is_volatile=True) + + # ---- Phase 2: reduce token `token` over the ranks, then the attn_res + RMSNorm epilogue. + cute.arch.sync_warp() + if reducer: + if cutlass.const_expr(not plain): + # res_w * rms_w of this thread's columns. + qv = [] + for q in cutlass.range_constexpr(4): + rword = cutlass.Int32(rw[q]) + mword = cutlass.Int32(mw[q]) + qv.append(_lo(rword) * _lo(mword)) + qv.append(_hi(rword) * _hi(mword)) + + # The snapshots' statistics, while the peers' rows are in flight: warp sums (value i of 16 in lanes + # 2 i, 2 i + 1), the CTA sum over the warps in order, sent to slot [this CTA] of every cluster CTA. + snap_stats = [] + for n in cutlass.range_constexpr(MAX_SNAPSHOTS): + sum_sq = cutlass.Float32(0.0) + dot = cutlass.Float32(0.0) + for e in cutlass.range_constexpr(ELTS): + word = cutlass.Int32(snap[n][e // 2]) + val = _lo(word) if e % 2 == 0 else _hi(word) + sum_sq = val * val + sum_sq + dot = val * qv[e] + dot + snap_stats.append(sum_sq) + snap_stats.append(dot) + snap_sum = _warp_sum_scatter(snap_stats, lane) + if lane % 2 == 0: + warp_snap.store( + snap_sum, idx=w * cutlass.Int32(SNAP_STATS) + lane // cutlass.Int32(2) + ) + prims.barrier_cta_sync(1, thread_count=EPI_THREADS) + if tid < cutlass.Int32(CLUSTER * SNAP_STATS): + peer = tid // cutlass.Int32(SNAP_STATS) + s = tid % cutlass.Int32(SNAP_STATS) + if s < cutlass.Int32(2) * (num_cand - cutlass.Int32(1)): + partial = cutlass.Float32(0.0) + for ww in cutlass.range_constexpr(EPI_WARPS): + partial = partial + warp_snap.load( + idx=cutlass.Int32(ww * SNAP_STATS) + s + ) + _st_async_f32( + _mapa_u32( + box_snap.subview(crank * cutlass.Int32(SNAP_STATS) + s).data_ptr(), + peer, + ), + partial, + _mapa_u32(mb_snap.data_ptr(), peer), + ) + # The snapshots' logits and their maximum need only the cluster's snapshot statistics: taken here, + # while the ranks' rows are in flight (lane n owns snapshot n). + while not _test_wait_cluster(mb_snap.data_ptr(), 0): + pass + cute.arch.sync_warp() + snap_lane = cutlass.Int32( + cutlass.select_(lane < cutlass.Int32(MAX_SNAPSHOTS), lane, cutlass.Int32(0)) + ) + box_s = [] + for r in cutlass.range_constexpr(CLUSTER): + box_s.append( + box_snap.load( + idx=cutlass.Int32(r * SNAP_STATS) + cutlass.Int32(2) * snap_lane, + vector_size=2, + alignment=8, + ) + ) + snap_tot_sq = cutlass.Float32(0.0) + snap_tot_dot = cutlass.Float32(0.0) + for r in cutlass.range_constexpr(CLUSTER): + snap_tot_sq = snap_tot_sq + cutlass.Float32(box_s[r][0]) + snap_tot_dot = snap_tot_dot + cutlass.Float32(box_s[r][1]) + is_snap = lane < num_cand - cutlass.Int32(1) + snap_logit = cutlass.Float32( + cutlass.select_( + is_snap, + snap_tot_dot * cute.math.rsqrt(snap_tot_sq / cutlass.Float32(H) + rms_eps), + cutlass.Float32(-3.4028234663852886e38), + ) + ) + max_snap = snap_logit + for offset in (16, 8, 4, 2, 1): + max_snap = _fmax(max_snap, cute.arch.shuffle_sync_bfly(max_snap, offset=offset)) + + slot0 = (cur * cutlass.Int32(MAX_TOKENS) + token) * cutlass.Int32( + world * WORDS_PER_ROW + ) + col_word + a0 = cutlass.Float32(0.0) + a1 = cutlass.Float32(0.0) + a2 = cutlass.Float32(0.0) + a3 = cutlass.Float32(0.0) + a4 = cutlass.Float32(0.0) + a5 = cutlass.Float32(0.0) + a6 = cutlass.Float32(0.0) + a7 = cutlass.Float32(0.0) + pending = cutlass.Boolean(True) + while pending: + dirty = cutlass.Boolean(False) + total = [cutlass.Float32(0.0)] * ELTS + for rb in cutlass.range_constexpr(0, world, RANK_CHUNK): + chunk = [cutlass.Float32(0.0)] * ELTS + for rr in cutlass.range_constexpr(min(RANK_CHUNK, world - rb)): + pv = ws_uc.load( + idx=slot0 + cutlass.Int32((rb + rr) * WORDS_PER_ROW), + vector_size=4, + alignment=16, + is_volatile=True, + ) + for q in cutlass.range_constexpr(4): + pword = cutlass.Int32(pv[q]) + dirty = dirty | (pword == cutlass.Int32(EMPTY_WORD)) + chunk[2 * q] = chunk[2 * q] + _lo(pword) + chunk[2 * q + 1] = chunk[2 * q + 1] + _hi(pword) + for e in cutlass.range_constexpr(ELTS): + total[e] = total[e] + chunk[e] + a0, a1, a2, a3, a4, a5, a6, a7 = total + pending = dirty + cute.arch.sync_warp() + acc_sum = [a0, a1, a2, a3, a4, a5, a6, a7] + + # updated = bf16(prefix + bf16(sum)) (bf16(sum) without the prefix). + upd_words = [] + for q in cutlass.range_constexpr(4): + red = _pack_bf16x2(acc_sum[2 * q + 1], acc_sum[2 * q]) + pw = cutlass.Int32(pre[q]) + with_prefix = _pack_bf16x2(_hi(pw) + _hi(red), _lo(pw) + _lo(red)) + upd_words.append( + cutlass.Int32(cutlass.select_(add_prefix != cutlass.Int32(0), with_prefix, red)) + ) + + if cutlass.const_expr(plain): + # The MNNVL one-shot's residual + RMSNorm (kARResidualRMSNorm), its arithmetic and reduction order. + updated.store( + (upd_words[0], upd_words[1], upd_words[2], upd_words[3]), + idx=row_word, + alignment=16, + ) + empty_p = cutlass.Int32(EMPTY_WORD) + for r in cutlass.range_constexpr(world): + ws_uc.store((empty_p, empty_p, empty_p, empty_p), idx=slot0 + cutlass.Int32(r * WORDS_PER_ROW), + alignment=16) # fmt: skip + # This thread's sum of squares in element order: bf16 products (x * x rounded to bf16; plain 2: fp32, + # exact for bf16 x). + t_sq = cutlass.Float32(0.0) + for e in cutlass.range_constexpr(ELTS): + word = upd_words[e // 2] + val = _lo(word) if e % 2 == 0 else _hi(word) + if cutlass.const_expr(plain == 2): + t_sq = t_sq + val * val + else: + sq_pair = _pack_bf16x2(val * val, val * val) + t_sq = t_sq + _lo(sq_pair) + # To slot [this thread's group] of every cluster CTA. + my_group = crank * cutlass.Int32(EPI_THREADS) + tid + for pc in cutlass.range_constexpr(CLUSTER): + _st_async_f32( + _mapa_u32(box_ts.subview(my_group).data_ptr(), cutlass.Int32(pc)), t_sq, + _mapa_u32(mb_ts.data_ptr(), cutlass.Int32(pc)), + ) # fmt: skip + while not _test_wait_cluster(mb_ts.data_ptr(), 0): + pass + cute.arch.sync_warp() + if w == cutlass.Int32(0): + # The one-shot's tree over the 896 thread sums. MNNVL (8 blocks of 112 threads): lane k takes warp + # k % 4 of block k // 4 (32 thread sums; the fourth warp has 16, the rest zero) through the xor + # butterfly; lane 4 b then forms block b's sum (s0 + s2) + (s1 + s3), and lane 0 adds the + # blocks. IPC (plain 2, 4 blocks of 224 threads): lane k < 28 takes warp k % 7 of block k // 7 + # (thread sums [32 k, 32 k + 32)); lane 7 b forms block b's sum as blockReduceSumV2's butterfly + # over the 7 warp sums and 25 zeros, ((s0 + s4) + (s2 + s6)) + ((s1 + s5) + s3); lane 0 adds the + # blocks. + vw_live = lane < cutlass.Int32(28 if plain == 2 else 32) + if cutlass.const_expr(plain == 2): + vw_base = lane * cutlass.Int32(32) + vw_count = cutlass.Int32(32) + else: + vw_base = (lane // cutlass.Int32(4)) * cutlass.Int32(112) + ( + lane % cutlass.Int32(4) + ) * cutlass.Int32(32) + vw_count = cutlass.Int32(cutlass.select_(lane % cutlass.Int32(4) == cutlass.Int32(3), + cutlass.Int32(16), cutlass.Int32(32))) # fmt: skip + vals = [] + for i in cutlass.range_constexpr(32): + v_i = box_ts.load( + idx=cutlass.Int32( + cutlass.select_( + vw_live, vw_base + cutlass.Int32(i), cutlass.Int32(0) + ) + ) + ) + vals.append( + cutlass.Float32( + cutlass.select_( + cutlass.Int32(i) < vw_count, v_i, cutlass.Float32(0.0) + ) + ) + ) + for m in (16, 8, 4, 2, 1): + vals = [vals[ln] + vals[ln ^ m] for ln in range(32)] + warp_sum = vals[0] + s_w = [warp_sum] + for j in cutlass.range_constexpr(1, 7 if plain == 2 else 4): + s_w.append( + cute.arch.shuffle_sync( + warp_sum, offset=(lane + cutlass.Int32(j)) % cutlass.Int32(32) + ) + ) + if cutlass.const_expr(plain == 2): + block_sum = ((s_w[0] + s_w[4]) + (s_w[2] + s_w[6])) + ( + (s_w[1] + s_w[5]) + s_w[3] + ) + else: + block_sum = (s_w[0] + s_w[2]) + (s_w[1] + s_w[3]) + full_sum = cutlass.Float32(0.0) + for b in cutlass.range_constexpr(4 if plain == 2 else 8): + full_sum = full_sum + cute.arch.shuffle_sync( + block_sum, offset=(7 if plain == 2 else 4) * b + ) + r_plain = cute.math.rsqrt(full_sum / cutlass.Float32(H) + out_eps) + if lane == cutlass.Int32(0): + row_scale.store(r_plain, idx=0) + prims.barrier_cta_sync(1, thread_count=EPI_THREADS) + rsigma_p = row_scale.load(idx=0) + out_words = [] + for q in cutlass.range_constexpr(4): + oword = cutlass.Int32(ow[q]) + x_word = upd_words[q] + out_words.append( + _pack_bf16x2( + _hi(x_word) * rsigma_p * _hi(oword), _lo(x_word) * rsigma_p * _lo(oword) + ) + ) + normed.store( + (out_words[0], out_words[1], out_words[2], out_words[3]), + idx=row_word, + alignment=16, + ) + if cutlass.const_expr(publish): + x_slab.store( + (_slab_word(out_words[0]), _slab_word(out_words[1]), _slab_word(out_words[2]), + _slab_word(out_words[3])), + idx=slab_buf * cutlass.Int32(SLAB_WORDS) + row_word, + alignment=16, + ) # fmt: skip + else: + # The updated sum's statistics: warp sums (sum of squares in lanes 0-15, projection in 16-31), sent per + # warp to slot [this CTA][warp] of every cluster CTA. + u_sq = cutlass.Float32(0.0) + u_dot = cutlass.Float32(0.0) + for e in cutlass.range_constexpr(ELTS): + word = upd_words[e // 2] + val = _lo(word) if e % 2 == 0 else _hi(word) + u_sq = val * val + u_sq + u_dot = val * qv[e] + u_dot + upd_sum = _warp_sum_scatter([u_sq, u_dot], lane) + u_peer = lane % cutlass.Int32(16) + if u_peer < cutlass.Int32(CLUSTER): + _st_async_f32( + _mapa_u32( + box_upd.subview( + (crank * cutlass.Int32(EPI_WARPS) + w) * cutlass.Int32(2) + + lane // cutlass.Int32(16) + ).data_ptr(), + u_peer, + ), + upd_sum, + _mapa_u32(mb_upd.data_ptr(), u_peer), + ) + updated.store( + (upd_words[0], upd_words[1], upd_words[2], upd_words[3]), + idx=row_word, + alignment=16, + ) + if cutlass.const_expr(tap_out == 2): + # The running prefix sum (a capture layer that taps the prefix rather than the mixture). + tap.store( + (upd_words[0], upd_words[1], upd_words[2], upd_words[3]), + idx=tap_word, + alignment=16, + ) + # Every rank's single push of this call into these words has landed; empty them for the call after next. + empty = cutlass.Int32(EMPTY_WORD) + for r in cutlass.range_constexpr(world): + ws_uc.store( + (empty, empty, empty, empty), + idx=slot0 + cutlass.Int32(r * WORDS_PER_ROW), + alignment=16, + ) + # Candidate n < num_cand - 1 is snapshot n, candidate num_cand - 1 is updated; later ones are masked. + cand = [] + for n in cutlass.range_constexpr(MAX_CANDIDATES): + vals = [] + for q in cutlass.range_constexpr(4): + word = upd_words[q] + if cutlass.const_expr(n < MAX_CANDIDATES - 1): + word = cutlass.Int32( + cutlass.select_( + cutlass.Int32(n) == num_cand - cutlass.Int32(1), + upd_words[q], + cutlass.Int32(snap[n][q]), + ) + ) + vals.append(_lo(word)) + vals.append(_hi(word)) + cand.append(vals) + + # Every warp: the candidates' logits (lane n owns candidate n: snapshot n < num_cand - 1, then updated), + # each statistic summed over warps, then over the cluster's CTAs in rank order. The snapshots' logits + # and their maximum were taken before the poll; now the updated sum's. + while not _test_wait_cluster(mb_upd.data_ptr(), 0): + pass + cute.arch.sync_warp() + box_u = [] + for r in cutlass.range_constexpr(CLUSTER): + for hv in cutlass.range_constexpr(EPI_WARPS * 2 // 4): + box_u.append( + box_upd.load( + idx=r * EPI_WARPS * 2 + hv * 4, vector_size=4, alignment=16 + ) + ) + upd_tot_sq = cutlass.Float32(0.0) + upd_tot_dot = cutlass.Float32(0.0) + for r in cutlass.range_constexpr(CLUSTER): + part_sq = cutlass.Float32(0.0) + part_dot = cutlass.Float32(0.0) + for ww in cutlass.range_constexpr(EPI_WARPS): + part_sq = part_sq + cutlass.Float32(box_u[r * 2 + ww // 2][(ww % 2) * 2]) + part_dot = part_dot + cutlass.Float32( + box_u[r * 2 + ww // 2][(ww % 2) * 2 + 1] + ) + upd_tot_sq = upd_tot_sq + part_sq + upd_tot_dot = upd_tot_dot + part_dot + upd_logit = upd_tot_dot * cute.math.rsqrt(upd_tot_sq / cutlass.Float32(H) + rms_eps) + # The maximum over every lane's logit is exact whatever the order it is taken in. + max_logit = _fmax(max_snap, upd_logit) + live = lane < num_cand + logit = cutlass.Float32( + cutlass.select_( + is_snap, + snap_logit, + cutlass.select_( + lane == num_cand - cutlass.Int32(1), + upd_logit, + cutlass.Float32(-3.4028234663852886e38), + ), + ) + ) + weight = cutlass.Float32(0.0) + if live: + weight = cute.math.exp2( + (logit - max_logit) * cutlass.Float32(LOG2E), fastmath=True + ) + denominator = weight + for offset in (16, 8, 4, 2, 1): + denominator = denominator + cute.arch.shuffle_sync_bfly( + denominator, offset=offset + ) + # Zero past num_cand, so the masked candidates add nothing below. + lane_weight = weight * (cutlass.Float32(1.0) / denominator) + weights = [ + cute.arch.shuffle_sync(lane_weight, offset=n) for n in range(MAX_CANDIDATES) + ] + + # Selection, rounded to bf16, and its RMS over the token (second cluster reduction). + mixed = [] + out_sq = cutlass.Float32(0.0) + for e in cutlass.range_constexpr(ELTS): + value = cutlass.Float32(0.0) + for n in cutlass.range_constexpr(MAX_CANDIDATES): + value = weights[n] * cand[n][e] + value + mixed.append(value) + mixed_words = [] + for q in cutlass.range_constexpr(4): + mixed_words.append(_pack_bf16x2(mixed[2 * q + 1], mixed[2 * q])) + if cutlass.const_expr(tap_out == 1): + # The mixture the RMSNorm below normalizes (a DSpark capture layer's tap). + tap.store( + (mixed_words[0], mixed_words[1], mixed_words[2], mixed_words[3]), + idx=tap_word, + alignment=16, + ) + for q in cutlass.range_constexpr(4): + lo = _lo(mixed_words[q]) + hi = _hi(mixed_words[q]) + out_sq = lo * lo + out_sq + out_sq = hi * hi + out_sq + out_sq = _warp_allsum(out_sq) + if lane < cutlass.Int32(CLUSTER): + _st_async_f32( + _mapa_u32( + box_sq.subview(crank * cutlass.Int32(EPI_WARPS) + w).data_ptr(), lane + ), + out_sq, + _mapa_u32(mb_sq.data_ptr(), lane), + ) + while not _test_wait_cluster(mb_sq.data_ptr(), 0): + pass + cute.arch.sync_warp() + box_q = [ + box_sq.load(idx=r * EPI_WARPS, vector_size=EPI_WARPS, alignment=16) + for r in range(CLUSTER) + ] + total_sq = cutlass.Float32(0.0) + for r in cutlass.range_constexpr(CLUSTER): + part = cutlass.Float32(0.0) + for ww in cutlass.range_constexpr(EPI_WARPS): + part = part + cutlass.Float32(box_q[r][ww]) + total_sq = total_sq + part + rsigma = cute.math.rsqrt(total_sq / cutlass.Float32(H) + out_eps) + out_words = [] + for q in cutlass.range_constexpr(4): + oword = cutlass.Int32(ow[q]) + # KimiK3RMSNorm: normalize in fp32, round to bf16, then apply the bf16 weight. + normed_pair = _pack_bf16x2( + _hi(mixed_words[q]) * rsigma, _lo(mixed_words[q]) * rsigma + ) + out_words.append( + _pack_bf16x2(_hi(normed_pair) * _hi(oword), _lo(normed_pair) * _lo(oword)) + ) + normed.store( + (out_words[0], out_words[1], out_words[2], out_words[3]), + idx=row_word, + alignment=16, + ) + if cutlass.const_expr(publish): + # The same 16 bytes into the slab, where the next kernel's poll takes them as ready. + x_slab.store( + ( + _slab_word(out_words[0]), + _slab_word(out_words[1]), + _slab_word(out_words[2]), + _slab_word(out_words[3]), + ), + idx=slab_buf * cutlass.Int32(SLAB_WORDS) + row_word, + alignment=16, + ) + if cutlass.const_expr(x_src >= 1): + # The grid completes only after its producer (the dependents' own transitive reads rely on that). + prims.griddepcontrol(prims.GridDepAction.WAIT) + if cutlass.const_expr(x_src == 2): + # The next call's count, kept in [0, 6) so that it never wraps: its parity picks the half, its value mod 3 + # the scale buffer. Every CTA of this call has read this one (CTA 0's phase 2 needed every CTA's push). + if (bx == cutlass.Int32(0)) & (tid == cutlass.Int32(0)): + count_period = cutlass.Int32(2 * LAT_SCALE_BUFS) + lat_flags.store((n_calls % count_period + cutlass.Int32(1)) % count_period, idx=0, is_volatile=True) + + +@cute.kernel +def k3_sandwich_kernel( + tma_desc_w: cutlass.GridConstant[cuda.TensorMap], + tma_desc_x: cutlass.GridConstant[cuda.TensorMap], + tma_desc_x2: cutlass.GridConstant[cuda.TensorMap], + rms_src: cutlass.Array, + ws_uc: cutlass.Array, + ws_mc: cutlass.Array, + ws_flags: cutlass.Array, + prefix: cutlass.Array, + snapshots: cutlass.Array, + res_w: cutlass.Array, + rms_w: cutlass.Array, + out_w: cutlass.Array, + updated: cutlass.Array, + normed: cutlass.Array, + x_slab: cutlass.Array, + lat_flags: cutlass.Array, + tap: cutlass.Array, + num_tokens: cutlass.Int32, + rank: cutlass.Int32, + num_cand: cutlass.Int32, + add_prefix: cutlass.Int32, + rms_eps: cutlass.Float32, + out_eps: cutlass.Float32, + x_col0: cutlass.Int32, + lat_eps: cutlass.Float32, + slab_buf: cutlass.Int32, + x_buf: cutlass.Int32, + tap_stride: cutlass.Int32, + world: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + x_tiles: cutlass.Constexpr[int], + rms_cols: cutlass.Constexpr[int], + publish: cutlass.Constexpr[int], + x_src: cutlass.Constexpr[int], + plain: cutlass.Constexpr[int], + tap_out: cutlass.Constexpr[int], + swiglu: cutlass.Constexpr[int], + k_split: cutlass.Constexpr[int], +): + """The sandwich alone: 56 CTAs, each running ``sandwich_role`` on its own shared memory (arguments as there).""" + smem = SmemAllocator().allocate(smem_bytes(k_in, rms_cols, swiglu), byte_alignment=1024) + sandwich_role( + tma_desc_w, tma_desc_x, tma_desc_x2, rms_src, ws_uc, ws_mc, ws_flags, prefix, snapshots, res_w, rms_w, out_w, + updated, normed, x_slab, lat_flags, tap, smem, num_tokens, rank, num_cand, add_prefix, rms_eps, out_eps, + x_col0, lat_eps, slab_buf, x_buf, tap_stride, world, k_in, x_tiles, rms_cols, publish, x_src, plain, tap_out, + swiglu, k_split, 0, None, + ) # fmt: skip + + +def weight_tensor_map(w, k_in): + """W [7168, k_in] as five TMA dimensions (64-element column chunk, row, chunk index, 1, 1) so one call per + k-tile lands both 128-byte-swizzled halves; strides in 16-byte units.""" + return cuda.create_tensor_map_tiled( + global_address=w.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[TMA_K_BOX, H, k_in // TMA_K_BOX, 1, 1], + global_strides=[ + (k_in * ELEM_BYTES) // 16, + (TMA_K_BOX * ELEM_BYTES) // 16, + (H * k_in * ELEM_BYTES) // 16, + (H * k_in * ELEM_BYTES) // 16, + ], + box_dims=[TMA_K_BOX, CTA_M, TMA_COPY_ITERS, 1, 1], + swizzle=cuda.TensorMapSwizzle.s128b, + ) + + +def activation_tensor_map(x, cols, num_tokens): + """x [M, cols] as (cols, M) with an 8-row box: rows past num_tokens arrive as zeros.""" + return cuda.create_tensor_map_tiled( + global_address=x.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[cols, num_tokens], + global_strides=[(cols * ELEM_BYTES) // 16], + box_dims=[TMA_K_BOX, MMA_N], + swizzle=cuda.TensorMapSwizzle.s128b, + ) + + +@cute.jit +def k3_sandwich_oproj( + w: cute.Tensor, + x: cute.Tensor, + ws_uc: cute.Tensor, + ws_mc: cute.Tensor, + ws_flags: cute.Tensor, + prefix: cute.Tensor, + snapshots: cute.Tensor, + res_w: cute.Tensor, + rms_w: cute.Tensor, + out_w: cute.Tensor, + updated: cute.Tensor, + normed: cute.Tensor, + x_slab: cute.Tensor, + src_slab: cute.Tensor, + num_tokens: cutlass.Int32, + rank: cutlass.Int32, + num_cand: cutlass.Int32, + add_prefix: cutlass.Int32, + rms_eps: cutlass.Float32, + out_eps: cutlass.Float32, + slab_buf: cutlass.Int32, + src_buf: cutlass.Int32, + world: cutlass.Constexpr[int], + publish: cutlass.Constexpr[int], + x_src: cutlass.Constexpr[int], + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + """The post-attention sandwich: phase 1 is ``core @ W_o^T`` (x_src 1: core polled from ``src_slab``).""" + tma_desc_w = weight_tensor_map(w, K_IN) + tma_desc_x = activation_tensor_map(x, K_IN, num_tokens) + k3_sandwich_kernel( + tma_desc_w, tma_desc_x, tma_desc_x, src_slab, ws_uc, ws_mc, ws_flags, prefix, snapshots, res_w, rms_w, out_w, + updated, normed, x_slab, ws_flags, ws_flags, num_tokens, rank, num_cand, add_prefix, rms_eps, out_eps, + cutlass.Int32(0), cutlass.Float32(0.0), slab_buf, src_buf, cutlass.Int32(WORDS_PER_ROW), world, K_IN, + MAX_K_TILES, 0, publish, x_src, 0, 0, 0, 1, + ).launch( + grid=(NUM_CTAS, 1, 1), block=(THREADS, 1, 1), cluster=(CLUSTER, 1, 1), stream=stream, use_pdl=use_pdl, + ) # fmt: skip + + +@cute.jit +def k3_sandwich_tail( + w: cute.Tensor, # [7168, 256 + 384] bf16: latent-up columns of this rank's slice (zero-padded) | shared down + latent: cute.Tensor, # [M, 3584] bf16, the whole reduced latent row + latent_words: cute.Tensor, # the same memory as int32 words, for the RMS + act: cute.Tensor, # [M, 384] bf16, the shared-expert activation + ws_uc: cute.Tensor, + ws_mc: cute.Tensor, + ws_flags: cute.Tensor, + prefix: cute.Tensor, + snapshots: cute.Tensor, + res_w: cute.Tensor, + rms_w: cute.Tensor, + out_w: cute.Tensor, + updated: cute.Tensor, + normed: cute.Tensor, + x_slab: cute.Tensor, + lat_flags: cute.Tensor, # x_src 2: int32 [LAT_FLAG_WORDS] (the call count and the scale slab) + tap: cute.Tensor, # tap_out: int32 words of the bf16 tap rows [M, 7168], tap_stride words apart + num_tokens: cutlass.Int32, + rank: cutlass.Int32, + num_cand: cutlass.Int32, + add_prefix: cutlass.Int32, + rms_eps: cutlass.Float32, + out_eps: cutlass.Float32, + lat_col0: cutlass.Int32, + lat_eps: cutlass.Float32, + slab_buf: cutlass.Int32, + src_buf: cutlass.Int32, + tap_stride: cutlass.Int32, + world: cutlass.Constexpr[int], + publish: cutlass.Constexpr[int], + x_src: cutlass.Constexpr[int], + tap_out: cutlass.Constexpr[int], # 1: the pre-norm mixture into ``tap``; 2: updated + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + """The pre-attention sandwich: phase 1 is the row-parallel MoE tail (x_src 1: ``latent_words`` is the reduced + latent's Lamport slab, int32 [3][8][1792], polled in buffer ``src_buf``; x_src 2: ``latent_words`` is this rank's + latent exchange buffer, int32 [2][8][world][1792], and ``latent`` gives only the shape).""" + k_in = TAIL_LAT + TAIL_ACT + tma_desc_w = weight_tensor_map(w, k_in) + tma_desc_x = activation_tensor_map(latent, LATENT, num_tokens) + tma_desc_x2 = activation_tensor_map(act, TAIL_ACT, num_tokens) + k3_sandwich_kernel( + tma_desc_w, tma_desc_x, tma_desc_x2, latent_words, ws_uc, ws_mc, ws_flags, prefix, snapshots, res_w, rms_w, + out_w, updated, normed, x_slab, lat_flags, tap, num_tokens, rank, num_cand, add_prefix, rms_eps, out_eps, + lat_col0, lat_eps, slab_buf, src_buf, tap_stride, world, k_in, TAIL_LAT // CTA_K, LATENT, publish, x_src, 0, + tap_out, 0, 1, + ).launch( + grid=(NUM_CTAS, 1, 1), block=(THREADS, 1, 1), cluster=(CLUSTER, 1, 1), stream=stream, use_pdl=use_pdl, + ) # fmt: skip + + +@cute.jit +def k3_sandwich_plain( + w: cute.Tensor, # [7168, k_in] bf16: this rank's slice of the row-parallel projection + x: cute.Tensor, # [M, k_in] bf16 + ws_uc: cute.Tensor, + ws_mc: cute.Tensor, + ws_flags: cute.Tensor, + residual: cute.Tensor, # int32 words of bf16 [M, 7168] + norm_w: cute.Tensor, # int32 words of bf16 [7168] + updated: cute.Tensor, + normed: cute.Tensor, + num_tokens: cutlass.Int32, + rank: cutlass.Int32, + eps: cutlass.Float32, + world: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + order: cutlass.Constexpr[ + int + ], # 1: the MNNVL one-shot's arithmetic, 2: the IPC one-shot's (world <= 8) + swiglu: cutlass.Constexpr[ + int + ], # 1: x is a gate_up output [M, 2 k_in] (gate columns first), B = silu_and_mul(x) + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + """The plain sandwich: ``x @ w^T``, the TP all-reduce, updated = sum + residual and normed = RMSNorm(updated), + as a row-parallel GEMV followed by the one-shot all-reduce's RESIDUAL_RMS_NORM. swiglu: ``silu_and_mul(x) @ w^T`` + with k3_ctm_gemv_swiglu split 2's arithmetic (the drafter MLP's down projection).""" + tma_desc_w = weight_tensor_map(w, k_in) + tma_desc_x = activation_tensor_map(x, 2 * k_in if swiglu else k_in, num_tokens) + k3_sandwich_kernel( + tma_desc_w, tma_desc_x, tma_desc_x, ws_flags, ws_uc, ws_mc, ws_flags, residual, residual, norm_w, norm_w, + norm_w, updated, normed, ws_flags, ws_flags, ws_flags, num_tokens, rank, cutlass.Int32(1), cutlass.Int32(1), + eps, eps, cutlass.Int32(k_in if swiglu else 0), cutlass.Float32(0.0), cutlass.Int32(0), cutlass.Int32(0), + cutlass.Int32(WORDS_PER_ROW), world, k_in, k_in // CTA_K, 0, 0, 0, order, 0, swiglu, 2 if swiglu else 1, + ).launch( + grid=(NUM_CTAS, 1, 1), block=(THREADS, 1, 1), cluster=(CLUSTER, 1, 1), stream=stream, use_pdl=use_pdl, + ) # fmt: skip diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py new file mode 100644 index 000000000000..d78692dbd862 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py @@ -0,0 +1,616 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Kimi K3's collective sandwiches in CuTe DSL, for M <= 8 decode tokens: a row-parallel projection, the TP +all-reduce of its output and the residual update (attention-residual selection + RMSNorm) in one kernel. +``trtllm::k3_sandwich_oproj`` is the post-attention step (the attention output projection); +``trtllm::k3_sandwich_tail`` the pre-attention step (the MoE tail, then the next layer's input norm); +``trtllm::k3_sandwich_plain`` a row-parallel projection with a plain residual add + RMSNorm (the drafter layers). + +The all-reduce runs over a dedicated multicast buffer per TP group (:func:`workspace`), not the model's MNNVL +all-reduce workspace, so it keeps its own call parity. The kernel compiles on the first call for each +(world, publish, input source, PDL), which must happen outside CUDA-graph capture; the number of snapshots, the prefix +and the slab buffer are runtime arguments. + +Publishing (``x_slab``, ``slab_buf``): with a slab (``slab_tensor``) the normed rows are also written into buffer +``slab_buf`` (0-2) of it, sentinel-armed Lamport words the next kernel polls, and buffer ``(slab_buf + 1) % 3`` is +re-armed. ``slab_buf`` is the ordinal of the call among this op's calls in the forward, mod 3. + +The latent all-reduce folded into the tail (``lat_uc``, ``lat_flags`` from :func:`latent_exchange`): k3_moe pushes +its routed latent partial into every rank's exchange buffer and exits; ``k3_sandwich_tail`` sums the ranks' partials +itself (bit-identical to the one-shot all-reduce) instead of reading a reduced ``latent``. Every push-only k3_moe call +must be followed by exactly one such tail call on the same exchange. + +Polling the input (``src_slab``, ``src_buf``): when the producer of phase 1's input (the attention output core for +``k3_sandwich_oproj``, the reduced latent for ``k3_sandwich_tail``) publishes it as such a slab (int32 [3][8][cols / +2], sentinel 0xFFFFFFFF) and launches its dependents only after its own grid wait, the kernel polls buffer +``src_buf`` of it instead of waiting for the producer's grid; ``core`` / ``latent`` then give only the shape. +""" + +from __future__ import annotations + +import os +import threading +from typing import Dict, List, Optional + +import torch + +MAX_TOKENS = 8 +EMPTY_WORD = -(2**31) +SLAB_BUFS = 3 +SLAB_SENTINEL = -1 + +_lock = threading.Lock() +_compiled: Dict[tuple, object] = {} +_workspaces: Dict[object, dict] = {} +_lat_exchanges: Dict[object, dict] = {} + + +def _arg(t: torch.Tensor): + from cutlass.cute.runtime import from_dlpack + + # detach(): DLPack refuses tensors that require grad (weights are parameters). + return from_dlpack(t.detach(), assumed_align=16).mark_layout_dynamic(leading_dim=t.dim() - 1) + + +def _words(t: torch.Tensor) -> torch.Tensor: + return t.reshape(-1).view(torch.int32) + + +def _tap_words(tap: torch.Tensor) -> torch.Tensor: + """A strided bf16 [M, 7168] view as the flat int32 words from its first row to the end of its last (the gaps + between rows are never touched), without a copy.""" + words = tap.view(torch.int32) + extent = (tap.shape[0] - 1) * words.stride(0) + words.shape[1] + return words.as_strided((extent,), (1,)) + + +def _use_pdl() -> bool: + return os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" + + +def slab_tensor(device) -> torch.Tensor: + """A publication slab for one producer edge: int32 [3, 8, 3584] (three buffers of 8 bf16 rows of 7168), every + word the sentinel.""" + from . import k3_sandwich_kernel as kernel + + return torch.full( + (SLAB_BUFS, MAX_TOKENS, kernel.WORDS_PER_ROW), + SLAB_SENTINEL, + dtype=torch.int32, + device=device, + ) + + +def _src_args(src_slab: Optional[torch.Tensor], src_buf: int, cols: int, fallback: torch.Tensor): + """(slab words, buffer, x_src) of phase 1's input; ``fallback`` and x_src 0 without a slab.""" + if src_slab is None: + return fallback, 0, 0 + words = SLAB_BUFS * MAX_TOKENS * cols // 2 + if src_slab.dtype != torch.int32 or src_slab.numel() != words or not src_slab.is_contiguous(): + raise ValueError( + f"k3_sandwich: src_slab must be a contiguous int32 [3, 8, {cols // 2}] slab, got {tuple(src_slab.shape)} " + f"{src_slab.dtype}" + ) + if not 0 <= src_buf < SLAB_BUFS: + raise ValueError(f"k3_sandwich: src_buf must be 0, 1 or 2, got {src_buf}") + return src_slab.reshape(-1), int(src_buf), 1 + + +def _slab_args(x_slab: Optional[torch.Tensor], slab_buf: int, fallback: torch.Tensor): + """(slab words, buffer, publish) for the kernel; a dummy view and publish 0 without a slab.""" + from . import k3_sandwich_kernel as kernel + + if x_slab is None: + return _words(fallback), 0, 0 + if ( + x_slab.dtype != torch.int32 + or x_slab.numel() != SLAB_BUFS * kernel.SLAB_WORDS + or not x_slab.is_contiguous() + ): + raise ValueError( + f"k3_sandwich: x_slab must be a contiguous int32 [3, 8, 3584] slab, " + f"got {tuple(x_slab.shape)} {x_slab.dtype}" + ) + if not 0 <= slab_buf < SLAB_BUFS: + raise ValueError(f"k3_sandwich: slab_buf must be 0, 1 or 2, got {slab_buf}") + return x_slab.reshape(-1), int(slab_buf), 1 + + +def workspace(mapping) -> dict: + """The sandwich's all-reduce buffer for ``mapping``'s TP group: ``uc`` (this rank's words), ``mc`` (their + multicast mapping) and ``flags`` (int32, the call count of each CTA). Collective on first use: every rank of + the group must make its first call at the same point, outside CUDA-graph capture.""" + ws = _workspaces.get(mapping) + if ws is not None: + return ws + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError("k3_sandwich: the all-reduce buffer must be allocated outside CUDA-graph capture") + from tensorrt_llm._torch.distributed.ops import ( + _get_mnnvl_workspace_comm, + _make_mnnvl_mcast_buffer, + _mnnvl_workspace_all_succeeded, + ) + + from . import k3_sandwich_kernel as kernel + + world = mapping.tp_size + words = kernel.buffer_words(world) + comm = _get_mnnvl_workspace_comm(mapping) + use_fabric_handle = ( + os.environ.get("TRTLLM_FORCE_MNNVL_AR", "0") == "1" or mapping.is_multi_node() + ) + error: Optional[Exception] = None + try: + handle = _make_mnnvl_mcast_buffer(comm, words * 4, mapping, use_fabric_handle) + uc = handle.get_uc_buffer(mapping.tp_rank, (words,), torch.int32, 0) + mc = handle.get_mc_buffer((words,), torch.int32, 0) + with torch.inference_mode(): + uc.fill_(EMPTY_WORD) + flags = torch.zeros(kernel.FLAG_WORDS, dtype=torch.int32, device=uc.device) + torch.cuda.synchronize() + ws = dict( + handle=handle, comm=comm, uc=uc, mc=mc, flags=flags, rank=mapping.tp_rank, world=world + ) + except Exception as exc: # noqa: BLE001 -- reported to every rank below, then re-raised + error = exc + # Also the barrier that keeps any rank from pushing into a peer's buffer before the peer has emptied it. + if not _mnnvl_workspace_all_succeeded(comm, error is None): + raise RuntimeError("k3_sandwich all-reduce buffer failed on at least one rank") from error + _workspaces[mapping] = ws + return ws + + +def latent_exchange(mapping) -> dict: + """The latent exchange of ``mapping``'s TP group, shared by k3_moe (push-only: its ``ar_uc`` / ``ar_mc`` / + ``ar_flags``) and ``k3_sandwich_tail`` (``lat_uc`` / ``lat_flags``): ``uc`` (this rank's int32 [2][8][world][1792] + words, every word empty), ``mc`` (their multicast mapping, where k3_moe pushes) and ``flags`` (int32 [64]: [0] the + tail's call count mod 6, whose parity picks the half; the tail's scale slab after it). Collective on first use, as + :func:`workspace`.""" + ex = _lat_exchanges.get(mapping) + if ex is not None: + return ex + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError("k3_sandwich: the latent exchange must be allocated outside CUDA-graph capture") + from tensorrt_llm._torch.distributed.ops import ( + _get_mnnvl_workspace_comm, + _make_mnnvl_mcast_buffer, + _mnnvl_workspace_all_succeeded, + ) + + from . import k3_sandwich_kernel as kernel + + world = mapping.tp_size + words = kernel.lat_buffer_words(world) + comm = _get_mnnvl_workspace_comm(mapping) + use_fabric_handle = ( + os.environ.get("TRTLLM_FORCE_MNNVL_AR", "0") == "1" or mapping.is_multi_node() + ) + error: Optional[Exception] = None + try: + handle = _make_mnnvl_mcast_buffer(comm, words * 4, mapping, use_fabric_handle) + uc = handle.get_uc_buffer(mapping.tp_rank, (words,), torch.int32, 0) + mc = handle.get_mc_buffer((words,), torch.int32, 0) + with torch.inference_mode(): + uc.fill_(EMPTY_WORD) + flags = torch.zeros(kernel.LAT_FLAG_WORDS, dtype=torch.int32, device=uc.device) + flags[kernel.LAT_SCALES : kernel.LAT_SCALES + kernel.LAT_SCALE_BUFS * MAX_TOKENS] = ( + kernel.SCALE_SENTINEL + ) + torch.cuda.synchronize() + ex = dict( + handle=handle, comm=comm, uc=uc, mc=mc, flags=flags, rank=mapping.tp_rank, world=world + ) + except Exception as exc: # noqa: BLE001 -- reported to every rank below, then re-raised + error = exc + # Also the barrier that keeps any rank from pushing into a peer's buffer before the peer has emptied it. + if not _mnnvl_workspace_all_succeeded(comm, error is None): + raise RuntimeError("k3_sandwich latent exchange failed on at least one rank") from error + _lat_exchanges[mapping] = ex + return ex + + +def supports(core: torch.Tensor, o_weight: torch.Tensor, num_snapshots: int) -> bool: + """Whether the kernel runs this call: bf16, M <= 8, the TP16 per-rank o_proj shape [7168, 768].""" + from . import k3_sandwich_kernel as kernel + + return ( + core.is_cuda + and core.dtype == o_weight.dtype == torch.bfloat16 + and core.dim() == 2 + and 0 < core.shape[0] <= MAX_TOKENS + and core.shape[1] == kernel.K_IN + and tuple(o_weight.shape) == (kernel.H, kernel.K_IN) + and core.is_contiguous() + and o_weight.is_contiguous() + and 0 <= num_snapshots < kernel.MAX_CANDIDATES + ) + + +def _compile_and_run(entry, key, args, runtime, consts, stream, name): + fn = _compiled.get(key) + if fn is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + f"trtllm::{name} must run once per configuration outside CUDA-graph capture first " + "(it compiles its kernel on the first call)." + ) + import cutlass.cute as cute + + with _lock: + fn = _compiled.get(key) + if fn is None: + fn = _compiled[key] = cute.compile(entry, *args, *runtime, *consts, stream) + # The compiled function takes the runtime arguments only. + fn(*args, *runtime, stream) + + +@torch.library.custom_op("trtllm::k3_sandwich_oproj", mutates_args=("x_slab",)) +def k3_sandwich_oproj( + core: torch.Tensor, + o_weight: torch.Tensor, + prefix: Optional[torch.Tensor], + block_residual: torch.Tensor, + res_weight: torch.Tensor, + rms_weight: torch.Tensor, + output_rms_weight: torch.Tensor, + rms_eps: float, + output_rms_eps: float, + ws_uc: torch.Tensor, + ws_mc: torch.Tensor, + ws_flags: torch.Tensor, + rank: int, + x_slab: Optional[torch.Tensor] = None, + slab_buf: int = 0, + src_slab: Optional[torch.Tensor] = None, + src_buf: int = 0, +) -> List[torch.Tensor]: + """``(normed, updated)`` of ``o_proj`` followed by ``allreduce_attn_res_rmsnorm``. + + ``core`` bf16 [M, 768] is this rank's attention output, ``o_weight`` [7168, 768] its o_proj slice; + ``prefix`` [M, 7168] (or None), ``block_residual`` [S, M, 7168] the S valid snapshots, the weights [7168]; + ``ws_*`` from :func:`workspace`. ``x_slab`` (from :func:`slab_tensor`) also receives normed in buffer + ``slab_buf``; ``src_slab`` (int32 [3, 8, 384]) supplies core in buffer ``src_buf``.""" + import cuda.bindings.driver as cuda_driver + + from . import k3_sandwich_kernel as kernel + + num_tokens = core.shape[0] + num_snapshots = block_residual.shape[0] + if not supports(core, o_weight, num_snapshots): + raise ValueError( + f"k3_sandwich_oproj: unsupported call core {tuple(core.shape)} {core.dtype}, o_weight " + f"{tuple(o_weight.shape)}, snapshots {num_snapshots}" + ) + world = ws_uc.numel() // kernel.buffer_words(1) + normed = torch.empty(num_tokens, kernel.H, dtype=torch.bfloat16, device=core.device) + updated = torch.empty_like(normed) + snaps = block_residual if num_snapshots > 0 else core.new_zeros(1, num_tokens, kernel.H) + add_prefix = prefix is not None + slab_words, buf, publish = _slab_args(x_slab, slab_buf, ws_flags) + src_words, sbuf, x_src = _src_args(src_slab, src_buf, kernel.K_IN, ws_flags) + args = ( + _arg(o_weight), + _arg(core), + _arg(ws_uc), + _arg(ws_mc), + _arg(ws_flags), + _arg(_words(prefix if add_prefix else snaps)), + _arg(_words(snaps)), + _arg(_words(res_weight)), + _arg(_words(rms_weight)), + _arg(_words(output_rms_weight)), + _arg(_words(updated)), + _arg(_words(normed)), + _arg(slab_words), + _arg(src_words), + ) + stream = cuda_driver.CUstream(torch.cuda.current_stream(core.device).cuda_stream) + use_pdl = _use_pdl() + key = (world, publish, x_src, use_pdl) + runtime = ( + num_tokens, + rank, + num_snapshots + 1, + int(add_prefix), + float(rms_eps), + float(output_rms_eps), + buf, + sbuf, + ) + consts = (world, publish, x_src, use_pdl) + _compile_and_run( + kernel.k3_sandwich_oproj, key, args, runtime, consts, stream, "k3_sandwich_oproj" + ) + return [normed, updated] + + +@k3_sandwich_oproj.register_fake +def _(core, o_weight, prefix, block_residual, res_weight, rms_weight, output_rms_weight, rms_eps, output_rms_eps, + ws_uc, ws_mc, ws_flags, rank, x_slab=None, slab_buf=0, src_slab=None, src_buf=0): # fmt: skip + normed = core.new_empty((core.shape[0], o_weight.shape[0]), dtype=torch.bfloat16) + return [normed, torch.empty_like(normed)] + + +def supports_tail( + latent: torch.Tensor, act: torch.Tensor, tail_weight: torch.Tensor, num_snapshots: int +) -> bool: + """Whether the pre-attention kernel runs this call: bf16, M <= 8, the TP16 shapes (latent [M, 3584], act + [M, 384], weight [7168, 256 + 384]).""" + from . import k3_sandwich_kernel as kernel + + return ( + latent.is_cuda + and latent.dtype == act.dtype == tail_weight.dtype == torch.bfloat16 + and latent.dim() == 2 + and act.dim() == 2 + and 0 < latent.shape[0] <= MAX_TOKENS + and act.shape[0] == latent.shape[0] + and latent.shape[1] == kernel.LATENT + and act.shape[1] == kernel.TAIL_ACT + and tuple(tail_weight.shape) == (kernel.H, kernel.TAIL_LAT + kernel.TAIL_ACT) + and latent.is_contiguous() + and act.is_contiguous() + and tail_weight.is_contiguous() + and 0 <= num_snapshots < kernel.MAX_CANDIDATES + ) + + +@torch.library.custom_op("trtllm::k3_sandwich_tail", mutates_args=("x_slab", "tap", "updated_out")) +def k3_sandwich_tail( + latent: torch.Tensor, + act: torch.Tensor, + tail_weight: torch.Tensor, + lo: int, + lat_eps: float, + prefix: Optional[torch.Tensor], + block_residual: torch.Tensor, + res_weight: torch.Tensor, + rms_weight: torch.Tensor, + output_rms_weight: torch.Tensor, + rms_eps: float, + output_rms_eps: float, + ws_uc: torch.Tensor, + ws_mc: torch.Tensor, + ws_flags: torch.Tensor, + rank: int, + x_slab: Optional[torch.Tensor] = None, + slab_buf: int = 0, + src_slab: Optional[torch.Tensor] = None, + src_buf: int = 0, + lat_uc: Optional[torch.Tensor] = None, + lat_flags: Optional[torch.Tensor] = None, + tap: Optional[torch.Tensor] = None, + tap_updated: bool = False, + updated_out: Optional[torch.Tensor] = None, +) -> List[torch.Tensor]: + """``(normed, updated)`` of the row-parallel MoE tail (as ``trtllm::pdl_gemv_tail``: ``[rmsnorm(latent)[:, + lo:lo+224] | act] @ tail_weight^T``) followed by ``allreduce_attn_res_rmsnorm`` of that partial. ``latent`` is + the whole reduced latent row (16-byte aligned: its rows are bulk-copied), ``act`` the shared-expert activation, + ``tail_weight`` [7168, 256 + 384] the latent up columns of the slice zero-padded to 256 and the shared down + projection; ``src_slab`` (int32 [3, 8, 1792]) supplies the latent in buffer ``src_buf``; with ``lat_uc`` / + ``lat_flags`` (:func:`latent_exchange`) the kernel sums the ranks' pushed partials itself and ``latent`` gives only + the shape; with ``tap`` (bf16 [M, 7168], unit column stride, rows a multiple of 8 elements apart, 16-byte aligned: + e.g. a column slice of a capture buffer) it also stores there the pre-norm attn_res mixture rows (a DSpark + capture layer's tap) or, with ``tap_updated``, ``updated``; with ``updated_out`` (bf16 [M, 7168], contiguous, + 16-byte aligned, e.g. the next row of the attention-residual snapshot bank, which this call does not read) it stores + ``updated`` there instead of a new tensor and returns an empty [0, 7168] in its place; the rest as + ``k3_sandwich_oproj``.""" + import cuda.bindings.driver as cuda_driver + + from . import k3_sandwich_kernel as kernel + + num_tokens = latent.shape[0] + num_snapshots = block_residual.shape[0] + if not supports_tail(latent, act, tail_weight, num_snapshots): + raise ValueError( + f"k3_sandwich_tail: unsupported call latent {tuple(latent.shape)}, act {tuple(act.shape)}, weight " + f"{tuple(tail_weight.shape)}, snapshots {num_snapshots}" + ) + world = ws_uc.numel() // kernel.buffer_words(1) + fold = lat_uc is not None + if fold: + if src_slab is not None: + raise ValueError( + "k3_sandwich_tail: the latent comes from src_slab or from the exchange, not both" + ) + if ( + lat_flags is None + or lat_flags.dtype != torch.int32 + or lat_flags.numel() != kernel.LAT_FLAG_WORDS + ): + raise ValueError(f"k3_sandwich_tail: lat_flags must be int32 [{kernel.LAT_FLAG_WORDS}]") + if (lat_uc.dtype != torch.int32 or lat_uc.numel() != kernel.lat_buffer_words(world) + or not lat_uc.is_contiguous()): # fmt: skip + raise ValueError( + f"k3_sandwich_tail: lat_uc must be the int32 [2, 8, {world}, 1792] exchange buffer" + ) + if world % 2 != 0 or not (world <= 8 or world % 8 == 0): + raise ValueError(f"k3_sandwich_tail: the latent exchange takes an even TP of at most 8 or a multiple of 8, " + f"not {world}") # fmt: skip + elif src_slab is None and latent.data_ptr() % 16 != 0: + raise ValueError( + "k3_sandwich_tail: the latent rows must be 16-byte aligned (they are bulk-copied)" + ) + if tap is not None and (tap.dtype != torch.bfloat16 or tuple(tap.shape) != (num_tokens, kernel.H) + or tap.stride(1) != 1 or tap.stride(0) % 8 != 0 or tap.data_ptr() % 16 != 0): # fmt: skip + raise ValueError(f"k3_sandwich_tail: tap must be a 16-byte aligned bf16 [{num_tokens}, {kernel.H}] view with " + f"unit column stride and a row stride that is a multiple of 8, got {tuple(tap.shape)} " + f"{tap.dtype} strides {tuple(tap.stride())}") # fmt: skip + if updated_out is not None and (updated_out.dtype != torch.bfloat16 + or tuple(updated_out.shape) != (num_tokens, kernel.H) + or not updated_out.is_contiguous() + or updated_out.data_ptr() % 16 != 0): # fmt: skip + raise ValueError(f"k3_sandwich_tail: updated_out must be a contiguous, 16-byte aligned bf16 [{num_tokens}, " + f"{kernel.H}] tensor, got {tuple(updated_out.shape)} {updated_out.dtype}") # fmt: skip + normed = torch.empty(num_tokens, kernel.H, dtype=torch.bfloat16, device=latent.device) + updated = updated_out if updated_out is not None else torch.empty_like(normed) + snaps = block_residual if num_snapshots > 0 else latent.new_zeros(1, num_tokens, kernel.H) + add_prefix = prefix is not None + slab_words, buf, publish = _slab_args(x_slab, slab_buf, ws_flags) + if fold: + src_words, sbuf, x_src = lat_uc.reshape(-1), 0, 2 + else: + src_words, sbuf, x_src = _src_args(src_slab, src_buf, kernel.LATENT, _words(latent)) + args = ( + _arg(tail_weight), + _arg(latent), + _arg(src_words), + _arg(act), + _arg(ws_uc), + _arg(ws_mc), + _arg(ws_flags), + _arg(_words(prefix if add_prefix else snaps)), + _arg(_words(snaps)), + _arg(_words(res_weight)), + _arg(_words(rms_weight)), + _arg(_words(output_rms_weight)), + _arg(_words(updated)), + _arg(_words(normed)), + _arg(slab_words), + _arg(lat_flags if fold else ws_flags), + _arg(_tap_words(tap) if tap is not None else ws_flags), + ) + stream = cuda_driver.CUstream(torch.cuda.current_stream(latent.device).cuda_stream) + use_pdl = _use_pdl() + tap_out = 0 if tap is None else 2 if tap_updated else 1 + tap_stride = tap.stride(0) // 2 if tap is not None else kernel.WORDS_PER_ROW + key = ("tail", world, publish, x_src, tap_out, use_pdl) + runtime = ( + num_tokens, rank, num_snapshots + 1, int(add_prefix), float(rms_eps), float(output_rms_eps), lo, + float(lat_eps), buf, sbuf, tap_stride, + ) # fmt: skip + consts = (world, publish, x_src, tap_out, use_pdl) + _compile_and_run( + kernel.k3_sandwich_tail, key, args, runtime, consts, stream, "k3_sandwich_tail" + ) + if updated_out is not None: + return [normed, normed.new_empty((0, kernel.H))] + return [normed, updated] + + +@k3_sandwich_tail.register_fake +def _(latent, act, tail_weight, lo, lat_eps, prefix, block_residual, res_weight, rms_weight, output_rms_weight, + rms_eps, output_rms_eps, ws_uc, ws_mc, ws_flags, rank, x_slab=None, slab_buf=0, src_slab=None, + src_buf=0, lat_uc=None, lat_flags=None, tap=None, tap_updated=False, updated_out=None): # fmt: skip + normed = latent.new_empty((latent.shape[0], tail_weight.shape[0]), dtype=torch.bfloat16) + if updated_out is not None: + return [normed, normed.new_empty((0, tail_weight.shape[0]))] + return [normed, torch.empty_like(normed)] + + +def supports_plain(x: torch.Tensor, weight: torch.Tensor, residual: torch.Tensor, norm_weight: torch.Tensor, + swiglu: bool = False) -> bool: # fmt: skip + """Whether the plain sandwich runs this call: bf16, M <= 8, a row-parallel slice [7168, K] with K a multiple of + 128 up to 896 (the drafter's TP16 o_proj: 384; its down projection: 896), residual [M, 7168]; x [M, K], or with + ``swiglu`` (K 896: one k-tile per cluster CTA) a gate_up output [M, 2 K].""" + from . import k3_sandwich_kernel as kernel + + k_in = weight.shape[1] if weight.dim() == 2 else -1 + return ( + x.is_cuda + and x.dtype == weight.dtype == residual.dtype == norm_weight.dtype == torch.bfloat16 + and x.dim() == 2 + and 0 < x.shape[0] <= MAX_TOKENS + and k_in % kernel.CTA_K == 0 + and 0 < k_in <= kernel.PLAIN_MAX_K + and (not swiglu or k_in == kernel.CLUSTER * kernel.CTA_K) + and x.shape[1] == (2 * k_in if swiglu else k_in) + and tuple(weight.shape) == (kernel.H, k_in) + and tuple(residual.shape) == (x.shape[0], kernel.H) + and tuple(norm_weight.shape) == (kernel.H,) + and x.is_contiguous() + and weight.is_contiguous() + and residual.is_contiguous() + ) + + +@torch.library.custom_op("trtllm::k3_sandwich_plain", mutates_args=()) +def k3_sandwich_plain( + x: torch.Tensor, + weight: torch.Tensor, + residual: torch.Tensor, + norm_weight: torch.Tensor, + eps: float, + ws_uc: torch.Tensor, + ws_mc: torch.Tensor, + ws_flags: torch.Tensor, + rank: int, + ipc_order: bool = False, + swiglu: bool = False, +) -> List[torch.Tensor]: + """``(normed, updated)`` of the row-parallel projection ``x @ weight^T`` followed by the TP all-reduce with the + residual add and RMSNorm (``AllReduceFusionOp.RESIDUAL_RMS_NORM``): updated = residual + the sum, + normed = RMSNorm(updated) * norm_weight. ``x`` bf16 [M, K], ``weight`` [7168, K] (K a multiple of 128 up to 896), + ``residual`` [M, 7168]; ``ws_*`` from :func:`workspace`. With ``swiglu``, ``x`` is a gate_up output [M, 2 K] + (gate columns first) and the projection is ``silu_and_mul(x) @ weight^T`` with ``k3_ctm_gemv_swiglu`` split 2's + arithmetic (the drafter MLP's down projection). The arithmetic and summation order are those of the + all-reduce kernel the call replaces: the MNNVL one-shot's, or with ``ipc_order`` the IPC one-shot's + (``allreduce_fusion_kernel_oneshot_lamport`` with fp32 accumulation, TP <= 8 within one node).""" + import cuda.bindings.driver as cuda_driver + + from . import k3_sandwich_kernel as kernel + + if not supports_plain(x, weight, residual, norm_weight, swiglu): + raise ValueError( + f"k3_sandwich_plain: unsupported call x {tuple(x.shape)} {x.dtype}, weight {tuple(weight.shape)}, " + f"residual {tuple(residual.shape)}, norm_weight {tuple(norm_weight.shape)}, swiglu {swiglu}" + ) + num_tokens, k_in = x.shape[0], weight.shape[1] + world = ws_uc.numel() // kernel.buffer_words(1) + if ipc_order and world > 8: + raise ValueError( + f"k3_sandwich_plain: the IPC one-shot's order is defined for TP <= 8, not {world}" + ) + order = 2 if ipc_order else 1 + normed = torch.empty(num_tokens, kernel.H, dtype=torch.bfloat16, device=x.device) + updated = torch.empty_like(normed) + args = ( + _arg(weight), + _arg(x), + _arg(ws_uc), + _arg(ws_mc), + _arg(ws_flags), + _arg(_words(residual)), + _arg(_words(norm_weight)), + _arg(_words(updated)), + _arg(_words(normed)), + ) + stream = cuda_driver.CUstream(torch.cuda.current_stream(x.device).cuda_stream) + use_pdl = _use_pdl() + key = ("plain", world, k_in, order, int(swiglu), use_pdl) + runtime = (num_tokens, rank, float(eps)) + consts = (world, k_in, order, int(swiglu), use_pdl) + _compile_and_run( + kernel.k3_sandwich_plain, key, args, runtime, consts, stream, "k3_sandwich_plain" + ) + return [normed, updated] + + +@k3_sandwich_plain.register_fake +def _( + x, + weight, + residual, + norm_weight, + eps, + ws_uc, + ws_mc, + ws_flags, + rank, + ipc_order=False, + swiglu=False, +): + normed = x.new_empty((x.shape[0], weight.shape[0]), dtype=torch.bfloat16) + return [normed, torch.empty_like(normed)] diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py new file mode 100644 index 000000000000..37f2d3652a4f --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py @@ -0,0 +1,341 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""trtllm::k3_fused_moe (k3_route_quant + the persistent k3_moe kernel, M <= 8) on one GPU, at the Kimi K3 TP16 +deployment's routed-expert rank layout (experts TP4 x EP4: 224 local experts, intermediate 768 per rank), at every M +in 1..8 with random routing, with 0, 4 and 16 of each token's experts local, and with 16 local experts per token none +shared (16 M groups: the kernel's group capacity at M = 8): against the stock path +(trtllm::kimi_k3_noaux_tc_mxfp8_quant then the TRTLLM-Gen W4A8_MXFP4_MXFP8 MoE runner with those ids) and an fp32 +reference over the dequantized MXFP4 experts (op-catalog gates: 8 ulp of the row max per element, 4 ulp relative RMS), +run-to-run identical bits, each M's rows within one bf16 ulp of the same rows of the 8-token call (bit-identity +reported), and the kernel's scratch (intermediate slab, layer counters) re-armed after every call. Weights are random +checkpoint-format MXFP4 experts put through TRT-LLM's own loader.""" + +import functools +import math +from types import SimpleNamespace + +import pytest +import torch + + +def _is_sm100() -> bool: + if not torch.cuda.is_available(): + return False + return torch.cuda.get_device_capability() == (10, 0) + + +pytestmark = pytest.mark.skipif(not _is_sm100(), reason="k3_fused_moe needs sm_100") + +H, TOP_K, NUM_EXPERTS, SV = 3584, 16, 896, 32 +I_TP, E_LOCAL, MOE_TP, TP_RANK, EP_RANK = 768, 224, 4, 1, 1 # one rank of experts TP4 x EP4 +OFFSET = EP_RANK * E_LOCAL +GATE_CAP, LINEAR_CAP = 4.0, 25.0 # the SiTU caps (activation_situ_beta, activation_situ_linear_beta) +RSF = 2.827 +ULP = 2.0**-8 +E4M3_MAX = 448.0 +M_ALL = list(range(1, 9)) + + +def _ops(): + import tensorrt_llm # noqa: F401 + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op # noqa: F401 + + return torch.ops.trtllm + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def _rand_mxfp4(rows, k, gen): + """Random checkpoint-format MXFP4: packed [rows, k / 2] (low nibble = even k), E8M0 per 32 k, scaled so a k-long + dot product lands near std 3.""" + codes = torch.randint(0, 16, (rows, k), dtype=torch.uint8, device="cuda", generator=gen) + base = 127 + round(0.5 * math.log2(0.01057 / k)) + exps = torch.randint(base, base + 6, (rows, k // SV), dtype=torch.uint8, device="cuda", generator=gen) + return (codes[:, 0::2] | (codes[:, 1::2] << 4)).contiguous(), exps + + +@functools.lru_cache(maxsize=None) +def _experts(seed: int = 20260928): + """This rank's experts through W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod's loader (the buffers both kernels read), + and the rank's logical slices of the checkpoint tensors (what the reference reads).""" + from tensorrt_llm._torch.moe.fused_moe.quantization import W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod + + i_full = I_TP * MOE_TP + method = W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod() + module = SimpleNamespace(tp_size=MOE_TP, tp_rank=TP_RANK, scaling_vector_size=SV, intermediate_size=i_full, + intermediate_size_per_partition=I_TP, hidden_size=H) # fmt: skip + kw = dict(dtype=torch.uint8, device="cuda") + proc = dict( + w31=torch.empty(E_LOCAL, 2 * I_TP, H // 2, **kw), + w31s=torch.empty(E_LOCAL, 2 * I_TP, H // SV, **kw), + w2=torch.empty(E_LOCAL, H, I_TP // 2, **kw), + w2s=torch.empty(E_LOCAL, H, I_TP // SV, **kw), + ) + raw = {name: [] for name in ("up", "up_s", "gate", "gate_s", "down", "down_s")} + gen = torch.Generator(device="cuda").manual_seed(seed) + lo, hi = TP_RANK * I_TP, (TP_RANK + 1) * I_TP + for e in range(E_LOCAL): + w1, w1s = _rand_mxfp4(i_full, H, gen) # gate + w3, w3s = _rand_mxfp4(i_full, H, gen) # up + w2, w2s = _rand_mxfp4(H, i_full, gen) # down + method.load_expert_w3_w1_weight(module, w1, w3, proc["w31"][e]) + method.load_expert_w2_weight(module, w2, proc["w2"][e]) + method.load_expert_w3_w1_weight_scale_mxfp4(module, w1s, w3s, proc["w31s"][e]) + method.load_expert_w2_weight_scale_mxfp4(module, w2s, proc["w2s"][e]) + raw["up"].append(w3[lo:hi]) + raw["up_s"].append(w3s[lo:hi]) + raw["gate"].append(w1[lo:hi]) + raw["gate_s"].append(w1s[lo:hi]) + raw["down"].append(w2[:, lo // 2 : hi // 2].contiguous()) + raw["down_s"].append(w2s[:, lo // SV : hi // SV].contiguous()) + torch.cuda.synchronize() + gen_b = torch.Generator(device="cuda").manual_seed(seed + 1) + bias = (torch.randn(NUM_EXPERTS, generator=gen_b, device="cuda") * 0.05).float() + return proc, raw, bias + + +_E2M1 = [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0, -0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0] + + +def _deq_w(packed, sf): + lut = torch.tensor(_E2M1, device=packed.device) + vals = torch.empty(packed.shape[0], packed.shape[1] * 2, device=packed.device) + vals[:, 0::2] = lut[(packed & 0xF).long()] + vals[:, 1::2] = lut[(packed >> 4).long()] + return vals * torch.exp2(sf.float() - 127.0).repeat_interleave(SV, dim=1) + + +def _deq_x(x_fp8, x_sf): + rows, k = x_fp8.shape + return x_fp8.float() * torch.exp2(x_sf.reshape(rows, k // SV).float() - 127.0).repeat_interleave(SV, dim=1) + + +def _requant(act): + """The FC1 epilogue's MXFP8 requantization per 32 columns (round-up scale), dequantized.""" + rows, cols = act.shape + blocks = act.reshape(rows, cols // SV, SV) + amax = blocks.abs().amax(dim=-1, keepdim=True) + ex = torch.ceil(torch.log2(amax / E4M3_MAX)) + ex = torch.where(amax == 0, torch.full_like(amax, -127.0), ex).clamp(-127.0, 127.0) + scale = torch.exp2(ex) + q8 = (blocks / scale).clamp(-E4M3_MAX, E4M3_MAX).to(torch.float8_e4m3fn) + return (q8.float() * scale).reshape(rows, cols) + + +def _reference(raw, x_deq, ids, weights): + """fp32 routed MoE over this rank's experts from the checkpoint tensors (TF32 off): SiTU, the MXFP8 intermediate, + the down projection, the routing-weighted sum.""" + allow = torch.backends.cuda.matmul.allow_tf32 + torch.backends.cuda.matmul.allow_tf32 = False + try: + out = torch.zeros(x_deq.shape[0], H, device="cuda") + for e in range(E_LOCAL): + tok, slot = (ids == OFFSET + e).nonzero(as_tuple=True) + if tok.numel() == 0: + continue + xe = x_deq[tok].double() + up = (xe @ _deq_w(raw["up"][e], raw["up_s"][e]).double().t()).float() + gate = (xe @ _deq_w(raw["gate"][e], raw["gate_s"][e]).double().t()).float() + act = GATE_CAP * torch.tanh(gate / GATE_CAP) * torch.sigmoid(gate) * (LINEAR_CAP * torch.tanh(up / LINEAR_CAP)) + y = (_requant(act).double() @ _deq_w(raw["down"][e], raw["down_s"][e]).double().t()).float() + out.index_add_(0, tok, y * weights[tok, slot].float().unsqueeze(1)) + return out + finally: + torch.backends.cuda.matmul.allow_tf32 = allow + + +def _compare(y, ref): + """Op-catalog gates: |d| <= 8 ulp of the row's max |ref| per element, relative RMS <= 4 ulp; finite.""" + o, r = y.float(), ref.float() + row = r.abs().amax(dim=1, keepdim=True).clamp_min(1e-12) + elt = ((o - r).abs() / row).max().item() / ULP + rms = ((o - r).pow(2).mean().sqrt() / r.pow(2).mean().sqrt().clamp_min(1e-12)).item() / ULP + finite = bool(torch.isfinite(o).all()) + return dict(elt_ulp=elt, rms_ulp=rms, ok=finite and elt <= 8.0 and rms <= 4.0) + + +def _stock(proc, bias, x, logits): + """The model's base path for these experts: fused route + MXFP8 quant, then the TRTLLM-Gen runner, pre-routed.""" + from tensorrt_llm._torch.moe.fused_moe.routing import RoutingMethodType + from tensorrt_llm._torch.utils import ActType_TrtllmGen + + ops = _ops() + ids, w, x_fp8, x_sf = ops.kimi_k3_noaux_tc_mxfp8_quant(logits, bias, x, RSF) + alpha = torch.full((E_LOCAL,), GATE_CAP, dtype=torch.float32, device="cuda") + beta = torch.full((E_LOCAL,), LINEAR_CAP, dtype=torch.float32, device="cuda") + y = ops.mxe4m3_mxe2m1_block_scale_moe_runner( + None, None, x_fp8, x_sf.view(-1), proc["w31"], proc["w31s"], None, alpha, beta, None, proc["w2"], + proc["w2s"], None, NUM_EXPERTS, TOP_K, 1, 1, I_TP, H, I_TP, OFFSET, E_LOCAL, 1.0, + int(RoutingMethodType.DeepSeekV3), int(ActType_TrtllmGen.SiTu), topk_weights=w, topk_ids=ids, + ) # fmt: skip + return y, ids, w, x_fp8, x_sf + + +def _fused(proc, bias, x, logits): + return _ops().k3_fused_moe(x, logits, bias, proc["w31"], proc["w31s"], proc["w2"], proc["w2s"], OFFSET, E_LOCAL, + RSF) # fmt: skip + + +def _scratch_rearmed(): + """The intermediate slab armed again (FP8 -0.0 codes, E8M0 NaN scale words) and every layer's counters zero.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op + + st = op._state(torch.device("cuda", torch.cuda.current_device()), I_TP, E_LOCAL) + mod = st.mod + cs = st.cs.view(mod.G_CAP, 8, mod.K2_TILES, mod.SFB_GROUP_BYTES) + armed = bool((st.c == -128).all()) and bool((cs[..., :4] == -1).all()) + return armed and all(bool((layer[4] == 0).all()) for layer in st.layers.values()) + + +CASES = ["random", "4_local", "16_local", "16_local_disjoint", "none_local"] + + +def _max_ulp(a: torch.Tensor, b: torch.Tensor) -> int: + """Largest distance in bf16 ulps between two bf16 tensors (bit patterns as ordered integers).""" + + def ordered(x): + i = x.contiguous().view(torch.int16).int() + return torch.where(i < 0, -(i & 0x7FFF), i) + + return int((ordered(a) - ordered(b)).abs().max().item()) if a.numel() else 0 + + +@functools.lru_cache(maxsize=None) +def _tokens(case: str, seed: int = 7): + """8 tokens: router logits and the latent rows. "_local": k of this rank's experts in each token's top-16; + "16_local_disjoint": 16 local experts per token, no two tokens sharing one (16 M groups: 128 = the kernel's group + capacity at M = 8); "none_local": no local expert.""" + salt = CASES.index(case) + gen = torch.Generator(device="cuda").manual_seed(seed + salt) + cpu = torch.Generator().manual_seed(seed + salt) + logits = (torch.randn(8, NUM_EXPERTS, generator=gen, device="cuda") * 3.0).float() + if case.endswith("_local") and case != "none_local": + k = int(case.split("_")[0]) + logits = logits - 30.0 + for t in range(8): + logits[t, OFFSET + torch.randperm(E_LOCAL, generator=cpu)[:k].cuda()] = 30.0 + elif case == "16_local_disjoint": + logits = logits - 30.0 + perm = torch.randperm(E_LOCAL, generator=cpu)[: 8 * TOP_K].view(8, TOP_K) + for t in range(8): + logits[t, OFFSET + perm[t].cuda()] = 30.0 + elif case == "none_local": + logits[:, OFFSET : OFFSET + E_LOCAL] = -30.0 + x = torch.randn(8, H, generator=gen, device="cuda").bfloat16() + return logits, x + + +@pytest.mark.parametrize("m", M_ALL) +@pytest.mark.parametrize("case", CASES) +def test_k3_fused_moe(case, m): + """Each M's rows against the 8-token call: the slice FC2 adds a token's expert terms in slices whose bounds + follow the step's group count, so a row may round differently with other tokens present (at most 1 bf16 ulp); + bit-identity is reported.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op + + _ops() + proc, raw, bias = _experts() + ok, why = op.is_supported(proc["w31"], proc["w31s"], proc["w2"], proc["w2s"], E_LOCAL) + assert ok, why + logits8, x8 = _tokens(case) + logits, x = logits8[:m].contiguous(), x8[:m].contiguous() + y = _fused(proc, bias, x, logits) + torch.cuda.synchronize() + rearmed = _scratch_rearmed() + reruns = [_fused(proc, bias, x, logits) for _ in range(2)] + y8 = _fused(proc, bias, x8, logits8) + y_stock, ids, w, x_fp8, x_sf = _stock(proc, bias, x, logits) + det = all(torch.equal(_bits(r), _bits(y)) for r in reruns) + minv = torch.equal(_bits(y), _bits(y8[:m])) + ulp_m8 = _max_ulp(y, y8[:m]) + local = int(((ids >= OFFSET) & (ids < OFFSET + E_LOCAL)).sum()) + if local == 0: + zeros = bool((y.float() == 0).all()) + print(f"OPCHECK op=k3_fused_moe case={case} M={m} zeros={zeros} det={det} rows_as_m8={minv} " + f"scratch_rearmed={rearmed}") # fmt: skip + assert zeros and det and minv and rearmed + return + ref = _reference(raw, _deq_x(x_fp8, x_sf), ids, w).bfloat16() + c_ref, c_stock, c_stock_ref = _compare(y, ref), _compare(y, y_stock), _compare(y_stock, ref) + groups = int(torch.unique(ids[(ids >= OFFSET) & (ids < OFFSET + E_LOCAL)]).numel()) + print(f"OPCHECK op=k3_fused_moe case={case} M={m} local_pairs={local} local_experts={groups} " + f"vs_ref_elt_ulp={c_ref['elt_ulp']:.2f} vs_ref_rms_ulp={c_ref['rms_ulp']:.2f} " + f"vs_stock_elt_ulp={c_stock['elt_ulp']:.2f} vs_stock_rms_ulp={c_stock['rms_ulp']:.2f} " + f"stock_vs_ref_elt_ulp={c_stock_ref['elt_ulp']:.2f} det={det} rows_as_m8={minv} max_ulp_vs_m8={ulp_m8} " + f"scratch_rearmed={rearmed}") # fmt: skip + assert c_ref["ok"] and c_stock["ok"] and c_stock_ref["ok"] + assert det and ulp_m8 <= 1 and rearmed + + +def test_k3_fused_moe_mixed_m_sequence(): + """M 8, 1, 5, 2, 8, 7 back to back on one stream: each call the bits of the same call alone (the scratch the + calls share is re-armed by each).""" + _ops() + proc, _, bias = _experts() + logits8, x8 = _tokens("random") + alone = {m: _fused(proc, bias, x8[:m].contiguous(), logits8[:m].contiguous()) for m in (1, 2, 5, 7, 8)} + seq = [(m, _fused(proc, bias, x8[:m].contiguous(), logits8[:m].contiguous())) for m in (8, 1, 5, 2, 8, 7)] + assert all(torch.equal(_bits(y), _bits(alone[m])) for m, y in seq) + assert _scratch_rearmed() + + +def test_k3_fused_moe_partial_rows_past_m(): + """The combine loads the FC2 partial rows of all 8 token slots, and a call writes only its M tokens' rows: a + fresh state's partials start zeroed, and each M's output is the same with the rows past M poisoned (NaN) as with + them zeroed, so no output reads a row its call did not write.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op + + _ops() + proc, _, bias = _experts() + fresh = op._K3FusedMoE(torch.device("cuda", torch.cuda.current_device()), I_TP, E_LOCAL) + assert bool((fresh.part == 0).all()) + st = op._state(torch.device("cuda", torch.cuda.current_device()), I_TP, E_LOCAL) + logits8, x8 = _tokens("random") + for m in M_ALL: + logits, x = logits8[:m].contiguous(), x8[:m].contiguous() + st.part.fill_(float("nan")) + y_poisoned = _fused(proc, bias, x, logits) + st.part.zero_() + y_zeroed = _fused(proc, bias, x, logits) + assert torch.equal(_bits(y_poisoned), _bits(y_zeroed)), m + + +def test_token_limit(): + _ops() + proc, _, bias = _experts() + logits = torch.zeros(9, NUM_EXPERTS, device="cuda") + x = torch.zeros(9, H, dtype=torch.bfloat16, device="cuda") + with pytest.raises(ValueError): + _fused(proc, bias, x, logits) + + +def test_collective_workspaces_refuse_graph_capture(): + """The head all-gather's buffers (the front's) and the fused all-reduce's are collective on first use (an MNNVL + multicast allocation over the TP group): a first use under CUDA-graph capture raises instead of entering the + collective, and nothing is cached for the group.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op + from tensorrt_llm.mapping import Mapping + + _ops() + mapping = Mapping(world_size=1, rank=0, tp_size=1) + graph, stream = torch.cuda.CUDAGraph(), torch.cuda.Stream() + with torch.cuda.graph(graph, stream=stream): + with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): + op.head_workspace(mapping) + with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): + op.ar_workspace(mapping) + assert mapping not in op._head_workspaces and mapping not in op._ar_workspaces diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_latent_reduce.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_latent_reduce.py new file mode 100644 index 000000000000..5a475cacdf8a --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_latent_reduce.py @@ -0,0 +1,259 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""trtllm::k3_latent_reduce (the Kimi K3 latent all-reduce as the consumer of the push-only k3_moe), one process per +GPU over the TP group of this run (4 on one GB200 tray, 16 on four), at every M in 1..8 and every CTA split: + each rank's partial is pushed as the push-only k3_moe stores it (bf16 pairs through the multicast mapping into slot + [rank] of the call's half, -0.0 as +0.0); the reduce must equal MNNVLAllReduce's one-shot of the partials bit for + bit, on every rank, run to run, with one rank's perturbed partial changing the result; after every call the whole + buffer is empty again, the call count is +1 and the arrival word 0; a CUDA graph of four push + reduce pairs, + replayed three times with new partials, matches the eager calls. + +Run under pytest (a pool of 4 MPI workers) or directly, one process per GPU: + srun -N1 -n4 --mpi=pmix python3 test_k3_latent_reduce.py +""" + +import hashlib +import os +import pickle +import sys +import traceback +from types import SimpleNamespace + +import pytest +import torch + +try: + import cloudpickle + from mpi4py import MPI +except ImportError: # the test is skipped below + cloudpickle = MPI = None + +if cloudpickle is not None: + cloudpickle.register_pickle_by_value(sys.modules[__name__]) + MPI.pickle.__init__(cloudpickle.dumps, cloudpickle.loads, pickle.HIGHEST_PROTOCOL) + +WORLD = 4 +LATENT = 3584 +M_CASES = list(range(1, 9)) +CTAS = (4, 14, 28) +EMPTY_WORD = -(2**31) + + +def _supported() -> bool: + if MPI is None or not torch.cuda.is_available() or torch.cuda.device_count() < WORLD: + return False + return torch.cuda.get_device_capability()[0] == 10 + + +pytestmark = [ + pytest.mark.threadleak(enabled=False), + pytest.mark.skipif(not _supported(), reason=f"needs {WORLD} SM100 GPUs with MNNVL and mpi4py"), +] + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def _same(a: torch.Tensor, b: torch.Tensor) -> bool: + return a.shape == b.shape and a.dtype == b.dtype and torch.equal(_bits(a), _bits(b)) + + +def _digest(t: torch.Tensor) -> str: + return hashlib.sha256(_bits(t).cpu().numpy().tobytes()).hexdigest() + + +def _context(): + os.environ.setdefault("TRTLLM_FORCE_MNNVL_AR", "1") + comm = MPI.COMM_WORLD + rank, world = comm.Get_rank(), comm.Get_size() + gpus = torch.cuda.device_count() + torch.cuda.set_device(rank % gpus) + import tensorrt_llm # noqa: F401 + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import latent_op + from tensorrt_llm._torch.distributed.ops import MNNVLAllReduce + from tensorrt_llm.mapping import Mapping + + mapping = Mapping(world_size=world, rank=rank, gpus_per_node=gpus, tp_size=world) + mnnvl = MNNVLAllReduce(mapping, torch.bfloat16) + ex = latent_op.LatentExchange(mapping) + return SimpleNamespace(comm=comm, rank=rank, world=world, mnnvl=mnnvl, ex=ex, count=0) + + +def _allreduce(ctx, x): + """MNNVLAllReduce sent one-shot (the order the reduce reproduces).""" + from tensorrt_llm._torch.distributed import AllReduceParams + + return ctx.mnnvl( + x, AllReduceParams(), one_shot_max_bytes=x.numel() * ctx.world * x.element_size() + ) + + +def _partial(m, seed, rank): + gen = torch.Generator(device="cuda").manual_seed(seed * 131 + rank) + x = (torch.randn(8, LATENT, generator=gen, device="cuda") * 0.5).bfloat16() + x[:, ::97] = -0.0 # the pushes store these as +0.0 + return x[:m].contiguous() + + +def _push(ctx, x, half): + """What the push-only k3_moe stores: the rows (-0.0 as +0.0) into slot [rank] of ``half`` of every rank's buffer, + through the multicast mapping.""" + m = x.shape[0] + words = x.clone().view(torch.int16) + words[words == -32768] = 0 + rows = ctx.ex.mc.view(2, 8, ctx.world, LATENT // 2) + rows[half, :m, ctx.rank].copy_(words.view(torch.int32).view(m, LATENT // 2)) + + +def _reduce(ctx, m, ctas): + out = torch.ops.trtllm.k3_latent_reduce(ctx.ex.uc, ctx.ex.flags, m, ctas) + ctx.count += 1 + return out + + +def _call(ctx, x, ctas): + _push(ctx, x, ctx.count & 1) + return _reduce(ctx, x.shape[0], ctas) + + +def _state_ok(ctx) -> bool: + """Every rank's last reduce done and nothing of the next call pushed yet: the whole buffer is empty, the count + advanced, the arrival word cleared.""" + torch.cuda.synchronize() + ctx.comm.Barrier() + flags = ctx.ex.flags.tolist() + ok = bool((ctx.ex.uc == EMPTY_WORD).all().item()) and flags[0] == ctx.count and flags[2] == 0 + ctx.comm.Barrier() + return ok + + +def check_reduce(ctx): + results = [] + for ctas in CTAS: + for m in M_CASES: + x = _partial(m, 1 + m + 10 * ctas, ctx.rank) + ref = _allreduce(ctx, x) + got = _call(ctx, x, ctas) + state = _state_ok(ctx) + again = [_call(ctx, x, ctas) for _ in range(2)] + bad = x.clone() + if ctx.rank == ctx.world - 1: + bad[0, 1] += 1.0 + bad_out = _call(ctx, bad, ctas) + row = dict(op="k3_latent_reduce", case=f"ctas{ctas}", M=m, exact=_same(got, ref), + det=all(_same(a, got) for a in again), control=not _same(bad_out, got), + ranks_agree=len(set(ctx.comm.allgather(_digest(got)))) == 1, + state=state and _state_ok(ctx)) # fmt: skip + row["ok"] = all(ctx.comm.allgather(all(row[k] for k in ("exact", "det", "control", "ranks_agree", + "state")))) # fmt: skip + results.append(row) + return results + + +def check_graph(ctx): + """Four push + reduce pairs (M 8, 3, 8, 1) captured in one CUDA graph, replayed three times with new partials; an + even number of pairs, so every replay starts on the half the capture pushed into.""" + results = [] + ms = (8, 3, 8, 1) + inputs = [_partial(m, 500 + i, ctx.rank) for i, m in enumerate(ms)] + outs = [None] * len(ms) + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(stream): + # eager first: every configuration compiled outside capture + for i, m in enumerate(ms): + outs[i] = _call(ctx, inputs[i], 0) + torch.cuda.synchronize() + base = ctx.count + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + for i, m in enumerate(ms): + # the halves alternate, so the captured pairs follow the eager calls' parity + _push(ctx, inputs[i], (base + i) & 1) + outs[i] = torch.ops.trtllm.k3_latent_reduce(ctx.ex.uc, ctx.ex.flags, m, 0) + torch.cuda.synchronize() + ctx.comm.Barrier() + for rep in range(3): + fresh = [_partial(m, 900 + 10 * rep + i, ctx.rank) for i, m in enumerate(ms)] + for i in range(len(ms)): + inputs[i].copy_(fresh[i]) + refs = [_allreduce(ctx, x) for x in fresh] + torch.cuda.synchronize() + ctx.comm.Barrier() + graph.replay() + ctx.count += len(ms) + torch.cuda.synchronize() + exact = all(_same(o, r) for o, r in zip(outs, refs)) + row = dict( + op="k3_latent_reduce", case=f"graph_replay{rep}", M=8, exact=exact, state=_state_ok(ctx) + ) + row["ok"] = all(ctx.comm.allgather(exact and row["state"])) + results.append(row) + del graph + return results + + +CHECKS = {"reduce": check_reduce, "graph": check_graph} + + +def _run_checks(names): + try: + ctx = _context() + with torch.inference_mode(): + return [row for name in names for row in CHECKS[name](ctx)] + except Exception: + traceback.print_exc() + raise + + +def _report(rows): + for row in rows: + fields = " ".join(f"{k}={v}" for k, v in row.items() if k not in ("op", "case", "M")) + print(f"OPCHECK op={row['op']} case={row['case']} M={row['M']} {fields}", flush=True) + + +@pytest.mark.parametrize("mpi_pool_executor", [WORLD], indirect=True) +def test_k3_latent_reduce(mpi_pool_executor): + per_rank = list(mpi_pool_executor.map(_run_checks, [list(CHECKS)] * WORLD)) + _report(per_rank[0]) + assert all(row["ok"] for rows in per_rank for row in rows) + + +def test_latent_exchange_refuses_graph_capture(): + """The exchange is collective on construction (an MNNVL multicast allocation over the TP group): constructing it + under CUDA-graph capture raises instead of entering the collective. One process, a group of one.""" + import tensorrt_llm # noqa: F401 + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import latent_op + from tensorrt_llm.mapping import Mapping + + mapping = Mapping(world_size=1, rank=0, tp_size=1) + graph, stream = torch.cuda.CUDAGraph(), torch.cuda.Stream() + with torch.cuda.graph(graph, stream=stream): + with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): + latent_op.LatentExchange(mapping) + + +def main() -> int: + names = sys.argv[1:] or list(CHECKS) + rows = _run_checks(names) + if MPI.COMM_WORLD.Get_rank() == 0: + _report(rows) + print("PASS" if all(r["ok"] for r in rows) else "FAIL", flush=True) + return 0 if all(r["ok"] for r in rows) else 1 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py new file mode 100644 index 000000000000..6d7e6292fd84 --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py @@ -0,0 +1,555 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""trtllm::k3_moe_front and trtllm::k3_fused_moe_front (the Kimi K3 MoE front: sharded head GEMV, head all-gather, +top-16 routing, MXFP8 latent, shared gate_up + SiTU; then k3_moe on its grid), one process per GPU over the TP group of +this run, at every M in 1..8. The head is sharded over the group (TP W: 3584 / W latent + 896 / W router rows and +2 x 6144 / W shared rows per rank; W = 4 on one GB200 tray, the model's TP16 shapes with 16 processes); the routed +experts are one rank of experts TP4 x EP4 (224 local experts, intermediate 768), as in the TP16 deployment. + front : against the unfused chain (the head GEMV pdl_gemv in fp32 -> the gather -> + trtllm::kimi_k3_noaux_tc_mxfp8_quant; shared: cuBLAS gate_up -> trtllm::situ_and_mul): top-16 ids per + token (a mismatch only at a reference + 16th / 17th key margin below 1e-4: the split-K head sums in another order), routing weights, MXFP8 codes and + scales (> 99.9 % equal, dequantized within one block-scale unit), the shared activation within 2e-2; the same + routing and latent bits on every rank; the head buffers empty and the buffer index flipped after each call; + fused : y against the TRTLLM-Gen W4A8_MXFP4_MXFP8 runner on the front's own routing and MXFP8 latent (op-catalog + gates), the shared activation the front's bits, k3_moe's scratch re-armed; + head_flags : fused with the ready-word handoff (k3_moe built with head_flags) across the head epoch's int32 wrap: + no word a call polls already holds the value it waits for, the plain call's bits, every ready word left + at the next call's epoch (check_head_flags); + publish_order : the handoff with k3_moe's epoch advance racing the front's epoch read: rank 1 routed nothing, k3_moe + on half the SMs, each quantization CTA held before its flags read until the epoch moves; every ready word + left at the next call's epoch (check_publish_order); +front and fused each with run-to-run identical bits and each M's rows bit-identical to the same rows of the 8-token +call (the fused y within one bf16 ulp: k3_moe's slice FC2 groups a token's expert terms by the step's group count). +The buffer and scratch checks run between barriers (peers write this rank's buffers in their next call); a failing +rank's rows are printed after rank 0's. + +Run under pytest (a pool of 4 MPI workers) or directly, one process per GPU: + srun -N1 -n4 --mpi=pmix python3 test_k3_moe_front.py [front fused head_flags publish_order] +""" + +import importlib.util +import math +import os +import pickle +import shutil +import sys +import tempfile +import time +import traceback +from types import SimpleNamespace + +import pytest +import torch +import torch.nn.functional as F + +try: + import cloudpickle + from mpi4py import MPI +except ImportError: # the test is skipped below + cloudpickle = MPI = None + +if cloudpickle is not None: + cloudpickle.register_pickle_by_value(sys.modules[__name__]) + MPI.pickle.__init__(cloudpickle.dumps, cloudpickle.loads, pickle.HIGHEST_PROTOCOL) + +WORLD = 4 +HIDDEN, LATENT, EXPERTS, TOP_K, SV = 7168, 3584, 896, 16, 32 +SHARED_INTER = 6144 # two shared experts of 3072 +GATE_CAP, LINEAR_CAP = 4.0, 25.0 +RSF = 2.827 +I_TP, E_LOCAL, MOE_TP = 768, 224, 4 # one rank of the routed experts' TP4 x EP4 +EMPTY = -(2**31) +ULP = 2.0**-8 +M_ALL = list(range(1, 9)) + + +def _supported() -> bool: + if MPI is None or not torch.cuda.is_available() or torch.cuda.device_count() < WORLD: + return False + return torch.cuda.get_device_capability() == (10, 0) + + +pytestmark = [ + pytest.mark.threadleak(enabled=False), + pytest.mark.skipif(not _supported(), reason=f"needs {WORLD} sm_100 GPUs with MNNVL and mpi4py"), +] + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.uint8) + + +def _same(a: torch.Tensor, b: torch.Tensor) -> bool: + return a.shape == b.shape and torch.equal(_bits(a), _bits(b)) + + +def _rand_mxfp4(rows, k, gen): + codes = torch.randint(0, 16, (rows, k), dtype=torch.uint8, device="cuda", generator=gen) + base = 127 + round(0.5 * math.log2(0.01057 / k)) + exps = torch.randint( + base, base + 6, (rows, k // SV), dtype=torch.uint8, device="cuda", generator=gen + ) + return (codes[:, 0::2] | (codes[:, 1::2] << 4)).contiguous(), exps + + +def _experts(seed): + """224 random MXFP4 experts through TRT-LLM's TRTLLM-Gen loader (this rank's buffers).""" + from tensorrt_llm._torch.moe.fused_moe.quantization import W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod + + i_full = I_TP * MOE_TP + method = W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod() + module = SimpleNamespace(tp_size=MOE_TP, tp_rank=1, scaling_vector_size=SV, intermediate_size=i_full, + intermediate_size_per_partition=I_TP, hidden_size=LATENT) # fmt: skip + kw = dict(dtype=torch.uint8, device="cuda") + proc = dict( + w31=torch.empty(E_LOCAL, 2 * I_TP, LATENT // 2, **kw), + w31s=torch.empty(E_LOCAL, 2 * I_TP, LATENT // SV, **kw), + w2=torch.empty(E_LOCAL, LATENT, I_TP // 2, **kw), + w2s=torch.empty(E_LOCAL, LATENT, I_TP // SV, **kw), + ) + gen = torch.Generator(device="cuda").manual_seed(seed) + for e in range(E_LOCAL): + w1, w1s = _rand_mxfp4(i_full, LATENT, gen) + w3, w3s = _rand_mxfp4(i_full, LATENT, gen) + w2, w2s = _rand_mxfp4(LATENT, i_full, gen) + method.load_expert_w3_w1_weight(module, w1, w3, proc["w31"][e]) + method.load_expert_w2_weight(module, w2, proc["w2"][e]) + method.load_expert_w3_w1_weight_scale_mxfp4(module, w1s, w3s, proc["w31s"][e]) + method.load_expert_w2_weight_scale_mxfp4(module, w2s, proc["w2s"][e]) + torch.cuda.synchronize() + return proc + + +def _context(with_experts): + os.environ.setdefault( + "TRTLLM_FORCE_MNNVL_AR", "1" + ) # fabric handles within one tray, as across trays + comm = MPI.COMM_WORLD + rank, world = comm.Get_rank(), comm.Get_size() + gpus = torch.cuda.device_count() + torch.cuda.set_device(rank % gpus) + import tensorrt_llm # noqa: F401 + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import front_op + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op as moe_op + from tensorrt_llm._torch.modules import situ # noqa: F401 (registers trtllm::situ_and_mul) + from tensorrt_llm.mapping import Mapping + + mapping = Mapping(world_size=world, rank=rank, gpus_per_node=gpus, tp_size=world) + wl, we, inter = LATENT // world, EXPERTS // world, SHARED_INTER // world + gen = torch.Generator(device="cuda").manual_seed(11) # the same on every rank + bias = (torch.randn(EXPERTS, device="cuda", generator=gen) * 0.05).float() + x8 = torch.randn(8, HIDDEN, device="cuda", generator=gen).bfloat16() + wgen = torch.Generator(device="cuda").manual_seed(111 + rank) + head = (torch.randn(wl + we, HIDDEN, device="cuda", generator=wgen) * 0.02).bfloat16() + head[wl:] *= 8.0 # router rows: logits of a few units + gate_up = (torch.randn(2 * inter, HIDDEN, device="cuda", generator=wgen) * 0.02).bfloat16() + assert front_op.weight_supported( + world, inter, HIDDEN, torch.device("cuda", torch.cuda.current_device()) + ) + ws = moe_op.head_workspace(mapping) + return SimpleNamespace( + comm=comm, rank=rank, world=world, wl=wl, we=we, inter=inter, bias=bias, x8=x8, head=head, gate_up=gate_up, + front=front_op.front_weight(head, gate_up), ws=ws, ag=(ws["uc"], ws["mc"], ws["flags"], ws["rank"]), + offset=(rank % 4) * E_LOCAL, experts=_experts(20260928 + rank) if with_experts else None, + ) # fmt: skip + + +def _all_ranks(ctx, good) -> bool: + return all(ctx.comm.allgather(bool(good))) + + +def _quiet_check(ctx, fn): + """fn() with every rank's kernels done before it and no rank's next collective call started until every rank + has run it: the checks read this rank's all-gather buffers and scratch, which peers write.""" + ctx.comm.Barrier() + value = fn() + ctx.comm.Barrier() + return value + + +def _front(ctx, x): + return torch.ops.trtllm.k3_moe_front(x, ctx.front, ctx.bias, RSF, ctx.inter, GATE_CAP, LINEAR_CAP, *ctx.ag, + ctx.world) # fmt: skip + + +def _reference(ctx, x): + """The unfused chain: this rank's head rows in fp32 (pdl_gemv), every rank's gathered (latent columns rounded to + bf16), the fused C++ routing + MXFP8 quantization; the shared expert's gate_up (cuBLAS) and SiTU-and-mul.""" + head = torch.ops.trtllm.pdl_gemv(x, ctx.head, True, False) + parts = [torch.from_numpy(a).cuda() for a in ctx.comm.allgather(head.cpu().numpy())] + latent = torch.cat([p[:, : ctx.wl] for p in parts], dim=1).bfloat16().contiguous() + logits = torch.cat([p[:, ctx.wl :] for p in parts], dim=1).contiguous() + ids, w, q, s = torch.ops.trtllm.kimi_k3_noaux_tc_mxfp8_quant(logits, ctx.bias, latent, RSF) + shared = torch.ops.trtllm.situ_and_mul(F.linear(x, ctx.gate_up), GATE_CAP, LINEAR_CAP) + key = torch.sigmoid(logits) + ctx.bias + top = key.sort(dim=1, descending=True).values + return ids, w, q, s, shared, top[:, TOP_K - 1] - top[:, TOP_K] + + +def _dequant(q, s): + scale = torch.pow(2.0, s.float() - 127.0).repeat_interleave(SV, dim=1) + return q.float() * scale, scale + + +def _buffers_empty(ctx) -> bool: + return bool((ctx.ws["uc"] == EMPTY).all()) and int(ctx.ws["flags"][1].item()) == 0 + + +def check_front(ctx): + results = [] + out8 = _front(ctx, ctx.x8) + for m in M_ALL: + x = ctx.x8[:m].contiguous() + flag0 = int(ctx.ws["flags"][0].item()) + ids, w, q, s, shared = out = _front(ctx, x) + torch.cuda.synchronize() + flag1 = int(ctx.ws["flags"][0].item()) + empty = _quiet_check(ctx, lambda: _buffers_empty(ctx)) + r_ids, r_w, r_q, r_s, r_shared, margin = _reference(ctx, x) + # The selected experts per token (their order inside the top 16 may differ at near-equal keys) and each + # selected expert's weight. + sets_equal = (ids.sort(dim=1).values == r_ids.sort(dim=1).values).all(dim=1) + near_tie = margin < 1e-4 + dense = torch.zeros(m, EXPERTS, device="cuda").scatter_(1, ids.long(), w.float()) + r_dense = torch.zeros(m, EXPERTS, device="cuda").scatter_(1, r_ids.long(), r_w.float()) + same_tok = sets_equal.nonzero().flatten() + dq, scale = _dequant(q, s) + rdq, rscale = _dequant(r_q, r_s) + again = [_front(ctx, x) for _ in range(2)] + gathered = ctx.comm.allgather([_bits(t).cpu() for t in (ids, w, q, s)]) + row = dict( + op="k3_moe_front", case=f"tp{ctx.world}", M=m, expert_sets_equal=f"{int(sets_equal.sum())}/{m}", + order_equal=f"{int((ids == r_ids).all(dim=1).sum())}/{m}", + mismatch_not_near_tie=int((~sets_equal & ~near_tie).sum()), + weight_max_err=(dense[same_tok] - r_dense[same_tok]).abs().max().item() if same_tok.numel() else 0.0, + codes_equal=(_bits(q) == _bits(r_q)).float().mean().item(), scales_equal=(s == r_s).float().mean().item(), + latent_err_scale_units=((dq - rdq).abs() / torch.maximum(scale, rscale)).max().item(), + shared_rel=((shared.float() - r_shared.float()).abs().max() / r_shared.float().abs().max()).item(), + det=all(all(_same(a, b) for a, b in zip(r, out)) for r in again), + rows_as_m8=all(_same(a, b[:m]) for a, b in zip(out, out8)), + ranks_agree=all(all(torch.equal(a, b) for a, b in zip(g, gathered[0])) for g in gathered), + buffers_empty=empty and _quiet_check(ctx, lambda: _buffers_empty(ctx)), flag_flipped=flag1 != flag0, + ) # fmt: skip + good = (row["mismatch_not_near_tie"] == 0 and row["weight_max_err"] <= 0.01 * RSF and row["codes_equal"] > 0.999 + and row["scales_equal"] > 0.999 and row["latent_err_scale_units"] <= 1.0 and row["shared_rel"] <= 2e-2 + and row["det"] and row["rows_as_m8"] and row["ranks_agree"] and row["buffers_empty"] + and row["flag_flipped"]) # fmt: skip + row["rank"], row["good"] = ctx.rank, bool(good) + row["ok"] = _all_ranks(ctx, good) + results.append(row) + return results + + +def _fused(ctx, x, ready=None): + p = ctx.experts + return torch.ops.trtllm.k3_fused_moe_front( + x, ctx.front, ctx.bias, p["w31"], p["w31s"], p["w2"], p["w2s"], ctx.offset, E_LOCAL, RSF, ctx.inter, GATE_CAP, + LINEAR_CAP, *ctx.ag, ctx.world, ag_ready=ready) # fmt: skip + + +def _runner(ctx, ids, w, q, s): + """The TRTLLM-Gen W4A8_MXFP4_MXFP8 MoE, pre-routed (the model's base path for these experts).""" + from tensorrt_llm._torch.moe.fused_moe.routing import RoutingMethodType + from tensorrt_llm._torch.utils import ActType_TrtllmGen + + p = ctx.experts + alpha = torch.full((E_LOCAL,), GATE_CAP, dtype=torch.float32, device="cuda") + beta = torch.full((E_LOCAL,), LINEAR_CAP, dtype=torch.float32, device="cuda") + return torch.ops.trtllm.mxe4m3_mxe2m1_block_scale_moe_runner( + None, None, q, s.view(-1), p["w31"], p["w31s"], None, alpha, beta, None, p["w2"], p["w2s"], None, EXPERTS, + TOP_K, 1, 1, I_TP, LATENT, I_TP, ctx.offset, E_LOCAL, 1.0, int(RoutingMethodType.DeepSeekV3), + int(ActType_TrtllmGen.SiTu), topk_weights=w, topk_ids=ids) # fmt: skip + + +def _compare(y, ref): + o, r = y.float(), ref.float() + row = r.abs().amax(dim=1, keepdim=True).clamp_min(1e-12) + elt = ((o - r).abs() / row).max().item() / ULP + rms = ((o - r).pow(2).mean().sqrt() / r.pow(2).mean().sqrt().clamp_min(1e-12)).item() / ULP + return elt, rms, bool(torch.isfinite(o).all()) and elt <= 8.0 and rms <= 4.0 + + +def _scratch_rearmed(ctx): + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op + + st = op._state(torch.device("cuda", torch.cuda.current_device()), I_TP, E_LOCAL) + mod = st.mod + cs = st.cs.view(mod.G_CAP, 8, mod.K2_TILES, mod.SFB_GROUP_BYTES) + armed = bool((st.c == -128).all()) and bool((cs[..., :4] == -1).all()) + return armed and all(bool((layer[4] == 0).all()) for layer in st.layers.values()) + + +def _max_ulp(a: torch.Tensor, b: torch.Tensor) -> int: + def ordered(x): + i = x.contiguous().view(torch.int16).int() + return torch.where(i < 0, -(i & 0x7FFF), i) + + return int((ordered(a) - ordered(b)).abs().max().item()) if a.numel() else 0 + + +def check_fused(ctx): + """y's rows against the 8-token call within one bf16 ulp: k3_moe's slice FC2 groups a token's expert terms by the + step's group count (bit-identity reported); the shared activation (per token) bit for bit.""" + results = [] + y8, sh8 = _fused(ctx, ctx.x8) + for m in M_ALL: + x = ctx.x8[:m].contiguous() + y, shared = _fused(ctx, x) + torch.cuda.synchronize() + rearmed = _quiet_check(ctx, lambda: _scratch_rearmed(ctx) and _buffers_empty(ctx)) + ids, w, q, s, f_shared = _front(ctx, x) + y_stock = _runner(ctx, ids, w, q, s) + local = int(((ids >= ctx.offset) & (ids < ctx.offset + E_LOCAL)).sum()) + again = [_fused(ctx, x) for _ in range(2)] + if local: + elt, rms, close = _compare(y, y_stock) + else: + elt, rms, close = 0.0, 0.0, bool((y.float() == 0).all()) + row = dict( + op="k3_fused_moe_front", case=f"tp{ctx.world}", M=m, local_pairs=local, vs_stock_elt_ulp=elt, + vs_stock_rms_ulp=rms, shared_eq_front=_same(shared, f_shared), + det=all(_same(a, y) and _same(b, shared) for a, b in again), + rows_as_m8=_same(y, y8[:m]) and _same(shared, sh8[:m]), max_ulp_vs_m8=_max_ulp(y, y8[:m]), + shared_rows_as_m8=_same(shared, sh8[:m]), scratch_rearmed=rearmed, + ) # fmt: skip + good = (close and row["shared_eq_front"] and row["det"] and row["max_ulp_vs_m8"] <= 1 + and row["shared_rows_as_m8"] and rearmed) # fmt: skip + row["rank"], row["good"] = ctx.rank, bool(good) + row["ok"] = _all_ranks(ctx, good) + results.append(row) + return results + + +def _i32(v: int) -> int: + return (v + 2**31) % 2**32 - 2**31 + + +def check_head_flags(ctx): + """k3_fused_moe_front with the ready-word handoff (``ag_ready``: k3_moe built with head_flags acquires the front's + ready words, ready[t] / ready[8 + t] = the head epoch flags[2] + 1 for token t, instead of waiting for its grid) + across the epoch's int32 wrap. From a new workspace's state (epoch 0, ready words 0), two calls at M 1, then the + epoch preset to -2, then calls at M 1, 8, 3, 8: the M 8 call at epoch -1 waits for 0, the value of the words that + no call has published. Per call: no word the call polls already holds its epoch + 1 (such a word would let k3_moe + read the routing before the front writes it; the call is then not run); y and the shared activation the bits of + the plain call; afterwards the epoch and every ready word hold the next call's epoch, the head buffers empty.""" + flags, ready = ctx.ws["flags"], ctx.ws["ready"] + plain = {m: _fused(ctx, ctx.x8[:m].contiguous()) for m in (1, 3, 8)} + torch.cuda.synchronize() + + def set_epoch(epoch): + flags[2] = epoch + + _quiet_check(ctx, lambda: (ready.zero_(), set_epoch(0))) + results = [] + for m in (1, 1, None, 1, 8, 3, 8): + if m is None: # two calls before the epoch reaches 0 + _quiet_check(ctx, lambda: set_epoch(-2)) + continue + x = ctx.x8[:m].contiguous() + epoch, words = _quiet_check(ctx, lambda: (int(flags[2].item()), ready[:16].tolist())) + want = _i32(epoch + 1) + pre_matched = [i for i in [*range(m), *range(8, 8 + m)] if words[i] == want] + row = dict(op="k3_fused_moe_front_ready", case=f"tp{ctx.world}", M=m, epoch=epoch, + pre_matched=pre_matched) # fmt: skip + if not _all_ranks(ctx, not pre_matched): + row["rank"], row["good"], row["ok"] = ctx.rank, False, False + results.append(row) + break + y, shared = _fused(ctx, x, ready) + torch.cuda.synchronize() + after, words = _quiet_check(ctx, lambda: (int(flags[2].item()), ready[:16].tolist())) + row.update( + epoch_advanced=after == want, ready_rearmed=all(w == want for w in words), + y_as_plain=_same(y, plain[m][0]), shared_as_plain=_same(shared, plain[m][1]), + buffers_empty=_quiet_check(ctx, lambda: _buffers_empty(ctx)), + ) # fmt: skip + good = all(row[k] for k in ("epoch_advanced", "ready_rearmed", "y_as_plain", "shared_as_plain", + "buffers_empty")) # fmt: skip + row["rank"], row["good"] = ctx.rank, bool(good) + row["ok"] = _all_ranks(ctx, good) + results.append(row) + return results + + +# The role CTAs' read of the head workspace's flags (buffer index, epoch) in k3_moe_front.py, and what +# check_publish_order inserts before it: each quantization CTA's thread 0 waits until the epoch moves, or ~20 ms. +_FLAGS_READ = ( + " if tx == 0:\n" + " s_flags.store(flags.load(idx=0, is_volatile=True), idx=0)\n" + " s_flags.store(flags.load(idx=2, is_volatile=True), idx=1)\n" +) +_QUANT_HOLD = ( + " if role >= num_tokens:\n" + " if tx == 0:\n" + " hold_epoch = flags.load(idx=2, is_volatile=True)\n" + " hold_polls = cutlass.Int32(0)\n" + " while (flags.load(idx=2, is_volatile=True) == hold_epoch) & (\n" + " hold_polls < cutlass.Int32(10000)\n" + " ):\n" + " prims.nanosleep(2000)\n" + " hold_polls = hold_polls + cutlass.Int32(1)\n" +) + + +def _held_front(tmp_dir): + """k3_moe_front.py with _QUANT_HOLD before the role CTAs' flags read, loaded from a file as a sibling module.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import front_op + + with open(os.path.join(os.path.dirname(front_op.__file__), "k3_moe_front.py")) as f: + src = f.read() + assert src.count(_FLAGS_READ) == 1, "k3_moe_front.py's role CTAs read the flags elsewhere now" + path = os.path.join(tmp_dir, "k3_moe_front_quant_hold.py") + with open(path, "w") as f: + f.write(src.replace(_FLAGS_READ, _QUANT_HOLD + _FLAGS_READ)) + name = front_op.__name__.rsplit(".", 1)[0] + ".k3_moe_front_quant_hold" + spec = importlib.util.spec_from_file_location(name, path) + mod = importlib.util.module_from_spec(spec) + sys.modules[name] = mod + spec.loader.exec_module(mod) + return mod + + +def check_publish_order(ctx, num_ctas=None): + """The ready-word handoff when k3_moe's epoch advance races the front's epoch read. k3_moe (head_flags) launches + once every front CTA has triggered and advances the head epoch at its last claim, which on a rank without a routed + expert waits for no ready word of the quantization CTAs, only for every k3_moe CTA to start. Here the routing bias + keeps every token off rank 1's experts, k3_moe runs ``num_ctas`` CTAs (default half the SMs, so that all of them + start beside the front's role CTAs; 0: one per SM, as the op builds it), and the front is k3_moe_front.py with each + quantization CTA held before its flags read until the epoch moves (_QUANT_HOLD). A role CTA that triggers before + reading the epoch then publishes the advanced epoch + 1, which the next call's poll takes for its own. After a + warm-up call, from epoch 0, calls at M 1, 8, 3; per call: no word the call polls already holds its epoch + 1 (else + the call is not run), afterwards the epoch and every ready word at the next call's epoch, the head buffers empty, + rank 1's y all zeros.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import front_op + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op as moe_op + + flags, ready = ctx.ws["flags"], ctx.ws["ready"] + idle = 1 + bias = ctx.bias.clone() + bias[idle * E_LOCAL : (idle + 1) * E_LOCAL] = -8.0 # below every other expert's selection key + device = torch.device("cuda", torch.cuda.current_device()) + ctas = torch.cuda.get_device_properties(device).multi_processor_count + if num_ctas != 0: + ctas = num_ctas or ctas // 2 + key = (device.index, I_TP, E_LOCAL, 0, True, False, False) # op._state's key of the head_flags build + saved_state, saved_kernel, saved_compiled = moe_op._states.get(key), front_op._kernel, dict(front_op._compiled) + tmp_dir = tempfile.mkdtemp(prefix="k3_moe_front_") + p = ctx.experts + + def fused(x): + return torch.ops.trtllm.k3_fused_moe_front( + x, ctx.front, bias, p["w31"], p["w31s"], p["w2"], p["w2s"], ctx.offset, E_LOCAL, RSF, ctx.inter, GATE_CAP, + LINEAR_CAP, *ctx.ag, ctx.world, ag_ready=ready) # fmt: skip + + def set_epoch(epoch): + flags[2] = epoch + + results = [] + try: + held = _held_front(tmp_dir) + moe_op._states[key] = moe_op._K3FusedMoE( + device, I_TP, E_LOCAL, {"head_flags": 1, "lat_slab": 0, "num_ctas": ctas} + ) + front_op._kernel = lambda: held + front_op._compiled.clear() + # Compiles both kernels; the ranks then start each call within the hold. + _quiet_check(ctx, lambda: (fused(ctx.x8[:1].contiguous()), torch.cuda.synchronize())) + _quiet_check(ctx, lambda: (ready.zero_(), set_epoch(0))) + for m in (1, 8, 3): + x = ctx.x8[:m].contiguous() + epoch, words = _quiet_check(ctx, lambda: (int(flags[2].item()), ready[:16].tolist())) + want = _i32(epoch + 1) + pre_matched = [i for i in [*range(m), *range(8, 8 + m)] if words[i] == want] + row = dict(op="k3_fused_moe_front_ready_order", case=f"tp{ctx.world}", M=m, num_ctas=ctas, epoch=epoch, + pre_matched=pre_matched) # fmt: skip + if not _all_ranks(ctx, not pre_matched): + row["rank"], row["good"], row["ok"] = ctx.rank, False, False + results.append(row) + break + t0 = time.perf_counter() + y, _ = fused(x) + torch.cuda.synchronize() + ms = (time.perf_counter() - t0) * 1e3 # ~20 ms: the hold ran out, no epoch moved during it + after, words = _quiet_check(ctx, lambda: (int(flags[2].item()), ready[:16].tolist())) + row.update( + epoch_advanced=after == want, off_epoch=[(i, w) for i, w in enumerate(words) if w != want], + buffers_empty=_quiet_check(ctx, lambda: _buffers_empty(ctx)), + y_zero=bool((y == 0).all()) if ctx.rank == idle else True, + call_ms=[round(v, 2) for v in ctx.comm.allgather(ms)], + ) # fmt: skip + good = row["epoch_advanced"] and not row["off_epoch"] and row["buffers_empty"] and row["y_zero"] + row["rank"], row["good"] = ctx.rank, bool(good) + row["ok"] = _all_ranks(ctx, good) + results.append(row) + finally: + front_op._kernel = saved_kernel + front_op._compiled.clear() + front_op._compiled.update(saved_compiled) + if saved_state is None: + moe_op._states.pop(key, None) + else: + moe_op._states[key] = saved_state + shutil.rmtree(tmp_dir, ignore_errors=True) + return results + + +CHECKS = { + "front": check_front, + "fused": check_fused, + "head_flags": check_head_flags, + "publish_order": check_publish_order, +} + + +def _run_checks(names): + try: + ctx = _context(with_experts=bool({"fused", "head_flags", "publish_order"} & set(names))) + with torch.inference_mode(): + return [row for name in names for row in CHECKS[name](ctx)] + except Exception: + traceback.print_exc() + raise + + +def _report(per_rank): + """Rank 0's rows, then every other rank's rows that failed there.""" + rows = list(per_rank[0]) + [row for rows in per_rank[1:] for row in rows if not row["good"]] + for row in rows: + fields = " ".join(f"{k}={(f'{v:.3e}' if isinstance(v, float) else v)}" for k, v in row.items() + if k not in ("op", "case", "M")) # fmt: skip + print(f"OPCHECK op={row['op']} case={row['case']} M={row['M']} {fields}", flush=True) + + +@pytest.mark.parametrize("mpi_pool_executor", [WORLD], indirect=True) +@pytest.mark.parametrize("check", list(CHECKS)) +def test_k3_moe_front(mpi_pool_executor, check): + per_rank = list(mpi_pool_executor.map(_run_checks, [[check]] * WORLD)) + _report(per_rank) + assert all(row["ok"] for rows in per_rank for row in rows) + + +def main() -> int: + names = sys.argv[1:] or list(CHECKS) + rows = _run_checks(names) + per_rank = MPI.COMM_WORLD.gather(rows, root=0) + if MPI.COMM_WORLD.Get_rank() == 0: + _report(per_rank) + print("PASS" if all(r["ok"] for r in rows) else "FAIL", flush=True) + return 0 if all(r["ok"] for r in rows) else 1 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_wide.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_wide.py new file mode 100644 index 000000000000..aced0e0d330d --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_wide.py @@ -0,0 +1,459 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""k3_moe for steps of up to 64 tokens (K3MoeWideState: trtllm::k3_route_quant, then the m_max 64 build of k3_moe) +on one GPU at the Kimi K3 TP16 deployment's routed-expert rank layout (experts TP4 x EP4: 224 local experts, +intermediate 768 per rank), at every M in 1..64, for these routings: +- random router logits; +- 16 local experts per token (1024 local pairs at M = 64); +- disjoint: 3 local experts per token, no two tokens sharing one; +- hot: one local expert in every token's top-16 (8 groups of it at M = 64); +- group_cap: 100 experts with 9 of the 64 tokens and 124 with one (324 groups at M = 64: the kernel's group capacity); +- none_local: no local expert; +- rtb2 (with K3_ROUTING_DUMP naming the campaign's router dump): 4 draws of 8 decode forwards of 8 tokens each, the + step's busiest EP group mapped onto this rank (M tokens: the first M of the draw). +Checks: against the stock path (trtllm::kimi_k3_noaux_tc_mxfp8_quant, then the TRTLLM-Gen W4A8_MXFP4_MXFP8 MoE runner +with those ids) and an fp64 reference over the dequantized MXFP4 experts (op-catalog gates: 8 ulp of the row max per +element, 4 ulp relative RMS); run-to-run identical bits; the slab armed and the layer's counters zero after every +call; at M <= 8, within one bf16 ulp of trtllm::k3_fused_moe (the decode build; bit-identity reported). Then: calls +of two layers at mixed M on one stream and replayed from a CUDA graph give each call's bits alone; 0 and 65 tokens are +refused. Weights are random checkpoint-format MXFP4 experts put through TRT-LLM's own loader.""" + +import functools +import math +import os +from types import SimpleNamespace + +import pytest +import torch + + +def _is_sm100() -> bool: + if not torch.cuda.is_available(): + return False + return torch.cuda.get_device_capability() == (10, 0) + + +pytestmark = pytest.mark.skipif(not _is_sm100(), reason="k3_moe needs sm_100") + +H, TOP_K, NUM_EXPERTS, SV = 3584, 16, 896, 32 +I_TP, E_LOCAL, MOE_TP, TP_RANK, EP_RANK = 768, 224, 4, 1, 1 # one rank of experts TP4 x EP4 +OFFSET = EP_RANK * E_LOCAL +GATE_CAP, LINEAR_CAP = ( + 4.0, + 25.0, +) # the SiTU caps (activation_situ_beta, activation_situ_linear_beta) +RSF = 2.827 +ULP = 2.0**-8 +E4M3_MAX = 448.0 +M_MAX = 64 +M_ALL = list(range(1, M_MAX + 1)) +ROUTING_DUMP = os.environ.get("K3_ROUTING_DUMP") + + +def _ops(): + import tensorrt_llm # noqa: F401 + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op # noqa: F401 + from tensorrt_llm._torch.cute_dsl_kernels.k3_route_quant import op as _rq # noqa: F401 + + return torch.ops.trtllm + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def _rand_mxfp4(rows, k, gen): + """Random checkpoint-format MXFP4: packed [rows, k / 2] (low nibble = even k), E8M0 per 32 k, scaled so a k-long + dot product lands near std 3.""" + codes = torch.randint(0, 16, (rows, k), dtype=torch.uint8, device="cuda", generator=gen) + base = 127 + round(0.5 * math.log2(0.01057 / k)) + exps = torch.randint( + base, base + 6, (rows, k // SV), dtype=torch.uint8, device="cuda", generator=gen + ) + return (codes[:, 0::2] | (codes[:, 1::2] << 4)).contiguous(), exps + + +@functools.lru_cache(maxsize=None) +def _experts(seed: int = 20260928): + """This rank's experts through W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod's loader (the buffers both kernels read), + and the rank's logical slices of the checkpoint tensors (what the reference reads).""" + from tensorrt_llm._torch.moe.fused_moe.quantization import W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod + + i_full = I_TP * MOE_TP + method = W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod() + module = SimpleNamespace(tp_size=MOE_TP, tp_rank=TP_RANK, scaling_vector_size=SV, intermediate_size=i_full, + intermediate_size_per_partition=I_TP, hidden_size=H) # fmt: skip + kw = dict(dtype=torch.uint8, device="cuda") + proc = dict( + w31=torch.empty(E_LOCAL, 2 * I_TP, H // 2, **kw), + w31s=torch.empty(E_LOCAL, 2 * I_TP, H // SV, **kw), + w2=torch.empty(E_LOCAL, H, I_TP // 2, **kw), + w2s=torch.empty(E_LOCAL, H, I_TP // SV, **kw), + ) + raw = {name: [] for name in ("up", "up_s", "gate", "gate_s", "down", "down_s")} + gen = torch.Generator(device="cuda").manual_seed(seed) + lo, hi = TP_RANK * I_TP, (TP_RANK + 1) * I_TP + for e in range(E_LOCAL): + w1, w1s = _rand_mxfp4(i_full, H, gen) # gate + w3, w3s = _rand_mxfp4(i_full, H, gen) # up + w2, w2s = _rand_mxfp4(H, i_full, gen) # down + method.load_expert_w3_w1_weight(module, w1, w3, proc["w31"][e]) + method.load_expert_w2_weight(module, w2, proc["w2"][e]) + method.load_expert_w3_w1_weight_scale_mxfp4(module, w1s, w3s, proc["w31s"][e]) + method.load_expert_w2_weight_scale_mxfp4(module, w2s, proc["w2s"][e]) + raw["up"].append(w3[lo:hi]) + raw["up_s"].append(w3s[lo:hi]) + raw["gate"].append(w1[lo:hi]) + raw["gate_s"].append(w1s[lo:hi]) + raw["down"].append(w2[:, lo // 2 : hi // 2].contiguous()) + raw["down_s"].append(w2s[:, lo // SV : hi // SV].contiguous()) + torch.cuda.synchronize() + gen_b = torch.Generator(device="cuda").manual_seed(seed + 1) + bias = (torch.randn(NUM_EXPERTS, generator=gen_b, device="cuda") * 0.05).float() + return proc, raw, bias + + +@functools.lru_cache(maxsize=None) +def _wide(): + """One wide state and two layers on it (the same experts; two counter sets).""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op + + proc, _, _ = _experts() + state = op.K3MoeWideState(torch.device("cuda", torch.cuda.current_device()), I_TP, E_LOCAL) + weights = (proc["w31"], proc["w31s"], proc["w2"], proc["w2s"]) + return state, state.layer(*weights), state.layer(*weights) + + +_E2M1 = [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0, -0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0] + + +def _deq_w(packed, sf): + lut = torch.tensor(_E2M1, device=packed.device) + vals = torch.empty(packed.shape[0], packed.shape[1] * 2, device=packed.device) + vals[:, 0::2] = lut[(packed & 0xF).long()] + vals[:, 1::2] = lut[(packed >> 4).long()] + return vals * torch.exp2(sf.float() - 127.0).repeat_interleave(SV, dim=1) + + +def _deq_x(x_fp8, x_sf): + rows, k = x_fp8.shape + return x_fp8.float() * torch.exp2( + x_sf.reshape(rows, k // SV).float() - 127.0 + ).repeat_interleave(SV, dim=1) + + +def _requant(act): + """The FC1 epilogue's MXFP8 requantization per 32 columns (round-up scale), dequantized.""" + rows, cols = act.shape + blocks = act.reshape(rows, cols // SV, SV) + amax = blocks.abs().amax(dim=-1, keepdim=True) + ex = torch.ceil(torch.log2(amax / E4M3_MAX)) + ex = torch.where(amax == 0, torch.full_like(amax, -127.0), ex).clamp(-127.0, 127.0) + scale = torch.exp2(ex) + q8 = (blocks / scale).clamp(-E4M3_MAX, E4M3_MAX).to(torch.float8_e4m3fn) + return (q8.float() * scale).reshape(rows, cols) + + +def _reference(raw, x_deq, ids, weights): + """fp64 routed MoE over this rank's experts from the checkpoint tensors: SiTU, the MXFP8 intermediate, the down + projection, the routing-weighted sum.""" + out = torch.zeros(x_deq.shape[0], H, device="cuda") + for e in range(E_LOCAL): + tok, slot = (ids == OFFSET + e).nonzero(as_tuple=True) + if tok.numel() == 0: + continue + xe = x_deq[tok].double() + up = (xe @ _deq_w(raw["up"][e], raw["up_s"][e]).double().t()).float() + gate = (xe @ _deq_w(raw["gate"][e], raw["gate_s"][e]).double().t()).float() + act = ( + GATE_CAP + * torch.tanh(gate / GATE_CAP) + * torch.sigmoid(gate) + * (LINEAR_CAP * torch.tanh(up / LINEAR_CAP)) + ) + y = (_requant(act).double() @ _deq_w(raw["down"][e], raw["down_s"][e]).double().t()).float() + out.index_add_(0, tok, y * weights[tok, slot].float().unsqueeze(1)) + return out + + +def _compare(y, ref): + """Op-catalog gates: |d| <= 8 ulp of the row's max |ref| per element, relative RMS <= 4 ulp; finite.""" + o, r = y.float(), ref.float() + row = r.abs().amax(dim=1, keepdim=True).clamp_min(1e-12) + elt = ((o - r).abs() / row).max().item() / ULP + rms = ((o - r).pow(2).mean().sqrt() / r.pow(2).mean().sqrt().clamp_min(1e-12)).item() / ULP + finite = bool(torch.isfinite(o).all()) + return dict(elt_ulp=elt, rms_ulp=rms, ok=finite and elt <= 8.0 and rms <= 4.0) + + +def _max_ulp(a: torch.Tensor, b: torch.Tensor) -> int: + """Largest distance in bf16 ulps between two bf16 tensors (bit patterns as ordered integers).""" + + def ordered(x): + i = x.contiguous().view(torch.int16).int() + return torch.where(i < 0, -(i & 0x7FFF), i) + + return int((ordered(a) - ordered(b)).abs().max().item()) if a.numel() else 0 + + +def _stock(proc, bias, x, logits): + """The model's base path for these experts: fused route + MXFP8 quant, then the TRTLLM-Gen runner, pre-routed.""" + from tensorrt_llm._torch.moe.fused_moe.routing import RoutingMethodType + from tensorrt_llm._torch.utils import ActType_TrtllmGen + + ops = _ops() + ids, w, x_fp8, x_sf = ops.kimi_k3_noaux_tc_mxfp8_quant(logits, bias, x, RSF) + alpha = torch.full((E_LOCAL,), GATE_CAP, dtype=torch.float32, device="cuda") + beta = torch.full((E_LOCAL,), LINEAR_CAP, dtype=torch.float32, device="cuda") + y = ops.mxe4m3_mxe2m1_block_scale_moe_runner( + None, None, x_fp8, x_sf.view(-1), proc["w31"], proc["w31s"], None, alpha, beta, None, proc["w2"], + proc["w2s"], None, NUM_EXPERTS, TOP_K, 1, 1, I_TP, H, I_TP, OFFSET, E_LOCAL, 1.0, + int(RoutingMethodType.DeepSeekV3), int(ActType_TrtllmGen.SiTu), topk_weights=w, topk_ids=ids, + ) # fmt: skip + return y, ids, w, x_fp8, x_sf + + +def _wide_call(layer, bias, x, logits, keep=None): + """The wide chain: k3_route_quant (early trigger, as k3_moe's PDL producer), then k3_moe. With keep (a list), the + route outputs are appended to it.""" + ids, w, x_fp8, x_sf = _ops().k3_route_quant(logits, bias, x, RSF, early_trigger=True) + if keep is not None: + keep.append((ids, w, x_fp8, x_sf)) + return layer(x_fp8, x_sf, ids, w, OFFSET) + + +def _scratch_rearmed(state, *layers): + """The intermediate slab armed again (FP8 -0.0 codes, E8M0 NaN scale words) and the layers' counters zero.""" + mod = state.mod + cs = state.cs.view(mod.G_CAP, 8, mod.K2_TILES, mod.SFB_GROUP_BYTES) + armed = bool((state.c == -128).all()) and bool((cs[..., :4] == -1).all()) + return armed and all(bool((layer.counters == 0).all()) for layer in layers) + + +CASES = ["random", "16_local", "disjoint", "hot", "group_cap", "none_local"] +RTB2_DRAWS = 4 + + +def _chosen_logits(chosen, gen): + """Logits whose top-16 (sigmoid + bias) is exactly each row's chosen experts: chosen ~ N(2, 0.5) (sigmoid + 0.62..0.97, so varied routing weights), the rest -20.""" + logits = torch.full((len(chosen), NUM_EXPERTS), -20.0, device="cuda") + for t, experts in enumerate(chosen): + idx = torch.tensor(sorted(experts), device="cuda") + logits[t, idx] = 2.0 + 0.5 * torch.randn(len(experts), generator=gen, device="cuda") + return logits + + +@functools.lru_cache(maxsize=None) +def _rtb2_steps(path, draws, seed=11): + """R decode forwards of 8 tokens each from the router dump (one random layer per draw), R = 8 (64 tokens): each + step's ids rotated so that its busiest EP group (most distinct experts) is this rank's.""" + import numpy as np + + z = np.load(path) + ntok = z["fwd_ntok"] + starts = np.concatenate([[0], np.cumsum(ntok)[:-1]]) + decode = [i for i, n in enumerate(ntok) if n == 8] + layers = [int(v) for v in z["layers"]] + rng = np.random.default_rng(seed) + steps = [] + for _ in range(draws): + fwds = rng.choice(decode, size=M_MAX // 8, replace=False) + toks = np.concatenate([np.arange(starts[f], starts[f] + 8) for f in fwds]) + ids = z[f"ids_{rng.choice(layers)}"][toks].astype(np.int64) + counts = np.bincount(ids.reshape(-1), minlength=NUM_EXPERTS).reshape(4, E_LOCAL) + busiest = int(np.argmax((counts > 0).sum(axis=1))) + ids = (ids + (EP_RANK - busiest) * E_LOCAL) % NUM_EXPERTS + steps.append([set(int(e) for e in row) for row in ids]) + return steps + + +@functools.lru_cache(maxsize=None) +def _tokens(case: str, seed: int = 7): + """64 tokens: router logits and the latent rows (M tokens use the first M).""" + salt = CASES.index(case) if case in CASES else len(CASES) + int(case[5:]) + gen = torch.Generator(device="cuda").manual_seed(seed + salt) + cpu = torch.Generator().manual_seed(seed + salt) + x = torch.randn(M_MAX, H, generator=gen, device="cuda").bfloat16() + local = list(range(OFFSET, OFFSET + E_LOCAL)) + remote = [e for e in range(NUM_EXPERTS) if e not in set(local)] + + def pick(pool, k): + return [pool[i] for i in torch.randperm(len(pool), generator=cpu)[:k].tolist()] + + if case == "random": + return (torch.randn(M_MAX, NUM_EXPERTS, generator=gen, device="cuda") * 3.0).float(), x + if case == "none_local": + logits = (torch.randn(M_MAX, NUM_EXPERTS, generator=gen, device="cuda") * 3.0).float() + logits[:, OFFSET : OFFSET + E_LOCAL] = -30.0 + return logits, x + if case == "16_local": + chosen = [pick(local, TOP_K) for _ in range(M_MAX)] + elif case == "disjoint": + perm = pick(local, 3 * M_MAX) + chosen = [perm[3 * t : 3 * t + 3] + pick(remote, TOP_K - 3) for t in range(M_MAX)] + elif case == "hot": + hot = local[17] + chosen = [[hot] + pick([e for e in local if e != hot], TOP_K - 1) for _ in range(M_MAX)] + elif case == "group_cap": + # 100 experts on 9 tokens and 124 on one: 1024 pairs, 2 * 100 + 124 = 324 groups. Experts are dealt in order + # of their count to the tokens with the most free slots, so every token gets 16 distinct experts. + order = pick(local, E_LOCAL) + free = [TOP_K] * M_MAX + chosen = [[] for _ in range(M_MAX)] + for i, e in enumerate(order): + need = 9 if i < 100 else 1 + toks = sorted(range(M_MAX), key=lambda t: (-free[t], t))[:need] + for t in toks: + chosen[t].append(e) + free[t] -= 1 + assert all(f == 0 for f in free) + elif case.startswith("rtb2."): + chosen = [sorted(e) for e in _rtb2_steps(ROUTING_DUMP, RTB2_DRAWS)[int(case[5:])]] + else: + raise ValueError(case) + return _chosen_logits(chosen, gen), x + + +def _groups(ids): + """The kernel's groups for these ids: sum over this rank's experts of ceil(tokens / 8).""" + local = ids[(ids >= OFFSET) & (ids < OFFSET + E_LOCAL)] + counts = torch.bincount(local - OFFSET, minlength=E_LOCAL) + return int(((counts + 7) // 8).sum()) + + +def _check(case, m): + """One call of M tokens of a routing case: against the stock path and the fp64 reference, re-runs, scratch, and + at M <= 8 the decode build.""" + _ops() + proc, raw, bias = _experts() + state, layer, _ = _wide() + logits64, x64 = _tokens(case) + logits, x = logits64[:m].contiguous(), x64[:m].contiguous() + y = _wide_call(layer, bias, x, logits) + torch.cuda.synchronize() + rearmed = _scratch_rearmed(state, layer) + reruns = [_wide_call(layer, bias, x, logits) for _ in range(2)] + det = all(torch.equal(_bits(r), _bits(y)) for r in reruns) + y_stock, ids, w, x_fp8, x_sf = _stock(proc, bias, x, logits) + if case not in ("random", "none_local"): + # The constructed logits route exactly the chosen experts. + want = torch.topk(logits, TOP_K, dim=1).indices.sort(dim=1).values + assert torch.equal(ids.long().sort(dim=1).values, want) + groups = _groups(ids) + if case == "group_cap" and m == M_MAX: + assert groups == state.mod.G_CAP == 324 + decode = "" + if m <= 8: + y8 = _ops().k3_fused_moe(x, logits, bias, proc["w31"], proc["w31s"], proc["w2"], proc["w2s"], OFFSET, + E_LOCAL, RSF) # fmt: skip + ulp_dec = _max_ulp(y, y8) + decode = f" max_ulp_vs_decode={ulp_dec} bits_as_decode={torch.equal(_bits(y), _bits(y8))}" + assert ulp_dec <= 1 + local = int(((ids >= OFFSET) & (ids < OFFSET + E_LOCAL)).sum()) + if local == 0: + zeros = bool((y.float() == 0).all()) + print(f"OPCHECK op=k3_moe_wide case={case} M={m} zeros={zeros} det={det} " + f"scratch_rearmed={rearmed}{decode}") # fmt: skip + assert zeros and det and rearmed + return + ref = _reference(raw, _deq_x(x_fp8, x_sf), ids, w).bfloat16() + c_ref, c_stock, c_stock_ref = _compare(y, ref), _compare(y, y_stock), _compare(y_stock, ref) + print(f"OPCHECK op=k3_moe_wide case={case} M={m} local_pairs={local} groups={groups} " + f"vs_ref_elt_ulp={c_ref['elt_ulp']:.2f} vs_ref_rms_ulp={c_ref['rms_ulp']:.2f} " + f"vs_stock_elt_ulp={c_stock['elt_ulp']:.2f} vs_stock_rms_ulp={c_stock['rms_ulp']:.2f} " + f"stock_vs_ref_elt_ulp={c_stock_ref['elt_ulp']:.2f} det={det} " + f"scratch_rearmed={rearmed}{decode}") # fmt: skip + assert c_ref["ok"] and c_stock["ok"] and c_stock_ref["ok"] + assert det and rearmed + + +@pytest.mark.parametrize("m", M_ALL) +@pytest.mark.parametrize("case", CASES) +def test_k3_moe_wide(case, m): + _check(case, m) + + +@pytest.mark.skipif(not ROUTING_DUMP, reason="K3_ROUTING_DUMP names no router dump") +@pytest.mark.parametrize("m", M_ALL) +@pytest.mark.parametrize("draw", range(RTB2_DRAWS)) +def test_k3_moe_wide_rtb2(draw, m): + _check(f"rtb2.{draw}", m) + + +def test_k3_moe_wide_mixed_sequence(): + """Two layers of one state at M 64, 16, 1, 40, 8, 64, 23 back to back on one stream, then the same calls replayed + from one CUDA graph with refilled inputs: each call the bits of the same call alone, the scratch re-armed.""" + _ops() + _, _, bias = _experts() + state, layer_a, layer_b = _wide() + logits64, x64 = _tokens("random") + calls = [ + (layer_a, 64), + (layer_b, 16), + (layer_a, 1), + (layer_b, 40), + (layer_a, 8), + (layer_b, 64), + (layer_a, 23), + ] + + def run(layer, m): + return _wide_call(layer, bias, x64[:m].contiguous(), logits64[:m].contiguous()) + + alone = {} + for layer, m in calls: + key = (id(layer), m) + if key not in alone: + alone[key] = run(layer, m) + torch.cuda.synchronize() + seq = [run(layer, m) for layer, m in calls] + torch.cuda.synchronize() + assert all( + torch.equal(_bits(y), _bits(alone[(id(layer), m)])) for (layer, m), y in zip(calls, seq) + ) + assert _scratch_rearmed(state, layer_a, layer_b) + + # The same sequence from a graph: inputs in static buffers, filled after capture. The route outputs are held with + # the graph, so that no allocation inside the graph reuses their addresses. + static_x = torch.zeros_like(x64) + static_logits = torch.zeros_like(logits64) + routed = [] + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + graph = torch.cuda.CUDAGraph() + with torch.cuda.stream(stream): + with torch.cuda.graph(graph, stream=stream): + outs = [_wide_call(layer, bias, static_x[:m], static_logits[:m], routed) for layer, m in calls] + static_x.copy_(x64) + static_logits.copy_(logits64) + graph.replay() + torch.cuda.synchronize() + assert all( + torch.equal(_bits(y), _bits(alone[(id(layer), m)])) for (layer, m), y in zip(calls, outs) + ) + assert _scratch_rearmed(state, layer_a, layer_b) + + +@pytest.mark.parametrize("m", [0, M_MAX + 1]) +def test_token_limit(m): + _ops() + _, layer, _ = _wide() + x_fp8 = torch.zeros(m, H, dtype=torch.float8_e4m3fn, device="cuda") + x_sf = torch.zeros(m, H // SV, dtype=torch.uint8, device="cuda") + ids = torch.zeros(m, TOP_K, dtype=torch.int32, device="cuda") + w = torch.zeros(m, TOP_K, dtype=torch.bfloat16, device="cuda") + with pytest.raises(ValueError): + layer(x_fp8, x_sf, ids, w, OFFSET) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_route_quant.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_route_quant.py new file mode 100644 index 000000000000..72f646f0ed89 --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_route_quant.py @@ -0,0 +1,161 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""trtllm::k3_route_quant (CuTe DSL top-16 routing + MXFP8 latent, M <= 64) at every M in 1..8 and at 16, 32, 64: +top-16 ids, routing weights, MXFP8 codes and scales bit for bit against trtllm::kimi_k3_noaux_tc_mxfp8_quant and +against the unfused chain (trtllm::noaux_tc_op + trtllm::mxfp8_quantize), run-to-run and early-trigger bits, each M's +rows the bits of the same rows of the 64-row call; plus ties, saturation, lane overflow and zero / large / denormal +latent rows at M = 1, 3, 8. The ids are also compared with a stable PyTorch sort of sigmoid + bias (reported).""" + +import functools + +import pytest +import torch + + +def _is_sm100() -> bool: + if not torch.cuda.is_available(): + return False + major, _ = torch.cuda.get_device_capability() + return major == 10 + + +pytestmark = pytest.mark.skipif(not _is_sm100(), reason="needs an SM100-family GPU") + +E, K, H = 896, 16, 3584 +SCALE = 2.827 +M_CASES = list(range(1, 9)) + [16, 32, 64] + + +def _ops(): + import tensorrt_llm # noqa: F401 + from tensorrt_llm._torch.cute_dsl_kernels.k3_route_quant import op # noqa: F401 + + return torch.ops.trtllm + + +def _bits(t: torch.Tensor) -> torch.Tensor: + view = {torch.bfloat16: torch.int16, torch.float32: torch.int32, torch.float8_e4m3fn: torch.uint8} + return t.contiguous().view(view.get(t.dtype, t.dtype)) + + +def _same(a: torch.Tensor, b: torch.Tensor) -> bool: + return a.shape == b.shape and a.dtype == b.dtype and torch.equal(_bits(a), _bits(b)) + + +def _torch_ids(scores, bias): + """Top-16 of sigmoid + bias, descending, ties to the lower id (stable sort of -key).""" + key = torch.sigmoid(scores) + bias + return torch.sort(-key, dim=1, stable=True).indices[:, :K].int() + + +def _unfused(scores, bias, hidden): + """The stock chain the fused C++ op replaces: noaux_tc routing, then the MXFP8 quantization of the latent.""" + ops = _ops() + weights, ids = ops.noaux_tc_op(scores, bias, 1, 1, K, SCALE) + quantized, scales = ops.mxfp8_quantize(hidden, False, alignment=256) + m = scores.shape[0] + return ids.int(), weights.to(torch.bfloat16), quantized.view(torch.float8_e4m3fn), scales.view(m, -1) + + +@functools.lru_cache(maxsize=None) +def _random(m_max: int = 64): + gen = torch.Generator(device="cuda").manual_seed(20260928) + s = torch.randn(m_max, E, generator=gen, device="cuda") * 2.5 + b = torch.randn(E, generator=gen, device="cuda") * 0.1 + h = (torch.randn(m_max, H, generator=gen, device="cuda") * 0.7).bfloat16() + return s, b, h + + +def _check(name, s, b, h): + ops = _ops() + got = ops.k3_route_quant(s, b, h, SCALE) + cpp = ops.kimi_k3_noaux_tc_mxfp8_quant(s, b, h, SCALE) + unf = _unfused(s, b, h) + vs_cpp = [_same(g, w) for g, w in zip(got, cpp)] + vs_unfused = [_same(g, w) for g, w in zip(got, unf)] + torch_ids = torch.equal(got[0], _torch_ids(s, b)) + rerun = all(_same(a, c) for a, c in zip(ops.k3_route_quant(s, b, h, SCALE), got)) + early = all(_same(a, c) for a, c in zip(ops.k3_route_quant(s, b, h, SCALE, True), got)) + print(f"OPCHECK op=k3_route_quant case={name} M={s.shape[0]} vs_cpp(ids,w,q,sf)={vs_cpp} " + f"vs_unfused(ids,w,q,sf)={vs_unfused} torch_ids={torch_ids} det={rerun} early_same={early}") # fmt: skip + return got, all(vs_cpp), all(vs_unfused), torch_ids, rerun, early + + +@pytest.mark.parametrize("m", M_CASES) +def test_k3_route_quant(m): + s64, b, h64 = _random() + s, h = s64[:m].contiguous(), h64[:m].contiguous() + got, vs_cpp, vs_unfused, torch_ids, rerun, early = _check("random", s, b, h) + full = _ops().k3_route_quant(s64, b, h64, SCALE) + rows = all(_same(g, f[:m]) for g, f in zip(got, full)) + print(f"OPCHECK op=k3_route_quant case=random M={m} rows_as_m64={rows}") + assert vs_cpp and vs_unfused and rerun and early and rows + + +def _edge_cases(): + gen = torch.Generator(device="cuda").manual_seed(20260929) + m = 8 + b = torch.randn(E, generator=gen, device="cuda") * 0.1 + h = (torch.randn(m, H, generator=gen, device="cuda") * 0.7).bfloat16() + s = torch.randn(m, E, generator=gen, device="cuda") * 2.5 + s[:, 100:140] = 3.0 + b_tied = b.clone() + b_tied[100:140] = 0.25 # 40 equal keys compete for the top 16: ties go to the lower id + yield "40_tied_keys", s, b_tied, h + yield "all_equal_logits", torch.full((m, E), 0.3, device="cuda"), torch.zeros(E, device="cuda"), h + yield "huge_logits", torch.randn(m, E, generator=gen, device="cuda") * 40.0, b, h + s = torch.randn(m, E, generator=gen, device="cuda") + s[:, 3::32] += 20.0 # every winner in one selection lane (the exact fallback) + yield "16_winners_one_lane", s, b, h + s = torch.randn(m, E, generator=gen, device="cuda") + s[:, [5, 37, 69, 101, 133]] += 20.0 + yield "5_winners_one_lane", s, b, h + h2 = h.clone() + h2[0] = 0 + h2[1, :64] = 0 + h2[2] = h2[2] * 3e4 + h2[3, ::7] = torch.tensor(1e-39).bfloat16() + h2[4, 5] = torch.tensor(-3e38).bfloat16() + yield "zero_large_denormal_rows", torch.randn(m, E, generator=gen, device="cuda"), b, h2 + yield "bias_parameter", torch.randn(m, E, generator=gen, device="cuda"), torch.nn.Parameter(b.clone()), h + + +EDGE_CASES = [ + "40_tied_keys", + "all_equal_logits", + "huge_logits", + "16_winners_one_lane", + "5_winners_one_lane", + "zero_large_denormal_rows", + "bias_parameter", +] + + +@pytest.mark.parametrize("case", EDGE_CASES) +def test_k3_route_quant_edge_cases(case): + """The PyTorch sort is reported, not asserted: its sigmoid may tie or split keys the kernels' sigmoid does not.""" + name, s, b, h = next(c for c in _edge_cases() if c[0] == case) + for m in (1, 3, 8): + _, vs_cpp, vs_unfused, _, rerun, early = _check(name, s[:m].contiguous(), b, h[:m].contiguous()) + assert vs_cpp and vs_unfused and rerun and early + + +@pytest.mark.parametrize("m", [0, 65]) +def test_token_limit(m): + ops = _ops() + s = torch.zeros(m, E, device="cuda") + h = torch.zeros(m, H, dtype=torch.bfloat16, device="cuda") + with pytest.raises(ValueError): + ops.k3_route_quant(s, torch.zeros(E, device="cuda"), h, SCALE) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py new file mode 100644 index 000000000000..d64ebed69b52 --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py @@ -0,0 +1,524 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""trtllm::k3_sandwich_oproj / k3_sandwich_tail / k3_sandwich_plain (projection + MNNVL all-reduce + residual update in +one kernel, M <= 8) at the Kimi K3 TP16 per-rank shapes, one process per GPU over the TP group of this run (4 on one +GB200 tray; the per-rank shapes do not depend on the group size, the reduction is 4-way instead of 16-way), at every M +in 1..8: + oproj against o_proj (cuBLAS, rows as in an 8-row call) -> MNNVLAllReduce.allreduce_attn_res_rmsnorm, bit for bit, + 0 / 1 / 3 / 8 snapshots, with and without the prefix sum; + tail against pdl_gemv_tail -> allreduce_attn_res_rmsnorm (fp32 tolerance: different GEMV kernels), and with the + DSpark capture tap (the pre-norm mixture against trtllm::attn_res_fwd) and updated_out (a snapshot bank row); + plain against k3_ctm_gemv -> the MNNVL one-shot RESIDUAL_RMS_NORM all-reduce, bit for bit (drafter o_proj, K 384), + and the SwiGLU form against k3_ctm_gemv_swiglu split 2 -> the same all-reduce (drafter down, K 896); +each with run-to-run identical bits, each M's rows bit-identical to the same rows of the 8-row call, every rank's +result changed by one rank's perturbed input; then CUDA graphs captured per M and replayed in mixed order with refilled +inputs, against eager calls; and the call counters (the buffer's, the folded latent all-reduce's) across the int32 +wrap. + +Run under pytest (a pool of 4 MPI workers) or directly, one process per GPU: + srun -N1 -n4 --mpi=pmix python3 test_k3_sandwich.py [oproj tail plain swiglu replay wrap fold_wrap] +""" + +import os +import pickle +import sys +import traceback +from types import SimpleNamespace + +import pytest +import torch +import torch.nn.functional as F + +try: + import cloudpickle + from mpi4py import MPI +except ImportError: # the test is skipped below + cloudpickle = MPI = None + +if cloudpickle is not None: + cloudpickle.register_pickle_by_value(sys.modules[__name__]) + MPI.pickle.__init__(cloudpickle.dumps, cloudpickle.loads, pickle.HIGHEST_PROTOCOL) + +WORLD = 4 +H, K_O, LATENT, WIDTH, PAD, ACT = 7168, 768, 3584, 224, 256, 384 +PLAIN_K, DOWN_K = 384, 896 +EPS, LAT_EPS = 1e-5, 1e-6 +M_ALL = list(range(1, 9)) +SNAPSHOTS = (0, 1, 3, 8) +REPLAY_ORDER = (8, 1, 5, 2, 7, 3, 6, 4, 8, 1, 3, 8) + + +def _supported() -> bool: + if MPI is None or not torch.cuda.is_available() or torch.cuda.device_count() < WORLD: + return False + return torch.cuda.get_device_capability()[0] == 10 + + +pytestmark = [ + pytest.mark.threadleak(enabled=False), + pytest.mark.skipif(not _supported(), reason=f"needs {WORLD} SM100 GPUs with MNNVL and mpi4py"), +] + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def _same(a: torch.Tensor, b: torch.Tensor) -> bool: + return a.shape == b.shape and torch.equal(_bits(a), _bits(b)) + + +def _rel(a: torch.Tensor, b: torch.Tensor) -> float: + return ((a.float() - b.float()).abs().max() / b.float().abs().max().clamp_min(1e-30)).item() + + +def _rand(shape, seed, scale=1.0): + gen = torch.Generator(device="cuda").manual_seed(seed) + return (torch.randn(*shape, generator=gen, device="cuda") * scale).bfloat16().contiguous() + + +def _norm_w(seed): + return (1.0 + _rand((H,), seed, 0.1).float()).bfloat16() + + +def _nan(*shape): + return torch.full(shape, float("nan"), dtype=torch.bfloat16, device="cuda") + + +def _context(): + """This process's rank, the MNNVL all-reduce of the TP group (the unfused path) and the sandwiches' workspace. + One process per GPU, ranks filling the nodes in order (gpus_per_node = the node's GPU count, so a rank's + local_rank is its device on every node). Within one tray the multicast buffers use fabric handles as across + trays.""" + os.environ.setdefault("TRTLLM_FORCE_MNNVL_AR", "1") + comm = MPI.COMM_WORLD + rank, world = comm.Get_rank(), comm.Get_size() + gpus = torch.cuda.device_count() + torch.cuda.set_device(rank % gpus) + import tensorrt_llm # noqa: F401 + from tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv import op as _ctm # noqa: F401 + from tensorrt_llm._torch.cute_dsl_kernels.k3_sandwich import op as sw_op + from tensorrt_llm._torch.distributed.ops import MNNVLAllReduce + from tensorrt_llm.mapping import Mapping + + mapping = Mapping(world_size=world, rank=rank, gpus_per_node=gpus, tp_size=world) + return SimpleNamespace(comm=comm, rank=rank, world=world, mapping=mapping, + mnnvl=MNNVLAllReduce(mapping, torch.bfloat16), ws=sw_op.workspace(mapping)) # fmt: skip + + +def _all_ranks(ctx, good) -> bool: + return all(ctx.comm.allgather(bool(good))) + + +def _oproj(ctx, core, w, prefix, block, res_w, rms_w, out_w): + ws = ctx.ws + return torch.ops.trtllm.k3_sandwich_oproj(core, w, prefix, block, res_w, rms_w, out_w, EPS, EPS, ws["uc"], + ws["mc"], ws["flags"], ws["rank"]) # fmt: skip + + +def _tail(ctx, latent, act, w, lo, prefix, block, res_w, rms_w, out_w, **extra): + ws = ctx.ws + return torch.ops.trtllm.k3_sandwich_tail(latent, act, w, lo, LAT_EPS, prefix, block, res_w, rms_w, out_w, EPS, + EPS, ws["uc"], ws["mc"], ws["flags"], ws["rank"], **extra) # fmt: skip + + +def _plain(ctx, x, w, residual, norm_w, swiglu=False): + ws = ctx.ws + return torch.ops.trtllm.k3_sandwich_plain(x, w, residual, norm_w, EPS, ws["uc"], ws["mc"], ws["flags"], + ws["rank"], swiglu=swiglu) # fmt: skip + + +def _attn_res_ar(ctx, partial, prefix, block, res_w, rms_w, out_w): + """The unfused post-projection step: the MNNVL one-shot all-reduce with the attention-residual epilogue.""" + return ctx.mnnvl.allreduce_attn_res_rmsnorm(partial, prefix, block, res_w, rms_w, out_w, EPS, EPS) + + +def _residual_rms_ar(ctx, partial, residual, norm_w): + """The unfused drafter step: the MNNVL all-reduce with the residual add + RMSNorm fusion, sent one-shot (the + sandwich reproduces the one-shot kernel's order; above 8 ranks two-shot sums the ranks in another order).""" + from tensorrt_llm._torch.distributed import AllReduceFusionOp, AllReduceParams + + params = AllReduceParams(fusion_op=AllReduceFusionOp.RESIDUAL_RMS_NORM, residual=residual, norm_weight=norm_w, + eps=EPS) # fmt: skip + out = ctx.mnnvl(partial, params, one_shot_max_bytes=partial.numel() * ctx.world * partial.element_size()) + return out[0], out[1] + + +def _row8(fn, x, m): + """fn on x padded to 8 rows, first m rows: the rows an 8-row call computes (cuBLAS picks its kernel by M).""" + pad = torch.zeros(8, x.shape[1], dtype=x.dtype, device=x.device) + pad[:m] = x + return fn(pad)[:m].contiguous() + + +def _oproj_inputs(ctx, snapshots, seed): + r = ctx.rank + core8 = _rand((8, K_O), seed + 1000 * r + 1) + w = _rand((H, K_O), seed + 1000 * r + 2, 0.03) + prefix8 = _rand((8, H), seed + 3) + block8 = _rand((snapshots, 8, H), seed + 4) + return core8, w, prefix8, block8, _rand((H,), seed + 5, 0.05), _norm_w(seed + 6), _norm_w(seed + 7) + + +def _tail_inputs(ctx, snapshots, seed): + r = ctx.rank + latent8 = _rand((8, LATENT), seed + 11, 0.8) # the reduced latent: the same on every rank + act8 = _rand((8, ACT), seed + 1000 * r + 12, 0.5) + w = _rand((H, PAD + ACT), seed + 1000 * r + 13, 0.03) + w[:, WIDTH:PAD] = 0 + prefix8 = _rand((8, H), seed + 14) + block8 = _rand((snapshots, 8, H), seed + 15) + return latent8, act8, w, r * WIDTH, prefix8, block8, _rand((H,), seed + 16, 0.05), _norm_w(seed + 17), _norm_w( + seed + 18) + + +def _first(m, prefix8, block8, with_prefix): + return (prefix8[:m].contiguous() if with_prefix else None), block8[:, :m].contiguous() + + +def _perturbed(ctx, t, col=0): + """t with one element changed on the last rank only.""" + bad = t.clone() + if ctx.rank == ctx.world - 1: + bad[0, col] += 1.0 + return bad + + +def check_oproj(ctx): + results = [] + for snapshots in SNAPSHOTS: + for with_prefix in (True, False): + core8, w, prefix8, block8, res_w, rms_w, out_w = _oproj_inputs(ctx, snapshots, 100 * snapshots) + pre8, _ = _first(8, prefix8, block8, with_prefix) + n8, u8 = _oproj(ctx, core8, w, pre8, block8, res_w, rms_w, out_w) + for m in M_ALL: + core = core8[:m].contiguous() + pre, block = _first(m, prefix8, block8, with_prefix) + n, u = _oproj(ctx, core, w, pre, block, res_w, rms_w, out_w) + want_n, want_u = _attn_res_ar(ctx, _row8(lambda x: F.linear(x, w), core, m), pre, block, res_w, rms_w, + out_w) # fmt: skip + again = [_oproj(ctx, core, w, pre, block, res_w, rms_w, out_w) for _ in range(2)] + bad_n, bad_u = _oproj(ctx, _perturbed(ctx, core), w, pre, block, res_w, rms_w, out_w) + row = dict( + op="k3_sandwich_oproj", case=f"S{snapshots}_{'prefix' if with_prefix else 'noprefix'}", M=m, + eq_unfused=_same(n, want_n) and _same(u, want_u), rel_updated=_rel(u, want_u), + rel_normed=_rel(n, want_n), det=all(_same(a, n) and _same(b, u) for a, b in again), + rows_as_m8=_same(n, n8[:m]) and _same(u, u8[:m]), control=not _same(bad_u, u), + ) # fmt: skip + row["ok"] = _all_ranks(ctx, row["eq_unfused"] and row["det"] and row["rows_as_m8"] and row["control"]) + results.append(row) + return results + + +def _tap_mixture(updated, block, res_w, rms_w): + """The unfused path's pre-norm attention-residual mixture: trtllm::attn_res_fwd on the updated row and the bank.""" + m, s = updated.shape[0], block.shape[0] + out, _, _, _ = torch.ops.trtllm.attn_res_fwd(updated.reshape(m, 1, H).contiguous(), + block.reshape(s, m, 1, H).contiguous(), res_w.reshape(-1).contiguous(), + rms_w.contiguous(), EPS) # fmt: skip + return out.reshape(m, H) + + +def check_tail(ctx): + results = [] + for snapshots in SNAPSHOTS: + for with_prefix in (True, False): + latent8, act8, w, lo, prefix8, block8, res_w, rms_w, out_w = _tail_inputs(ctx, snapshots, 50 + snapshots) + pre8, _ = _first(8, prefix8, block8, with_prefix) + n8, u8 = _tail(ctx, latent8, act8, w, lo, pre8, block8, res_w, rms_w, out_w) + for m in M_ALL: + latent, act = latent8[:m].contiguous(), act8[:m].contiguous() + pre, block = _first(m, prefix8, block8, with_prefix) + n, u = _tail(ctx, latent, act, w, lo, pre, block, res_w, rms_w, out_w) + part = torch.ops.trtllm.pdl_gemv_tail(latent, act, w, lo, WIDTH, LAT_EPS) + want_n, want_u = _attn_res_ar(ctx, part, pre, block, res_w, rms_w, out_w) + again = [_tail(ctx, latent, act, w, lo, pre, block, res_w, rms_w, out_w) for _ in range(2)] + bad_n, bad_u = _tail(ctx, latent, _perturbed(ctx, act), w, lo, pre, block, res_w, rms_w, out_w) + row = dict( + op="k3_sandwich_tail", case=f"S{snapshots}_{'prefix' if with_prefix else 'noprefix'}", M=m, + rel_updated=_rel(u, want_u), rel_normed=_rel(n, want_n), + det=all(_same(a, n) and _same(b, u) for a, b in again), rows_as_m8=_same(n, n8[:m]) and _same( + u, u8[:m]), control=not _same(bad_u, u), + ) # fmt: skip + row["ok"] = _all_ranks(ctx, row["rel_updated"] <= 8e-3 and row["rel_normed"] <= 2e-2 and row["det"] + and row["rows_as_m8"] and row["control"]) # fmt: skip + results.append(row) + # The DSpark capture tap and the snapshot bank row (updated_out), at every M. + latent8, act8, w, lo, prefix8, block8, res_w, rms_w, out_w = _tail_inputs(ctx, 3, 90) + layers = 5 + for m in M_ALL: + latent, act = latent8[:m].contiguous(), act8[:m].contiguous() + pre, block = _first(m, prefix8, block8, True) + args = (latent, act, w, lo, pre, block, res_w, rms_w, out_w) + n0, u0 = _tail(ctx, *args) + cap = _nan(m, layers * H) + tap = cap[:, 2 * H : 3 * H] + n1, u1 = _tail(ctx, *args, tap=tap) + capu = _nan(m, layers * H) + n2, u2 = _tail(ctx, *args, tap=capu[:, 4 * H :], tap_updated=True) + bank = _nan(5, m, H) + n3, u3 = _tail(ctx, *args, updated_out=bank[2]) + mix = _tap_mixture(u0, block, res_w, rms_w) + torch.cuda.synchronize() + rest_cap = torch.cat([cap[:, : 2 * H], cap[:, 3 * H :]], dim=1) + rest_bank = torch.cat([bank[:2], bank[3:]]) + frac = (_bits(tap) != _bits(mix)).float().mean().item() + row = dict( + op="k3_sandwich_tail", case="tap_updated_out", M=m, + unchanged=_same(n1, n0) and _same(u1, u0) and _same(n2, n0) and _same(u2, u0) and _same(n3, n0) + and u3.numel() == 0, tap_rel=_rel(tap, mix), tap_frac_diff=frac, + tap_updated=_same(capu[:, 4 * H :], u0), bank_row=_same(bank[2], u0), + untouched=bool(torch.isnan(rest_cap.float()).all()) and bool(torch.isnan(capu[:, : 4 * H].float()).all()) + and bool(torch.isnan(rest_bank.float()).all()), + ) # fmt: skip + row["ok"] = _all_ranks(ctx, row["unchanged"] and row["tap_rel"] <= 4e-3 and frac <= 1e-3 + and row["tap_updated"] and row["bank_row"] and row["untouched"]) # fmt: skip + results.append(row) + return results + + +def _plain_inputs(ctx, k, seed): + r = ctx.rank + x8 = _rand((8, k), seed + 1000 * r + 1) + w = _rand((H, k if k != 2 * DOWN_K else DOWN_K), seed + 1000 * r + 2, 0.03) + return x8, w, _rand((8, H), seed + 3), _norm_w(seed + 4) + + +def _check_plain(ctx, swiglu): + results = [] + name = "k3_sandwich_plain_swiglu" if swiglu else "k3_sandwich_plain" + for seed in (0, 1): + x8, w, res8, norm_w = _plain_inputs(ctx, 2 * DOWN_K if swiglu else PLAIN_K, 300 + 10 * seed + int(swiglu)) + + def gemv(x): + if swiglu: + return torch.ops.trtllm.k3_ctm_gemv_swiglu(x, w, True, 2, True) + return torch.ops.trtllm.k3_ctm_gemv(x, w, True, 1) + + n8, u8 = _plain(ctx, x8, w, res8, norm_w, swiglu) + for m in M_ALL: + x, res = x8[:m].contiguous(), res8[:m].contiguous() + n, u = _plain(ctx, x, w, res, norm_w, swiglu) + want_n, want_u = _residual_rms_ar(ctx, gemv(x), res, norm_w) + again = [_plain(ctx, x, w, res, norm_w, swiglu) for _ in range(2)] + bad_n, bad_u = _plain(ctx, _perturbed(ctx, x, DOWN_K if swiglu else 0), w, res, norm_w, swiglu) + row = dict( + op=name, case=f"seed{seed}", M=m, eq_unfused=_same(n, want_n) and _same(u, want_u), + rel_updated=_rel(u, want_u), rel_normed=_rel(n, want_n), + det=all(_same(a, n) and _same(b, u) for a, b in again), + rows_as_m8=_same(n, n8[:m]) and _same(u, u8[:m]), control=not _same(bad_u, u), + ) # fmt: skip + row["ok"] = _all_ranks(ctx, row["eq_unfused"] and row["det"] and row["rows_as_m8"] and row["control"]) + results.append(row) + return results + + +def check_plain(ctx): + return _check_plain(ctx, swiglu=False) + + +def check_swiglu(ctx): + return _check_plain(ctx, swiglu=True) + + +def check_replay(ctx): + """One graph per M of [oproj, tail, plain] on static inputs, replayed in mixed M order with refilled inputs, + against eager calls of the same ops (the engine replays the graph of each step's batch size in any order).""" + graphs = {} + stream = torch.cuda.Stream() + for m in M_ALL: + core8, w_o, prefix8, block8, res_w, rms_w, out_w = _oproj_inputs(ctx, 3, 700) + latent8, act8, w_t, lo, t_prefix8, t_block8, t_res, t_rms, t_out = _tail_inputs(ctx, 3, 710) + x8, w_p, res8, norm_w = _plain_inputs(ctx, PLAIN_K, 720) + a_in = [core8[:m].clone(), w_o, prefix8[:m].clone(), block8[:, :m].clone(), res_w, rms_w, out_w] + b_in = [latent8[:m].clone(), act8[:m].clone(), w_t, lo, t_prefix8[:m].clone(), t_block8[:, :m].clone(), t_res, + t_rms, t_out] # fmt: skip + c_in = [x8[:m].clone(), w_p, res8[:m].clone(), norm_w] + _oproj(ctx, *a_in) + _tail(ctx, *b_in) + _plain(ctx, *c_in) + torch.cuda.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.stream(stream): + with torch.cuda.graph(graph, stream=stream): + outs = [_oproj(ctx, *a_in), _tail(ctx, *b_in), _plain(ctx, *c_in)] + torch.cuda.synchronize() + graphs[m] = (graph, a_in, b_in, c_in, outs) + results = [] + for it, m in enumerate(REPLAY_ORDER): + graph, a_in, b_in, c_in, outs = graphs[m] + seed = 9000 + 97 * it + a_in[0].copy_(_rand(tuple(a_in[0].shape), seed + ctx.rank)) + a_in[2].copy_(_rand(tuple(a_in[2].shape), seed + 1000)) + a_in[3].copy_(_rand(tuple(a_in[3].shape), seed + 2000)) + b_in[0].copy_(_rand(tuple(b_in[0].shape), seed + 3000, 0.8)) + b_in[1].copy_(_rand(tuple(b_in[1].shape), seed + 4000 + ctx.rank, 0.5)) + b_in[4].copy_(_rand(tuple(b_in[4].shape), seed + 5000)) + c_in[0].copy_(_rand(tuple(c_in[0].shape), seed + 6000 + ctx.rank)) + c_in[2].copy_(_rand(tuple(c_in[2].shape), seed + 7000)) + torch.cuda.synchronize() + ctx.comm.Barrier() + graph.replay() + torch.cuda.synchronize() + got = [[t.clone() for t in o] for o in outs] + want = [_oproj(ctx, *a_in), _tail(ctx, *b_in), _plain(ctx, *c_in)] + torch.cuda.synchronize() + same = all(_same(g, x) for go, wo in zip(got, want) for g, x in zip(go, wo)) + results.append(dict(op="k3_sandwich_replay", case=f"replay{it}", M=m, eq_eager=same, ok=_all_ranks(ctx, same))) + del graphs + return results + + +def check_wrap(ctx): + """The all-reduce buffer's per-CTA call counters (``flags``, whose parity picks the buffer half) across the int32 + wrap: a sequence of oproj / tail / plain calls from counters preset just below 2**31 gives the bits of the same + sequence from the counters as they were. Every CTA counts every call, so the counters stay equal; the preset keeps + their parity, which decides the half the previous call emptied.""" + core8, w_o, prefix8, block8, res_w, rms_w, out_w = _oproj_inputs(ctx, 3, 800) + latent8, act8, w_t, lo, t_prefix8, t_block8, t_res, t_rms, t_out = _tail_inputs(ctx, 3, 810) + x8, w_p, res8, norm_w = _plain_inputs(ctx, PLAIN_K, 820) + calls = [] + for m in (8, 1, 5): + calls += [ + lambda m=m: _oproj(ctx, core8[:m].contiguous(), w_o, prefix8[:m].contiguous(), block8[:, :m].contiguous(), + res_w, rms_w, out_w), + lambda m=m: _tail(ctx, latent8[:m].contiguous(), act8[:m].contiguous(), w_t, lo, + t_prefix8[:m].contiguous(), t_block8[:, :m].contiguous(), t_res, t_rms, t_out), + lambda m=m: _plain(ctx, x8[:m].contiguous(), w_p, res8[:m].contiguous(), norm_w), + ] # fmt: skip + + def run(): + outs = [[t.clone() for t in call()] for call in calls] + torch.cuda.synchronize() + return outs + + flags = ctx.ws["flags"] + fresh = run() + count = int(flags[0].item()) + ctx.comm.Barrier() + flags.fill_(2**31 - 4 + (count & 1)) + torch.cuda.synchronize() + ctx.comm.Barrier() + wrapped = run() + after = int(flags[0].item()) + same = all(_same(a, b) for fo, wo in zip(fresh, wrapped) for a, b in zip(fo, wo)) + row = dict(op="k3_sandwich_wrap", case="int32_wrap", M=8, eq_fresh=same, crossed=after < 0, calls=len(calls)) + row["ok"] = _all_ranks(ctx, same and row["crossed"]) + return [row] + + +def check_fold_wrap(ctx): + """The tail with the latent all-reduce folded in (``lat_uc`` / ``lat_flags``; every rank pushes its partial rows + into slot [rank] of half ``n & 1`` of every rank's exchange buffer through the multicast mapping, as k3_moe does) + across the wrap of its call count n: the same calls from n preset just below 2**31 give the bits of the calls from + n as it was, and no call writes a ``lat_flags`` word other than the count and the scale slab.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_sandwich import k3_sandwich_kernel as kernel + from tensorrt_llm._torch.cute_dsl_kernels.k3_sandwich import op as sw_op + + ex = sw_op.latent_exchange(ctx.mapping) + flags = ex["flags"] + lanes = ex["mc"].view(2, 8, ctx.world, LATENT // 2) + latent8, act8, w, lo, prefix8, block8, res_w, rms_w, out_w = _tail_inputs(ctx, 3, 830) + part8 = _rand((8, LATENT), 840 + 1000 * ctx.rank, 0.3) + part8[part8 == 0] = 0.0 # pushes never send -0.0: with +0.0 beside it, that word would read as empty + slab = slice(kernel.LAT_SCALES, kernel.LAT_SCALES + kernel.LAT_SCALE_BUFS * 8) + + def run(): + outs, others = [], [] + for m in (8, 1, 5, 3, 8, 2): + torch.cuda.synchronize() + ctx.comm.Barrier() # every rank's previous call has emptied the half this call's pushes go to + lanes[int(flags[0].item()) & 1, :m, ctx.rank].copy_(part8[:m].view(torch.int32)) + torch.cuda.synchronize() + ctx.comm.Barrier() + before = torch.cat([flags[1 : slab.start], flags[slab.stop :]]).clone() + pre, block = _first(m, prefix8, block8, True) + out = _tail(ctx, latent8[:m].contiguous(), act8[:m].contiguous(), w, lo, pre, block, res_w, rms_w, out_w, + lat_uc=ex["uc"], lat_flags=flags) # fmt: skip + torch.cuda.synchronize() + outs.append([t.clone() for t in out]) + others.append(torch.equal(before, torch.cat([flags[1 : slab.start], flags[slab.stop :]]))) + return outs, others + + fresh, fresh_kept = run() + count = int(flags[0].item()) + ctx.comm.Barrier() + flags[0] = 2**31 - 4 + (count & 1) + torch.cuda.synchronize() + wrapped, wrapped_kept = run() + same = all(_same(a, b) for fo, wo in zip(fresh, wrapped) for a, b in zip(fo, wo)) + kept = all(fresh_kept) and all(wrapped_kept) + row = dict(op="k3_sandwich_tail_fold", case="count_wrap", M=8, eq_fresh=same, other_words_kept=kept, + count_after=int(flags[0].item())) # fmt: skip + row["ok"] = _all_ranks(ctx, same and kept) + return [row] + + +CHECKS = {"oproj": check_oproj, "tail": check_tail, "plain": check_plain, "swiglu": check_swiglu, + "replay": check_replay, "wrap": check_wrap, "fold_wrap": check_fold_wrap} # fmt: skip + + +def _run_checks(names): + """Every rank runs the same checks in the same order (they are collectives); returns this rank's result rows.""" + try: + ctx = _context() + with torch.inference_mode(): + return [row for name in names for row in CHECKS[name](ctx)] + except Exception: + traceback.print_exc() + raise + + +def _report(rows): + for row in rows: + fields = " ".join(f"{k}={(f'{v:.3e}' if isinstance(v, float) else v)}" for k, v in row.items() + if k not in ("op", "case", "M")) # fmt: skip + print(f"OPCHECK op={row['op']} case={row['case']} M={row['M']} {fields}", flush=True) + + +@pytest.mark.parametrize("mpi_pool_executor", [WORLD], indirect=True) +@pytest.mark.parametrize("check", list(CHECKS)) +def test_k3_sandwich(mpi_pool_executor, check): + per_rank = list(mpi_pool_executor.map(_run_checks, [[check]] * WORLD)) + _report(per_rank[0]) + assert all(row["ok"] for rows in per_rank for row in rows) + + +def test_workspaces_refuse_graph_capture(): + """The all-reduce buffer and the latent exchange are collective on first use: allocating either under CUDA-graph + capture raises instead of entering the collective, which could hang the group.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_sandwich import op + from tensorrt_llm.mapping import Mapping + + mapping = Mapping(world_size=1, rank=0, tp_size=1) + graph, stream = torch.cuda.CUDAGraph(), torch.cuda.Stream() + with torch.cuda.graph(graph, stream=stream): + with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): + op.workspace(mapping) + with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): + op.latent_exchange(mapping) + + +def main() -> int: + names = sys.argv[1:] or list(CHECKS) + rows = _run_checks(names) + if MPI.COMM_WORLD.Get_rank() == 0: + _report(rows) + print("PASS" if all(r["ok"] for r in rows) else "FAIL", flush=True) + return 0 if all(r["ok"] for r in rows) else 1 + + +if __name__ == "__main__": + sys.exit(main()) From 5181297081660b36f0ce36a13a0f533da0d28d71 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 18:03:46 -0700 Subject: [PATCH 057/161] [None][test] Kimi K3 MoE front and tail sandwich tests: torch references main has no trtllm::pdl_gemv (the K3 stack retires it at CTM gate 3, d461ee592d), so the two tests that used it as a reference compute it in torch instead: - test_k3_moe_front: the head GEMV of the unfused chain is an fp32 torch product; - test_k3_sandwich: the tail's reference partial is computed with the kernel's arithmetic (fp32 accumulators of the latent slice and of the activation, the latent one scaled by the RMS of the whole latent row, one bf16 rounding), compared with the fp32 tolerance as before. The k3_sandwich_tail docstring states the arithmetic without naming pdl_gemv_tail. These hunks are d461ee592d's for these files. Signed-off-by: Vasanth Sabavat --- .../_torch/cute_dsl_kernels/k3_sandwich/op.py | 4 ++-- .../kimi_k3/test_k3_moe_front.py | 6 +++--- .../cute_dsl_kernels/kimi_k3/test_k3_sandwich.py | 16 +++++++++++++--- 3 files changed, 18 insertions(+), 8 deletions(-) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py index d78692dbd862..03661fea8ed5 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py @@ -395,8 +395,8 @@ def k3_sandwich_tail( tap_updated: bool = False, updated_out: Optional[torch.Tensor] = None, ) -> List[torch.Tensor]: - """``(normed, updated)`` of the row-parallel MoE tail (as ``trtllm::pdl_gemv_tail``: ``[rmsnorm(latent)[:, - lo:lo+224] | act] @ tail_weight^T``) followed by ``allreduce_attn_res_rmsnorm`` of that partial. ``latent`` is + """``(normed, updated)`` of the row-parallel MoE tail (``[rmsnorm(latent)[:, lo:lo+224] | act] @ tail_weight^T``, + the RMS on the fp32 latent accumulator) followed by ``allreduce_attn_res_rmsnorm`` of that partial. ``latent`` is the whole reduced latent row (16-byte aligned: its rows are bulk-copied), ``act`` the shared-expert activation, ``tail_weight`` [7168, 256 + 384] the latent up columns of the slice zero-padded to 256 and the shared down projection; ``src_slab`` (int32 [3, 8, 1792]) supplies the latent in buffer ``src_buf``; with ``lat_uc`` / diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py index 6d7e6292fd84..7ec33b188ace 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py @@ -17,7 +17,7 @@ this run, at every M in 1..8. The head is sharded over the group (TP W: 3584 / W latent + 896 / W router rows and 2 x 6144 / W shared rows per rank; W = 4 on one GB200 tray, the model's TP16 shapes with 16 processes); the routed experts are one rank of experts TP4 x EP4 (224 local experts, intermediate 768), as in the TP16 deployment. - front : against the unfused chain (the head GEMV pdl_gemv in fp32 -> the gather -> + front : against the unfused chain (the head GEMV in fp32 torch -> the gather -> trtllm::kimi_k3_noaux_tc_mxfp8_quant; shared: cuBLAS gate_up -> trtllm::situ_and_mul): top-16 ids per token (a mismatch only at a reference 16th / 17th key margin below 1e-4: the split-K head sums in another order), routing weights, MXFP8 codes and @@ -186,9 +186,9 @@ def _front(ctx, x): def _reference(ctx, x): - """The unfused chain: this rank's head rows in fp32 (pdl_gemv), every rank's gathered (latent columns rounded to + """The unfused chain: this rank's head rows in fp32 (torch), every rank's gathered (latent columns rounded to bf16), the fused C++ routing + MXFP8 quantization; the shared expert's gate_up (cuBLAS) and SiTU-and-mul.""" - head = torch.ops.trtllm.pdl_gemv(x, ctx.head, True, False) + head = x.float() @ ctx.head.float().t() parts = [torch.from_numpy(a).cuda() for a in ctx.comm.allgather(head.cpu().numpy())] latent = torch.cat([p[:, : ctx.wl] for p in parts], dim=1).bfloat16().contiguous() logits = torch.cat([p[:, ctx.wl :] for p in parts], dim=1).contiguous() diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py index d64ebed69b52..f5f4c6d26a69 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py @@ -18,8 +18,9 @@ in 1..8: oproj against o_proj (cuBLAS, rows as in an 8-row call) -> MNNVLAllReduce.allreduce_attn_res_rmsnorm, bit for bit, 0 / 1 / 3 / 8 snapshots, with and without the prefix sum; - tail against pdl_gemv_tail -> allreduce_attn_res_rmsnorm (fp32 tolerance: different GEMV kernels), and with the - DSpark capture tap (the pre-norm mixture against trtllm::attn_res_fwd) and updated_out (a snapshot bank row); + tail against the tail in torch (fp32 accumulators, the latent RMS on the latent one) -> allreduce_attn_res_rmsnorm + (fp32 tolerance: another summation order), and with the DSpark capture tap (the pre-norm mixture against + trtllm::attn_res_fwd) and updated_out (a snapshot bank row); plain against k3_ctm_gemv -> the MNNVL one-shot RESIDUAL_RMS_NORM all-reduce, bit for bit (drafter o_proj, K 384), and the SwiGLU form against k3_ctm_gemv_swiglu split 2 -> the same all-reduce (drafter down, K 896); each with run-to-run identical bits, each M's rows bit-identical to the same rows of the 8-row call, every rank's @@ -145,6 +146,15 @@ def _attn_res_ar(ctx, partial, prefix, block, res_w, rms_w, out_w): return ctx.mnnvl.allreduce_attn_res_rmsnorm(partial, prefix, block, res_w, rms_w, out_w, EPS, EPS) +def _tail_partial(latent, act, w, lo): + """This rank's row-parallel tail partial with the kernel's arithmetic: fp32 accumulators of the latent slice and of + the activation, the latent one scaled by the RMS of the whole latent row, one bf16 rounding.""" + lat = latent.float() + scale = torch.rsqrt(lat.pow(2).mean(dim=1, keepdim=True) + LAT_EPS) + acc_lat = lat[:, lo : lo + WIDTH] @ w[:, :WIDTH].float().t() + return (acc_lat * scale + act.float() @ w[:, PAD:].float().t()).bfloat16() + + def _residual_rms_ar(ctx, partial, residual, norm_w): """The unfused drafter step: the MNNVL all-reduce with the residual add + RMSNorm fusion, sent one-shot (the sandwich reproduces the one-shot kernel's order; above 8 ranks two-shot sums the ranks in another order).""" @@ -242,7 +252,7 @@ def check_tail(ctx): latent, act = latent8[:m].contiguous(), act8[:m].contiguous() pre, block = _first(m, prefix8, block8, with_prefix) n, u = _tail(ctx, latent, act, w, lo, pre, block, res_w, rms_w, out_w) - part = torch.ops.trtllm.pdl_gemv_tail(latent, act, w, lo, WIDTH, LAT_EPS) + part = _tail_partial(latent, act, w, lo) want_n, want_u = _attn_res_ar(ctx, part, pre, block, res_w, rms_w, out_w) again = [_tail(ctx, latent, act, w, lo, pre, block, res_w, rms_w, out_w) for _ in range(2)] bad_n, bad_u = _tail(ctx, latent, _perturbed(ctx, act), w, lo, pre, block, res_w, rms_w, out_w) From 6fbd618f0655ce8f2855d699576ca72cb9b72e2a Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 18:18:59 -0700 Subject: [PATCH 058/161] [None][refactor] Kimi K3 collectives: caller-owned state objects Every stateful op of these kernels now runs on state that the caller creates and passes in (STAIRCASE_ASKS A1). No module-level dict holds buffers, counters or Lamport words any more: - K3SandwichWorkspace (k3_sandwich/op.py) replaces workspace(mapping): the sandwich all-reduce buffer of a TP group, shared by k3_sandwich_oproj, _tail and _plain. K3SandwichLatentExchange replaces latent_exchange() for the tail's folded latent all-reduce. - K3MoeHeadWorkspace (k3_fused_moe/op.py) replaces head_workspace(): the head all-gather buffers and ready words of trtllm::k3_moe_front. - K3LatentExchange.create (latent_op.py) replaces the LatentExchange constructor. - K3MoeState / K3MoeLayer replace the per-process states behind trtllm::k3_fused_moe and k3_fused_moe_front (scratch per device, intermediate size, local experts and options; counters per layer, keyed by the weight buffer). As with K3MoeWideState, the caller builds one state per device and one layer per MoE layer. K3MoeLayer() runs k3_route_quant + k3_moe and K3MoeLayer.front() runs k3_moe_front + k3_moe, with the same launches as the two ops they replace. The collective objects are built by create(mapping, fabric_handle=None): every rank of the TP group calls it at the same point, outside CUDA-graph capture, and it either returns on every rank or raises on every rank. Whether to share the memory by fabric handle is an argument, no longer read from TRTLLM_FORCE_MNNVL_AR. The per-rank objects (K3MoeState, K3MoeLayer) refuse capture too. Each kernel still compiles on its first call (a result-neutral compile cache), which must come before capture. mutates_args now names every buffer the ops write: the sandwich workspace (uc, mc, flags), plus the folded exchange for the tail; the front's head workspace (uc, mc, flags, ready). Removed together with their module dicts, since nothing calls them: trtllm::k3_fused_moe_ar and ar_workspace (the all-reduce fused into k3_moe), and trtllm::k3_fused_moe_head and trtllm::k3_route_quant_ag (the head all-gather as its own kernel, which k3_moe_front replaces). The kernel files are unchanged. The op tests use the new objects; what each test checks is unchanged. Signed-off-by: Vasanth Sabavat --- .../cute_dsl_kernels/k3_fused_moe/front_op.py | 22 +- .../k3_fused_moe/latent_op.py | 89 +- .../cute_dsl_kernels/k3_fused_moe/op.py | 915 +++++------------- .../_torch/cute_dsl_kernels/k3_sandwich/op.py | 213 ++-- .../kimi_k3/test_k3_fused_moe.py | 45 +- .../kimi_k3/test_k3_latent_reduce.py | 8 +- .../kimi_k3/test_k3_moe_front.py | 80 +- .../kimi_k3/test_k3_moe_wide.py | 15 +- .../kimi_k3/test_k3_sandwich.py | 29 +- 9 files changed, 555 insertions(+), 861 deletions(-) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/front_op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/front_op.py index 994c9715b7b4..dcb359241cf3 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/front_op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/front_op.py @@ -14,10 +14,11 @@ # limitations under the License. """``trtllm::k3_moe_front``: the Kimi K3 MoE front at decode size in one kernel (``k3_moe_front.py``). -Replaces, per MoE layer: the sharded head GEMV, ``trtllm::k3_route_quant_ag`` (head all-gather, top-16 routing, MXFP8 -latent) and, on the shared-expert stream, the shared gate_up GEMV and SiTU-and-mul. The head all-gather uses -``op.head_workspace``'s buffers and protocol, so the front and ``k3_route_quant_ag`` must not both serve one layer. -The kernel compiles on the first call for each configuration, which must happen outside CUDA-graph capture. +Replaces, per MoE layer: the sharded head GEMV, the head all-gather, the top-16 routing and the MXFP8 latent, and, on +the shared-expert stream, the shared gate_up GEMV and SiTU-and-mul. The head all-gather runs over the TP group's +:class:`~.op.K3MoeHeadWorkspace` (the caller's, created collectively before CUDA-graph capture), with the buffer +layout and protocol of ``k3_route_quant_ag.py``. The kernel compiles on the first call for each configuration, which +must happen outside CUDA-graph capture. """ from __future__ import annotations @@ -110,7 +111,9 @@ def supports( ) -@torch.library.custom_op("trtllm::k3_moe_front", mutates_args=()) +@torch.library.custom_op( + "trtllm::k3_moe_front", mutates_args=("ag_uc", "ag_mc", "ag_flags", "ag_ready") +) def k3_moe_front( x: torch.Tensor, w_front: torch.Tensor, @@ -128,9 +131,10 @@ def k3_moe_front( ag_ready: Optional[torch.Tensor] = None, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: """``x`` bf16 ``[M <= 8, 7168]`` (the MoE input, the same on every rank), ``w_front`` from ``front_weight``, - ``bias`` the routing bias fp32 ``[896]``. Returns ``(topk_ids, topk_weights, quantized, scales, shared)``: what - ``trtllm::k3_route_quant_ag`` returns for the gathered head, and the shared experts' activation bf16 - ``[M, shared_cols]``. With ``ag_ready`` it also releases the per-token ready words as route_quant_ag does.""" + ``bias`` the routing bias fp32 ``[896]``, ``ag_*`` the fields of the TP group's ``K3MoeHeadWorkspace``. Returns + ``(topk_ids, topk_weights, quantized, scales, shared)``: what ``trtllm::k3_route_quant`` returns for the gathered + head's router logits and latent, and the shared experts' activation bf16 ``[M, shared_cols]``. With ``ag_ready`` + it also releases the per-token ready words a ``head_flags`` build of k3_moe acquires.""" import cuda.bindings.driver as cuda_driver import cutlass.cute as cute from cutlass.cute.runtime import from_dlpack @@ -147,7 +151,7 @@ def k3_moe_front( if ag_uc.numel() < ag_words or ag_mc.numel() < ag_words: raise ValueError( f"k3_moe_front: the head workspace holds {ag_uc.numel()} words, the front needs {ag_words} " - "(op.head_workspace: the all-gather's buffers, then the router partials)" + "(K3MoeHeadWorkspace: the all-gather's buffers, then the router partials)" ) num_tokens, k_in = x.shape device = x.device diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py index 220a7b8a0d8a..9e5777306915 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py @@ -15,9 +15,9 @@ """``trtllm::k3_latent_reduce``: the Kimi K3 latent all-reduce at decode size as the consumer of the push-only k3_moe (``k3_latent_reduce.py``). -``LatentExchange`` owns a TP group's buffers: the push-only ops (``trtllm::k3_fused_moe_push`` / -``trtllm::k3_fused_moe_front_push``) store every rank's routed partial into them, and ``trtllm::k3_latent_reduce`` -returns the sum, bit-identical to ``MNNVLAllReduce``'s one-shot of the partials. Each push must be followed by exactly +A :class:`K3LatentExchange` holds a TP group's buffers; the caller creates it (collectively, before CUDA-graph capture). +The push-only k3_moe builds store every rank's routed partial into them, and ``trtllm::k3_latent_reduce`` returns the +sum, bit-identical to ``MNNVLAllReduce``'s one-shot of the partials. Each push must be followed by exactly one reduce of the same token count on the same exchange before the next push, on every rank in the same order. A push reads the half from the call count after its grid-dependency wait and triggers its dependents only after that wait, and the reduce reads the count before its own wait, so every kernel from a reduce to the next push must end only after @@ -29,7 +29,8 @@ import os import threading -from typing import Dict, Optional +from dataclasses import dataclass +from typing import Any, Dict, Optional import torch @@ -52,48 +53,72 @@ def default_ctas(world: int) -> int: return 4 if world <= 8 else 14 -class LatentExchange: - """The latent all-reduce buffers of ``mapping``'s TP group: int32 ``[2][8][world][1792]`` per rank behind one - multicast mapping, every word ``0x80000000``, and ``flags`` (int32 ``[4]``: the consumer's call count and its - CTA arrivals). Collective: every rank of the group constructs it at the same point (outside graph capture). - Separate from the MNNVL all-reduce workspace.""" - - def __init__(self, mapping): +@dataclass(eq=False) +class K3LatentExchange: + """The latent all-reduce buffers of a TP group: int32 ``[2][8][world][1792]`` per rank behind one multicast + mapping, every word ``0x80000000`` (empty), and ``flags`` (int32 ``[4]``: the consumer's call count, whose parity + selects the half, and its CTA arrivals). Separate from the MNNVL all-reduce workspace. Pass ``uc`` and ``flags`` + as ``trtllm::k3_latent_reduce``'s ``lat_uc`` and ``lat_flags``; :meth:`push_args` gives the producers' arguments.""" + + uc: torch.Tensor + """int32 [2 * 8 * world * 1792]: this rank's words.""" + mc: torch.Tensor + """The same words through the multicast mapping (where the producers push).""" + flags: torch.Tensor + """int32 [4]: [0] the consumer's call count, then its CTA arrivals.""" + rank: int + world_size: int + handle: Any + """The ``McastGPUBuffer`` that owns the memory; the exchange is valid while this object lives.""" + comm: Any + """The TP-group communicator the handles were exchanged over.""" + + @classmethod + def create(cls, mapping, fabric_handle: Optional[bool] = None) -> "K3LatentExchange": + """Allocate and arm an exchange for ``mapping``'s TP group. Collective: every rank of the group calls it at the + same point, eagerly (not under CUDA-graph capture); it returns on every rank or raises on every rank. + ``fabric_handle``: share the memory by fabric handle (required across nodes) rather than POSIX file + descriptor; default ``mapping.is_multi_node()``.""" + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "K3LatentExchange.create is collective and allocates: call it outside CUDA-graph capture" + ) from tensorrt_llm._torch.distributed.ops import ( _get_mnnvl_workspace_comm, _make_mnnvl_mcast_buffer, _mnnvl_workspace_all_succeeded, ) - if torch.cuda.is_current_stream_capturing(): - raise RuntimeError("Kimi K3 latent exchange buffers must be built outside CUDA-graph capture") - self.world = mapping.tp_size - self.rank = mapping.tp_rank - words = _kernel().buffer_words(self.world) + words = _kernel().buffer_words(mapping.tp_size) + use_fabric_handle = mapping.is_multi_node() if fabric_handle is None else bool(fabric_handle) comm = _get_mnnvl_workspace_comm(mapping) - use_fabric_handle = ( - os.environ.get("TRTLLM_FORCE_MNNVL_AR", "0") == "1" or mapping.is_multi_node() - ) error: Optional[Exception] = None + exchange = None try: - self.handle = _make_mnnvl_mcast_buffer(comm, words * 4, mapping, use_fabric_handle) - self.uc = self.handle.get_uc_buffer(self.rank, (words,), torch.int32, 0) - self.mc = self.handle.get_mc_buffer((words,), torch.int32, 0) - with torch.inference_mode(): - self.uc.fill_(EMPTY_WORD) - self.flags = torch.zeros(4, dtype=torch.int32, device=self.uc.device) + handle = _make_mnnvl_mcast_buffer(comm, words * 4, mapping, use_fabric_handle) + uc = handle.get_uc_buffer(mapping.tp_rank, (words,), torch.int32, 0) + mc = handle.get_mc_buffer((words,), torch.int32, 0) + uc.fill_(EMPTY_WORD) + flags = torch.zeros(4, dtype=torch.int32, device=uc.device) torch.cuda.synchronize() + exchange = cls( + uc=uc, + mc=mc, + flags=flags, + rank=mapping.tp_rank, + world_size=mapping.tp_size, + handle=handle, + comm=comm, + ) except Exception as exc: # noqa: BLE001 -- reported to every rank below, then re-raised error = exc - # Also the barrier that keeps any rank from pushing into a peer's buffer before the peer has filled it. + # Also the barrier that keeps any rank from pushing into a peer's buffer before the peer has emptied it. if not _mnnvl_workspace_all_succeeded(comm, error is None): - raise RuntimeError( - "Kimi K3 latent exchange buffers failed on at least one rank" - ) from error - self.comm = comm + raise RuntimeError("K3LatentExchange: allocation failed on at least one rank") from error + return exchange def push_args(self): - """``(ar_uc, ar_mc, ar_flags, ar_rank)`` of the push-only ops.""" + """``(ar_uc, ar_mc, ar_flags, ar_rank)`` of the push-only producers.""" return self.uc, self.mc, self.flags, self.rank @@ -106,7 +131,7 @@ def k3_latent_reduce( lat_uc: torch.Tensor, lat_flags: torch.Tensor, num_tokens: int, ctas_per_token: int = 0 ) -> torch.Tensor: """The latent rows ``[num_tokens, 3584]`` bf16: the sum over the ranks of the routed partials the push-only k3_moe - stored into ``lat_uc`` (``LatentExchange.uc``) since the last call, in the MNNVL one-shot's order. Empties the + stored into ``lat_uc`` (``K3LatentExchange.uc``) since the last call, in the MNNVL one-shot's order. Empties the words it read and advances ``lat_flags``' call count. ``ctas_per_token``: 4, 14 or 28 (0: ``default_ctas``).""" import cuda.bindings.driver as cuda_driver import cutlass.cute as cute diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py index 2f97a353a35d..94c9ea5e759b 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py @@ -12,38 +12,30 @@ # 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. -"""``trtllm::k3_fused_moe``: Kimi K3 routed experts for decode (M <= 8 tokens). - -Two kernels on the current stream, no host synchronization: - -1. ``trtllm::k3_route_quant`` -- the routing and MXFP8 input quantization the TRTLLM-Gen - path uses under separated routing (sigmoid, top-16 of sigmoid + bias, unbiased scores - renormalized times the routed scaling factor; ties to the lower id), the CuTe DSL form of - ``trtllm::kimi_k3_noaux_tc_mxfp8_quant`` with the same outputs bit for bit. -2. ``k3_moe`` -- one persistent CuTe DSL kernel, launched as a programmatic dependent of - the first: this rank's (expert, token) groups in its prologue, then FC1 + SiTU + FC2 - with the routing-weighted, deterministic combine. - -With the ``fold`` option there is one kernel: ``k3_moe`` computes the routing and the quantization in -its prologue (``trtllm::k3_route_quant``'s device code, so the same ids, weights and MXFP8 bits) -from the router logits and the latent, and its PDL predecessor is whatever produced them. - -The result is this rank's routed partial ``[M, 3584]`` bf16, the tensor the TRTLLM-Gen -W4A8_MXFP4_MXFP8 op returns, so the routed-latent all-reduce and the latent-up tail are -unchanged. Weights are the TRTLLM-Gen buffers, read in place. - -``trtllm::k3_fused_moe_ar`` also performs the all-reduce of that partial over the TP group -inside ``k3_moe`` (buffers from ``ar_workspace``) and returns the reduced latent. - -The CuTe DSL kernel is compiled on the first call for each intermediate size (for the -model: the warmup that precedes CUDA-graph capture) and the scratch buffers are allocated -then too, so captured calls only launch. Each layer (keyed by its weight buffer) owns a -few counters that the kernel returns to zero; the intermediate slab, shared by all layers, -is left armed by every call. - -Steps of up to 64 tokens use :class:`K3MoeWideState` (the m_max 64 build of ``k3_moe``, -launched after ``trtllm::k3_route_quant``), whose scratch and per-layer counters the caller -owns. +"""Kimi K3 routed experts for decode: the persistent CuTe DSL kernel ``k3_moe`` (``k3_moe_kernel.py``). + +For M <= 8 tokens (:class:`K3MoeState`, :class:`K3MoeLayer`), two kernels on the current stream, no host +synchronization: + +1. the routing and the MXFP8 input quantization: ``trtllm::k3_route_quant`` from the router logits and the latent + (the routing the TRTLLM-Gen path uses under separated routing: sigmoid, top-16 of sigmoid + bias, unbiased scores + renormalized times the routed scaling factor, ties to the lower id; the CuTe DSL form of + ``trtllm::kimi_k3_noaux_tc_mxfp8_quant`` with the same outputs bit for bit), or ``trtllm::k3_moe_front`` from the + MoE input (:meth:`K3MoeLayer.front`: head GEMV, head all-gather, routing, MXFP8 latent, shared gate_up + SiTU); +2. ``k3_moe``, launched as a programmatic dependent of the first: this rank's (expert, token) groups in its prologue, + then FC1 + SiTU + FC2 with the routing-weighted, deterministic combine. + +The result is this rank's routed partial ``[M, 3584]`` bf16, the tensor the TRTLLM-Gen W4A8_MXFP4_MXFP8 op returns, +so the routed-latent all-reduce and the latent-up tail are unchanged. Weights are the TRTLLM-Gen buffers, read in +place. + +Steps of up to 64 tokens use :class:`K3MoeWideState` (the m_max 64 build of ``k3_moe``, launched after +``trtllm::k3_route_quant``). + +The caller owns all state: the scratch a state's layers share (the intermediate slab, left armed by every call, and +the FC2 partial rows), each layer's counters (left zero by every call), and the head all-gather buffers of the front +(:class:`K3MoeHeadWorkspace`, collective over the TP group). Build them before CUDA-graph capture; each kernel compiles +on its first call, which must also come before capture. """ from __future__ import annotations @@ -51,8 +43,8 @@ import importlib.util import os import sys -import threading -from typing import Dict, Optional, Tuple +from dataclasses import dataclass +from typing import Any, Dict, Optional, Tuple import torch @@ -62,11 +54,10 @@ MAX_TOKENS = 8 _TOKEN_SLOTS = 8 _SF_VEC = 32 +EMPTY_WORD = -(2**31) _KERNEL_PATH = os.path.join(os.path.dirname(os.path.abspath(__file__)), "k3_moe_kernel.py") -_lock = threading.Lock() _modules: Dict[tuple, object] = {} -_states: Dict[Tuple[int, int, int], "_K3FusedMoE"] = {} def _kernel_module(config: dict): @@ -126,21 +117,120 @@ def is_supported(w3_w1_weight: torch.Tensor, w3_w1_weight_scale: torch.Tensor, w return True, "" -class _K3FusedMoE: - """Compiled kernel and scratch for one (device, intermediate size, local experts). +@dataclass(eq=False) +class K3MoeHeadWorkspace: + """One TP group's MoE head all-gather buffers, read and written by ``trtllm::k3_moe_front`` (alone, or as the + producer of :meth:`K3MoeLayer.front`): two alternating Lamport buffers of every rank's head slice per token behind + one multicast mapping, then the front's router partials; the flag words that rotate them; and the per-token ready + words a ``head_flags`` build of k3_moe acquires. Every front call on it takes the next buffer, so all of a group's + ranks make the same front calls on it in the same order. Pass ``uc``, ``mc``, ``flags``, ``rank`` and + ``world_size`` as the front's ``ag_uc``, ``ag_mc``, ``ag_flags``, ``ag_rank`` and ``ag_world``, and ``ready`` as + its ``ag_ready``. Separate from the MNNVL all-reduce workspace.""" + + uc: torch.Tensor + """int32 [workspace_words(world)]: this rank's words (0x80000000 = empty).""" + mc: torch.Tensor + """The same words through the multicast mapping (where the peers push).""" + flags: torch.Tensor + """int32 [4]: [0] the buffer of the next call; [2] the ready words' epoch (advanced by k3_moe's head_flags build); + [3] the front's sign-ins after its reads of [0], counted by the CTA that flips [0]; [1] unused.""" + ready: torch.Tensor + """int32 [32]: [t] token t's ids and weights, [8 + t] its MXFP8 row, released as ``epoch + 1``.""" + rank: int + world_size: int + handle: Any + """The ``McastGPUBuffer`` that owns the memory; the workspace is valid while this object lives.""" + comm: Any + """The TP-group communicator the handles were exchanged over.""" + + @classmethod + def create(cls, mapping, fabric_handle: Optional[bool] = None) -> "K3MoeHeadWorkspace": + """Allocate and arm a workspace for ``mapping``'s TP group. Collective: every rank of the group calls it at the + same point, eagerly (not under CUDA-graph capture); it returns on every rank or raises on every rank. + ``fabric_handle``: share the memory by fabric handle (required across nodes) rather than POSIX file + descriptor; default ``mapping.is_multi_node()``.""" + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "K3MoeHeadWorkspace.create is collective and allocates: call it outside CUDA-graph capture" + ) + from tensorrt_llm._torch.distributed.ops import ( + _get_mnnvl_workspace_comm, + _make_mnnvl_mcast_buffer, + _mnnvl_workspace_all_succeeded, + ) + + from . import k3_route_quant_ag as layout + + words = layout.workspace_words(mapping.tp_size) + use_fabric_handle = mapping.is_multi_node() if fabric_handle is None else bool(fabric_handle) + comm = _get_mnnvl_workspace_comm(mapping) + error: Optional[Exception] = None + workspace = None + try: + handle = _make_mnnvl_mcast_buffer(comm, words * 4, mapping, use_fabric_handle) + uc = handle.get_uc_buffer(mapping.tp_rank, (words,), torch.int32, 0) + mc = handle.get_mc_buffer((words,), torch.int32, 0) + uc.fill_(EMPTY_WORD) + flags = torch.zeros(4, dtype=torch.int32, device=uc.device) + ready = torch.zeros(32, dtype=torch.int32, device=uc.device) + torch.cuda.synchronize() + workspace = cls( + uc=uc, + mc=mc, + flags=flags, + ready=ready, + rank=mapping.tp_rank, + world_size=mapping.tp_size, + handle=handle, + comm=comm, + ) + except Exception as exc: # noqa: BLE001 -- reported to every rank below, then re-raised + error = exc + # Also the barrier that keeps any rank from pushing into a peer's buffer before the peer has emptied it. + if not _mnnvl_workspace_all_succeeded(comm, error is None): + raise RuntimeError("K3MoeHeadWorkspace: allocation failed on at least one rank") from error + return workspace - ``config`` overrides kernel options (tests and A/B runs: ``pdl``, - ``num_ctas``, ...); anything it leaves out takes the kernel's default.""" + +class K3MoeState: + """``k3_moe`` for 1..8 decode tokens on one device: its build and the scratch its layers share, i.e. the FC1 -> + FC2 intermediate slab (armed between calls: FP8 -0.0 values, E8M0 NaN scale words) and the FC2 partial rows. Build + it eagerly before CUDA-graph capture and keep it with the model; every layer takes its own counters from + :meth:`layer`. The layers of one state run in one stream order (they share the scratch). The kernel compiles on the + first call, which must therefore come before capture. + + ``head_flags``: the build in which k3_moe acquires the front's routing and MXFP8 rows through the head workspace's + ready words instead of waiting for the front's grid (:meth:`K3MoeLayer.front` only). ``config`` overrides kernel + options (tests and A/B runs: ``pdl``, ``num_ctas``, ...); anything it leaves out takes the kernel's default.""" def __init__( - self, device: torch.device, i_tp: int, num_local: int, config: Optional[dict] = None + self, + device: torch.device, + i_tp: int, + num_local: int, + head_flags: bool = False, + config: Optional[dict] = None, ): - # One persistent CTA per SM (config "num_ctas" caps it, e.g. for a grid-size A/B). + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError("K3MoeState allocates its scratch: build it outside CUDA-graph capture") + # One persistent CTA per SM (config "num_ctas" caps it, e.g. for a grid-size A/B). head_flags always explicit: + # its environment fallback must not reach the plain build. num_ctas = torch.cuda.get_device_properties(device).multi_processor_count - cfg = {"i_tp": i_tp, "num_ctas": num_ctas, "num_local": num_local} + cfg = { + "i_tp": i_tp, + "num_ctas": num_ctas, + "num_local": num_local, + "head_flags": int(head_flags), + "lat_slab": 0, + } cfg.update(config or {}) self.mod = mod = _kernel_module(cfg) + if mod.FUSED_AR or mod.FOLD or mod.LAT_SLAB or mod.WIDE: + raise ValueError("K3MoeState is the M <= 8 build without the fused all-reduce, the fold or the slab") self.device = device + self.i_tp = i_tp + self.num_local = num_local + self.head_flags = mod.HEAD_FLAGS g_cap = mod.G_CAP kw = dict(device=device) # Lamport slab, armed: FP8 -0.0 values; E8M0 NaN in bytes 0..3 of each 16-byte scale @@ -163,631 +253,156 @@ def __init__( ) self.sfb2_t = _view(sfb2, 16, 0, mod.sf_dtype) self.part_t = _view(self.part, 16, 1) - # Stand-ins for the all-reduce buffers of builds without the fused all-reduce, and for the - # inputs a build does not read (fold: the top-k; otherwise the logits and the latent). - self.no_ar = torch.zeros(4, dtype=torch.int32, **kw) - self.no_ar_t = _view(self.no_ar, 16, 0) - self.no_w = torch.zeros(8, dtype=torch.bfloat16, **kw) - self.no_w_t = _view(self.no_w, 16, 0) - self.ar_views: Dict[Tuple[int, int, int], tuple] = {} - self.fold = mod.FOLD - if self.fold: - # Each CTA's MXFP8 latent rows [8 * cta, 8 * cta + M) and their linear scales. - rows = mod.NUM_CTAS * mod.M_MAX - self.xq = torch.empty(rows, HIDDEN_SIZE, dtype=torch.uint8, **kw) - self.xsf = torch.empty(rows, HIDDEN_SIZE // _SF_VEC, dtype=torch.uint8, **kw) - self.b1_t = _view(self.xq.permute(1, 0), 16, 0, mod.b_dtype) - self.sfb1_t = _view(self.xsf, 16, 1, mod.sf_dtype) - self.xq_words_t = _view(self.xq.view(-1).view(torch.int32), 16, 0) - self.xsf_t = _view(self.xsf.view(-1), 16, 0) + # Stand-in for the buffers of the options this build does not have (fused all-reduce, fold inputs, latent + # slab, and the ready words without head_flags). + self.unused = torch.zeros(4, dtype=torch.int32, **kw) + self.unused_t = _view(self.unused, 16, 0) # The route+quant kernel triggers k3_moe's launch right after its own grid dependency: # k3_moe waits for the whole route+quant grid before reading its outputs. self.route_kwargs = {"early_trigger": True} if mod.USE_PDL else {} - from ..k3_route_quant import op as _k3_route_quant_op # noqa: F401 - - self.route_quant = torch.ops.trtllm.k3_route_quant - # Per layer (keyed by its w3_w1 buffer): weight views and the kernel's counters. - self.layers: Dict[Tuple[int, int, int, int], tuple] = {} - self.moe = None + from ..k3_route_quant import op as _k3_route_quant_op # noqa: F401 (registers trtllm::k3_route_quant) - def _layer(self, w31, w31s, w2, w2s): - key = (w31.data_ptr(), w31s.data_ptr(), w2.data_ptr(), w2s.data_ptr()) - layer = self.layers.get(key) - if layer is None: - if torch.cuda.is_current_stream_capturing(): - raise RuntimeError( - "trtllm::k3_fused_moe must run once per layer outside CUDA-graph capture " - "first (it allocates the layer's counters on the first call)." - ) - mod = self.mod - e, two_i, _ = w31.shape - i_tp = two_i // 2 - sfa1 = w31s.view(e, two_i // 128, HIDDEN_SIZE // 128, 512).permute(3, 2, 1, 0) - sfa2 = w2s.view(e, HIDDEN_SIZE // 128, i_tp // 128, 512).permute(3, 2, 1, 0) - state = torch.zeros(mod.NUM_STATE, dtype=torch.int32, device=self.device) - layer = ( - _view(w31.view(torch.int8).permute(2, 1, 0), 16, 0), - _view(sfa1, 16, 0, mod.sf_dtype), - _view(w2.view(torch.int8).permute(2, 1, 0), 16, 0), - _view(sfa2, 16, 0, mod.sf_dtype), - state, - _view(state, 4, 0), - ) - self.layers[key] = layer - return layer - - def _flat_view(self, t: torch.Tensor): - key = ("flat", t.data_ptr(), t.numel()) - v = self.ar_views.get(key) - if v is None: - v = self.ar_views[key] = _view(t.view(-1), 16, 0) - return v - - def _ar_views(self, ar): - if ar is None: - return self.no_ar_t, self.no_ar_t, self.no_ar_t, 0 - uc, mc, flags, rank = ar - key = (uc.data_ptr(), mc.data_ptr(), flags.data_ptr()) - views = self.ar_views.get(key) - if views is None: - views = self.ar_views[key] = (_view(uc, 16, 0), _view(mc, 16, 0), _view(flags, 4, 0)) - return (*views, rank) + self.compiled = None - def __call__( + def layer( self, - x, - router_logits, - bias, - w31, - w31s, - w2, - w2s, - local_offset: int, - num_local: int, - scale: float, - ar: Optional[tuple] = None, - head_ag: Optional[tuple] = None, - head_ready: Optional[torch.Tensor] = None, - front: Optional[tuple] = None, - lat_slab: Optional[tuple] = None, - ): - """``ar``: (uc words, mc words, flags, rank) of the fused all-reduce, for a build - with ``ar_world`` set; the result is then the all-reduced sum. With ``ar_push_only`` the - rows only go to every rank's buffer (half ``flags[0] & 1``) and the result has no rows. - - ``head_ag``: (uc words, mc words, flags, rank) of the head all-gather's buffers; ``x`` - is then this rank's slice of the sharded MoE head (fp32 ``[M, (H + E) / world]``), - ``router_logits`` is None, and ``trtllm::k3_route_quant_ag`` gathers, routes and - quantizes before ``k3_moe``. ``head_ready``: route_quant_ag's ready words, for a build with - ``head_flags`` (k3_moe then acquires them instead of waiting for that grid). - - ``front``: (front weight, shared columns, gate cap, linear cap, world) of - ``trtllm::k3_moe_front``, with ``head_ag``; ``x`` is then the MoE input (bf16 ``[M, 7168]``, - the same on every rank), the front stands in for the head GEMV and route_quant_ag, and - the call returns ``(y, shared activation)``. - - ``lat_slab``: (slab, buffer, re-arm buffer 0) for a build with ``lat_slab`` (and the fused - all-reduce): the reduced latent rows also go into ``slab`` (int32 ``[3, 8, 1792]``, all-ones - empty) in ``buffer`` (the call's ordinal mod 3); the call re-arms the next buffer, and buffer - 0 too when the flag is set (the step's last call, whose own buffer must not be 0). - """ - import cuda.bindings.driver as cuda_driver - import cutlass.cute as cute - - mod = self.mod - if (ar is not None) != mod.FUSED_AR: - raise ValueError("the all-reduce buffers go with an ar_world build, and only with one") - if head_ag is not None and self.fold: - raise ValueError("the head all-gather feeds the unfolded route + quant") - if (head_ready is not None) != mod.HEAD_FLAGS or (mod.HEAD_FLAGS and head_ag is None): - raise ValueError( - "the ready words go with a head_flags build and the head all-gather, and only there" - ) - if front is not None and head_ag is None: - raise ValueError("the front pushes its head into the head all-gather's buffers") - if (lat_slab is not None) != mod.LAT_SLAB: - raise ValueError("the latent slab goes with a lat_slab build, and only with one") - lat_buf, lat_rearm0 = 0, 0 - if lat_slab is not None: - slab, lat_buf, rearm0 = lat_slab - if ( - slab.dtype != torch.int32 - or not slab.is_contiguous() - or slab.numel() != mod.LAT_SLAB_BUFS * MAX_TOKENS * HIDDEN_SIZE // 2 - or not 0 <= lat_buf < mod.LAT_SLAB_BUFS - or (rearm0 and lat_buf == 0) - ): - raise ValueError( - f"k3_moe latent slab: int32 [3, 8, 1792] contiguous, buffer in [0, 3), and not buffer 0 " - f"with the buffer-0 re-arm; got {tuple(slab.shape)} {slab.dtype}, buffer {lat_buf}, " - f"re-arm 0 {rearm0}" - ) - lat_rearm0 = int(bool(rearm0)) - num_tokens = x.shape[0] - shared = None - a1, sfa1, a2, sfa2, _, state_t = self._layer(w31, w31s, w2, w2s) - ar_uc_t, ar_mc_t, ar_flags_t, ar_rank = self._ar_views(ar) - y = torch.empty(num_tokens, HIDDEN_SIZE, dtype=torch.bfloat16, device=x.device) - stream = cuda_driver.CUstream(torch.cuda.current_stream().cuda_stream) - if self.fold: - ids_t, weights_t, b1, sfb1 = self.no_ar_t, self.no_w_t, self.b1_t, self.sfb1_t - fold_in = [ - _view(router_logits.contiguous().view(-1), 16, 0), - _view(bias.detach().contiguous().view(-1), 16, 0), - _view(x.contiguous().view(-1).view(torch.int32), 16, 0), - self.xq_words_t, - self.xsf_t, - ] - else: - if front is not None: - # Registers trtllm::k3_moe_front. - from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import front_op # noqa: F401 - - w_front, shared_cols, gate_cap, linear_cap, world = front - ids, weights, x_fp8, x_sf, shared = torch.ops.trtllm.k3_moe_front( - x, w_front, bias, scale, shared_cols, gate_cap, linear_cap, *head_ag, world, - ag_ready=head_ready, - ) # fmt: skip - elif head_ag is not None: - ids, weights, x_fp8, x_sf = torch.ops.trtllm.k3_route_quant_ag( - x, bias, scale, *head_ag, early_trigger=mod.USE_PDL, ag_ready=head_ready - ) - else: - ids, weights, x_fp8, x_sf = self.route_quant( - router_logits, bias, x, scale, **self.route_kwargs - ) - ids_t, weights_t = _view(ids, 4, 1), _view(weights, 4, 1) - b1 = _view(x_fp8.view(torch.uint8).permute(1, 0), 16, 0, mod.b_dtype) - sfb1 = _view( - x_sf.view(torch.uint8).view(num_tokens, HIDDEN_SIZE // _SF_VEC), 16, 1, mod.sf_dtype - ) - fold_in = [self.no_ar_t] * 5 - if mod.HEAD_FLAGS: - flag_in = [self._flat_view(head_ready), self._flat_view(head_ag[2])] - else: - flag_in = [self.no_ar_t, self.no_ar_t] - args = [ - a1, b1, sfa1, sfb1, self.c_t, self.cs_t, self.c_words_t, self.cs_words_t, a2, - self.b2_t, sfa2, self.sfb2_t, _view(y, 16, 1), _view(y.view(torch.int32), 16, 1), - self.part_t, ids_t, weights_t, state_t, ar_uc_t, ar_mc_t, ar_flags_t, - *fold_in, *flag_in, - self._flat_view(lat_slab[0]) if lat_slab is not None else self.no_ar_t, - ] # fmt: skip - scalars = (num_tokens, local_offset, num_local, ar_rank, float(scale), lat_buf, lat_rearm0) - if self.moe is None: - self.moe = cute.compile(mod.k3_moe, *args, *scalars, stream) - self.moe(*args, *scalars, stream) - if mod.AR_PUSH_ONLY: - y = y[:0] # the rows went to every rank's buffer; the consumer reduces them - return y if shared is None else (y, shared) - - -def _state( - device: torch.device, - i_tp: int, - num_local: int, - ar_world: int = 0, - head_flags: bool = False, - lat_slab: bool = False, - ar_push_only: bool = False, -) -> _K3FusedMoE: - key = ( - device.index if device.index is not None else torch.cuda.current_device(), - i_tp, - num_local, - ar_world, - head_flags, - lat_slab, - ar_push_only, - ) - st = _states.get(key) - if st is None: - if torch.cuda.is_current_stream_capturing(): - raise RuntimeError( - "trtllm::k3_fused_moe must run once outside CUDA-graph capture first " - "(it compiles its kernels and allocates its scratch on the first call)." - ) - with _lock: - st = _states.get(key) - if st is None: - # head_flags always explicit: its environment fallback must not reach the plain op. - config = {"head_flags": int(head_flags), "lat_slab": int(lat_slab)} - if ar_world: - config["ar_world"] = ar_world - config["ar_push_only"] = int(ar_push_only) - st = _states[key] = _K3FusedMoE(device, i_tp, num_local, config) - return st - + w3_w1_weight: torch.Tensor, + w3_w1_weight_scale: torch.Tensor, + w2_weight: torch.Tensor, + w2_weight_scale: torch.Tensor, + ) -> "K3MoeLayer": + """A layer's handle: its experts' TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers (read in place) and its counters.""" + return K3MoeLayer(self, w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale) -# --------------------------------------------------------------------------- -# Fused all-reduce: this TP group's Lamport buffers, 2 x [MAX_TOKENS][group][3584] bf16 per -# rank behind one multicast mapping (the MNNVL all-reduce's allocator), and a local flag -# holding the buffer the next call uses. The kernel leaves both buffers empty (every word -# 0x80000000) after each call. Separate from the model's all-reduce workspace, which the -# shared expert's all-reduce uses concurrently on its auxiliary stream. -# --------------------------------------------------------------------------- -AR_BUFFERS = 2 -AR_EMPTY_WORD = -(2**31) -_ar_workspaces: Dict[object, dict] = {} -_head_workspaces: Dict[object, dict] = {} +class K3MoeLayer: + """One MoE layer on a :class:`K3MoeState`: its weights as the kernel reads them and its counters (int32, zero + between calls; every call leaves them zero).""" -def ar_buffer_words(world: int) -> int: - return AR_BUFFERS * MAX_TOKENS * world * HIDDEN_SIZE // 2 + def __init__( + self, + state: K3MoeState, + w3_w1_weight: torch.Tensor, + w3_w1_weight_scale: torch.Tensor, + w2_weight: torch.Tensor, + w2_weight_scale: torch.Tensor, + ): + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError("K3MoeLayer allocates its counters: build it outside CUDA-graph capture") + ok, why = is_supported( + w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale, state.num_local + ) + if not ok or w3_w1_weight.shape[1] != 2 * state.i_tp: + raise ValueError(f"k3_moe layer: {why or 'intermediate size differs from the state'}") + mod = state.mod + e, two_i, _ = w3_w1_weight.shape + i_tp = two_i // 2 + sfa1 = w3_w1_weight_scale.view(e, two_i // 128, HIDDEN_SIZE // 128, 512).permute(3, 2, 1, 0) + sfa2 = w2_weight_scale.view(e, HIDDEN_SIZE // 128, i_tp // 128, 512).permute(3, 2, 1, 0) + self.state = state + self.counters = torch.zeros(mod.NUM_STATE, dtype=torch.int32, device=state.device) + self.weights = ( + _view(w3_w1_weight.view(torch.int8).permute(2, 1, 0), 16, 0), + _view(sfa1, 16, 0, mod.sf_dtype), + _view(w2_weight.view(torch.int8).permute(2, 1, 0), 16, 0), + _view(sfa2, 16, 0, mod.sf_dtype), + ) + self.counters_t = _view(self.counters, 4, 0) + def __call__( + self, + hidden_states: torch.Tensor, + router_logits: torch.Tensor, + e_score_correction_bias: torch.Tensor, + local_expert_offset: int, + routed_scaling_factor: float, + ) -> torch.Tensor: + """This rank's routed partial ``[M, 3584]`` bf16 for M <= 8 decode tokens: ``trtllm::k3_route_quant`` (the + routing and the MXFP8 latent), then ``k3_moe`` launched as its programmatic dependent. -def ar_workspace(mapping) -> dict: - """The fused all-reduce buffers of ``mapping``'s TP group; collective on first use, so - every rank of the group must make the first call at the same point, outside CUDA-graph - capture.""" - return _lamport_workspace(mapping, ar_buffer_words(mapping.tp_size), _ar_workspaces) + ``hidden_states``: bf16 ``[M, 3584]`` latent; ``router_logits``: fp32 ``[M, 896]``; + ``e_score_correction_bias``: fp32 ``[896]``; the layer's experts hold global ids + ``[local_expert_offset, local_expert_offset + num_local)``.""" + st = self.state + if st.head_flags: + raise ValueError("a head_flags build takes the front's ready words: call K3MoeLayer.front") + _check_tokens(hidden_states) + ids, weights, x_fp8, x_sf = torch.ops.trtllm.k3_route_quant( + router_logits.contiguous(), e_score_correction_bias, hidden_states.contiguous(), + float(routed_scaling_factor), **st.route_kwargs, + ) # fmt: skip + return self._launch(ids, weights, x_fp8, x_sf, local_expert_offset, routed_scaling_factor) + def front( + self, + x: torch.Tensor, + w_front: torch.Tensor, + e_score_correction_bias: torch.Tensor, + local_expert_offset: int, + routed_scaling_factor: float, + shared_cols: int, + gate_cap: float, + linear_cap: float, + head: K3MoeHeadWorkspace, + ) -> Tuple[torch.Tensor, torch.Tensor]: + """``trtllm::k3_moe_front`` (head GEMV, head all-gather over ``head``, routing, MXFP8 latent, shared gate_up + + SiTU) then ``k3_moe`` for the MoE input ``x`` bf16 ``[M, 7168]`` (the same on every rank). Returns ``(y [M, + 3584] bf16, shared activation [M, shared_cols] bf16)``. A ``head_flags`` build acquires the front's routing and + MXFP8 rows through ``head.ready`` instead of waiting for the front's grid, whose shared tiles may still be + running.""" + from . import front_op # noqa: F401 (registers trtllm::k3_moe_front) -def _rqag_module(): - """k3_route_quant_ag.py next to this file (also when op.py is loaded outside the package).""" - mod = _modules.get("rqag") - if mod is None: - path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "k3_route_quant_ag.py") - spec = importlib.util.spec_from_file_location(f"{__name__}_k3_route_quant_ag", path) - mod = importlib.util.module_from_spec(spec) - sys.modules[spec.name] = mod - spec.loader.exec_module(mod) - _modules["rqag"] = mod - return mod + st = self.state + _check_tokens(x) + ready = head.ready if st.head_flags else None + ids, weights, x_fp8, x_sf, shared = torch.ops.trtllm.k3_moe_front( + x.contiguous(), w_front, e_score_correction_bias, float(routed_scaling_factor), shared_cols, gate_cap, + linear_cap, head.uc, head.mc, head.flags, head.rank, head.world_size, ag_ready=ready, + ) # fmt: skip + flag_in = None + if st.head_flags: + flag_in = (_view(head.ready.view(-1), 16, 0), _view(head.flags.view(-1), 16, 0)) + y = self._launch(ids, weights, x_fp8, x_sf, local_expert_offset, routed_scaling_factor, flag_in) + return y, shared + def _launch(self, ids, weights, x_fp8, x_sf, local_offset, scale, flag_in=None) -> torch.Tensor: + import cuda.bindings.driver as cuda_driver + import cutlass.cute as cute -def head_workspace(mapping) -> dict: - """The head all-gather's buffers of ``mapping``'s TP group (``trtllm::k3_route_quant_ag``), then - ``trtllm::k3_moe_front``'s router partials; collective on first use like ``ar_workspace``. - ``ready``: route_quant_ag's per-token ready words for k3_moe's flag handoff; flags[2] is their - epoch.""" - ws = _lamport_workspace( - mapping, _rqag_module().workspace_words(mapping.tp_size), _head_workspaces - ) - if "ready" not in ws: - ws["ready"] = torch.zeros(32, dtype=torch.int32, device=ws["uc"].device) - return ws - - -def _lamport_workspace(mapping, words: int, cache: Dict[object, dict]) -> dict: - ws = cache.get(mapping) - if ws is not None: - return ws - if torch.cuda.is_current_stream_capturing(): - raise RuntimeError( - "k3_fused_moe: the head all-gather and fused all-reduce buffers are collective on first " - "use and must be allocated outside CUDA-graph capture" - ) - from tensorrt_llm._torch.distributed.ops import ( - _get_mnnvl_workspace_comm, - _make_mnnvl_mcast_buffer, - _mnnvl_workspace_all_succeeded, - ) - - world = mapping.tp_size - comm = _get_mnnvl_workspace_comm(mapping) - use_fabric_handle = ( - os.environ.get("TRTLLM_FORCE_MNNVL_AR", "0") == "1" or mapping.is_multi_node() - ) - error: Optional[Exception] = None - ws = None - try: - handle = _make_mnnvl_mcast_buffer(comm, words * 4, mapping, use_fabric_handle) - uc = handle.get_uc_buffer(mapping.tp_rank, (words,), torch.int32, 0) - mc = handle.get_mc_buffer((words,), torch.int32, 0) - with torch.inference_mode(): - uc.fill_(AR_EMPTY_WORD) - flags = torch.zeros(4, dtype=torch.int32, device=uc.device) - torch.cuda.synchronize() - ws = dict( - handle=handle, comm=comm, uc=uc, mc=mc, flags=flags, rank=mapping.tp_rank, world=world + st = self.state + mod = st.mod + num_tokens = ids.shape[0] + a1, sfa1, a2, sfa2 = self.weights + y = torch.empty(num_tokens, HIDDEN_SIZE, dtype=torch.bfloat16, device=x_fp8.device) + stream = cuda_driver.CUstream(torch.cuda.current_stream().cuda_stream) + b1 = _view(x_fp8.view(torch.uint8).permute(1, 0), 16, 0, mod.b_dtype) + sfb1 = _view( + x_sf.view(torch.uint8).view(num_tokens, HIDDEN_SIZE // _SF_VEC), 16, 1, mod.sf_dtype ) - except Exception as exc: # noqa: BLE001 -- reported to every rank below, then re-raised - error = exc - # Also the barrier that keeps any rank from pushing into a peer's buffer before the - # peer has emptied it. - if not _mnnvl_workspace_all_succeeded(comm, error is None): - raise RuntimeError("k3_fused_moe Lamport buffers failed on at least one rank") from error - cache[mapping] = ws - return ws - + u = st.unused_t + args = [ + a1, b1, sfa1, sfb1, st.c_t, st.cs_t, st.c_words_t, st.cs_words_t, a2, st.b2_t, sfa2, st.sfb2_t, + _view(y, 16, 1), _view(y.view(torch.int32), 16, 1), st.part_t, _view(ids, 4, 1), _view(weights, 4, 1), + self.counters_t, + u, u, u, # the fused all-reduce's buffers + u, u, u, u, u, # the fold's inputs + *(flag_in or (u, u)), # the ready words and the head flags (head_flags) + u, # the latent slab + ] # fmt: skip + scalars = (num_tokens, local_offset, st.num_local, 0, float(scale), 0, 0) + if st.compiled is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "k3_moe compiles on its first call: call it once before CUDA-graph capture" + ) + st.compiled = cute.compile(mod.k3_moe, *args, *scalars, stream) + st.compiled(*args, *scalars, stream) + return y -# --------------------------------------------------------------------------- -# The head all-gather + route + quant (k3_route_quant_ag.py): one kernel instead of the MNNVL -# all-gather of the sharded head followed by trtllm::k3_route_quant. -# --------------------------------------------------------------------------- -_rqag_compiled: Dict[Tuple[int, bool, bool, bool], object] = {} - - -@torch.library.custom_op("trtllm::k3_route_quant_ag", mutates_args=()) -def k3_route_quant_ag( - head: torch.Tensor, - bias: torch.Tensor, - routed_scaling_factor: float, - ag_uc: torch.Tensor, - ag_mc: torch.Tensor, - ag_flags: torch.Tensor, - ag_rank: int, - early_trigger: bool = False, - ag_ready: Optional[torch.Tensor] = None, -) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - """Gathers every rank's slice of the sharded MoE head (``head``: this rank's fp32 - ``[M, (3584 + 896) / world]``, latent columns then router logits) through ``head_workspace``'s - buffers and returns ``trtllm::k3_route_quant``'s outputs for the gathered logits and latent: - ``(topk_ids, topk_weights, quantized, scales)``. With ``ag_ready`` it also releases the - per-token ready words for a k3_moe built with head_flags.""" - import cuda.bindings.driver as cuda_driver - import cutlass.cute as cute - from cutlass.cute.runtime import from_dlpack - rqag = _rqag_module() - num_tokens, width = head.shape - world = (HIDDEN_SIZE + NUM_EXPERTS) // width - if ( - head.dtype != torch.float32 - or width * world != HIDDEN_SIZE + NUM_EXPERTS - or not 0 < num_tokens <= MAX_TOKENS - ): - raise ValueError( - f"k3_route_quant_ag: head must be fp32 [M <= {MAX_TOKENS}, (3584 + 896) / world]" - ) - device = head.device - topk_ids = torch.empty(num_tokens, TOP_K, dtype=torch.int32, device=device) - topk_weights = torch.empty(num_tokens, TOP_K, dtype=torch.bfloat16, device=device) - quantized = torch.empty(num_tokens, HIDDEN_SIZE, dtype=torch.float8_e4m3fn, device=device) - scales = torch.empty(num_tokens, HIDDEN_SIZE // _SF_VEC, dtype=torch.uint8, device=device) - - def arg(t): - return from_dlpack(t.detach(), assumed_align=16).mark_layout_dynamic(leading_dim=0) - - args = ( - arg(head.contiguous().view(-1).view(torch.int32)), - arg(bias.contiguous().view(-1)), - arg(ag_uc.view(-1)), - arg(ag_mc.view(-1)), - arg(ag_flags.view(-1)), - arg(topk_ids.view(-1)), - arg(topk_weights.view(-1).view(torch.int16)), - arg(quantized.view(-1).view(torch.int32)), - arg(scales.view(-1)), - arg(ag_ready.view(-1) if ag_ready is not None else ag_flags.view(-1)), - ) - publish = ag_ready is not None - stream = cuda_driver.CUstream(torch.cuda.current_stream(device).cuda_stream) - use_pdl = os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" - key = (world, bool(early_trigger), publish, use_pdl) - fn = _rqag_compiled.get(key) - if fn is None: - if torch.cuda.is_current_stream_capturing(): - raise RuntimeError( - "trtllm::k3_route_quant_ag must run once outside CUDA-graph capture first" - ) - with _lock: - fn = _rqag_compiled.get(key) - if fn is None: - fn = _rqag_compiled[key] = cute.compile( - rqag.k3_route_quant_ag, *args, num_tokens, ag_rank, float(routed_scaling_factor), world, - bool(early_trigger), publish, use_pdl, stream, - ) # fmt: skip - fn(*args, num_tokens, ag_rank, float(routed_scaling_factor), stream) - return topk_ids, topk_weights, quantized, scales - - -@k3_route_quant_ag.register_fake -def _( - head, - bias, - routed_scaling_factor, - ag_uc, - ag_mc, - ag_flags, - ag_rank, - early_trigger=False, - ag_ready=None, -): - num_tokens = head.shape[0] - return ( - head.new_empty((num_tokens, TOP_K), dtype=torch.int32), - head.new_empty((num_tokens, TOP_K), dtype=torch.bfloat16), - head.new_empty((num_tokens, HIDDEN_SIZE), dtype=torch.float8_e4m3fn), - head.new_empty((num_tokens, HIDDEN_SIZE // _SF_VEC), dtype=torch.uint8), - ) - - -@torch.library.custom_op("trtllm::k3_fused_moe_head", mutates_args=()) -def k3_fused_moe_head( - head: torch.Tensor, - e_score_correction_bias: torch.Tensor, - w3_w1_weight: torch.Tensor, - w3_w1_weight_scale: torch.Tensor, - w2_weight: torch.Tensor, - w2_weight_scale: torch.Tensor, - local_expert_offset: int, - local_num_experts: int, - routed_scaling_factor: float, - ag_uc: torch.Tensor, - ag_mc: torch.Tensor, - ag_flags: torch.Tensor, - ag_rank: int, - ar_uc: Optional[torch.Tensor] = None, - ar_mc: Optional[torch.Tensor] = None, - ar_flags: Optional[torch.Tensor] = None, - ar_rank: int = -1, - ag_ready: Optional[torch.Tensor] = None, -) -> torch.Tensor: - """``trtllm::k3_fused_moe`` (or ``_ar`` with the ``ar_*`` buffers) from this rank's slice of - the sharded MoE head: ``trtllm::k3_route_quant_ag`` gathers it, routes and quantizes, then - ``k3_moe``. Returns ``[M, 3584]`` bf16 (the partial, or the reduced latent). With ``ag_ready`` - (``head_workspace``'s ready words) k3_moe acquires route_quant_ag's outputs through them - instead of waiting for its grid.""" - if head.shape[0] > MAX_TOKENS: - raise ValueError(f"k3_fused_moe handles at most {MAX_TOKENS} tokens, got {head.shape[0]}") - ar = None - world = 0 - if ar_uc is not None: - ar = (ar_uc, ar_mc, ar_flags, ar_rank) - world = ar_uc.numel() // ar_buffer_words(1) - st = _state( - head.device, w3_w1_weight.shape[1] // 2, local_num_experts, world, ag_ready is not None - ) - return st( - head, None, e_score_correction_bias, w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale, - local_expert_offset, local_num_experts, routed_scaling_factor, ar=ar, - head_ag=(ag_uc, ag_mc, ag_flags, ag_rank), head_ready=ag_ready, - ) # fmt: skip - - -@k3_fused_moe_head.register_fake -def _(head, e_score_correction_bias, w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale, - local_expert_offset, local_num_experts, routed_scaling_factor, ag_uc, ag_mc, ag_flags, ag_rank, - ar_uc=None, ar_mc=None, ar_flags=None, ar_rank=-1, ag_ready=None): # fmt: skip - return head.new_empty((head.shape[0], HIDDEN_SIZE), dtype=torch.bfloat16) - - -@torch.library.custom_op("trtllm::k3_fused_moe_front", mutates_args=()) -def k3_fused_moe_front( - x: torch.Tensor, - w_front: torch.Tensor, - e_score_correction_bias: torch.Tensor, - w3_w1_weight: torch.Tensor, - w3_w1_weight_scale: torch.Tensor, - w2_weight: torch.Tensor, - w2_weight_scale: torch.Tensor, - local_expert_offset: int, - local_num_experts: int, - routed_scaling_factor: float, - shared_cols: int, - gate_cap: float, - linear_cap: float, - ag_uc: torch.Tensor, - ag_mc: torch.Tensor, - ag_flags: torch.Tensor, - ag_rank: int, - ag_world: int, - ar_uc: Optional[torch.Tensor] = None, - ar_mc: Optional[torch.Tensor] = None, - ar_flags: Optional[torch.Tensor] = None, - ar_rank: int = -1, - ag_ready: Optional[torch.Tensor] = None, -) -> Tuple[torch.Tensor, torch.Tensor]: - """``trtllm::k3_moe_front`` (head GEMV, head all-gather, routing, MXFP8 latent, shared gate_up + SiTU) then - ``k3_moe`` (or its fused all-reduce build with the ``ar_*`` buffers) for the MoE input ``x`` bf16 ``[M, 7168]``. - Returns ``(y [M, 3584] bf16, shared activation [M, shared_cols] bf16)``. With ``ag_ready`` - (``head_workspace``'s ready words) k3_moe acquires the front's routing and MXFP8 rows through them instead of - waiting for the front's grid, whose shared tiles may still be running.""" - if x.shape[0] > MAX_TOKENS: - raise ValueError(f"k3_fused_moe handles at most {MAX_TOKENS} tokens, got {x.shape[0]}") - ar = None - world = 0 - if ar_uc is not None: - ar = (ar_uc, ar_mc, ar_flags, ar_rank) - world = ar_uc.numel() // ar_buffer_words(1) - st = _state( - x.device, w3_w1_weight.shape[1] // 2, local_num_experts, world, ag_ready is not None - ) - return st( - x, None, e_score_correction_bias, w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale, - local_expert_offset, local_num_experts, routed_scaling_factor, ar=ar, - head_ag=(ag_uc, ag_mc, ag_flags, ag_rank), head_ready=ag_ready, - front=(w_front, shared_cols, gate_cap, linear_cap, ag_world), - ) # fmt: skip - - -@k3_fused_moe_front.register_fake -def _(x, w_front, e_score_correction_bias, w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale, - local_expert_offset, local_num_experts, routed_scaling_factor, shared_cols, gate_cap, linear_cap, ag_uc, ag_mc, - ag_flags, ag_rank, ag_world, ar_uc=None, ar_mc=None, ar_flags=None, ar_rank=-1, ag_ready=None): # fmt: skip - return ( - x.new_empty((x.shape[0], HIDDEN_SIZE), dtype=torch.bfloat16), - x.new_empty((x.shape[0], shared_cols), dtype=torch.bfloat16), - ) - - -@torch.library.custom_op("trtllm::k3_fused_moe", mutates_args=()) -def k3_fused_moe( - hidden_states: torch.Tensor, - router_logits: torch.Tensor, - e_score_correction_bias: torch.Tensor, - w3_w1_weight: torch.Tensor, - w3_w1_weight_scale: torch.Tensor, - w2_weight: torch.Tensor, - w2_weight_scale: torch.Tensor, - local_expert_offset: int, - local_num_experts: int, - routed_scaling_factor: float, -) -> torch.Tensor: - """Kimi K3 routed experts for M <= 8 decode tokens; returns ``[M, 3584]`` bf16. - - ``hidden_states``: bf16 ``[M, 3584]`` latent; ``router_logits``: fp32 ``[M, 896]``; - ``e_score_correction_bias``: fp32 ``[896]``; weights: the TRTLLM-Gen - W4A8_MXFP4_MXFP8 buffers of this rank's ``local_num_experts`` experts, which hold - global ids ``[local_expert_offset, local_expert_offset + local_num_experts)``. - """ - if hidden_states.shape[0] > MAX_TOKENS: - raise ValueError( - f"k3_fused_moe handles at most {MAX_TOKENS} tokens, got {hidden_states.shape[0]}" - ) - st = _state(hidden_states.device, w3_w1_weight.shape[1] // 2, local_num_experts) - return st( - hidden_states, router_logits, e_score_correction_bias, w3_w1_weight, w3_w1_weight_scale, - w2_weight, w2_weight_scale, local_expert_offset, local_num_experts, routed_scaling_factor, - ) # fmt: skip - - -@k3_fused_moe.register_fake -def _(hidden_states, router_logits, e_score_correction_bias, w3_w1_weight, w3_w1_weight_scale, w2_weight, - w2_weight_scale, local_expert_offset, local_num_experts, routed_scaling_factor): # fmt: skip - return hidden_states.new_empty((hidden_states.shape[0], HIDDEN_SIZE), dtype=torch.bfloat16) - - -@torch.library.custom_op("trtllm::k3_fused_moe_ar", mutates_args=("lat_slab",)) -def k3_fused_moe_ar( - hidden_states: torch.Tensor, - router_logits: torch.Tensor, - e_score_correction_bias: torch.Tensor, - w3_w1_weight: torch.Tensor, - w3_w1_weight_scale: torch.Tensor, - w2_weight: torch.Tensor, - w2_weight_scale: torch.Tensor, - local_expert_offset: int, - local_num_experts: int, - routed_scaling_factor: float, - ar_uc: torch.Tensor, - ar_mc: torch.Tensor, - ar_flags: torch.Tensor, - ar_rank: int, - lat_slab: Optional[torch.Tensor] = None, - lat_buf: int = 0, - lat_rearm0: bool = False, -) -> torch.Tensor: - """``trtllm::k3_fused_moe`` followed by the all-reduce of the routed partial over the - group of ``ar_workspace``'s buffers (``ar_uc``, ``ar_mc``, ``ar_flags``), in one kernel. - Returns the reduced ``[M, 3584]`` bf16 latent, identical on every rank. With ``lat_slab`` - (int32 ``[3, 8, 1792]``, all-ones empty) the reduced rows are also published into buffer - ``lat_buf`` for a consumer that polls them, and the call re-arms buffer ``(lat_buf + 1) % 3`` - (and buffer 0 with ``lat_rearm0``, the step's last call).""" - if hidden_states.shape[0] > MAX_TOKENS: - raise ValueError( - f"k3_fused_moe handles at most {MAX_TOKENS} tokens, got {hidden_states.shape[0]}" - ) - world = ar_uc.numel() // ar_buffer_words(1) - st = _state( - hidden_states.device, w3_w1_weight.shape[1] // 2, local_num_experts, world, - lat_slab=lat_slab is not None, - ) # fmt: skip - return st( - hidden_states, router_logits, e_score_correction_bias, w3_w1_weight, w3_w1_weight_scale, - w2_weight, w2_weight_scale, local_expert_offset, local_num_experts, routed_scaling_factor, - ar=(ar_uc, ar_mc, ar_flags, ar_rank), - lat_slab=(lat_slab, lat_buf, lat_rearm0) if lat_slab is not None else None, - ) # fmt: skip - - -@k3_fused_moe_ar.register_fake -def _(hidden_states, router_logits, e_score_correction_bias, w3_w1_weight, w3_w1_weight_scale, w2_weight, - w2_weight_scale, local_expert_offset, local_num_experts, routed_scaling_factor, ar_uc, ar_mc, ar_flags, - ar_rank, lat_slab=None, lat_buf=0, lat_rearm0=False): # fmt: skip - return hidden_states.new_empty((hidden_states.shape[0], HIDDEN_SIZE), dtype=torch.bfloat16) +def _check_tokens(x: torch.Tensor) -> None: + if not 0 < x.shape[0] <= MAX_TOKENS: + raise ValueError(f"k3_moe handles 1 to {MAX_TOKENS} tokens, got {x.shape[0]}") # --------------------------------------------------------------------------- diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py index 03661fea8ed5..b902e2cbc68f 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py @@ -18,19 +18,20 @@ ``trtllm::k3_sandwich_tail`` the pre-attention step (the MoE tail, then the next layer's input norm); ``trtllm::k3_sandwich_plain`` a row-parallel projection with a plain residual add + RMSNorm (the drafter layers). -The all-reduce runs over a dedicated multicast buffer per TP group (:func:`workspace`), not the model's MNNVL -all-reduce workspace, so it keeps its own call parity. The kernel compiles on the first call for each -(world, publish, input source, PDL), which must happen outside CUDA-graph capture; the number of snapshots, the prefix -and the slab buffer are runtime arguments. +The all-reduce runs over a dedicated multicast buffer per TP group, a :class:`K3SandwichWorkspace` that the caller +creates (collectively, before CUDA-graph capture) and passes to every call; it is not the model's MNNVL all-reduce +workspace, so it keeps its own call parity. The kernel compiles on the first call for each (world, publish, input +source, PDL), which must happen outside CUDA-graph capture; the number of snapshots, the prefix and the slab buffer are +runtime arguments. Publishing (``x_slab``, ``slab_buf``): with a slab (``slab_tensor``) the normed rows are also written into buffer ``slab_buf`` (0-2) of it, sentinel-armed Lamport words the next kernel polls, and buffer ``(slab_buf + 1) % 3`` is re-armed. ``slab_buf`` is the ordinal of the call among this op's calls in the forward, mod 3. -The latent all-reduce folded into the tail (``lat_uc``, ``lat_flags`` from :func:`latent_exchange`): k3_moe pushes -its routed latent partial into every rank's exchange buffer and exits; ``k3_sandwich_tail`` sums the ranks' partials -itself (bit-identical to the one-shot all-reduce) instead of reading a reduced ``latent``. Every push-only k3_moe call -must be followed by exactly one such tail call on the same exchange. +The latent all-reduce folded into the tail (``lat_uc``, ``lat_flags`` of a :class:`K3SandwichLatentExchange`): k3_moe +pushes its routed latent partial into every rank's exchange buffer and exits; ``k3_sandwich_tail`` sums the ranks' +partials itself (bit-identical to the one-shot all-reduce) instead of reading a reduced ``latent``. Every push-only +k3_moe call must be followed by exactly one such tail call on the same exchange. Polling the input (``src_slab``, ``src_buf``): when the producer of phase 1's input (the attention output core for ``k3_sandwich_oproj``, the reduced latent for ``k3_sandwich_tail``) publishes it as such a slab (int32 [3][8][cols / @@ -42,7 +43,8 @@ import os import threading -from typing import Dict, List, Optional +from dataclasses import dataclass +from typing import Any, Callable, Dict, List, Optional import torch @@ -53,8 +55,6 @@ _lock = threading.Lock() _compiled: Dict[tuple, object] = {} -_workspaces: Dict[object, dict] = {} -_lat_exchanges: Dict[object, dict] = {} def _arg(t: torch.Tensor): @@ -128,97 +128,119 @@ def _slab_args(x_slab: Optional[torch.Tensor], slab_buf: int, fallback: torch.Te return x_slab.reshape(-1), int(slab_buf), 1 -def workspace(mapping) -> dict: - """The sandwich's all-reduce buffer for ``mapping``'s TP group: ``uc`` (this rank's words), ``mc`` (their - multicast mapping) and ``flags`` (int32, the call count of each CTA). Collective on first use: every rank of - the group must make its first call at the same point, outside CUDA-graph capture.""" - ws = _workspaces.get(mapping) - if ws is not None: - return ws +def _create_buffer(cls, mapping, words: int, flag_words: int, fabric_handle: Optional[bool], + arm_flags: Optional[Callable[[torch.Tensor], None]] = None): # fmt: skip + """A ``cls`` over a new multicast buffer of ``words`` int32 per rank of ``mapping``'s TP group, every word empty, + and ``flag_words`` int32 flags, zero (then ``arm_flags``). Collective and eager: every rank of the group calls it at + the same point, outside CUDA-graph capture; it returns on every rank or raises on every rank.""" if torch.cuda.is_current_stream_capturing(): - raise RuntimeError("k3_sandwich: the all-reduce buffer must be allocated outside CUDA-graph capture") + raise RuntimeError( + f"{cls.__name__}.create is collective and allocates: call it outside CUDA-graph capture" + ) from tensorrt_llm._torch.distributed.ops import ( _get_mnnvl_workspace_comm, _make_mnnvl_mcast_buffer, _mnnvl_workspace_all_succeeded, ) - from . import k3_sandwich_kernel as kernel - - world = mapping.tp_size - words = kernel.buffer_words(world) + use_fabric_handle = mapping.is_multi_node() if fabric_handle is None else bool(fabric_handle) comm = _get_mnnvl_workspace_comm(mapping) - use_fabric_handle = ( - os.environ.get("TRTLLM_FORCE_MNNVL_AR", "0") == "1" or mapping.is_multi_node() - ) error: Optional[Exception] = None + state = None try: handle = _make_mnnvl_mcast_buffer(comm, words * 4, mapping, use_fabric_handle) uc = handle.get_uc_buffer(mapping.tp_rank, (words,), torch.int32, 0) mc = handle.get_mc_buffer((words,), torch.int32, 0) - with torch.inference_mode(): - uc.fill_(EMPTY_WORD) - flags = torch.zeros(kernel.FLAG_WORDS, dtype=torch.int32, device=uc.device) + uc.fill_(EMPTY_WORD) + flags = torch.zeros(flag_words, dtype=torch.int32, device=uc.device) + if arm_flags is not None: + arm_flags(flags) torch.cuda.synchronize() - ws = dict( - handle=handle, comm=comm, uc=uc, mc=mc, flags=flags, rank=mapping.tp_rank, world=world + state = cls( + uc=uc, + mc=mc, + flags=flags, + rank=mapping.tp_rank, + world_size=mapping.tp_size, + handle=handle, + comm=comm, ) except Exception as exc: # noqa: BLE001 -- reported to every rank below, then re-raised error = exc # Also the barrier that keeps any rank from pushing into a peer's buffer before the peer has emptied it. if not _mnnvl_workspace_all_succeeded(comm, error is None): - raise RuntimeError("k3_sandwich all-reduce buffer failed on at least one rank") from error - _workspaces[mapping] = ws - return ws - - -def latent_exchange(mapping) -> dict: - """The latent exchange of ``mapping``'s TP group, shared by k3_moe (push-only: its ``ar_uc`` / ``ar_mc`` / - ``ar_flags``) and ``k3_sandwich_tail`` (``lat_uc`` / ``lat_flags``): ``uc`` (this rank's int32 [2][8][world][1792] - words, every word empty), ``mc`` (their multicast mapping, where k3_moe pushes) and ``flags`` (int32 [64]: [0] the - tail's call count mod 6, whose parity picks the half; the tail's scale slab after it). Collective on first use, as - :func:`workspace`.""" - ex = _lat_exchanges.get(mapping) - if ex is not None: - return ex - if torch.cuda.is_current_stream_capturing(): - raise RuntimeError("k3_sandwich: the latent exchange must be allocated outside CUDA-graph capture") - from tensorrt_llm._torch.distributed.ops import ( - _get_mnnvl_workspace_comm, - _make_mnnvl_mcast_buffer, - _mnnvl_workspace_all_succeeded, - ) + raise RuntimeError(f"{cls.__name__}: allocation failed on at least one rank") from error + return state + + +@dataclass(eq=False) +class K3SandwichWorkspace: + """One TP group's sandwich all-reduce buffer, shared by ``k3_sandwich_oproj``, ``k3_sandwich_tail`` and + ``k3_sandwich_plain``: two alternating halves of [8 tokens][world][7168] bf16 per rank behind one multicast mapping, + and one call counter per CTA whose parity selects the half. Every sandwich call on it advances every counter, so all + of a group's ranks make the same calls on it in the same order. Pass ``uc``, ``mc``, ``flags`` and ``rank`` as the + ops' ``ws_uc``, ``ws_mc``, ``ws_flags`` and ``rank``.""" + + uc: torch.Tensor + """int32 [2 * 8 * world * 3584]: this rank's words (0x80000000 = empty).""" + mc: torch.Tensor + """The same words through the multicast mapping (where the peers push).""" + flags: torch.Tensor + """int32 [64]: the call count of each of the kernel's 56 CTAs.""" + rank: int + world_size: int + handle: Any + """The ``McastGPUBuffer`` that owns the memory; the workspace is valid while this object lives.""" + comm: Any + """The TP-group communicator the handles were exchanged over.""" + + @classmethod + def create(cls, mapping, fabric_handle: Optional[bool] = None) -> "K3SandwichWorkspace": + """Allocate and arm a workspace for ``mapping``'s TP group. Collective: every rank of the group calls it at the + same point, eagerly (not under CUDA-graph capture); it returns on every rank or raises on every rank. + ``fabric_handle``: share the memory by fabric handle (required across nodes) rather than POSIX file + descriptor; default ``mapping.is_multi_node()``.""" + from . import k3_sandwich_kernel as kernel + + return _create_buffer( + cls, mapping, kernel.buffer_words(mapping.tp_size), kernel.FLAG_WORDS, fabric_handle + ) - from . import k3_sandwich_kernel as kernel - world = mapping.tp_size - words = kernel.lat_buffer_words(world) - comm = _get_mnnvl_workspace_comm(mapping) - use_fabric_handle = ( - os.environ.get("TRTLLM_FORCE_MNNVL_AR", "0") == "1" or mapping.is_multi_node() - ) - error: Optional[Exception] = None - try: - handle = _make_mnnvl_mcast_buffer(comm, words * 4, mapping, use_fabric_handle) - uc = handle.get_uc_buffer(mapping.tp_rank, (words,), torch.int32, 0) - mc = handle.get_mc_buffer((words,), torch.int32, 0) - with torch.inference_mode(): - uc.fill_(EMPTY_WORD) - flags = torch.zeros(kernel.LAT_FLAG_WORDS, dtype=torch.int32, device=uc.device) - flags[kernel.LAT_SCALES : kernel.LAT_SCALES + kernel.LAT_SCALE_BUFS * MAX_TOKENS] = ( - kernel.SCALE_SENTINEL - ) - torch.cuda.synchronize() - ex = dict( - handle=handle, comm=comm, uc=uc, mc=mc, flags=flags, rank=mapping.tp_rank, world=world +@dataclass(eq=False) +class K3SandwichLatentExchange: + """One TP group's latent exchange for ``k3_sandwich_tail`` with the latent all-reduce folded in: the push-only + k3_moe stores every rank's routed partial into two alternating halves of [8 tokens][world][3584] bf16 per rank + behind one multicast mapping, and the tail sums them. ``flags``: [0] the tail's call count mod 6 (its parity selects + the half), then the tail's latent scale slab. Pass ``uc`` and ``flags`` as the tail's ``lat_uc`` and ``lat_flags``. + Separate from :class:`K3SandwichWorkspace`.""" + + uc: torch.Tensor + """int32 [2 * 8 * world * 1792]: this rank's words (0x80000000 = empty).""" + mc: torch.Tensor + """The same words through the multicast mapping (where the producers push).""" + flags: torch.Tensor + """int32 [64]: [0] the call count mod 6, [32 + 8 b + t] buffer b of the latent scales (sentinel 0xFFFFFFFF).""" + rank: int + world_size: int + handle: Any + """The ``McastGPUBuffer`` that owns the memory; the exchange is valid while this object lives.""" + comm: Any + """The TP-group communicator the handles were exchanged over.""" + + @classmethod + def create(cls, mapping, fabric_handle: Optional[bool] = None) -> "K3SandwichLatentExchange": + """Allocate and arm an exchange for ``mapping``'s TP group; collective and eager, as + :meth:`K3SandwichWorkspace.create`.""" + from . import k3_sandwich_kernel as kernel + + def arm(flags: torch.Tensor) -> None: + scales = slice(kernel.LAT_SCALES, kernel.LAT_SCALES + kernel.LAT_SCALE_BUFS * MAX_TOKENS) + flags[scales] = kernel.SCALE_SENTINEL + + return _create_buffer( + cls, mapping, kernel.lat_buffer_words(mapping.tp_size), kernel.LAT_FLAG_WORDS, fabric_handle, arm ) - except Exception as exc: # noqa: BLE001 -- reported to every rank below, then re-raised - error = exc - # Also the barrier that keeps any rank from pushing into a peer's buffer before the peer has emptied it. - if not _mnnvl_workspace_all_succeeded(comm, error is None): - raise RuntimeError("k3_sandwich latent exchange failed on at least one rank") from error - _lat_exchanges[mapping] = ex - return ex def supports(core: torch.Tensor, o_weight: torch.Tensor, num_snapshots: int) -> bool: @@ -256,7 +278,9 @@ def _compile_and_run(entry, key, args, runtime, consts, stream, name): fn(*args, *runtime, stream) -@torch.library.custom_op("trtllm::k3_sandwich_oproj", mutates_args=("x_slab",)) +@torch.library.custom_op( + "trtllm::k3_sandwich_oproj", mutates_args=("ws_uc", "ws_mc", "ws_flags", "x_slab") +) def k3_sandwich_oproj( core: torch.Tensor, o_weight: torch.Tensor, @@ -280,7 +304,7 @@ def k3_sandwich_oproj( ``core`` bf16 [M, 768] is this rank's attention output, ``o_weight`` [7168, 768] its o_proj slice; ``prefix`` [M, 7168] (or None), ``block_residual`` [S, M, 7168] the S valid snapshots, the weights [7168]; - ``ws_*`` from :func:`workspace`. ``x_slab`` (from :func:`slab_tensor`) also receives normed in buffer + ``ws_*`` from a :class:`K3SandwichWorkspace`. ``x_slab`` (from :func:`slab_tensor`) also receives normed in buffer ``slab_buf``; ``src_slab`` (int32 [3, 8, 384]) supplies core in buffer ``src_buf``.""" import cuda.bindings.driver as cuda_driver @@ -367,7 +391,10 @@ def supports_tail( ) -@torch.library.custom_op("trtllm::k3_sandwich_tail", mutates_args=("x_slab", "tap", "updated_out")) +@torch.library.custom_op( + "trtllm::k3_sandwich_tail", + mutates_args=("ws_uc", "ws_mc", "ws_flags", "x_slab", "lat_uc", "lat_flags", "tap", "updated_out"), +) def k3_sandwich_tail( latent: torch.Tensor, act: torch.Tensor, @@ -400,12 +427,12 @@ def k3_sandwich_tail( the whole reduced latent row (16-byte aligned: its rows are bulk-copied), ``act`` the shared-expert activation, ``tail_weight`` [7168, 256 + 384] the latent up columns of the slice zero-padded to 256 and the shared down projection; ``src_slab`` (int32 [3, 8, 1792]) supplies the latent in buffer ``src_buf``; with ``lat_uc`` / - ``lat_flags`` (:func:`latent_exchange`) the kernel sums the ranks' pushed partials itself and ``latent`` gives only - the shape; with ``tap`` (bf16 [M, 7168], unit column stride, rows a multiple of 8 elements apart, 16-byte aligned: - e.g. a column slice of a capture buffer) it also stores there the pre-norm attn_res mixture rows (a DSpark - capture layer's tap) or, with ``tap_updated``, ``updated``; with ``updated_out`` (bf16 [M, 7168], contiguous, - 16-byte aligned, e.g. the next row of the attention-residual snapshot bank, which this call does not read) it stores - ``updated`` there instead of a new tensor and returns an empty [0, 7168] in its place; the rest as + ``lat_flags`` (a :class:`K3SandwichLatentExchange`) the kernel sums the ranks' pushed partials itself and + ``latent`` gives only the shape; with ``tap`` (bf16 [M, 7168], unit column stride, rows a multiple of 8 elements + apart, 16-byte aligned: e.g. a column slice of a capture buffer) it also stores there the pre-norm attn_res mixture + rows (a DSpark capture layer's tap) or, with ``tap_updated``, ``updated``; with ``updated_out`` (bf16 [M, 7168], + contiguous, 16-byte aligned, e.g. the next row of the attention-residual snapshot bank, which this call does not + read) it stores ``updated`` there instead of a new tensor and returns an empty [0, 7168] in its place; the rest as ``k3_sandwich_oproj``.""" import cuda.bindings.driver as cuda_driver @@ -536,7 +563,7 @@ def supports_plain(x: torch.Tensor, weight: torch.Tensor, residual: torch.Tensor ) -@torch.library.custom_op("trtllm::k3_sandwich_plain", mutates_args=()) +@torch.library.custom_op("trtllm::k3_sandwich_plain", mutates_args=("ws_uc", "ws_mc", "ws_flags")) def k3_sandwich_plain( x: torch.Tensor, weight: torch.Tensor, @@ -553,9 +580,9 @@ def k3_sandwich_plain( """``(normed, updated)`` of the row-parallel projection ``x @ weight^T`` followed by the TP all-reduce with the residual add and RMSNorm (``AllReduceFusionOp.RESIDUAL_RMS_NORM``): updated = residual + the sum, normed = RMSNorm(updated) * norm_weight. ``x`` bf16 [M, K], ``weight`` [7168, K] (K a multiple of 128 up to 896), - ``residual`` [M, 7168]; ``ws_*`` from :func:`workspace`. With ``swiglu``, ``x`` is a gate_up output [M, 2 K] - (gate columns first) and the projection is ``silu_and_mul(x) @ weight^T`` with ``k3_ctm_gemv_swiglu`` split 2's - arithmetic (the drafter MLP's down projection). The arithmetic and summation order are those of the + ``residual`` [M, 7168]; ``ws_*`` from a :class:`K3SandwichWorkspace`. With ``swiglu``, ``x`` is a gate_up output + [M, 2 K] (gate columns first) and the projection is ``silu_and_mul(x) @ weight^T`` with ``k3_ctm_gemv_swiglu`` + split 2's arithmetic (the drafter MLP's down projection). The arithmetic and summation order are those of the all-reduce kernel the call replaces: the MNNVL one-shot's, or with ``ipc_order`` the IPC one-shot's (``allreduce_fusion_kernel_oneshot_lamport`` with fp32 accumulation, TP <= 8 within one node).""" import cuda.bindings.driver as cuda_driver diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py index 37f2d3652a4f..a056a8e88f00 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py @@ -12,7 +12,7 @@ # 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. -"""trtllm::k3_fused_moe (k3_route_quant + the persistent k3_moe kernel, M <= 8) on one GPU, at the Kimi K3 TP16 +"""K3MoeLayer (trtllm::k3_route_quant + the persistent k3_moe kernel, M <= 8) on one GPU, at the Kimi K3 TP16 deployment's routed-expert rank layout (experts TP4 x EP4: 224 local experts, intermediate 768 per rank), at every M in 1..8 with random routing, with 0, 4 and 16 of each token's experts local, and with 16 local experts per token none shared (16 M groups: the kernel's group capacity at M = 8): against the stock path @@ -185,20 +185,29 @@ def _stock(proc, bias, x, logits): return y, ids, w, x_fp8, x_sf +@functools.lru_cache(maxsize=None) +def _layer(): + """One K3MoeState on this GPU and the layer of _experts()' buffers on it.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op + + _ops() + proc, _, _ = _experts() + state = op.K3MoeState(torch.device("cuda", torch.cuda.current_device()), I_TP, E_LOCAL) + return state.layer(proc["w31"], proc["w31s"], proc["w2"], proc["w2s"]) + + def _fused(proc, bias, x, logits): - return _ops().k3_fused_moe(x, logits, bias, proc["w31"], proc["w31s"], proc["w2"], proc["w2s"], OFFSET, E_LOCAL, - RSF) # fmt: skip + return _layer()(x, logits, bias, OFFSET, RSF) def _scratch_rearmed(): - """The intermediate slab armed again (FP8 -0.0 codes, E8M0 NaN scale words) and every layer's counters zero.""" - from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op - - st = op._state(torch.device("cuda", torch.cuda.current_device()), I_TP, E_LOCAL) + """The intermediate slab armed again (FP8 -0.0 codes, E8M0 NaN scale words) and the layer's counters zero.""" + layer = _layer() + st = layer.state mod = st.mod cs = st.cs.view(mod.G_CAP, 8, mod.K2_TILES, mod.SFB_GROUP_BYTES) armed = bool((st.c == -128).all()) and bool((cs[..., :4] == -1).all()) - return armed and all(bool((layer[4] == 0).all()) for layer in st.layers.values()) + return armed and bool((layer.counters == 0).all()) CASES = ["random", "4_local", "16_local", "16_local_disjoint", "none_local"] @@ -301,9 +310,9 @@ def test_k3_fused_moe_partial_rows_past_m(): _ops() proc, _, bias = _experts() - fresh = op._K3FusedMoE(torch.device("cuda", torch.cuda.current_device()), I_TP, E_LOCAL) + fresh = op.K3MoeState(torch.device("cuda", torch.cuda.current_device()), I_TP, E_LOCAL) assert bool((fresh.part == 0).all()) - st = op._state(torch.device("cuda", torch.cuda.current_device()), I_TP, E_LOCAL) + st = _layer().state logits8, x8 = _tokens("random") for m in M_ALL: logits, x = logits8[:m].contiguous(), x8[:m].contiguous() @@ -324,18 +333,22 @@ def test_token_limit(): def test_collective_workspaces_refuse_graph_capture(): - """The head all-gather's buffers (the front's) and the fused all-reduce's are collective on first use (an MNNVL - multicast allocation over the TP group): a first use under CUDA-graph capture raises instead of entering the - collective, and nothing is cached for the group.""" + """The head all-gather's buffers (the front's) are created collectively (an MNNVL multicast allocation over the TP + group): creating them under CUDA-graph capture raises instead of entering the collective. The per-rank state and + a layer's counters refuse capture too (they allocate).""" from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op from tensorrt_llm.mapping import Mapping _ops() + proc, _, _ = _experts() + device = torch.device("cuda", torch.cuda.current_device()) + state = op.K3MoeState(device, I_TP, E_LOCAL) mapping = Mapping(world_size=1, rank=0, tp_size=1) graph, stream = torch.cuda.CUDAGraph(), torch.cuda.Stream() with torch.cuda.graph(graph, stream=stream): with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): - op.head_workspace(mapping) + op.K3MoeHeadWorkspace.create(mapping) + with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): + op.K3MoeState(device, I_TP, E_LOCAL) with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): - op.ar_workspace(mapping) - assert mapping not in op._head_workspaces and mapping not in op._ar_workspaces + state.layer(proc["w31"], proc["w31s"], proc["w2"], proc["w2s"]) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_latent_reduce.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_latent_reduce.py index 5a475cacdf8a..f5cf86348935 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_latent_reduce.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_latent_reduce.py @@ -88,7 +88,7 @@ def _context(): mapping = Mapping(world_size=world, rank=rank, gpus_per_node=gpus, tp_size=world) mnnvl = MNNVLAllReduce(mapping, torch.bfloat16) - ex = latent_op.LatentExchange(mapping) + ex = latent_op.K3LatentExchange.create(mapping, fabric_handle=True) return SimpleNamespace(comm=comm, rank=rank, world=world, mnnvl=mnnvl, ex=ex, count=0) @@ -233,8 +233,8 @@ def test_k3_latent_reduce(mpi_pool_executor): def test_latent_exchange_refuses_graph_capture(): - """The exchange is collective on construction (an MNNVL multicast allocation over the TP group): constructing it - under CUDA-graph capture raises instead of entering the collective. One process, a group of one.""" + """The exchange is created collectively (an MNNVL multicast allocation over the TP group): creating it under + CUDA-graph capture raises instead of entering the collective. One process, a group of one.""" import tensorrt_llm # noqa: F401 from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import latent_op from tensorrt_llm.mapping import Mapping @@ -243,7 +243,7 @@ def test_latent_exchange_refuses_graph_capture(): graph, stream = torch.cuda.CUDAGraph(), torch.cuda.Stream() with torch.cuda.graph(graph, stream=stream): with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): - latent_op.LatentExchange(mapping) + latent_op.K3LatentExchange.create(mapping) def main() -> int: diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py index 7ec33b188ace..56a5772639c5 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py @@ -12,9 +12,9 @@ # 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. -"""trtllm::k3_moe_front and trtllm::k3_fused_moe_front (the Kimi K3 MoE front: sharded head GEMV, head all-gather, -top-16 routing, MXFP8 latent, shared gate_up + SiTU; then k3_moe on its grid), one process per GPU over the TP group of -this run, at every M in 1..8. The head is sharded over the group (TP W: 3584 / W latent + 896 / W router rows and +"""trtllm::k3_moe_front and K3MoeLayer.front (the Kimi K3 MoE front: sharded head GEMV, head all-gather, top-16 +routing, MXFP8 latent, shared gate_up + SiTU; then k3_moe on its grid), one process per GPU over the TP group of this +run, at every M in 1..8, over one K3MoeHeadWorkspace. The head is sharded over the group (TP W: 3584 / W latent + 896 / W router rows and 2 x 6144 / W shared rows per rank; W = 4 on one GB200 tray, the model's TP16 shapes with 16 processes); the routed experts are one rank of experts TP4 x EP4 (224 local experts, intermediate 768), as in the TP16 deployment. front : against the unfused chain (the head GEMV in fp32 torch -> the gather -> @@ -134,9 +134,6 @@ def _experts(seed): def _context(with_experts): - os.environ.setdefault( - "TRTLLM_FORCE_MNNVL_AR", "1" - ) # fabric handles within one tray, as across trays comm = MPI.COMM_WORLD rank, world = comm.Get_rank(), comm.Get_size() gpus = torch.cuda.device_count() @@ -159,12 +156,26 @@ def _context(with_experts): assert front_op.weight_supported( world, inter, HIDDEN, torch.device("cuda", torch.cuda.current_device()) ) - ws = moe_op.head_workspace(mapping) - return SimpleNamespace( + # Fabric handles within one tray, as across trays. + ws = moe_op.K3MoeHeadWorkspace.create(mapping, fabric_handle=True) + ctx = SimpleNamespace( comm=comm, rank=rank, world=world, wl=wl, we=we, inter=inter, bias=bias, x8=x8, head=head, gate_up=gate_up, - front=front_op.front_weight(head, gate_up), ws=ws, ag=(ws["uc"], ws["mc"], ws["flags"], ws["rank"]), - offset=(rank % 4) * E_LOCAL, experts=_experts(20260928 + rank) if with_experts else None, + front=front_op.front_weight(head, gate_up), ws=ws, ag=(ws.uc, ws.mc, ws.flags, ws.rank), + offset=(rank % 4) * E_LOCAL, experts=_experts(20260928 + rank) if with_experts else None, layer=None, ) # fmt: skip + if with_experts: + ctx.layer = _layer(ctx) + return ctx + + +def _layer(ctx, head_flags=False, config=None): + """This rank's experts as a layer of a new K3MoeState: the plain build, or the head_flags build.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op as moe_op + + p = ctx.experts + state = moe_op.K3MoeState(torch.device("cuda", torch.cuda.current_device()), I_TP, E_LOCAL, head_flags=head_flags, + config=config) # fmt: skip + return state.layer(p["w31"], p["w31s"], p["w2"], p["w2s"]) def _all_ranks(ctx, good) -> bool: @@ -205,7 +216,7 @@ def _dequant(q, s): def _buffers_empty(ctx) -> bool: - return bool((ctx.ws["uc"] == EMPTY).all()) and int(ctx.ws["flags"][1].item()) == 0 + return bool((ctx.ws.uc == EMPTY).all()) and int(ctx.ws.flags[1].item()) == 0 def check_front(ctx): @@ -213,10 +224,10 @@ def check_front(ctx): out8 = _front(ctx, ctx.x8) for m in M_ALL: x = ctx.x8[:m].contiguous() - flag0 = int(ctx.ws["flags"][0].item()) + flag0 = int(ctx.ws.flags[0].item()) ids, w, q, s, shared = out = _front(ctx, x) torch.cuda.synchronize() - flag1 = int(ctx.ws["flags"][0].item()) + flag1 = int(ctx.ws.flags[0].item()) empty = _quiet_check(ctx, lambda: _buffers_empty(ctx)) r_ids, r_w, r_q, r_s, r_shared, margin = _reference(ctx, x) # The selected experts per token (their order inside the top 16 may differ at near-equal keys) and each @@ -253,11 +264,11 @@ def check_front(ctx): return results -def _fused(ctx, x, ready=None): - p = ctx.experts - return torch.ops.trtllm.k3_fused_moe_front( - x, ctx.front, ctx.bias, p["w31"], p["w31s"], p["w2"], p["w2s"], ctx.offset, E_LOCAL, RSF, ctx.inter, GATE_CAP, - LINEAR_CAP, *ctx.ag, ctx.world, ag_ready=ready) # fmt: skip +def _fused(ctx, x, layer=None, bias=None): + """K3MoeLayer.front on ``layer`` (default: the plain build's ``ctx.layer``).""" + layer = layer or ctx.layer + return layer.front(x, ctx.front, ctx.bias if bias is None else bias, ctx.offset, RSF, ctx.inter, GATE_CAP, + LINEAR_CAP, ctx.ws) # fmt: skip def _runner(ctx, ids, w, q, s): @@ -283,13 +294,12 @@ def _compare(y, ref): def _scratch_rearmed(ctx): - from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op - - st = op._state(torch.device("cuda", torch.cuda.current_device()), I_TP, E_LOCAL) + """The plain build's intermediate slab armed again and its layer's counters zero.""" + st = ctx.layer.state mod = st.mod cs = st.cs.view(mod.G_CAP, 8, mod.K2_TILES, mod.SFB_GROUP_BYTES) armed = bool((st.c == -128).all()) and bool((cs[..., :4] == -1).all()) - return armed and all(bool((layer[4] == 0).all()) for layer in st.layers.values()) + return armed and bool((ctx.layer.counters == 0).all()) def _max_ulp(a: torch.Tensor, b: torch.Tensor) -> int: @@ -338,14 +348,15 @@ def _i32(v: int) -> int: def check_head_flags(ctx): - """k3_fused_moe_front with the ready-word handoff (``ag_ready``: k3_moe built with head_flags acquires the front's + """K3MoeLayer.front with the ready-word handoff (``ag_ready``: k3_moe built with head_flags acquires the front's ready words, ready[t] / ready[8 + t] = the head epoch flags[2] + 1 for token t, instead of waiting for its grid) across the epoch's int32 wrap. From a new workspace's state (epoch 0, ready words 0), two calls at M 1, then the epoch preset to -2, then calls at M 1, 8, 3, 8: the M 8 call at epoch -1 waits for 0, the value of the words that no call has published. Per call: no word the call polls already holds its epoch + 1 (such a word would let k3_moe read the routing before the front writes it; the call is then not run); y and the shared activation the bits of the plain call; afterwards the epoch and every ready word hold the next call's epoch, the head buffers empty.""" - flags, ready = ctx.ws["flags"], ctx.ws["ready"] + flags, ready = ctx.ws.flags, ctx.ws.ready + flag_layer = _layer(ctx, head_flags=True) plain = {m: _fused(ctx, ctx.x8[:m].contiguous()) for m in (1, 3, 8)} torch.cuda.synchronize() @@ -368,7 +379,7 @@ def set_epoch(epoch): row["rank"], row["good"], row["ok"] = ctx.rank, False, False results.append(row) break - y, shared = _fused(ctx, x, ready) + y, shared = _fused(ctx, x, flag_layer) torch.cuda.synchronize() after, words = _quiet_check(ctx, lambda: (int(flags[2].item()), ready[:16].tolist())) row.update( @@ -434,9 +445,8 @@ def check_publish_order(ctx, num_ctas=None): the call is not run), afterwards the epoch and every ready word at the next call's epoch, the head buffers empty, rank 1's y all zeros.""" from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import front_op - from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op as moe_op - flags, ready = ctx.ws["flags"], ctx.ws["ready"] + flags, ready = ctx.ws.flags, ctx.ws.ready idle = 1 bias = ctx.bias.clone() bias[idle * E_LOCAL : (idle + 1) * E_LOCAL] = -8.0 # below every other expert's selection key @@ -444,15 +454,12 @@ def check_publish_order(ctx, num_ctas=None): ctas = torch.cuda.get_device_properties(device).multi_processor_count if num_ctas != 0: ctas = num_ctas or ctas // 2 - key = (device.index, I_TP, E_LOCAL, 0, True, False, False) # op._state's key of the head_flags build - saved_state, saved_kernel, saved_compiled = moe_op._states.get(key), front_op._kernel, dict(front_op._compiled) + saved_kernel, saved_compiled = front_op._kernel, dict(front_op._compiled) tmp_dir = tempfile.mkdtemp(prefix="k3_moe_front_") - p = ctx.experts + held_layer = _layer(ctx, head_flags=True, config={"num_ctas": ctas}) def fused(x): - return torch.ops.trtllm.k3_fused_moe_front( - x, ctx.front, bias, p["w31"], p["w31s"], p["w2"], p["w2s"], ctx.offset, E_LOCAL, RSF, ctx.inter, GATE_CAP, - LINEAR_CAP, *ctx.ag, ctx.world, ag_ready=ready) # fmt: skip + return _fused(ctx, x, held_layer, bias) def set_epoch(epoch): flags[2] = epoch @@ -460,9 +467,6 @@ def set_epoch(epoch): results = [] try: held = _held_front(tmp_dir) - moe_op._states[key] = moe_op._K3FusedMoE( - device, I_TP, E_LOCAL, {"head_flags": 1, "lat_slab": 0, "num_ctas": ctas} - ) front_op._kernel = lambda: held front_op._compiled.clear() # Compiles both kernels; the ranks then start each call within the hold. @@ -498,10 +502,6 @@ def set_epoch(epoch): front_op._kernel = saved_kernel front_op._compiled.clear() front_op._compiled.update(saved_compiled) - if saved_state is None: - moe_op._states.pop(key, None) - else: - moe_op._states[key] = saved_state shutil.rmtree(tmp_dir, ignore_errors=True) return results diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_wide.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_wide.py index aced0e0d330d..56f753e5fb2b 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_wide.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_wide.py @@ -26,7 +26,7 @@ Checks: against the stock path (trtllm::kimi_k3_noaux_tc_mxfp8_quant, then the TRTLLM-Gen W4A8_MXFP4_MXFP8 MoE runner with those ids) and an fp64 reference over the dequantized MXFP4 experts (op-catalog gates: 8 ulp of the row max per element, 4 ulp relative RMS); run-to-run identical bits; the slab armed and the layer's counters zero after every -call; at M <= 8, within one bf16 ulp of trtllm::k3_fused_moe (the decode build; bit-identity reported). Then: calls +call; at M <= 8, within one bf16 ulp of K3MoeLayer (the decode build; bit-identity reported). Then: calls of two layers at mixed M on one stream and replayed from a CUDA graph give each call's bits alone; 0 and 65 tokens are refused. Weights are random checkpoint-format MXFP4 experts put through TRT-LLM's own loader.""" @@ -136,6 +136,16 @@ def _wide(): return state, state.layer(*weights), state.layer(*weights) +@functools.lru_cache(maxsize=None) +def _decode_layer(): + """The same experts as a layer of the M <= 8 decode build (K3MoeState).""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op + + proc, _, _ = _experts() + state = op.K3MoeState(torch.device("cuda", torch.cuda.current_device()), I_TP, E_LOCAL) + return state.layer(proc["w31"], proc["w31s"], proc["w2"], proc["w2s"]) + + _E2M1 = [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0, -0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0] @@ -357,8 +367,7 @@ def _check(case, m): assert groups == state.mod.G_CAP == 324 decode = "" if m <= 8: - y8 = _ops().k3_fused_moe(x, logits, bias, proc["w31"], proc["w31s"], proc["w2"], proc["w2s"], OFFSET, - E_LOCAL, RSF) # fmt: skip + y8 = _decode_layer()(x, logits, bias, OFFSET, RSF) ulp_dec = _max_ulp(y, y8) decode = f" max_ulp_vs_decode={ulp_dec} bits_as_decode={torch.equal(_bits(y), _bits(y8))}" assert ulp_dec <= 1 diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py index f5f4c6d26a69..a111e9d20ff0 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py @@ -116,7 +116,8 @@ def _context(): mapping = Mapping(world_size=world, rank=rank, gpus_per_node=gpus, tp_size=world) return SimpleNamespace(comm=comm, rank=rank, world=world, mapping=mapping, - mnnvl=MNNVLAllReduce(mapping, torch.bfloat16), ws=sw_op.workspace(mapping)) # fmt: skip + mnnvl=MNNVLAllReduce(mapping, torch.bfloat16), + ws=sw_op.K3SandwichWorkspace.create(mapping, fabric_handle=True)) # fmt: skip def _all_ranks(ctx, good) -> bool: @@ -125,20 +126,20 @@ def _all_ranks(ctx, good) -> bool: def _oproj(ctx, core, w, prefix, block, res_w, rms_w, out_w): ws = ctx.ws - return torch.ops.trtllm.k3_sandwich_oproj(core, w, prefix, block, res_w, rms_w, out_w, EPS, EPS, ws["uc"], - ws["mc"], ws["flags"], ws["rank"]) # fmt: skip + return torch.ops.trtllm.k3_sandwich_oproj(core, w, prefix, block, res_w, rms_w, out_w, EPS, EPS, ws.uc, ws.mc, + ws.flags, ws.rank) # fmt: skip def _tail(ctx, latent, act, w, lo, prefix, block, res_w, rms_w, out_w, **extra): ws = ctx.ws return torch.ops.trtllm.k3_sandwich_tail(latent, act, w, lo, LAT_EPS, prefix, block, res_w, rms_w, out_w, EPS, - EPS, ws["uc"], ws["mc"], ws["flags"], ws["rank"], **extra) # fmt: skip + EPS, ws.uc, ws.mc, ws.flags, ws.rank, **extra) # fmt: skip def _plain(ctx, x, w, residual, norm_w, swiglu=False): ws = ctx.ws - return torch.ops.trtllm.k3_sandwich_plain(x, w, residual, norm_w, EPS, ws["uc"], ws["mc"], ws["flags"], - ws["rank"], swiglu=swiglu) # fmt: skip + return torch.ops.trtllm.k3_sandwich_plain(x, w, residual, norm_w, EPS, ws.uc, ws.mc, ws.flags, ws.rank, + swiglu=swiglu) # fmt: skip def _attn_res_ar(ctx, partial, prefix, block, res_w, rms_w, out_w): @@ -414,7 +415,7 @@ def run(): torch.cuda.synchronize() return outs - flags = ctx.ws["flags"] + flags = ctx.ws.flags fresh = run() count = int(flags[0].item()) ctx.comm.Barrier() @@ -437,9 +438,9 @@ def check_fold_wrap(ctx): from tensorrt_llm._torch.cute_dsl_kernels.k3_sandwich import k3_sandwich_kernel as kernel from tensorrt_llm._torch.cute_dsl_kernels.k3_sandwich import op as sw_op - ex = sw_op.latent_exchange(ctx.mapping) - flags = ex["flags"] - lanes = ex["mc"].view(2, 8, ctx.world, LATENT // 2) + ex = sw_op.K3SandwichLatentExchange.create(ctx.mapping, fabric_handle=True) + flags = ex.flags + lanes = ex.mc.view(2, 8, ctx.world, LATENT // 2) latent8, act8, w, lo, prefix8, block8, res_w, rms_w, out_w = _tail_inputs(ctx, 3, 830) part8 = _rand((8, LATENT), 840 + 1000 * ctx.rank, 0.3) part8[part8 == 0] = 0.0 # pushes never send -0.0: with +0.0 beside it, that word would read as empty @@ -456,7 +457,7 @@ def run(): before = torch.cat([flags[1 : slab.start], flags[slab.stop :]]).clone() pre, block = _first(m, prefix8, block8, True) out = _tail(ctx, latent8[:m].contiguous(), act8[:m].contiguous(), w, lo, pre, block, res_w, rms_w, out_w, - lat_uc=ex["uc"], lat_flags=flags) # fmt: skip + lat_uc=ex.uc, lat_flags=flags) # fmt: skip torch.cuda.synchronize() outs.append([t.clone() for t in out]) others.append(torch.equal(before, torch.cat([flags[1 : slab.start], flags[slab.stop :]]))) @@ -507,7 +508,7 @@ def test_k3_sandwich(mpi_pool_executor, check): def test_workspaces_refuse_graph_capture(): - """The all-reduce buffer and the latent exchange are collective on first use: allocating either under CUDA-graph + """The all-reduce buffer and the latent exchange are created collectively: creating either under CUDA-graph capture raises instead of entering the collective, which could hang the group.""" from tensorrt_llm._torch.cute_dsl_kernels.k3_sandwich import op from tensorrt_llm.mapping import Mapping @@ -516,9 +517,9 @@ def test_workspaces_refuse_graph_capture(): graph, stream = torch.cuda.CUDAGraph(), torch.cuda.Stream() with torch.cuda.graph(graph, stream=stream): with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): - op.workspace(mapping) + op.K3SandwichWorkspace.create(mapping) with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): - op.latent_exchange(mapping) + op.K3SandwichLatentExchange.create(mapping) def main() -> int: From 1db2a80a578b16ba2860e64900db445bf7b126ba Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 19:26:21 -0700 Subject: [PATCH 059/161] [None][fix] Kimi K3 MoE wide state: refuse graph capture; flag word docs K3MoeWideState and K3MoeWideLayer allocate and arm their scratch and counters, so like K3MoeState and K3MoeLayer they raise under CUDA-graph capture instead of capturing the arming. The docs of K3LatentExchange's flags now name each word ([0] the count, [2] the CTA arrivals, [1] and [3] unused). A stale comment about head_flags' environment fallback is gone, since the kernel reads its options only from K3_CONFIG. Signed-off-by: Vasanth Sabavat --- .../_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py | 9 +++++---- tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py | 7 +++++-- 2 files changed, 10 insertions(+), 6 deletions(-) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py index 9e5777306915..b43918d9dc51 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py @@ -56,16 +56,17 @@ def default_ctas(world: int) -> int: @dataclass(eq=False) class K3LatentExchange: """The latent all-reduce buffers of a TP group: int32 ``[2][8][world][1792]`` per rank behind one multicast - mapping, every word ``0x80000000`` (empty), and ``flags`` (int32 ``[4]``: the consumer's call count, whose parity - selects the half, and its CTA arrivals). Separate from the MNNVL all-reduce workspace. Pass ``uc`` and ``flags`` - as ``trtllm::k3_latent_reduce``'s ``lat_uc`` and ``lat_flags``; :meth:`push_args` gives the producers' arguments.""" + mapping, every word ``0x80000000`` (empty), and ``flags`` (int32 ``[4]``: [0] the consumer's call count, whose + parity selects the half, [2] the CTAs of its running call that have counted in). Separate from the MNNVL all-reduce + workspace. Pass ``uc`` and ``flags`` as ``trtllm::k3_latent_reduce``'s ``lat_uc`` and ``lat_flags``; + :meth:`push_args` gives the producers' arguments.""" uc: torch.Tensor """int32 [2 * 8 * world * 1792]: this rank's words.""" mc: torch.Tensor """The same words through the multicast mapping (where the producers push).""" flags: torch.Tensor - """int32 [4]: [0] the consumer's call count, then its CTA arrivals.""" + """int32 [4]: [0] the consumer's call count, [2] its CTA arrivals; [1] and [3] unused.""" rank: int world_size: int handle: Any diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py index 94c9ea5e759b..3fc6d6c760fe 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py @@ -213,8 +213,7 @@ def __init__( ): if torch.cuda.is_current_stream_capturing(): raise RuntimeError("K3MoeState allocates its scratch: build it outside CUDA-graph capture") - # One persistent CTA per SM (config "num_ctas" caps it, e.g. for a grid-size A/B). head_flags always explicit: - # its environment fallback must not reach the plain build. + # One persistent CTA per SM (config "num_ctas" caps it, e.g. for a grid-size A/B). num_ctas = torch.cuda.get_device_properties(device).multi_processor_count cfg = { "i_tp": i_tp, @@ -427,6 +426,8 @@ class K3MoeWideState: may read once that grid has completed).""" def __init__(self, device: torch.device, i_tp: int, num_local: int, use_pdl: bool = True): + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError("K3MoeWideState allocates its scratch: build it outside CUDA-graph capture") num_ctas = torch.cuda.get_device_properties(device).multi_processor_count config = { "i_tp": i_tp, @@ -506,6 +507,8 @@ def __init__( w2_weight: torch.Tensor, w2_weight_scale: torch.Tensor, ): + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError("K3MoeWideLayer allocates its counters: build it outside CUDA-graph capture") ok, why = is_supported( w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale, state.num_local ) From bdad277f8e51965a35e88ed67a7d40691c9397ce Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 19:26:32 -0700 Subject: [PATCH 060/161] [None][feat] modeling_v2 catalog: Kimi K3 collective and MoE entries Eight catalog entries for the Kimi K3 decode kernels and MNNVL collectives of this PR. Each has a contract with a ## State section, a wrapper that makes one call on caller-owned state, and a GPU test that drives call sequences on real state objects: - comm/k3_sandwich_oproj, comm/k3_sandwich_tail, comm/k3_sandwich_plain: the three sandwiches over one K3SandwichWorkspace per TP group; - comm/mnnvl_fusion_allreduce (the sum, or with residual + RMSNorm; one-shot size per call) and comm/mnnvl_allgather_split, both over the TP group's MnnvlWorkspace, which they share with comm/mnnvl_allreduce_attn_res; - comm/k3_latent_reduce over a K3LatentExchange (its producers' pushes are emulated until the push builds land); - moe/k3_moe_front over a K3MoeHeadWorkspace; - moe/k3_moe: K3MoeState / K3MoeLayer for up to 8 tokens (from the router logits, or after the front), K3MoeWideState / K3MoeWideLayer for up to 64, with k3_route_quant. The collective matrices (comm/__op_matrix.py, started by the collected test_modeling_v2_* files) take --world-size and --launcher. They run single calls against native-torch references, steps of layers with the token count dipping and growing back and a random rank late, two state objects interleaved, CUDA-graph capture and replay with eager calls in between, refusals that leave the state untouched, create() under capture, and last a negative control showing what a wrong call order does. moe/k3_moe's test runs on one GPU. Signed-off-by: Vasanth Sabavat --- .../catalog/comm/k3_latent_reduce.md | 205 ++++ .../catalog/comm/k3_latent_reduce.py | 19 + .../catalog/comm/k3_sandwich_oproj.md | 186 ++++ .../catalog/comm/k3_sandwich_oproj.py | 47 + .../catalog/comm/k3_sandwich_plain.md | 159 ++++ .../catalog/comm/k3_sandwich_plain.py | 41 + .../catalog/comm/k3_sandwich_tail.md | 206 +++++ .../catalog/comm/k3_sandwich_tail.py | 61 ++ .../catalog/comm/mnnvl_allgather_split.md | 175 ++++ .../catalog/comm/mnnvl_allgather_split.py | 46 + .../catalog/comm/mnnvl_fusion_allreduce.md | 229 +++++ .../catalog/comm/mnnvl_fusion_allreduce.py | 68 ++ .../modeling_v2/catalog/moe/k3_moe.md | 281 ++++++ .../modeling_v2/catalog/moe/k3_moe.py | 90 ++ .../modeling_v2/catalog/moe/k3_moe_front.md | 235 +++++ .../modeling_v2/catalog/moe/k3_moe_front.py | 46 + .../comm/_k3_latent_reduce_op_matrix.py | 539 +++++++++++ .../comm/_k3_moe_front_op_matrix.py | 874 ++++++++++++++++++ .../modeling_v2/comm/_k3_sandwich_common.py | 548 +++++++++++ .../comm/_k3_sandwich_oproj_op_matrix.py | 214 +++++ .../comm/_k3_sandwich_plain_op_matrix.py | 177 ++++ .../comm/_k3_sandwich_tail_op_matrix.py | 209 +++++ .../comm/_mnnvl_allgather_split_op_matrix.py | 636 +++++++++++++ .../comm/_mnnvl_fusion_allreduce_op_matrix.py | 702 ++++++++++++++ ..._modeling_v2_k3_latent_reduce_op_matrix.py | 23 + ...modeling_v2_k3_sandwich_oproj_op_matrix.py | 23 + ...modeling_v2_k3_sandwich_plain_op_matrix.py | 23 + ..._modeling_v2_k3_sandwich_tail_op_matrix.py | 23 + ...ling_v2_mnnvl_allgather_split_op_matrix.py | 19 + ...ing_v2_mnnvl_fusion_allreduce_op_matrix.py | 19 + .../moe/test_modeling_v2_k3_moe.py | 870 +++++++++++++++++ ...test_modeling_v2_k3_moe_front_op_matrix.py | 32 + 32 files changed, 7025 insertions(+) create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.py create mode 100644 tests/unittest/_torch/modeling_v2/comm/_k3_latent_reduce_op_matrix.py create mode 100644 tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py create mode 100644 tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_common.py create mode 100644 tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_oproj_op_matrix.py create mode 100644 tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_plain_op_matrix.py create mode 100644 tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_tail_op_matrix.py create mode 100644 tests/unittest/_torch/modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py create mode 100644 tests/unittest/_torch/modeling_v2/comm/_mnnvl_fusion_allreduce_op_matrix.py create mode 100644 tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_latent_reduce_op_matrix.py create mode 100644 tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_sandwich_oproj_op_matrix.py create mode 100644 tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_sandwich_plain_op_matrix.py create mode 100644 tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_sandwich_tail_op_matrix.py create mode 100644 tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_allgather_split_op_matrix.py create mode 100644 tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_fusion_allreduce_op_matrix.py create mode 100644 tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe.py create mode 100644 tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_front_op_matrix.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md new file mode 100644 index 000000000000..89663c308c43 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md @@ -0,0 +1,205 @@ +--- +receipts: + sm_100: {status: pending, world_size: 4} +--- + +# k3_latent_reduce + +**Wraps** `torch.ops.trtllm.k3_latent_reduce` (one call). + +A stateful entry: its correctness depends on a caller-owned state object, `K3LatentExchange`, passed as the last +argument and described under *State*. The op is the consumer half of a collective; the producers that push into the +exchange are not part of this entry (see *Notes*). + +## Semantics + +Kimi K3's latent all-reduce at decode size (at most 8 tokens), split in two. On every rank of a TP group of `W` +ranks, the routed experts' push-only kernel (the producer) stores this rank's routed partial rows into every rank's +exchange through a multicast mapping, and exits; this op then waits until every rank's rows of the same call are in +its own copy and sums them. Every rank gets back the same rows. Per call of `M` tokens: + +``` +p_r = rank r's routed partial [M, 3584] bf16, as r's producer pushed it (-0.0 stored as +0.0) +c_k = fp32 sum of p_r over r = 8k .. min(8k + 8, W) - 1, in rank order, from +0 # chunks of 8 ranks +out = bf16(c_0 + c_1 + ..., in order, from +0) # round to nearest even; [M, 3584] +``` + +This is the order of the MNNVL one-shot all-reduce (`reduceOneshotLamport`, +`cpp/tensorrt_llm/kernels/communicationKernels/mnnvlAllreduceKernels.cu`), so `out` equals the one-shot all-reduce of +the partials bit for bit (the kernel's statement). Certified at every `M` from 1 to 8: `out` equals +`comm/mnnvl_fusion_allreduce` sent one-shot over an `MnnvlWorkspace`, and a torch reference in that order, bit for +bit, including for partials whose sum another order changes in at least a quarter of the elements. The result is +bitwise identical on every rank (certified for every call of the test outside its negative controls). + +Fusion boundary. Inside: waiting for every rank's rows of this call, the sum, the output, emptying the words it read +and advancing the exchange's call count. Outside: computing the routed partial and pushing it (the producer), the +shared experts, and whatever consumes `out` (the routed latent's up-projection). Steps of more than 8 tokens need +another all-reduce of the partial. + +## Signature + +```python +def k3_latent_reduce(num_tokens: int, exchange: K3LatentExchange) -> torch.Tensor +``` + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `num_tokens` | `M`, 1-8 (all certified) | Python int | — | host | +| `exchange` | a `K3LatentExchange` of this rank's TP group, `W` = 4 certified (see *State*) | — | — | — | +| returns | `[M, 3584]` | bf16 | contiguous, newly allocated | = `exchange.uc.device` | + +The summands are not arguments: they are rows `0..M-1` of every rank's slot in this call's half of the exchange, +written by the producers. The op runs on the current stream of `exchange.uc`'s device. Inert (not exposed by the +wrapper): the op's `ctas_per_token` = 0, i.e. 4 CTAs per token row up to 8 ranks and 14 at 16. The op also takes 4, +14 or 28, which its own test (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_latent_reduce.py`) runs; only +the default is certified here. + +## State + +**Object.** `K3LatentExchange`, defined beside the buffer layout it allocates +(`cute_dsl_kernels/k3_fused_moe/latent_op.py`) and re-exported by this entry's wrapper; one per TP group, owned by +the caller. A state type, not an entry: it launches nothing per call. + +**Contents and size.** + +- `uc`: int32 `[2 x 8 x W x 1792]`, this rank's buffer: two halves of 8 token rows of `W` rank slots of 3584 bf16, + stored in pairs as int32 words (the even column in the low 16 bits). The word `0x80000000` is empty (not written). + Every word is empty after `create` and again after every call (certified). +- `mc`: the same words through the multicast mapping; a store through `mc` lands in every rank's `uc`. +- `flags`: int32 `[4]`, this rank's own: `[0]` the op's call count, whose parity is the half the next push and + reduce use; `[2]` the arrivals of the running reduce's CTAs, 0 between calls; `[1]` and `[3]` unused, 0. After + `create` all four are 0, and after every call they read `[count, 0, 0, 0]` (certified). +- `handle`, the `McastGPUBuffer` that owns the memory (the exchange is valid while it lives); `comm`, the TP-group + communicator the handles were exchanged over; `rank`, `world_size`. + +The views are `2 x 8 x W x 1792 x 4` bytes per rank (certified): 448 KiB at `W` = 4, 896 KiB at 8, 1.75 MiB at 16; +the `McastGPUBuffer` rounds the allocation up to its multicast granularity. The size depends on `W` only: every call +of up to 8 tokens fits. + +**Who creates it, and when.** The target, in `post_load_weights`, with +`K3LatentExchange.create(mapping, fabric_handle=None)`: + +- collective over `mapping`'s TP group: every rank calls it at the same point; it returns on every rank or raises on + every rank (each rank's success is agreed before anyone proceeds); +- eager: it allocates and exchanges handles, so under CUDA-graph capture it raises `RuntimeError` before any + collective step (certified on every rank at once, and on one rank alone while its peers do not call it); +- it empties every word and zeroes `flags`, synchronizes, and returns only once every rank has done so (the success + agreement is the barrier), so no producer can push into a buffer before its rank has armed it; armed and sized on + every rank right after `create` is certified; +- `fabric_handle`: share the memory by fabric handle (required across nodes) or POSIX file descriptor; default + `mapping.is_multi_node()`. No environment variable is read. + +`create` does not check `W`: it builds an exchange for any TP size, and the op then raises `ValueError` at every +call unless `W` is 4, 8 or 16 (the op's code). + +**Which ops may share one object.** One call's producers and this op. The producers are the push-only builds of the +routed experts (`k3_moe_m1` / `k3_moe_m2` push, `trtllm::k3_fused_moe_push`, `trtllm::k3_fused_moe_front_push`; +not in this tree yet): they take `exchange.push_args()` (`uc`, `mc`, `flags`, `rank`), read `flags[0]`, store +through `mc` and write no flag. One exchange serves every MoE layer of a TP group: all the layers' calls form one +sequence. It is not the MNNVL workspace (`MnnvlWorkspace`), the sandwich workspace, or the sandwich tail's latent +exchange (`K3SandwichLatentExchange`: the same buffer layout, but its own flags); a push must be reduced from the +exchange it went into. Two exchanges keep independent counts and buffers: calls alternating between two exchanges +in an irregular pattern, so that the halves they use differ from call to call, are all correct, and each exchange +ends clean at its own count (certified, 20 calls). They are not independent orders: a reduce spins until every +rank's push of the same call has landed, and a stream runs its kernels one after the other, so ranks that issue +calls on two exchanges (or this op and another polling collective) in different relative orders on one stream +deadlock (the kernel's behaviour; not run). + +**Call-order invariant.** On one exchange every rank makes the same sequence of calls, across layers and decode +steps, eager calls and graph replays alike. A call is one push of `M` rows by the rank's producer followed by exactly +one reduce of the same `M`, before the next push; the `k`-th call has the same `M` on every rank. A push never waits +for the peers; only the reduce does. A producer reads the half from `flags[0]` after its grid-dependency wait, +and the op reads it before its own (its CTA 0 advances the count at the very end, once every CTA has read it). So +every kernel launched between a reduce and the next push on a stream must end only after its predecessor has ended: +it calls `griddepcontrol.wait`, or it launches without programmatic dependent launch. Otherwise a push could read +the count before that reduce has advanced it and write into the half the reduce is still reading (the op's +statement, `latent_op.py`). + +**What a later launch reads.** Before its grid-dependency wait the op reads `flags[0]` (its half) and counts each +CTA into `flags[2]`. After the wait it reads rows `0..M-1` of every rank's slot of that half, polling until none of +the words is empty. At its very end CTA 0 waits until all `M x CTAs` CTAs have counted in, then sets `flags[2]` back +to 0 and `flags[0]` to the count plus one. Rows `M..7` are neither summed nor emptied (certified by the second +negative control below). The next producer reads the advanced `flags[0]`. The count is int32 and only its parity is +read: it passes from 2^31 - 1 to -2^31 without consequence (certified: four calls across the wrap from a preset +count). + +**How it is re-armed.** By the op: each thread empties every word it read (its 16-byte vector of every rank's slot) +after storing its output. Nobody pushes into that half again before every rank's reduce of it has ended: the next +push into it is two calls later, and on each rank it follows that rank's reduce of the call in between, which waits +for every rank's push of that call, each launched after that rank's reduce of this call (the kernel's statement; it +relies on the condition on programmatic dependent launch above). So no separate clear and no record of an earlier +call's size is needed, as long as every push is reduced with its own `M`: rows a reduce does not sum stay as they +are (*What a wrong order does*). Certified by the sequences below: after every call of the single-call check and +after every decode step, both halves are empty on every rank. + +**Why the test drives call sequences.** A re-arm that depends on the call's size passes every single-call test: it +fails only after a smaller call, when an older, larger call's words are still in the buffer and a later larger call +reads them as fresh rows. This entry's test therefore runs 11 decode steps of 12 layers at `M` = 8, 8, 8, 2, 7, 8, 1, +1, 8, 3, 8, each step's calls queued without host synchronization, a random rank 5 ms late before every call (its +push lands while the others' reduces poll), every call against the reference. + +**What a wrong order does.** Two negative controls at `W` = 4, both certified: + +- Swapped calls. Rank 0 makes two same-shaped calls in swapped order (it pushes its partial of the second call + first). Nothing raises and nothing hangs, since the counts still agree, but every rank's two results are wrong: + each reduce sums rank 0's partial of the other call, and more than half of the elements differ on every rank. The + exchange is clean afterwards and a plain call right after is correct. +- A token count that differs from the push. Every rank pushes 8 rows and rank 0 reduces 4: its 4 rows are right, + but rows 4-7 of that half stay full in its buffer. Two calls later every rank pushes 4 rows into that half and rank + 0 reduces 8: it does not wait for rows 4-7, which are already full, and returns the older call's sums for them, + bit for bit. Nothing raises or hangs; that reduce empties all 8 rows, the exchange is clean again and the next call + is correct. + +Not run, because each waits forever (the op polls until no word it reads is empty): a reduce of rows nobody pushed +(with nothing stale there), a push into the other half, and a rank making one call more or fewer than its peers +(from then on its count's parity differs from theirs). + +## Metadata consumed + +Besides `exchange` (an explicit argument), one process-wide cache inside the op: compiled kernels keyed by (`W`, +CTAs per token, PDL). `M` is a runtime argument, so one compile serves every `M` (the op's code). The first call for +a key compiles (seconds) and must be eager: under capture it raises `RuntimeError` ("must run once outside +CUDA-graph capture first") before launching anything (certified: the test's first call is made under capture on +every rank, and the exchange is untouched). The cache is result-neutral. `TRTLLM_ENABLE_PDL` (default `1`), read at +every call, is part of the key: it decides whether the kernel launches as a programmatic dependent, which changes +scheduling, not results (the op's statement; the test runs the default). + +## Preconditions + +- `M` 1-8 and `W` 4, 8 or 16 (the op's `supports`). `M` = 0 or 9 raises `ValueError` on every rank before the op + touches the exchange: certified, with the count unchanged, every word still empty and the next call correct. +- Every rank calls with the same `M`, the token count its producer pushed in this call; the call order is the + *State* invariant, including its condition on programmatic dependent launch. +- The producers push rows `0..M-1` of this rank's partial into slot `[rank]` of half `flags[0] & 1` of every rank's + buffer through `mc`, as bf16 pairs in int32 words with -0.0 stored as +0.0. A word `0x80000000` (-0.0 in the odd + column, +0.0 in the even one) reads as not written, so a producer that stored one would make every rank's reduce + wait forever (the kernel's code). The test's emulated pushes store exactly this, -0.0 entries included; columns + that are -0.0 on every rank sum to +0.0 (certified). +- `exchange` was created, and the kernel compiled by one eager call, before any capture. Calls may be captured: + certified with two captured steps on one exchange, 12 push + reduce pairs at `M` = 8 and at `M` = 3, replayed + alternately 4 times each with rewritten partials and a random rank late, two eager calls of other `M` between + replays, every replayed and eager call against the reference. + +## Notes + +- Certified path: 4 ranks of one GB200 tray (sm_100), one rank per GPU, POSIX-fd handles, 4 CTAs per token row. + Test: `tests/unittest/_torch/modeling_v2/comm/_k3_latent_reduce_op_matrix.py`, collected by + `test_modeling_v2_k3_latent_reduce_op_matrix.py`. The reference is native torch in the one-shot's order. Most + checks use partials that are multiples of 1/16, whose sum is exact in any order; the bit-identity check adds + partials with a +B / -B pair per element, whose sum depends on the order (summing the ranks in reverse, or at `W` > + 8 without chunks, must change at least a quarter of the elements), and normal-distributed ones. Every comparison is + bit for bit, and every output is also compared across the ranks. +- The pushes are emulated. The producers are not in this tree yet, so the test stores each rank's partial as they + do, with a tensor copy through `mc` into the call's half. Not reached by the emulation: a producer reading the half + on the device (the test takes it from its own count of the calls, checked against `flags[0]` each time it checks + the exchange clean), and a producer still running under programmatic dependent launch when the reduce starts (the + copies launch without it). The emulated pushes also fix their half when captured, so each captured step holds an + even number of calls and an even number of eager calls runs between replays; a producer reads the half at replay + time. +- World size: the matrix takes `--world-size` and `--launcher` (`mpirun` on one node, `srun` across nodes). Kimi K3 + runs the op over 16 ranks on four trays, where the kernel sums the ranks in two chunks of 8 and uses 14 CTAs per + token row; a 4-rank run reaches neither. The 16-rank receipt is pending; `W` = 8 is not run. +- The op writes `lat_uc` (it empties words) and `lat_flags`, and its schema declares both mutable; the producers' + writes through `mc` are theirs to declare. The compile cache is a module-level dict (result-neutral, above). diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.py new file mode 100644 index 000000000000..63ab33d84835 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.py @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3's latent all-reduce at decode size as the consumer of pushed partials: the sum over the TP group of the +routed partials every rank's producer stored into a caller-owned :class:`K3LatentExchange`, in the MNNVL one-shot's +order.""" + +import torch + +# Importing the op module registers trtllm::k3_latent_reduce; the state type is the one its buffers come from. +from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.latent_op import K3LatentExchange + +__all__ = ["K3LatentExchange", "k3_latent_reduce"] + + +def k3_latent_reduce(num_tokens: int, exchange: K3LatentExchange) -> torch.Tensor: + """Return the latent rows ``[num_tokens, 3584]`` bf16: the sum over ``exchange``'s TP group of the partial rows + every rank pushed into it since the previous reduce. Empties the words it read and advances ``exchange`` by one + call: every rank pushes, then reduces, the same token count on it in the same order.""" + return torch.ops.trtllm.k3_latent_reduce(exchange.uc, exchange.flags, num_tokens) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md new file mode 100644 index 000000000000..333f925cf38e --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md @@ -0,0 +1,186 @@ +--- +receipts: + sm_100: {status: pending, world_size: 4} +--- + +# k3_sandwich_oproj + +**Wraps** `torch.ops.trtllm.k3_sandwich_oproj` (one call). + +## Semantics + +Kimi K3's post-attention step for a decode batch of at most 8 tokens, in one kernel: the row-parallel attention +output projection, the TP all-reduce of its output, and the residual update. Every rank of the workspace's TP group +calls with its own attention output `core` `[M, 768]` and its slice of the output projection `o_weight` +`[7168, 768]`; every rank gets back the same two tensors. Per token, with `W` ranks: + +``` +partial_r = bf16(core_r @ o_weight_r^T) # this rank's [M, 7168] share, fp32 accumulator +updated = bf16(prefix + bf16(sum over the W ranks of partial_r)) # the sum alone when prefix is None +normed = RMSNorm(attn_res(block_residual[0], ..., block_residual[S-1], updated); output_rms_weight, output_rms_eps) +``` + +and returns `(normed, updated)`. `attn_res` is the attention-residual selection of `comm/mnnvl_allreduce_attn_res` +(candidates scored by `rmsnorm(v) . (rms_weight * res_weight)` with `rms_eps`, softmax, mix); the RMSNorm is Kimi +K3's (normalize in fp32, round to bf16, apply the weight). The kernel follows `oneshotAllreduceAttnResKernel`'s +arithmetic — ranks summed in chunks of 8 in rank order, the same statistics, summation orders and roundings — so +its outputs are bit-identical to `o_proj` followed by `trtllm::mnnvl_allreduce_attn_res` (its kernel's statement). +The result is bitwise identical on every rank (certified, every call of the test). + +Fusion boundary. Inside: the projection, the all-reduce, the residual add, the selection, the RMSNorm. Outside: the +attention that produced `core`, keeping the snapshot bank, chaining `prefix` from layer to layer, and the MoE that +consumes `normed`. The next pre-attention step (the MoE tail, its all-reduce and the next layer's residual update) +is `comm/k3_sandwich_tail`, on the same workspace. + +## Signature + +```python +def k3_sandwich_oproj( + core: torch.Tensor, + o_weight: torch.Tensor, + prefix: Optional[torch.Tensor], + block_residual: torch.Tensor, + res_weight: torch.Tensor, + rms_weight: torch.Tensor, + output_rms_weight: torch.Tensor, + rms_eps: float, + output_rms_eps: float, + workspace: K3SandwichWorkspace, +) -> Tuple[torch.Tensor, torch.Tensor] +``` + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `core` | `[M, 768]`, `M` 1-8 (TP16's per-rank shape, whatever `W`) | bf16 | contiguous | CUDA, this rank's device | +| `o_weight` | `[7168, 768]` (this rank's slice) | bf16 | contiguous | CUDA | +| `prefix` | `None`, or `[M, 7168]` | bf16 | contiguous | CUDA | +| `block_residual` | `[S, M, 7168]`, `S` 0-8 | bf16 | contiguous | CUDA | +| `res_weight`, `rms_weight`, `output_rms_weight` | `[7168]` | bf16 | contiguous | CUDA | +| `rms_eps`, `output_rms_eps` | scalar (certified at 1e-6) | Python float | — | — | +| `workspace` | a `K3SandwichWorkspace` of this rank's TP group (see *State*) | — | — | — | +| returns | `(normed, updated)`, each `[M, 7168]` | bf16 | contiguous, newly allocated | = `core.device` | + +Single calls are certified at every `M` x `S` in {0, 1, 4, 8} x with and without `prefix`; the call sequences below +take `S` 0-8. The inputs are read only. + +Inert (not exposed by the wrapper): `x_slab=None, slab_buf=0, src_slab=None, src_buf=0` — the op can also publish +`normed` into a Lamport slab the next kernel polls, and poll `core` from its producer's slab. Each slab is cross-call +state of its own (three sentinel-armed buffers the caller rotates by the call's ordinal), not part of the workspace; +see *Notes*. + +## State + +**Object.** `K3SandwichWorkspace`, defined beside the buffer layout it allocates +(`tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py`) and re-exported by this entry's wrapper; one per TP +group, owned by the caller. + +**Contents and size.** The all-reduce buffer: two alternating halves of `[8 tokens][W ranks][7168]` bf16 per rank, +as int32 words (`0x80000000` = empty; a push writes bf16 `-0.0` as `+0.0`, so a pushed word never reads as empty), +in one multicast allocation (`uc`: this rank's words; `mc`: the same words through the multicast mapping, where the +peers push) — `2 x 8 x W x 7168 x 2` bytes, 0.92 MB at `W` = 4 and 3.67 MB at `W` = 16; `flags`, int32 `[64]`, the +call count of each of the kernel's 56 CTAs (the other 8 words unused); `rank` and `world_size`; the `McastGPUBuffer` +handle that owns the memory (the workspace is valid while the object lives); and the communicator the handles were +exchanged over. The size depends on `W` only, not on `M`: every call fits. Certified after `create`: these sizes, +every word empty, every counter zero. + +**Who creates it, and when.** The target, in `post_load_weights`, with +`K3SandwichWorkspace.create(mapping, fabric_handle=None)`: + +- collective over `mapping`'s TP group: every rank calls it at the same point; it returns on every rank or raises on + every rank (each rank's success is agreed before any returns, which is also the barrier that keeps a peer from + pushing into a buffer its owner has not emptied yet); +- eager: it allocates and exchanges handles, so under CUDA-graph capture it raises `RuntimeError` before it enters + any collective — certified on every rank at once, and on one rank alone while its peers do not call it; +- every word emptied and every counter zeroed before any rank returns; +- `fabric_handle`: share the memory by fabric handle (required across nodes) or POSIX file descriptor; default + `mapping.is_multi_node()`. No environment variable is read. + +**Which ops may share one object.** The three sandwich entries take the same object — `comm/k3_sandwich_oproj`, +`comm/k3_sandwich_tail` and the drafter's `comm/k3_sandwich_plain` (their wrappers export the one type) — and in the +model they run on one workspace per TP group, target layers and drafter layers alike, so they form one call +sequence. Certified: decode steps of three target layers (this op, then `k3_sandwich_tail`, one of the tails storing +`updated` into a bank row) with two drafter layers (`k3_sandwich_plain`, then its SwiGLU form) interleaved, at +(target `M`, drafter `M`) = (8, 3), (2, 8), (7, 7), a random rank late at every call, every call against its +reference; then that step captured at (8, 4) and replayed 6 times with rewritten inputs, one or two eager calls of +the three ops at other token counts between replays. The sandwich buffer is not the MNNVL workspace: it keeps its own +counters, and a sandwich call does not advance an `MnnvlWorkspace`. Two objects keep independent counters: calls +alternating between two workspaces in an irregular pattern, so that their counters differ, are all correct +(certified, 20 calls). + +**Call-order invariant.** Every rank of the group makes the same sequence of sandwich calls on one workspace — the +same number of calls, the `k`-th with the same op and `M` — across layers and decode steps, eager calls and graph +replays alike; and on one stream the same order of calls across objects: each call spins until its peers' rows of +the same call arrive, so two ranks issuing calls on two objects in different orders on one stream deadlock (measured +for `mnnvl_allreduce_attn_res`, `runs/drafter/u4-mnnvl-srun-2`; this kernel waits the same way). On each rank the +calls on one workspace run one after another: a call reads the counters after its grid-dependency wait, which covers +the previous call because every kernel between two calls waits for its predecessor (the kernel's statement) — a +kernel launched under PDL that skips that wait must not sit between two calls. + +**What a later launch reads.** `flags`: each CTA reads its counter after its grid-dependency wait and stores it plus +one after its push, so every call adds one to every CTA's counter (certified: all 56 advanced by exactly one per +call, the 8 spare words untouched), and the counter's parity before the call selects the half this call pushes into +and polls. And that half's words, which must be empty except for this call's pushes. + +**How it is re-armed.** By its readers: a thread empties every word it read right after reading it. The next push +into that half comes from a peer's call after next, which starts only after this call has ended on this rank (the +kernel's statement), so no separate clear and no record of the previous call's size is needed — the +`k3_spec_accept` failure (a re-arm sized by the current call) cannot occur. The test still drives the sequence that +exposed it (below). + +**Why the test drives call sequences.** See `mnnvl_allreduce_attn_res.md` (*State*): Phase 0's `k3_spec_accept` +re-armed its Lamport buffer for the current call's rows only, every single-call test passed, and a call sequence +whose row count dipped and grew back caught it; in serving it hung the ranks. This entry's test runs 11 decode steps +of 12 chained layers at `M` = 8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8, a random rank 5 ms late at every call, each call +against the reference. + +**What a wrong order does.** Certified by the test's negative control: rank 0 issues two same-shaped calls on one +workspace in swapped order. Nothing raises and nothing hangs — the counters still agree — but every rank's two +results are wrong (each call paired with the peers' call at the same position; more than half the `updated` elements +differ on every rank). A plain call right after is correct again. A rank making one call more or fewer than its peers +was not exercised: its later calls would no longer meet their peers' calls of the same position. + +## Metadata consumed + +Besides `workspace` (an explicit argument), one process-wide cache inside the op module: compiled kernels keyed by +(`W`, publish, input source, PDL), `W` read off the workspace's size. The first call for a key compiles (seconds) and +must be made eagerly — under capture the op raises "must run once per configuration outside CUDA-graph capture +first". The cache is result-neutral. PDL (`TRTLLM_ENABLE_PDL`, read at every call, default on) is part of the key; +it changes scheduling, not results. + +## Preconditions + +- bf16, contiguous; `core` `[M, 768]` with `M` 1-8, `o_weight` `[7168, 768]`, at most 8 snapshots: the op's + `supports`. Anything else (e.g. `M` = 9) raises `ValueError` on every rank before any launch: certified (every + counter unchanged), and the next call is correct. The op does not check the other tensors: `prefix`, + `block_residual` and the weights must be bf16 `[M, 7168]`, `[S, M, 7168]` and `[7168]` (certified contiguous; the + op flattens them with `reshape`, which copies a non-contiguous view such as a slice of the snapshot bank). +- Every rank calls with the same `M`, `S` and `prefix` presence, and the same `prefix`, `block_residual` and weights + (the replicated residual stream) for the same result on every rank; the call order is the *State* invariant. +- `workspace` was created, and the kernel compiled (one eager call per key), before any capture. Calls may be + captured: certified with a captured step of 12 chained calls at `M` = 8 replayed 8 times with rewritten inputs, an + eager call of another `M` on the same workspace between replays, every replayed and eager call against the + reference; and with the shared step above. +- Under PDL the kernel launches its dependents once every CTA has pushed, before its outputs are written (the + kernel's statement): a kernel launched after it reads `normed` / `updated` only after its own grid-dependency wait. + +## Notes + +- Certified path: 4 ranks of one GB200 tray (sm_100), one rank per GPU, POSIX-fd handles, Kimi K3 TP16 per-rank + shapes. Test: `tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_oproj_op_matrix.py`, with the helpers of the + three sandwich entries in `_k3_sandwich_common.py` beside it. The reference is native torch: `core` and `o_weight` + are small multiples of 1/8 and 1/16, so every partial sum is exact in fp32 and `updated` is compared bit for bit; + `normed` against an fp32 reference within 2e-2 of its largest magnitude. +- A2: the matrix takes `--world-size` and `--launcher` (`mpirun` on one tray, `srun` across trays). The kernel sums + ranks in chunks of 8, so a run at `W` <= 8 exercises one chunk; Kimi K3 runs `W` = 16 over four trays. This + entry's 16-rank receipt is pending. The op's kernel test + (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py`), in its B7 form — this op bit for bit + against `o_proj` and the MNNVL one-shot — passed 180 / 180 at 16 ranks + (`runs/session-7594608/pre/test_k3_sandwich.log`): the kernel's record, not this entry's receipt. +- A1: `mutates_args` names `ws_uc`, `ws_mc` and `ws_flags` — every call pushes through `ws_mc` into every rank's + `ws_uc`, empties the words it read in `ws_uc` and advances `ws_flags` — and `x_slab`. The op module keeps no + workspace registry: the caller passes the object it created. The compile cache is the documented process-wide cache + above. The published / polled slabs (`x_slab`, `src_slab`) are cross-call state without a state object, so they + stay inert here: proposed, a `K3Slab` state type (three buffers, sentinel-armed) whose rotation index the object + owns instead of the caller's `slab_buf` ordinal, certified in its own entry. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.py new file mode 100644 index 000000000000..628b3e774a9f --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.py @@ -0,0 +1,47 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3's post-attention sandwich: the row-parallel attention output projection, its TP all-reduce and the +residual update (attention-residual selection + RMSNorm) in one kernel, over a caller-owned +:class:`K3SandwichWorkspace`.""" + +from typing import Optional, Tuple + +import torch + +# Importing the op module registers trtllm::k3_sandwich_*; the state type is the one the ops' buffers come from. +from tensorrt_llm._torch.cute_dsl_kernels.k3_sandwich.op import K3SandwichWorkspace + +__all__ = ["K3SandwichWorkspace", "k3_sandwich_oproj"] + + +def k3_sandwich_oproj( + core: torch.Tensor, + o_weight: torch.Tensor, + prefix: Optional[torch.Tensor], + block_residual: torch.Tensor, + res_weight: torch.Tensor, + rms_weight: torch.Tensor, + output_rms_weight: torch.Tensor, + rms_eps: float, + output_rms_eps: float, + workspace: K3SandwichWorkspace, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Return ``(normed, updated)``: ``updated = prefix + allreduce(core @ o_weight^T)`` over ``workspace``'s TP group + (the sum alone without ``prefix``), ``normed`` = RMSNorm(attn_res(block_residual..., updated)). Advances + ``workspace`` by one call: every rank of the group makes the same calls on it in the same order.""" + normed, updated = torch.ops.trtllm.k3_sandwich_oproj( + core, + o_weight, + prefix, + block_residual, + res_weight, + rms_weight, + output_rms_weight, + rms_eps, + output_rms_eps, + workspace.uc, + workspace.mc, + workspace.flags, + workspace.rank, + ) + return normed, updated diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md new file mode 100644 index 000000000000..2cbe28ac06e3 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md @@ -0,0 +1,159 @@ +--- +receipts: + sm_100: {status: pending, world_size: 4} +--- + +# k3_sandwich_plain + +**Wraps** `torch.ops.trtllm.k3_sandwich_plain` (one call). + +## Semantics + +A row-parallel projection, its TP all-reduce, the residual add and an RMSNorm in one kernel, for a decode batch of at +most 8 tokens: Kimi K3's drafter (DSpark) layers' attention output projection, or with `swiglu` their MLP's +SiLU-and-mul and down projection. Every rank of the workspace's TP group calls with its own `x` and weight slice and +the same `residual`; every rank gets back the same two tensors. Per token, with `W` ranks and `K` the slice's width: + +``` +a_r = x_r, or with swiglu silu_and_mul(x_r) = bf16(silu(x_r[:, :K]) * x_r[:, K:]) # gate columns first +partial_r = bf16(a_r @ weight_r^T) # this rank's [M, 7168] share, fp32 accumulator +updated = bf16(residual + bf16(sum over the W ranks of partial_r)) +normed = bf16(updated * rsqrt(mean over the 7168 columns of updated^2 + eps) * norm_weight) +``` + +and returns `(normed, updated)`. The arithmetic and summation order are those of the all-reduce the call replaces, +the MNNVL one-shot with `AllReduceFusionOp.RESIDUAL_RMS_NORM` (`kARResidualRMSNorm`): ranks summed in chunks of 8 in +rank order, the residual added to the rounded sum, the mean square taken over bf16-rounded squares in that kernel's +reduction tree, `normed` rounded once — so the outputs are bit-identical to the projection followed by that +all-reduce (the kernel's statement). With `swiglu`, `silu_and_mul` follows `k3_ctm_gemv_swiglu` (fp32, the sigmoid +as the correctly rounded reciprocal of `1 + exp(-gate)`, one bf16 rounding), and the projection accumulates the even +and the odd k-tiles in two fp32 accumulators added as `(0 + even) + odd` before its rounding, split 2's order (the +kernel's statement). The op's kernel test checks the plain form against `k3_ctm_gemv` (split 1) and the SwiGLU form +against `k3_ctm_gemv_swiglu` (split 2), each followed by the MNNVL one-shot RESIDUAL_RMS_NORM all-reduce, bit for +bit. The result is bitwise identical on every rank (certified, every call of the test). + +Fusion boundary. Inside: (with `swiglu`) the SiLU-and-mul, the projection, the all-reduce, the residual add, the +RMSNorm. Outside: what produced `x` (the drafter's attention, or its gate_up projection), chaining `residual` from +call to call, and what consumes `normed`. The drafter's calls run on the workspace of the target's sandwiches +(`comm/k3_sandwich_oproj`, `comm/k3_sandwich_tail`). + +## Signature + +```python +def k3_sandwich_plain( + x: torch.Tensor, + weight: torch.Tensor, + residual: torch.Tensor, + norm_weight: torch.Tensor, + eps: float, + workspace: K3SandwichWorkspace, + swiglu: bool = False, +) -> Tuple[torch.Tensor, torch.Tensor] +``` + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `x` | `[M, 384]`, `M` 1-8; with `swiglu` `[M, 1792]` (gate first) | bf16 | contiguous | CUDA, this rank's device | +| `weight` | `[7168, 384]`; with `swiglu` `[7168, 896]` (this rank's slice) | bf16 | contiguous | CUDA | +| `residual` | `[M, 7168]` | bf16 | contiguous | CUDA | +| `norm_weight` | `[7168]` | bf16 | contiguous | CUDA | +| `eps` | scalar (certified at 1e-6) | Python float | — | — | +| `workspace` | a `K3SandwichWorkspace` of this rank's TP group (see *State*) | — | — | — | +| `swiglu` | `False` (o_proj, `K` 384) or `True` (down, `K` 896) | bool | — | — | +| returns | `(normed, updated)`, each `[M, 7168]` | bf16 | contiguous, newly allocated | = `x.device` | + +Single calls are certified at every `M` in both forms. Without `swiglu` the op also takes `K` any multiple of 128 up +to 896 (its `supports_plain`); only Kimi K3 TP16's 384 is certified. The inputs are read only. + +Inert (not exposed by the wrapper): `ipc_order=False`. With it the op follows the IPC one-shot's order instead +(`allreduce_fusion_kernel_oneshot_lamport`: fp32 squares, its own reduction tree; TP <= 8, `ValueError` above), for a +drafter whose all-reduce is the IPC one — another compiled kernel, not certified here. + +## State + +**Object.** `K3SandwichWorkspace`, the object `comm/k3_sandwich_oproj` and `comm/k3_sandwich_tail` take, defined in +`tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py` and re-exported by this entry's wrapper too; one per TP +group, owned by the caller. + +**Contents and size.** As in `k3_sandwich_oproj.md`: two alternating halves of `[8 tokens][W ranks][7168]` bf16 per +rank, as int32 words (`0x80000000` = empty), behind one multicast mapping (`uc`, `mc`) — `2 x 8 x W x 7168 x 2` +bytes, 0.92 MB at `W` = 4 and 3.67 MB at `W` = 16 — and `flags`, int32 `[64]`, the call counts of the kernel's 56 +CTAs, with the handle that owns the memory and the communicator. This op's partial rows go into the same slots as the +target's sandwiches'. Certified after `create`: sized for the group, every word empty, every counter zero. + +**Who creates it, and when.** The target, in `post_load_weights`, with +`K3SandwichWorkspace.create(mapping, fabric_handle=None)`: collective over the TP group (it returns on every rank or +raises on every rank), eager, every word emptied and every counter zeroed before any rank returns. Under CUDA-graph +capture it raises `RuntimeError` before it enters any collective: certified on every rank at once, and on one rank +alone while its peers do not call it. Details in `k3_sandwich_oproj.md`. The drafter does not create its own: it +takes the target's. + +**Which ops may share one object.** The three sandwich entries — `comm/k3_sandwich_oproj`, `comm/k3_sandwich_tail` +and this one — take the same object, and in the model the drafter's calls run on the target's workspace of their TP +group: one call sequence. That sharing is certified by `comm/k3_sandwich_oproj`'s matrix: decode steps of target +layers with two drafter layers (this op, then its SwiGLU form) interleaved at the drafter's own token count, eager +with a random rank late and captured with eager calls between replays. The two forms of this op share one sequence +too (certified here: every sequence of this entry's test alternates them). Two workspaces keep independent +counters: calls of both forms alternating between two in an irregular pattern are all correct (certified, 20 +calls). + +**Call-order invariant.** As in `k3_sandwich_oproj.md`: every rank makes the same sequence of sandwich calls on one +workspace (the same number of calls, the `k`-th with the same op, form and `M`) across layers, decode steps, eager +calls and graph replays; the same order of calls across objects on one stream; and no kernel between two calls that +skips its grid-dependency wait. + +**What a later launch reads.** `flags`: every call adds one to every CTA's counter (certified for both forms: all 56 +by exactly one per call, the spare words untouched), and the counter's parity before the call selects the half this +call pushes into and polls; and that half's words, which must be empty except for this call's pushes. + +**How it is re-armed.** By its readers, as for `k3_sandwich_oproj`: a thread empties every word it read right after +reading it, so no separate clear and no record of the previous call's size is needed (the kernel's statement). + +**Why the test drives call sequences.** See `k3_sandwich_oproj.md`. This entry's test runs 11 decode steps of 6 +drafter layers (12 calls chained through the residual, the two forms alternating) at `M` = 8, 8, 8, 2, 7, 8, 1, 1, 8, +3, 8, a random rank 5 ms late at every call, each call against the reference. + +**What a wrong order does.** Certified by the test's negative control: rank 0 issues two same-shaped calls on one +workspace in swapped order. Nothing raises and nothing hangs, but every rank's two results are wrong (more than half +the `updated` elements differ on every rank). A plain call right after is correct again. + +## Metadata consumed + +Besides `workspace`, the op module's process-wide compile cache, keyed by (`W`, `K`, order, `swiglu`, PDL): the two +forms are two kernels, each compiled (seconds) on its first call, which must be eager — under capture the op raises +"must run once per configuration outside CUDA-graph capture first". The cache is result-neutral. PDL +(`TRTLLM_ENABLE_PDL`, read at every call, default on) changes scheduling, not results. + +## Preconditions + +- bf16, contiguous; `M` 1-8; `weight` `[7168, K]` with `K` a multiple of 128 up to 896 and `x` `[M, K]`, or with + `swiglu` `K` = 896 exactly and `x` `[M, 1792]`; `residual` `[M, 7168]`, `norm_weight` `[7168]`: the op's + `supports_plain`. Anything else raises `ValueError` on every rank before any launch — certified for `M` = 9, the + SwiGLU form on a `K` 384 slice and an `x` whose `K` is not the weight's (every counter unchanged) — and the next + call is correct. +- Every rank calls with the same `M` and form, and the same `residual` and `norm_weight` (the replicated residual + stream) for the same result on every rank; the call order is the *State* invariant. +- `workspace` was created, and each kernel compiled (one eager call per key), before any capture. Calls may be + captured: certified with a captured drafter step of 6 layers (12 chained calls, both forms) at `M` = 8 replayed 8 + times with rewritten inputs, an eager call of another `M` on the same workspace between replays, every replayed + and eager call against the reference. +- Under PDL the kernel launches its dependents once every CTA has pushed, before its outputs are written (the + kernel's statement): a kernel launched after it reads the outputs only after its own grid-dependency wait. + +## Notes + +- Certified path: 4 ranks of one GB200 tray (sm_100), one rank per GPU, POSIX-fd handles, Kimi K3 TP16 per-rank + shapes. Test: `tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_plain_op_matrix.py`, helpers in + `_k3_sandwich_common.py`. The reference is native torch: `x` and the weights are small multiples of 1/8, 1/16 and + 1/64, and the SwiGLU gates are 0, 32 or 64, on which silu is exact in fp32 (silu(0) = 0; from 32 up, + `1 + exp(-gate)` rounds to 1), so `silu_and_mul(x)` is `gate * up` exactly and every partial sum is exact: `updated` + is compared bit for bit in both forms. `normed` is compared with torch's fp32 RMSNorm within 2e-2 of its largest + magnitude (the one-shot rounds the squares to bf16 and sums them in its own tree). SiLU-and-mul's rounding on + general inputs is the kernel test's to certify (bit for bit against `k3_ctm_gemv_swiglu`), not this matrix's. +- A2: as for `k3_sandwich_oproj`; this entry's 16-rank receipt is pending. The op's kernel test + (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py`), in its B7 form, passed 180 / 180 at 16 + ranks (`runs/session-7594608/pre/test_k3_sandwich.log`): the kernel's record, not this entry's receipt. +- A1: `mutates_args` names `ws_uc`, `ws_mc` and `ws_flags`, every buffer the op writes. The compile cache is the + documented process-wide cache above. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.py new file mode 100644 index 000000000000..19a6bf69aa4f --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.py @@ -0,0 +1,41 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""A row-parallel projection, its TP all-reduce, the residual add and an RMSNorm in one kernel (Kimi K3's drafter +layers: the attention output projection, or with ``swiglu`` the MLP's SiLU-and-mul and down projection), over a +caller-owned :class:`K3SandwichWorkspace`.""" + +from typing import Tuple + +import torch + +# Importing the op module registers trtllm::k3_sandwich_*; the state type is the one the ops' buffers come from. +from tensorrt_llm._torch.cute_dsl_kernels.k3_sandwich.op import K3SandwichWorkspace + +__all__ = ["K3SandwichWorkspace", "k3_sandwich_plain"] + + +def k3_sandwich_plain( + x: torch.Tensor, + weight: torch.Tensor, + residual: torch.Tensor, + norm_weight: torch.Tensor, + eps: float, + workspace: K3SandwichWorkspace, + swiglu: bool = False, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Return ``(normed, updated)``: ``updated = residual + allreduce(x @ weight^T)`` over ``workspace``'s TP group + (with ``swiglu``, of ``silu_and_mul(x) @ weight^T``), ``normed = RMSNorm(updated) * norm_weight``. Advances + ``workspace`` by one call: every rank of the group makes the same calls on it in the same order.""" + normed, updated = torch.ops.trtllm.k3_sandwich_plain( + x, + weight, + residual, + norm_weight, + eps, + workspace.uc, + workspace.mc, + workspace.flags, + workspace.rank, + swiglu=swiglu, + ) + return normed, updated diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md new file mode 100644 index 000000000000..3b31a813036f --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md @@ -0,0 +1,206 @@ +--- +receipts: + sm_100: {status: pending, world_size: 4} +--- + +# k3_sandwich_tail + +**Wraps** `torch.ops.trtllm.k3_sandwich_tail` (one call). + +## Semantics + +Kimi K3's pre-attention step for a decode batch of at most 8 tokens, in one kernel: the row-parallel MoE tail (this +rank's slice of the latent up projection of the normed latent, and its slice of the shared experts' down +projection), the TP all-reduce of its output, and the next layer's residual update. Every rank of the workspace's TP +group calls with the whole reduced latent `latent` `[M, 3584]` (the same on every rank), its slice of the +shared-expert activation `act` `[M, 384]`, its `tail_weight` `[7168, 640]` and its slice's first latent column `lo`; +every rank gets back the same two tensors. Per token, with `W` ranks: + +``` +scale = 1 / sqrt(mean over the 3584 columns of latent^2 + lat_eps) # fp32, the whole reduced latent row +partial_r = bf16(scale * (latent[:, lo_r : lo_r + 224] @ tail_weight_r[:, 0:224]^T) + + act_r @ tail_weight_r[:, 256:640]^T) # fp32 accumulators, one rounding +updated = bf16(prefix + bf16(sum over the W ranks of partial_r)) # the sum alone when prefix is None +normed = RMSNorm(attn_res(block_residual[0], ..., block_residual[S-1], updated); output_rms_weight, output_rms_eps) +``` + +and returns `(normed, updated)`: `[rmsnorm(latent)[:, lo:lo+224] | act] @ tail_weight^T` with the latent's RMS +applied to the fp32 latent accumulator rather than to the latent, then the epilogue of `comm/k3_sandwich_oproj` — +the tail's partial followed by `trtllm::mnnvl_allreduce_attn_res` (the op's statement). The latent RMSNorm carries +no weight here: a norm weight belongs folded into the latent columns of `tail_weight` (Kimi K3's model folds it). +The kernel multiplies latent columns `[lo, lo + 256)` by `tail_weight[:, 0:256]`; columns 224-255 of `tail_weight` +are zero padding, so the next slice's columns (or, past the end of the row, zeros) add nothing. The result is bitwise +identical on every rank (certified, every call of the test). + +Two optional outputs. With `tap` the kernel also stores into it the pre-norm attention-residual mixture — the bf16 +`attn_res(...)` that the RMSNorm normalizes, a DSpark capture layer's tap — or, with `tap_updated`, `updated`. With +`updated_out` it stores `updated` there instead of into a new tensor (e.g. the next row of the snapshot bank), and the +wrapper returns that tensor as `updated` (the op itself returns an empty `[0, 7168]` in its place). The options move +outputs, not values: certified, the same inputs give bit-identical `normed` and `updated` with each option and +without. + +Fusion boundary. Inside: the latent's RMS, the tail projection, the all-reduce, the residual add, the selection, the +RMSNorm, and the tap. Outside: the latent all-reduce that produced `latent`, the shared experts' gate_up and +activation that produced `act`, folding the latent norm's weight into `tail_weight`, keeping the snapshot bank, +chaining `prefix`, and the attention that consumes `normed` (its post-attention step is `comm/k3_sandwich_oproj`, on +the same workspace). + +## Signature + +```python +def k3_sandwich_tail( + latent: torch.Tensor, + act: torch.Tensor, + tail_weight: torch.Tensor, + lo: int, + lat_eps: float, + prefix: Optional[torch.Tensor], + block_residual: torch.Tensor, + res_weight: torch.Tensor, + rms_weight: torch.Tensor, + output_rms_weight: torch.Tensor, + rms_eps: float, + output_rms_eps: float, + workspace: K3SandwichWorkspace, + tap: Optional[torch.Tensor] = None, + tap_updated: bool = False, + updated_out: Optional[torch.Tensor] = None, +) -> Tuple[torch.Tensor, torch.Tensor] +``` + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `latent` | `[M, 3584]`, `M` 1-8 (below) | bf16 | contiguous, 16-byte aligned | CUDA, this rank's device | +| `act` | `[M, 384]` (this rank's slice) | bf16 | contiguous | CUDA | +| `tail_weight` | `[7168, 640]` (below) | bf16 | contiguous | CUDA | +| `lo` | `0, 224, ..., 3360` (below) | Python int | — | — | +| `lat_eps`, `rms_eps`, `output_rms_eps` | scalar (certified at 1e-6) | Python float | — | — | +| `prefix` | `None`, or `[M, 7168]` | bf16 | contiguous | CUDA | +| `block_residual` | `[S, M, 7168]`, `S` 0-8 | bf16 | contiguous | CUDA | +| `res_weight`, `rms_weight`, `output_rms_weight` | `[7168]` | bf16 | contiguous | CUDA | +| `workspace` | a `K3SandwichWorkspace` of this rank's TP group (see *State*) | — | — | — | +| `tap` | `None`, or `[M, 7168]` (below) | bf16 | unit column stride | CUDA | +| `tap_updated` | `False` (the mixture) or `True` (`updated`) | bool | — | — | +| `updated_out` | `None`, or `[M, 7168]` | bf16 | contiguous, 16-byte aligned | CUDA | +| returns | `(normed, updated)`, each `[M, 7168]` | bf16 | contiguous | = `latent.device` | + +- `latent`: the whole reduced latent row, the same on every rank (the latent all-reduce's output). +- `tail_weight`: columns 0-223 this rank's slice of the latent up projection (the latent norm's weight folded in), + 224-255 zero, 256-639 its slice of the shared experts' down projection. +- `lo`: the slice's first latent column, 224 x the slice index in Kimi K3. All 16 slices are certified (every rank + takes each of them across the single calls), including the last, whose padding columns run past the row. +- `tap`: rows a multiple of 8 elements apart, 16-byte aligned; certified as a column slice of an `[M, 5 x 7168]` + capture buffer, nothing outside the slice written. +- `updated_out`: certified as a row of a `[3, M, 7168]` bank, the other rows untouched; `updated` is then that + tensor. Otherwise both outputs are newly allocated. +- Single calls are certified at every `M` x `S` in {0, 1, 4, 8} x with and without `prefix`, and the options at every + `M`; the call sequences below take `S` 0-8. The inputs are read only. + +Inert (not exposed by the wrapper): `x_slab`, `slab_buf`, `src_slab`, `src_buf` — publishing `normed` into a +Lamport slab and polling the reduced latent from its producer's slab, cross-call state of their own as for +`comm/k3_sandwich_oproj`; and `lat_uc`, `lat_flags` — the latent all-reduce folded into this op over a second state +object, a `K3SandwichLatentExchange` (see *Notes*). + +## State + +**Object.** `K3SandwichWorkspace`, the object `comm/k3_sandwich_oproj` takes, defined in +`tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py` and re-exported by this entry's wrapper too; one per TP +group, owned by the caller. + +**Contents and size.** As in `k3_sandwich_oproj.md`: two alternating halves of `[8 tokens][W ranks][7168]` bf16 per +rank, as int32 words (`0x80000000` = empty), behind one multicast mapping (`uc`, `mc`) — `2 x 8 x W x 7168 x 2` +bytes, 0.92 MB at `W` = 4 and 3.67 MB at `W` = 16 — and `flags`, int32 `[64]`, the call counts of the kernel's 56 +CTAs, with the handle that owns the memory and the communicator. This op's partial rows go into the same slots as +`k3_sandwich_oproj`'s. Certified after `create`: sized for the group, every word empty, every counter zero. + +**Who creates it, and when.** The target, in `post_load_weights`, with +`K3SandwichWorkspace.create(mapping, fabric_handle=None)`: collective over the TP group (it returns on every rank or +raises on every rank), eager, every word emptied and every counter zeroed before any rank returns. Under CUDA-graph +capture it raises `RuntimeError` before it enters any collective: certified on every rank at once, and on one rank +alone while its peers do not call it. Details in `k3_sandwich_oproj.md`. + +**Which ops may share one object.** The three sandwich entries — `comm/k3_sandwich_oproj`, this one and +`comm/k3_sandwich_plain` — take the same object, and in the model the target's and the drafter's calls run on one +workspace per TP group, one call sequence. That sharing is certified by `comm/k3_sandwich_oproj`'s matrix: decode +steps of target layers (`k3_sandwich_oproj`, then this op) with drafter layers interleaved, eager with a random rank +late and captured with eager calls between replays. The folded latent exchange (`K3SandwichLatentExchange`) is a +separate object with its own counter. Two workspaces keep independent counters: calls of this op alternating between +two in an irregular pattern are all correct (certified, 20 calls). + +**Call-order invariant.** As in `k3_sandwich_oproj.md`: every rank makes the same sequence of sandwich calls on one +workspace (the same number of calls, the `k`-th with the same op and `M`) across layers, decode steps, eager calls +and graph replays; the same order of calls across objects on one stream; and no kernel between two calls that skips +its grid-dependency wait. + +**What a later launch reads.** `flags`: every call adds one to every CTA's counter (certified for this op: all 56 by +exactly one per call, the spare words untouched), and the counter's parity before the call selects the half this +call pushes into and polls; and that half's words, which must be empty except for this call's pushes. + +**How it is re-armed.** By its readers, as for `k3_sandwich_oproj`: a thread empties every word it read right after +reading it, so no separate clear and no record of the previous call's size is needed (the kernel's statement). + +**Why the test drives call sequences.** See `k3_sandwich_oproj.md`. This entry's test runs 11 decode steps of 12 +chained layers at `M` = 8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8, a random rank 5 ms late at every call, each call against +the reference. + +**What a wrong order does.** Certified by the test's negative control: rank 0 issues two same-shaped calls (no +prefix) on one workspace in swapped order. Nothing raises and nothing hangs, but every rank's two results are wrong, +far outside the 8e-3 tolerance below: more than half the `updated` elements are off by more than it, the largest by +over 10 times it. A plain call right after is correct again. + +## Metadata consumed + +Besides `workspace`, the op module's process-wide compile cache, keyed by (`W`, publish, input source, tap kind, +PDL): a call without a tap, one tapping the mixture and one tapping `updated` run three kernels, each compiled +(seconds) on its first call, which must be eager — under capture the op raises "must run once per configuration +outside CUDA-graph capture first". `updated_out` is not part of the key. The cache is result-neutral. PDL +(`TRTLLM_ENABLE_PDL`, read at every call, default on) changes scheduling, not results. + +## Preconditions + +- bf16; `latent` `[M, 3584]`, `act` `[M, 384]` and `tail_weight` `[7168, 640]` contiguous; `M` 1-8; at most 8 + snapshots: the op's `supports_tail`. `latent`'s rows 16-byte aligned (the kernel bulk-copies them); `tap` with unit + column stride, a row stride that is a multiple of 8 elements and a 16-byte aligned start; `updated_out` contiguous + and 16-byte aligned. Each violation raises `ValueError` on every rank before any launch — certified for `M` = 9, + a latent 2 bytes off alignment and a tap whose rows are 7172 elements apart (every counter unchanged) — and the + next call is correct. +- Not checked by the op, so the caller's: `lo` a slice start (a multiple of 224 from 0 to 3360 in Kimi K3), the + zero columns 224-255 of `tail_weight`, and `prefix`, `block_residual` and the weights being bf16 `[M, 7168]`, + `[S, M, 7168]` and `[7168]` (certified contiguous; the op flattens them with `reshape`, which copies a + non-contiguous view). +- `updated_out` is not a tensor the call reads (the op's example: the next row of the snapshot bank, which this call + does not read). +- Every rank calls with the same `M`, `S` and `prefix` presence, and the same `latent`, `prefix`, `block_residual` + and weights (the replicated residual stream); the call order is the *State* invariant. +- `workspace` was created, and each kernel compiled (one eager call per key), before any capture. Calls may be + captured: certified with a captured step of 12 chained calls at `M` = 8 — one tapping the mixture, one tapping + `updated`, one storing `updated` into a bank row — replayed 8 times with rewritten inputs, an eager call of + another `M` on the same workspace between replays, every replayed and eager call against the reference. +- Under PDL the kernel launches its dependents once every CTA has pushed, before its outputs are written (the + kernel's statement): a kernel launched after it reads the outputs only after its own grid-dependency wait. + +## Notes + +- Certified path: 4 ranks of one GB200 tray (sm_100), one rank per GPU, POSIX-fd handles, Kimi K3 TP16 per-rank + shapes. Test: `tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_tail_op_matrix.py`, helpers in + `_k3_sandwich_common.py`. The reference is native torch, in fp64 for the partials: `latent`, `act` and + `tail_weight` are small multiples of 1/8 and 1/16, so both accumulators are exact, but the RMS scale is the + kernel's fp32 rsqrt, not torch's, so a partial element can round to the neighbouring bf16. `updated` is compared + within 8e-3 of its largest magnitude (about one bf16 ulp of its largest elements); `normed` and the tapped mixture + within 2e-2 of an fp32 reference; the tapped `updated` bit for bit with the returned one; every output bitwise + across the ranks. The negative control's wrong pairing puts more than half the elements outside the 8e-3 bound, + the largest error over 10 times it. +- A2: as for `k3_sandwich_oproj`; this entry's 16-rank receipt is pending. The op's kernel test + (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py`), in its B7 form, passed 180 / 180 at 16 + ranks (`runs/session-7594608/pre/test_k3_sandwich.log`): the kernel's record, not this entry's receipt. +- A1: `mutates_args` names every buffer the op can write: `ws_uc`, `ws_mc`, `ws_flags`, `x_slab`, `lat_uc`, + `lat_flags`, `tap` and `updated_out`. The compile cache is the documented process-wide cache above. +- Not certified (an op option the wrapper does not expose): the folded latent all-reduce. With `lat_uc` / `lat_flags` + of a `K3SandwichLatentExchange` — its own collective `create(mapping, fabric_handle=None)`; two halves of + `[8 tokens][W ranks][3584]` bf16 per rank, a call count mod 6 and a latent-scale slab — `k3_moe` pushes its + routed latent partial into every rank's exchange and exits, and this op sums the ranks' partials itself, in the + one-shot's order, instead of reading a reduced `latent`, which then gives only the shape; every push-only `k3_moe` + call must be followed by exactly one such tail call on the same exchange (the op's statement). That is a second + state object with its own call-order rule, outside this entry. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.py new file mode 100644 index 000000000000..c0a0da60f889 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.py @@ -0,0 +1,61 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3's pre-attention sandwich: the row-parallel MoE tail (the latent up projection of the normed latent slice and +the shared experts' down projection), its TP all-reduce and the next layer's residual update (attention-residual +selection + RMSNorm) in one kernel, over a caller-owned :class:`K3SandwichWorkspace`.""" + +from typing import Optional, Tuple + +import torch + +# Importing the op module registers trtllm::k3_sandwich_*; the state type is the one the ops' buffers come from. +from tensorrt_llm._torch.cute_dsl_kernels.k3_sandwich.op import K3SandwichWorkspace + +__all__ = ["K3SandwichWorkspace", "k3_sandwich_tail"] + + +def k3_sandwich_tail( + latent: torch.Tensor, + act: torch.Tensor, + tail_weight: torch.Tensor, + lo: int, + lat_eps: float, + prefix: Optional[torch.Tensor], + block_residual: torch.Tensor, + res_weight: torch.Tensor, + rms_weight: torch.Tensor, + output_rms_weight: torch.Tensor, + rms_eps: float, + output_rms_eps: float, + workspace: K3SandwichWorkspace, + tap: Optional[torch.Tensor] = None, + tap_updated: bool = False, + updated_out: Optional[torch.Tensor] = None, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Return ``(normed, updated)``: ``updated = prefix + allreduce([rmsnorm(latent)[:, lo:lo+224] | act] @ + tail_weight^T)`` over ``workspace``'s TP group (the sum alone without ``prefix``), ``normed`` = + RMSNorm(attn_res(block_residual..., updated)). ``tap``: also store the pre-norm attention-residual mixture there + (``updated`` with ``tap_updated``). ``updated_out``: store ``updated`` there; it is then the returned ``updated``. + Advances ``workspace`` by one call: every rank of the group makes the same calls on it in the same order.""" + normed, updated = torch.ops.trtllm.k3_sandwich_tail( + latent, + act, + tail_weight, + lo, + lat_eps, + prefix, + block_residual, + res_weight, + rms_weight, + output_rms_weight, + rms_eps, + output_rms_eps, + workspace.uc, + workspace.mc, + workspace.flags, + workspace.rank, + tap=tap, + tap_updated=tap_updated, + updated_out=updated_out, + ) + return normed, (updated if updated_out is None else updated_out) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md new file mode 100644 index 000000000000..d3efd95a22ae --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md @@ -0,0 +1,175 @@ +--- +receipts: + sm_100: {status: pending, world_size: 4} +--- + +# mnnvl_allgather_split + +**Wraps** `torch.ops.trtllm.mnnvl_allgather_split` (one call). + +## Semantics + +A one-shot all-gather over a TP group's caller-owned `MnnvlWorkspace` of fp32 rows whose leading columns travel, and +are gathered, as bf16. Every rank of the workspace's group calls with its own rows `input` `[T, B + F]` (fp32) and the +same `bf16_columns` = `B`; every rank gets back the same two tensors, the ranks' slices in rank order. With `W` ranks: + +``` +bf16_out[t, r * B + j] = bf16(input_r[t, j]) j < B # round to nearest, ties to even +fp32_out[t, r * F + j] = input_r[t, B + j] j < F # the fp32 word, unchanged +``` + +and `-0.0` arrives as `+0.0` in both outputs: the Lamport buffers' empty word is the fp32 `-0.0`, so the kernel +replaces a bf16 `-0.0` half and an fp32 `-0.0` word by `+0.0` before sending. Nothing else is computed. Certified bit +for bit against torch (`input[:, :B].bfloat16()` and the fp32 columns copied, in rank order, `-0.0` made `+0.0`) on +rows that exercise each rule: normal values over exponents `2^-16` to `2^16` in the bf16 columns (inexact in bf16, so +rounded), exact midpoints between two bf16 values in every fourth of them (ties to even decide), `-0.0` in every +fourth column of each part, and in the fp32 part infinities, a NaN, signed denormals and the largest float, which +arrive unchanged. The result is bitwise the same on every rank (certified, every call of the test). + +Kimi K3's use (B7 82a110a92a, its code): the row-sharded MoE head of a wide decode step (9 to 64 tokens), and of a +decode step of at most 8 tokens where the fused MoE front kernel does not run. Rank `r`'s GEMV gives fp32 +`[T, 3584/W + 896/W]`: the latent down projection's columns `[r * 3584/W, (r+1) * 3584/W)`, then the router logits of +experts `[r * 896/W, (r+1) * 896/W)`. This call assembles the bf16 latent `[T, 3584]` (`bf16_out`) and the fp32 +logits `[T, 896]` (`fp32_out`) on every rank in one exchange, on the workspace of the routed experts' all-reduce. + +Fusion boundary. Inside: the bf16 rounding of the leading columns and the exchange. Outside: the GEMV that produced +`input`, and the routing and quantization that consume the outputs (`trtllm::k3_route_quant`; +`k3_fused_moe/k3_route_quant_ag.py` fuses this exchange with them and states it is bit for bit this op followed by +`k3_route_quant`). + +The kernel releases its programmatic dependents as soon as it starts; its outputs are complete only when its grid is, +so a kernel launched as its programmatic dependent must wait for the grid before reading them (the kernel's +statement). + +## Signature + +```python +def mnnvl_allgather_split( + input: torch.Tensor, + bf16_columns: int, + workspace: MnnvlWorkspace, +) -> Tuple[torch.Tensor, torch.Tensor] + +def required_buffer_bytes(num_tokens: int, bf16_columns: int, fp32_columns: int, world_size: int) -> int +``` + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `input` | `[T, B + F]`; `(B, F)` = (1792, 448), (896, 224), (448, 112), (224, 56) — Kimi K3's per-rank head at `W` = 2, 4, 8, 16, each certified at the world size of the run — and (8, 4); `T` 1-8, 16, 32, 64 | fp32 | contiguous, 16-byte aligned | CUDA, this rank's device | +| `bf16_columns` | `B`, a multiple of 8, with `F` = `input.shape[1] - B` a multiple of 4 | Python int | — | — | +| `workspace` | an `MnnvlWorkspace` of this rank's TP group whose buffers hold the call (see *State*) | — | — | — | +| returns | `(bf16_out, fp32_out)`: `[T, W x B]` and `[T, W x F]` | bf16, fp32 | contiguous, newly allocated | = `input.device` | + +`input` is read only. `required_buffer_bytes` = `T x W x (2B + 4F)`, the bytes the call writes into one Lamport +buffer (certified equal to what the call records, every call of the split grid). + +## State + +**Object.** `MnnvlWorkspace` (`catalog/comm/mnnvl_workspace.py`), one per TP group, owned by the caller; the object's +own contract is the *State* section of `mnnvl_allreduce_attn_res.md`. This section states what this op does with it. + +**Contents and size.** Three Lamport buffers of `buffer_bytes` each behind one multicast mapping (every word `-0.0` +when armed) and the flag words `buffer_flags` (uint32 `[9]`: current buffer, dirty buffer, bytes per buffer, dirty +stage count, bytes to clear x 4, arrival count). A call writes, from the start of one buffer, one slot per token and +rank: `B / 8` vectors of 8 bf16 values, then `F / 4` vectors of 4 fp32 words, `T x W x (2B + 4F)` bytes in all — +688128 bytes at Kimi K3's split for 64 tokens, at any `W`. A call that needs more than `buffer_bytes` is refused by +the wrapper with `ValueError` on every rank before it touches the workspace (certified at the first `T` over one +buffer; the flags do not move and the next call is correct; the op has the same check of its own, raising +`RuntimeError`). The workspace is never grown. The test's buffer is the largest of its calls, 1.75 MiB at `W` = 4 (a +two-shot `[64, 7168]` all-reduce of its shared-workspace sequence). + +**Who creates it, and when.** The target, in `post_load_weights`, with +`MnnvlWorkspace.create(mapping, buffer_bytes, fabric_handle=None)`: collective over the TP group, eager, every word +and flag armed before any rank returns (see `mnnvl_allreduce_attn_res.md`). It refuses CUDA-graph capture: certified +with every rank capturing, each raising `RuntimeError`. The check runs before any communication, so a rank that is +not capturing while its peers are would go on into the communicator split and wait for them (code). + +**Which ops may share one object.** Every MNNVL op of the group takes the same `comm_buffer` / `buffer_flags`: +`comm/mnnvl_allreduce_attn_res`, `comm/mnnvl_fusion_allreduce` on either path, and this entry. Their calls form one +sequence and each takes exactly one turn of the one rotation, whatever the op (certified: after every eager call of +the test the flags equal this test's model of one turn per call). Certified on one workspace in Kimi K3's order, a +random rank late at every call: decode steps of [attention-residual all-reduce, head all-gather, latent all-reduce +`[T, 3584]`, fused all-reduce `[T, 7168]`] and wide steps of [all-reduce `[T, 7168]`, head all-gather, latent +all-reduce] at Kimi K3's one-shot ceilings, `T` = 8, 2, 16, 64, 1, 7, 32, 8, 3, 16, two layers each, so all-gathers +follow one-shot and two-shot all-reduces. Two objects are two independent rotations: 20 all-gathers of mixed token +counts and splits alternating irregularly between two workspaces are all correct, and each workspace's flags move +with its own calls only (certified). They are not independent orders: on one stream every rank must issue its +collectives in the same order, whatever object each belongs to (measured for `mnnvl_allreduce_attn_res`: ranks +issuing B-then-A against A-then-B deadlocked, `runs/drafter/u4-mnnvl-srun-2`; this op waits for its peers the same +way). + +**Call-order invariant.** Every rank of the group makes the same sequence of calls on one workspace — the same +number, the `k`-th with the same op, `T`, `B` and `F` — across layers and decode steps, eager calls and graph replays +alike; and on one stream the same order of calls across workspaces. + +**What a later launch reads.** `buffer_flags`, which every call leaves as: current = its own buffer plus one, mod 3; +dirty = its own buffer; bytes per buffer unchanged; dirty stage count 1; bytes to clear `(T x W x (2B + 4F), 0, 0, +0)`; arrival count 0 (certified after every eager call of the test). The next call, of any MNNVL op, takes the current +buffer, whose words must all be `-0.0` but for its own pushes, and clears the dirty one by that size. The kernel waits +for the previous kernel on the stream before it reads the flags (its code). + +**How it is re-armed.** Each call clears the previous call's buffer stage by stage, by the bytes the previous call +recorded and in the previous call's stage layout (`cpp/tensorrt_llm/common/lamportUtils.cuh`, +`LamportFlags::clearDirtyLamportBuf`): after a two-shot all-reduce both of its stages, after any other call the first. +Certified by the split grid (55 calls of different sizes back to back on one workspace) and by the sequences below. + +**Why the test drives call sequences.** See `mnnvl_allreduce_attn_res.md` (*State*): Phase 0's `k3_spec_accept` +re-armed its Lamport buffer for the current call's rows only; every single-call test passed, and a sequence whose row +count dipped and grew back caught it. This entry's test runs 16 decode steps of 6 layers, one head all-gather per +layer at Kimi K3's split, at `T` = 8, 8, 8, 2, 7, 8, 1, 1, 64, 3, 32, 8, 16, 1, 64, 8, a random rank 5 ms late at +every call, each call against the reference and its flags against the model. + +**What a wrong order does.** Certified (the test's negative control): rank 0 issues two same-shaped all-gathers on +one workspace in swapped order. Nothing raises and nothing hangs — the two calls write and wait for the same words of +the same buffers — but every rank's two results are wrong in rank 0's columns, which hold rank 0's rows of the other +call, while every other rank's columns are right: the `k`-th call on every rank gathers what every rank sent at +position `k` (bit for bit). A plain call right after is correct again: a swapped pair realigns the positions. Two +calls that write different words (another `T`, `B` or `F`) cannot pair like that: a rank would wait for words its +peers do not write at that position, and hang (the protocol, not exercised). A rank making one call more or fewer +than its peers was not exercised; its positions never realign. + +## Metadata consumed + +None besides `workspace`, an explicit argument. The op keeps no cache and compiles nothing (a precompiled kernel). It +finds the multicast mapping by looking `comm_buffer`'s address up in a process registry of multicast buffers, which +the workspace's handle keeps registered. `TRTLLM_ENABLE_PDL` (read once per process, default on at SM 90 and newer) +launches the kernel as a programmatic dependent (see *Semantics* for what that asks of its consumers); results do not +depend on it. + +## Preconditions + +- `input` fp32, contiguous, 2-D, 16-byte aligned; `B` a multiple of 8 and `F` a multiple of 4; `T` at least 1. + Otherwise the op raises `RuntimeError` on every rank before it touches the workspace (certified: `B` = 12, `F` = 2, + a bf16 input, `T` = 0; the flags do not move and the next call is correct). +- `required_buffer_bytes(T, B, F, W) <= workspace.buffer_bytes` (*State*). +- Every rank calls with the same `T`, `B` and `F`; the call order is the *State* invariant. +- `workspace` was created before any capture. Calls may be captured: certified with a captured step of five calls on + one workspace — the head all-gather at `T` = 8, the latent all-reduce one-shot, the all-gather again, a `[32, 3584]` + all-reduce sent two-shot and the all-gather at `T` = 32 — replayed 8 times with rewritten rows and an eager + all-gather of another `T` or split, or an all-reduce, on the same workspace between replays, every replayed and + eager result against the reference and the flags after each. +- The kernel takes any `W`; the all-reduces sharing the workspace take `W` in {2, 4, 8, 16, 32, 64}. + +## Notes + +- Certified path: 4 ranks of one GB200 tray (sm_100), one rank per GPU, POSIX-fd handles, PDL on (the default). + Test: `tests/unittest/_torch/modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py`. The reference is native torch + and exact; this op's outputs are compared bit for bit. The all-reduce and attention-residual calls of its sequences + are checked as in their own matrices (sums bit for bit, normed outputs within a tolerance). +- Design choices this entry follows (U4U5_PLAN §1): A1, "a typed state object per stateful op ..., built by an + explicit, collective, eager `create()` ... Tests drive real-state call sequences (layers x steps, capture + replay, + two objects interleaved) plus a negative control. `mutates_args` names every written buffer" (the op falls short of + the last, see the gaps below); A2, the matrix takes `--world-size` and `--launcher` (`mpirun` on one node, `srun` + across trays), CI runs it at 4 ranks on one GB200 tray; A4, one `MnnvlWorkspace` shared by every MNNVL entry of the + TP group. The 16-rank receipt is pending. +- Not exercised: `B` = 0 or `F` = 0 (the op accepts both), denormal and non-finite values in the bf16 columns, an + accepted call of more than 64 tokens. +- Gaps against A1 (the op is unchanged by this entry): the schema marks `comm_buffer` mutable `(a!)` but not + `buffer_flags`, which every call advances (A1 item 4); the op has no `register_fake` + (`tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py` registers one for the other two MNNVL ops), so fake-tensor + tracing, e.g. `torch.compile`, cannot run it. +- In the model today the call is `MNNVLAllReduce.allgather_split(input, bf16_columns)` on `MNNVLAllReduce`'s workspace + (a dict keyed by `Mapping`, grown to the call's footprint by the first eager call that needs more). This entry takes + the explicit object instead (A4), sized at construction. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.py new file mode 100644 index 000000000000..6a4ac034f113 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.py @@ -0,0 +1,46 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""One-shot all-gather over a caller-owned :class:`MnnvlWorkspace` of fp32 rows whose leading columns travel, and are +gathered, as bf16 (Kimi K3's sharded MoE head: the latent columns in bf16, the router logits in fp32).""" + +from typing import Tuple + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + +from .mnnvl_workspace import MnnvlWorkspace + +__all__ = ["MnnvlWorkspace", "mnnvl_allgather_split", "required_buffer_bytes"] + + +def required_buffer_bytes(num_tokens: int, bf16_columns: int, fp32_columns: int, world_size: int) -> int: + """Bytes of one Lamport buffer a call occupies: every rank's rows, the bf16 columns at 2 bytes and the fp32 columns + at 4.""" + return num_tokens * world_size * (bf16_columns * 2 + fp32_columns * 4) + + +def mnnvl_allgather_split( + input: torch.Tensor, + bf16_columns: int, + workspace: MnnvlWorkspace, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Gather ``input`` (fp32 ``[num_tokens, columns]``, this rank's slice) from every rank of ``workspace``'s TP group. + Returns ``(bf16_out, fp32_out)``: ``bf16_out[t, r * bf16_columns + j] = bf16(input_r[t, j])`` and ``fp32_out[t, r + * fp32_columns + j] = input_r[t, bf16_columns + j]``, ``fp32_columns = columns - bf16_columns``, in rank order. + Takes one turn of the workspace's Lamport rotation, as an all-reduce on it does: every rank of the group makes the + same MNNVL calls on it in the same order.""" + num_tokens, columns = input.shape + need = required_buffer_bytes(num_tokens, bf16_columns, columns - bf16_columns, workspace.world_size) + if need > workspace.buffer_bytes: + raise ValueError( + f"mnnvl_allgather_split: the call needs {need} bytes per Lamport buffer, the workspace has " + f"{workspace.buffer_bytes}" + ) + bf16_out, fp32_out = torch.ops.trtllm.mnnvl_allgather_split( + input, + bf16_columns, + workspace.comm_buffer(torch.bfloat16), + workspace.buffer_flags, + ) + return bf16_out, fp32_out diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md new file mode 100644 index 000000000000..55486ef6d808 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md @@ -0,0 +1,229 @@ +--- +receipts: + sm_100: {status: pending, world_size: 4} +--- + +# mnnvl_fusion_allreduce + +**Wraps** `torch.ops.trtllm.mnnvl_fusion_allreduce` (one call). + +## Semantics + +The MNNVL all-reduce of a TP group over the group's caller-owned `MnnvlWorkspace`, optionally followed by a residual +add and an RMSNorm. Every rank of the workspace's group calls with its own rows `input` `[T, H]`; every rank gets back +the same result. With `W` ranks, per token: + +``` +s = sum over the ranks r = 0 .. W-1 of input_r # fp32, a fixed rank order, then one rounding to bf16 +plain: returns s # residual, norm_weight, eps all None +fused: updated = bf16(s + residual) # AllReduceFusionOp.RESIDUAL_RMS_NORM + rcp = rsqrt(sum_h bf16(updated_h * updated_h) / H + eps) # fp32; each square rounded to bf16 + normed = bf16(updated * rcp * norm_weight) # fp32 products + returns (normed, updated) +``` + +An input `-0.0` travels as `+0.0` (the Lamport buffers' empty word is `-0.0`), so a sum is never `-0.0`: certified on +the plain sum (every rank's input holds `-0.0` in the same columns of every call of the test, and the plain result +there is `+0.0`, bit for bit). The squares inside the RMSNorm are rounded to bf16 before they are summed (both +kernels' code; the one-shot kernel marks it `FIXME: Use float square if accuracy issue`). + +**Two paths, chosen per call.** With `one_shot = T x H x W x 2` bytes, the call is sent **one-shot** when +`one_shot <= one_shot_max_bytes` and **two-shot** otherwise. Certified at the exact boundary for every certified +shape: `one_shot_max_bytes` equal to the footprint goes one-shot, one byte less goes two-shot, as the stage count the +call leaves in the workspace's flags shows (*State*). + +- One-shot: one kernel. Each rank writes its rows into every rank's buffer through the multicast mapping, waits for + all `W` rows of every token in its own copy and sums them; the residual add and the RMSNorm run in the same kernel. +- Two-shot: each rank writes token `t`'s row into rank `t mod W`'s buffer, which sums the `W` rows of its tokens and + writes the bf16 sums into every rank's buffer through the multicast mapping. Plain, the same kernel then waits for + every token's sum and copies it out; fused, a second kernel does that and adds the residual and normalizes. + +Both paths return the same plain sum and `updated` on the test's inputs (certified: every certified shape is sent +one-shot and then two-shot with the same inputs, each bit for bit against the exact reference). Beyond exact inputs +(the kernels' code): up to 8 ranks both paths add the ranks in ascending order with the same fp32 operations, so the +sum and `updated` agree bit for bit; at 16 ranks the one-shot adds two partial sums of 8 ranks while the two-shot adds +all 16 in sequence, so an inexact sum can differ in the last bit between the paths. `normed` can differ in the last +bit at any `W`: the two RMSNorm kernels form `updated * rcp * norm_weight` in different orders (`x * rcp * g` +one-shot, `g * x * rcp` two-shot); the test bounds both against fp32. Whatever the path, the result is bitwise the +same on every rank (certified, every call of the test). + +Fusion boundary. Inside: the exchange, the sum, and with `residual` the add and the RMSNorm. Outside: whatever +produced `input` (a row-parallel projection's partial output), the choice of `one_shot_max_bytes` (the caller's, per +call), the workspace (the target's). + +Inert (not exposed by the wrapper): the op's quantizing epilogues (`fusion_op` `RESIDUAL_RMS_NORM_QUANT_FP8`, +`_OUT_QUANT_FP8`, `_QUANT_NVFP4`, `_OUT_QUANT_NVFP4`, with `scale`). The wrapper passes `scale=None` and `fusion_op` +`NONE` or `RESIDUAL_RMS_NORM` only. + +## Signature + +```python +def mnnvl_fusion_allreduce( + input: torch.Tensor, + workspace: MnnvlWorkspace, + one_shot_max_bytes: int, + residual: Optional[torch.Tensor] = None, + norm_weight: Optional[torch.Tensor] = None, + eps: Optional[float] = None, +) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]] + +def required_buffer_bytes( + num_tokens: int, hidden: int, world_size: int, dtype: torch.dtype, one_shot_max_bytes: int +) -> int +``` + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `input` | `[T, H]`; `H` 3584 and 7168; `T` 1-8, 16, 32, 64 on both paths (and 65 two-shot, above 2 ranks) | bf16 | contiguous | CUDA, this rank's device | +| `workspace` | an `MnnvlWorkspace` of this rank's TP group whose buffers hold the call (see *State*) | — | — | — | +| `one_shot_max_bytes` | int >= 0; for every certified `(T, H)` its one-shot footprint `T x H x W x 2` and one less; Kimi K3's 4 MiB and 1 MiB (*Notes*) | Python int | — | — | +| `residual` | `None` (plain), or `[T, H]` | bf16 | contiguous | CUDA | +| `norm_weight` | `None`, or `[H]`; with `residual` | bf16 | contiguous | CUDA | +| `eps` | `None`, or a float with `residual`; 1e-5 certified | Python float | — | — | +| returns | plain: the sum `[T, H]`; fused: `(normed, updated)`, each `[T, H]` | bf16 | contiguous, newly allocated | = `input.device` | + +`input`, `residual` and `norm_weight` are read only. `required_buffer_bytes` is the space one Lamport buffer must +have for the call: `T x H x W x 2` one-shot, `2 x ceil(T / W) x W x H x 2` two-shot (two stages); certified equal to +the space the call's stages take, every call of the shape grid. + +## State + +**Object.** `MnnvlWorkspace` (`catalog/comm/mnnvl_workspace.py`), one per TP group, owned by the caller; the object's +own contract is the *State* section of `mnnvl_allreduce_attn_res.md`. This section states what this op does with it. + +**Contents and size.** Three Lamport buffers of `buffer_bytes` each behind one multicast mapping (every word `-0.0` +when armed) and the flag words `buffer_flags` (uint32 `[9]`: current buffer, dirty buffer, bytes per buffer, dirty +stage count, bytes to clear x 4, arrival count). A one-shot call writes `T x H x W x 2` bytes from the start of one +buffer. A two-shot call splits the buffer into two stages of `buffer_bytes / 2`: the scatter stage (`ceil(T / W) x W +x H x 2` bytes from the start) and the broadcast stage (`T x H x 2` bytes from the middle). A call that needs more +than `buffer_bytes` (`required_buffer_bytes`) is refused by the wrapper with `ValueError` on every rank before it +touches the workspace (certified: `T` = 65 at `H` = 7168 one-shot; the flags do not move and the next call is +correct; the op itself does not check, see *Notes*). Two-shot needs about `2 / W` of the one-shot space, so a call +too large for one buffer one-shot can fit two-shot (certified: that `T` = 65 call sent two-shot is correct, above 2 +ranks). The workspace is never grown. With `one_shot_max_bytes <= buffer_bytes` every one-shot call fits; the +two-shot calls Kimi K3 makes (at most 64 tokens of 7168) take at most 1.75 MiB, so a 4 MiB buffer holds all its +calls of this op (arithmetic, not a test). The test's buffer is the one-shot footprint of `[64, 7168]`, 3.5 MiB at +`W` = 4. + +**Who creates it, and when.** The target, in `post_load_weights`, with +`MnnvlWorkspace.create(mapping, buffer_bytes, fabric_handle=None)`: collective over the TP group, eager, every word +and flag armed before any rank returns (see `mnnvl_allreduce_attn_res.md`). It refuses CUDA-graph capture: certified +with every rank capturing, each raising `RuntimeError`. The check runs before any communication, so a rank that is +not capturing while its peers are would go on into the communicator split and wait for them (code). + +**Which ops may share one object.** Every MNNVL op of the group takes the same `comm_buffer` / `buffer_flags`: +`comm/mnnvl_allreduce_attn_res`, this entry on either path, and `comm/mnnvl_allgather_split`. Their calls form one +sequence and each takes exactly one turn of the one rotation, whatever the op and path (certified: after every eager +call of the test the flags equal this test's model of one turn per call). Certified on one workspace in Kimi K3's +order, a random rank late at every call: decode steps of [attention-residual all-reduce, head all-gather, latent +all-reduce `[T, 3584]`, fused all-reduce `[T, 7168]`] and wide steps of [all-reduce `[T, 7168]`, head all-gather, +latent all-reduce] at Kimi K3's ceilings, `T` = 8, 2, 16, 64, 1, 7, 32, 8, 3, 16, two layers each. Two objects are +two independent rotations: 20 calls of mixed shapes and paths alternating irregularly between two workspaces are all +correct, and each workspace's flags move with its own calls only (certified). They are not independent orders: on one +stream every rank must issue its collectives in the same order, whatever object each belongs to (measured for +`mnnvl_allreduce_attn_res`: ranks issuing B-then-A against A-then-B deadlocked, `runs/drafter/u4-mnnvl-srun-2`; this +op waits for its peers the same way). + +**Call-order invariant.** Every rank of the group makes the same sequence of calls on one workspace — the same +number, the `k`-th with the same op, `T`, `H`, fusion and path — across layers and decode steps, eager calls and +graph replays alike; and on one stream the same order of calls across workspaces. The path is part of the sequence: +the two paths write and wait for different words, so the ranks' `one_shot_max_bytes` must pick the same path (they do +when they pass the same value). + +**What a later launch reads.** `buffer_flags`, which every call leaves as: current = its own buffer plus one, mod 3; +dirty = its own buffer; bytes per buffer unchanged; dirty stage count 1 (one-shot) or 2 (two-shot); bytes to clear +`(T x H x W x 2, 0, 0, 0)` one-shot or `(ceil(T / W) x W x H x 2, T x H x 2, 0, 0)` two-shot; arrival count 0 +(certified after every eager call of the test). The next call, of any MNNVL op, takes the current buffer, whose words +must all be `-0.0` but for its own pushes, and clears the dirty one by those sizes. Each call's first kernel waits +for the previous kernel on the stream before it reads the flags (the kernels' code). + +**How it is re-armed.** Each call clears the previous call's buffer stage by stage, by the bytes the previous call +recorded and in the previous call's stage layout (`cpp/tensorrt_llm/common/lamportUtils.cuh`, +`LamportFlags::clearDirtyLamportBuf`): after a two-shot call both stages, after a one-shot call, an all-gather or an +attention-residual all-reduce the first stage, whichever path the clearing call itself takes. Certified by the shape +grid (every shape one-shot and then two-shot back to back on one workspace, then the next shape) and by the +sequences below. + +**Why the test drives call sequences.** See `mnnvl_allreduce_attn_res.md` (*State*): Phase 0's `k3_spec_accept` +re-armed its Lamport buffer for the current call's rows only; every single-call test passed, and a sequence whose row +count dipped and grew back caught it; in serving it made the ranks disagree and hang. This entry's test runs 16 decode +steps of 8 layers, each layer the latent all-reduce `[T, 3584]` and the fused all-reduce `[T, 7168]` chained through +`updated`, at `T` = 8, 8, 8, 2, 7, 8, 1, 1, 64, 3, 32, 8, 16, 1, 64, 8 with Kimi K3's ceilings, so the path changes +inside the sequence (at `W` = 4 `[64, 3584]` and `[32 or 64, 7168]` go two-shot, at `W` = 16 every step above 8 +tokens), a random rank 5 ms late at every call, each call against the reference and its flags against the model. + +**What a wrong order does.** Certified (the test's negative control, one-shot and two-shot): rank 0 issues two +same-shaped calls on one workspace in swapped order. Nothing raises and nothing hangs — the two calls write and wait +for the same words of the same buffers — but every rank's two results are wrong, and are exactly the +position-paired sums: the `k`-th call on every rank adds what every rank sent at position `k` (rank 0's second call +with the others' first; more than half the elements differ from the intended sums). A plain call right after is +correct again: a swapped pair realigns the positions. Two calls that write different words (another `T` or `H`, or +the other path) cannot pair like that: a rank would wait for words its peers do not write at that position, and hang +(the protocol, not exercised). A rank making one call more or fewer than its peers was not exercised; its positions +never realign. + +## Metadata consumed + +None besides `workspace` and `one_shot_max_bytes`, both explicit. The op keeps no cache and compiles nothing +(precompiled kernels). It finds the multicast mapping by looking `comm_buffer`'s address up in a process registry of +multicast buffers, which the workspace's handle keeps registered. `TRTLLM_ENABLE_PDL` (read once per process, +default on at SM 90 and newer) launches the kernels as programmatic dependents: the one-shot kernel releases its own +dependents as it starts, the two-shot one after its scatter; consumers of the outputs wait for the grid (the kernels' +statement); results do not depend on it. The launch shape (CTAs per token, cluster size) follows `T`, `H` and the +device's SM count, and in the fused form it sets the order in which the squares are summed. + +## Preconditions + +- bf16, contiguous, `input` 2-D (the op flattens all but the last dimension of an N-D input; not certified; it also + takes fp16 and fp32, not certified). `H` a multiple of 8; otherwise the op raises `RuntimeError` before it touches + the workspace (certified at `H` = 3580, on every rank, the flags unmoved, the next call correct). The one-shot + kernel's launch check names 65536 as its largest `H` (bf16). +- `W` in {2, 4, 8, 16, 32, 64} (the kernels' dispatch). +- `required_buffer_bytes(T, H, W, bf16, one_shot_max_bytes) <= workspace.buffer_bytes` (*State*). +- `residual`, `norm_weight` and `eps` together or not at all; otherwise the wrapper raises `ValueError` on every rank + before it touches the workspace (certified). +- Every rank calls with the same `T`, `H`, fusion and `one_shot_max_bytes` decision; the call order is the *State* + invariant. +- `workspace` was created before any capture. Calls may be captured: certified with a captured step of six calls on + one workspace — the attention-residual all-reduce, this op plain and fused one-shot at `T` = 8, the head all-gather, + a plain `[32, 7168]` and a fused `[16, 7168]` sent two-shot (the fused two-shot call is two kernels, both in the + graph) — replayed 8 times with rewritten inputs and an eager call of another shape or op on the same workspace + between replays, every replayed and eager result against the reference and the flags after each. +- SM 90 or newer (the kernels' check). + +## Notes + +- Certified path: 4 ranks of one GB200 tray (sm_100), one rank per GPU, POSIX-fd handles, PDL on (the default). + Test: `tests/unittest/_torch/modeling_v2/comm/_mnnvl_fusion_allreduce_op_matrix.py`. The reference is native + torch: the inputs are multiples of 1/16 with at most 4/16 per rank, so every sum over the ranks is exact in fp32 and + bf16 and the residual add is one bf16 rounding of an exact fp32 value in the op and in the reference; the sum and + `updated` are compared bit for bit, `normed` against the fp32 RMSNorm of `updated` within 1e-2 of its largest + magnitude (the bf16 squares cost at most 2^-9 on `rcp`, the bf16 output 2^-8). +- Design choices this entry follows (U4U5_PLAN §1): A1, "a typed state object per stateful op ..., built by an + explicit, collective, eager `create()` ... Tests drive real-state call sequences (layers x steps, capture + replay, + two objects interleaved) plus a negative control. `mutates_args` names every written buffer" (the op falls short of + the last, see the gaps below); A2, the matrix takes `--world-size` and `--launcher` (`mpirun` on one node, `srun` + across trays), CI runs it at 4 ranks on one GB200 tray; A4, "one caller-owned `MnnvlWorkspace` ..., shared by" every + MNNVL entry of the TP group, "`one_shot_max_bytes` per call". +- At `W` = 16 the one-shot kernel adds the ranks in two chunks of 8, a branch a 4-rank run never reaches. The 16-rank + receipt is pending. +- Kimi K3's calls (B7 82a110a92a, its code, not this test): the model sets every `MNNVLAllReduce` of the target, its + LM head and a drafter to `one_shot_max_bytes` = 4 MiB (`DECODE_AR_ONE_SHOT_MAX_BYTES`, against main's 1 MiB), and + a wide decode step (9 to 64 tokens) passes 1 MiB per call (`WIDE_AR_ONE_SHOT_MAX_BYTES`). Plain: the routed-latent + all-reduce `[T, 3584]` (decode steps where the latent exchange push is not used; wide steps) and a wide step's + attention all-reduces `[T, 7168]`. Fused: the DSpark drafter's residual + RMSNorm all-reduces `[T, 7168]` where + `k3_sandwich_plain` does not take the call. At 4 MiB a `[T, 7168]` call goes one-shot up to `T` = 18 at `W` = 16 + (73 at `W` = 4) and a `[T, 3584]` one up to 36 (146); at 1 MiB `[T, 7168]` up to 4 (18) and `[T, 3584]` up to 9 + (36). +- Gaps against A1 (the op is unchanged by this entry): the schema marks `comm_buffer` mutable `(a!)` but not + `buffer_flags`, which every call advances (A1 item 4); the op does not check that the call fits `comm_buffer` (the + attention-residual and all-gather ops do), so a direct op call over one buffer writes past it — the wrapper's + `required_buffer_bytes` check is the guard; `MnnvlWorkspace.create` accepts any multiple of 16 bytes, but the + two-shot broadcast stage starts at `buffer_bytes / 2` and is accessed in 16-byte vectors, so a two-shot call needs + `buffer_bytes` to be a multiple of 32 (code; every buffer in the test is); the schema's default + `one_shot_max_bytes=1048576` applies to a direct op call (the wrapper always passes one). +- In the model today the workspace is `MNNVLAllReduce`'s (a dict keyed by `Mapping`, grown on demand by the first + eager call that needs more, in 8 MiB steps) and the one-shot ceiling is a module attribute with a per-call override. + This entry takes both explicitly: the workspace sized at construction, the ceiling per call. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.py new file mode 100644 index 000000000000..2bde4b87db50 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.py @@ -0,0 +1,68 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""The MNNVL all-reduce over a caller-owned :class:`MnnvlWorkspace`: the sum over the TP group, or with a residual the +sum + residual add + RMSNorm; one-shot up to ``one_shot_max_bytes``, two-shot above.""" + +from typing import Optional, Tuple, Union + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* +from tensorrt_llm.functional import AllReduceFusionOp + +from .mnnvl_workspace import MnnvlWorkspace + +__all__ = ["MnnvlWorkspace", "mnnvl_fusion_allreduce", "required_buffer_bytes"] + + +def required_buffer_bytes( + num_tokens: int, hidden: int, world_size: int, dtype: torch.dtype, one_shot_max_bytes: int +) -> int: + """Bytes of one Lamport buffer a call of ``num_tokens`` rows of ``hidden`` needs: ``num_tokens * hidden * world * + element size`` when that is at most ``one_shot_max_bytes`` (one-shot), else two stages of ``num_tokens`` rounded up + to a multiple of ``world`` (two-shot).""" + itemsize = torch.empty((), dtype=dtype).element_size() + one_shot = num_tokens * hidden * world_size * itemsize + if one_shot <= one_shot_max_bytes: + return one_shot + return 2 * -(-num_tokens // world_size) * world_size * hidden * itemsize + + +def mnnvl_fusion_allreduce( + input: torch.Tensor, + workspace: MnnvlWorkspace, + one_shot_max_bytes: int, + residual: Optional[torch.Tensor] = None, + norm_weight: Optional[torch.Tensor] = None, + eps: Optional[float] = None, +) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: + """The sum of ``input`` over ``workspace``'s TP group (a new tensor of ``input``'s shape), or with ``residual``, + ``norm_weight`` and ``eps`` the pair ``(RMSNorm(sum + residual) * norm_weight, sum + residual)``. Sent one-shot + when ``num_tokens * hidden * world * element size <= one_shot_max_bytes``, else two-shot; the workspace's buffers + must hold the call (:func:`required_buffer_bytes`). Advances ``workspace`` by one call: every rank of the group + makes the same MNNVL calls on it in the same order.""" + fused = residual is not None + if fused != (norm_weight is not None) or fused != (eps is not None): + raise ValueError("mnnvl_fusion_allreduce: residual, norm_weight and eps go together") + hidden = input.shape[-1] + num_tokens = input.numel() // hidden + need = required_buffer_bytes(num_tokens, hidden, workspace.world_size, input.dtype, one_shot_max_bytes) + if need > workspace.buffer_bytes: + raise ValueError( + f"mnnvl_fusion_allreduce: the call needs {need} bytes per Lamport buffer, the workspace has " + f"{workspace.buffer_bytes}" + ) + fusion_op = AllReduceFusionOp.RESIDUAL_RMS_NORM if fused else AllReduceFusionOp.NONE + outputs = torch.ops.trtllm.mnnvl_fusion_allreduce( + input, + norm_weight, + residual, + eps, + workspace.comm_buffer(input.dtype), + workspace.buffer_flags, + fused, + None, + int(fusion_op), + one_shot_max_bytes, + ) + return (outputs[0], outputs[1]) if fused else outputs[0] diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md new file mode 100644 index 000000000000..781b5e447381 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md @@ -0,0 +1,281 @@ +--- +receipts: + sm_100: {status: pending, world_size: 4} +--- + +# k3_moe + +**Wraps** three calls of the caller-owned layer objects in `cute_dsl_kernels/k3_fused_moe/op.py`, one each. The +expert kernel `k3_moe` is a CuTe DSL kernel launched through its compiled function, not a torch op. + +| Function | Runs | Tokens `M` | +|---|---|---| +| `k3_moe` | `K3MoeLayer.__call__`: `torch.ops.trtllm.k3_route_quant`, then `k3_moe` as its programmatic dependent | 1-8 | +| `k3_moe_fused_front` | `K3MoeLayer.front`: `torch.ops.trtllm.k3_moe_front` (entry `moe/k3_moe_front`), then `k3_moe` | 1-8 | +| `k3_moe_wide` | `torch.ops.trtllm.k3_route_quant(early_trigger=True)`, then `K3MoeWideLayer.__call__`: the m_max 64 build of `k3_moe` | 1-64 | + +## Semantics + +Kimi K3's routed experts at decode size: this rank's routed partial, i.e. for each token the sum over its top-16 +experts that this rank holds of the expert's output times its routing weight. Two launches on the current stream, no +host synchronization (the op module's statement). + +Routing and MXFP8 quantization of the latent, `k3_route_quant` (inside `k3_moe` and `k3_moe_wide`): the CuTe DSL +form of `trtllm::kimi_k3_noaux_tc_mxfp8_quant`, whose four outputs it returns bit for bit (certified at `M` 1-8, 16, +33 and 64, with and without the early dependent trigger). Per token, as the kernel module states it: + +``` +s = 0.5 * tanh(0.5 * router_logits) + 0.5 # fp32, [896] +ids = the 16 experts with the largest s + e_score_correction_bias, descending, ties to the lower id +weights = bf16(s[ids] * routed_scaling_factor / (sum of the 16 s[ids] + 1e-20)) # in fp64 +xq, xs = MXFP8(latent): e4m3 codes, one UE8M0 scale per 32 columns, scale 2^ceil(log2(amax / 448)) +``` + +The experts, `k3_moe`, per token `t`: + +``` +y[t] = bf16( sum over the slots j with offset <= ids[t, j] < offset + num_local of + weights[t, j] * FC2_e(q8(SiTU(FC1_e(x_t)))) ), e = ids[t, j] - offset +x_t = xq[t] * 2^(xs[t] - 127) # the dequantized latent row, [3584] +FC1_e(x) = (x @ W_up[e]^T, x @ W_gate[e]^T) # [i_tp] each, fp32 accumulation +SiTU = 4 tanh(gate / 4) sigmoid(gate) * 25 tanh(up / 25) # caps 4.0 / 25.0 +q8 = MXFP8 per 32 intermediate columns, scale 2^ceil(log2(amax / 448)), e4m3 round to nearest even +FC2_e(a) = a @ W_down[e]^T # [3584], fp32 +``` + +`W_up[e]`, `W_gate[e]` (`[i_tp, 3584]`) and `W_down[e]` (`[3584, i_tp]`) are the values of the layer's MXFP4 expert +`e` (E2M1 codes times `2^(E8M0 - 127)` per 32 K elements), read in place from the TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers +(*Preconditions*). `offset` is `local_expert_offset`, `num_local` the state's. The SiTU caps are constants of the +`k3_moe` build (Kimi K3's `activation_situ_beta` 4.0 and `activation_situ_linear_beta` 25.0), not arguments; the +`gate_cap` / `linear_cap` arguments of `k3_moe_fused_front` are the shared activation's (the front's). + +Numerics, certified for `k3_moe` at every `M` 1-8 and for `k3_moe_wide` at `M` 1, 2, 7, 8, 9, 16, 33, 40 and 64, in +three routing cases each: random logits and no local expert for both; 16 local experts per token, none shared (128 +groups at `M` 8, the M <= 8 build's group capacity) for `k3_moe`; 100 local experts with 9 of 64 tokens and 124 with +one (324 groups at `M` 64, the wide build's capacity) for `k3_moe_wide`: + +- within the op-catalog gates (8 bf16 ulp of the token row's largest magnitude per element, 4 ulp relative RMS) of an + fp64 reference of the formula above over the dequantized experts, and of the stock path + (`kimi_k3_noaux_tc_mxfp8_quant`, then `mxe4m3_mxe2m1_block_scale_moe_runner` pre-routed with SiTU); +- run-to-run bit identical; +- a token's row may round differently when other tokens share the call: FC2 adds a token's expert terms in slices + whose bounds follow the call's group count. At most 1 bf16 ulp: each `M`'s rows against the same rows of the + 8-token call, and `k3_moe_wide`'s rows at `M` <= 8 against `k3_moe`'s; +- a token with no expert on this rank gets a zero row. + +`k3_moe_fused_front` returns `(y, shared)`: `y` is `k3_moe` applied to the front's routing and MXFP8 latent, `shared` +the front's shared activation (`moe/k3_moe_front`). Certified in that entry's 4-rank matrix: `y` within the op-catalog +gates of the stock runner on the front's own routing and latent; `shared` bit for bit the front's; on a head_flags +state, `y` and `shared` bit for bit the plain state's. + +Fusion boundary. Inside: the routing and the MXFP8 latent (`k3_moe`, `k3_moe_wide`) or the whole MoE front +(`k3_moe_fused_front`); the grouping of (expert, token) pairs; FC1, SiTU, the MXFP8 intermediate, FC2; the +routing-weighted combine. Outside: the router and latent-down GEMMs that produce `router_logits` and `latent` +(`k3_moe`, `k3_moe_wide`); the sum of the routed partials over the ranks that hold the other experts and +intermediate slices (the routed-latent all-reduce); the latent-up projection; the shared experts' down projection +(and, except in `k3_moe_fused_front`, their gate_up and activation); the residual. + +## Signature + +```python +def k3_moe( + latent: torch.Tensor, + router_logits: torch.Tensor, + e_score_correction_bias: torch.Tensor, + local_expert_offset: int, + routed_scaling_factor: float, + layer: K3MoeLayer, +) -> torch.Tensor + +def k3_moe_fused_front( + x: torch.Tensor, + w_front: torch.Tensor, + e_score_correction_bias: torch.Tensor, + local_expert_offset: int, + routed_scaling_factor: float, + shared_cols: int, + gate_cap: float, + linear_cap: float, + head: K3MoeHeadWorkspace, + layer: K3MoeLayer, +) -> Tuple[torch.Tensor, torch.Tensor] + +def k3_moe_wide( + latent: torch.Tensor, + router_logits: torch.Tensor, + e_score_correction_bias: torch.Tensor, + local_expert_offset: int, + routed_scaling_factor: float, + layer: K3MoeWideLayer, + out: Optional[torch.Tensor] = None, +) -> torch.Tensor +``` + +The wrapper module re-exports the state types (`K3MoeState`, `K3MoeLayer`, `K3MoeWideState`, `K3MoeWideLayer`, +`K3MoeHeadWorkspace`) and `is_supported`. + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `latent` | `[M, 3584]`; `M` 1-8 (`k3_moe`), 1-64 (`k3_moe_wide`) | bf16 | contiguous | CUDA, the state's device | +| `router_logits` | `[M, 896]` | fp32 | as `latent` | CUDA | +| `e_score_correction_bias` | `[896]` | fp32 | contiguous | CUDA | +| `local_expert_offset` | scalar: the global id of the layer's expert 0 (224 certified: experts `[224, 448)`) | Python int | — | — | +| `routed_scaling_factor` | scalar (2.827 certified) | Python float | — | — | +| `layer` | a `K3MoeLayer` of a plain `K3MoeState` (`k3_moe`); of a plain or head_flags one (`k3_moe_fused_front`); a `K3MoeWideLayer` (`k3_moe_wide`) | — | — | — | +| `out` (`k3_moe_wide`) | `None`, or `[>= M, 3584]`: the result is `out[:M]` (the same storage), rows past `M` untouched (certified at `M` 9 and 64) | bf16 | contiguous | CUDA | +| `x`, `w_front`, `shared_cols`, `gate_cap`, `linear_cap`, `head` (`k3_moe_fused_front`) | as `moe/k3_moe_front`; `head` that entry's `K3MoeHeadWorkspace` | — | — | — | +| returns | `y [M, 3584]` (`k3_moe_fused_front`: `(y, shared [M, shared_cols])`) | bf16 | contiguous, newly allocated (or `out[:M]`) | the inputs' device | + +The state objects' certified construction: `K3MoeState(device, 768, 224)` (`head_flags` False or True, `config` None), +`K3MoeWideState(device, 768, 224)` (`use_pdl` True), and `state.layer(w3_w1_weight, w3_w1_weight_scale, w2_weight, +w2_weight_scale)` over this rank's 224 experts in the TRTLLM-Gen W4A8_MXFP4_MXFP8 layout at intermediate 768: +`[224, 1536, 1792]`, `[224, 1536, 112]`, `[224, 3584, 384]`, `[224, 3584, 24]`, all uint8. + +## State + +**Objects.** Per rank and per device, owned by the caller (`cute_dsl_kernels/k3_fused_moe/op.py`, re-exported by +`catalog/moe/k3_moe.py`): + +- `K3MoeState` (`M` <= 8) and one `K3MoeLayer` per MoE layer from `state.layer(...)`; +- `K3MoeWideState` (`M` <= 64) and one `K3MoeWideLayer` per MoE layer; +- `k3_moe_fused_front` also takes the TP group's `K3MoeHeadWorkspace` (collective; its *State* is in + `moe/k3_moe_front`). + +**Contents and size.** At the certified layout (`i_tp` 768, `num_local` 224), certified right after construction: + +| Object | Tensor | Shape and dtype | Bytes | Between calls | +|---|---|---|---|---| +| `K3MoeState` | `c`: the FC1 -> FC2 intermediate slab | int8 `[G, 8, i_tp]`, `G = min(num_local, 128)` = 128 | 786,432 | armed: every byte 0x80 (FP8 -0.0) | +| | `cs`: its E8M0 scales | int8 `[G, 8, i_tp / 8]` | 98,304 | armed: bytes 0-3 of every 16-byte group 0xFF (E8M0 NaN), bytes 4-15 zero | +| | `part`: the FC2 partial rows | fp32 `[8 G, 3584]` | 14,680,064 | zero when built; then what the last call left | +| `K3MoeLayer` | `counters` | int32 `[32 + 2 G]` = `[288]` | 1,152 | zero | +| `K3MoeWideState` | `c`, `cs` | int8 `[G, 8, i_tp]`, `[G, 8, i_tp / 8]`, `G = 224 + (1024 - 224) / 8` = 324 | 1,990,656 + 248,832 | armed, as above | +| | `part` | fp32 `[1024, 3584]` | 14,680,064 | not initialized | +| `K3MoeWideLayer` | `counters` | int32 `[32 + 2 G]` = `[680]` | 2,720 | zero | + +`G` is the build's group capacity (`group_capacity` in the kernel module): an (expert, up to 8 tokens) group per local +expert a call routes to, at most `min(num_local, 8 x 16)` for `M` <= 8; the wide build gives an expert with `t` tokens +`ceil(t / 8)` groups. Each state also holds `mod` (its configuration's kernel module, from a process-wide cache keyed +by the configuration), `compiled` (the compiled `k3_moe`, built by the state's first call) and, for the M <= 8 build, +`head_flags`. A layer holds views of its four weight buffers (no copy) and its counters; it does not hold +`local_expert_offset`. + +**Who creates it, and when.** The target, in `post_load_weights`, after the expert weights are final (a layer reads +the buffers it was built over; see *What a wrong order does*): `K3MoeState(device, i_tp, num_local, +head_flags=False)` once per device and `state.layer(w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale)` +once per MoE layer; likewise `K3MoeWideState(device, i_tp, num_local)` and its layers. Not collective. Eager: +`K3MoeState()` and `K3MoeState.layer()` raise `RuntimeError` under CUDA-graph capture (certified); +`K3MoeWideState()` and `K3MoeWideState.layer()` do not check (*Notes*). The kernel compiles on the first call of each +state object (seconds), which must be eager: under capture that call raises `RuntimeError` before the `k3_moe` launch +and leaves the state uncompiled (certified for both builds). `config` (kernel options for tests and A/B runs) stays +`None`. No environment variable selects anything here except PDL (*Metadata consumed*). + +**Which ops may share one object.** All layers of a state share its slab and partial rows; each layer has its own +counters. Layers of one state may be built over the same weight buffers (two counter sets) or over different ones +(certified, both). On a plain `K3MoeState`, `k3_moe` and `k3_moe_fused_front` calls may be mixed (one build). A +head_flags `K3MoeState` serves `k3_moe_fused_front` only: `k3_moe` on its layers raises `ValueError` before any +launch (certified). `K3MoeWideState` serves `k3_moe_wide` only. Separate states share nothing: with calls on two +`K3MoeState`s and a `K3MoeWideState` interleaved in an irregular pattern (26 calls, back to back), each returns the +bits of the same call made alone (certified). + +**Call-order invariant.** The calls on all layers of one state run one after the other in one stream order. Every +call needs the slab armed and its layer's counters at zero, which only the end of the previous call on the state +guarantees; and (the kernel's statement) a call issues its first tile claim on its layer's counters before its grid +dependency wait, relying on the previous call of that layer having completed before the producer kernel launched this +one, which holds on one stream. Calls on one state from two streams at once are not exercised: nothing in the state +keeps two concurrent calls' groups apart. The state is per rank: there is no cross-rank order, except through the +head workspace for `k3_moe_fused_front` (`moe/k3_moe_front`). + +**What a later launch reads.** The slab armed: FC2 treats a group's intermediate as written only once its FP8 -0.0 +and E8M0 NaN sentinels are gone. The layer's counters at zero: the tile-queue cursor, the FC2 m-tile arrivals, the +per-group FC1 and FC2 counts. The partial rows: the M <= 8 build's combine also loads the rows of the token slots past +`M`, whose sums it drops; they are zero in a new state, so those loads never read unwritten memory (the op module's +statement). With head_flags: the head workspace's epoch `flags[2]` and its ready words (below). + +**How it is re-armed.** By `k3_moe` itself, on every call: the last FC2 task that reads a group's intermediate puts +its sentinels back, and each counter is reset by its last user (the cursor by the grid's last claim). Nothing is +cleared between calls and nothing records a call's size, so a call after a smaller one reads nothing an older, +larger call left. Certified: the slab armed and every counter zero after every single call and every sequence of +the test; decode steps of three layers at `M` 8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8 and of two wide layers at `M` 64, 64, +16, 1, 40, 64, 8, 23, 64, 9, 64, new inputs every call, back to back, each call the bits of the same call made alone +(itself within the gates of the stock path); a captured step (three `k3_moe` and one `k3_moe_wide` call) replayed 6 +times with rewritten inputs and eager calls of other `M` on the same states between replays, every replayed and +eager call the bits of the same call alone. + +With head_flags (`k3_moe_fused_front`), the ready-word handoff. The front releases `ready[t]` (token `t`'s ids and +weights) and `ready[8 + t]` (its MXFP8 row) as `epoch + 1`, with `epoch = head.flags[2]`; the head_flags `k3_moe` +reads the epoch, acquires those words for its `M` tokens instead of waiting for the front's grid, waits for the grid +only before its FC2 phase (which writes the output, memory it does not own), and at its last tile claim writes +`epoch + 1` into the ready words of the tokens past `M` and into `flags[2]` (the kernel's statement). So after every +head_flags call `flags[2]` and all 16 words hold one value, and the next call waits for a value no word holds, across +the int32 wraps too. Certified in `moe/k3_moe_front`'s matrix: before every head_flags call made alone no polled word +already holds `epoch + 1`, and after it the epoch advanced by one with all 16 words at it, across -1 -> 0 and +2^31 - 1 -> -2^31 too; after every back-to-back sequence and every replayed step, the epoch advanced by the number of +head_flags calls with all 16 words at it. + +**What a wrong order does.** The negative control (certified): a layer reads the weight buffers it was built over. +Layers of both builds are built over one set of experts whose weights are then "reloaded" by rebinding to new tensors +(the experts rolled by one), as a loader that replaces its parameters would. Nothing raises, and each layer keeps +returning the old experts' partial bit for bit, outside the op-catalog gates of the new experts' stock path: +silently stale. Copying the new weights into the old buffers in place is seen by the next call (the new experts' +partial, bit for bit), and layers built over the new tensors are correct. So build the layers once the weights are +final, and reload weights in place. Not exercised: one state on two streams at once (a race, no deterministic +control); a head_flags call whose polled ready words already hold `epoch + 1`, e.g. after a caller resets +`flags[2]` (`k3_moe` would read the front's output buffers before the front writes them; the matrix checks this +precondition before every head_flags call instead). + +## Metadata consumed + +Besides the state objects (explicit arguments): + +- `TRTLLM_ENABLE_PDL` (default on), read by `k3_route_quant` on every call (part of its compile key), by + `k3_moe_front` on every call, and by the M <= 8 build's kernel module when a configuration is first loaded (the + wide build takes `use_pdl` from its constructor). It changes scheduling, not results (the ops' statement); the + tests run with the default. +- Process-wide caches, result-neutral: the kernel modules keyed by configuration (`op._modules`), `k3_route_quant`'s + compiled kernels keyed by (early trigger, PDL), `k3_moe_front`'s keyed by its configuration. The compiled `k3_moe` + is per state object (`state.compiled`): every new state compiles on its first call. + +## Preconditions + +- sm_100, with the CuTe DSL package (`is_supported`): the layer constructor raises `ValueError` otherwise. +- The weights are the four contiguous uint8 buffers W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod writes (rows `[up; gate]` + interleaved and shuffled in 32-row blocks, scales block-interleaved 128 x 4): `w3_w1_weight [E, 2 i_tp, 1792]`, + `w3_w1_weight_scale [E, 2 i_tp, 112]`, `w2_weight [E, 3584, i_tp / 2]`, `w2_weight_scale [E, 3584, i_tp / 32]`, + with `E` the state's `num_local` (at most 896) and `i_tp` the state's, a multiple of 128. Anything else raises + `ValueError` in `state.layer()`. +- `M` 1-8 (`k3_moe`, `k3_moe_fused_front`) or 1-64 (`k3_moe_wide`): `M` 0 and 9, and 0 and 65, raise `ValueError` + before any launch, every state's slab, partial rows and counters keep their bits, and the next call returns the + bits of the same call made before (certified). +- `local_expert_offset` is the global id of the expert the layer's buffers start with. The layer does not record it: + another value computes other global ids with these weights, without an error. +- `k3_moe_wide`'s `latent`, `router_logits` and `e_score_correction_bias`, and `k3_moe`'s bias, contiguous: + `k3_route_quant` raises `ValueError` otherwise (its check). `k3_moe` makes its `latent` and `router_logits` + contiguous itself; only contiguous inputs are certified. +- Inputs on the state's device, with that device current (not checked). +- Calls may be captured once the state's first call has run eagerly: certified with the captured step above, and in + `moe/k3_moe_front`'s matrix for `k3_moe_fused_front` on both states. + +## Notes + +- Certified path: one GB200 GPU (sm_100) for `k3_moe` and `k3_moe_wide`, test + `tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe.py`, at the Kimi K3 TP16 deployment's routed-expert + rank layout (experts TP4 x EP4: 224 local experts of 896, intermediate 768), with random checkpoint-format MXFP4 + experts put through TRT-LLM's own loader. `k3_moe_fused_front`: 4 ranks of one GB200 tray in + `tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py` (entry point + `moe/test_modeling_v2_k3_moe_front_op_matrix.py`); its 16-rank receipt (A2) is pending with `moe/k3_moe_front`'s. +- References: an fp64 reference over the dequantized experts (from the checkpoint-format tensors) and the stock path. + The kernel tests (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py`, `test_k3_moe_wide.py`, + `test_k3_route_quant.py`) remain the exhaustive numerics; this entry's test copies their references. +- Gaps against A1 (the ops are unchanged by this entry): + - `K3MoeWideState()` and `K3MoeWideState.layer()` do not refuse CUDA-graph capture. Built under capture, their + allocations come from the graph's pool and the slab's arming fills and the counters' zeroing are captured instead + of run, so the first eager call would find the slab unarmed and the counters undefined. + - The `k3_moe` launch is not a torch op, so nothing declares what it writes: the state's slab and partial rows, the + layer's counters, and with head_flags the head workspace's `flags[2]` and ready words (A1: `mutates_args` names + every written buffer). `trtllm::k3_route_quant` writes only its new outputs (`mutates_args=()`). + - A layer does not record `local_expert_offset`, and nothing ties a state to a stream or checks the inputs' device. + - The compiled kernel is per state object, not per configuration: a second state of the same configuration + compiles again. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.py new file mode 100644 index 000000000000..31b5daeb99c8 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.py @@ -0,0 +1,90 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3's routed experts at decode size: this rank's routed partial from the persistent CuTe DSL kernel ``k3_moe`` +(FC1 + SiTU + FC2 with the routing-weighted combine over the TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers, read in place), +after the routing and MXFP8 quantization of its producer, on caller-owned state: a :class:`K3MoeState` and one +:class:`K3MoeLayer` per MoE layer for up to 8 tokens, a :class:`K3MoeWideState` and one :class:`K3MoeWideLayer` per +layer for up to 64.""" + +from typing import Optional, Tuple + +import torch + +# The state types; is_supported reads metadata only. Importing k3_route_quant's op registers trtllm::k3_route_quant. +from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.op import ( + K3MoeHeadWorkspace, + K3MoeLayer, + K3MoeState, + K3MoeWideLayer, + K3MoeWideState, + is_supported, +) +from tensorrt_llm._torch.cute_dsl_kernels.k3_route_quant import op as _k3_route_quant_op # noqa: F401 + +__all__ = [ + "K3MoeHeadWorkspace", + "K3MoeLayer", + "K3MoeState", + "K3MoeWideLayer", + "K3MoeWideState", + "is_supported", + "k3_moe", + "k3_moe_fused_front", + "k3_moe_wide", +] + + +def k3_moe( + latent: torch.Tensor, + router_logits: torch.Tensor, + e_score_correction_bias: torch.Tensor, + local_expert_offset: int, + routed_scaling_factor: float, + layer: K3MoeLayer, +) -> torch.Tensor: + """This rank's routed partial ``[M, 3584]`` bf16 for ``M <= 8`` tokens: ``trtllm::k3_route_quant`` of + ``router_logits`` (fp32 ``[M, 896]``) and ``latent`` (bf16 ``[M, 3584]``), then ``k3_moe`` on ``layer``'s experts + (global ids ``[local_expert_offset, local_expert_offset + num_local)``). Writes ``layer``'s state's scratch (left + armed) and ``layer``'s counters (left zero).""" + return layer(latent, router_logits, e_score_correction_bias, local_expert_offset, routed_scaling_factor) + + +def k3_moe_fused_front( + x: torch.Tensor, + w_front: torch.Tensor, + e_score_correction_bias: torch.Tensor, + local_expert_offset: int, + routed_scaling_factor: float, + shared_cols: int, + gate_cap: float, + linear_cap: float, + head: K3MoeHeadWorkspace, + layer: K3MoeLayer, +) -> Tuple[torch.Tensor, torch.Tensor]: + """``(routed partial [M, 3584] bf16, shared activation [M, shared_cols] bf16)`` for the MoE input ``x`` (bf16 + ``[M <= 8, 7168]``): ``trtllm::k3_moe_front`` over ``head`` (see ``moe/k3_moe_front``), then ``k3_moe`` on + ``layer``'s experts. Advances ``head`` by one front call; writes ``layer``'s scratch and counters as + :func:`k3_moe`.""" + return layer.front( + x, w_front, e_score_correction_bias, local_expert_offset, routed_scaling_factor, shared_cols, gate_cap, + linear_cap, head, + ) # fmt: skip + + +def k3_moe_wide( + latent: torch.Tensor, + router_logits: torch.Tensor, + e_score_correction_bias: torch.Tensor, + local_expert_offset: int, + routed_scaling_factor: float, + layer: K3MoeWideLayer, + out: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """This rank's routed partial ``[M, 3584]`` bf16 for ``1 <= M <= 64`` tokens: ``trtllm::k3_route_quant`` (its + dependents launched early, as ``k3_moe``'s PDL producer), then the m_max 64 build of ``k3_moe`` on ``layer``'s + experts. ``out``: bf16, contiguous, at least ``[M, 3584]``; its first M rows are the result (a new tensor without + it). Writes ``layer``'s state's scratch (left armed) and ``layer``'s counters (left zero).""" + ids, weights, x_fp8, x_sf = torch.ops.trtllm.k3_route_quant( + router_logits, e_score_correction_bias, latent, routed_scaling_factor, early_trigger=True + ) + return layer(x_fp8, x_sf, ids, weights, local_expert_offset, out=out) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md new file mode 100644 index 000000000000..30438984f096 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md @@ -0,0 +1,235 @@ +--- +receipts: + sm_100: {status: pending, world_size: 4} +--- + +# k3_moe_front + +**Wraps** `torch.ops.trtllm.k3_moe_front` (one call), over a caller-owned `K3MoeHeadWorkspace`. + +## Semantics + +Kimi K3's MoE front for a decode batch of at most 8 tokens, in one kernel: the row-sharded MoE head GEMV (this rank's +slice of the latent-down projection and of the router), the all-gather of every rank's head slice over the TP group's +head workspace, the top-16 routing and the MXFP8 quantization of the gathered latent, and this rank's slice of the +shared experts' gate_up GEMV with its SiTU-and-mul. Every rank of the workspace's TP group (`W` ranks) calls with the +same MoE input `x` `[M, K]` (`K` = 7168) and its own front weight. With `WL = 3584 / W` and `WE = 896 / W`: + +``` +head_r = x @ head_weight_r^T # fp32 [M, WL + WE] on rank r: latent columns [r WL, (r + 1) WL), + # then the logits of experts [r WE, (r + 1) WE) +latent = bf16(head_0[:, :WL] | head_1[:, :WL] | ... | head_{W-1}[:, :WL]) # [M, 3584]; -0.0 becomes +0.0 +logits = head_0[:, WL:] | head_1[:, WL:] | ... | head_{W-1}[:, WL:] # fp32 [M, 896] +(topk_ids, topk_weights, quantized, scales) = k3_route_quant(logits, e_score_correction_bias, latent, rsf) +g, u = bf16(x @ gate^T), bf16(x @ up^T) # this rank's shared_cols columns +shared = bf16(gate_cap * tanh(g / gate_cap) * sigmoid(g) * linear_cap * tanh(u / linear_cap)) +``` + +`head_weight_r` is rank `r`'s `[WL + WE, K]` head slice and `gate`, `up` this rank's `[shared_cols, K]` shared rows, +all inside the ranks' `w_front`. `k3_route_quant` is the routing and quantization of `moe/k3_moe` (top-16 of the +sigmoid plus the bias, ties to the lower id, weights renormalized times `routed_scaling_factor`; MXFP8 with one UE8M0 +scale per 32 columns): the front selects with all warps of a CTA but returns `top16_warp`'s experts, order and weight +bits, and quantizes with the same device code (the kernel's statement). Every head and shared value is an fp32 sum +over `K` split across the 8 CTAs of a cluster, the 8 partials added in cluster-rank order from +0.0 (the kernel's +statement): deterministic, but not the summation order of another GEMM. + +Certified at every `M` 1-8 and for every front call the matrix makes alone, with payloads whose head and shared sums +are exact in fp32 in any order (`x` a multiple of 1/8 in [-1/4, 1/4], the latent-down and shared rows multiples of +1/16 in [-1/8, 1/8], the router rows multiples of 1/8 in [-1/4, 1/4]): + +- `topk_ids`, `topk_weights`, `quantized` and `scales` bit for bit those of `trtllm::kimi_k3_noaux_tc_mxfp8_quant` + on the exactly gathered head (logits in fp32, latent rounded to bf16), and bit for bit the same on every rank; +- `shared` within 2e-2 (of its largest magnitude) of the fp32 SiTU of the bf16-rounded exact sums: the kernel's tanh + and sigmoid are fast approximations; +- every output run-to-run bit identical, and each `M`'s outputs bit for bit the same rows of the 8-token call. + +With general inputs the head sums round in the kernel's own order, so the logits and the latent can differ from +another GEMM's in the last bit; the kernel test (`test_k3_moe_front.py`) bounds the effect: other experts only where +the 16th and 17th selection keys are within 1e-4, more than 99.9 % of the MXFP8 codes and scales equal. + +Fusion boundary. Inside: the head GEMV, the head all-gather, the routing, the MXFP8 latent, the shared gate_up and +its SiTU. Outside: the producer of the MoE input `x`; the routed experts (`moe/k3_moe`'s `k3_moe_fused_front` runs +`k3_moe` on this op's outputs within the same call); the shared experts' down projection and its all-reduce; the +routed partials' all-reduce and the latent-up projection. + +## Signature + +```python +def k3_moe_front( + x: torch.Tensor, + w_front: torch.Tensor, + e_score_correction_bias: torch.Tensor, + routed_scaling_factor: float, + shared_cols: int, + gate_cap: float, + linear_cap: float, + workspace: K3MoeHeadWorkspace, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor] +``` + +The wrapper module also exports the load-time helpers `front_weight(head_weight, gate_up_weight)`, which packs +`w_front`, and `weight_supported(world, shared_cols, k_in, device)`. + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `x` | `[M, 7168]`, `M` 1-8; the same values on every rank | bf16 | contiguous | CUDA, this rank's device | +| `w_front` | `front_weight(head_weight, gate_up_weight)`: `[head_tiles(W) x 128 + 2 shared_cols, 7168]`, i.e. `[1920, 7168]` at `W` 4 | bf16 | contiguous, as `front_weight` returns it | CUDA | +| `head_weight` (to `front_weight`) | this rank's `[WL + WE, 7168]`: its latent-down rows, then its router rows (1120 rows at `W` 4; 280 at `W` 16) | bf16 | contiguous | CUDA | +| `gate_up_weight` (to `front_weight`) | `[2 shared_cols, 7168]`: gate rows, then up rows | bf16 | contiguous | CUDA | +| `e_score_correction_bias` | `[896]` | fp32 | contiguous | CUDA | +| `routed_scaling_factor` | scalar (2.827 certified) | Python float | — | — | +| `shared_cols` | 384 (Kimi K3 TP16's per-rank width: two shared experts of 3072 over 16 ranks) | Python int | — | — | +| `gate_cap`, `linear_cap` | 4.0, 25.0 (Kimi K3's SiTU caps) | Python float | — | — | +| `workspace` | the `K3MoeHeadWorkspace` of this rank's TP group, `W` = 4 certified (see *State*) | — | — | — | +| returns | `topk_ids [M, 16]` int32, `topk_weights [M, 16]` bf16, `quantized [M, 3584]` float8_e4m3fn, `scales [M, 112]` uint8 (linear: one byte per 32 columns), `shared [M, shared_cols]` bf16 | — | contiguous, newly allocated | `x.device` | + +Inert (not exposed by the wrapper): `ring` (4, the weight ring's stages) and `ag_ready` (`None`: the standalone front +publishes no ready words; `K3MoeLayer.front` on a head_flags state passes the workspace's `ready`, see *State*). +`gate_cap` and `linear_cap` are compile-time constants of the kernel: each pair compiles once (*Metadata consumed*). + +## State + +**Object.** `K3MoeHeadWorkspace` (`cute_dsl_kernels/k3_fused_moe/op.py`, re-exported by `catalog/moe/k3_moe_front.py` +and `catalog/moe/k3_moe.py`), one per TP group, owned by the caller. + +**Contents and size.** One multicast allocation of `workspace_words(W)` int32 words per rank, 157,696 words (616 KiB) +at `W` 4, 8 and 16 (certified at 4): + +- `uc`: this rank's words, every one 0x80000000 (empty) when armed. First the two alternating Lamport buffers + `[buffer 2][token 8][rank W][slot]` (43,008 words), a slot holding `3584 / W / 8` latent vectors of 8 bf16 (4 words + each), then `896 / W / 4` logit vectors; the front pushes each rank's latent rows into the latent vectors, and the + logit vectors (the layout of `k3_route_quant_ag.py`) stay empty. Then the router partials, + `[buffer 2][token 8][rank W][cluster rank 8][896 / W]` fp32 words (114,688 words): each CTA of a GEMV cluster + pushes its split-K partial of its rank's router rows there. +- `mc`: the same words through the multicast mapping, where every rank pushes. +- `flags`, int32 `[4]`: `[0]` the buffer of the next call; `[1]` unused (zero); `[2]` the ready words' epoch, + advanced only by a head_flags build of `k3_moe` (`moe/k3_moe`); `[3]` the sign-ins of the CTAs that read `[0]` in + the current call. +- `ready`, int32 `[32]`: `[t]` token `t`'s routing and `[8 + t]` its MXFP8 row, released as `epoch + 1` by a + publishing front (`K3MoeLayer.front` on a head_flags state); `[16, 32)` unused. +- `rank`, `world_size`; `handle`, the `McastGPUBuffer` that owns the memory (the workspace is valid while this object + lives); `comm`, the TP-group communicator the handles were exchanged over. + +The size depends on `W` only, not on `M`: every call fits. + +**Who creates it, and when.** The target, in `post_load_weights`, with `K3MoeHeadWorkspace.create(mapping, +fabric_handle=None)`: collective over `mapping`'s TP group (every rank calls it at the same point; it returns on every +rank or raises on every rank, the agreement also being the barrier that keeps any rank from pushing into a peer's +buffer before the peer has emptied it); eager: under CUDA-graph capture it raises `RuntimeError` before entering the +collective (certified, on every rank). It empties every word and zeroes `flags` and `ready` (certified, both of the +matrix's workspaces). `fabric_handle`: share the memory by fabric handle (required across nodes) or POSIX file +descriptor; default `mapping.is_multi_node()`. No environment variable is read. + +**Which ops may share one object.** Every MoE front call of the TP group: this entry and `moe/k3_moe`'s +`k3_moe_fused_front`, on a plain or a head_flags `K3MoeState`. They form one sequence on the workspace: mixed in one +step (the fused front on each state, then the front alone) they all stay correct (certified, below). Separate from the +MNNVL all-reduce workspace and the sandwich workspace: a front call advances neither. Two workspaces are two +independent rotations and two independent epochs: 20 calls alternating between two workspaces in an irregular +pattern, kinds mixed, each return the bits of the same call made alone, and each workspace's epoch advances by its +own head_flags calls only (certified). They are not independent orders: each call spins until its peers' pushes of +the same call arrive and the calls of one stream run one after the other, so ranks that order calls on two +workspaces differently on one stream would deadlock (the kernel's design; not exercised). + +**Call-order invariant.** Every rank of the group makes the same sequence of front calls on one workspace (the same +number of calls, the `k`-th with the same `M`), eager calls and graph replays alike, with the same `x`; and on one +stream the same order of calls across workspaces and other collectives. Each call reads `flags[0]` (its buffer), +pushes its latent slice and router partials into that buffer on every rank, polls its own copy until every rank's +pushes of this call are there, and empties what it read. Every CTA that reads `flags[0]` signs in on `flags[3]` right +after the read; the CTA that flips `flags[0]` for the next call waits for all of them and zeroes `flags[3]` (the +kernel's statement; certified: `flags[0]` flips once per call, and `flags[3]` is zero after each single call and +each sequence). + +**What a later launch reads.** `flags[0]`, and its buffer's words, which must be empty except for this call's pushes; +`flags[3]` at zero. A publishing front also reads the epoch `flags[2]`, before it lets its dependent launch (the +kernel's statement), and releases its ready words as `epoch + 1`; `k3_moe`, not the front, reads the ready words. + +**How it is re-armed.** By its readers: each CTA writes the empty word back over every word it read, which are the +words of the same call's tokens. The next push into a buffer comes from a call two calls later, which starts only +after this call has ended on this rank (the kernel's statement, from the stream order and the grid-dependency waits). +So no separate clear and no record of the previous call's size is needed, and a call after a smaller one finds no +word of an older, larger one. Certified: every word of this rank's buffers empty and `flags[1]`, `flags[3]` zero +after each single call and after each sequence below; decode steps of three layers (the fused front on the plain +state, the fused front on the head_flags state, the front alone) at `M` 8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8 with new +inputs every call, back to back, a random rank 5 ms late before every call: every call the bits of the same call made +alone (itself checked against the references). + +The ready words are re-armed by the head_flags `k3_moe` (`moe/k3_moe`, *State*): at its last tile claim it writes +`epoch + 1` into the words of the tokens past its `M` and into `flags[2]`, so after every head_flags call `flags[2]` +and all 16 words hold one value and the next call waits for a value no word holds. Certified: before every head_flags +call made alone no word it polls already holds `epoch + 1`, and after it the epoch advanced by one with all 16 words +at it; after every back-to-back sequence and every replayed step, the epoch advanced by the number of head_flags +calls with all 16 words at it; across -1 -> 0 (from zeroed words and epoch 0, calls at `M` 1, 1, then the epoch +preset to -2 and calls at `M` 1, 8, 3, 8: the `M` 8 call at epoch -1 waits for 0, the value a word past an earlier +call's tokens would still hold without that re-arm) and across 2^31 - 1 -> -2^31 (epoch preset to 2^31 - 2, calls at +`M` 8, 1, 8), each call's outputs the plain build's bits. The standalone front leaves `flags[2]` and the ready words +untouched (certified). Nothing else may write them: a caller that resets `flags[2]` while the words hold +`epoch + 1` would let the next head_flags `k3_moe` read the front's outputs before the front writes them (not +exercised). + +**What a wrong order does.** Certified at `W` = 4 (the negative control): rank 0 issues two same-shaped front calls +on one workspace in swapped order. Nothing raises and nothing hangs, the rotation positions still agree, but each call +pairs with the peers' call at the same position: every rank's gathered head mixes rank 0's slice of its input with +the peers' slices of theirs. Every rank returns the same wrong routing and latent, bit for bit those of the reference +of the mixed inputs; the MXFP8 latent's columns `[0, 3584 / W)` and their scales are bit for bit those of rank 0's +input's call and the other columns those of the peers' input's call, so each rank's latent codes differ from those of +the call it made in more than half of the mixed-in columns. The shared activation, local to each rank, is right for +each rank's own input, bit for bit. A plain call right after is correct. Ranks calling with different `M` at one +position, or one rank making a call more or fewer, were not exercised; a rank would then poll for words its peers +never push. + +## Metadata consumed + +Besides `workspace` (an explicit argument): + +- A process-wide cache of compiled kernels keyed by (`W`, `shared_cols`, tiles, clusters, `K`, ring, `gate_cap`, + `linear_cap`, publish, PDL, half-tile head). The first call of a key compiles (seconds) and must be eager: under + capture it raises `RuntimeError` ("must run once per configuration outside CUDA-graph capture first") before any + launch, the workspace untouched (certified with other SiTU caps). The cache is result-neutral. +- `TRTLLM_ENABLE_PDL` (default on), read on every call and part of the key; it changes scheduling, not results (the + op's statement). +- The device's cluster capacity (`max_clusters`, cached per device), which sizes the grid and picks the head's + geometry: 128-row tiles, or one round of 64-row half-tiles when they fit next to the shared tiles + (`half_geometry`; the op's statement: TP16). At `W` 4 the head (1120 rows per rank, 18 half-tiles) does not fit in + one round, so the 4-rank run uses 128-row tiles. + +## Preconditions + +- sm_100 (tcgen05, clusters of 8 CTAs, one per SM) and `weight_supported(W, shared_cols, K, device)`: `W` in {4, 8, + 16}, `K` a multiple of 1024, `shared_cols` a multiple of 64, and the tiles fitting the device's clusters (true for + `W` 4, 384 columns, `K` 7168 on the certified device). Any unsupported call (`M` outside 1-8, a `w_front` of + another shape or dtype, another `W`) raises `ValueError` before it touches the workspace (the op's check; certified + for `M` 0 and 9 on every rank, the workspace's words, flags and ready words unchanged, the next call correct). A + workspace smaller than `workspace_words(W)` raises `ValueError` too. +- Every rank calls with the same `x` and `M`, its own head slice in `w_front`; the call order is the *State* + invariant. Nothing checks that `x` agrees across ranks: a rank with another `x` mixes its slice into every rank's + result (the negative control shows that effect). +- `workspace` was created, and the kernel compiled (one eager call per configuration), before any capture. Calls may + be captured: certified with a captured step of three calls at `M` 8 (the fused front on the plain and on the + head_flags state, the front alone) replayed 6 times with rewritten inputs, an eager call of another `M` on the same + workspace between replays, every replayed and eager call the bits of the same call made alone, the epoch advanced + once per head_flags call, replayed or eager. + +## Notes + +- Certified path: 4 ranks of one GB200 tray (sm_100), one rank per GPU, the head sharded over those 4 ranks (1120 + rows per rank, 128-row tiles), the shared activation at TP16's per-rank width (384). Kimi K3 TP16 shards the head + over 16 ranks on four trays (280 rows per rank), which runs the half-tile geometry and which only a 16-rank run + reaches; the matrix takes `--world-size` and `--launcher` (A2) and that receipt is pending. +- Test: `tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py` (rank body), collected by + `tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_front_op_matrix.py`. It also certifies + `moe/k3_moe`'s `k3_moe_fused_front` cells. +- Reference: native torch for the head (every rank's slice in fp64, exact for the payloads, then fp32; the latent + columns rounded to bf16) and the shared gate_up (fp64, exact, rounded to bf16) with SiTU in fp32; the stock + `trtllm::kimi_k3_noaux_tc_mxfp8_quant` for the routing and quantization of the gathered head. Every rank draws every + rank's head slice from one seed, so each holds the whole reference. The routed experts of the fused cells are + random checkpoint-format MXFP4 (224 per rank, rank `r` at global ids `[224 (r % 4), 224 (r % 4) + 224)`), the + reference for `y` the stock TRTLLM-Gen runner. The kernel test + (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py`) covers Gaussian payloads and races the + publishing front's epoch read against `k3_moe`'s epoch advance (`check_publish_order`); neither is repeated here. +- `mutates_args` names every buffer the op writes (`ag_uc`, `ag_mc`, `ag_flags`, `ag_ready`). Gaps against A1 (the op + is unchanged by this entry): the compile cache is a module-level dict (result-neutral, documented above); the slots' + logit vectors are dead space for this kernel; `ring` is not exposed. +- In the model the head workspace was `head_workspace(mapping)`, a module dict created on the first eager call; this + entry takes the explicit object instead (A1). diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.py new file mode 100644 index 000000000000..c88abf6d0982 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.py @@ -0,0 +1,46 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3's MoE front at decode size in one kernel: the sharded MoE head GEMV, its all-gather over a caller-owned +:class:`K3MoeHeadWorkspace`, the top-16 routing, the MXFP8 latent, and the shared experts' gate_up + SiTU.""" + +from typing import Tuple + +import torch + +# Importing front_op registers trtllm::k3_moe_front. front_weight packs the front's one weight at load time; +# weight_supported reads metadata only. +from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.front_op import front_weight, weight_supported +from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.op import K3MoeHeadWorkspace + +__all__ = ["K3MoeHeadWorkspace", "front_weight", "k3_moe_front", "weight_supported"] + + +def k3_moe_front( + x: torch.Tensor, + w_front: torch.Tensor, + e_score_correction_bias: torch.Tensor, + routed_scaling_factor: float, + shared_cols: int, + gate_cap: float, + linear_cap: float, + workspace: K3MoeHeadWorkspace, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Return ``(topk_ids, topk_weights, quantized, scales, shared)`` for the MoE input ``x`` (bf16 ``[M <= 8, 7168]``, + the same on every rank): the top-16 routing of the gathered router logits with ``e_score_correction_bias`` and the + MXFP8 latent with its UE8M0 scales, as ``trtllm::k3_route_quant`` returns them for the gathered head, and the shared + experts' activation (bf16 ``[M, shared_cols]``). ``w_front`` from :func:`front_weight`. Advances ``workspace`` by + one call: every rank of the group makes the same front calls on it in the same order.""" + return torch.ops.trtllm.k3_moe_front( + x, + w_front, + e_score_correction_bias, + routed_scaling_factor, + shared_cols, + gate_cap, + linear_cap, + workspace.uc, + workspace.mc, + workspace.flags, + workspace.rank, + workspace.world_size, + ) diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_latent_reduce_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_k3_latent_reduce_op_matrix.py new file mode 100644 index 000000000000..228606e27021 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_latent_reduce_op_matrix.py @@ -0,0 +1,539 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU certification matrix for the ``comm/k3_latent_reduce`` catalog entry and its ``K3LatentExchange``. + +The op is the consumer half of Kimi K3's latent all-reduce at decode size: every rank's producer pushes its routed +partial into the exchange, and the op, on every rank, sums the ranks' rows. Its correctness depends on state that +outlives a call (the call count in ``flags[0]``, whose parity picks the half every push and reduce use, and the words +a reduce empties for the push two calls later), so beyond single calls this drives call *sequences*: layers x steps +with the token count dipping and growing back and a random rank late, two exchanges interleaved, CUDA-graph capture +and replay mixed with eager calls, the count across its int32 wrap, and two negative controls in which ranks break +the call order and get a wrong answer without an error. + + CUDA_VISIBLE_DEVICES=0,1,2,3 python _k3_latent_reduce_op_matrix.py [--world-size 4] + srun -N 4 --ntasks-per-node 4 --mpi=pmix python _k3_latent_reduce_op_matrix.py --launcher srun --world-size 16 + +Not a pytest module: one fixed sequence of checks inside one W-rank job (they share the exchanges and their counts). +The collected entry point is ``test_modeling_v2_k3_latent_reduce_op_matrix.py``. + +The producers (the push-only routed-expert kernels) are not part of this entry, so every push is emulated as they +store it: this rank's partial rows, -0.0 as +0.0, as int32 words (bf16 pairs) copied through the multicast mapping +into slot [rank] of the call's half of every rank's buffer. The emulation takes the half from the host's count of the +exchange's calls, which ``assert_clean`` checks against ``flags[0]``; a producer reads ``flags[0]`` on the device. +The copies launch without programmatic dependent launch, so the op's PDL condition (every kernel between a reduce +and the next push ends only after its predecessor has ended) holds throughout. + +Every rank draws every rank's partial from one seed, so each holds the whole reference: the sum in the MNNVL +one-shot's order (fp32 over chunks of 8 ranks in rank order, each from +0, the chunks added in order, then bf16), +computed with torch. Most checks use multiples of 1/16, whose sum is exact in any order. The bit-identity check adds +partials whose sum depends on the order (per element a +B / -B pair that absorbs the small values summed beside it; +another order is shown to change at least a quarter of the elements) and compares the op bit for bit with that +reference and with the MNNVL one-shot all-reduce itself (``comm/mnnvl_fusion_allreduce`` over an ``MnnvlWorkspace``). +Every output is also compared bitwise across the ranks. +""" + +import random +import sys +from pathlib import Path +from typing import List, Optional, Sequence + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parent)) +import _lockstep as ls # noqa: E402 + +assert torch.cuda.is_available(), "k3_latent_reduce requires CUDA devices" + +DEADLINE_S = 900 +H = 3584 # the routed latent: the all-reduce's row width +ROW_WORDS = H // 2 # int32 words (bf16 pairs) of one row +MAX_TOKENS = 8 +RANK_CHUNK = 8 # the one-shot sums the ranks in chunks of 8 +EMPTY_WORD = -(2**31) # 0x80000000: a word no producer has written +NEG_ZERO = -(2**15) # bf16 -0.0 as int16, which the producers store as +0.0 +WORLDS = (4, 8, 16) # the op's TP sizes +TOKENS = (1, 2, 3, 4, 5, 6, 7, 8) +LAYERS = 12 # even: a captured step keeps the halves' parity (check_graph_capture_and_replay) +DIP_STEPS = (8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8) +GRAPH_TOKENS = (8, 3) # one captured step per batch size, as an engine keeps one graph per size +EAGER_BETWEEN = (5, 1) # an even number of eager calls between two replays +REPLAYS = 4 +MNNVL_BUFFER_BYTES = 1 << 20 # one Lamport buffer: 8 rows of 3584 bf16 from 16 ranks go one-shot +# The least share of the elements that another summation order must change in the "cancel" partials. +ORDER_GUARD = 0.25 +WRAP_PRESET = 2**31 - 3 + +R = None +entry = None +create_exchange = None +mnnvl = None +EX_A = None +EX_B = None +MNNVL_WS = None +STATS = {"order_guard": 1.0} + + +def bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def same_bits(a: torch.Tensor, b: torch.Tensor) -> bool: + return a.shape == b.shape and a.dtype == b.dtype and torch.equal(bits(a), bits(b)) + + +def differing(a: torch.Tensor, b: torch.Tensor) -> float: + """The share of the elements whose bits differ (same shapes).""" + return (bits(a) != bits(b)).float().mean().item() + + +def int32(n: int) -> int: + """``n`` as the int32 count holds it (two's complement wrap).""" + return (n + 2**31) % 2**32 - 2**31 + + +def ordered_sum(parts: Sequence[torch.Tensor], chunk: int) -> torch.Tensor: + """fp32 sums of ``chunk`` consecutive partials each, in the given order and from +0; those sums added in order to + +0; then bf16 (round to nearest even). With ``chunk`` = 8 and the ranks in order, the MNNVL one-shot's order.""" + total = torch.zeros(parts[0].shape, dtype=torch.float32, device=parts[0].device) + for base in range(0, len(parts), chunk): + acc = torch.zeros_like(total) + for part in parts[base : base + chunk]: + acc = acc + part.float() + total = total + acc + return total.bfloat16() + + +def exact_partials(g: torch.Generator, tokens: int) -> List[torch.Tensor]: + """Multiples of 1/16 in [-1/4, 1/4]: a sum of up to 16 of them is exact in fp32 and bf16, whatever the order.""" + return [ls.exact_bf16(g, (tokens, H), -4, 5, 1 / 16) for _ in range(R.world)] + + +def randn_partials(g: torch.Generator, tokens: int) -> List[torch.Tensor]: + return [ + (torch.randn(tokens, H, generator=g, device="cuda") * 0.5).bfloat16() + for _ in range(R.world) + ] + + +def cancel_partials(g: torch.Generator, tokens: int) -> List[torch.Tensor]: + """Small values (multiples of 1/8, at most 31/8 in magnitude) on every rank, then per element +B on one rank and + -B on another, B = 2^30 (1 + k/128) with k in 1..127. Half an fp32 ulp at B is 64, more than any sum of the small + values (at most 14 of them), so a running fp32 sum that holds B absorbs every small value added to it, and the + pair cancels exactly: the result depends on where the pair falls in the order of summation, and another order + gives another result in most elements.""" + shape = (tokens, H) + parts = torch.stack( + [ + (torch.randint(-31, 32, shape, generator=g, device="cuda") / 8).bfloat16() + for _ in range(R.world) + ] + ) + big = (128 + torch.randint(1, 128, shape, generator=g, device="cuda")).float() * 2.0**23 + sign = torch.randint(0, 2, shape, generator=g, device="cuda").float() * 2 - 1 + first = torch.randint(0, R.world, shape, generator=g, device="cuda") + second = (first + torch.randint(1, R.world, shape, generator=g, device="cuda")) % R.world + parts.scatter_(0, first.unsqueeze(0), (sign * big).bfloat16().unsqueeze(0)) + parts.scatter_(0, second.unsqueeze(0), (-sign * big).bfloat16().unsqueeze(0)) + return list(parts.unbind(0)) + + +PARTIALS = {"exact": exact_partials, "randn": randn_partials, "cancel": cancel_partials} + + +class Exchange: + """A ``K3LatentExchange`` with the host's count of its calls, from which the emulated pushes take their half.""" + + def __init__(self, state): + self.state = state + self.calls = 0 + + def push(self, partial: torch.Tensor, half: Optional[int] = None) -> None: + """What a push-only producer stores: this rank's rows, -0.0 as +0.0, as int32 words into slot [rank] of + ``half`` (default: this call's, the count's parity) of every rank's buffer, through the multicast mapping.""" + rows = partial.shape[0] + words = partial.contiguous().view(torch.int16) + words = words.masked_fill(words == NEG_ZERO, 0).view(torch.int32) + dest = self.state.mc.view(2, MAX_TOKENS, R.world, ROW_WORDS) + dest[self.calls & 1 if half is None else half, :rows, R.rank].copy_(words) + + def reduce(self, tokens: int) -> torch.Tensor: + out = entry(tokens, self.state) + self.calls += 1 + return out + + +class Call: + """One push + reduce of ``tokens`` rows: every rank's partial (each rank draws them all from one seed, so each + holds the whole reference) and the reference. Every partial has -0.0 entries (pushed as +0.0), and in some columns + every rank's is -0.0, whose sum is +0.0.""" + + def __init__(self, seed: int, tokens: int, kind: str = "exact"): + g = torch.Generator(device="cuda").manual_seed(seed) + self.tokens = tokens + self.parts = PARTIALS[kind](g, tokens) + for r, part in enumerate(self.parts): + part[:, r::97] = -0.0 + part[:, 5::101] = -0.0 + + def ref(self) -> torch.Tensor: + """The MNNVL one-shot's order (``reduceOneshotLamport``), which the kernel states it follows.""" + return ordered_sum(self.parts, RANK_CHUNK) + + def run(self, ex: Exchange) -> torch.Tensor: + ex.push(self.parts[R.rank]) + return ex.reduce(self.tokens) + + +def verify(call: Call, got: torch.Tensor, where: str) -> None: + want = call.ref() + assert got.dtype == torch.bfloat16 and got.is_contiguous(), f"{where}: {got.dtype} output" + assert got.shape == want.shape, ( + f"{where}: shape {tuple(got.shape)}, expected {tuple(want.shape)}" + ) + assert same_bits(got, want), f"{where}: {differing(got, want):.2%} of the elements differ" + assert R.same_on_ranks(got), f"{where}: ranks disagree" + + +def assert_clean(ex: Exchange, where: str) -> None: + """Every rank's last reduce on ``ex`` has ended and no rank has pushed the next call: both halves empty and flags + [count, 0, 0, 0] (the count, the arrivals word back to 0, two unused words) on every rank.""" + R.barrier() + flags = ex.state.flags.tolist() + left = int((ex.state.uc != EMPTY_WORD).sum()) + expected = [int32(ex.calls), 0, 0, 0] + # The allgather is also the barrier that keeps every rank from pushing until every rank has looked. + assert R.all_true(left == 0 and flags == expected), ( + f"{where}: rank {R.rank} flags {flags} (expected {expected}), {left} words not empty" + ) + + +def raised_under_capture(fn) -> str: + """Run ``fn`` under CUDA-graph capture; the message of the RuntimeError it raised, or '' if it raised none.""" + graph, stream = torch.cuda.CUDAGraph(), torch.cuda.Stream() + resting = torch.cuda.current_stream() + stream.wait_stream(resting) + message = "" + try: + with torch.cuda.graph(graph, stream=stream): + try: + fn() + except RuntimeError as exc: + message = str(exc) + finally: + # A capture that fails when it ends leaves its own stream current; put the resting one back. + torch.cuda.set_stream(resting) + del graph + return message + + +def check_exchange_is_armed_and_sized() -> None: + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import k3_latent_reduce as kernel + + words = 2 * MAX_TOKENS * R.world * ROW_WORDS + assert kernel.buffer_words(R.world) == words + for ex in (EX_A, EX_B): + s = ex.state + assert (s.rank, s.world_size) == (R.rank, R.world) + assert s.uc.dtype == s.mc.dtype == s.flags.dtype == torch.int32 + assert s.uc.numel() == s.mc.numel() == words, f"{s.uc.numel()} words, expected {words}" + assert s.uc.device.index == torch.cuda.current_device() and s.flags.device == s.uc.device + assert bool((s.uc == EMPTY_WORD).all()), "every word empty (0x80000000)" + assert s.flags.tolist() == [0, 0, 0, 0], f"flags {s.flags.tolist()}" + uc, mc, flags, rank = s.push_args() + assert uc is s.uc and mc is s.mc and flags is s.flags and rank == R.rank + assert EX_A.state.uc.data_ptr() != EX_B.state.uc.data_ptr(), "two exchanges, two buffers" + assert EX_A.state.flags.data_ptr() != EX_B.state.flags.data_ptr(), "two exchanges, two counts" + + +def check_capture_refusals() -> None: + """Under CUDA-graph capture ``K3LatentExchange.create`` raises RuntimeError on every rank at once, and on one rank + alone while its peers do not call it: before any collective step, where that rank would wait for its peers. The + op's first call, which would compile the kernel, raises on every rank before it launches anything. The exchange + is untouched. Runs before every eager call of the op: the compile cache must still be cold.""" + # Imported by the op's first call; imported here so that nothing is imported inside the capture. + import cutlass.cute.runtime # noqa: F401 + + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import k3_latent_reduce # noqa: F401 + + def refused() -> bool: + message = raised_under_capture(lambda: create_exchange(R.mapping, fabric_handle=R.fabric)) + return "outside CUDA-graph capture" in message + + every = refused() + R.barrier() + alone = refused() if R.rank == R.world - 1 else True + assert R.all_true(every and alone), f"create: every rank {every}, one rank alone {alone}" + first = raised_under_capture(lambda: entry(MAX_TOKENS, EX_A.state)) + assert R.all_true("outside CUDA-graph capture first" in first), f"first call: {first!r}" + assert_clean(EX_A, "after the refused calls") + + +def check_single_calls() -> None: + """One push + reduce at every token count 1-8, exact partials: the reference bit for bit on every rank; after each + call both halves are empty again, the count is up by one and the arrivals word is 0.""" + for t in TOKENS: + call = Call(1000 + t, t) + verify(call, call.run(EX_A), f"M {t}") + assert_clean(EX_A, f"after M {t}") + + +def check_bit_identical_to_mnnvl_oneshot() -> None: + """The op's statement: its rows are the MNNVL one-shot all-reduce of the partials, bit for bit. At every token + count, with partials whose sum depends on the order (summing the ranks in reverse, or at W > 8 without chunks, + changes at least ORDER_GUARD of the elements) and with normal-distributed ones, the op, ``mnnvl_fusion_allreduce`` + sent one-shot and the torch reference in the one-shot's order agree bit for bit.""" + for t in TOKENS: + assert t * H * R.world * 2 <= MNNVL_BUFFER_BYTES, "the reference call must go one-shot" + for k, kind in enumerate(("cancel", "randn")): + call = Call(2000 + 10 * t + k, t, kind) + want = call.ref() + if kind == "cancel": + others = [ordered_sum(call.parts[::-1], R.world)] + if R.world > RANK_CHUNK: + others.append(ordered_sum(call.parts, R.world)) + share = min(differing(other, want) for other in others) + STATS["order_guard"] = min(STATS["order_guard"], share) + assert share >= ORDER_GUARD, f"M {t}: another order changes only {share:.2%}" + oneshot = mnnvl(call.parts[R.rank], MNNVL_WS, MNNVL_BUFFER_BYTES) + got = call.run(EX_A) + assert same_bits(got, oneshot), f"M {t} {kind}: differs from the MNNVL one-shot" + assert same_bits(oneshot, want), ( + f"M {t} {kind}: the MNNVL one-shot differs from the reference" + ) + verify(call, got, f"M {t} {kind}") + assert_clean(EX_A, "after the bit-identity calls") + + +def check_unsupported_token_counts_raise_on_every_rank() -> None: + """Token counts 0 and 9 raise ValueError on every rank before the op touches the exchange (no producer pushes + them: a half holds 8 rows); the count and the words are unchanged and the next call is correct.""" + for t in (0, MAX_TOKENS + 1): + try: + entry(t, EX_A.state) + raised = False + except ValueError: + raised = True + assert R.all_true(raised), f"M {t} did not raise ValueError on every rank" + assert_clean(EX_A, "after the rejected calls") + call = Call(3000, MAX_TOKENS) + verify(call, call.run(EX_A), "after the rejected calls") + + +def run_step(ex: Exchange, seed: int, tokens: int, late: Optional[random.Random] = None) -> None: + """One decode step: LAYERS push + reduce pairs queued back to back (no host synchronization between them, as in a + step), a random rank late before each when ``late`` is given; every result checked after the step.""" + calls = [Call(seed + layer, tokens) for layer in range(LAYERS)] + R.barrier() + outs = [] + for call in calls: + if late is not None: + R.late(late.randrange(R.world)) + outs.append(call.run(ex)) + for layer, (call, got) in enumerate(zip(calls, outs)): + verify(call, got, f"step seed {seed} M {tokens} layer {layer}") + + +def check_dip_and_regrow_sequence() -> None: + """Decode steps of LAYERS layers at M 8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8, a random rank 5 ms late before every call + (its push lands while the others' reduces poll): a call after a smaller one must not read words an older, larger + call left. After every step both halves are empty and the count is right on every rank.""" + late = random.Random(7) + for i, t in enumerate(DIP_STEPS): + run_step(EX_A, 10_000 + 100 * i, t, late) + assert_clean(EX_A, f"after step {i} (M {t})") + + +def check_two_exchanges_interleaved() -> None: + """Two exchanges are two counts and two buffers: calls alternate between them in an irregular pattern (A A B A B B + ...), so the halves they use differ from call to call, a random rank late before each; every call is correct and + each exchange ends clean at its own count. The pattern is the same on every rank: a reduce waits for its peers' + pushes of the same call and a stream runs its kernels in order, so ranks issuing calls on two exchanges in + different orders would deadlock (not run).""" + pattern = "AABABBAAAB" * 2 + late = random.Random(9) + for i, which in enumerate(pattern): + ex = EX_A if which == "A" else EX_B + call = Call(20_000 + i, (3, 8, 1, 8, 5)[i % 5]) + R.late(late.randrange(R.world)) + verify(call, call.run(ex), f"interleaved {which} {i}") + assert_clean(EX_A, "exchange A after the interleaving") + assert_clean(EX_B, "exchange B after the interleaving") + + +def check_graph_capture_and_replay() -> None: + """Two captured steps on one exchange, as an engine keeps one graph per batch size: LAYERS push + reduce pairs at + M 8 in one graph and at M 3 in another, replayed alternately REPLAYS times each with rewritten partials and a + random rank late, two eager calls of other sizes between replays. Replays and eager calls advance one count. + + The emulated pushes are copies whose half is chosen on the host when they are captured, so each graph holds an + even number of pairs and an even number of eager calls runs between replays: every replay starts on the parity + its capture assumed. The reduce reads the count on the device at replay time, as a producer does.""" + ex = EX_B + for t in GRAPH_TOKENS: + run_step(ex, 30_000 + 100 * t, t) # each size eagerly first, as an engine's warm-up does + R.barrier() + base = ex.calls + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + graphs = {} + for t in GRAPH_TOKENS: + bufs = [torch.zeros(t, H, dtype=torch.bfloat16, device="cuda") for _ in range(LAYERS)] + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + outs = [] + for layer, buf in enumerate(bufs): + ex.push(buf, half=(base + layer) & 1) + outs.append(entry(t, ex.state)) + assert len(outs) == LAYERS + graphs[t] = (graph, bufs, outs) + R.barrier() + late = random.Random(11) + for rep in range(REPLAYS): + for t in GRAPH_TOKENS: + graph, bufs, outs = graphs[t] + calls = [Call(40_000 + 1000 * rep + 100 * t + layer, t) for layer in range(LAYERS)] + for buf, call in zip(bufs, calls): + buf.copy_(call.parts[R.rank]) + assert (ex.calls - base) % 2 == 0, "a replay must start on the parity of its capture" + R.barrier() + R.late(late.randrange(R.world)) + graph.replay() + ex.calls += LAYERS + for layer, (call, got) in enumerate(zip(calls, outs)): + verify(call, got, f"replay {rep} of M {t}, layer {layer}") + for j, m in enumerate(EAGER_BETWEEN): + call = Call(50_000 + 100 * rep + 10 * t + j, m) + verify(call, call.run(ex), f"eager M {m} after replay {rep} of M {t}") + assert_clean(ex, "after the replays") + del graphs + + +def check_count_parity_across_the_int32_wrap() -> None: + """The count is int32 and only its parity is read: preset to 2^31 - 3 on every rank while the exchange is clean + and idle, it crosses 2^31 - 1 -> -2^31 in the next calls, which stay correct, and then reads -2^31 + 1.""" + ex = EX_A + R.barrier() + ex.state.flags[0] = WRAP_PRESET + ex.calls = WRAP_PRESET + R.barrier() + for i, t in enumerate((8, 2, 5, 8)): + call = Call(60_000 + i, t) + verify(call, call.run(ex), f"call {i} across the wrap") + assert int32(ex.calls) == -(2**31) + 1 + assert_clean(ex, "after the wrap") + + +def check_swapped_calls_are_wrong() -> None: + """Negative control: rank 0 makes two same-shaped calls in swapped order (it pushes its partial of the second call + first). Every push is still followed by one reduce of its token count, so the counts agree and nothing raises or + hangs, but every rank's two results are wrong: each reduce sums rank 0's partial of the other call. The exchange + is clean afterwards and a plain call is correct.""" + first, second = Call(70_000, MAX_TOKENS), Call(70_001, MAX_TOKENS) + R.barrier() + if R.rank == 0: + got_second, got_first = second.run(EX_A), first.run(EX_A) + else: + got_first, got_second = first.run(EX_A), second.run(EX_A) + torch.cuda.synchronize() + wrong = [differing(got, call.ref()) for call, got in ((first, got_first), (second, got_second))] + assert R.all_true(min(wrong) > 0.5), ( + f"rank {R.rank}: the swap went unnoticed, wrong shares {wrong}" + ) + assert_clean(EX_A, "after the swapped pair") + call = Call(70_002, MAX_TOKENS) + verify(call, call.run(EX_A), "after the swapped pair") + + +def check_token_count_mismatch_returns_stale_rows() -> None: + """Negative control: rank 0's reduces disagree with the pushed token count. Every rank pushes 8 rows and rank 0 + reduces 4: its 4 rows are right, but rows 4-7 of that half stay full in its buffer (a reduce empties only the rows + it sums). Two calls later, on the same half, every rank pushes 4 rows and rank 0 reduces 8: it does not wait for + rows 4-7, which are already full, and returns the older call's sums for them. Nothing raises or hangs; that reduce + empties all 8 rows, so the exchange is clean again and the next call is correct. + + Not run, because they wait forever: a reduce of rows nobody pushed (with nothing stale there), and a push into the + other half.""" + ex = EX_A + old, between, short = Call(80_000, 8), Call(80_001, 8), Call(80_002, 4) + mine = 4 if R.rank == 0 else 8 + R.barrier() + half = ex.calls & 1 + ex.push(old.parts[R.rank]) + got = ex.reduce(mine) + R.barrier() + rows = ex.state.uc.view(2, MAX_TOKENS, R.world, ROW_WORDS) + left = int((ex.state.uc != EMPTY_WORD).sum()) + stuck = 4 * R.world * ROW_WORDS if R.rank == 0 else 0 + full = R.rank != 0 or bool((rows[half, 4:] != EMPTY_WORD).all()) + ok = same_bits(got, old.ref()[:mine]) and left == stuck and full + assert R.all_true(ok), ( + f"rank {R.rank}: {left} words not empty (expected {stuck}), rows 4-7 full {full}" + ) + verify(between, between.run(ex), "the call on the other half") + assert ex.calls & 1 == half + ex.push(short.parts[R.rank]) + got = ex.reduce(8 if R.rank == 0 else 4) + torch.cuda.synchronize() + if R.rank == 0: + ok = same_bits(got[:4], short.ref()) and same_bits(got[4:], old.ref()[4:]) + else: + ok = same_bits(got, short.ref()) + assert R.all_true(ok), f"rank {R.rank}: the reduce of 8 rows did not return the 4 stale ones" + assert_clean(ex, "after the mismatched calls") + call = Call(80_003, MAX_TOKENS) + verify(call, call.run(ex), "after the mismatched calls") + + +CHECKS = [ + check_exchange_is_armed_and_sized, + # Before every eager call of the op: it needs a cold compile cache. + check_capture_refusals, + check_single_calls, + check_bit_identical_to_mnnvl_oneshot, + check_unsupported_token_counts_raise_on_every_rank, + check_dip_and_regrow_sequence, + check_two_exchanges_interleaved, + check_graph_capture_and_replay, + check_count_parity_across_the_int32_wrap, + # Stay last: they deliberately break the call-order invariant. + check_swapped_calls_are_wrong, + check_token_count_mismatch_returns_stale_rows, +] + + +def _run_one_rank(args) -> int: + global R, entry, create_exchange, mnnvl, EX_A, EX_B, MNNVL_WS + R = ls.Rank(args) + assert R.world in WORLDS, f"k3_latent_reduce runs on TP groups of {WORLDS} ranks, not {R.world}" + from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm import ( + k3_latent_reduce as module, + ) + from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm import ( + mnnvl_fusion_allreduce as mnnvl_module, + ) + + entry = module.k3_latent_reduce + create_exchange = module.K3LatentExchange.create + mnnvl = mnnvl_module.mnnvl_fusion_allreduce + with torch.inference_mode(): + EX_A = Exchange(create_exchange(R.mapping, fabric_handle=R.fabric)) + EX_B = Exchange(create_exchange(R.mapping, fabric_handle=R.fabric)) + MNNVL_WS = mnnvl_module.MnnvlWorkspace.create( + R.mapping, MNNVL_BUFFER_BYTES, fabric_handle=R.fabric + ) + code = ls.run_checks(R, CHECKS) + if R.rank == 0: + print( + f"[rank 0] world {R.world}; calls on A {EX_A.calls} (count {int32(EX_A.calls)}), on B {EX_B.calls}; " + f"least share of elements another summation order changes {STATS['order_guard']:.2%}", + flush=True, + ) + return code + + +if __name__ == "__main__": + ARGS = ls.parse_args(sys.argv[1:]) + if ARGS.rank_worker or ARGS.launcher == "srun": + sys.exit(_run_one_rank(ARGS)) + ls.spawn(__file__, ARGS, DEADLINE_S) + print("OK") diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py new file mode 100644 index 000000000000..487fc265ce20 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py @@ -0,0 +1,874 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU certification matrix for the ``moe/k3_moe_front`` catalog entry and its ``K3MoeHeadWorkspace``, with the +fused-front cells of ``moe/k3_moe`` (``k3_moe_fused_front`` on a plain and on a head_flags ``K3MoeState``). + +The front's correctness depends on state that outlives a call: the head workspace's two alternating Lamport buffers +(flags[0] says which one a call uses; every call flips it), the readers' re-arm of every word they read, and, with a +head_flags build of k3_moe, the ready words and their epoch (flags[2]), which that k3_moe advances. So beyond single +calls this drives call sequences: layers x steps with the token count dipping and growing back and a random rank late +at every call, two workspaces interleaved, CUDA-graph capture and replay mixed with eager calls, the epoch across both +int32 wraps, the capture refusals, and last a negative control in which one rank swaps two calls and every rank gets a +wrong (and exactly predictable) answer without an error. + + CUDA_VISIBLE_DEVICES=0,1,2,3 python _k3_moe_front_op_matrix.py [--world-size 4] + srun -N 4 --ntasks-per-node 4 --mpi=pmix python _k3_moe_front_op_matrix.py --launcher srun --world-size 16 + +Not a pytest module: one fixed sequence of checks inside one W-rank job (they share the workspaces, the states and +their counters). The collected entry point is ``moe/test_modeling_v2_k3_moe_front_op_matrix.py``. + +Shapes. The head is row-sharded over the run's W ranks: 3584 / W latent rows and 896 / W router rows per rank (1120 at +W 4; 280 at W 16, Kimi K3 TP16's). The shared activation is TP16's per-rank width at every W (384 columns: two shared +experts of 3072 over 16 ranks). The routed experts are one rank of experts TP4 x EP4 (224 local experts, intermediate +768; rank r's are global ids [224 (r % 4), 224 (r % 4) + 224)), random checkpoint-format MXFP4 put through TRT-LLM's +own TRTLLM-Gen loader. + +References. Every rank draws every rank's head slice and every call's input from one seed, so each rank holds the +whole reference. The payloads are exact: x is a multiple of 1/8 in [-1/4, 1/4], the latent-down and shared rows +multiples of 1/16 in [-1/8, 1/8], the router rows multiples of 1/8 in [-1/4, 1/4] (_lockstep.exact_bf16), so every +head and shared sum over K = 7168 is a multiple of 1/128 below 2^9 and exact in fp32 in any summation order: the +front's split-K sums equal the reference's. The reference is the head in fp64 (exact, then fp32; the latent columns +rounded to bf16), gathered, and the stock trtllm::kimi_k3_noaux_tc_mxfp8_quant on it: the front's routing ids and +weights, MXFP8 codes and scales must equal it bit for bit (the front routes with k3_route_quant's selection, weights +and quantization: the kernel's statement), and be bit for bit the same on every rank. The shared activation: this +rank's gate_up in fp64 (exact) rounded to bf16, SiTU in fp32, within 2e-2 of its largest magnitude (the kernel's tanh +and sigmoid are fast approximations). k3_moe_fused_front's routed partial against the stock TRTLLM-Gen +W4A8_MXFP4_MXFP8 runner on the front's own routing and latent (op-catalog gates: 8 bf16 ulp of the row's max per +element, 4 ulp relative RMS); the head_flags build bit for bit against the plain one. Call sequences compare every +call bit for bit with the same call made alone (itself checked against the references first). +""" + +import math +import random +import sys +from pathlib import Path +from types import SimpleNamespace + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parent)) +import _lockstep as ls # noqa: E402 + +assert torch.cuda.is_available(), "k3_moe_front requires CUDA devices" + +DEADLINE_S = 1200 +HIDDEN = 7168 # the MoE input's width, the front's K +LATENT = 3584 +EXPERTS = 896 +TOP_K = 16 +SV = 32 +# Kimi K3 TP16's per-rank shared activation: two shared experts of 3072 over 16 ranks. +SHARED_COLS = 384 +# Kimi K3's SiTU caps: the front's shared-activation arguments, and k3_moe's build constants. +GATE_CAP, LINEAR_CAP = 4.0, 25.0 +RSF = 2.827 +# One rank of the routed experts' TP4 x EP4. +I_TP, E_LOCAL, MOE_TP = 768, 224, 4 +# 0x80000000: an empty head workspace word. +EMPTY = -(2**31) +# The head workspace per rank: 157,696 int32 words at W 4, 8 and 16. +WORKSPACE_WORDS = 2 * 8 * (LATENT // 8 + EXPERTS // 4) * 4 + 2 * 8 * 8 * EXPERTS +ULP = 2.0**-8 +SHARED_TOL = 2e-2 +TOKENS = (1, 2, 3, 4, 5, 6, 7, 8) +LAYERS = 3 +# Layer l of a step: k3_moe_fused_front on the plain state, on the head_flags state, and k3_moe_front alone. +KINDS = ("fused", "flags", "front") +DIP_STEPS = (8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8) +REPLAYS = 6 + +R = None +OPS = None +WS_A = None +WS_B = None +PLAIN = None # the plain K3MoeState and its layers +FLAGS = None # the head_flags K3MoeState and its layers +LAYER_WEIGHTS = [] # per layer: every rank's head slice, this rank's gate_up and front weight +BIAS = None +EXPERT_BUFFERS = None +OFFSET = 0 +WL = WE = 0 # latent columns and experts per rank +STATS = { + "y_elt_ulp": 0.0, + "y_rms_ulp": 0.0, + "shared_err": 0.0, + "fronts_checked": 0, + "rows_bits_as_m8": 0, + "rows_checked": 0, +} + + +def i32(v: int) -> int: + return (v + 2**31) % 2**32 - 2**31 + + +def bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.uint8) + + +def same(a: torch.Tensor, b: torch.Tensor) -> bool: + return a.shape == b.shape and a.dtype == b.dtype and torch.equal(bits(a), bits(b)) + + +def same_result(got, want) -> bool: + return len(got) == len(want) and all(same(a, b) for a, b in zip(got, want)) + + +def same_bits_on_ranks(*tensors: torch.Tensor) -> bool: + """Every tensor bit for bit the same on every rank (the raw bytes, gathered).""" + mine = [bits(t).cpu().numpy().tobytes() for t in tensors] + every = R.comm.allgather(mine) + return all(other == every[0] for other in every[1:]) + + +def max_ulp(a: torch.Tensor, b: torch.Tensor) -> int: + """Largest distance in bf16 ulps between two bf16 tensors (bit patterns as ordered integers).""" + + def ordered(x): + i = x.contiguous().view(torch.int16).int() + return torch.where(i < 0, -(i & 0x7FFF), i) + + return int((ordered(a) - ordered(b)).abs().max().item()) if a.numel() else 0 + + +class Call: + """One call: the MoE input ``x`` (the same on every rank), the layer whose weights it uses, and its kind. + + Kinds: "front" (k3_moe_front), "fused" (k3_moe_fused_front on the plain K3MoeState's layer), "flags" + (k3_moe_fused_front on the head_flags K3MoeState's layer). + """ + + def __init__(self, seed: int, tokens: int, layer: int = 0, kind: str = "front", x=None): + if x is None: + g = torch.Generator(device="cuda").manual_seed(seed) + x = ls.exact_bf16(g, (tokens, HIDDEN), -2, 3, 1 / 8) + self.x, self.layer, self.kind = x, layer, kind + + def first(self, m: int, kind=None) -> "Call": + """The first ``m`` tokens of this call's input, as a call of ``kind`` (default: this call's).""" + return Call(0, m, self.layer, kind or self.kind, x=self.x[:m].contiguous()) + + def as_kind(self, kind: str) -> "Call": + return Call(0, 0, self.layer, kind, x=self.x) + + def run(self, ws, x=None): + x = self.x if x is None else x + lw = LAYER_WEIGHTS[self.layer] + if self.kind == "front": + return OPS.front(x, lw.front, BIAS, RSF, SHARED_COLS, GATE_CAP, LINEAR_CAP, ws) + layer = (PLAIN if self.kind == "fused" else FLAGS).layers[self.layer] + return OPS.fused( + x, lw.front, BIAS, OFFSET, RSF, SHARED_COLS, GATE_CAP, LINEAR_CAP, ws, layer + ) + + +# ── references ──────────────────────────────────────────────────────────── + + +def front_reference(x, layer: int, xs=None): + """The unfused front for input ``x`` with layer ``layer``'s weights. + + Every rank's head slice of its input (``xs[r]``; default ``x`` on every rank) in fp64 (exact for these payloads), + then fp32, gathered (latent columns rounded to bf16); the stock routing + MXFP8 quantization of the gathered head; + this rank's shared activation. Returns (ids, weights, quantized, scales, shared). + """ + lw = LAYER_WEIGHTS[layer] + xs = xs or [x] * R.world + heads = [(xr.double() @ w.double().t()).float() for xr, w in zip(xs, lw.heads)] + latent = torch.cat([h[:, :WL] for h in heads], dim=1).bfloat16().contiguous() + logits = torch.cat([h[:, WL:] for h in heads], dim=1).contiguous() + ids, w, q, s = torch.ops.trtllm.kimi_k3_noaux_tc_mxfp8_quant(logits, BIAS, latent, RSF) + gate_up = (xs[R.rank].double() @ lw.gate_up.double().t()).float().bfloat16().float() + gate, up = gate_up[:, :SHARED_COLS], gate_up[:, SHARED_COLS:] + shared = ( + GATE_CAP + * torch.tanh(gate / GATE_CAP) + * torch.sigmoid(gate) + * (LINEAR_CAP * torch.tanh(up / LINEAR_CAP)) + ) + return ids, w, q, s, shared.bfloat16() + + +def verify_front(got, ref, where: str) -> None: + """The front's five outputs against ``ref`` (see the module docstring), and bit for bit across the ranks.""" + ids, w, q, s, shared = got + m = ids.shape[0] + assert ids.dtype == torch.int32 and tuple(ids.shape) == (m, TOP_K), where + assert w.dtype == torch.bfloat16 and tuple(w.shape) == (m, TOP_K), where + assert q.dtype == torch.float8_e4m3fn and tuple(q.shape) == (m, LATENT), where + assert s.dtype == torch.uint8 and tuple(s.shape) == (m, LATENT // SV), where + assert shared.dtype == torch.bfloat16 and tuple(shared.shape) == (m, SHARED_COLS), where + names = ("topk_ids", "topk_weights", "quantized", "scales") + for name, a, b in zip(names, got[:4], ref[:4]): + if not same(a, b): + unequal = int((bits(a) != bits(b)).sum()) if a.shape == b.shape else -1 + raise AssertionError(f"{where}: {name} differs from the reference in {unequal} bytes") + err = ls.rel_err(shared, ref[4]) + STATS["shared_err"] = max(STATS["shared_err"], err) + STATS["fronts_checked"] += 1 + assert err <= SHARED_TOL, f"{where}: shared activation rel err {err:.3e} > {SHARED_TOL}" + assert same_bits_on_ranks(ids, w, q, s), f"{where}: ranks disagree on the routing or the latent" + + +def runner(ids, w, q, s): + """The stock TRTLLM-Gen W4A8_MXFP4_MXFP8 MoE over this rank's experts, pre-routed (SiTU with Kimi K3's caps).""" + from tensorrt_llm._torch.moe.fused_moe.routing import RoutingMethodType + from tensorrt_llm._torch.utils import ActType_TrtllmGen + + p = EXPERT_BUFFERS + alpha = torch.full((E_LOCAL,), GATE_CAP, dtype=torch.float32, device="cuda") + beta = torch.full((E_LOCAL,), LINEAR_CAP, dtype=torch.float32, device="cuda") + return torch.ops.trtllm.mxe4m3_mxe2m1_block_scale_moe_runner( + None, None, q, s.view(-1), p["w31"], p["w31s"], None, alpha, beta, None, p["w2"], p["w2s"], None, EXPERTS, + TOP_K, 1, 1, I_TP, LATENT, I_TP, OFFSET, E_LOCAL, 1.0, int(RoutingMethodType.DeepSeekV3), + int(ActType_TrtllmGen.SiTu), topk_weights=w, topk_ids=ids, + ) # fmt: skip + + +def compare(y, ref): + """Op-catalog gates: |d| <= 8 ulp of the row's max |ref| per element, relative RMS <= 4 ulp; finite.""" + o, r = y.float(), ref.float() + row = r.abs().amax(dim=1, keepdim=True).clamp_min(1e-12) + elt = ((o - r).abs() / row).max().item() / ULP + rms = ((o - r).pow(2).mean().sqrt() / r.pow(2).mean().sqrt().clamp_min(1e-12)).item() / ULP + return elt, rms, bool(torch.isfinite(o).all()) and elt <= 8.0 and rms <= 4.0 + + +def verify_fused(got, call: Call, ws, where: str) -> None: + """k3_moe_fused_front's (y, shared) for ``call``. + + The front alone on the same input (made here, on ``ws``, by every rank) gives the routing and MXFP8 latent k3_moe + consumed: y within the op-catalog gates of the stock runner on them (all zeros on a rank no token routes to), and + shared the front's bits. + """ + y, shared = got + m = call.x.shape[0] + assert y.dtype == torch.bfloat16 and tuple(y.shape) == (m, LATENT) and y.is_contiguous(), where + ids, w, q, s, f_shared = call.as_kind("front").run(ws) + assert same(shared, f_shared), f"{where}: the shared activation differs from the front's" + if not bool(((ids >= OFFSET) & (ids < OFFSET + E_LOCAL)).any()): + assert bool((y == 0).all()), f"{where}: no token routes to this rank, y is not zero" + return + elt, rms, ok = compare(y, runner(ids, w, q, s)) + STATS["y_elt_ulp"] = max(STATS["y_elt_ulp"], elt) + STATS["y_rms_ulp"] = max(STATS["y_rms_ulp"], rms) + assert ok, f"{where}: y against the stock runner {elt:.2f} ulp per element, {rms:.2f} ulp RMS" + + +# ── state ───────────────────────────────────────────────────────────────── + + +def between_barriers(fn): + """fn() with every rank's kernels done before it and no rank's next call started until every rank has run it. + + Peers push into this rank's head buffers in their calls, so its words are read only there. + """ + R.barrier() + value = fn() + R.barrier() + return value + + +def workspace_rearmed(ws) -> bool: + """Every word of this rank's head buffers empty, flags[1] and the sign-in count flags[3] zero.""" + return between_barriers( + lambda: bool((ws.uc == EMPTY).all()) and int(ws.flags[1]) == 0 and int(ws.flags[3]) == 0 + ) + + +def epoch_and_words(ws): + """This rank's head epoch (flags[2]) and its 16 ready words.""" + return between_barriers(lambda: (int(ws.flags[2]), ws.ready[:16].tolist())) + + +def assert_ready_at(ws, want: int, where: str) -> None: + epoch, words = epoch_and_words(ws) + assert epoch == want and all(v == want for v in words), ( + f"{where}: epoch {epoch}, ready {words}, want {want}" + ) + + +def workspace_snapshot(ws): + return between_barriers(lambda: [ws.uc.clone(), ws.flags.clone(), ws.ready.clone()]) + + +def moe_rearmed() -> bool: + """Both K3MoeStates' slabs armed and every layer's counters zero (rank-local state).""" + torch.cuda.synchronize() + for st in (PLAIN, FLAGS): + mod = st.state.mod + cs = st.state.cs.view(mod.G_CAP, 8, mod.K2_TILES, mod.SFB_GROUP_BYTES) + if not (bool((st.state.c == -128).all()) and bool((cs[..., :4] == -1).all())): + return False + if not all(bool((layer.counters == 0).all()) for layer in st.layers): + return False + return True + + +def flags_call_checked(call: Call, ws, where: str): + """A head_flags call with its handoff checked on every rank, before and after. + + Before: no ready word the call polls (ids [t], MXFP8 row [8 + t], t < M) already holds epoch + 1; else k3_moe could + read the routing before the front writes it, and the call is not made. After: the epoch is epoch + 1 and every one + of the 16 ready words holds it. + """ + epoch, words = epoch_and_words(ws) + want = i32(epoch + 1) + m = call.x.shape[0] + early = [i for i in [*range(m), *range(8, 8 + m)] if words[i] == want] + assert R.all_true(not early), ( + f"{where}: ready words {early} already hold {want}; the call would race" + ) + got = call.run(ws) + torch.cuda.synchronize() + assert_ready_at(ws, want, where) + return got + + +def alone_results(calls, ws, where: str): + """Each call made alone (after a barrier) and checked; returns their outputs. + + The front against the unfused reference; the fused calls by verify_fused; head_flags calls also by + flags_call_checked. + """ + out = [] + for k, call in enumerate(calls): + tag = f"{where} {k} ({call.kind}, layer {call.layer}, M {call.x.shape[0]}) alone" + R.barrier() + if call.kind == "flags": + got = flags_call_checked(call, ws, tag) + else: + got = call.run(ws) + torch.cuda.synchronize() + if call.kind == "front": + verify_front(got, front_reference(call.x, call.layer), tag) + else: + verify_fused(got, call, ws, tag) + out.append(got) + return out + + +# ── checks ──────────────────────────────────────────────────────────────── + + +def check_workspace_is_armed_and_sized() -> None: + """After create(): 157,696 int32 words per rank behind one multicast mapping, every word empty. + + Flags and ready words zero; the front supports this W at the certified widths (checked at setup); the two + K3MoeStates are the plain and the head_flags build. + """ + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import k3_route_quant_ag as layout + + assert layout.workspace_words(R.world) == WORKSPACE_WORDS == 157_696 + for ws in (WS_A, WS_B): + assert ws.rank == R.rank and ws.world_size == R.world + assert ws.uc.dtype == ws.mc.dtype == ws.flags.dtype == ws.ready.dtype == torch.int32 + assert ws.uc.numel() == ws.mc.numel() == WORKSPACE_WORDS + assert tuple(ws.flags.shape) == (4,) and tuple(ws.ready.shape) == (32,) + assert between_barriers(lambda ws=ws: bool((ws.uc == EMPTY).all())), "every word empty" + assert ws.flags.tolist() == [0, 0, 0, 0] and not bool(ws.ready.any()), ( + "flags and ready words zero" + ) + assert not PLAIN.state.head_flags and FLAGS.state.head_flags + + +def check_front_single_calls() -> None: + """The front alone at every M 1-8 (the first M tokens of one 8-token batch) against the unfused reference. + + Each call's outputs the bits of the same rows of the 8-token call and run-to-run identical (routing and latent + also bit for bit on every rank, in verify_front); after each call this rank's head buffers empty, the sign-in + count zero, flags[0] flipped once, and the epoch and ready words untouched: the standalone front publishes nothing. + """ + base = Call(100, 8, layer=0, kind="front") + out8 = base.run(WS_A) + torch.cuda.synchronize() + verify_front(out8, front_reference(base.x, 0), "front M 8 (base)") + for m in TOKENS: + call = base.first(m) + epoch_words = epoch_and_words(WS_A) + flag0 = between_barriers(lambda: int(WS_A.flags[0])) + got = call.run(WS_A) + torch.cuda.synchronize() + flag1 = between_barriers(lambda: int(WS_A.flags[0])) + assert flag1 == flag0 ^ 1, f"front M {m}: buffer index {flag0} -> {flag1}" + assert workspace_rearmed(WS_A), f"front M {m}: head buffers not re-armed after the call" + verify_front(got, front_reference(call.x, 0), f"front M {m}") + again = [call.run(WS_A) for _ in range(2)] + torch.cuda.synchronize() + assert all(same_result(a, got) for a in again), f"front M {m}: not run-to-run identical" + assert same_result(got, [t[:m] for t in out8]), ( + f"front M {m}: rows differ from the 8-token call's" + ) + assert workspace_rearmed(WS_A), f"front M {m}: head buffers not re-armed after the re-runs" + assert epoch_and_words(WS_A) == epoch_words, ( + f"front M {m}: the standalone front moved the ready words" + ) + + +def check_fused_front_single_calls() -> None: + """k3_moe_fused_front at every M 1-8 on the plain K3MoeState and on the head_flags one. + + Layer 0's weights, the first M tokens of one batch. The plain y within the op-catalog gates of the stock runner + on the front's own routing and latent (all zeros on a rank no token routes to), the shared activation the front's + bits; the head_flags build's y and shared the plain build's bits, its handoff checked before and after + (flags_call_checked); the plain y's rows within one bf16 ulp of the same rows of the 8-token call (k3_moe's slice + FC2 groups a token's expert terms by the step's group count; bit-identity counted) and the shared rows their bits; + run-to-run identical bits; afterwards both states armed with every counter zero and the head buffers empty. + """ + base = Call(200, 8, layer=0, kind="fused") + y8, sh8 = base.run(WS_A) + torch.cuda.synchronize() + verify_fused((y8, sh8), base, WS_A, "fused M 8 (base)") + for m in TOKENS: + plain = base.first(m) + got = plain.run(WS_A) + torch.cuda.synchronize() + verify_fused(got, plain, WS_A, f"fused M {m}") + y, sh = got + y_flags, sh_flags = flags_call_checked(plain.as_kind("flags"), WS_A, f"head_flags M {m}") + assert same(y_flags, y) and same(sh_flags, sh), ( + f"M {m}: the head_flags build differs from the plain one" + ) + ulp = max_ulp(y, y8[:m]) + STATS["rows_checked"] += 1 + STATS["rows_bits_as_m8"] += int(same(y, y8[:m])) + assert ulp <= 1, f"fused M {m}: rows {ulp} ulp from the 8-token call's" + assert same(sh, sh8[:m]), f"fused M {m}: shared rows differ from the 8-token call's" + rerun = plain.run(WS_A) + torch.cuda.synchronize() + assert same_result(rerun, got), f"fused M {m}: not run-to-run identical" + assert moe_rearmed(), f"fused M {m}: a k3_moe slab is not armed or a counter is not zero" + assert workspace_rearmed(WS_A), f"fused M {m}: head buffers not re-armed" + + +def check_head_flags_epoch_wraps() -> None: + """The ready-word handoff across both int32 wraps of the epoch (check_head_flags of the kernel test, extended). + + From a new workspace's state (ready words zeroed, epoch 0): calls at M 1, 1; the epoch preset to -2 (as 2^32 - 2 + calls later): calls at M 1, 8, 3, 8, so that the M 8 call at epoch -1 waits for 0, the value a word past an earlier + call's tokens would still hold had k3_moe not re-armed it; then the epoch preset to 2^31 - 2: calls at M 8, 1, 8, + across 2^31 - 1 -> -2^31. Every call checked by flags_call_checked, its y and shared the plain build's bits; the + head buffers empty afterwards. + """ + base = Call(300, 8, layer=0, kind="fused") + plain = {m: base.first(m).run(WS_A) for m in (1, 3, 8)} + torch.cuda.synchronize() + + def preset(epoch, zero_words=False): + R.barrier() + if zero_words: + WS_A.ready.zero_() + WS_A.flags[2] = epoch + R.barrier() + + preset(0, zero_words=True) + for item in (1, 1, ("epoch", -2), 1, 8, 3, 8, ("epoch", 2**31 - 2), 8, 1, 8): + if isinstance(item, tuple): + preset(item[1]) + continue + epoch = epoch_and_words(WS_A)[0] + got = flags_call_checked(base.first(item, "flags"), WS_A, f"epoch {epoch} M {item}") + assert same_result(got, plain[item]), ( + f"epoch {epoch} M {item}: differs from the plain build" + ) + assert workspace_rearmed(WS_A) + + +def check_dip_and_regrow_sequence() -> None: + """Decode steps of three layers on one workspace, the token count dipping and growing back, a random rank late. + + Each step: k3_moe_fused_front on the plain state (layer 0's weights), on the head_flags state (layer 1's) and the + front alone (layer 2's), at M 8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8 with new inputs every call, the layers back to back + with a random rank 5 ms late before every call. Every call returns the bits of the same call made alone (each + checked first). Afterwards the head buffers are empty, the epoch advanced once per head_flags call with every + ready word at it, and both states armed: a call after a smaller one reads nothing an older, larger call left. + """ + calls = [ + Call(4000 + 10 * s + layer, m, layer=layer, kind=KINDS[layer]) + for s, m in enumerate(DIP_STEPS) + for layer in range(LAYERS) + ] + alone = alone_results(calls, WS_A, "dip") + epoch0 = epoch_and_words(WS_A)[0] + late = random.Random(7) + got = [] + for s in range(len(DIP_STEPS)): + R.barrier() + for layer in range(LAYERS): + R.late(late.randrange(R.world)) + got.append(calls[s * LAYERS + layer].run(WS_A)) + torch.cuda.synchronize() + bad = [k for k, (g, a) in enumerate(zip(got, alone)) if not same_result(g, a)] + assert not bad, f"calls {bad} of the sequence differ from the same calls alone" + assert workspace_rearmed(WS_A), "head buffers not re-armed after the sequence" + assert_ready_at(WS_A, i32(epoch0 + len(DIP_STEPS)), "after the sequence") + assert moe_rearmed(), "a k3_moe slab is not armed or a counter is not zero after the sequence" + + +def check_two_workspaces_interleaved() -> None: + """Two workspaces are two rotations and two epochs. + + 20 calls alternate between WS_A and WS_B in an irregular pattern (A A B A B B A A A B, twice), kinds and layers + cycling, M in 3, 8, 1, 8, 5, back to back: every call returns the bits of the same call alone; each workspace's + epoch advanced by its own head_flags calls only, every ready word at it; both workspaces' buffers empty. The + pattern is the same on every rank: calls on one stream run one after the other and each waits for its peers, so + ranks ordering calls on two workspaces differently would deadlock (not exercised). + """ + pattern = "AABABBAAAB" * 2 + calls = [ + Call(5000 + i, (3, 8, 1, 8, 5)[i % 5], layer=i % LAYERS, kind=KINDS[i % LAYERS]) + for i in range(len(pattern)) + ] + alone = alone_results(calls, WS_A, "interleaved") + epoch_a, epoch_b = epoch_and_words(WS_A)[0], epoch_and_words(WS_B)[0] + R.barrier() + got = [c.run(WS_A if which == "A" else WS_B) for c, which in zip(calls, pattern)] + torch.cuda.synchronize() + bad = [k for k, (g, a) in enumerate(zip(got, alone)) if not same_result(g, a)] + assert not bad, f"interleaved calls {bad} differ from the same calls alone" + flag_calls = { + w: sum(c.kind == "flags" and p == w for c, p in zip(calls, pattern)) for w in "AB" + } + assert_ready_at(WS_A, i32(epoch_a + flag_calls["A"]), "WS_A after the interleaving") + assert_ready_at(WS_B, i32(epoch_b + flag_calls["B"]), "WS_B after the interleaving") + assert workspace_rearmed(WS_A) and workspace_rearmed(WS_B) + assert moe_rearmed() + + +def check_graph_capture_and_replay() -> None: + """A captured step replayed with rewritten inputs, eager calls of other token counts between replays. + + The step: the three layers at M 8 on WS_B (k3_moe_fused_front plain and head_flags, the front alone), captured + once and replayed 6 times with new inputs copied into its static buffers; between replays an eager call of + another M on WS_B. Every replayed and eager call returns the bits of the same call alone; afterwards the head + buffers are empty, WS_B's epoch advanced once per head_flags call, replayed or eager, with every ready word at it, + and both states armed. + """ + statics = [torch.zeros(8, HIDDEN, dtype=torch.bfloat16, device="cuda") for _ in range(LAYERS)] + shells = [ + Call(0, 0, layer=layer, kind=KINDS[layer], x=statics[layer]) for layer in range(LAYERS) + ] + + def step(): + return [shell.run(WS_B) for shell in shells] + + warm = [Call(6000 + layer, 8, layer, KINDS[layer]) for layer in range(LAYERS)] + reps = [ + [Call(6100 + 10 * r + layer, 8, layer, KINDS[layer]) for layer in range(LAYERS)] + for r in range(REPLAYS) + ] + eagers = [ + Call(6500 + r, (3, 1, 6, 5, 2, 7)[r], r % LAYERS, KINDS[r % LAYERS]) for r in range(REPLAYS) + ] + alone_reps = [alone_results(rep, WS_A, f"replay {r}") for r, rep in enumerate(reps)] + alone_eager = alone_results(eagers, WS_A, "eager between replays") + + epoch0 = epoch_and_words(WS_B)[0] + for static, c in zip(statics, warm): + static.copy_(c.x) + R.barrier() + step() # every first call of this step eager + R.barrier() + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + outs = step() + R.barrier() + flag_calls = 1 # the warm-up step's + for r in range(REPLAYS): + for static, c in zip(statics, reps[r]): + static.copy_(c.x) + R.barrier() + graph.replay() + torch.cuda.synchronize() + bad = [ + layer for layer in range(LAYERS) if not same_result(outs[layer], alone_reps[r][layer]) + ] + assert not bad, f"replay {r}: layers {bad} differ from the same calls alone" + got = eagers[r].run(WS_B) + torch.cuda.synchronize() + assert same_result(got, alone_eager[r]), ( + f"eager call after replay {r} differs from the same call alone" + ) + flag_calls += 1 + int(eagers[r].kind == "flags") + del graph + assert workspace_rearmed(WS_B), "head buffers not re-armed after the replays" + assert_ready_at(WS_B, i32(epoch0 + flag_calls), "WS_B after the replays") + assert moe_rearmed() + + +def check_unsupported_shape_raises_on_every_rank() -> None: + """M 0 and 9 raise ValueError on every rank before touching anything, and the next call is correct. + + The front alone and k3_moe_fused_front (plain and head_flags) at M 0 and 9: the head workspace's words, flags and + ready words keep their bits; the next call returns the bits of the same call made alone. + """ + nxt = Call(7000, 8, layer=0, kind="fused") + want = alone_results([nxt], WS_A, "before the unsupported calls")[0] + before = workspace_snapshot(WS_A) + raised = [] + for m in (0, 9): + for kind in KINDS: + try: + Call(7100 + m, m, layer=0, kind=kind).run(WS_A) + raised.append(False) + except ValueError: + raised.append(True) + assert R.all_true(all(raised)), ( + f"M 0 / 9 did not raise ValueError on every rank and kind: {raised}" + ) + after = workspace_snapshot(WS_A) + assert all(same(a, b) for a, b in zip(before, after)), ( + "a refused call touched the head workspace" + ) + R.barrier() + got = nxt.run(WS_A) + torch.cuda.synchronize() + assert same_result(got, want), ( + "the call after the refused ones differs from the same call alone" + ) + + +def check_create_and_first_compile_refuse_capture() -> None: + """Under CUDA-graph capture on every rank, create() and a not yet compiled front configuration raise. + + K3MoeHeadWorkspace.create raises RuntimeError before entering the collective (no rank waits for another), and a + front call of a configuration not yet compiled (other SiTU caps) raises RuntimeError before any launch. The + workspace keeps its bits and the next call returns the bits of the same call made alone. + """ + nxt = Call(7500, 4, layer=1, kind="front") + want = alone_results([nxt], WS_A, "before the capture refusals")[0] + before = workspace_snapshot(WS_A) + messages = [] + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + try: + OPS.workspace.create(R.mapping, fabric_handle=R.fabric) + except RuntimeError as exc: + messages.append(str(exc)) + try: + OPS.front( + nxt.x, + LAYER_WEIGHTS[1].front, + BIAS, + RSF, + SHARED_COLS, + GATE_CAP + 1.0, + LINEAR_CAP, + WS_A, + ) + except RuntimeError as exc: + messages.append(str(exc)) + del graph + refused = len(messages) == 2 and all("outside CUDA-graph capture" in msg for msg in messages) + assert R.all_true(refused), f"under capture: {messages}" + after = workspace_snapshot(WS_A) + assert all(same(a, b) for a, b in zip(before, after)), ( + "a refused call touched the head workspace" + ) + R.barrier() + got = nxt.run(WS_A) + torch.cuda.synchronize() + assert same_result(got, want), ( + "the call after the capture refusals differs from the same call alone" + ) + + +def check_wrong_call_order_is_detected() -> None: + """Negative control: rank 0 swaps two same-shaped front calls on one workspace. + + Rank 0 runs c2 then c1, its peers c1 then c2. Nothing raises or hangs (the rotation positions still agree), but + each call pairs with the peers' call at the same position: every rank's gathered head mixes rank 0's slice of the + other input with the peers' slices of this one. Every rank returns the same wrong routing and latent, exactly the + unfused reference of those mixed inputs; the MXFP8 latent's columns [0, 3584 / W) and their scales are bit for bit + those of rank 0's input's call and the rest those of the peers' input's call; each rank's latent differs from that + of the call it made in more than half of the mixed-in columns. The shared activation (rank-local) is each rank's + own input's, bit for bit. A plain call right after is correct again. + """ + c1, c2 = Call(8000, 8, layer=0, kind="front"), Call(8001, 8, layer=0, kind="front") + o1, o2 = alone_results([c1, c2], WS_A, "the ordered pair") + R.barrier() + if R.rank == 0: + first, second = c2.run(WS_A), c1.run(WS_A) + else: + first, second = c1.run(WS_A), c2.run(WS_A) + torch.cuda.synchronize() + sc = WL // SV + for pos, got, (x0, xp), (o0, op) in ( + (1, first, (c2.x, c1.x), (o2, o1)), + (2, second, (c1.x, c2.x), (o1, o2)), + ): + where = f"swapped pair, position {pos}" + xs = [x0] + [xp] * (R.world - 1) # each rank's input at this position + verify_front(got, front_reference(xp, 0, xs=xs), where) + q, s = got[2], got[3] + assert same(q[:, :WL], o0[2][:, :WL]) and same(q[:, WL:], op[2][:, WL:]), ( + f"{where}: latent not the mix" + ) + assert same(s[:, :sc], o0[3][:, :sc]) and same(s[:, sc:], op[3][:, sc:]), ( + f"{where}: scales not the mix" + ) + # The call this rank made at this position, correctly ordered, and the columns the swap mixed into it. + made = o0 if R.rank == 0 else op + mixed_in = slice(WL, LATENT) if R.rank == 0 else slice(0, WL) + wrong = ( + (q[:, mixed_in].view(torch.uint8) != made[2][:, mixed_in].view(torch.uint8)) + .float() + .mean() + .item() + ) + assert wrong > 0.5, ( + f"{where}: only {wrong:.3f} of the mixed-in latent differs from the call made" + ) + assert same(got[4], made[4]), ( + f"{where}: the shared activation is not this rank's own input's" + ) + if R.rank == 0 and pos == 1: + print( + f"[rank 0] swapped pair: {wrong:.3f} of the mixed-in latent codes wrong", flush=True + ) + c3 = Call(8002, 8, layer=0, kind="front") + R.barrier() + got = c3.run(WS_A) + torch.cuda.synchronize() + verify_front(got, front_reference(c3.x, 0), "after the swapped pair") + + +CHECKS = [ + check_workspace_is_armed_and_sized, + check_front_single_calls, + check_fused_front_single_calls, + check_head_flags_epoch_wraps, + check_dip_and_regrow_sequence, + check_two_workspaces_interleaved, + check_graph_capture_and_replay, + check_unsupported_shape_raises_on_every_rank, + check_create_and_first_compile_refuse_capture, + # Stays last: it deliberately disagrees on call order. + check_wrong_call_order_is_detected, +] + + +# ── setup ───────────────────────────────────────────────────────────────── + + +def make_experts(seed: int): + """224 random MXFP4 experts through TRT-LLM's W4A8_MXFP4_MXFP8 TRTLLM-Gen loader: this rank's buffers.""" + from tensorrt_llm._torch.moe.fused_moe.quantization import W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod + + def rand_mxfp4(rows, k, gen): + codes = torch.randint(0, 16, (rows, k), dtype=torch.uint8, device="cuda", generator=gen) + base = 127 + round(0.5 * math.log2(0.01057 / k)) + exps = torch.randint( + base, base + 6, (rows, k // SV), dtype=torch.uint8, device="cuda", generator=gen + ) + return (codes[:, 0::2] | (codes[:, 1::2] << 4)).contiguous(), exps + + i_full = I_TP * MOE_TP + method = W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod() + module = SimpleNamespace( + tp_size=MOE_TP, + tp_rank=1, + scaling_vector_size=SV, + intermediate_size=i_full, + intermediate_size_per_partition=I_TP, + hidden_size=LATENT, + ) + kw = dict(dtype=torch.uint8, device="cuda") + proc = dict( + w31=torch.empty(E_LOCAL, 2 * I_TP, LATENT // 2, **kw), + w31s=torch.empty(E_LOCAL, 2 * I_TP, LATENT // SV, **kw), + w2=torch.empty(E_LOCAL, LATENT, I_TP // 2, **kw), + w2s=torch.empty(E_LOCAL, LATENT, I_TP // SV, **kw), + ) + gen = torch.Generator(device="cuda").manual_seed(seed) + for e in range(E_LOCAL): + w1, w1s = rand_mxfp4(i_full, LATENT, gen) + w3, w3s = rand_mxfp4(i_full, LATENT, gen) + w2, w2s = rand_mxfp4(LATENT, i_full, gen) + method.load_expert_w3_w1_weight(module, w1, w3, proc["w31"][e]) + method.load_expert_w2_weight(module, w2, proc["w2"][e]) + method.load_expert_w3_w1_weight_scale_mxfp4(module, w1s, w3s, proc["w31s"][e]) + method.load_expert_w2_weight_scale_mxfp4(module, w2s, proc["w2s"][e]) + torch.cuda.synchronize() + return proc + + +def make_layer_weights(layer: int, front_weight): + """Layer ``layer``'s front weights, exact payloads (see the module docstring). + + Every rank's head slice: its latent-down rows (multiples of 1/16), then its router rows (multiples of 1/8, logits + of a few units); one seed per rank, so each rank holds them all for the reference. This rank's shared gate_up + (multiples of 1/16; gate rows, then up rows), and this rank's front weight built from them. + """ + heads = [] + for r in range(R.world): + g = torch.Generator(device="cuda").manual_seed(1000 * (layer + 1) + r) + latent_rows = ls.exact_bf16(g, (WL, HIDDEN), -2, 3, 1 / 16) + router_rows = ls.exact_bf16(g, (WE, HIDDEN), -2, 3, 1 / 8) + heads.append(torch.cat([latent_rows, router_rows]).contiguous()) + g = torch.Generator(device="cuda").manual_seed(50_000 + 1000 * layer + R.rank) + gate_up = ls.exact_bf16(g, (2 * SHARED_COLS, HIDDEN), -2, 3, 1 / 16) + return SimpleNamespace(heads=heads, gate_up=gate_up, front=front_weight(heads[R.rank], gate_up)) + + +def make_state(head_flags: bool): + """A K3MoeState on this device and one layer per front layer, all over this rank's experts (own counters each).""" + device = torch.device("cuda", torch.cuda.current_device()) + state = OPS.K3MoeState(device, I_TP, E_LOCAL, head_flags=head_flags) + p = EXPERT_BUFFERS + layers = [state.layer(p["w31"], p["w31s"], p["w2"], p["w2s"]) for _ in range(LAYERS)] + return SimpleNamespace(state=state, layers=layers) + + +def _run_one_rank(args) -> int: + global R, OPS, WS_A, WS_B, PLAIN, FLAGS, LAYER_WEIGHTS, BIAS, EXPERT_BUFFERS, OFFSET, WL, WE + R = ls.Rank(args) + assert torch.cuda.get_device_capability() == (10, 0), "k3_moe_front and k3_moe need sm_100" + import tensorrt_llm._torch.custom_ops # noqa: F401 -- registers the stock path's ops + from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe import k3_moe as moe + from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe import k3_moe_front as front + + OPS = SimpleNamespace( + front=front.k3_moe_front, + fused=moe.k3_moe_fused_front, + workspace=front.K3MoeHeadWorkspace, + K3MoeState=moe.K3MoeState, + ) + WL, WE = LATENT // R.world, EXPERTS // R.world + OFFSET = (R.rank % MOE_TP) * E_LOCAL + with torch.inference_mode(): + device = torch.device("cuda", torch.cuda.current_device()) + assert front.weight_supported(R.world, SHARED_COLS, HIDDEN, device), ( + f"k3_moe_front does not support W {R.world} with {SHARED_COLS} shared columns on this device" + ) + g = torch.Generator(device="cuda").manual_seed(11) + BIAS = (torch.randn(EXPERTS, generator=g, device="cuda") * 0.05).float() + LAYER_WEIGHTS = [make_layer_weights(layer, front.front_weight) for layer in range(LAYERS)] + EXPERT_BUFFERS = make_experts(20260928 + R.rank) + WS_A = OPS.workspace.create(R.mapping, fabric_handle=R.fabric) + WS_B = OPS.workspace.create(R.mapping, fabric_handle=R.fabric) + PLAIN = make_state(head_flags=False) + FLAGS = make_state(head_flags=True) + code = ls.run_checks(R, CHECKS) + if R.rank == 0: + s = STATS + print( + f"[rank 0] world {R.world}; {s['fronts_checked']} front calls bit for bit the reference's routing and " + f"latent; shared max rel err {s['shared_err']:.3e}; y vs stock runner max {s['y_elt_ulp']:.2f} ulp " + f"per element, {s['y_rms_ulp']:.2f} ulp RMS; fused rows bit-identical to the 8-token call's " + f"{s['rows_bits_as_m8']}/{s['rows_checked']}", + flush=True, + ) + return code + + +if __name__ == "__main__": + ARGS = ls.parse_args(sys.argv[1:]) + if ARGS.rank_worker or ARGS.launcher == "srun": + sys.exit(_run_one_rank(ARGS)) + ls.spawn(__file__, ARGS, DEADLINE_S) + print("OK") diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_common.py b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_common.py new file mode 100644 index 000000000000..5f3dca684c54 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_common.py @@ -0,0 +1,548 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Shared pieces of the Kimi K3 sandwich matrices (``_k3_sandwich_oproj_op_matrix.py``, +``_k3_sandwich_tail_op_matrix.py``, ``_k3_sandwich_plain_op_matrix.py``): one class per op holding one call's +arguments on every rank and its native-torch reference, a driver for call sequences (chained the way the model chains +its residual streams; eager with a random rank late, or captured and replayed), and the A1 checks the three matrices +run the same way on a ``K3SandwichWorkspace``. + +Started by file path like ``_lockstep`` (this tree is not a package). Importing it pulls in torch only; ``bind`` +(called by the rank body) imports the catalog wrappers. + +Payloads and references. Every rank draws every rank's inputs from one seed, so each rank holds the whole reference. +The projections' inputs are small multiples of powers of two (``_lockstep.exact_bf16``), so every partial product +sum is exact in fp32 in whatever order the kernel accumulates: the reference is the fp64 product rounded once to bf16 +per rank (the kernel rounds its fp32 accumulator once), the ranks' partials added in fp32 and rounded to bf16, then +the prefix sum (or residual) added in fp32 and rounded -- the kernel's own steps, so ``updated`` of +``k3_sandwich_oproj`` and ``k3_sandwich_plain`` is compared bit for bit. ``k3_sandwich_tail`` scales its latent +accumulator by the latent row's RMS, an fp32 rsqrt that is not torch's, so a partial element can round to the +neighbouring bf16: its ``updated`` is compared within ``TOL_TAIL`` of an fp64 reference of the kernel's arithmetic. +``normed`` and the tail's tap (softmax, rsqrt) are compared within ``TOL`` of an fp32 reference. Every output is also +compared bitwise across the ranks. +""" + +from __future__ import annotations + +import copy +import random +from typing import Callable, Dict, List, Optional, Sequence, Tuple + +import _lockstep as ls +import torch + +H = 7168 +K_O = 768 # core / o_weight columns at TP16: 96 heads x 128 / 16 +LATENT = 3584 # the reduced latent row +LAT_SLICE = 224 # a rank's latent columns at TP16 +LAT_PAD = 256 # the latent slice zero-padded to whole k-tiles inside tail_weight +SLICES = LATENT // LAT_SLICE # 16 +ACT = 384 # the shared-expert activation at TP16 +PLAIN_K = 384 # the drafter's o_proj slice at TP16 +DOWN_K = 896 # the drafter's down projection slice at TP16; its gate_up output is [M, 2 x 896] +NUM_CTAS = 56 # the kernel's CTAs: one call counter each, in 64 flag words +FLAG_WORDS = 64 +RMS_EPS = 1e-6 +OUT_EPS = 1e-6 +LAT_EPS = 1e-6 +EPS = 1e-6 +TOL = 2e-2 # normed and the tap: max |err| / max |ref| +TOL_TAIL = 8e-3 # the tail's updated: max |err| / max |ref|, about one bf16 ulp of its largest elements +GATES = (0.0, 32.0, 64.0) # swiglu gate values on which silu is exact in fp32 (see PlainCall) +WEIGHT_SETS = 3 +TOKENS = (1, 2, 3, 4, 5, 6, 7, 8) +DIP_STEPS = (8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8) +INTERLEAVE = "AABABBAAAB" * 2 +REPLAYS = 8 +WEIGHT_SEEDS = {"oproj": 900, "tail": 910, "plain": 920, "down": 930} + +R = None # this rank (a _lockstep.Rank), set by bind() +OPS: Dict[str, Callable] = {} +WEIGHTS: Dict[str, List[List[torch.Tensor]]] = {} +STATS = {"normed": 0.0, "tail_updated": 0.0, "tap": 0.0} + + +def bind(rank: ls.Rank, weight_sets: Dict[str, int]) -> None: + """Bind this rank and the three catalog wrappers, and draw ``weight_sets[kind]`` weight sets of each kind ("oproj", + "tail", "plain", "down"), every rank's slice, from fixed seeds.""" + global R + from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm import ( + k3_sandwich_oproj, + k3_sandwich_plain, + k3_sandwich_tail, + ) + + R = rank + OPS["oproj"] = k3_sandwich_oproj.k3_sandwich_oproj + OPS["tail"] = k3_sandwich_tail.k3_sandwich_tail + OPS["plain"] = k3_sandwich_plain.k3_sandwich_plain + for kind, sets in weight_sets.items(): + WEIGHTS[kind] = [] + for s in range(sets): + g = _gen(WEIGHT_SEEDS[kind] + s) + WEIGHTS[kind].append([_weight(kind, g) for _ in range(R.world)]) + + +def report() -> None: + if R.rank == 0: + print(f"[rank 0] world {R.world}; max rel err: normed {STATS['normed']:.3e}, tail updated " + f"{STATS['tail_updated']:.3e}, tap {STATS['tap']:.3e}", flush=True) # fmt: skip + + +def _gen(seed: int) -> torch.Generator: + return torch.Generator(device="cuda").manual_seed(seed) + + +def _weight(kind: str, g: torch.Generator) -> torch.Tensor: + """One rank's weight slice, entries in {-2, ..., 2} / 16; tail_weight = [latent up (224) | zeros (32) | shared + down (384)].""" + if kind == "tail": + lat = ls.exact_bf16(g, (H, LAT_SLICE), -2, 3, 1 / 16) + pad = torch.zeros(H, LAT_PAD - LAT_SLICE, dtype=torch.bfloat16, device="cuda") + act = ls.exact_bf16(g, (H, ACT), -2, 3, 1 / 16) + return torch.cat([lat, pad, act], dim=1).contiguous() + k = {"oproj": K_O, "plain": PLAIN_K, "down": DOWN_K}[kind] + return ls.exact_bf16(g, (H, k), -2, 3, 1 / 16) + + +def _running_sum(g: torch.Generator, tokens: int) -> torch.Tensor: + """A drawn prefix sum or residual: {-32, ..., 32} / 16.""" + return ls.exact_bf16(g, (tokens, H), -32, 33, 1 / 16) + + +def _attn_res_inputs(g: torch.Generator, snapshots: int, tokens: int): + """``block_residual`` [S, M, 7168] and the epilogue's three [7168] weights.""" + block = torch.randn(snapshots, tokens, H, generator=g, device="cuda").bfloat16() + res_w = (torch.randn(H, generator=g, device="cuda") * 0.05).bfloat16() + rms_w = (1.0 + 0.1 * torch.randn(H, generator=g, device="cuda")).bfloat16() + out_w = (1.0 + 0.1 * torch.randn(H, generator=g, device="cuda")).bfloat16() + return block, res_w, rms_w, out_w + + +def _gates(g: torch.Generator, shape) -> torch.Tensor: + idx = torch.randint(0, len(GATES), tuple(shape), generator=g, device="cuda") + return torch.tensor(GATES, device="cuda")[idx].bfloat16() + + +def bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def bf16_partial(x: torch.Tensor, w: torch.Tensor) -> torch.Tensor: + """bf16(x @ w^T) from the exact fp64 product: the kernel's one rounding of an fp32 accumulator that is exact for + these payloads.""" + return (x.double() @ w.double().t()).float().bfloat16() + + +def reduce_ref(partials: Sequence[torch.Tensor], carry: Optional[torch.Tensor]) -> torch.Tensor: + """bf16(carry + bf16(sum of the partials)): the partials added in fp32 in rank order (the kernel's order up to 8 + ranks; exact sums make any order the same), then the carry -- the prefix sum or the residual -- added in fp32; + the rounded sum alone without a carry.""" + total = partials[0].float() + for p in partials[1:]: + total = total + p.float() + updated = total.bfloat16() + if carry is not None: + updated = (carry.float() + updated.float()).bfloat16() + return updated + + +def attn_res_mixture(updated, block, res_w, rms_w) -> torch.Tensor: + """The attention-residual mixture of [block..., updated] before the output RMSNorm, in fp32 and rounded to bf16: + the tensor ``_lockstep.residual_update_ref`` normalizes, and the tail's tap.""" + v = torch.cat([block, updated.unsqueeze(0)], dim=0).float() + rs = (v.square().mean(dim=-1) + RMS_EPS).rsqrt() + logits = (v * rs[..., None] * (rms_w.float() * res_w.float())).sum(dim=-1) + probs = torch.softmax(logits, dim=0) + return (probs[..., None] * v).sum(dim=0).bfloat16() + + +def _within(got: torch.Tensor, want: torch.Tensor, tol: float, key: str, where: str, what: str) -> None: + err = ls.rel_err(got, want) + STATS[key] = max(STATS[key], err) + assert err <= tol, f"{where}: {what} rel err {err:.3e} > {tol}" + + +class Call: + """One call's arguments on every rank and its reference. ``carry`` is the running sum the call adds to: the prefix + sum of oproj / tail (None: the sum alone) or the residual of plain. Subclasses define ``ref``, ``run`` and + ``refill``.""" + + tokens: int + carry: Optional[torch.Tensor] + + def ref(self) -> Tuple[torch.Tensor, torch.Tensor]: + raise NotImplementedError + + def run(self, ws) -> Tuple[torch.Tensor, torch.Tensor]: + raise NotImplementedError + + def refill(self, fresh: "Call", carry: bool) -> None: + """Copy ``fresh``'s inputs into this call's tensors in place (a captured call reads them there), the carry too + when ``carry`` (a chain's first call; later calls read the previous call's output). Weights and scalar + arguments stay: the graph holds them.""" + raise NotImplementedError + + def verify(self, got, where: str) -> None: + normed, updated = got + want_normed, want_updated = self.ref() + assert torch.equal(updated, want_updated), f"{where}: updated differs from the exact reference" + _within(normed, want_normed, TOL, "normed", where, "normed") + assert R.same_on_ranks(normed, updated), f"{where}: ranks disagree" + + def wrong_fraction(self, got) -> float: + """The fraction of ``updated``'s elements that differ from the reference.""" + return (got[1] != self.ref()[1]).float().mean().item() + + def wrong_margin(self, got) -> float: + """``updated``'s largest error over the comparison's bound: any difference fails an exact comparison.""" + return float("inf") if self.wrong_fraction(got) > 0 else 0.0 + + +class OprojCall(Call): + """k3_sandwich_oproj: core {-2, ..., 2} / 8 [M, 768] per rank, o_weight from weight set ``weights``. ``prefix``: + True (drawn), None, or a tensor.""" + + def __init__(self, seed, tokens, snapshots, prefix=True, weights=0): + g = _gen(seed) + self.tokens = tokens + self.cores = [ls.exact_bf16(g, (tokens, K_O), -2, 3, 1 / 8) for _ in range(R.world)] + self.carry = _running_sum(g, tokens) if prefix is True else prefix + self.weights = WEIGHTS["oproj"][weights] + self.block, self.res_w, self.rms_w, self.out_w = _attn_res_inputs(g, snapshots, tokens) + + def ref(self): + updated = reduce_ref([bf16_partial(c, w) for c, w in zip(self.cores, self.weights)], self.carry) + normed = ls.residual_update_ref( + updated, self.block, self.res_w, self.rms_w, RMS_EPS, self.out_w, OUT_EPS + ) + return normed, updated + + def run(self, ws): + r = R.rank + return OPS["oproj"](self.cores[r], self.weights[r], self.carry, self.block, self.res_w, self.rms_w, + self.out_w, RMS_EPS, OUT_EPS, ws) # fmt: skip + + def refill(self, fresh, carry): + for mine, new in zip(self.cores, fresh.cores): + mine.copy_(new) + self.block.copy_(fresh.block) + if carry: + self.carry.copy_(fresh.carry) + + +class TailCall(Call): + """k3_sandwich_tail: the reduced latent {-16, ..., 16} / 8 [M, 3584] (the same on every rank), act + {-2, ..., 2} / 8 [M, 384] per rank, rank r's latent slice at lo = ((5 r + shift) % 16) x 224 (``shift``: the seed + unless given), tail_weight from weight set ``weights``. ``tap``: None, "mix" (the pre-norm mixture into columns + [2 H, 3 H) of a NaN-filled [M, 5 H] capture buffer) or "updated" (``updated`` into its columns [4 H, 5 H)); + ``updated_out``: store ``updated`` into row 1 of a NaN-filled [3, M, H] bank.""" + + def __init__(self, seed, tokens, snapshots, prefix=True, weights=0, shift=None, tap=None, updated_out=False): + g = _gen(seed) + self.tokens = tokens + self.latent = ls.exact_bf16(g, (tokens, LATENT), -16, 17, 1 / 8) + self.acts = [ls.exact_bf16(g, (tokens, ACT), -2, 3, 1 / 8) for _ in range(R.world)] + shift = seed if shift is None else shift + self.los = [((5 * r + shift) % SLICES) * LAT_SLICE for r in range(R.world)] + self.carry = _running_sum(g, tokens) if prefix is True else prefix + self.weights = WEIGHTS["tail"][weights] + self.block, self.res_w, self.rms_w, self.out_w = _attn_res_inputs(g, snapshots, tokens) + self._set_options(tap, updated_out) + + def _set_options(self, tap, updated_out) -> None: + self.tap_kind, self.cap, self.tap = tap, None, None + if tap is not None: + self.cap = torch.full((self.tokens, 5 * H), float("nan"), dtype=torch.bfloat16, device="cuda") + self.tap = self.cap[:, self._tap_col() * H : (self._tap_col() + 1) * H] + self.bank, self.updated_out = None, None + if updated_out: + self.bank = torch.full((3, self.tokens, H), float("nan"), dtype=torch.bfloat16, device="cuda") + self.updated_out = self.bank[1] + + def _tap_col(self) -> int: + return 2 if self.tap_kind == "mix" else 4 + + def with_options(self, tap=None, updated_out=False) -> "TailCall": + """The same call -- the same input tensors -- with other output options.""" + other = copy.copy(self) + other._set_options(tap, updated_out) + return other + + def partials(self) -> List[torch.Tensor]: + """Every rank's partial in fp64: scale_t (acc_lat) + acc_act, one rounding to bf16; the padding columns of + tail_weight are zero, so the latent slice is the 224 columns from lo.""" + lat = self.latent.double() + scale = (lat.square().mean(dim=1, keepdim=True) + LAT_EPS).rsqrt() + out = [] + for act, w, lo in zip(self.acts, self.weights, self.los): + w64 = w.double() + acc_lat = lat[:, lo : lo + LAT_SLICE] @ w64[:, :LAT_SLICE].t() + acc_act = act.double() @ w64[:, LAT_PAD:].t() + out.append((acc_lat * scale + acc_act).float().bfloat16()) + return out + + def ref(self): + updated = reduce_ref(self.partials(), self.carry) + normed = ls.residual_update_ref( + updated, self.block, self.res_w, self.rms_w, RMS_EPS, self.out_w, OUT_EPS + ) + return normed, updated + + def run(self, ws): + r = R.rank + return OPS["tail"](self.latent, self.acts[r], self.weights[r], self.los[r], LAT_EPS, self.carry, self.block, + self.res_w, self.rms_w, self.out_w, RMS_EPS, OUT_EPS, ws, tap=self.tap, + tap_updated=self.tap_kind == "updated", updated_out=self.updated_out) # fmt: skip + + def refill(self, fresh, carry): + self.latent.copy_(fresh.latent) + for mine, new in zip(self.acts, fresh.acts): + mine.copy_(new) + self.block.copy_(fresh.block) + if carry: + self.carry.copy_(fresh.carry) + + def verify(self, got, where): + normed, updated = got + want_normed, want_updated = self.ref() + _within(updated, want_updated, TOL_TAIL, "tail_updated", where, "updated") + _within(normed, want_normed, TOL, "normed", where, "normed") + outputs = [normed, updated] + if self.updated_out is not None: + assert updated.data_ptr() == self.updated_out.data_ptr(), f"{where}: updated is not updated_out" + assert bool(torch.isnan(self.bank[0::2].float()).all()), f"{where}: another bank row was written" + if self.tap is not None: + col = self._tap_col() + rest = torch.cat([self.cap[:, : col * H], self.cap[:, (col + 1) * H :]], dim=1) + assert bool(torch.isnan(rest.float()).all()), f"{where}: the capture buffer was written outside the tap" + if self.tap_kind == "mix": + mix = attn_res_mixture(want_updated, self.block, self.res_w, self.rms_w) + _within(self.tap, mix, TOL, "tap", where, "tapped mixture") + else: + assert torch.equal(bits(self.tap), bits(updated)), f"{where}: the tap is not updated" + outputs.append(self.tap) + assert R.same_on_ranks(*outputs), f"{where}: ranks disagree" + + def wrong_fraction(self, got) -> float: + """The fraction of ``updated``'s elements off the reference by more than the tolerance.""" + want = self.ref()[1].float() + return ((got[1].float() - want).abs() > TOL_TAIL * want.abs().max()).float().mean().item() + + def wrong_margin(self, got) -> float: + """``updated``'s largest error in units of the tolerance.""" + want = self.ref()[1].float() + return ((got[1].float() - want).abs().max() / (TOL_TAIL * want.abs().max())).item() + + +class PlainCall(Call): + """k3_sandwich_plain: x {-2, ..., 2} / 8 [M, 384] per rank and a [7168, 384] slice; with ``swiglu`` the gate_up + output [M, 1792] per rank (gate columns first: gates from GATES, up {-2, ..., 2} / 64) and a [7168, 896] slice. + silu is exact in fp32 on these gates -- silu(0) = 0, and for g >= 32 the sum 1 + exp(-g) rounds to 1, so the + sigmoid is 1 -- hence silu_and_mul(x) = gate * up exactly, in the kernel (bf16((g * sigmoid(g)) * up)) as in the + reference. ``residual``: True (drawn) or a tensor.""" + + def __init__(self, seed, tokens, swiglu=False, residual=True, weights=0): + g = _gen(seed) + self.tokens = tokens + self.swiglu = swiglu + if swiglu: + self.xs = [ + torch.cat([_gates(g, (tokens, DOWN_K)), ls.exact_bf16(g, (tokens, DOWN_K), -2, 3, 1 / 64)], dim=1) + for _ in range(R.world) + ] + self.weights = WEIGHTS["down"][weights] + else: + self.xs = [ls.exact_bf16(g, (tokens, PLAIN_K), -2, 3, 1 / 8) for _ in range(R.world)] + self.weights = WEIGHTS["plain"][weights] + self.carry = _running_sum(g, tokens) if residual is True else residual + self.norm_w = (1.0 + 0.1 * torch.randn(H, generator=g, device="cuda")).bfloat16() + + def operand(self, x: torch.Tensor) -> torch.Tensor: + """The projection's input: x, or silu_and_mul(x) = gate * up (exact on these payloads).""" + if not self.swiglu: + return x + return (x[:, :DOWN_K].float() * x[:, DOWN_K:].float()).bfloat16() + + def ref(self): + partials = [bf16_partial(self.operand(x), w) for x, w in zip(self.xs, self.weights)] + updated = reduce_ref(partials, self.carry) + u = updated.float() + normed = (u * (u.square().mean(dim=-1, keepdim=True) + EPS).rsqrt() * self.norm_w.float()).bfloat16() + return normed, updated + + def run(self, ws): + r = R.rank + return OPS["plain"](self.xs[r], self.weights[r], self.carry, self.norm_w, EPS, ws, swiglu=self.swiglu) + + def refill(self, fresh, carry): + for mine, new in zip(self.xs, fresh.xs): + mine.copy_(new) + if carry: + self.carry.copy_(fresh.carry) + + +def _label(i: int, call: Call) -> str: + return f"call {i} ({type(call).__name__} M {call.tokens})" + + +def run_sequence(seq, ws, verify: bool = True, late: Optional[random.Random] = None, where: str = "sequence"): + """Run ``seq`` -- (chain, call) pairs -- in order on ``ws``. A call's carry is the ``updated`` of the previous call + of its chain (a chain's first call keeps its own), as the model chains its residual streams. With ``late`` (a + random.Random seeded alike on every rank) every call starts after a barrier, a random rank 5 ms late; with + ``verify`` every call is checked against its reference as it returns. Returns the outputs.""" + last: Dict[str, torch.Tensor] = {} + outs = [] + for i, (chain, call) in enumerate(seq): + if chain in last: + call.carry = last[chain] + if late is not None: + R.barrier() + R.late(late.randrange(R.world)) + got = call.run(ws) + if verify: + call.verify(got, f"{where} {_label(i, call)}") + last[chain] = got[1] + outs.append(got) + return outs + + +def _chain_heads(seq) -> set: + seen, heads = set(), set() + for i, (chain, _) in enumerate(seq): + if chain not in seen: + seen.add(chain) + heads.add(i) + return heads + + +def armed_and_sized(ws) -> None: + """``create`` returned this rank's view of a buffer sized for the group (two halves of [8 tokens][W ranks][7168] + bf16 as int32 words), every word empty, every counter zero.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_sandwich import k3_sandwich_kernel as kernel + + assert ws.world_size == R.world and ws.rank == R.rank + assert ws.uc.dtype == ws.mc.dtype == ws.flags.dtype == torch.int32 + words = 2 * 8 * R.world * H // 2 + assert ws.uc.numel() == ws.mc.numel() == words == kernel.buffer_words(R.world) + assert ws.flags.numel() == FLAG_WORDS + assert bool((ws.uc == kernel.EMPTY_WORD).all()), "every word empty" + assert int(ws.flags.abs().sum()) == 0, "every counter 0" + + +def create_refuses_capture(workspace_type, ws, next_call: Call) -> None: + """``create`` is collective and allocates: under CUDA-graph capture it raises RuntimeError on every rank at once, + and on one rank alone while its peers do not call it -- before any collective, since a rank inside one would wait + for its peers there. The workspace in use is untouched: the next call is correct.""" + + def attempt() -> bool: + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + graph = torch.cuda.CUDAGraph() + try: + with torch.cuda.graph(graph, stream=stream): + workspace_type.create(R.mapping, fabric_handle=R.fabric) + except RuntimeError as exc: + return "outside CUDA-graph capture" in str(exc) + return False + + every = attempt() + R.barrier() + alone = attempt() if R.rank == R.world - 1 else True + assert R.all_true(every and alone), f"create under capture: every rank raised {every}, one rank alone {alone}" + next_call.verify(next_call.run(ws), "after the refused creates") + + +def counters_advance_once(ws, call: Call) -> None: + """One call advances every one of the 56 CTAs' counters by exactly one (the parity, the half the next call uses, + is shared by all of them) and leaves the spare flag words alone.""" + before = ws.flags.clone() + call.verify(call.run(ws), "counter call") + torch.cuda.synchronize() + delta = (ws.flags - before)[:NUM_CTAS] + assert bool((delta == 1).all()), f"counters advanced by {delta.unique().tolist()}" + assert int(ws.flags[NUM_CTAS:].abs().sum()) == 0, "a spare flag word was written" + + +def unsupported_raises(ws, bad_calls, next_call: Call) -> None: + """Each of ``bad_calls`` ((label, thunk) pairs) raises ValueError on every rank before any launch: every counter + is as it was. The next call is correct.""" + before = ws.flags.clone() + raised = {} + for label, thunk in bad_calls: + try: + thunk() + raised[label] = False + except ValueError: + raised[label] = True + torch.cuda.synchronize() + kept = torch.equal(ws.flags, before) + assert R.all_true(all(raised.values()) and kept), f"raised {raised}, counters kept {kept}" + next_call.verify(next_call.run(ws), "after the rejected calls") + + +def dip_and_regrow(ws, make_step: Callable[[int, int], list], where: str = "dip") -> None: + """Decode steps (``make_step(seed, M)``) at M = 8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8, a random rank 5 ms late at every + call, every call against its reference: a call after a smaller one must not see words an older, larger call + left.""" + late = random.Random(7) + for i, t in enumerate(DIP_STEPS): + run_sequence(make_step(3000 + 100 * i, t), ws, late=late, where=f"{where} step {i} M {t}") + + +def interleaved(ws_a, ws_b, make_call: Callable[[int], Call], where: str = "interleaved") -> None: + """Calls alternate between two workspaces in an irregular pattern (A A B A B B ...), so that the two objects' + counters differ, every call against its reference. The pattern is the same on every rank: calls on one stream + run in order and each waits for its peers, so ranks issuing calls on two workspaces in different orders deadlock + (measured for the MNNVL entry; not exercised).""" + for i, which in enumerate(INTERLEAVE): + call = make_call(i) + call.verify(call.run(ws_a if which == "A" else ws_b), f"{where} {which} {i}") + + +def capture_and_replay(ws, make_seq: Callable[[int], list], eager_between: Callable[[int], List[Call]], where: str, + replays: int = REPLAYS) -> None: # fmt: skip + """Capture the call sequence ``make_seq(seed)`` on ``ws`` -- after one eager, verified run of it: the kernels + compile on their first call -- and replay it ``replays`` times with every input rewritten in place from + ``make_seq(another seed)``, every replayed call against its reference, with the eager calls ``eager_between(rep)`` + on the same workspace between replays, each against its reference. Replays and eager calls advance the same + counters, in the same order on every rank.""" + seq = make_seq(5000) + heads = _chain_heads(seq) + run_sequence(seq, ws, where=f"{where} eager run") + R.barrier() + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + outs = run_sequence(seq, ws, verify=False) + R.barrier() + for rep in range(replays): + fresh = make_seq(6000 + 100 * rep) + for i, ((_, call), (_, new)) in enumerate(zip(seq, fresh)): + call.refill(new, carry=i in heads) + R.barrier() + graph.replay() + for i, ((_, call), got) in enumerate(zip(seq, outs)): + call.verify(got, f"{where} replay {rep} {_label(i, call)}") + for j, call in enumerate(eager_between(rep)): + call.verify(call.run(ws), f"{where} eager {_label(j, call)} after replay {rep}") + del graph + + +def swapped_pair_is_wrong(ws, c1: Call, c2: Call, c3: Call) -> None: + """Negative control: rank 0 makes two same-shaped calls on one workspace in swapped order. Every call returns and + nothing raises or hangs (the counters agree), but every rank's two results are wrong: a call pairs with the + peers' call at the same position. Wrong means more than half of ``updated``'s elements fail the comparison and + the largest error is over 10 times its bound. Then a plain call is correct again.""" + R.barrier() + if R.rank == 0: + got2, got1 = c2.run(ws), c1.run(ws) + else: + got1, got2 = c1.run(ws), c2.run(ws) + torch.cuda.synchronize() + wrong = [c.wrong_fraction(got) for c, got in ((c1, got1), (c2, got2))] + margin = [c.wrong_margin(got) for c, got in ((c1, got1), (c2, got2))] + assert R.all_true(min(wrong) > 0.5 and min(margin) > 10), ( + f"the swap went unnoticed: wrong fractions {wrong}, largest errors over the bound {margin}" + ) + c3.verify(c3.run(ws), "after the swapped pair") diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_oproj_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_oproj_op_matrix.py new file mode 100644 index 000000000000..333ce767fd3b --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_oproj_op_matrix.py @@ -0,0 +1,214 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU certification matrix for the ``comm/k3_sandwich_oproj`` catalog entry and its ``K3SandwichWorkspace``. + +The kernel's correctness depends on state that outlives a call (each CTA's call count, whose parity picks the +buffer half every rank pushes into), so beyond single calls this drives call *sequences*: layers x steps with the +token count dipping and growing back and a random rank late, two workspaces interleaved, CUDA-graph capture and +replay mixed with eager calls, and a negative control in which one rank swaps two calls and every rank gets a wrong +answer without an error. The workspace is also shared the way the model shares it: one decode step of target layers +(this op, then ``k3_sandwich_tail``) with the drafter's ``k3_sandwich_plain`` calls interleaved, call by call, then +that step captured and replayed between eager calls of the three ops. And ``create`` must refuse CUDA-graph capture. + + CUDA_VISIBLE_DEVICES=0,1,2,3 python _k3_sandwich_oproj_op_matrix.py [--world-size 4] + srun -n 16 --mpi=pmix python _k3_sandwich_oproj_op_matrix.py --launcher srun --world-size 16 + +Not a pytest module: one fixed sequence of checks inside one W-rank job (they share the workspaces and their +counters). The collected entry point is ``test_modeling_v2_k3_sandwich_oproj_op_matrix.py``. + +Shapes are TP16's per-rank shapes (core [M, 768], o_weight [7168, 768]; the shared step's tail and plain calls at +theirs) whatever W; the kernel sums the ranks in chunks of 8, so W <= 8 exercises one chunk. Every rank draws every +rank's core and weight slice from one seed, so each rank holds the whole reference. core and o_weight are small +multiples of 1/8 and 1/16: every partial product sum is exact in fp32, so ``updated`` = bf16(prefix + bf16(sum_r +bf16(core_r @ o_weight_r^T))) is compared bit for bit; ``normed`` against the fp32 reference within 2e-2; every +output bitwise across the ranks (``_k3_sandwich_common``). +""" + +import random +import sys +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parent)) +import _k3_sandwich_common as cm # noqa: E402 +import _lockstep as ls # noqa: E402 + +assert torch.cuda.is_available(), "k3_sandwich_oproj requires CUDA devices" + +DEADLINE_S = 900 +LAYERS = 12 +SNAPSHOTS = (0, 1, 4, 8) # candidates 1..9 +TARGET_LAYERS = 3 # the shared step's target layers (one o_weight set each) +SHARED_STEPS = ((8, 3), (2, 8), (7, 7)) # (target M, drafter M) of the eager shared steps +SHARED_CAPTURED = (8, 4) +SHARED_REPLAYS = 6 + +R = None +WORKSPACE = None # the K3SandwichWorkspace type, as this entry's wrapper exports it +WS_A = None +WS_B = None + + +def check_workspace_is_armed_and_sized() -> None: + cm.armed_and_sized(WS_A) + cm.armed_and_sized(WS_B) + + +def check_create_refuses_capture() -> None: + """``create`` under CUDA-graph capture raises on every rank at once and on one rank alone (before any + collective); the workspace in use is untouched.""" + cm.create_refuses_capture(WORKSPACE, WS_A, cm.OprojCall(600, 8, 2)) + + +def check_single_calls() -> None: + for t in cm.TOKENS: + for s in SNAPSHOTS: + for with_prefix in (True, False): + call = cm.OprojCall( + 1000 + 37 * t + s, t, s, prefix=with_prefix or None, weights=t % cm.WEIGHT_SETS + ) + call.verify(call.run(WS_A), f"M {t} snapshots {s} prefix {with_prefix}") + + +def check_counters_advance_once_per_call() -> None: + cm.counters_advance_once(WS_A, cm.OprojCall(1500, 4, 2)) + + +def check_unsupported_shape_raises_on_every_rank() -> None: + """M 9 raises ValueError on every rank before any launch; the next call is correct.""" + bad = cm.OprojCall(2000, 9, 1) + cm.unsupported_raises(WS_A, [("M 9", lambda: bad.run(WS_A))], cm.OprojCall(2001, 8, 2)) + + +def _step(seed, tokens): + """One decode step: ``LAYERS`` chained calls, each layer's prefix the previous layer's ``updated``.""" + return [ + ("target", cm.OprojCall(seed + 1 + layer, tokens, layer % 9, prefix=True if layer == 0 else None, + weights=layer % cm.WEIGHT_SETS)) # fmt: skip + for layer in range(LAYERS) + ] + + +def check_dip_and_regrow_sequence() -> None: + """Steps of 12 layers at M 8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8, a random rank late at every call: a call after a + smaller one must not see words an older, larger call left.""" + cm.dip_and_regrow(WS_A, _step) + + +def check_two_workspaces_interleaved() -> None: + """Two workspaces are two sets of counters: 20 calls alternate between them in an irregular pattern.""" + cm.interleaved( + WS_A, + WS_B, + lambda i: cm.OprojCall(4000 + i, (3, 8, 1, 8, 5)[i % 5], i % 9, weights=i % cm.WEIGHT_SETS), + ) + + +def check_graph_capture_and_replay() -> None: + """A captured step of 12 chained calls (M 8) replayed 8 times with rewritten inputs, an eager call of another M on + the same workspace between replays.""" + cm.capture_and_replay( + WS_B, + lambda seed: _step(seed, 8), + lambda rep: [cm.OprojCall(7000 + rep, (3, 1, 6, 5)[rep % 4], rep % 9, weights=rep % cm.WEIGHT_SETS)], + "captured step", + ) + + +def _shared_step(seed, t_tokens, d_tokens): + """One decode step as the model runs it on one workspace. Target layer l: k3_sandwich_oproj (post-attention), then + k3_sandwich_tail (the MoE tail and the next pre-attention step; layer 1's stores ``updated`` into a bank row), + chained through the target's prefix sum. The drafter's layers -- k3_sandwich_plain (o_proj, K 384), then its + SwiGLU form (down, K 896), chained through the drafter's residual -- run after target layers 0 and 2 at the + drafter's own token count.""" + seq = [] + for layer in range(TARGET_LAYERS): + s = seed + 10 * layer + seq.append(("target", cm.OprojCall(s + 1, t_tokens, (4 * layer) % 9, prefix=True if layer == 0 else None, + weights=layer))) # fmt: skip + seq.append(("target", cm.TailCall(s + 2, t_tokens, (4 * layer + 1) % 9, prefix=None, updated_out=layer == 1))) + if layer != 1: + seq.append(("drafter", cm.PlainCall(s + 3, d_tokens, residual=True if layer == 0 else None))) + seq.append(("drafter", cm.PlainCall(s + 4, d_tokens, swiglu=True, residual=None))) + return seq + + +def check_shared_sequence() -> None: + """The three sandwich ops on one workspace, as the model runs them (``_shared_step``, 10 calls): steps at (target + M, drafter M) = (8, 3), (2, 8), (7, 7), a random rank late at every call, every call against its reference. + The three wrappers export one workspace type.""" + from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm import k3_sandwich_plain, k3_sandwich_tail + + assert k3_sandwich_tail.K3SandwichWorkspace is WORKSPACE is k3_sandwich_plain.K3SandwichWorkspace + late = random.Random(11) + for i, (t, d) in enumerate(SHARED_STEPS): + cm.run_sequence(_shared_step(9000 + 100 * i, t, d), WS_A, late=late, where=f"shared step {i} M {t} / {d}") + + +def check_shared_sequence_captured() -> None: + """That step captured at (target M, drafter M) = (8, 4) and replayed 6 times with rewritten inputs, eager calls of + the three ops at other token counts on the same workspace between replays (one or two, so the replays start at + either counter parity).""" + eager = ( + lambda rep: [cm.TailCall(9700 + rep, 3, 2)], + lambda rep: [cm.PlainCall(9710 + rep, 5), cm.OprojCall(9720 + rep, 1, 4)], + lambda rep: [cm.PlainCall(9730 + rep, 8, swiglu=True)], + ) + cm.capture_and_replay( + WS_A, + lambda seed: _shared_step(seed, *SHARED_CAPTURED), + lambda rep: eager[rep % len(eager)](rep), + "captured shared step", + replays=SHARED_REPLAYS, + ) + + +def check_wrong_call_order_is_detected() -> None: + """Negative control: rank 0 swaps two same-shaped calls on one workspace. Every call returns and nothing raises + or hangs, but every rank's two results are wrong (more than half the ``updated`` elements differ); then a plain + call is correct again.""" + cm.swapped_pair_is_wrong( + WS_A, cm.OprojCall(8000, 8, 3), cm.OprojCall(8001, 8, 3, weights=1), cm.OprojCall(8002, 8, 3) + ) + + +CHECKS = [ + check_workspace_is_armed_and_sized, + check_create_refuses_capture, + check_single_calls, + check_counters_advance_once_per_call, + check_unsupported_shape_raises_on_every_rank, + check_dip_and_regrow_sequence, + check_two_workspaces_interleaved, + check_graph_capture_and_replay, + check_shared_sequence, + check_shared_sequence_captured, + # Stays last: it deliberately disagrees on call order. + check_wrong_call_order_is_detected, +] + + +def _run_one_rank(args) -> int: + global R, WORKSPACE, WS_A, WS_B + R = ls.Rank(args) + from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.k3_sandwich_oproj import ( + K3SandwichWorkspace, + ) + + WORKSPACE = K3SandwichWorkspace + with torch.inference_mode(): + cm.bind(R, {"oproj": cm.WEIGHT_SETS, "tail": 1, "plain": 1, "down": 1}) + WS_A = K3SandwichWorkspace.create(R.mapping, fabric_handle=R.fabric) + WS_B = K3SandwichWorkspace.create(R.mapping, fabric_handle=R.fabric) + code = ls.run_checks(R, CHECKS) + cm.report() + return code + + +if __name__ == "__main__": + ARGS = ls.parse_args(sys.argv[1:]) + if ARGS.rank_worker or ARGS.launcher == "srun": + sys.exit(_run_one_rank(ARGS)) + ls.spawn(__file__, ARGS, DEADLINE_S) + print("OK") diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_plain_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_plain_op_matrix.py new file mode 100644 index 000000000000..c9adbff4760f --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_plain_op_matrix.py @@ -0,0 +1,177 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU certification matrix for the ``comm/k3_sandwich_plain`` catalog entry on a ``K3SandwichWorkspace``. + +The drafter's sandwich: a row-parallel projection (the o_proj slice, K 384; or with ``swiglu`` SiLU-and-mul of the +gate_up output and the down projection slice, K 896), the TP all-reduce, the residual add and an RMSNorm. The +kernel's correctness depends on state that outlives a call (each CTA's call count, whose parity picks the buffer half +every rank pushes into), so beyond single calls this drives call *sequences*: drafter layers x steps with the token +count dipping and growing back and a random rank late, two workspaces interleaved, CUDA-graph capture and replay mixed +with eager calls, and a negative control in which one rank swaps two calls and every rank gets a wrong answer without +an error. The workspace shared with ``k3_sandwich_oproj`` and ``k3_sandwich_tail`` is certified by +``_k3_sandwich_oproj_op_matrix.py``. + + CUDA_VISIBLE_DEVICES=0,1,2,3 python _k3_sandwich_plain_op_matrix.py [--world-size 4] + srun -n 16 --mpi=pmix python _k3_sandwich_plain_op_matrix.py --launcher srun --world-size 16 + +Not a pytest module: one fixed sequence of checks inside one W-rank job (they share the workspaces and their +counters). The collected entry point is ``test_modeling_v2_k3_sandwich_plain_op_matrix.py``. + +Shapes are TP16's per-rank shapes (x [M, 384] and weight [7168, 384]; with swiglu x [M, 1792] and weight [7168, 896]) +whatever W; the kernel sums the ranks in chunks of 8, so W <= 8 exercises one chunk. Every rank draws every rank's +inputs from one seed, so each rank holds the whole reference. x and the weights are small multiples of powers of two, +and the swiglu gates are 0, 32 or 64, on which silu is exact in fp32: every partial sum is exact, so ``updated`` = +bf16(residual + bf16(sum_r bf16(x_r @ weight_r^T))) is compared bit for bit in both forms; ``normed`` (the one-shot's +RMSNorm sums bf16-rounded squares in its own tree) against the fp32 RMSNorm within 2e-2; every output bitwise across +the ranks (``_k3_sandwich_common``). +""" + +import sys +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parent)) +import _k3_sandwich_common as cm # noqa: E402 +import _lockstep as ls # noqa: E402 + +assert torch.cuda.is_available(), "k3_sandwich_plain requires CUDA devices" + +DEADLINE_S = 900 +DRAFTER_LAYERS = 6 # two calls each: 12 per step + +R = None +WORKSPACE = None # the K3SandwichWorkspace type, as this entry's wrapper exports it +WS_A = None +WS_B = None + + +def check_workspace_is_armed_and_sized() -> None: + cm.armed_and_sized(WS_A) + cm.armed_and_sized(WS_B) + + +def check_create_refuses_capture() -> None: + """``create`` under CUDA-graph capture raises on every rank at once and on one rank alone (before any + collective); the workspace in use is untouched.""" + cm.create_refuses_capture(WORKSPACE, WS_A, cm.PlainCall(600, 8)) + + +def check_single_calls() -> None: + """M 1-8 x both forms x two draws.""" + for t in cm.TOKENS: + for swiglu in (False, True): + for draw in (0, 1): + call = cm.PlainCall(1000 + 37 * t + 7 * draw + int(swiglu), t, swiglu=swiglu, + weights=(t + draw) % cm.WEIGHT_SETS) # fmt: skip + call.verify(call.run(WS_A), f"M {t} swiglu {swiglu} draw {draw}") + + +def check_counters_advance_once_per_call() -> None: + cm.counters_advance_once(WS_A, cm.PlainCall(1500, 4)) + cm.counters_advance_once(WS_A, cm.PlainCall(1501, 4, swiglu=True)) + + +def check_unsupported_calls_raise_on_every_rank() -> None: + """M 9, the SwiGLU form on a K 384 slice (it takes K 896 only) and an x whose K is not the weight's each raise + ValueError on every rank before any launch; the next call is correct.""" + big = cm.PlainCall(2000, 9) + narrow = cm.PlainCall(2001, 4) + narrow.swiglu, narrow.xs = True, [torch.cat([x, x], dim=1) for x in narrow.xs] # [4, 768] on a K 384 slice + wide = cm.PlainCall(2002, 4) + wide.xs = [torch.zeros(4, 512, dtype=torch.bfloat16, device="cuda") for _ in wide.xs] + cm.unsupported_raises( + WS_A, + [ + ("M 9", lambda: big.run(WS_A)), + ("swiglu on K 384", lambda: narrow.run(WS_A)), + ("x K 512, weight K 384", lambda: wide.run(WS_A)), + ], + cm.PlainCall(2003, 8), + ) + + +def _step(seed, tokens): + """One drafter step: ``DRAFTER_LAYERS`` layers of k3_sandwich_plain (o_proj, K 384) then its SwiGLU form (down, + K 896), chained through the residual (each call's residual the previous call's ``updated``).""" + seq = [] + for layer in range(DRAFTER_LAYERS): + s = seed + 2 * layer + w = layer % cm.WEIGHT_SETS + seq.append(("drafter", cm.PlainCall(s + 1, tokens, residual=True if layer == 0 else None, weights=w))) + seq.append(("drafter", cm.PlainCall(s + 2, tokens, swiglu=True, residual=None, weights=w))) + return seq + + +def check_dip_and_regrow_sequence() -> None: + """Steps of 6 drafter layers (12 calls) at M 8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8, a random rank late at every call: a + call after a smaller one must not see words an older, larger call left.""" + cm.dip_and_regrow(WS_A, _step) + + +def check_two_workspaces_interleaved() -> None: + """Two workspaces are two sets of counters: 20 calls of both forms alternate between them in an irregular + pattern.""" + cm.interleaved( + WS_A, + WS_B, + lambda i: cm.PlainCall(4000 + i, (3, 8, 1, 8, 5)[i % 5], swiglu=i % 3 == 1, weights=i % cm.WEIGHT_SETS), + ) + + +def check_graph_capture_and_replay() -> None: + """A captured drafter step (6 layers, 12 chained calls, M 8) replayed 8 times with rewritten inputs, an eager call + of another M on the same workspace between replays.""" + cm.capture_and_replay( + WS_B, + lambda seed: _step(seed, 8), + lambda rep: [cm.PlainCall(7000 + rep, (3, 1, 6, 5)[rep % 4], swiglu=rep % 2 == 1, + weights=rep % cm.WEIGHT_SETS)], # fmt: skip + "captured step", + ) + + +def check_wrong_call_order_is_detected() -> None: + """Negative control: rank 0 swaps two same-shaped calls on one workspace. Every call returns and nothing raises + or hangs, but every rank's two results are wrong (more than half the ``updated`` elements differ); then a plain + call is correct again.""" + cm.swapped_pair_is_wrong(WS_A, cm.PlainCall(8000, 8), cm.PlainCall(8001, 8, weights=1), cm.PlainCall(8002, 8)) + + +CHECKS = [ + check_workspace_is_armed_and_sized, + check_create_refuses_capture, + check_single_calls, + check_counters_advance_once_per_call, + check_unsupported_calls_raise_on_every_rank, + check_dip_and_regrow_sequence, + check_two_workspaces_interleaved, + check_graph_capture_and_replay, + # Stays last: it deliberately disagrees on call order. + check_wrong_call_order_is_detected, +] + + +def _run_one_rank(args) -> int: + global R, WORKSPACE, WS_A, WS_B + R = ls.Rank(args) + from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.k3_sandwich_plain import ( + K3SandwichWorkspace, + ) + + WORKSPACE = K3SandwichWorkspace + with torch.inference_mode(): + cm.bind(R, {"plain": cm.WEIGHT_SETS, "down": cm.WEIGHT_SETS}) + WS_A = K3SandwichWorkspace.create(R.mapping, fabric_handle=R.fabric) + WS_B = K3SandwichWorkspace.create(R.mapping, fabric_handle=R.fabric) + code = ls.run_checks(R, CHECKS) + cm.report() + return code + + +if __name__ == "__main__": + ARGS = ls.parse_args(sys.argv[1:]) + if ARGS.rank_worker or ARGS.launcher == "srun": + sys.exit(_run_one_rank(ARGS)) + ls.spawn(__file__, ARGS, DEADLINE_S) + print("OK") diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_tail_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_tail_op_matrix.py new file mode 100644 index 000000000000..718747f2823e --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_tail_op_matrix.py @@ -0,0 +1,209 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU certification matrix for the ``comm/k3_sandwich_tail`` catalog entry on a ``K3SandwichWorkspace``. + +The pre-attention sandwich: this rank's row-parallel MoE tail ``[rmsnorm(latent)[:, lo:lo+224] | act] @ +tail_weight^T`` (the latent RMS taken over the whole reduced latent row and applied to the fp32 accumulator), the TP +all-reduce and the next layer's residual update, with the optional tap (the pre-norm mixture, or ``updated``) and +``updated_out``. The kernel's correctness depends on state that outlives a call (each CTA's call count, whose parity +picks the buffer half every rank pushes into), so beyond single calls this drives call *sequences*: layers x steps +with the token count dipping and growing back and a random rank late, two workspaces interleaved, CUDA-graph capture +and replay (with a tapping layer and a bank-row layer inside) mixed with eager calls, and a negative control in which +one rank swaps two calls and every rank gets a wrong answer without an error. The workspace shared with +``k3_sandwich_oproj`` and ``k3_sandwich_plain`` is certified by ``_k3_sandwich_oproj_op_matrix.py``. + + CUDA_VISIBLE_DEVICES=0,1,2,3 python _k3_sandwich_tail_op_matrix.py [--world-size 4] + srun -n 16 --mpi=pmix python _k3_sandwich_tail_op_matrix.py --launcher srun --world-size 16 + +Not a pytest module: one fixed sequence of checks inside one W-rank job (they share the workspaces and their +counters). The collected entry point is ``test_modeling_v2_k3_sandwich_tail_op_matrix.py``. + +Shapes are TP16's per-rank shapes (latent [M, 3584], act [M, 384], tail_weight [7168, 256 + 384]) whatever W; rank +r's latent slice ``lo`` moves from call to call over the 16 slices. The kernel sums the ranks in chunks of 8, so +W <= 8 exercises one chunk. Every rank draws every rank's inputs from one seed, so each rank holds the whole +reference: the fp64 tail partial rounded once to bf16 per rank, the ranks' sum and the prefix in fp32 as the kernel +adds them. The kernel's latent rsqrt is not torch's, so a partial element may round to the neighbouring bf16: +``updated`` is compared within 8e-3 of its largest magnitude (the negative control's wrong pairing is checked against +the same bound), ``normed`` and the tapped mixture within 2e-2, the tapped ``updated`` bit for bit, every output +bitwise across the ranks (``_k3_sandwich_common``). +""" + +import sys +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parent)) +import _k3_sandwich_common as cm # noqa: E402 +import _lockstep as ls # noqa: E402 + +assert torch.cuda.is_available(), "k3_sandwich_tail requires CUDA devices" + +DEADLINE_S = 900 +LAYERS = 12 +SNAPSHOTS = (0, 1, 4, 8) # candidates 1..9 +CAPTURED_OPTIONS = {3: {"tap": "mix"}, 7: {"updated_out": True}, 10: {"tap": "updated"}} + +R = None +WORKSPACE = None # the K3SandwichWorkspace type, as this entry's wrapper exports it +WS_A = None +WS_B = None + + +def check_workspace_is_armed_and_sized() -> None: + cm.armed_and_sized(WS_A) + cm.armed_and_sized(WS_B) + + +def check_create_refuses_capture() -> None: + """``create`` under CUDA-graph capture raises on every rank at once and on one rank alone (before any + collective); the workspace in use is untouched.""" + cm.create_refuses_capture(WORKSPACE, WS_A, cm.TailCall(600, 8, 2)) + + +def check_single_calls() -> None: + """M 1-8 x snapshots 0, 1, 4, 8 x prefix or none. Call i puts rank r's slice at ((5 r + i) % 16) x 224, so every + rank takes every slice, lo = 0 to 3360 (the last one's padding columns run past the latent row).""" + i = 0 + for t in cm.TOKENS: + for s in SNAPSHOTS: + for with_prefix in (True, False): + call = cm.TailCall(1000 + 37 * t + s, t, s, prefix=with_prefix or None, weights=t % cm.WEIGHT_SETS, + shift=i) # fmt: skip + call.verify(call.run(WS_A), f"M {t} snapshots {s} prefix {with_prefix} lo {call.los[R.rank]}") + i += 1 + + +def check_tap_and_updated_out() -> None: + """At every M, one call's inputs four times: no option, the mixture tapped into a column slice of a capture + buffer, ``updated`` tapped there, ``updated`` stored into a bank row (then returned as ``updated``). Every call + against the reference, nothing written outside the tap and the bank row, and the four calls' ``normed`` and + ``updated`` bit-identical.""" + for t in cm.TOKENS: + base = cm.TailCall(1300 + t, t, 3) + got0 = base.run(WS_A) + base.verify(got0, f"M {t} no options") + for options in ({"tap": "mix"}, {"tap": "updated"}, {"updated_out": True}): + call = base.with_options(**options) + got = call.run(WS_A) + call.verify(got, f"M {t} {options}") + same = all(torch.equal(cm.bits(a), cm.bits(b)) for a, b in zip(got, got0)) + assert same, f"M {t} {options}: outputs differ from the call without options" + + +def check_counters_advance_once_per_call() -> None: + cm.counters_advance_once(WS_A, cm.TailCall(1500, 4, 2)) + + +def check_unsupported_calls_raise_on_every_rank() -> None: + """M 9, a latent 2 bytes off 16-byte alignment (its rows are bulk-copied) and a tap whose rows are 7172 elements + apart (not a multiple of 8) each raise ValueError on every rank before any launch; the next call is correct.""" + big = cm.TailCall(2000, 9, 1) + misaligned = cm.TailCall(2001, 4, 1) + store = torch.zeros(4 * cm.LATENT + 8, dtype=torch.bfloat16, device="cuda") + shifted = store[1 : 1 + 4 * cm.LATENT].view(4, cm.LATENT) # contiguous, 2 bytes past an aligned address + shifted.copy_(misaligned.latent) + misaligned.latent = shifted + strided = cm.TailCall(2002, 4, 1) + rows = torch.zeros(4, cm.H + 4, dtype=torch.bfloat16, device="cuda") + strided.tap_kind, strided.tap = "mix", rows[:, : cm.H] # only run, never verified + cm.unsupported_raises( + WS_A, + [ + ("M 9", lambda: big.run(WS_A)), + ("misaligned latent", lambda: misaligned.run(WS_A)), + ("tap row stride 7172", lambda: strided.run(WS_A)), + ], + cm.TailCall(2003, 8, 2), + ) + + +def _step(seed, tokens, options=None): + """One decode step: ``LAYERS`` chained calls, each layer's prefix the previous layer's ``updated``; ``options`` + maps a layer to its tap / updated_out.""" + options = options or {} + return [ + ("target", cm.TailCall(seed + 1 + layer, tokens, layer % 9, prefix=True if layer == 0 else None, + weights=layer % cm.WEIGHT_SETS, **options.get(layer, {}))) # fmt: skip + for layer in range(LAYERS) + ] + + +def check_dip_and_regrow_sequence() -> None: + """Steps of 12 layers at M 8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8, a random rank late at every call: a call after a + smaller one must not see words an older, larger call left.""" + cm.dip_and_regrow(WS_A, _step) + + +def check_two_workspaces_interleaved() -> None: + """Two workspaces are two sets of counters: 20 calls alternate between them in an irregular pattern.""" + cm.interleaved( + WS_A, + WS_B, + lambda i: cm.TailCall(4000 + i, (3, 8, 1, 8, 5)[i % 5], i % 9, weights=i % cm.WEIGHT_SETS), + ) + + +def check_graph_capture_and_replay() -> None: + """A captured step of 12 chained calls (M 8; layer 3 taps the mixture, layer 7 stores ``updated`` into a bank row, + layer 10 taps ``updated``) replayed 8 times with rewritten inputs, an eager call of another M on the same + workspace between replays.""" + cm.capture_and_replay( + WS_B, + lambda seed: _step(seed, 8, CAPTURED_OPTIONS), + lambda rep: [cm.TailCall(7000 + rep, (3, 1, 6, 5)[rep % 4], rep % 9, weights=rep % cm.WEIGHT_SETS)], + "captured step", + ) + + +def check_wrong_call_order_is_detected() -> None: + """Negative control: rank 0 swaps two same-shaped calls (no prefix) on one workspace. Every call returns and + nothing raises or hangs, but every rank's two results are wrong, far outside the 8e-3 tolerance: more than half + the ``updated`` elements are off by more than it, the largest by over 10 times it. Then a plain call is correct + again.""" + cm.swapped_pair_is_wrong( + WS_A, + cm.TailCall(8000, 8, 3, prefix=None), + cm.TailCall(8001, 8, 3, prefix=None, weights=1), + cm.TailCall(8002, 8, 3), + ) + + +CHECKS = [ + check_workspace_is_armed_and_sized, + check_create_refuses_capture, + check_single_calls, + check_tap_and_updated_out, + check_counters_advance_once_per_call, + check_unsupported_calls_raise_on_every_rank, + check_dip_and_regrow_sequence, + check_two_workspaces_interleaved, + check_graph_capture_and_replay, + # Stays last: it deliberately disagrees on call order. + check_wrong_call_order_is_detected, +] + + +def _run_one_rank(args) -> int: + global R, WORKSPACE, WS_A, WS_B + R = ls.Rank(args) + from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.k3_sandwich_tail import ( + K3SandwichWorkspace, + ) + + WORKSPACE = K3SandwichWorkspace + with torch.inference_mode(): + cm.bind(R, {"tail": cm.WEIGHT_SETS}) + WS_A = K3SandwichWorkspace.create(R.mapping, fabric_handle=R.fabric) + WS_B = K3SandwichWorkspace.create(R.mapping, fabric_handle=R.fabric) + code = ls.run_checks(R, CHECKS) + cm.report() + return code + + +if __name__ == "__main__": + ARGS = ls.parse_args(sys.argv[1:]) + if ARGS.rank_worker or ARGS.launcher == "srun": + sys.exit(_run_one_rank(ARGS)) + ls.spawn(__file__, ARGS, DEADLINE_S) + print("OK") diff --git a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py new file mode 100644 index 000000000000..a8ecfd1fcedd --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py @@ -0,0 +1,636 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU certification matrix for the ``comm/mnnvl_allgather_split`` catalog entry on its ``MnnvlWorkspace``. + +The op takes one turn of the workspace's Lamport rotation, which every MNNVL op of the TP group advances (this op, +``comm/mnnvl_fusion_allreduce`` on either path and ``comm/mnnvl_allreduce_attn_res``). So beyond single calls (every +certified split at every certified token count against a bit-exact reference) this drives call *sequences*: decode +steps whose token count dips and grows back with a random rank late at every call; two workspaces interleaved; the +three MNNVL ops interleaved on one workspace; CUDA-graph capture and replay mixed with eager calls; and a negative +control in which one rank swaps two calls and every rank gets a wrong answer without an error. After every eager call +the workspace's ``buffer_flags`` are compared with this file's model of the rotation (``Rotation``): one turn per call +whatever the op, one stage, the bytes the call wrote. + + CUDA_VISIBLE_DEVICES=0,1,2,3 python _mnnvl_allgather_split_op_matrix.py [--world-size 4] + srun -N 4 --ntasks-per-node 4 --mpi=pmix python _mnnvl_allgather_split_op_matrix.py --launcher srun \ + --world-size 16 + +Not a pytest module: one fixed sequence of checks inside one W-rank job (they share the workspaces and their +rotations). The collected entry point is ``test_modeling_v2_mnnvl_allgather_split_op_matrix.py``. + +Every rank draws every rank's rows from one seed, so each rank holds the whole reference: the bf16 columns rounded by +torch (round to nearest even), the fp32 columns copied, -0.0 made +0.0, compared bit for bit. The rows exercise the +rounding (normal values over exponents 2^-16..2^16, inexact in bf16, and exact midpoints between two bf16 values), +the -0.0 rule in both parts, and special fp32 words. Every output is also compared bitwise across the ranks. The +all-reduce and attention-residual calls of the shared-workspace, capture and sequence checks are checked as in their +own matrices: sums bit for bit (small multiples of 1/16), normed outputs within a tolerance. +""" + +import copy +import random +import sys +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parent)) +import _lockstep as ls # noqa: E402 + +assert torch.cuda.is_available(), "mnnvl_allgather_split requires CUDA devices" + +DEADLINE_S = 900 +# Kimi K3's MoE head, sharded over the TP group: the routed latent (the bf16 columns) and the router logits of the +# routed experts (the fp32 columns). +LATENT = 3584 +EXPERTS = 896 +H_MODEL = 7168 # Kimi K3's hidden size (the all-reduces this op shares the workspace with) +TOKENS = (1, 2, 3, 4, 5, 6, 7, 8, 16, 32, 64) +# (bf16 columns, fp32 columns) per rank: Kimi K3's MoE head at W = 2, 4, 8, 16, then the smallest split. +SPLITS = tuple((LATENT // w, EXPERTS // w) for w in (2, 4, 8, 16)) + ((8, 4),) +DECODE_MAX_TOKENS = 8 +# Kimi K3's one-shot ceilings: DECODE_AR_ONE_SHOT_MAX_BYTES on decode steps (at most 8 tokens), +# WIDE_AR_ONE_SHOT_MAX_BYTES (main's default) on wide decode steps. +DECODE_ONE_SHOT_MAX_BYTES = 4 << 20 +WIDE_ONE_SHOT_MAX_BYTES = 1 << 20 +EPS = 1e-5 +TOL = 1e-2 # the fused all-reduce's normed, as its own matrix bounds it +ATTN_RES_TOL = 2e-2 # mnnvl_allreduce_attn_res's normed, as its own matrix bounds it +ATTN_RES_MAX_TOKENS = 16 # Kimi K3 sends the attention-residual all-reduce up to 16 tokens +INT16_MIN = torch.iinfo(torch.int16).min # the bf16 -0.0 word +INT32_MIN = torch.iinfo(torch.int32).min # the fp32 -0.0 word, the Lamport buffers' empty word +NEG_ZERO_STRIDE = 37 # every rank's all-reduce input holds -0.0 in these columns +# fp32 words the all-gather carries unchanged: +inf, -inf, NaN, +-the smallest denormal, the largest float. +SPECIAL_WORDS = (0x7F800000, -0x00800000, 0x7FC00000, 1, -0x7FFFFFFF, 0x7F7FFFFF) +LAYERS = 6 +DIP_STEPS = (8, 8, 8, 2, 7, 8, 1, 1, 64, 3, 32, 8, 16, 1, 64, 8) +SHARED_STEPS = (8, 2, 16, 64, 1, 7, 32, 8, 3, 16) +INTERLEAVED_TOKENS = (3, 8, 1, 64, 5, 16, 2) + +R = None +allgather_split = None +required_buffer_bytes = None +fusion_allreduce = None +MnnvlWorkspace = None +BUFFER_BYTES = None +WS_A = None +WS_B = None +ROT = {} +STATS = {"normed_err": 0.0, "attn_res_err": 0.0} + + +def buffer_bytes(world: int) -> int: + """One Lamport buffer holds the largest call this matrix makes: the all-gather at the widest certified split and + 64 tokens, a wide step's [64, 7168] all-reduce sent two-shot, or the attention-residual all-reduce at 16 + tokens.""" + widest = max(2 * b + 4 * f for b, f in SPLITS) + return max( + max(TOKENS) * world * widest, + 2 * -(-max(TOKENS) // world) * world * H_MODEL * 2, + ATTN_RES_MAX_TOKENS * H_MODEL * world * 2, + ) + + +def bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16 if t.element_size() == 2 else torch.int32) + + +def bits_equal(a: torch.Tensor, b: torch.Tensor) -> bool: + return a.shape == b.shape and a.dtype == b.dtype and torch.equal(bits(a), bits(b)) + + +def positive_zero(t: torch.Tensor) -> torch.Tensor: + """``t`` with every -0.0 made +0.0, by bit pattern: the Lamport buffers' empty word never travels.""" + words = bits(t) + sign = torch.iinfo(words.dtype).min + return torch.where(words == sign, torch.zeros_like(words), words).view(t.dtype) + + +def one_shot_bytes(tokens: int, hidden: int) -> int: + return tokens * hidden * R.world * 2 + + +def ceiling(tokens: int, hidden: int, path: str) -> int: + """An all-reduce's one_shot_max_bytes: at the boundary for "one" (the call's one-shot footprint itself, so + one-shot) and "two" (one byte less, so two-shot), or Kimi K3's ceiling for a step of ``tokens`` ("k3").""" + if path == "one": + return one_shot_bytes(tokens, hidden) + if path == "two": + return one_shot_bytes(tokens, hidden) - 1 + return DECODE_ONE_SHOT_MAX_BYTES if tokens <= DECODE_MAX_TOKENS else WIDE_ONE_SHOT_MAX_BYTES + + +def k3_split(): + """Kimi K3's MoE head shard at this world size: (latent columns, router-logit columns) per rank.""" + return LATENT // R.world, EXPERTS // R.world + + +class Rotation: + """This test's model of a workspace's ``buffer_flags`` (``mnnvl_workspace.FLAG_WORDS``): the buffer the next call + takes, the one the last call used, the bytes per buffer, the last call's stage count and the bytes it wrote per + stage (what the next call clears), and the arrival counter (0 between calls).""" + + def __init__(self, ws): + self.ws = ws + self.flags = [0, 2, ws.buffer_bytes, 0, 0, 0, 0, 0, 0] # as create() arms them + + def read(self): + torch.cuda.synchronize() + return self.ws.buffer_flags.view(torch.int32).tolist() + + def unchanged(self, where: str) -> None: + got = self.read() + assert got == self.flags, f"{where}: buffer_flags {got} moved from {self.flags}" + + def advance(self, record, where: str, calls: int = 1) -> None: + """``calls`` more calls ran on the workspace, the last of which recorded ``record`` = (stages, bytes).""" + stages, written = record + current = (self.flags[0] + calls) % 3 + self.flags = [current, (current - 1) % 3, self.ws.buffer_bytes, stages, *written, 0] + got = self.read() + assert got == self.flags, f"{where}: buffer_flags {got}, expected {self.flags}" + + +def rot(ws) -> Rotation: + return ROT[id(ws)] + + +def gather_payload(g, tokens, bf16_columns, fp32_columns): + """One rank's fp32 rows for the all-gather: normal values over exponents 2^-16..2^16 in the bf16 columns + (inexact in bf16), exact midpoints between two bf16 values in every fourth of them (ties go to even), -0.0 in + every fourth column of each part, and special words (infinities, NaN, denormals, the largest float) in the first + row of the fp32 part.""" + b = bf16_columns + x = torch.randn(tokens, b + fp32_columns, generator=g, device="cuda") + x[:, :b] *= torch.exp2(((torch.arange(b, device="cuda") % 5) - 2).float() * 8) + nearest = x[:, 1:b:4].bfloat16().float() + x[:, 1:b:4] = (nearest.view(torch.int32) | 0x8000).view(torch.float32) + words = x.view(torch.int32) + words[:, 3:b:4] = INT32_MIN + words[:, b + 2 :: 4] = INT32_MIN + special = SPECIAL_WORDS[: min(len(SPECIAL_WORDS), fp32_columns)] + words[0, b : b + len(special)] = torch.tensor(special, dtype=torch.int32, device="cuda") + return x + + +class AG: + """One mnnvl_allgather_split call's arguments on every rank and its reference.""" + + def __init__(self, seed, tokens, bf16_columns, fp32_columns): + g = torch.Generator(device="cuda").manual_seed(seed) + self.tokens, self.bf16_columns, self.fp32_columns = tokens, bf16_columns, fp32_columns + self.inputs = [ + gather_payload(g, tokens, bf16_columns, fp32_columns) for _ in range(R.world) + ] + + def fresh(self, seed): + return AG(seed, self.tokens, self.bf16_columns, self.fp32_columns) + + def run(self, ws, x=None): + return allgather_split(self.inputs[R.rank] if x is None else x, self.bf16_columns, ws) + + def ref(self): + """Every rank's leading columns rounded to bf16 by torch (round to nearest even) and its other columns as + they are, in rank order, -0.0 made +0.0.""" + b = self.bf16_columns + bf16_out = torch.cat([x[:, :b] for x in self.inputs], dim=1).bfloat16() + fp32_out = torch.cat([x[:, b:] for x in self.inputs], dim=1) + return positive_zero(bf16_out), positive_zero(fp32_out) + + def verify(self, got, where: str) -> None: + want = self.ref() + assert all(bits_equal(a, b) for a, b in zip(got, want)), ( + f"{where}: the gather differs from the reference" + ) + assert R.same_on_ranks(*got), f"{where}: ranks disagree" + + def record(self): + """What the call leaves in buffer_flags: one stage, the bytes it wrote.""" + written = self.tokens * R.world * (2 * self.bf16_columns + 4 * self.fp32_columns) + return 1, (written, 0, 0, 0) + + def static(self): + return {"x": self.inputs[R.rank].clone()} + + def refill(self, fresh, static) -> None: + """Take ``fresh``'s rows (same shape) into this call and its static buffer.""" + self.inputs = fresh.inputs + static["x"].copy_(fresh.inputs[R.rank]) + + +class AR: + """One comm/mnnvl_fusion_allreduce call's arguments on every rank and its reference. ``residual``: a tensor + (chained), True (drawn) or None (the plain sum); ``path``: see ``ceiling``.""" + + def __init__(self, seed, tokens, hidden, residual=None, path="k3"): + g = torch.Generator(device="cuda").manual_seed(seed) + self.tokens, self.hidden, self.path = tokens, hidden, path + self.one_shot_max_bytes = ceiling(tokens, hidden, path) + self.inputs = [] + for _ in range(R.world): + x = ls.exact_bf16(g, (tokens, hidden), -4, 5, 1 / 16) + x.view(torch.int16)[:, ::NEG_ZERO_STRIDE] = INT16_MIN + self.inputs.append(x) + if residual is True: + residual = ls.exact_bf16(g, (tokens, hidden), -32, 33, 1 / 16) + self.residual = residual + self.gamma = None + if residual is not None: + self.gamma = (1.0 + 0.1 * torch.randn(hidden, generator=g, device="cuda")).bfloat16() + + @property + def fused(self) -> bool: + return self.residual is not None + + @property + def one_shot(self) -> bool: + return one_shot_bytes(self.tokens, self.hidden) <= self.one_shot_max_bytes + + def fresh(self, seed): + return AR( + seed, self.tokens, self.hidden, residual=True if self.fused else None, path=self.path + ) + + def run(self, ws, x=None, residual=None): + x = self.inputs[R.rank] if x is None else x + if not self.fused: + return fusion_allreduce(x, ws, self.one_shot_max_bytes) + residual = self.residual if residual is None else residual + return fusion_allreduce(x, ws, self.one_shot_max_bytes, residual, self.gamma, EPS) + + def ref(self): + total = positive_zero(torch.stack([x.float() for x in self.inputs]).sum(dim=0)) + if not self.fused: + return total.bfloat16() + updated = (total + self.residual.float()).bfloat16() + x = updated.float() + normed = x * torch.rsqrt(x.square().mean(dim=-1, keepdim=True) + EPS) * self.gamma.float() + return normed.bfloat16(), updated + + def verify(self, got, where: str) -> None: + want = self.ref() + if not self.fused: + assert bits_equal(got, want), f"{where}: the sum differs from the exact reference" + assert R.same_on_ranks(got), f"{where}: ranks disagree" + return + (normed, updated), (want_normed, want_updated) = got, want + assert bits_equal(updated, want_updated), ( + f"{where}: updated differs from the exact reference" + ) + err = ls.rel_err(normed, want_normed) + STATS["normed_err"] = max(STATS["normed_err"], err) + assert err <= TOL, f"{where}: normed rel err {err:.3e} > {TOL}" + assert R.same_on_ranks(normed, updated), f"{where}: ranks disagree" + + def record(self): + if self.one_shot: + return 1, (one_shot_bytes(self.tokens, self.hidden), 0, 0, 0) + rows = -(-self.tokens // R.world) * R.world + return 2, (rows * self.hidden * 2, self.tokens * self.hidden * 2, 0, 0) + + def static(self): + static = {"x": self.inputs[R.rank].clone()} + if self.fused: + static["residual"] = self.residual.clone() + return static + + def refill(self, fresh, static) -> None: + self.inputs = fresh.inputs + static["x"].copy_(fresh.inputs[R.rank]) + if self.fused: + self.residual = fresh.residual + static["residual"].copy_(fresh.residual) + + +class AttnRes: + """One comm/mnnvl_allreduce_attn_res call (G1's entry, called through its op on the workspace's comm_buffer and + buffer_flags, as that entry's wrapper does) and its reference. ``prefix``: a tensor (chained) or True (drawn).""" + + def __init__(self, seed, tokens, snapshots, prefix=True): + g = torch.Generator(device="cuda").manual_seed(seed) + self.tokens, self.snapshots = tokens, snapshots + self.inputs = [ls.exact_bf16(g, (tokens, H_MODEL), -4, 5, 1 / 16) for _ in range(R.world)] + if prefix is True: + prefix = ls.exact_bf16(g, (tokens, H_MODEL), -32, 33, 1 / 16) + self.prefix = prefix + self.block = torch.randn(snapshots, tokens, H_MODEL, generator=g, device="cuda").bfloat16() + self.res_w = (torch.randn(H_MODEL, generator=g, device="cuda") * 0.05).bfloat16() + self.rms_w = (1.0 + 0.1 * torch.randn(H_MODEL, generator=g, device="cuda")).bfloat16() + self.out_w = (1.0 + 0.1 * torch.randn(H_MODEL, generator=g, device="cuda")).bfloat16() + + def run(self, ws): + normed, updated = torch.ops.trtllm.mnnvl_allreduce_attn_res( + self.inputs[R.rank], + self.prefix, + self.block, + self.res_w, + self.rms_w, + self.out_w, + EPS, + EPS, + ws.comm_buffer(torch.bfloat16), + ws.buffer_flags, + ) + return normed, updated + + def ref(self): + updated = (sum(x.float() for x in self.inputs) + self.prefix.float()).bfloat16() + normed = ls.residual_update_ref( + updated, self.block, self.res_w, self.rms_w, EPS, self.out_w, EPS + ) + return normed, updated + + def verify(self, got, where: str) -> None: + (normed, updated), (want_normed, want_updated) = got, self.ref() + assert bits_equal(updated, want_updated), ( + f"{where}: attn_res updated differs from the exact sum" + ) + err = ls.rel_err(normed, want_normed) + STATS["attn_res_err"] = max(STATS["attn_res_err"], err) + assert err <= ATTN_RES_TOL, f"{where}: attn_res normed rel err {err:.3e} > {ATTN_RES_TOL}" + assert R.same_on_ranks(normed, updated), f"{where}: ranks disagree" + + def record(self): + return 1, (self.tokens * H_MODEL * R.world * 2, 0, 0, 0) + + +def call_and_check(call, ws, where: str, late=None): + """Run ``call`` eagerly on ``ws`` (rank ``late``, if given, 5 ms late), check its result and what it left in the + workspace's flags; return the result.""" + if late is not None: + R.barrier() + R.late(late) + got = call.run(ws) + call.verify(got, where) + rot(ws).advance(call.record(), where) + return got + + +def call_late(call, ws, where: str, late_rng: random.Random): + """``call_and_check`` with a rank drawn from ``late_rng`` (the same draw on every rank) 5 ms late.""" + return call_and_check(call, ws, where, late=late_rng.randrange(R.world)) + + +def expect_refusal(fn, error, where: str) -> None: + """``fn`` raises ``error`` on every rank and WS_A's flags do not move.""" + try: + fn() + raised = False + except error: + raised = True + assert R.all_true(raised), f"{where}: not refused with {error.__name__} on every rank" + rot(WS_A).unchanged(where) + + +def with_inputs(call, inputs): + """``call`` as if the ranks had sent ``inputs`` (one tensor per rank).""" + mixed = copy.copy(call) + mixed.inputs = inputs + return mixed + + +def check_workspaces_are_armed_and_sized() -> None: + """create() armed both workspaces (every Lamport word -0.0; flags at buffer 0, buffer 2 dirty with nothing to + clear), and one buffer holds the largest certified all-gather and every other call of this matrix.""" + for ws in (WS_A, WS_B): + assert ws.world_size == R.world and ws.rank == R.rank + assert ws.buffer_bytes == BUFFER_BYTES and BUFFER_BYTES % 32 == 0 + assert ws.comm_buffer(torch.bfloat16).shape == (3, BUFFER_BYTES // 2) + armed = ws.lamport.view(torch.int32) + assert bool((armed == torch.tensor(INT32_MIN, dtype=torch.int32, device="cuda")).all()), ( + "every word -0.0" + ) + rot(ws).unchanged("armed") + for b, f in SPLITS: + assert required_buffer_bytes(max(TOKENS), b, f, R.world) <= BUFFER_BYTES + + +def check_single_calls() -> None: + """Every certified split at every certified token count, one call after the other on one workspace, bit for bit + against the reference (the rounding, ties to even, -0.0, special fp32 words), the flags after each recording one + stage and the bytes the call wrote, which equal required_buffer_bytes.""" + for b, f in SPLITS: + for t in TOKENS: + where = f"T {t} split {b} + {f}" + call = AG(1000 + 37 * t + b, t, b, f) + call_and_check(call, WS_A, where) + need = required_buffer_bytes(t, b, f, R.world) + assert need == call.record()[1][0], f"{where}: required_buffer_bytes {need}" + + +def check_unsupported_calls_raise_on_every_rank() -> None: + """Refused on every rank before the workspace is touched (its flags do not move), and the next call is correct: + more rows than one Lamport buffer holds (the wrapper's ValueError); bf16 columns not a multiple of 8, remaining + columns not a multiple of 4, a bf16 input, no rows (the op's RuntimeError).""" + b, f = k3_split() + t = BUFFER_BYTES // (R.world * (2 * b + 4 * f)) + 1 + over = AG(2000, t, b, f) + expect_refusal(lambda: over.run(WS_A), ValueError, f"T {t}: more rows than one buffer holds") + rows = torch.randn(4, 16, device="cuda") + narrow = rows[:, :10].contiguous() + expect_refusal(lambda: allgather_split(rows, 12, WS_A), RuntimeError, "12 bf16 columns") + expect_refusal(lambda: allgather_split(narrow, 8, WS_A), RuntimeError, "2 fp32 columns") + expect_refusal(lambda: allgather_split(rows.bfloat16(), 8, WS_A), RuntimeError, "bf16 rows") + expect_refusal(lambda: allgather_split(rows[:0], 8, WS_A), RuntimeError, "no rows") + call_and_check(AG(2001, 8, b, f), WS_A, "after the refused calls") + + +def check_dip_and_regrow_sequence() -> None: + """The k3_spec_accept failure mode: a call after a smaller one must not read what an older, larger call left. + 16 steps of 6 layers (one MoE head all-gather per layer, Kimi K3's split) at T 8, 8, 8, 2, 7, 8, 1, 1, 64, 3, 32, + 8, 16, 1, 64, 8, a random rank late at every call.""" + late = random.Random(7) + b, f = k3_split() + for i, t in enumerate(DIP_STEPS): + for layer in range(LAYERS): + call = AG(3000 + 100 * i + layer, t, b, f) + call_late(call, WS_A, f"step {i} T {t} layer {layer}", late) + + +def check_two_workspaces_interleaved() -> None: + """Two workspaces are two rotations: all-gathers of mixed token counts and splits alternate between them in an + irregular pattern (A A B A B B ...), every call is correct and each workspace's flags move with its own calls + only. The pattern is the same on every rank: calls on one stream are serialized and each waits for its peers, so + ranks issuing calls on two workspaces in different orders deadlock (measured for mnnvl_allreduce_attn_res, + runs/drafter/u4-mnnvl-srun-2).""" + pattern = "AABABBAAAB" * 2 + for i, which in enumerate(pattern): + t = INTERLEAVED_TOKENS[i % len(INTERLEAVED_TOKENS)] + call = AG(4700 + i, t, *SPLITS[i % len(SPLITS)]) + call_and_check(call, WS_A if which == "A" else WS_B, f"interleaved {which} {i}") + + +def check_one_workspace_three_ops() -> None: + """The three MNNVL entries on one workspace in Kimi K3's order, a random rank late at every call. A step of at + most 16 tokens: per layer the pre-attention all-reduce with the attention-residual epilogue + (comm/mnnvl_allreduce_attn_res, prefix chained), the MoE head all-gather (this entry, K3's split), the + routed-latent all-reduce and the fused all-reduce (comm/mnnvl_fusion_allreduce, chained). A wide step (32, 64 + tokens): the wide all-reduce [T, 7168], the all-gather and the routed-latent all-reduce, at the 1 MiB ceiling. + Every call takes one turn of the rotation whatever the op (the flags after each), and every result is right.""" + late = random.Random(11) + b, f = k3_split() + for i, t in enumerate(SHARED_STEPS): + seed = 5000 + 100 * i + g = torch.Generator(device="cuda").manual_seed(seed) + prefix = ls.exact_bf16(g, (t, H_MODEL), -32, 33, 1 / 16) + residual = ls.exact_bf16(g, (t, H_MODEL), -32, 33, 1 / 16) + for layer in range(2): + s = seed + 10 * layer + 1 + where = f"shared step {i} T {t} layer {layer}" + if t <= ATTN_RES_MAX_TOKENS: + attn = AttnRes(s, t, (0, 2, 5)[(i + layer) % 3], prefix=prefix) + prefix = call_late(attn, WS_A, f"{where} attn_res", late)[1] + else: + call_late(AR(s, t, H_MODEL), WS_A, f"{where} wide", late) + call_late(AG(s + 1, t, b, f), WS_A, f"{where} all-gather", late) + call_late(AR(s + 2, t, LATENT), WS_A, f"{where} latent", late) + if t <= ATTN_RES_MAX_TOKENS: + fused = AR(s + 3, t, H_MODEL, residual=residual) + residual = call_late(fused, WS_A, f"{where} fused", late)[1] + + +def check_graph_capture_and_replay() -> None: + """A captured step of five calls on WS_B, the MoE layers of a decode step: the head all-gather (T 8), the + routed-latent all-reduce one-shot, the head all-gather again, a [32, 3584] all-reduce sent two-shot and the head + all-gather at T 32. Replayed 8 times with rewritten rows, an eager all-gather of another token count or split, or + an all-reduce, on the same workspace between replays: replays and eager calls take turns of one rotation (the + flags after each), in the same order on every rank, and every replayed and eager result is right.""" + b, f = k3_split() + calls = [ + AG(6000, 8, b, f), + AR(6001, 8, LATENT), + AG(6002, 8, b, f), + AR(6003, 32, LATENT, path="two"), + AG(6004, 32, b, f), + ] + statics = [c.static() for c in calls] + + def step(): + return [c.run(WS_B, **s) for c, s in zip(calls, statics)] + + outs = step() # every call once eagerly, outside capture + for i, (c, got) in enumerate(zip(calls, outs)): + c.verify(got, f"eager step call {i}") + rot(WS_B).advance(calls[-1].record(), "eager step", calls=len(calls)) + R.barrier() + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + outs = step() + R.barrier() + rot(WS_B).unchanged("captured, not replayed") + eager = ( + lambda seed: AG(seed, 1, b, f), + lambda seed: AG(seed, 64, b, f), + lambda seed: AR(seed, 3, LATENT, path="one"), + lambda seed: AG(seed, 16, *SPLITS[-1]), + ) + for rep in range(8): + for i, (c, s) in enumerate(zip(calls, statics)): + c.refill(c.fresh(7000 + 100 * rep + i), s) + R.barrier() + graph.replay() + for i, (c, got) in enumerate(zip(calls, outs)): + c.verify(got, f"replay {rep} call {i}") + rot(WS_B).advance(calls[-1].record(), f"replay {rep}", calls=len(calls)) + call_and_check(eager[rep % len(eager)](7500 + rep), WS_B, f"eager after replay {rep}") + del graph + + +def check_create_under_capture_raises() -> None: + """MnnvlWorkspace.create allocates and exchanges handles, so it refuses CUDA-graph capture: with every rank + capturing, it raises RuntimeError on every rank (before any communication), and the next call on a workspace in + use is correct.""" + graph, stream = torch.cuda.CUDAGraph(), torch.cuda.Stream() + raised = False + with torch.cuda.graph(graph, stream=stream): + try: + MnnvlWorkspace.create(R.mapping, BUFFER_BYTES, fabric_handle=R.fabric) + except RuntimeError as exc: + raised = "before capture" in str(exc) + del graph + assert R.all_true(raised), "create() under capture did not raise on every rank" + call_and_check(AG(8500, 8, *k3_split()), WS_A, "after the refused create") + + +def check_wrong_call_order_is_detected() -> None: + """Negative control: rank 0 issues two same-shaped all-gathers on one workspace in swapped order. Nothing raises + and nothing hangs (the two calls write the same slots of the same buffers), but every rank's two results are + wrong: rank 0's columns hold rank 0's rows of the other call, every other rank's columns are right -- the k-th + call on every rank gathers what every rank sent at position k. A plain call right after is correct again: a + swapped pair realigns the positions.""" + b, f = k3_split() + c1, c2 = AG(9000, 8, b, f), AG(9001, 8, b, f) + R.barrier() + if R.rank == 0: + got2, got1 = c2.run(WS_A), c1.run(WS_A) + else: + got1, got2 = c1.run(WS_A), c2.run(WS_A) + torch.cuda.synchronize() + # Position 1 gathered rank 0's c2 rows with the others' c1 rows, position 2 the other way round. + first = with_inputs(c1, [c2.inputs[0]] + c1.inputs[1:]).ref() + second = with_inputs(c2, [c1.inputs[0]] + c2.inputs[1:]).ref() + at1, at2 = (got2, got1) if R.rank == 0 else (got1, got2) + paired = all(bits_equal(a, w) for a, w in zip(at1 + at2, first + second)) + assert paired, "not the position-paired gathers" + right = [ + all(bits_equal(a, w) for a, w in zip(got, c.ref())) for c, got in ((c1, got1), (c2, got2)) + ] + assert R.all_true(not any(right)), f"the swap went unnoticed: results right {right}" + rot(WS_A).advance(c2.record(), "swapped pair", calls=2) + call_and_check(AG(9002, 8, b, f), WS_A, "after the swapped pair") + + +CHECKS = [ + check_workspaces_are_armed_and_sized, + check_single_calls, + check_unsupported_calls_raise_on_every_rank, + check_dip_and_regrow_sequence, + check_two_workspaces_interleaved, + check_one_workspace_three_ops, + check_graph_capture_and_replay, + check_create_under_capture_raises, + # Stays last: it deliberately disagrees on call order. + check_wrong_call_order_is_detected, +] + + +def _run_one_rank(args) -> int: + global R, allgather_split, required_buffer_bytes, fusion_allreduce, MnnvlWorkspace + global BUFFER_BYTES, WS_A, WS_B + R = ls.Rank(args) + assert R.world in (2, 4, 8, 16), f"K3's shapes shard over 2, 4, 8 or 16 ranks, not {R.world}" + from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm import ( + mnnvl_allgather_split as module, + ) + from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm import ( + mnnvl_fusion_allreduce as reduce_module, + ) + + allgather_split = module.mnnvl_allgather_split + required_buffer_bytes = module.required_buffer_bytes + fusion_allreduce = reduce_module.mnnvl_fusion_allreduce + MnnvlWorkspace = module.MnnvlWorkspace + BUFFER_BYTES = buffer_bytes(R.world) + with torch.inference_mode(): + WS_A = MnnvlWorkspace.create(R.mapping, BUFFER_BYTES, fabric_handle=R.fabric) + WS_B = MnnvlWorkspace.create(R.mapping, BUFFER_BYTES, fabric_handle=R.fabric) + for ws in (WS_A, WS_B): + ROT[id(ws)] = Rotation(ws) + code = ls.run_checks(R, CHECKS) + if R.rank == 0: + print( + f"[rank 0] world {R.world}; buffer {BUFFER_BYTES} B; " + f"max all-reduce normed rel err {STATS['normed_err']:.3e}; " + f"max attn_res normed rel err {STATS['attn_res_err']:.3e}", + flush=True, + ) + return code + + +if __name__ == "__main__": + ARGS = ls.parse_args(sys.argv[1:]) + if ARGS.rank_worker or ARGS.launcher == "srun": + sys.exit(_run_one_rank(ARGS)) + ls.spawn(__file__, ARGS, DEADLINE_S) + print("OK") diff --git a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_fusion_allreduce_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_fusion_allreduce_op_matrix.py new file mode 100644 index 000000000000..75057e65e335 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_fusion_allreduce_op_matrix.py @@ -0,0 +1,702 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU certification matrix for the ``comm/mnnvl_fusion_allreduce`` catalog entry on its ``MnnvlWorkspace``. + +The op's correctness depends on state that outlives a call: the workspace's Lamport rotation, which every MNNVL op of +the TP group advances (this op on either path, ``comm/mnnvl_allgather_split`` and ``comm/mnnvl_allreduce_attn_res``). +So beyond single calls (every certified shape sent one-shot and two-shot back to back, the path chosen per call by +``one_shot_max_bytes`` at the exact boundary) this drives call *sequences*: decode steps whose token count dips and +grows back with Kimi K3's one-shot ceilings flipping the path inside the sequence, a random rank late at every call; +two workspaces interleaved; the three MNNVL ops interleaved on one workspace; CUDA-graph capture and replay mixed with +eager calls; and a negative control in which one rank swaps two calls and every rank gets a wrong answer without an +error. After every eager call the workspace's ``buffer_flags`` are compared with this file's model of the rotation +(``Rotation``): one turn per call whatever the op, the path the call took, the bytes it wrote. + + CUDA_VISIBLE_DEVICES=0,1,2,3 python _mnnvl_fusion_allreduce_op_matrix.py [--world-size 4] + srun -N 4 --ntasks-per-node 4 --mpi=pmix python _mnnvl_fusion_allreduce_op_matrix.py --launcher srun \ + --world-size 16 + +Not a pytest module: one fixed sequence of checks inside one W-rank job (they share the workspaces and their +rotations). The collected entry point is ``test_modeling_v2_mnnvl_fusion_allreduce_op_matrix.py``. + +Every rank draws every rank's inputs from one seed, so each rank holds the whole reference. The inputs of the sums +are small multiples of 1/16 (``_lockstep.exact_bf16``), so every sum over the ranks is exact in fp32 and in bf16 +whatever the summation order, and the residual add is one bf16 rounding of an exact fp32 value in the op and in the +reference alike: the sum and ``updated`` are compared bit for bit, ``normed`` (an RMSNorm) against the fp32 reference +within ``TOL``. Every output is also compared bitwise across the ranks. +""" + +import copy +import random +import sys +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parent)) +import _lockstep as ls # noqa: E402 + +assert torch.cuda.is_available(), "mnnvl_fusion_allreduce requires CUDA devices" + +DEADLINE_S = 900 +H_LATENT = 3584 # Kimi K3's routed latent: the MoE all-reduce +H_MODEL = 7168 # Kimi K3's hidden size: the drafter's residual + RMSNorm all-reduces, a wide step's attention ones +HIDDENS = (H_LATENT, H_MODEL) +TOKENS = (1, 2, 3, 4, 5, 6, 7, 8, 16, 32, 64) +EXPERTS = 896 # the router logits the MoE head all-gather carries beside the latent +DECODE_MAX_TOKENS = 8 +# Kimi K3's one-shot ceilings: DECODE_AR_ONE_SHOT_MAX_BYTES on decode steps (at most 8 tokens), +# WIDE_AR_ONE_SHOT_MAX_BYTES (main's default) on wide decode steps. +DECODE_ONE_SHOT_MAX_BYTES = 4 << 20 +WIDE_ONE_SHOT_MAX_BYTES = 1 << 20 +EPS = 1e-5 +# normed: max |err| / max |ref|. The kernels round each square to bf16 before summing (at most 2^-9 on the rsqrt) +# and round the output to bf16 (2^-8); the fp32 reference does neither. +TOL = 1e-2 +ATTN_RES_TOL = 2e-2 # mnnvl_allreduce_attn_res's normed, as its own matrix bounds it +ATTN_RES_MAX_TOKENS = 16 # Kimi K3 sends the attention-residual all-reduce up to 16 tokens +INT16_MIN = torch.iinfo(torch.int16).min # the bf16 -0.0 word +INT32_MIN = torch.iinfo(torch.int32).min # the fp32 -0.0 word, the Lamport buffers' empty word +NEG_ZERO_STRIDE = 37 # every rank's all-reduce input holds -0.0 in these columns +# fp32 words the all-gather carries unchanged: +inf, -inf, NaN, +-the smallest denormal, the largest float. +SPECIAL_WORDS = (0x7F800000, -0x00800000, 0x7FC00000, 1, -0x7FFFFFFF, 0x7F7FFFFF) +LAYERS = 8 +DIP_STEPS = (8, 8, 8, 2, 7, 8, 1, 1, 64, 3, 32, 8, 16, 1, 64, 8) +SHARED_STEPS = (8, 2, 16, 64, 1, 7, 32, 8, 3, 16) +INTERLEAVED = ( # (T, H, fused, path) + (3, H_LATENT, False, "one"), + (8, H_MODEL, True, "two"), + (1, H_MODEL, False, "two"), + (16, H_LATENT, True, "one"), + (64, H_LATENT, False, "two"), + (5, H_MODEL, True, "k3"), + (32, H_MODEL, False, "k3"), +) + +R = None +fusion_allreduce = None +required_buffer_bytes = None +allgather_split = None +MnnvlWorkspace = None +BUFFER_BYTES = None +WS_A = None +WS_B = None +ROT = {} +STATS = {"normed_err": 0.0, "normed_paths": 0.0, "attn_res_err": 0.0} + + +def buffer_bytes(world: int) -> int: + """One Lamport buffer holds the largest certified call sent one-shot, [64, 7168]; every other call is smaller.""" + return max(TOKENS) * H_MODEL * world * 2 + + +def bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16 if t.element_size() == 2 else torch.int32) + + +def bits_equal(a: torch.Tensor, b: torch.Tensor) -> bool: + return a.shape == b.shape and a.dtype == b.dtype and torch.equal(bits(a), bits(b)) + + +def positive_zero(t: torch.Tensor) -> torch.Tensor: + """``t`` with every -0.0 made +0.0, by bit pattern: the Lamport buffers' empty word never travels.""" + words = bits(t) + sign = torch.iinfo(words.dtype).min + return torch.where(words == sign, torch.zeros_like(words), words).view(t.dtype) + + +def one_shot_bytes(tokens: int, hidden: int) -> int: + return tokens * hidden * R.world * 2 + + +def ceiling(tokens: int, hidden: int, path: str) -> int: + """The call's one_shot_max_bytes: at the boundary for "one" (the call's one-shot footprint itself, so one-shot) + and "two" (one byte less, so two-shot), or Kimi K3's ceiling for a step of ``tokens`` ("k3").""" + if path == "one": + return one_shot_bytes(tokens, hidden) + if path == "two": + return one_shot_bytes(tokens, hidden) - 1 + return DECODE_ONE_SHOT_MAX_BYTES if tokens <= DECODE_MAX_TOKENS else WIDE_ONE_SHOT_MAX_BYTES + + +def k3_split(): + """Kimi K3's MoE head shard at this world size: (latent columns, router-logit columns) per rank.""" + return H_LATENT // R.world, EXPERTS // R.world + + +class Rotation: + """This test's model of a workspace's ``buffer_flags`` (``mnnvl_workspace.FLAG_WORDS``): the buffer the next call + takes, the one the last call used, the bytes per buffer, the last call's stage count and the bytes it wrote per + stage (what the next call clears), and the arrival counter (0 between calls).""" + + def __init__(self, ws): + self.ws = ws + self.flags = [0, 2, ws.buffer_bytes, 0, 0, 0, 0, 0, 0] # as create() arms them + + def read(self): + torch.cuda.synchronize() + return self.ws.buffer_flags.view(torch.int32).tolist() + + def unchanged(self, where: str) -> None: + got = self.read() + assert got == self.flags, f"{where}: buffer_flags {got} moved from {self.flags}" + + def advance(self, record, where: str, calls: int = 1) -> None: + """``calls`` more calls ran on the workspace, the last of which recorded ``record`` = (stages, bytes).""" + stages, written = record + current = (self.flags[0] + calls) % 3 + self.flags = [current, (current - 1) % 3, self.ws.buffer_bytes, stages, *written, 0] + got = self.read() + assert got == self.flags, f"{where}: buffer_flags {got}, expected {self.flags}" + + +def rot(ws) -> Rotation: + return ROT[id(ws)] + + +class AR: + """One mnnvl_fusion_allreduce call's arguments on every rank and its reference. ``residual``: a tensor (chained + from an earlier call), True (drawn) or None (the plain sum). ``path``: see ``ceiling``. Every rank's input holds + -0.0 in every NEG_ZERO_STRIDE-th column.""" + + def __init__(self, seed, tokens, hidden, residual=None, path="k3"): + g = torch.Generator(device="cuda").manual_seed(seed) + self.tokens, self.hidden, self.path = tokens, hidden, path + self.one_shot_max_bytes = ceiling(tokens, hidden, path) + self.inputs = [] + for _ in range(R.world): + x = ls.exact_bf16(g, (tokens, hidden), -4, 5, 1 / 16) + x.view(torch.int16)[:, ::NEG_ZERO_STRIDE] = INT16_MIN + self.inputs.append(x) + if residual is True: + residual = ls.exact_bf16(g, (tokens, hidden), -32, 33, 1 / 16) + self.residual = residual + self.gamma = None + if residual is not None: + self.gamma = (1.0 + 0.1 * torch.randn(hidden, generator=g, device="cuda")).bfloat16() + + @property + def fused(self) -> bool: + return self.residual is not None + + @property + def one_shot(self) -> bool: + return one_shot_bytes(self.tokens, self.hidden) <= self.one_shot_max_bytes + + def fresh(self, seed): + """A call of the same shape, fusion and path with new inputs.""" + return AR( + seed, self.tokens, self.hidden, residual=True if self.fused else None, path=self.path + ) + + def run(self, ws, x=None, residual=None): + x = self.inputs[R.rank] if x is None else x + if not self.fused: + return fusion_allreduce(x, ws, self.one_shot_max_bytes) + residual = self.residual if residual is None else residual + return fusion_allreduce(x, ws, self.one_shot_max_bytes, residual, self.gamma, EPS) + + def ref(self): + """The exact sum (+0.0 where every rank sent -0.0); with a residual ``(normed, updated)``: ``updated`` exact, + ``normed`` the fp32 RMSNorm of it.""" + total = positive_zero(torch.stack([x.float() for x in self.inputs]).sum(dim=0)) + if not self.fused: + return total.bfloat16() + updated = (total + self.residual.float()).bfloat16() + x = updated.float() + normed = x * torch.rsqrt(x.square().mean(dim=-1, keepdim=True) + EPS) * self.gamma.float() + return normed.bfloat16(), updated + + def verify(self, got, where: str) -> None: + want = self.ref() + if not self.fused: + assert bits_equal(got, want), f"{where}: the sum differs from the exact reference" + assert R.same_on_ranks(got), f"{where}: ranks disagree" + return + (normed, updated), (want_normed, want_updated) = got, want + assert bits_equal(updated, want_updated), ( + f"{where}: updated differs from the exact reference" + ) + err = ls.rel_err(normed, want_normed) + STATS["normed_err"] = max(STATS["normed_err"], err) + assert err <= TOL, f"{where}: normed rel err {err:.3e} > {TOL}" + assert R.same_on_ranks(normed, updated), f"{where}: ranks disagree" + + def record(self): + """What the call leaves in buffer_flags: (stages, bytes it wrote into each).""" + if self.one_shot: + return 1, (one_shot_bytes(self.tokens, self.hidden), 0, 0, 0) + rows = -(-self.tokens // R.world) * R.world + return 2, (rows * self.hidden * 2, self.tokens * self.hidden * 2, 0, 0) + + def static(self): + static = {"x": self.inputs[R.rank].clone()} + if self.fused: + static["residual"] = self.residual.clone() + return static + + def refill(self, fresh, static) -> None: + """Take ``fresh``'s inputs (same shape) into this call and its static buffers.""" + self.inputs = fresh.inputs + static["x"].copy_(fresh.inputs[R.rank]) + if self.fused: + self.residual = fresh.residual + static["residual"].copy_(fresh.residual) + + +def gather_payload(g, tokens, bf16_columns, fp32_columns): + """One rank's fp32 rows for the all-gather: normal values over exponents 2^-16..2^16 in the bf16 columns + (inexact in bf16), exact midpoints between two bf16 values in every fourth of them (ties go to even), -0.0 in + every fourth column of each part, and special words (infinities, NaN, denormals, the largest float) in the first + row of the fp32 part.""" + b = bf16_columns + x = torch.randn(tokens, b + fp32_columns, generator=g, device="cuda") + x[:, :b] *= torch.exp2(((torch.arange(b, device="cuda") % 5) - 2).float() * 8) + nearest = x[:, 1:b:4].bfloat16().float() + x[:, 1:b:4] = (nearest.view(torch.int32) | 0x8000).view(torch.float32) + words = x.view(torch.int32) + words[:, 3:b:4] = INT32_MIN + words[:, b + 2 :: 4] = INT32_MIN + special = SPECIAL_WORDS[: min(len(SPECIAL_WORDS), fp32_columns)] + words[0, b : b + len(special)] = torch.tensor(special, dtype=torch.int32, device="cuda") + return x + + +class AG: + """One comm/mnnvl_allgather_split call's arguments on every rank and its reference.""" + + def __init__(self, seed, tokens, bf16_columns, fp32_columns): + g = torch.Generator(device="cuda").manual_seed(seed) + self.tokens, self.bf16_columns, self.fp32_columns = tokens, bf16_columns, fp32_columns + self.inputs = [ + gather_payload(g, tokens, bf16_columns, fp32_columns) for _ in range(R.world) + ] + + def fresh(self, seed): + return AG(seed, self.tokens, self.bf16_columns, self.fp32_columns) + + def run(self, ws, x=None): + return allgather_split(self.inputs[R.rank] if x is None else x, self.bf16_columns, ws) + + def ref(self): + b = self.bf16_columns + bf16_out = torch.cat([x[:, :b] for x in self.inputs], dim=1).bfloat16() + fp32_out = torch.cat([x[:, b:] for x in self.inputs], dim=1) + return positive_zero(bf16_out), positive_zero(fp32_out) + + def verify(self, got, where: str) -> None: + want = self.ref() + assert all(bits_equal(a, b) for a, b in zip(got, want)), ( + f"{where}: the gather differs from the reference" + ) + assert R.same_on_ranks(*got), f"{where}: ranks disagree" + + def record(self): + written = self.tokens * R.world * (2 * self.bf16_columns + 4 * self.fp32_columns) + return 1, (written, 0, 0, 0) + + def static(self): + return {"x": self.inputs[R.rank].clone()} + + def refill(self, fresh, static) -> None: + self.inputs = fresh.inputs + static["x"].copy_(fresh.inputs[R.rank]) + + +class AttnRes: + """One comm/mnnvl_allreduce_attn_res call (G1's entry, called through its op on the workspace's comm_buffer and + buffer_flags, as that entry's wrapper does) and its reference. ``prefix``: a tensor (chained) or True (drawn).""" + + def __init__(self, seed, tokens, snapshots, prefix=True): + g = torch.Generator(device="cuda").manual_seed(seed) + self.tokens, self.snapshots = tokens, snapshots + self.inputs = [ls.exact_bf16(g, (tokens, H_MODEL), -4, 5, 1 / 16) for _ in range(R.world)] + if prefix is True: + prefix = ls.exact_bf16(g, (tokens, H_MODEL), -32, 33, 1 / 16) + self.prefix = prefix + self.block = torch.randn(snapshots, tokens, H_MODEL, generator=g, device="cuda").bfloat16() + self.res_w = (torch.randn(H_MODEL, generator=g, device="cuda") * 0.05).bfloat16() + self.rms_w = (1.0 + 0.1 * torch.randn(H_MODEL, generator=g, device="cuda")).bfloat16() + self.out_w = (1.0 + 0.1 * torch.randn(H_MODEL, generator=g, device="cuda")).bfloat16() + + def fresh(self, seed): + return AttnRes(seed, self.tokens, self.snapshots) + + def run(self, ws, x=None, prefix=None, block=None): + normed, updated = torch.ops.trtllm.mnnvl_allreduce_attn_res( + self.inputs[R.rank] if x is None else x, + self.prefix if prefix is None else prefix, + self.block if block is None else block, + self.res_w, + self.rms_w, + self.out_w, + EPS, + EPS, + ws.comm_buffer(torch.bfloat16), + ws.buffer_flags, + ) + return normed, updated + + def ref(self): + updated = (sum(x.float() for x in self.inputs) + self.prefix.float()).bfloat16() + normed = ls.residual_update_ref( + updated, self.block, self.res_w, self.rms_w, EPS, self.out_w, EPS + ) + return normed, updated + + def verify(self, got, where: str) -> None: + (normed, updated), (want_normed, want_updated) = got, self.ref() + assert bits_equal(updated, want_updated), ( + f"{where}: attn_res updated differs from the exact sum" + ) + err = ls.rel_err(normed, want_normed) + STATS["attn_res_err"] = max(STATS["attn_res_err"], err) + assert err <= ATTN_RES_TOL, f"{where}: attn_res normed rel err {err:.3e} > {ATTN_RES_TOL}" + assert R.same_on_ranks(normed, updated), f"{where}: ranks disagree" + + def record(self): + return 1, (self.tokens * H_MODEL * R.world * 2, 0, 0, 0) + + def static(self): + return { + "x": self.inputs[R.rank].clone(), + "prefix": self.prefix.clone(), + "block": self.block.clone(), + } + + def refill(self, fresh, static) -> None: + self.inputs, self.prefix, self.block = fresh.inputs, fresh.prefix, fresh.block + static["x"].copy_(fresh.inputs[R.rank]) + static["prefix"].copy_(fresh.prefix) + static["block"].copy_(fresh.block) + + +def call_and_check(call, ws, where: str, late=None): + """Run ``call`` eagerly on ``ws`` (rank ``late``, if given, 5 ms late), check its result and what it left in the + workspace's flags; return the result.""" + if late is not None: + R.barrier() + R.late(late) + got = call.run(ws) + call.verify(got, where) + rot(ws).advance(call.record(), where) + return got + + +def call_late(call, ws, where: str, late_rng: random.Random): + """``call_and_check`` with a rank drawn from ``late_rng`` (the same draw on every rank) 5 ms late.""" + return call_and_check(call, ws, where, late=late_rng.randrange(R.world)) + + +def expect_refusal(fn, error, where: str) -> None: + """``fn`` raises ``error`` on every rank and WS_A's flags do not move.""" + try: + fn() + raised = False + except error: + raised = True + assert R.all_true(raised), f"{where}: not refused with {error.__name__} on every rank" + rot(WS_A).unchanged(where) + + +def with_inputs(call, inputs): + """``call`` as if the ranks had sent ``inputs`` (one tensor per rank).""" + mixed = copy.copy(call) + mixed.inputs = inputs + return mixed + + +def check_workspaces_are_armed_and_sized() -> None: + """create() armed both workspaces (every Lamport word -0.0; flags at buffer 0, buffer 2 dirty with nothing to + clear), and one buffer holds the largest certified call sent one-shot and every other call of this matrix.""" + for ws in (WS_A, WS_B): + assert ws.world_size == R.world and ws.rank == R.rank + assert ws.buffer_bytes == BUFFER_BYTES and BUFFER_BYTES % 32 == 0 + assert ws.comm_buffer(torch.bfloat16).shape == (3, BUFFER_BYTES // 2) + armed = ws.lamport.view(torch.int32) + assert bool((armed == torch.tensor(INT32_MIN, dtype=torch.int32, device="cuda")).all()), ( + "every word -0.0" + ) + rot(ws).unchanged("armed") + one_shot = required_buffer_bytes(max(TOKENS), H_MODEL, R.world, torch.bfloat16, 1 << 62) + assert one_shot == BUFFER_BYTES, f"the largest one-shot call needs {one_shot} bytes" + b, f = k3_split() + assert max(TOKENS) * R.world * (2 * b + 4 * f) <= BUFFER_BYTES + assert ATTN_RES_MAX_TOKENS * H_MODEL * R.world * 2 <= BUFFER_BYTES + + +def check_single_calls_both_paths() -> None: + """Every certified shape, plain and fused, sent one-shot with one_shot_max_bytes = its one-shot footprint (the + boundary is inclusive) and right after two-shot with one byte less, on one workspace, the same inputs both + times: each against the reference, the flags after each recording the path and the bytes the formula says, and + required_buffer_bytes equal to the space the call's stages take.""" + for hidden in HIDDENS: + for t in TOKENS: + for fused in (False, True): + outs = [] + for path in ("one", "two"): + call = AR( + 1000 + 37 * t + 3 * hidden + int(fused), + t, + hidden, + residual=fused or None, + path=path, + ) + where = f"T {t} H {hidden} fused {fused} {path}-shot" + assert call.one_shot == (path == "one"), where + outs.append(call_and_check(call, WS_A, where)) + stages, written = call.record() + need = required_buffer_bytes( + t, hidden, R.world, torch.bfloat16, call.one_shot_max_bytes + ) + assert need == (written[0] if stages == 1 else 2 * written[0]), ( + f"{where}: needs {need}" + ) + if fused: + STATS["normed_paths"] = max( + STATS["normed_paths"], ls.rel_err(outs[0][0], outs[1][0]) + ) + + +def check_unsupported_calls_raise_on_every_rank() -> None: + """Refused on every rank before the workspace is touched (its flags do not move), and the next call is correct: + more tokens than one Lamport buffer holds one-shot (the wrapper's ValueError; the same rows sent two-shot fit, + above 2 ranks), a residual without norm_weight and eps (the wrapper's ValueError), a hidden size not a multiple + of 8 (the op's RuntimeError).""" + t = max(TOKENS) + 1 + over = AR(2000, t, H_MODEL, path="one") + expect_refusal(lambda: over.run(WS_A), ValueError, f"T {t} one-shot over one buffer") + if R.world > 2: + # Two-shot takes about 2 / W of the one-shot space. + call_and_check(AR(2000, t, H_MODEL, path="two"), WS_A, f"T {t} sent two-shot") + pair = AR(2001, 4, H_LATENT, residual=True) + expect_refusal( + lambda: fusion_allreduce(pair.inputs[R.rank], WS_A, pair.one_shot_max_bytes, pair.residual), + ValueError, + "a residual without norm_weight and eps", + ) + odd = AR(2002, 2, H_LATENT - 4) + expect_refusal(lambda: odd.run(WS_A), RuntimeError, f"hidden {H_LATENT - 4}") + call_and_check(AR(2003, 8, H_MODEL, residual=True), WS_A, "after the refused calls") + + +def run_step(ws, seed, tokens, late): + """One decode step of LAYERS layers, a random rank late at every call: per layer the routed-latent all-reduce + [T, 3584] and the fused all-reduce [T, 7168] chained through ``updated``, both at Kimi K3's ceiling.""" + residual = ls.exact_bf16( + torch.Generator(device="cuda").manual_seed(seed), (tokens, H_MODEL), -32, 33, 1 / 16 + ) + for layer in range(LAYERS): + where = f"step seed {seed} T {tokens} layer {layer}" + call_late(AR(seed + 2 * layer + 1, tokens, H_LATENT), ws, f"{where} latent", late) + fused = AR(seed + 2 * layer + 2, tokens, H_MODEL, residual=residual) + residual = call_late(fused, ws, f"{where} fused", late)[1] + + +def check_dip_and_regrow_sequence() -> None: + """The k3_spec_accept failure mode: a call after a smaller one must not read what an older, larger call left. + 16 steps of 8 layers at T 8, 8, 8, 2, 7, 8, 1, 1, 64, 3, 32, 8, 16, 1, 64, 8 at Kimi K3's ceilings, so the path + flips inside the sequence (at W = 4 [64, 3584] and [32 or 64, 7168] go two-shot, at W = 16 every step above 8 + tokens), a random rank late at every call.""" + late = random.Random(7) + for i, t in enumerate(DIP_STEPS): + run_step(WS_A, 3000 + 100 * i, t, late) + + +def check_two_workspaces_interleaved() -> None: + """Two workspaces are two rotations: calls of mixed shapes and paths alternate between them in an irregular + pattern (A A B A B B ...), every call is correct and each workspace's flags move with its own calls only. The + pattern is the same on every rank: calls on one stream are serialized and each waits for its peers, so ranks + issuing calls on two workspaces in different orders deadlock (measured for mnnvl_allreduce_attn_res, + runs/drafter/u4-mnnvl-srun-2).""" + pattern = "AABABBAAAB" * 2 + for i, which in enumerate(pattern): + t, hidden, fused, path = INTERLEAVED[i % len(INTERLEAVED)] + call = AR(4700 + i, t, hidden, residual=fused or None, path=path) + call_and_check(call, WS_A if which == "A" else WS_B, f"interleaved {which} {i}") + + +def check_one_workspace_three_ops() -> None: + """The three MNNVL entries on one workspace in Kimi K3's order, a random rank late at every call. A step of at + most 16 tokens: per layer the pre-attention all-reduce with the attention-residual epilogue + (comm/mnnvl_allreduce_attn_res, prefix chained), the MoE head all-gather (comm/mnnvl_allgather_split, K3's + split), the routed-latent all-reduce and the fused all-reduce (chained). A wide step (32, 64 tokens): the wide + all-reduce [T, 7168], the all-gather and the routed-latent all-reduce, at the 1 MiB ceiling. Every call takes one + turn of the rotation whatever the op (the flags after each), and every result is right.""" + late = random.Random(11) + b, f = k3_split() + for i, t in enumerate(SHARED_STEPS): + seed = 5000 + 100 * i + g = torch.Generator(device="cuda").manual_seed(seed) + prefix = ls.exact_bf16(g, (t, H_MODEL), -32, 33, 1 / 16) + residual = ls.exact_bf16(g, (t, H_MODEL), -32, 33, 1 / 16) + for layer in range(2): + s = seed + 10 * layer + 1 + where = f"shared step {i} T {t} layer {layer}" + if t <= ATTN_RES_MAX_TOKENS: + attn = AttnRes(s, t, (0, 2, 5)[(i + layer) % 3], prefix=prefix) + prefix = call_late(attn, WS_A, f"{where} attn_res", late)[1] + else: + call_late(AR(s, t, H_MODEL), WS_A, f"{where} wide", late) + call_late(AG(s + 1, t, b, f), WS_A, f"{where} all-gather", late) + call_late(AR(s + 2, t, H_LATENT), WS_A, f"{where} latent", late) + if t <= ATTN_RES_MAX_TOKENS: + fused = AR(s + 3, t, H_MODEL, residual=residual) + residual = call_late(fused, WS_A, f"{where} fused", late)[1] + + +def check_graph_capture_and_replay() -> None: + """A captured step of six calls on WS_B: the attention-residual all-reduce, the plain and fused all-reduces + one-shot (T 8, Kimi K3's ceiling), the head all-gather, a plain [32, 7168] and a fused [16, 7168] sent two-shot + (the fused two-shot call is two kernels, the exchange and the residual + RMSNorm one; both are in the graph). + Replayed 8 times with rewritten inputs, an eager call of another shape or op on the same workspace between + replays: replays and eager calls take turns of one rotation (the flags after each), in the same order on every + rank, and every replayed and eager result is right.""" + b, f = k3_split() + calls = [ + AttnRes(6000, 8, 2), + AR(6001, 8, H_LATENT), + AG(6002, 8, b, f), + AR(6003, 8, H_MODEL, residual=True), + AR(6004, 32, H_MODEL, path="two"), + AR(6005, 16, H_MODEL, residual=True, path="two"), + ] + statics = [c.static() for c in calls] + + def step(): + return [c.run(WS_B, **s) for c, s in zip(calls, statics)] + + outs = step() # every call once eagerly, outside capture + for i, (c, got) in enumerate(zip(calls, outs)): + c.verify(got, f"eager step call {i}") + rot(WS_B).advance(calls[-1].record(), "eager step", calls=len(calls)) + R.barrier() + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + outs = step() + R.barrier() + rot(WS_B).unchanged("captured, not replayed") + eager = ( + lambda seed: AR(seed, 3, H_LATENT, path="one"), + lambda seed: AR(seed, 64, H_MODEL, residual=True, path="two"), + lambda seed: AG(seed, 5, b, f), + lambda seed: AR(seed, 64, H_LATENT, path="one"), + ) + for rep in range(8): + for i, (c, s) in enumerate(zip(calls, statics)): + c.refill(c.fresh(7000 + 100 * rep + i), s) + R.barrier() + graph.replay() + for i, (c, got) in enumerate(zip(calls, outs)): + c.verify(got, f"replay {rep} call {i}") + rot(WS_B).advance(calls[-1].record(), f"replay {rep}", calls=len(calls)) + call_and_check(eager[rep % len(eager)](7500 + rep), WS_B, f"eager after replay {rep}") + del graph + + +def check_create_under_capture_raises() -> None: + """MnnvlWorkspace.create allocates and exchanges handles, so it refuses CUDA-graph capture: with every rank + capturing, it raises RuntimeError on every rank (before any communication), and the next call on a workspace in + use is correct.""" + graph, stream = torch.cuda.CUDAGraph(), torch.cuda.Stream() + raised = False + with torch.cuda.graph(graph, stream=stream): + try: + MnnvlWorkspace.create(R.mapping, BUFFER_BYTES, fabric_handle=R.fabric) + except RuntimeError as exc: + raised = "before capture" in str(exc) + del graph + assert R.all_true(raised), "create() under capture did not raise on every rank" + call_and_check(AR(8500, 8, H_MODEL, residual=True), WS_A, "after the refused create") + + +def check_wrong_call_order_is_detected() -> None: + """Negative control: rank 0 issues two same-shaped calls on one workspace in swapped order, once sent one-shot + and once two-shot. Nothing raises and nothing hangs (the two calls write the same slots of the same buffers), but + every rank's two results are wrong, and exactly the position-paired sums: the k-th call on every rank adds what + every rank sent at position k. A plain call right after is correct again: a swapped pair realigns the + positions.""" + for path in ("one", "two"): + c1 = AR(9000, 8, H_LATENT, path=path) + c2 = AR(9001, 8, H_LATENT, path=path) + R.barrier() + if R.rank == 0: + got2, got1 = c2.run(WS_A), c1.run(WS_A) + else: + got1, got2 = c1.run(WS_A), c2.run(WS_A) + torch.cuda.synchronize() + # Position 1 summed rank 0's c2 with the others' c1, position 2 rank 0's c1 with the others' c2. + first = with_inputs(c1, [c2.inputs[0]] + c1.inputs[1:]).ref() + second = with_inputs(c2, [c1.inputs[0]] + c2.inputs[1:]).ref() + at1, at2 = (got2, got1) if R.rank == 0 else (got1, got2) + assert bits_equal(at1, first) and bits_equal(at2, second), ( + f"{path}-shot: not the position-paired sums" + ) + wrong = [ + (bits(got) != bits(c.ref())).float().mean().item() + for c, got in ((c1, got1), (c2, got2)) + ] + assert R.all_true(min(wrong) > 0.5), ( + f"{path}-shot: the swap went unnoticed: wrong fractions {wrong}" + ) + rot(WS_A).advance(c2.record(), f"{path}-shot swapped pair", calls=2) + call_and_check( + AR(9002, 8, H_LATENT, path=path), WS_A, f"{path}-shot after the swapped pair" + ) + + +CHECKS = [ + check_workspaces_are_armed_and_sized, + check_single_calls_both_paths, + check_unsupported_calls_raise_on_every_rank, + check_dip_and_regrow_sequence, + check_two_workspaces_interleaved, + check_one_workspace_three_ops, + check_graph_capture_and_replay, + check_create_under_capture_raises, + # Stays last: it deliberately disagrees on call order. + check_wrong_call_order_is_detected, +] + + +def _run_one_rank(args) -> int: + global R, fusion_allreduce, required_buffer_bytes, allgather_split, MnnvlWorkspace + global BUFFER_BYTES, WS_A, WS_B + R = ls.Rank(args) + assert R.world in (2, 4, 8, 16), f"K3's shapes shard over 2, 4, 8 or 16 ranks, not {R.world}" + from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm import ( + mnnvl_allgather_split as gather_module, + ) + from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm import ( + mnnvl_fusion_allreduce as module, + ) + + fusion_allreduce = module.mnnvl_fusion_allreduce + required_buffer_bytes = module.required_buffer_bytes + allgather_split = gather_module.mnnvl_allgather_split + MnnvlWorkspace = module.MnnvlWorkspace + BUFFER_BYTES = buffer_bytes(R.world) + with torch.inference_mode(): + WS_A = MnnvlWorkspace.create(R.mapping, BUFFER_BYTES, fabric_handle=R.fabric) + WS_B = MnnvlWorkspace.create(R.mapping, BUFFER_BYTES, fabric_handle=R.fabric) + for ws in (WS_A, WS_B): + ROT[id(ws)] = Rotation(ws) + code = ls.run_checks(R, CHECKS) + if R.rank == 0: + print( + f"[rank 0] world {R.world}; buffer {BUFFER_BYTES} B; max normed rel err {STATS['normed_err']:.3e}; " + f"max normed one-shot vs two-shot {STATS['normed_paths']:.3e}; " + f"max attn_res normed rel err {STATS['attn_res_err']:.3e}", + flush=True, + ) + return code + + +if __name__ == "__main__": + ARGS = ls.parse_args(sys.argv[1:]) + if ARGS.rank_worker or ARGS.launcher == "srun": + sys.exit(_run_one_rank(ARGS)) + ls.spawn(__file__, ARGS, DEADLINE_S) + print("OK") diff --git a/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_latent_reduce_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_latent_reduce_op_matrix.py new file mode 100644 index 000000000000..2ca5da8b834f --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_latent_reduce_op_matrix.py @@ -0,0 +1,23 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Collected entry point for the k3_latent_reduce op's certification matrix. + +The matrix is ``_k3_latent_reduce_op_matrix.py`` beside this file, its own W-rank launcher (call sequences over +caller-owned exchanges, so one job, not independent cases); see ``_rank_job`` for why that is left intact. +""" + +import _rank_job +import pytest +import torch + +assert torch.cuda.is_available(), "k3_latent_reduce requires CUDA devices" + +if torch.cuda.get_device_capability() != (10, 0): + # The entry is certified on sm_100 (GB200) only; see its contract's receipts. + pytest.skip("k3_latent_reduce is certified on sm_100 only", allow_module_level=True) + + +# Each case starts its own W-rank mpirun over the visible devices; under xdist several workers would fight for them. +@pytest.mark.no_xdist +def test_k3_latent_reduce_op_matrix() -> None: + _rank_job.run("k3_latent_reduce") diff --git a/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_sandwich_oproj_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_sandwich_oproj_op_matrix.py new file mode 100644 index 000000000000..582cfa6d5578 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_sandwich_oproj_op_matrix.py @@ -0,0 +1,23 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Collected entry point for the k3_sandwich_oproj op's certification matrix. + +The matrix is ``_k3_sandwich_oproj_op_matrix.py`` beside this file, its own W-rank launcher (call sequences over one +caller-owned workspace, so one job, not independent cases); see ``_rank_job`` for why that is left intact. +""" + +import _rank_job +import pytest +import torch + +assert torch.cuda.is_available(), "k3_sandwich_oproj requires CUDA devices" + +if torch.cuda.get_device_capability() != (10, 0): + # The entry is certified on sm_100 (GB200) only; see its contract's receipts. + pytest.skip("k3_sandwich_oproj is certified on sm_100 only", allow_module_level=True) + + +# Each case starts its own W-rank mpirun over the visible devices; under xdist several workers would fight for them. +@pytest.mark.no_xdist +def test_k3_sandwich_oproj_op_matrix() -> None: + _rank_job.run("k3_sandwich_oproj") diff --git a/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_sandwich_plain_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_sandwich_plain_op_matrix.py new file mode 100644 index 000000000000..0737f54c3ebe --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_sandwich_plain_op_matrix.py @@ -0,0 +1,23 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Collected entry point for the k3_sandwich_plain op's certification matrix. + +The matrix is ``_k3_sandwich_plain_op_matrix.py`` beside this file, its own W-rank launcher (call sequences over one +caller-owned workspace, so one job, not independent cases); see ``_rank_job`` for why that is left intact. +""" + +import _rank_job +import pytest +import torch + +assert torch.cuda.is_available(), "k3_sandwich_plain requires CUDA devices" + +if torch.cuda.get_device_capability() != (10, 0): + # The entry is certified on sm_100 (GB200) only; see its contract's receipts. + pytest.skip("k3_sandwich_plain is certified on sm_100 only", allow_module_level=True) + + +# Each case starts its own W-rank mpirun over the visible devices; under xdist several workers would fight for them. +@pytest.mark.no_xdist +def test_k3_sandwich_plain_op_matrix() -> None: + _rank_job.run("k3_sandwich_plain") diff --git a/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_sandwich_tail_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_sandwich_tail_op_matrix.py new file mode 100644 index 000000000000..06110f2cee92 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_sandwich_tail_op_matrix.py @@ -0,0 +1,23 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Collected entry point for the k3_sandwich_tail op's certification matrix. + +The matrix is ``_k3_sandwich_tail_op_matrix.py`` beside this file, its own W-rank launcher (call sequences over one +caller-owned workspace, so one job, not independent cases); see ``_rank_job`` for why that is left intact. +""" + +import _rank_job +import pytest +import torch + +assert torch.cuda.is_available(), "k3_sandwich_tail requires CUDA devices" + +if torch.cuda.get_device_capability() != (10, 0): + # The entry is certified on sm_100 (GB200) only; see its contract's receipts. + pytest.skip("k3_sandwich_tail is certified on sm_100 only", allow_module_level=True) + + +# Each case starts its own W-rank mpirun over the visible devices; under xdist several workers would fight for them. +@pytest.mark.no_xdist +def test_k3_sandwich_tail_op_matrix() -> None: + _rank_job.run("k3_sandwich_tail") diff --git a/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_allgather_split_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_allgather_split_op_matrix.py new file mode 100644 index 000000000000..ae0e6db01dd7 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_allgather_split_op_matrix.py @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Collected entry point for the mnnvl_allgather_split op's certification matrix. + +The matrix is ``_mnnvl_allgather_split_op_matrix.py`` beside this file, its own W-rank launcher (call sequences +over caller-owned workspaces, so one job, not independent cases); see ``_rank_job`` for why that is left intact. +""" + +import _rank_job +import pytest +import torch + +assert torch.cuda.is_available(), "mnnvl_allgather_split requires CUDA devices" + + +# Each case starts its own W-rank mpirun over the visible devices; under xdist several workers would fight for them. +@pytest.mark.no_xdist +def test_mnnvl_allgather_split_op_matrix() -> None: + _rank_job.run("mnnvl_allgather_split") diff --git a/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_fusion_allreduce_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_fusion_allreduce_op_matrix.py new file mode 100644 index 000000000000..b0fd79635674 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_fusion_allreduce_op_matrix.py @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Collected entry point for the mnnvl_fusion_allreduce op's certification matrix. + +The matrix is ``_mnnvl_fusion_allreduce_op_matrix.py`` beside this file, its own W-rank launcher (call sequences +over caller-owned workspaces, so one job, not independent cases); see ``_rank_job`` for why that is left intact. +""" + +import _rank_job +import pytest +import torch + +assert torch.cuda.is_available(), "mnnvl_fusion_allreduce requires CUDA devices" + + +# Each case starts its own W-rank mpirun over the visible devices; under xdist several workers would fight for them. +@pytest.mark.no_xdist +def test_mnnvl_fusion_allreduce_op_matrix() -> None: + _rank_job.run("mnnvl_fusion_allreduce") diff --git a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe.py b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe.py new file mode 100644 index 000000000000..f38a1dadd2fe --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe.py @@ -0,0 +1,870 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the moe/k3_moe catalog entry: k3_moe and k3_moe_wide on their caller-owned state. + +One GPU (sm_100), one rank of the Kimi K3 TP16 deployment's routed experts: experts TP4 x EP4, so 224 of the 896 +experts are local (global ids [224, 448)) with intermediate 768 per rank. The experts are random checkpoint-format +MXFP4 tensors put through TRT-LLM's own W4A8_MXFP4_MXFP8 TRTLLM-Gen loader, which writes the buffers that both k3_moe +and the stock runner read. + +References: the stock path (trtllm::kimi_k3_noaux_tc_mxfp8_quant, then the TRTLLM-Gen W4A8_MXFP4_MXFP8 MoE runner on +those ids) and an fp64 reference over the dequantized MXFP4 experts (SiTU, the MXFP8 intermediate with the round-up +scale, the down projection, the routing-weighted sum), both under the op-catalog gates: 8 bf16 ulp of the token row's +largest magnitude per element, 4 ulp relative RMS. Everything about state is compared bit for bit. + +The checks run in file order and share the state objects: a K3MoeState with layers A, B, C (A and C over this rank's +experts, B over other experts), a second K3MoeState, a K3MoeWideState with layers A and B. Each check makes its own +first (compiling) call eagerly where it needs one, so it also runs alone (-k). The kernel tests under +tests/unittest/_torch/cute_dsl_kernels/kimi_k3/ remain the exhaustive numerics; this file copies what it needs. + +k3_moe_fused_front needs the TP group's head workspace: its cells run in moe/k3_moe_front's 4-rank matrix, +tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py. +""" + +import functools +import math +from types import SimpleNamespace + +import pytest +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 -- registers the stock path's ops +from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.k3_moe import ( + K3MoeState, + K3MoeWideState, + is_supported, + k3_moe, + k3_moe_wide, +) + +H, TOP_K, NUM_EXPERTS, SV = 3584, 16, 896, 32 +I_TP, E_LOCAL, MOE_TP, TP_RANK, EP_RANK = 768, 224, 4, 1, 1 # one rank of experts TP4 x EP4 +OFFSET = EP_RANK * E_LOCAL # this rank's experts: global ids [224, 448) +# The routed experts' SiTU caps: k3_moe's build constants, Kimi K3's values. +GATE_CAP, LINEAR_CAP = 4.0, 25.0 +RSF = 2.827 +ULP = 2.0**-8 +E4M3_MAX = 448.0 +DECODE_MAX, WIDE_MAX = 8, 64 +DEV = "cuda" +DECODE_CASES = ("random", "16_local_disjoint", "none_local") +WIDE_CASES = ("random", "group_cap", "none_local") +WIDE_M = (1, 2, 7, 8, 9, 16, 33, 40, 64) +ROUTE_M = (1, 2, 3, 4, 5, 6, 7, 8, 16, 33, 64) +DIP_STEPS = (8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8) +WIDE_DIP_STEPS = (64, 64, 16, 1, 40, 64, 8, 23, 64, 9, 64) +REPLAYS = 6 +_E2M1 = [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0, -0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0] +_CASE_SALT = {"random": 0, "16_local_disjoint": 1, "none_local": 2, "group_cap": 3} + + +def _is_sm100() -> bool: + return torch.cuda.is_available() and torch.cuda.get_device_capability() == (10, 0) + + +pytestmark = pytest.mark.skipif(not _is_sm100(), reason="k3_moe needs sm_100") + + +@pytest.fixture(autouse=True) +def _inference_mode(): + """Run every check as the model runs these ops, under inference mode.""" + with torch.inference_mode(): + yield + + +# ── operands ────────────────────────────────────────────────────────────── + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.uint8) + + +def _same(a: torch.Tensor, b: torch.Tensor) -> bool: + """Same shape, dtype and bits.""" + return a.shape == b.shape and a.dtype == b.dtype and torch.equal(_bits(a), _bits(b)) + + +def _device() -> torch.device: + return torch.device(DEV, torch.cuda.current_device()) + + +def _rand_mxfp4(rows, k, gen): + """Random checkpoint-format MXFP4. + + Packed [rows, k / 2] (low nibble = even k) and one E8M0 per 32 k, scaled so a k-long dot product lands near std 3. + """ + codes = torch.randint(0, 16, (rows, k), dtype=torch.uint8, device=DEV, generator=gen) + base = 127 + round(0.5 * math.log2(0.01057 / k)) + exps = torch.randint( + base, base + 6, (rows, k // SV), dtype=torch.uint8, device=DEV, generator=gen + ) + return (codes[:, 0::2] | (codes[:, 1::2] << 4)).contiguous(), exps + + +@functools.lru_cache(maxsize=None) +def _experts(seed: int = 20260928): + """This rank's experts through W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod's loader, and the reference's view of them. + + Returns (the loader's buffers, this rank's logical slices of the checkpoint tensors, the routing bias). + """ + from tensorrt_llm._torch.moe.fused_moe.quantization import W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod + + i_full = I_TP * MOE_TP + method = W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod() + module = SimpleNamespace( + tp_size=MOE_TP, + tp_rank=TP_RANK, + scaling_vector_size=SV, + intermediate_size=i_full, + intermediate_size_per_partition=I_TP, + hidden_size=H, + ) + kw = dict(dtype=torch.uint8, device=DEV) + proc = dict( + w31=torch.empty(E_LOCAL, 2 * I_TP, H // 2, **kw), + w31s=torch.empty(E_LOCAL, 2 * I_TP, H // SV, **kw), + w2=torch.empty(E_LOCAL, H, I_TP // 2, **kw), + w2s=torch.empty(E_LOCAL, H, I_TP // SV, **kw), + ) + raw = {name: [] for name in ("up", "up_s", "gate", "gate_s", "down", "down_s")} + gen = torch.Generator(device=DEV).manual_seed(seed) + lo, hi = TP_RANK * I_TP, (TP_RANK + 1) * I_TP + for e in range(E_LOCAL): + w1, w1s = _rand_mxfp4(i_full, H, gen) # gate + w3, w3s = _rand_mxfp4(i_full, H, gen) # up + w2, w2s = _rand_mxfp4(H, i_full, gen) # down + method.load_expert_w3_w1_weight(module, w1, w3, proc["w31"][e]) + method.load_expert_w2_weight(module, w2, proc["w2"][e]) + method.load_expert_w3_w1_weight_scale_mxfp4(module, w1s, w3s, proc["w31s"][e]) + method.load_expert_w2_weight_scale_mxfp4(module, w2s, proc["w2s"][e]) + raw["up"].append(w3[lo:hi]) + raw["up_s"].append(w3s[lo:hi]) + raw["gate"].append(w1[lo:hi]) + raw["gate_s"].append(w1s[lo:hi]) + raw["down"].append(w2[:, lo // 2 : hi // 2].contiguous()) + raw["down_s"].append(w2s[:, lo // SV : hi // SV].contiguous()) + torch.cuda.synchronize() + gen_b = torch.Generator(device=DEV).manual_seed(seed + 1) + bias = (torch.randn(NUM_EXPERTS, generator=gen_b, device=DEV) * 0.05).float() + return proc, raw, bias + + +@functools.lru_cache(maxsize=None) +def _rolled(): + """Other experts in the same layout: this rank's buffers rolled by one expert (expert e holds e - 1's weights).""" + proc, _, _ = _experts() + return {name: torch.roll(t, 1, dims=0).contiguous() for name, t in proc.items()} + + +def _bias() -> torch.Tensor: + return _experts()[2] + + +def _weights(p): + """A layer's four TRTLLM-Gen buffers, in K3MoeState.layer's argument order.""" + return p["w31"], p["w31s"], p["w2"], p["w2s"] + + +def _chosen_logits(chosen, gen): + """Logits whose top-16 (sigmoid + bias) is exactly each row's chosen experts. + + Chosen ~ N(2, 0.5) (sigmoid 0.62..0.97, so the routing weights vary), the rest -20. + """ + logits = torch.full((len(chosen), NUM_EXPERTS), -20.0, device=DEV) + for t, experts in enumerate(chosen): + idx = torch.tensor(sorted(experts), device=DEV) + logits[t, idx] = 2.0 + 0.5 * torch.randn(len(experts), generator=gen, device=DEV) + return logits + + +@functools.lru_cache(maxsize=None) +def _tokens(case: str, rows: int): + """``rows`` tokens of a routing case: router logits fp32 [rows, 896] and the latent bf16 [rows, 3584]. + + random: logits N(0, 3^2). none_local: no local expert in any token's top-16. 16_local_disjoint: 16 local experts + per token, no two tokens sharing one (16 M groups; 128 at M 8, the M <= 8 build's group capacity). group_cap: 100 + local experts with 9 of the 64 tokens and 124 with one (324 groups at M 64, the wide build's group capacity). + A call of M tokens takes the first M rows. + """ + seed = 7 + 1000 * _CASE_SALT[case] + rows + gen = torch.Generator(device=DEV).manual_seed(seed) + cpu = torch.Generator().manual_seed(seed) + x = torch.randn(rows, H, generator=gen, device=DEV).bfloat16() + logits = (torch.randn(rows, NUM_EXPERTS, generator=gen, device=DEV) * 3.0).float() + if case == "random": + return logits, x + if case == "none_local": + logits[:, OFFSET : OFFSET + E_LOCAL] = -30.0 + return logits, x + order = [OFFSET + i for i in torch.randperm(E_LOCAL, generator=cpu).tolist()] + if case == "16_local_disjoint": + assert rows * TOP_K <= E_LOCAL + chosen = [order[TOP_K * t : TOP_K * (t + 1)] for t in range(rows)] + elif case == "group_cap": + assert rows == WIDE_MAX + # Experts are dealt in order of their count to the tokens with the most free slots, so every token gets 16 + # distinct experts: 100 x 9 + 124 x 1 = 1024 pairs. + free = [TOP_K] * rows + chosen = [[] for _ in range(rows)] + for i, e in enumerate(order): + need = 9 if i < 100 else 1 + for t in sorted(range(rows), key=lambda tok: (-free[tok], tok))[:need]: + chosen[t].append(e) + free[t] -= 1 + assert all(f == 0 for f in free) + else: + raise ValueError(case) + return _chosen_logits(chosen, gen), x + + +def _draw(seed: int, rows: int): + """Random router logits fp32 [rows, 896] and latent bf16 [rows, 3584] from a seed: the call sequences' inputs.""" + gen = torch.Generator(device=DEV).manual_seed(seed) + logits = (torch.randn(rows, NUM_EXPERTS, generator=gen, device=DEV) * 3.0).float() + return logits, torch.randn(rows, H, generator=gen, device=DEV).bfloat16() + + +def _zeros(rows: int): + return ( + torch.zeros(rows, NUM_EXPERTS, device=DEV), + torch.zeros(rows, H, dtype=torch.bfloat16, device=DEV), + ) + + +# ── the state objects the checks share ──────────────────────────────────── + + +@functools.lru_cache(maxsize=None) +def _decode(): + """The shared K3MoeState and its layers A, C (this rank's experts, two counter sets) and B (the rolled experts).""" + proc, _, _ = _experts() + state = K3MoeState(_device(), I_TP, E_LOCAL) + return state, [state.layer(*_weights(p)) for p in (proc, _rolled(), proc)] + + +@functools.lru_cache(maxsize=None) +def _decode_b(): + """A second K3MoeState (its own scratch) with one layer over this rank's experts.""" + state = K3MoeState(_device(), I_TP, E_LOCAL) + return state, state.layer(*_weights(_experts()[0])) + + +@functools.lru_cache(maxsize=None) +def _wide(): + """The shared K3MoeWideState and its layers A (this rank's experts) and B (the rolled experts).""" + proc, _, _ = _experts() + state = K3MoeWideState(_device(), I_TP, E_LOCAL) + return state, [state.layer(*_weights(p)) for p in (proc, _rolled())] + + +def _experts_of(i: int): + """The buffers behind layer ``i`` of _decode (A, B, C) and _wide (A, B).""" + return (_experts()[0], _rolled(), _experts()[0])[i] + + +def _decode_call(layer, logits, x): + return k3_moe(x, logits, _bias(), OFFSET, RSF, layer) + + +def _wide_call(layer, logits, x, out=None): + return k3_moe_wide(x, logits, _bias(), OFFSET, RSF, layer, out=out) + + +def _armed(state, layers) -> bool: + """The slab armed and every layer's counters zero. + + Armed: every intermediate value byte 0x80 (FP8 -0.0), bytes 0-3 of every 16-byte scale group 0xFF (E8M0 NaN) and + bytes 4-15 zero. + """ + mod = state.mod + cs = state.cs.view(mod.G_CAP, 8, mod.K2_TILES, mod.SFB_GROUP_BYTES) + return ( + bool((state.c == -128).all()) + and bool((cs[..., :4] == -1).all()) + and bool((cs[..., 4:] == 0).all()) + and all(bool((layer.counters == 0).all()) for layer in layers) + ) + + +def _snapshot(state, layers): + return [ + t.clone() for t in (state.c, state.cs, state.part, *(layer.counters for layer in layers)) + ] + + +# ── references ──────────────────────────────────────────────────────────── + + +def _deq_w(packed, sf): + lut = torch.tensor(_E2M1, device=packed.device) + vals = torch.empty(packed.shape[0], packed.shape[1] * 2, device=packed.device) + vals[:, 0::2] = lut[(packed & 0xF).long()] + vals[:, 1::2] = lut[(packed >> 4).long()] + return vals * torch.exp2(sf.float() - 127.0).repeat_interleave(SV, dim=1) + + +def _deq_x(x_fp8, x_sf): + rows, k = x_fp8.shape + scale = torch.exp2(x_sf.reshape(rows, k // SV).float() - 127.0).repeat_interleave(SV, dim=1) + return x_fp8.float() * scale + + +def _requant(act): + """The FC1 epilogue's MXFP8 requantization per 32 columns (round-up scale), dequantized.""" + rows, cols = act.shape + blocks = act.reshape(rows, cols // SV, SV) + amax = blocks.abs().amax(dim=-1, keepdim=True) + ex = torch.ceil(torch.log2(amax / E4M3_MAX)) + ex = torch.where(amax == 0, torch.full_like(amax, -127.0), ex).clamp(-127.0, 127.0) + scale = torch.exp2(ex) + q8 = (blocks / scale).clamp(-E4M3_MAX, E4M3_MAX).to(torch.float8_e4m3fn) + return (q8.float() * scale).reshape(rows, cols) + + +def _reference(raw, x_deq, ids, weights): + """fp64 routed MoE over this rank's experts from the checkpoint tensors. + + SiTU, the MXFP8 intermediate, the down projection, the routing-weighted sum. + """ + out = torch.zeros(x_deq.shape[0], H, device=DEV) + for e in range(E_LOCAL): + tok, slot = (ids == OFFSET + e).nonzero(as_tuple=True) + if tok.numel() == 0: + continue + xe = x_deq[tok].double() + up = (xe @ _deq_w(raw["up"][e], raw["up_s"][e]).double().t()).float() + gate = (xe @ _deq_w(raw["gate"][e], raw["gate_s"][e]).double().t()).float() + act = ( + GATE_CAP + * torch.tanh(gate / GATE_CAP) + * torch.sigmoid(gate) + * (LINEAR_CAP * torch.tanh(up / LINEAR_CAP)) + ) + y = (_requant(act).double() @ _deq_w(raw["down"][e], raw["down_s"][e]).double().t()).float() + out.index_add_(0, tok, y * weights[tok, slot].float().unsqueeze(1)) + return out + + +def _compare(y, ref): + """Op-catalog gates: |d| <= 8 ulp of the row's max |ref| per element, relative RMS <= 4 ulp; finite.""" + o, r = y.float(), ref.float() + row = r.abs().amax(dim=1, keepdim=True).clamp_min(1e-12) + elt = ((o - r).abs() / row).max().item() / ULP + rms = ((o - r).pow(2).mean().sqrt() / r.pow(2).mean().sqrt().clamp_min(1e-12)).item() / ULP + finite = bool(torch.isfinite(o).all()) + return dict(elt_ulp=elt, rms_ulp=rms, ok=finite and elt <= 8.0 and rms <= 4.0) + + +def _max_ulp(a: torch.Tensor, b: torch.Tensor) -> int: + """Largest distance in bf16 ulps between two bf16 tensors (bit patterns as ordered integers).""" + + def ordered(x): + i = x.contiguous().view(torch.int16).int() + return torch.where(i < 0, -(i & 0x7FFF), i) + + return int((ordered(a) - ordered(b)).abs().max().item()) if a.numel() else 0 + + +def _stock(experts, bias, x, logits): + """The model's base path for these experts: the fused C++ route + MXFP8 quantization, then the TRTLLM-Gen runner. + + Pre-routed, SiTU with this entry's caps. Returns (y, ids, weights, x_fp8, x_sf). + """ + from tensorrt_llm._torch.moe.fused_moe.routing import RoutingMethodType + from tensorrt_llm._torch.utils import ActType_TrtllmGen + + ids, w, x_fp8, x_sf = torch.ops.trtllm.kimi_k3_noaux_tc_mxfp8_quant(logits, bias, x, RSF) + alpha = torch.full((E_LOCAL,), GATE_CAP, dtype=torch.float32, device=DEV) + beta = torch.full((E_LOCAL,), LINEAR_CAP, dtype=torch.float32, device=DEV) + y = torch.ops.trtllm.mxe4m3_mxe2m1_block_scale_moe_runner( + None, + None, + x_fp8, + x_sf.view(-1), + experts["w31"], + experts["w31s"], + None, + alpha, + beta, + None, + experts["w2"], + experts["w2s"], + None, + NUM_EXPERTS, + TOP_K, + 1, + 1, + I_TP, + H, + I_TP, + OFFSET, + E_LOCAL, + 1.0, + int(RoutingMethodType.DeepSeekV3), + int(ActType_TrtllmGen.SiTu), + topk_weights=w, + topk_ids=ids, + ) + return y, ids, w, x_fp8, x_sf + + +def _local_pairs(ids) -> int: + return int(((ids >= OFFSET) & (ids < OFFSET + E_LOCAL)).sum()) + + +def _groups_wide(ids) -> int: + """The wide build's groups for these ids: over this rank's experts, ceil(tokens / 8).""" + local = ids[(ids >= OFFSET) & (ids < OFFSET + E_LOCAL)] + counts = torch.bincount(local - OFFSET, minlength=E_LOCAL) + return int(((counts + 7) // 8).sum()) + + +def _check_vs_stock(y, experts, logits, x, where: str) -> None: + """Check y against the stock path on the same inputs (op-catalog gates); zeros when no expert here is routed.""" + y_stock, ids, *_ = _stock(experts, _bias(), x, logits) + if _local_pairs(ids) == 0: + assert bool((y == 0).all()), f"{where}: no local expert is routed, y is not zero" + return + c = _compare(y, y_stock) + assert c["ok"], f"{where}: against the stock path {c}" + + +def _alone(fn, layer, logits, x, experts, where: str): + """One call made alone (synchronized before and after), checked against the stock path.""" + torch.cuda.synchronize() + y = fn(layer, logits, x) + torch.cuda.synchronize() + _check_vs_stock(y, experts, logits, x, where) + return y + + +# ── checks ──────────────────────────────────────────────────────────────── + + +def test_state_armed_and_sized(): + """New state objects are armed and sized as the contract states for 224 local experts and intermediate 768. + + K3MoeState: slab [128, 8, 768] FP8 codes and [128, 8, 96] scale bytes armed, FC2 partial rows fp32 [1024, 3584] + zero, not compiled; a layer's counters int32 [288] zero. K3MoeWideState: slab [324, 8, 768] and [324, 8, 96] + armed, partials fp32 [1024, 3584]; a layer's counters int32 [680] zero. The loader's buffers fit the kernels. + """ + proc, _, _ = _experts() + ok, why = is_supported(*_weights(proc), E_LOCAL) + assert ok, why + state = K3MoeState(_device(), I_TP, E_LOCAL) + layer = state.layer(*_weights(proc)) + g = min(E_LOCAL, DECODE_MAX * TOP_K) + assert state.mod.G_CAP == g == 128 + assert state.c.dtype == state.cs.dtype == torch.int8 + assert tuple(state.c.shape) == (g, 8, I_TP) and tuple(state.cs.shape) == (g, 8, I_TP // 8) + assert state.part.dtype == torch.float32 and tuple(state.part.shape) == (8 * g, H) + assert bool((state.part == 0).all()) + assert layer.counters.dtype == torch.int32 and tuple(layer.counters.shape) == (32 + 2 * g,) + assert _armed(state, [layer]) and not state.head_flags and state.compiled is None + assert K3MoeState(_device(), I_TP, E_LOCAL, head_flags=True).head_flags + + wide = K3MoeWideState(_device(), I_TP, E_LOCAL) + wide_layer = wide.layer(*_weights(proc)) + gw = E_LOCAL + (WIDE_MAX * TOP_K - E_LOCAL) // 8 + assert wide.mod.G_CAP == gw == 324 + assert tuple(wide.c.shape) == (gw, 8, I_TP) and tuple(wide.cs.shape) == (gw, 8, I_TP // 8) + assert wide.part.dtype == torch.float32 and tuple(wide.part.shape) == (16 * WIDE_MAX, H) + assert tuple(wide_layer.counters.shape) == (32 + 2 * gw,) + assert _armed(wide, [wide_layer]) and wide.compiled is None + + +@pytest.mark.parametrize("m", ROUTE_M) +def test_k3_route_quant_is_the_stock_routing(m): + """k3_route_quant (inside k3_moe and k3_moe_wide) returns kimi_k3_noaux_tc_mxfp8_quant's four outputs, bit for bit. + + Top-16 ids, routing weights, MXFP8 codes and scales, with and without the early dependent trigger. + """ + logits64, x64 = _tokens("random", WIDE_MAX) + logits, x = logits64[:m].contiguous(), x64[:m].contiguous() + want = torch.ops.trtllm.kimi_k3_noaux_tc_mxfp8_quant(logits, _bias(), x, RSF) + for early in (False, True): + got = torch.ops.trtllm.k3_route_quant(logits, _bias(), x, RSF, early_trigger=early) + same = [_same(a, b) for a, b in zip(got, want)] + print(f"OPCHECK op=k3_route_quant M={m} early_trigger={early} same(ids,w,q,sf)={same}") + assert all(same), f"M {m} early_trigger {early}: {same}" + + +@pytest.mark.parametrize("m", range(1, DECODE_MAX + 1)) +@pytest.mark.parametrize("case", DECODE_CASES) +def test_k3_moe_single_call(case, m): + """One k3_moe call of M tokens on layer A against the fp64 reference and the stock path. + + Also: run-to-run identical bits; each row within one bf16 ulp of the same row of the 8-token call (the slice FC2 + adds a token's expert terms in slices whose bounds follow the step's group count; bit-identity is reported); the + slab armed and every counter zero after the call; at 16_local_disjoint M 8 the group count at the capacity, 128. + """ + proc, raw, bias = _experts() + state, layers = _decode() + logits8, x8 = _tokens(case, DECODE_MAX) + logits, x = logits8[:m].contiguous(), x8[:m].contiguous() + y = _decode_call(layers[0], logits, x) + torch.cuda.synchronize() + rearmed = _armed(state, layers) + assert ( + y.shape == (m, H) + and y.dtype == torch.bfloat16 + and y.is_contiguous() + and y.device == x.device + ) + det = all(_same(_decode_call(layers[0], logits, x), y) for _ in range(2)) + y8 = _decode_call(layers[0], logits8, x8) + ulp_m8 = _max_ulp(y, y8[:m]) + y_stock, ids, w, x_fp8, x_sf = _stock(proc, bias, x, logits) + local = _local_pairs(ids) + groups = int(torch.unique(ids[(ids >= OFFSET) & (ids < OFFSET + E_LOCAL)]).numel()) + if case == "16_local_disjoint": + assert local == TOP_K * m and groups == TOP_K * m + if m == DECODE_MAX: + assert groups == state.mod.G_CAP + if local == 0: + zeros = bool((y == 0).all()) + print( + f"OPCHECK op=k3_moe case={case} M={m} zeros={zeros} det={det} " + f"rows_as_m8={_same(y, y8[:m])} scratch_rearmed={rearmed}" + ) + assert zeros and det and ulp_m8 == 0 and rearmed + return + ref = _reference(raw, _deq_x(x_fp8, x_sf), ids, w).bfloat16() + c_ref, c_stock, c_stock_ref = _compare(y, ref), _compare(y, y_stock), _compare(y_stock, ref) + print( + f"OPCHECK op=k3_moe case={case} M={m} local_pairs={local} groups={groups} " + f"vs_ref_elt_ulp={c_ref['elt_ulp']:.2f} vs_ref_rms_ulp={c_ref['rms_ulp']:.2f} " + f"vs_stock_elt_ulp={c_stock['elt_ulp']:.2f} vs_stock_rms_ulp={c_stock['rms_ulp']:.2f} " + f"stock_vs_ref_elt_ulp={c_stock_ref['elt_ulp']:.2f} det={det} " + f"rows_as_m8={_same(y, y8[:m])} max_ulp_vs_m8={ulp_m8} scratch_rearmed={rearmed}" + ) + assert c_ref["ok"] and c_stock["ok"] and c_stock_ref["ok"] + assert det and ulp_m8 <= 1 and rearmed + + +@pytest.mark.parametrize("m", WIDE_M) +@pytest.mark.parametrize("case", WIDE_CASES) +def test_k3_moe_wide_single_call(case, m): + """One k3_moe_wide call of M tokens on its layer A against the fp64 reference and the stock path. + + Also: run-to-run identical bits; the slab armed and every counter zero after the call; at M <= 8 within one bf16 + ulp of k3_moe on the same experts (bit-identity reported); at group_cap M 64 the group count at the capacity, 324. + """ + proc, raw, bias = _experts() + state, layers = _wide() + _, decode_layers = _decode() + logits64, x64 = _tokens(case, WIDE_MAX) + logits, x = logits64[:m].contiguous(), x64[:m].contiguous() + y = _wide_call(layers[0], logits, x) + torch.cuda.synchronize() + rearmed = _armed(state, layers) + assert y.shape == (m, H) and y.dtype == torch.bfloat16 and y.is_contiguous() + det = all(_same(_wide_call(layers[0], logits, x), y) for _ in range(2)) + y_stock, ids, w, x_fp8, x_sf = _stock(proc, bias, x, logits) + groups = _groups_wide(ids) + if case == "group_cap" and m == WIDE_MAX: + assert groups == state.mod.G_CAP == 324 + decode = "" + if m <= DECODE_MAX: + y_dec = _decode_call(decode_layers[0], logits, x) + ulp_dec = _max_ulp(y, y_dec) + decode = f" max_ulp_vs_k3_moe={ulp_dec} bits_as_k3_moe={_same(y, y_dec)}" + assert ulp_dec <= 1 + local = _local_pairs(ids) + if local == 0: + zeros = bool((y == 0).all()) + print( + f"OPCHECK op=k3_moe_wide case={case} M={m} zeros={zeros} det={det} " + f"scratch_rearmed={rearmed}{decode}" + ) + assert zeros and det and rearmed + return + ref = _reference(raw, _deq_x(x_fp8, x_sf), ids, w).bfloat16() + c_ref, c_stock, c_stock_ref = _compare(y, ref), _compare(y, y_stock), _compare(y_stock, ref) + print( + f"OPCHECK op=k3_moe_wide case={case} M={m} local_pairs={local} groups={groups} " + f"vs_ref_elt_ulp={c_ref['elt_ulp']:.2f} vs_ref_rms_ulp={c_ref['rms_ulp']:.2f} " + f"vs_stock_elt_ulp={c_stock['elt_ulp']:.2f} vs_stock_rms_ulp={c_stock['rms_ulp']:.2f} " + f"stock_vs_ref_elt_ulp={c_stock_ref['elt_ulp']:.2f} det={det} scratch_rearmed={rearmed}{decode}" + ) + assert c_ref["ok"] and c_stock["ok"] and c_stock_ref["ok"] + assert det and rearmed + + +@pytest.mark.parametrize("m", [9, WIDE_MAX]) +def test_k3_moe_wide_out_buffer(m): + """k3_moe_wide with ``out``: the result is out[:M] (the same storage), the bits of the call without ``out``. + + The rows of ``out`` past M keep their bits. + """ + _, layers = _wide() + logits, x = _draw(900 + m, m) + want = _wide_call(layers[0], logits, x) + out = torch.full((WIDE_MAX + 3, H), -3.0, dtype=torch.bfloat16, device=DEV) + tail = out[m:].clone() + got = _wide_call(layers[0], logits, x, out=out) + torch.cuda.synchronize() + assert got.data_ptr() == out.data_ptr() and got.shape == (m, H) + assert _same(got, want) and _same(out[m:], tail) + + +def test_layers_by_steps_dip_and_regrow(): + """Decode steps of several layers on one state, the token count dipping and growing back, back to back. + + Steps of three k3_moe calls (layers A, B, C of one state) at M 8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8, then steps of two + k3_moe_wide calls (layers A, B of the wide state) at M 64, 64, 16, 1, 40, 64, 8, 23, 64, 9, 64, with new inputs + every call. Each call is first made alone (synchronized, checked against the stock path); in the sequence, back to + back on one stream, each returns the bits of its call alone, and afterwards both slabs are armed and every counter + is zero: a call after a smaller one reads nothing an older, larger call left. + """ + state, layers = _decode() + calls = [(i, *_draw(1000 + 10 * s + i, m)) for s, m in enumerate(DIP_STEPS) for i in range(3)] + alone = [ + _alone( + _decode_call, layers[i], lg, x, _experts_of(i), f"k3_moe layer {i} M {x.shape[0]} alone" + ) + for i, lg, x in calls + ] + seq = [_decode_call(layers[i], lg, x) for i, lg, x in calls] + torch.cuda.synchronize() + bad = [k for k, (y, a) in enumerate(zip(seq, alone)) if not _same(y, a)] + assert not bad, f"k3_moe calls {bad} of the sequence differ from the same calls alone" + assert _armed(state, layers) + + wide, wide_layers = _wide() + wide_calls = [ + (i, *_draw(2000 + 10 * s + i, m)) for s, m in enumerate(WIDE_DIP_STEPS) for i in range(2) + ] + alone = [ + _alone( + _wide_call, + wide_layers[i], + lg, + x, + _experts_of(i), + f"k3_moe_wide layer {i} M {x.shape[0]} alone", + ) + for i, lg, x in wide_calls + ] + seq = [_wide_call(wide_layers[i], lg, x) for i, lg, x in wide_calls] + torch.cuda.synchronize() + bad = [k for k, (y, a) in enumerate(zip(seq, alone)) if not _same(y, a)] + assert not bad, f"k3_moe_wide calls {bad} of the sequence differ from the same calls alone" + assert _armed(wide, wide_layers) + + +def test_two_states_interleaved(): + """Two K3MoeStates (each its own scratch) and the wide state, calls interleaved in an irregular pattern. + + A A B A B B A A A B, twice, with a k3_moe_wide call after every third, M varying, back to back on one stream: every + call returns the bits of the same call alone (checked against the stock path), and every slab is armed and every + counter zero afterwards. + """ + state_a, layers_a = _decode() + state_b, layer_b = _decode_b() + wide, wide_layers = _wide() + plan = [] + for i, which in enumerate("AABABBAAAB" * 2): + m = (3, 8, 1, 8, 5)[i % 5] + if which == "A": + plan.append((_decode_call, layers_a[i % 3], _experts_of(i % 3), *_draw(3000 + i, m))) + else: + plan.append((_decode_call, layer_b, _experts()[0], *_draw(3000 + i, m))) + if i % 3 == 2: + mw = (40, 64, 9)[(i // 3) % 3] + plan.append((_wide_call, wide_layers[i % 2], _experts_of(i % 2), *_draw(3500 + i, mw))) + alone = [ + _alone(fn, layer, lg, x, ex, f"interleaved call {k} alone") + for k, (fn, layer, ex, lg, x) in enumerate(plan) + ] + seq = [fn(layer, lg, x) for fn, layer, _, lg, x in plan] + torch.cuda.synchronize() + bad = [k for k, (y, a) in enumerate(zip(seq, alone)) if not _same(y, a)] + assert not bad, f"interleaved calls {bad} differ from the same calls alone" + assert _armed(state_a, layers_a) and _armed(state_b, [layer_b]) and _armed(wide, wide_layers) + + +def test_graph_capture_and_replay(): + """A captured step replayed with rewritten inputs, eager calls of other token counts between replays. + + The step: k3_moe on layers A, B, C at M 8, then k3_moe_wide on its layer A at M 64, captured once and replayed six + times with new inputs copied into its static buffers; between replays an eager k3_moe call (M 3, 1, 6, 5, 2, 7) + and an eager k3_moe_wide call (M 23, 9, 40, 1, 64, 16) on the same states. Every replayed and eager call returns + the bits of the same call alone (checked against the stock path), and the slabs are armed and every counter zero + afterwards. + """ + state, layers = _decode() + wide, wide_layers = _wide() + static = [_draw(4000 + i, DECODE_MAX) for i in range(3)] + [_draw(4003, WIDE_MAX)] + + def step(): + outs = [_decode_call(layers[i], *static[i]) for i in range(3)] + return outs + [_wide_call(wide_layers[0], *static[3])] + + step() # every first call eager: k3_route_quant and both k3_moe builds compile here if nothing has yet + torch.cuda.synchronize() + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + outs = step() + for rep in range(REPLAYS): + inputs = [_draw(5000 + 10 * rep + i, DECODE_MAX) for i in range(3)] + [ + _draw(5003 + 10 * rep, WIDE_MAX) + ] + alone = [ + _alone( + _decode_call, layers[i], *inputs[i], _experts_of(i), f"replay {rep} layer {i} alone" + ) + for i in range(3) + ] + alone.append( + _alone( + _wide_call, wide_layers[0], *inputs[3], _experts_of(0), f"replay {rep} wide alone" + ) + ) + eager = ( + _draw(6000 + rep, (3, 1, 6, 5, 2, 7)[rep]), + _draw(6100 + rep, (23, 9, 40, 1, 64, 16)[rep]), + ) + eager_layers = (layers[rep % 3], wide_layers[rep % 2]) + eager_alone = ( + _alone( + _decode_call, eager_layers[0], *eager[0], _experts_of(rep % 3), f"eager {rep} alone" + ), + _alone( + _wide_call, + eager_layers[1], + *eager[1], + _experts_of(rep % 2), + f"eager wide {rep} alone", + ), + ) + for (lg, x), (new_lg, new_x) in zip(static, inputs): + lg.copy_(new_lg) + x.copy_(new_x) + graph.replay() + torch.cuda.synchronize() + bad = [i for i, (y, a) in enumerate(zip(outs, alone)) if not _same(y, a)] + assert not bad, f"replay {rep}: calls {bad} differ from the same calls alone" + got = (_decode_call(eager_layers[0], *eager[0]), _wide_call(eager_layers[1], *eager[1])) + torch.cuda.synchronize() + assert all(_same(g, a) for g, a in zip(got, eager_alone)), f"eager calls after replay {rep}" + del graph + assert _armed(state, layers) and _armed(wide, wide_layers) + + +def test_unsupported_calls_refused_before_launch(): + """Unsupported calls raise ValueError before any launch, and the next call is correct. + + k3_moe at M 0 and 9, k3_moe_wide at M 0 and 65, and k3_moe on a head_flags state's layer (that build takes the + front's ready words: k3_moe_fused_front only). The slabs, the FC2 partial rows and every counter keep their bits; + the next calls return the bits of the same calls made before. + """ + state, layers = _decode() + wide, wide_layers = _wide() + lg8, x8 = _draw(7000, DECODE_MAX) + lg40, x40 = _draw(7001, 40) + want = _decode_call(layers[0], lg8, x8) + want_wide = _wide_call(wide_layers[0], lg40, x40) + flags_state = K3MoeState(_device(), I_TP, E_LOCAL, head_flags=True) + flags_layer = flags_state.layer(*_weights(_experts()[0])) + torch.cuda.synchronize() + objects = ((state, layers), (wide, wide_layers), (flags_state, [flags_layer])) + before = [_snapshot(st, lyrs) for st, lyrs in objects] + for m in (0, DECODE_MAX + 1): + with pytest.raises(ValueError): + _decode_call(layers[0], *_zeros(m)) + for m in (0, WIDE_MAX + 1): + with pytest.raises(ValueError): + _wide_call(wide_layers[0], *_zeros(m)) + with pytest.raises(ValueError, match="head_flags"): + _decode_call(flags_layer, lg8, x8) + torch.cuda.synchronize() + after = [_snapshot(st, lyrs) for st, lyrs in objects] + assert all(_same(a, b) for snap_a, snap_b in zip(before, after) for a, b in zip(snap_a, snap_b)) + assert flags_state.compiled is None + assert _same(_decode_call(layers[0], lg8, x8), want) + assert _same(_wide_call(wide_layers[0], lg40, x40), want_wide) + + +def test_construction_and_first_call_refuse_capture(): + """Under CUDA-graph capture the states refuse to allocate, and a new state refuses its first (compiling) call. + + K3MoeState() and K3MoeState.layer() raise RuntimeError; the first call on a new K3MoeState and on a new + K3MoeWideState raises RuntimeError before the k3_moe launch (k3_route_quant, compiled eagerly first, is captured + and discarded with the graph); neither state is compiled afterwards. K3MoeWideState() and K3MoeWideState.layer() + do not check for capture (see the contract's Notes), so they are built eagerly here. + """ + proc, _, bias = _experts() + fresh = K3MoeState(_device(), I_TP, E_LOCAL) + fresh_layer = fresh.layer(*_weights(proc)) + fresh_wide = K3MoeWideState(_device(), I_TP, E_LOCAL) + fresh_wide_layer = fresh_wide.layer(*_weights(proc)) + lg4, x4 = _draw(7100, 4) + lg16, x16 = _draw(7101, 16) + # Both k3_route_quant builds compiled, whichever PDL setting the layers use. + for early in (False, True): + torch.ops.trtllm.k3_route_quant(lg4, bias, x4, RSF, early_trigger=early) + torch.cuda.synchronize() + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): + K3MoeState(_device(), I_TP, E_LOCAL) + with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): + fresh.layer(*_weights(proc)) + with pytest.raises(RuntimeError, match="compiles on its first call"): + _decode_call(fresh_layer, lg4, x4) + with pytest.raises(RuntimeError, match="compiles on its first call"): + _wide_call(fresh_wide_layer, lg16, x16) + del graph + assert fresh.compiled is None and fresh_wide.compiled is None + + +def test_negative_control_weights_rebound_after_layer(): + """Negative control: a layer reads the weight buffers it was built over; rebinding the weights is not seen. + + Layers of both builds are built over one set of experts (E1, a copy of this rank's). The weights are then + "reloaded" by rebinding to new tensors (E2: the experts rolled by one), as a loader that replaces its parameters + would. Nothing raises, and each layer still returns E1's partial bit for bit, outside the op-catalog gates of the + stock path on E2: silently stale. Layers built over E2 match its stock path; and once E2 is copied into E1's + buffers in place, the stale layers return E2's partial bit for bit. + """ + proc, _, bias = _experts() + state, _ = _decode() + wide, _ = _wide() + e1 = {name: t.clone() for name, t in proc.items()} + stale = state.layer(*_weights(e1)) + stale_wide = wide.layer(*_weights(e1)) + lg8, x8 = _tokens("random", DECODE_MAX) + lg40, x40 = _draw(8000, 40) + y_e1 = _decode_call(stale, lg8, x8) + yw_e1 = _wide_call(stale_wide, lg40, x40) + torch.cuda.synchronize() + + # The reload: the caller's weights are now these new tensors; E1's buffers stay as they were. + e2 = _rolled() + y_stale = _decode_call(stale, lg8, x8) + yw_stale = _wide_call(stale_wide, lg40, x40) + stock_e2 = _stock(e2, bias, x8, lg8)[0] + stock_wide_e2 = _stock(e2, bias, x40, lg40)[0] + c, cw = _compare(y_stale, stock_e2), _compare(yw_stale, stock_wide_e2) + print( + f"OPCHECK op=k3_moe case=weights_rebound stale_bits_as_e1={_same(y_stale, y_e1)} " + f"stale_vs_e2_stock_elt_ulp={c['elt_ulp']:.1f} wide_stale_bits_as_e1={_same(yw_stale, yw_e1)} " + f"wide_stale_vs_e2_stock_elt_ulp={cw['elt_ulp']:.1f}" + ) + assert _same(y_stale, y_e1) and not c["ok"] + assert _same(yw_stale, yw_e1) and not cw["ok"] + + fresh = state.layer(*_weights(e2)) + fresh_wide = wide.layer(*_weights(e2)) + y_e2 = _decode_call(fresh, lg8, x8) + yw_e2 = _wide_call(fresh_wide, lg40, x40) + assert _compare(y_e2, stock_e2)["ok"] and _compare(yw_e2, stock_wide_e2)["ok"] + for name, t in e1.items(): + t.copy_(e2[name]) # in place: the buffers the stale layers read now hold E2 + assert _same(_decode_call(stale, lg8, x8), y_e2) + assert _same(_wide_call(stale_wide, lg40, x40), yw_e2) diff --git a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_front_op_matrix.py b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_front_op_matrix.py new file mode 100644 index 000000000000..ea93c75c4902 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_front_op_matrix.py @@ -0,0 +1,32 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Collected entry point for the k3_moe_front op's certification matrix. + +The matrix also certifies moe/k3_moe's fused-front cells (k3_moe_fused_front on a plain and a head_flags +K3MoeState). It is ``comm/_k3_moe_front_op_matrix.py``: rank bodies live beside ``_lockstep`` and ``_rank_job`` in +``comm/``, and it is its own W-rank launcher (call sequences over caller-owned workspaces and states, so one job, not +independent cases); see ``_rank_job`` for why that is left intact. This file puts ``comm/`` on the import path to +reach ``_rank_job``. +""" + +import sys +from pathlib import Path + +import pytest +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "comm")) +import _rank_job # noqa: E402 + + +def _is_sm100() -> bool: + return torch.cuda.is_available() and torch.cuda.get_device_capability() == (10, 0) + + +pytestmark = pytest.mark.skipif(not _is_sm100(), reason="k3_moe_front needs sm_100") + + +# Each case starts its own W-rank mpirun over the visible devices; under xdist several workers would fight for them. +@pytest.mark.no_xdist +def test_k3_moe_front_op_matrix() -> None: + _rank_job.run("k3_moe_front") From 20d29536ab81e44036d88989b16c82bc3ffd5c5a Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 19:32:46 -0700 Subject: [PATCH 061/161] [None][test] Kimi K3 catalog entries: index rows and CI lists - catalog/index.yaml: rows for the eight entries. - l0_b200 (pre-merge, one GPU): the k3_moe, wide k3_moe and k3_route_quant kernel tests and moe/k3_moe's catalog test. - l0_gb200_multi_gpus (pre-merge, 4 GPUs): the sandwich, MoE front and latent reduce kernel tests, and the 4-rank catalog matrices (sandwiches, MNNVL all-reduce and split all-gather, latent reduce, MoE front), plus comm/allgather's matrix, which gains its sm_100 run there. - The contracts state their design rules and records in their own terms rather than citing internal notes. moe/k3_moe's test checks that the wide state and its layers also refuse construction under capture. Signed-off-by: Vasanth Sabavat --- .../catalog/comm/k3_sandwich_oproj.md | 18 +++++----- .../catalog/comm/k3_sandwich_plain.md | 8 ++--- .../catalog/comm/k3_sandwich_tail.md | 8 ++--- .../catalog/comm/mnnvl_allgather_split.md | 26 +++++++------- .../catalog/comm/mnnvl_fusion_allreduce.md | 24 ++++++------- .../modeling_v2/catalog/index.yaml | 35 +++++++++++++++++++ .../modeling_v2/catalog/moe/k3_moe.md | 17 ++++----- .../modeling_v2/catalog/moe/k3_moe_front.md | 10 +++--- .../test_lists/test-db/l0_b200.yml | 5 +++ .../test-db/l0_gb200_multi_gpus.yml | 12 +++++++ .../moe/test_modeling_v2_k3_moe.py | 11 +++--- 11 files changed, 112 insertions(+), 62 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md index 333f925cf38e..34216e3e2f02 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md @@ -113,7 +113,7 @@ alternating between two workspaces in an irregular pattern, so that their counte same number of calls, the `k`-th with the same op and `M` — across layers and decode steps, eager calls and graph replays alike; and on one stream the same order of calls across objects: each call spins until its peers' rows of the same call arrive, so two ranks issuing calls on two objects in different orders on one stream deadlock (measured -for `mnnvl_allreduce_attn_res`, `runs/drafter/u4-mnnvl-srun-2`; this kernel waits the same way). On each rank the +for `comm/mnnvl_allreduce_attn_res`, see its contract; this kernel waits the same way). On each rank the calls on one workspace run one after another: a call reads the counters after its grid-dependency wait, which covers the previous call because every kernel between two calls waits for its predecessor (the kernel's statement) — a kernel launched under PDL that skips that wait must not sit between two calls. @@ -129,7 +129,7 @@ kernel's statement), so no separate clear and no record of the previous call's s `k3_spec_accept` failure (a re-arm sized by the current call) cannot occur. The test still drives the sequence that exposed it (below). -**Why the test drives call sequences.** See `mnnvl_allreduce_attn_res.md` (*State*): Phase 0's `k3_spec_accept` +**Why the test drives call sequences.** See `mnnvl_allreduce_attn_res.md` (*State*): Kimi K3's `k3_spec_accept` once re-armed its Lamport buffer for the current call's rows only, every single-call test passed, and a call sequence whose row count dipped and grew back caught it; in serving it hung the ranks. This entry's test runs 11 decode steps of 12 chained layers at `M` = 8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8, a random rank 5 ms late at every call, each call @@ -172,13 +172,13 @@ it changes scheduling, not results. three sandwich entries in `_k3_sandwich_common.py` beside it. The reference is native torch: `core` and `o_weight` are small multiples of 1/8 and 1/16, so every partial sum is exact in fp32 and `updated` is compared bit for bit; `normed` against an fp32 reference within 2e-2 of its largest magnitude. -- A2: the matrix takes `--world-size` and `--launcher` (`mpirun` on one tray, `srun` across trays). The kernel sums - ranks in chunks of 8, so a run at `W` <= 8 exercises one chunk; Kimi K3 runs `W` = 16 over four trays. This - entry's 16-rank receipt is pending. The op's kernel test - (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py`), in its B7 form — this op bit for bit - against `o_proj` and the MNNVL one-shot — passed 180 / 180 at 16 ranks - (`runs/session-7594608/pre/test_k3_sandwich.log`): the kernel's record, not this entry's receipt. -- A1: `mutates_args` names `ws_uc`, `ws_mc` and `ws_flags` — every call pushes through `ws_mc` into every rank's +- World sizes: the matrix takes `--world-size` and `--launcher` (`mpirun` on one tray, `srun` across trays). The + kernel sums ranks in chunks of 8, so a run at `W` <= 8 exercises one chunk; Kimi K3 runs `W` = 16 over four trays. + This entry's 16-rank receipt is pending. The op's kernel test + (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py`; this op bit for bit against `o_proj` and + the MNNVL one-shot) passed every case at 16 ranks on four trays in a recorded run: the kernel's record, not this + entry's receipt. +- State: `mutates_args` names `ws_uc`, `ws_mc` and `ws_flags` — every call pushes through `ws_mc` into every rank's `ws_uc`, empties the words it read in `ws_uc` and advances `ws_flags` — and `x_slab`. The op module keeps no workspace registry: the caller passes the object it created. The compile cache is the documented process-wide cache above. The published / polled slabs (`x_slab`, `src_slab`) are cross-call state without a state object, so they diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md index 2cbe28ac06e3..4dcd66c12552 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md @@ -152,8 +152,8 @@ forms are two kernels, each compiled (seconds) on its first call, which must be is compared bit for bit in both forms. `normed` is compared with torch's fp32 RMSNorm within 2e-2 of its largest magnitude (the one-shot rounds the squares to bf16 and sums them in its own tree). SiLU-and-mul's rounding on general inputs is the kernel test's to certify (bit for bit against `k3_ctm_gemv_swiglu`), not this matrix's. -- A2: as for `k3_sandwich_oproj`; this entry's 16-rank receipt is pending. The op's kernel test - (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py`), in its B7 form, passed 180 / 180 at 16 - ranks (`runs/session-7594608/pre/test_k3_sandwich.log`): the kernel's record, not this entry's receipt. -- A1: `mutates_args` names `ws_uc`, `ws_mc` and `ws_flags`, every buffer the op writes. The compile cache is the +- World sizes: as for `k3_sandwich_oproj`; this entry's 16-rank receipt is pending. The op's kernel test + (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py`) passed every case at 16 ranks on four + trays in a recorded run: the kernel's record, not this entry's receipt. +- State: `mutates_args` names `ws_uc`, `ws_mc` and `ws_flags`, every buffer the op writes. The compile cache is the documented process-wide cache above. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md index 3b31a813036f..a8eccdfb6fba 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md @@ -192,10 +192,10 @@ outside CUDA-graph capture first". `updated_out` is not part of the key. The cac within 2e-2 of an fp32 reference; the tapped `updated` bit for bit with the returned one; every output bitwise across the ranks. The negative control's wrong pairing puts more than half the elements outside the 8e-3 bound, the largest error over 10 times it. -- A2: as for `k3_sandwich_oproj`; this entry's 16-rank receipt is pending. The op's kernel test - (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py`), in its B7 form, passed 180 / 180 at 16 - ranks (`runs/session-7594608/pre/test_k3_sandwich.log`): the kernel's record, not this entry's receipt. -- A1: `mutates_args` names every buffer the op can write: `ws_uc`, `ws_mc`, `ws_flags`, `x_slab`, `lat_uc`, +- World sizes: as for `k3_sandwich_oproj`; this entry's 16-rank receipt is pending. The op's kernel test + (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py`) passed every case at 16 ranks on four + trays in a recorded run: the kernel's record, not this entry's receipt. +- State: `mutates_args` names every buffer the op can write: `ws_uc`, `ws_mc`, `ws_flags`, `x_slab`, `lat_uc`, `lat_flags`, `tap` and `updated_out`. The compile cache is the documented process-wide cache above. - Not certified (an op option the wrapper does not expose): the folded latent all-reduce. With `lat_uc` / `lat_flags` of a `K3SandwichLatentExchange` — its own collective `create(mapping, fabric_handle=None)`; two halves of diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md index d3efd95a22ae..576e54f7d960 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md @@ -26,7 +26,7 @@ rounded), exact midpoints between two bf16 values in every fourth of them (ties fourth column of each part, and in the fp32 part infinities, a NaN, signed denormals and the largest float, which arrive unchanged. The result is bitwise the same on every rank (certified, every call of the test). -Kimi K3's use (B7 82a110a92a, its code): the row-sharded MoE head of a wide decode step (9 to 64 tokens), and of a +Kimi K3's use: the row-sharded MoE head of a wide decode step (9 to 64 tokens), and of a decode step of at most 8 tokens where the fused MoE front kernel does not run. Rank `r`'s GEMV gives fp32 `[T, 3584/W + 896/W]`: the latent down projection's columns `[r * 3584/W, (r+1) * 3584/W)`, then the router logits of experts `[r * 896/W, (r+1) * 896/W)`. This call assembles the bf16 latent `[T, 3584]` (`bf16_out`) and the fp32 @@ -96,9 +96,8 @@ all-reduce] at Kimi K3's one-shot ceilings, `T` = 8, 2, 16, 64, 1, 7, 32, 8, 3, follow one-shot and two-shot all-reduces. Two objects are two independent rotations: 20 all-gathers of mixed token counts and splits alternating irregularly between two workspaces are all correct, and each workspace's flags move with its own calls only (certified). They are not independent orders: on one stream every rank must issue its -collectives in the same order, whatever object each belongs to (measured for `mnnvl_allreduce_attn_res`: ranks -issuing B-then-A against A-then-B deadlocked, `runs/drafter/u4-mnnvl-srun-2`; this op waits for its peers the same -way). +collectives in the same order, whatever object each belongs to (measured for `comm/mnnvl_allreduce_attn_res`: +ranks issuing B-then-A against A-then-B deadlocked; this op waits for its peers the same way). **Call-order invariant.** Every rank of the group makes the same sequence of calls on one workspace — the same number, the `k`-th with the same op, `T`, `B` and `F` — across layers and decode steps, eager calls and graph replays @@ -115,7 +114,7 @@ recorded and in the previous call's stage layout (`cpp/tensorrt_llm/common/lampo `LamportFlags::clearDirtyLamportBuf`): after a two-shot all-reduce both of its stages, after any other call the first. Certified by the split grid (55 calls of different sizes back to back on one workspace) and by the sequences below. -**Why the test drives call sequences.** See `mnnvl_allreduce_attn_res.md` (*State*): Phase 0's `k3_spec_accept` +**Why the test drives call sequences.** See `mnnvl_allreduce_attn_res.md` (*State*): Kimi K3's `k3_spec_accept` once re-armed its Lamport buffer for the current call's rows only; every single-call test passed, and a sequence whose row count dipped and grew back caught it. This entry's test runs 16 decode steps of 6 layers, one head all-gather per layer at Kimi K3's split, at `T` = 8, 8, 8, 2, 7, 8, 1, 1, 64, 3, 32, 8, 16, 1, 64, 8, a random rank 5 ms late at @@ -158,18 +157,17 @@ depend on it. Test: `tests/unittest/_torch/modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py`. The reference is native torch and exact; this op's outputs are compared bit for bit. The all-reduce and attention-residual calls of its sequences are checked as in their own matrices (sums bit for bit, normed outputs within a tolerance). -- Design choices this entry follows (U4U5_PLAN §1): A1, "a typed state object per stateful op ..., built by an - explicit, collective, eager `create()` ... Tests drive real-state call sequences (layers x steps, capture + replay, - two objects interleaved) plus a negative control. `mutates_args` names every written buffer" (the op falls short of - the last, see the gaps below); A2, the matrix takes `--world-size` and `--launcher` (`mpirun` on one node, `srun` - across trays), CI runs it at 4 ranks on one GB200 tray; A4, one `MnnvlWorkspace` shared by every MNNVL entry of the - TP group. The 16-rank receipt is pending. +- State and test design: a typed state object built by an explicit, collective, eager `create()`; a test that drives + call sequences on real state (layers x steps, capture + replay, two objects interleaved) plus a negative control; + every written buffer named in the schema (the op falls short there, see the gaps below); the matrix takes + `--world-size` and `--launcher` (`mpirun` on one node, `srun` across trays) and CI runs it at 4 ranks on one GB200 + tray; one `MnnvlWorkspace` shared by every MNNVL entry of the TP group. The 16-rank receipt is pending. - Not exercised: `B` = 0 or `F` = 0 (the op accepts both), denormal and non-finite values in the bf16 columns, an accepted call of more than 64 tokens. -- Gaps against A1 (the op is unchanged by this entry): the schema marks `comm_buffer` mutable `(a!)` but not - `buffer_flags`, which every call advances (A1 item 4); the op has no `register_fake` +- Gaps (the op is unchanged by this entry): the schema marks `comm_buffer` mutable `(a!)` but not `buffer_flags`, + which every call advances; the op has no `register_fake` (`tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py` registers one for the other two MNNVL ops), so fake-tensor tracing, e.g. `torch.compile`, cannot run it. - In the model today the call is `MNNVLAllReduce.allgather_split(input, bf16_columns)` on `MNNVLAllReduce`'s workspace (a dict keyed by `Mapping`, grown to the call's footprint by the first eager call that needs more). This entry takes - the explicit object instead (A4), sized at construction. + the explicit object instead, sized at construction. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md index 55486ef6d808..75dafff24a0c 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md @@ -123,8 +123,8 @@ latent all-reduce] at Kimi K3's ceilings, `T` = 8, 2, 16, 64, 1, 7, 32, 8, 3, 16 two independent rotations: 20 calls of mixed shapes and paths alternating irregularly between two workspaces are all correct, and each workspace's flags move with its own calls only (certified). They are not independent orders: on one stream every rank must issue its collectives in the same order, whatever object each belongs to (measured for -`mnnvl_allreduce_attn_res`: ranks issuing B-then-A against A-then-B deadlocked, `runs/drafter/u4-mnnvl-srun-2`; this -op waits for its peers the same way). +`comm/mnnvl_allreduce_attn_res`: ranks issuing B-then-A against A-then-B deadlocked; this op waits for its peers +the same way). **Call-order invariant.** Every rank of the group makes the same sequence of calls on one workspace — the same number, the `k`-th with the same op, `T`, `H`, fusion and path — across layers and decode steps, eager calls and @@ -146,7 +146,7 @@ attention-residual all-reduce the first stage, whichever path the clearing call grid (every shape one-shot and then two-shot back to back on one workspace, then the next shape) and by the sequences below. -**Why the test drives call sequences.** See `mnnvl_allreduce_attn_res.md` (*State*): Phase 0's `k3_spec_accept` +**Why the test drives call sequences.** See `mnnvl_allreduce_attn_res.md` (*State*): Kimi K3's `k3_spec_accept` once re-armed its Lamport buffer for the current call's rows only; every single-call test passed, and a sequence whose row count dipped and grew back caught it; in serving it made the ranks disagree and hang. This entry's test runs 16 decode steps of 8 layers, each layer the latent all-reduce `[T, 3584]` and the fused all-reduce `[T, 7168]` chained through @@ -201,15 +201,15 @@ device's SM count, and in the fused form it sets the order in which the squares bf16 and the residual add is one bf16 rounding of an exact fp32 value in the op and in the reference; the sum and `updated` are compared bit for bit, `normed` against the fp32 RMSNorm of `updated` within 1e-2 of its largest magnitude (the bf16 squares cost at most 2^-9 on `rcp`, the bf16 output 2^-8). -- Design choices this entry follows (U4U5_PLAN §1): A1, "a typed state object per stateful op ..., built by an - explicit, collective, eager `create()` ... Tests drive real-state call sequences (layers x steps, capture + replay, - two objects interleaved) plus a negative control. `mutates_args` names every written buffer" (the op falls short of - the last, see the gaps below); A2, the matrix takes `--world-size` and `--launcher` (`mpirun` on one node, `srun` - across trays), CI runs it at 4 ranks on one GB200 tray; A4, "one caller-owned `MnnvlWorkspace` ..., shared by" every - MNNVL entry of the TP group, "`one_shot_max_bytes` per call". +- State and test design: a typed state object built by an explicit, collective, eager `create()`; a test that drives + call sequences on real state (layers x steps, capture + replay, two objects interleaved) plus a negative control; + every written buffer named in the schema (the op falls short there, see the gaps below); the matrix takes + `--world-size` and `--launcher` (`mpirun` on one node, `srun` across trays) and CI runs it at 4 ranks on one GB200 + tray; one caller-owned `MnnvlWorkspace` shared by every MNNVL entry of the TP group, and + `one_shot_max_bytes` per call. - At `W` = 16 the one-shot kernel adds the ranks in two chunks of 8, a branch a 4-rank run never reaches. The 16-rank receipt is pending. -- Kimi K3's calls (B7 82a110a92a, its code, not this test): the model sets every `MNNVLAllReduce` of the target, its +- Kimi K3's calls (its decode path, not this test): the model sets every `MNNVLAllReduce` of the target, its LM head and a drafter to `one_shot_max_bytes` = 4 MiB (`DECODE_AR_ONE_SHOT_MAX_BYTES`, against main's 1 MiB), and a wide decode step (9 to 64 tokens) passes 1 MiB per call (`WIDE_AR_ONE_SHOT_MAX_BYTES`). Plain: the routed-latent all-reduce `[T, 3584]` (decode steps where the latent exchange push is not used; wide steps) and a wide step's @@ -217,8 +217,8 @@ device's SM count, and in the fused form it sets the order in which the squares `k3_sandwich_plain` does not take the call. At 4 MiB a `[T, 7168]` call goes one-shot up to `T` = 18 at `W` = 16 (73 at `W` = 4) and a `[T, 3584]` one up to 36 (146); at 1 MiB `[T, 7168]` up to 4 (18) and `[T, 3584]` up to 9 (36). -- Gaps against A1 (the op is unchanged by this entry): the schema marks `comm_buffer` mutable `(a!)` but not - `buffer_flags`, which every call advances (A1 item 4); the op does not check that the call fits `comm_buffer` (the +- Gaps (the op is unchanged by this entry): the schema marks `comm_buffer` mutable `(a!)` but not `buffer_flags`, + which every call advances; the op does not check that the call fits `comm_buffer` (the attention-residual and all-gather ops do), so a direct op call over one buffer writes past it — the wrapper's `required_buffer_bytes` check is the guard; `MnnvlWorkspace.create` accepts any multiple of 16 bytes, but the two-shot broadcast stage starts at `buffer_bytes / 2` and is accessed in 16-byte vectors, so a two-shot call needs diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml index b18ab85c9d70..2d628ad02393 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml @@ -277,3 +277,38 @@ entries: - path: comm/mnnvl_allreduce_attn_res.py impl: torch.ops.trtllm.mnnvl_allreduce_attn_res summary: "One-shot MNNVL all-reduce with Kimi K3's residual update as its epilogue (updated = prefix + sum over ranks, normed = RMSNorm of the attention-residual selection over the snapshot bank and updated), over a caller-owned MnnvlWorkspace: three Lamport buffers rotated by every call of every MNNVL op that shares the object, so every rank makes the same calls on it in the same order; a swapped pair of same-shaped calls is silently wrong on every rank, two ranks ordering calls on two workspaces differently on one stream deadlock" + + # Kimi K3 decode collectives and MoE: stateful entries with a ## State section and caller-owned state objects + # (K3SandwichWorkspace, MnnvlWorkspace, K3LatentExchange, K3MoeHeadWorkspace, K3MoeState / K3MoeWideState and their + # layers). The state types are not entries: they launch nothing per call. + - path: comm/k3_sandwich_oproj.py + impl: torch.ops.trtllm.k3_sandwich_oproj + summary: "Kimi K3's post-attention sandwich for at most 8 decode tokens in one kernel: the row-parallel attention output projection, its TP all-reduce and the residual update (attention-residual selection + RMSNorm), bit-identical to o_proj followed by the MNNVL one-shot attention-residual all-reduce, over a caller-owned K3SandwichWorkspace: two buffer halves selected by per-CTA call-count parity, shared by every sandwich op of the group, so every rank makes the same calls on it in the same order; a swapped pair of same-shaped calls is silently wrong on every rank" + + - path: comm/k3_sandwich_tail.py + impl: torch.ops.trtllm.k3_sandwich_tail + summary: "Kimi K3's pre-attention sandwich for at most 8 decode tokens in one kernel: the row-parallel MoE tail ([rmsnorm(latent)[:, lo:lo+224] | act] @ tail_weight^T, the latent RMS on the fp32 accumulator), its TP all-reduce and the next layer's residual update, optionally storing the pre-norm mixture or updated into a capture tap and updated into a snapshot-bank row, over the same caller-owned K3SandwichWorkspace as the other sandwiches and in one call sequence with them" + + - path: comm/k3_sandwich_plain.py + impl: torch.ops.trtllm.k3_sandwich_plain + summary: "A row-parallel projection (Kimi K3 drafter's o_proj, or with swiglu its MLP's SiLU-and-mul and down projection), its TP all-reduce, the residual add and RMSNorm in one kernel for at most 8 decode tokens, in the MNNVL one-shot RESIDUAL_RMS_NORM all-reduce's arithmetic, over the same caller-owned K3SandwichWorkspace as the target's sandwiches and in one call sequence with them" + + - path: comm/mnnvl_fusion_allreduce.py + impl: torch.ops.trtllm.mnnvl_fusion_allreduce + summary: "MNNVL all-reduce over a caller-owned MnnvlWorkspace: the sum over the TP group, or with a residual the residual add + RMSNorm; one-shot up to the call's one_shot_max_bytes, two-shot above, both deterministic; one turn of the workspace's three-buffer Lamport rotation per call, shared with every other MNNVL entry on the object, so every rank makes the same MNNVL calls on it in the same order and a swapped pair of same-shaped calls is silently wrong" + + - path: comm/mnnvl_allgather_split.py + impl: torch.ops.trtllm.mnnvl_allgather_split + summary: "One-shot MNNVL all-gather over a caller-owned MnnvlWorkspace of fp32 rows whose leading columns travel and are gathered as bf16 (round to nearest even, -0.0 as +0.0) and the rest as fp32, in rank order (Kimi K3's sharded MoE head: the latent in bf16, the router logits in fp32); one turn of the workspace's Lamport rotation, in one call order with the MNNVL all-reduces on the object" + + - path: comm/k3_latent_reduce.py + impl: torch.ops.trtllm.k3_latent_reduce + summary: "Kimi K3's latent all-reduce at decode size as the consumer of pushed partials: the sum over the TP group of the routed partial rows every rank's producer stored into a caller-owned K3LatentExchange, bit-identical to the MNNVL one-shot all-reduce of the partials; halves selected by the consumer's call parity, each push followed by exactly one reduce of the same token count on every rank in the same order" + + - path: moe/k3_moe_front.py + impl: torch.ops.trtllm.k3_moe_front + summary: "Kimi K3's MoE front at decode size in one kernel: the row-sharded MoE head GEMV, its all-gather over a caller-owned K3MoeHeadWorkspace (two alternating Lamport buffers rotated by every front call), the top-16 routing and MXFP8 latent as trtllm::k3_route_quant returns them for the gathered head, and the shared experts' gate_up + SiTU-and-mul; optionally publishing per-token ready words that a head_flags build of k3_moe acquires" + + - path: moe/k3_moe.py + impl: tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.op.K3MoeLayer + summary: "Kimi K3's routed experts at decode size: this rank's routed partial from the persistent CuTe DSL kernel k3_moe (FC1 + SiTU + FC2 with the routing-weighted combine over the TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers, read in place) after k3_route_quant or the MoE front, on caller-owned per-rank state: a K3MoeState and one K3MoeLayer per MoE layer for up to 8 tokens, a K3MoeWideState and one K3MoeWideLayer per layer for up to 64; a state's layers share its scratch (left armed by every call) and run in one stream order, each layer's counters are left zero" diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md index 781b5e447381..78ff45874b0f 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md @@ -165,9 +165,9 @@ by the configuration), `compiled` (the compiled `k3_moe`, built by the state's f **Who creates it, and when.** The target, in `post_load_weights`, after the expert weights are final (a layer reads the buffers it was built over; see *What a wrong order does*): `K3MoeState(device, i_tp, num_local, head_flags=False)` once per device and `state.layer(w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale)` -once per MoE layer; likewise `K3MoeWideState(device, i_tp, num_local)` and its layers. Not collective. Eager: -`K3MoeState()` and `K3MoeState.layer()` raise `RuntimeError` under CUDA-graph capture (certified); -`K3MoeWideState()` and `K3MoeWideState.layer()` do not check (*Notes*). The kernel compiles on the first call of each +once per MoE layer; likewise `K3MoeWideState(device, i_tp, num_local)` and its layers. Not collective. Eager: the +four constructors (`K3MoeState()`, `K3MoeState.layer()`, `K3MoeWideState()`, `K3MoeWideState.layer()`) raise +`RuntimeError` under CUDA-graph capture (certified). The kernel compiles on the first call of each state object (seconds), which must be eager: under capture that call raises `RuntimeError` before the `k3_moe` launch and leaves the state uncompiled (certified for both builds). `config` (kernel options for tests and A/B runs) stays `None`. No environment variable selects anything here except PDL (*Metadata consumed*). @@ -265,17 +265,14 @@ Besides the state objects (explicit arguments): rank layout (experts TP4 x EP4: 224 local experts of 896, intermediate 768), with random checkpoint-format MXFP4 experts put through TRT-LLM's own loader. `k3_moe_fused_front`: 4 ranks of one GB200 tray in `tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py` (entry point - `moe/test_modeling_v2_k3_moe_front_op_matrix.py`); its 16-rank receipt (A2) is pending with `moe/k3_moe_front`'s. + `moe/test_modeling_v2_k3_moe_front_op_matrix.py`); its 16-rank receipt is pending with `moe/k3_moe_front`'s. - References: an fp64 reference over the dequantized experts (from the checkpoint-format tensors) and the stock path. The kernel tests (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py`, `test_k3_moe_wide.py`, `test_k3_route_quant.py`) remain the exhaustive numerics; this entry's test copies their references. -- Gaps against A1 (the ops are unchanged by this entry): - - `K3MoeWideState()` and `K3MoeWideState.layer()` do not refuse CUDA-graph capture. Built under capture, their - allocations come from the graph's pool and the slab's arming fills and the counters' zeroing are captured instead - of run, so the first eager call would find the slab unarmed and the counters undefined. +- Gaps (the kernels are unchanged by this entry): - The `k3_moe` launch is not a torch op, so nothing declares what it writes: the state's slab and partial rows, the - layer's counters, and with head_flags the head workspace's `flags[2]` and ready words (A1: `mutates_args` names - every written buffer). `trtllm::k3_route_quant` writes only its new outputs (`mutates_args=()`). + layer's counters, and with head_flags the head workspace's `flags[2]` and ready words (a torch op would name them + in `mutates_args`). `trtllm::k3_route_quant` writes only its new outputs (`mutates_args=()`). - A layer does not record `local_expert_offset`, and nothing ties a state to a stream or checks the inputs' device. - The compiled kernel is per state object, not per configuration: a second state of the same configuration compiles again. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md index 30438984f096..7f4e669ce4d1 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md @@ -216,7 +216,7 @@ Besides `workspace` (an explicit argument): - Certified path: 4 ranks of one GB200 tray (sm_100), one rank per GPU, the head sharded over those 4 ranks (1120 rows per rank, 128-row tiles), the shared activation at TP16's per-rank width (384). Kimi K3 TP16 shards the head over 16 ranks on four trays (280 rows per rank), which runs the half-tile geometry and which only a 16-rank run - reaches; the matrix takes `--world-size` and `--launcher` (A2) and that receipt is pending. + reaches; the matrix takes `--world-size` and `--launcher`, and that receipt is pending. - Test: `tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py` (rank body), collected by `tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_front_op_matrix.py`. It also certifies `moe/k3_moe`'s `k3_moe_fused_front` cells. @@ -228,8 +228,8 @@ Besides `workspace` (an explicit argument): reference for `y` the stock TRTLLM-Gen runner. The kernel test (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py`) covers Gaussian payloads and races the publishing front's epoch read against `k3_moe`'s epoch advance (`check_publish_order`); neither is repeated here. -- `mutates_args` names every buffer the op writes (`ag_uc`, `ag_mc`, `ag_flags`, `ag_ready`). Gaps against A1 (the op - is unchanged by this entry): the compile cache is a module-level dict (result-neutral, documented above); the slots' +- `mutates_args` names every buffer the op writes (`ag_uc`, `ag_mc`, `ag_flags`, `ag_ready`). Gaps (the op is + unchanged by this entry): the compile cache is a module-level dict (result-neutral, documented above); the slots' logit vectors are dead space for this kernel; `ring` is not exposed. -- In the model the head workspace was `head_workspace(mapping)`, a module dict created on the first eager call; this - entry takes the explicit object instead (A1). +- Before this entry the head workspace was `head_workspace(mapping)`, a module dict created on the first eager call; + this entry takes the explicit object instead. diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index 9ff1f0bf3ab5..c5d5bcc35668 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -318,6 +318,11 @@ l0_b200: - kv_cache/test_kv_cache_iteration_stats.py::TestKvCacheIterationStats::test_field_completeness # ------------- Prefix-aware scheduling E2E tests --------------- - kv_cache/test_prefix_aware_scheduling.py::TestServePrefixAwareScheduling::test_multi_round_qa_shared_prefix_smoke + # ------------- Kimi K3 decode MoE kernels and their catalog entry (sm_100) --------------- + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_wide.py + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_route_quant.py + - unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe.py # ------------- Visual Gen tests --------------- - unittest/_torch/cute_dsl_kernels/test_nvfp4_conv3d.py - unittest/_torch/visual_gen/test_media_decode.py diff --git a/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml b/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml index c53cb87f897d..4477c83fd8c7 100644 --- a/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml +++ b/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml @@ -53,6 +53,18 @@ l0_gb200_multi_gpus: - unittest/_torch/multi_gpu/test_mnnvl_allreduce.py::test_mnnvl_workspace_growth_keeps_captured_graphs - unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_allreduce_attn_res_op_matrix.py - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mnnvl_comm.py + # Kimi K3 collective kernels (sm_100) and the catalog entries over caller-owned state, 4 ranks + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_latent_reduce.py + - unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_sandwich_oproj_op_matrix.py + - unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_sandwich_tail_op_matrix.py + - unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_sandwich_plain_op_matrix.py + - unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_fusion_allreduce_op_matrix.py + - unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_allgather_split_op_matrix.py + - unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_latent_reduce_op_matrix.py + - unittest/_torch/modeling_v2/comm/test_modeling_v2_allgather_op_matrix.py + - unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_front_op_matrix.py - unittest/_torch/moe/test_moe_comm.py::TestMoEComm::test_mnnvl_checkpoint_preserves_moe_graph_addresses - unittest/_torch/moe/test_moe_comm.py::TestMoEComm::test_mnnvl_engine_checkpoint_coordination - unittest/_torch/moe/test_moe_comm.py::TestMoEComm::test_mnnvl_checkpoint_failure_is_collective_and_bounded diff --git a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe.py b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe.py index f38a1dadd2fe..b4f8f8743980 100644 --- a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe.py +++ b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe.py @@ -791,10 +791,9 @@ def test_unsupported_calls_refused_before_launch(): def test_construction_and_first_call_refuse_capture(): """Under CUDA-graph capture the states refuse to allocate, and a new state refuses its first (compiling) call. - K3MoeState() and K3MoeState.layer() raise RuntimeError; the first call on a new K3MoeState and on a new - K3MoeWideState raises RuntimeError before the k3_moe launch (k3_route_quant, compiled eagerly first, is captured - and discarded with the graph); neither state is compiled afterwards. K3MoeWideState() and K3MoeWideState.layer() - do not check for capture (see the contract's Notes), so they are built eagerly here. + K3MoeState(), K3MoeState.layer(), K3MoeWideState() and K3MoeWideState.layer() raise RuntimeError; the first call + on a new K3MoeState and on a new K3MoeWideState raises RuntimeError before the k3_moe launch (k3_route_quant, + compiled eagerly first, is captured and discarded with the graph); neither state is compiled afterwards. """ proc, _, bias = _experts() fresh = K3MoeState(_device(), I_TP, E_LOCAL) @@ -815,6 +814,10 @@ def test_construction_and_first_call_refuse_capture(): K3MoeState(_device(), I_TP, E_LOCAL) with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): fresh.layer(*_weights(proc)) + with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): + K3MoeWideState(_device(), I_TP, E_LOCAL) + with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): + fresh_wide.layer(*_weights(proc)) with pytest.raises(RuntimeError, match="compiles on its first call"): _decode_call(fresh_layer, lg4, x4) with pytest.raises(RuntimeError, match="compiles on its first call"): From 84fc18f62c44bc06b9a9438b61b41abb6024ea25 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 19:38:25 -0700 Subject: [PATCH 062/161] [None][test] Kimi K3 MNNVL kernel test: the split all-gather, one-shot reference test_k3_mnnvl_comm.py gets back the split all-gather check it has in the K3 stack: - trtllm::mnnvl_allgather_split is compared bit for bit with the same gather done on the host, at the sharded MoE head's per-rank widths for TP16 and TP4, at every M; - six more calls are interleaved with all-reduces and attention-residual all-reduces on the same Lamport rotation; - one rank's perturbed input must change the result. The unfused attn_res reference now sends its all-reduce one-shot through one_shot_max_bytes, the fused op's order at any world size, instead of asserting at most 8 ranks. Signed-off-by: Vasanth Sabavat --- .../kimi_k3/test_k3_mnnvl_comm.py | 83 +++++++++++++++++-- 1 file changed, 75 insertions(+), 8 deletions(-) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mnnvl_comm.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mnnvl_comm.py index 0004cff4fe64..aaa750f960b5 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mnnvl_comm.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mnnvl_comm.py @@ -17,12 +17,16 @@ trtllm::mnnvl_allreduce_attn_res (MNNVLAllReduce.allreduce_attn_res_rmsnorm): against the unfused path it replaces, the MNNVL all-reduce then trtllm::attn_res_add_rmsnorm_fwd (attn_res_rmsnorm_fwd without a prefix sum): the updated prefix sum bit for bit, the normed rows within 1e-2 (max |d| / max |ref|), both against an fp32 port of the - attention-residual selection within 2e-2; 0, 1, 3, 8 and 11 snapshots, with and without the prefix; + attention-residual selection within 2e-2; 0, 1, 3, 8 and 11 snapshots, with and without the prefix; the reference + all-reduce sent one-shot (the fused op's order); + trtllm::mnnvl_allgather_split (MNNVLAllReduce.allgather_split): bit for bit against the same gather on the host (bf16 + columns rounded to nearest, -0.0 arriving as +0.0), at the K3 sharded MoE head's per-rank widths for TP16 (224 bf16 + + 56 fp32 columns) and TP4 (896 + 224), interleaved with all-reduces on the same Lamport rotation; each with run-to-run identical bits, the same bits on every rank, each M's rows bit-identical to the same rows of the 64-row call, and one rank's perturbed input changing every rank's result. Run under pytest (a pool of 4 MPI workers) or directly, one process per GPU: - srun -N1 -n4 --mpi=pmix python3 test_k3_mnnvl_comm.py [attn_res] + srun -N1 -n4 --mpi=pmix python3 test_k3_mnnvl_comm.py [attn_res allgather] """ import hashlib @@ -46,7 +50,7 @@ MPI.pickle.__init__(cloudpickle.dumps, cloudpickle.loads, pickle.HIGHEST_PROTOCOL) WORLD = 4 -H = 7168 +H, LATENT, EXPERTS = 7168, 3584, 896 EPS, OUT_EPS = 1e-6, 1e-5 M_CASES = list(range(1, 9)) + [16, 32, 64] M_MAX = 64 @@ -99,12 +103,13 @@ def _context(): def _allreduce(ctx, x): - """The plain MNNVL all-reduce. Up to 8 ranks its one-shot and two-shot kernels both sum the ranks in rank order in - fp32, the fused op's order, so the reference is exact whichever one the size picks.""" + """The plain MNNVL all-reduce, sent one-shot (the fused ops' order; above 8 ranks two-shot sums the ranks in + another order).""" from tensorrt_llm._torch.distributed import AllReduceParams - assert ctx.world <= 8, "above 8 ranks the two-shot kernel sums the ranks in another order" - return ctx.mnnvl(x, AllReduceParams()) + return ctx.mnnvl( + x, AllReduceParams(), one_shot_max_bytes=x.numel() * ctx.world * x.element_size() + ) def _all_ranks(ctx, good) -> bool: @@ -199,7 +204,69 @@ def fused(rows, part=None): return results -CHECKS = {"attn_res": check_attn_res} +def _gather_input(rows, bf16_cols, fp32_cols, rank): + gen = torch.Generator(device="cuda").manual_seed(7919 * bf16_cols + 131 * rank + fp32_cols) + x = torch.randn(rows, bf16_cols + fp32_cols, generator=gen, device="cuda") * 3 + flat = x.view(-1) + flat[0], flat[1], flat[2], flat[-1] = -0.0, 0.0, 3.0e38, -0.0 + return x.contiguous() + + +def _host_gather(inputs, bf16_cols): + """The same gather on the host: bf16 part rounded to nearest, -0.0 as +0.0.""" + return ( + torch.cat([x[:, :bf16_cols].bfloat16() + 0.0 for x in inputs], dim=1), + torch.cat([x[:, bf16_cols:] + 0.0 for x in inputs], dim=1), + ) + + +def check_allgather(ctx): + results = [] + # The sharded MoE head per rank: 3584 / TP latent columns (bf16 after the gather) + 896 / TP router logits. + heads = ( + ("tp16_head", (LATENT // 16, EXPERTS // 16)), + ("tp4_head", (LATENT // 4, EXPERTS // 4)), + ) + for label, (bf16_cols, fp32_cols) in heads: + mine64 = _gather_input(M_MAX, bf16_cols, fp32_cols, ctx.rank) + b64, f64 = ctx.mnnvl.allgather_split(mine64, bf16_cols) + for m in M_CASES: + mine = mine64[:m].contiguous() + everyone = [torch.from_numpy(a).cuda() for a in ctx.comm.allgather(mine.cpu().numpy())] + want_b, want_f = _host_gather(everyone, bf16_cols) + got_b, got_f = ctx.mnnvl.allgather_split(mine, bf16_cols) + exact = _same(got_b, want_b) and _same(got_f, want_f) + interleaved = True + partial = torch.randn(m, H, device="cuda").bfloat16() + for i in range(6): # two turns of the three-buffer Lamport rotation, mixed with the other collectives + _allreduce(ctx, partial) + if i % 2: + block = torch.randn(2, m, H, device="cuda").bfloat16() + ones = torch.ones(H, device="cuda").bfloat16() + ctx.mnnvl.allreduce_attn_res_rmsnorm(partial, None, block, ones, ones, ones, EPS, EPS) + b, f = ctx.mnnvl.allgather_split(mine, bf16_cols) + interleaved &= _same(b, want_b) and _same(f, want_f) + bad = mine.clone() + if ctx.rank == min(1, ctx.world - 1): + bad.view(-1)[3] += 1.0 + bad_b, bad_f = ctx.mnnvl.allgather_split(bad, bf16_cols) + row = dict( + op="mnnvl_allgather_split", + case=label, + M=m, + exact=exact, + interleaved_x6=interleaved, + rows_as_m64=_same(got_b, b64[:m]) and _same(got_f, f64[:m]), + control=not (_same(bad_b, want_b) and _same(bad_f, want_f)), + ) + row["ok"] = _all_ranks(ctx, row["exact"] and interleaved and row["rows_as_m64"]) and _all_ranks( + ctx, row["control"] + ) + results.append(row) + return results + + +CHECKS = {"attn_res": check_attn_res, "allgather": check_allgather} def _run_checks(names): From 7a16d03d4090ce80e471af95ae9af1c7b2aa4857 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 19:39:11 -0700 Subject: [PATCH 063/161] [None][test] Kimi K3 kernel lints: the MoE front, k3_moe and sandwich kernels The two source lints get this PR's kernels: - test_k3_cluster_waits.py: the mailboxes that other CTAs complete with st.async in k3_moe_front (mail_full) and k3_sandwich_kernel (b_ready, lat_full, stage 0 of rms_full, and the per-statistic mailboxes) must be waited on at cluster scope. - test_k3_tcgen05_fences.py: k3_moe_front, k3_moe_kernel and k3_sandwich_kernel keep their tcgen05 fences. A new check (GENERIC_FED) requires a fence.proxy.async between a wait on a barrier whose stages generic-proxy writes fill and the tcgen05 read after it: k3_moe's FC1 activation scales, copied in by cp.async and read by tcgen05.cp. The check comes with a test of the checker itself. Signed-off-by: Vasanth Sabavat --- .../kimi_k3/test_k3_cluster_waits.py | 13 ++++ .../kimi_k3/test_k3_tcgen05_fences.py | 62 ++++++++++++++++++- 2 files changed, 74 insertions(+), 1 deletion(-) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_cluster_waits.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_cluster_waits.py index f9433e11a165..e9860683aab6 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_cluster_waits.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_cluster_waits.py @@ -29,6 +29,19 @@ # Kernel module -> the barriers in it that other CTAs complete with st.async. MAILBOXES = { "tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv.k3_ctm_gemv_kernel": ("mail_full",), + "tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.k3_moe_front": ("mail_full",), + # rms_full: stage 0 only (filled by the cluster CTAs' st.async pushes); its other stages are completed by this + # CTA's own arrive or bulk copy. + "tensorrt_llm._torch.cute_dsl_kernels.k3_sandwich.k3_sandwich_kernel": ( + "b_ready", + "lat_full", + r"rms_full\.subview\(0", + "mb_rms", + "mb_snap", + "mb_ts", + "mb_upd", + "mb_sq", + ), } # Kernel module -> the header of the blocks whose mailbox waits are on this CTA's own arrival and keep CTA scope (a # kernel that arrives on its mailbox itself and acquires the other CTAs' data through a counter). diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py index 613e5d559a89..6abcb564c141 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py @@ -16,7 +16,9 @@ sync patterns): a tcgen05 operation issued after an mbarrier wait follows ``tcgen05.fence::after_thread_sync``, and a thread whose TMEM loads another thread's tcgen05 work must not overtake (an mbarrier arrive or a CTA barrier after ``tcgen05.wait::ld``) issues ``tcgen05.fence::before_thread_sync`` first. Without them ptxas may move the TMEM access -across the synchronization; no test of values can see it, so the kernels' sources are read. +across the synchronization; no test of values can see it, so the kernels' sources are read. A tcgen05 operation that +reads shared memory filled by generic-proxy writes (cp.async) also follows ``fence.proxy.async.shared::cta`` after the +wait that orders those writes. pytest test_k3_tcgen05_fences.py """ @@ -29,6 +31,9 @@ KERNELS = [ "tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv.k3_ctm_gemv_kernel", "tensorrt_llm._torch.cute_dsl_kernels.k3_decode_gemv.k3_decode_gemv_kernel", + "tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.k3_moe_front", + "tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.k3_moe_kernel", + "tensorrt_llm._torch.cute_dsl_kernels.k3_sandwich.k3_sandwich_kernel", ] # Thread syncs that order other threads' work before a tcgen05 operation of this thread. Waits on a barrier only TMA @@ -45,6 +50,13 @@ LOAD_WAIT = "Tcgen05Wait.LOAD" SYNC = re.compile(r"mbarrier_arrive\(|barrier_cta_sync\(|cute\.arch\.barrier\(") +# Kernel module -> the barriers whose stages hold shared memory that generic-proxy writes fill and a tcgen05 operation +# reads after the wait (k3_moe: the FC1 activation scales, copied in by cp.async and read by tcgen05.cp). +GENERIC_FED = { + "tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.k3_moe_kernel": ("ab_full_fc1",), +} +PROXY_FENCE = re.compile(r"fence_proxy\(\s*(\"async_shared\"|prims\.Proxy\.ASYNC_SHARED)") + def violations(lines): """(line number, rule) of every tcgen05 operation reached from a wait without the after-fence, and every arrive / @@ -79,6 +91,25 @@ def violations(lines): return out +def generic_reads(lines, barriers): + """(line number of the wait, whether it is fenced) for each wait on a listed barrier (named on the wait's line or + the next) that reaches a tcgen05 operation; fenced means ``fence.proxy.async`` lies between the two.""" + code = [ln.split("#", 1)[0] for ln in lines] + out = [] + for i, ln in enumerate(code): + window = ln + (code[i + 1] if i + 1 < len(code) else "") + if WAIT.search(ln) and any(re.search(rf"\b{name}\b", window) for name in barriers): + fenced = False + for later in code[i + 1 :]: + if later.lstrip().startswith("def "): + break + fenced = fenced or bool(PROXY_FENCE.search(later)) + if TCGEN05_OP.search(later): + out.append((i + 1, fenced)) + break + return out + + def test_checker_catches_the_patterns(): """The checker flags both missing fences (and accepts the fenced forms).""" bad = [ @@ -109,3 +140,32 @@ def test_kernel_fences(module): with open(path) as f: found = violations(f.read().split("\n")) assert not found, f"{path}: {found}" + + +def test_checker_catches_an_unfenced_generic_read(): + """The proxy check flags a tcgen05 read after the wait without fence.proxy.async (and accepts the fenced form).""" + wait = [ + "def f():", + " while not cute.arch.mbarrier_try_wait(", + " ab_full_fc1.subview(stage).data_ptr(), phase", + " ):", + " pass", + ] + copy = [" prims.tcgen05_cp(shape, tmem_ptr, desc)"] + fence = [' prims.fence_proxy("async_shared", space=prims.SharedSpace.shared_cta)'] + assert generic_reads(wait + copy, ("ab_full_fc1",)) == [(2, False)] + assert generic_reads(wait + fence + copy, ("ab_full_fc1",)) == [(2, True)] + + +@pytest.mark.parametrize( + "module", list(GENERIC_FED), ids=[m.rsplit(".", 1)[1] for m in GENERIC_FED] +) +def test_generic_proxy_reads_fenced(module): + path = importlib.util.find_spec(module).origin + with open(path) as f: + reads = generic_reads(f.read().split("\n"), GENERIC_FED[module]) + assert reads, f"{path}: no wait on {GENERIC_FED[module]} reaches a tcgen05 operation" + unfenced = [line for line, fenced in reads if not fenced] + assert not unfenced, ( + f"{path}: no fence.proxy.async between the waits at lines {unfenced} and their tcgen05 reads" + ) From a35362ef961207c1bc24118e519049728a46b22f Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 19:41:37 -0700 Subject: [PATCH 064/161] [None][test] Kimi K3 collective tests: names codespell accepts, lines under 120 Loop variables `fo` / `wo` in test_k3_sandwich.py become `f_out` / `w_out` and the matrices' `statics` become `static_bufs`, which codespell no longer flags. Two lines over 120 columns are wrapped (the SiTU reference in test_k3_fused_moe.py keeps its evaluation order). Signed-off-by: Vasanth Sabavat --- .../_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py | 3 ++- .../_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py | 7 ++++--- .../_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py | 4 ++-- .../_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py | 8 ++++---- .../modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py | 6 +++--- .../modeling_v2/comm/_mnnvl_fusion_allreduce_op_matrix.py | 6 +++--- 6 files changed, 18 insertions(+), 16 deletions(-) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py index a056a8e88f00..cd015657b644 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py @@ -150,7 +150,8 @@ def _reference(raw, x_deq, ids, weights): xe = x_deq[tok].double() up = (xe @ _deq_w(raw["up"][e], raw["up_s"][e]).double().t()).float() gate = (xe @ _deq_w(raw["gate"][e], raw["gate_s"][e]).double().t()).float() - act = GATE_CAP * torch.tanh(gate / GATE_CAP) * torch.sigmoid(gate) * (LINEAR_CAP * torch.tanh(up / LINEAR_CAP)) + act = GATE_CAP * torch.tanh(gate / GATE_CAP) * torch.sigmoid(gate) + act = act * (LINEAR_CAP * torch.tanh(up / LINEAR_CAP)) y = (_requant(act).double() @ _deq_w(raw["down"][e], raw["down_s"][e]).double().t()).float() out.index_add_(0, tok, y * weights[tok, slot].float().unsqueeze(1)) return out diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py index 56a5772639c5..da975bfcbc96 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py @@ -14,9 +14,10 @@ # limitations under the License. """trtllm::k3_moe_front and K3MoeLayer.front (the Kimi K3 MoE front: sharded head GEMV, head all-gather, top-16 routing, MXFP8 latent, shared gate_up + SiTU; then k3_moe on its grid), one process per GPU over the TP group of this -run, at every M in 1..8, over one K3MoeHeadWorkspace. The head is sharded over the group (TP W: 3584 / W latent + 896 / W router rows and -2 x 6144 / W shared rows per rank; W = 4 on one GB200 tray, the model's TP16 shapes with 16 processes); the routed -experts are one rank of experts TP4 x EP4 (224 local experts, intermediate 768), as in the TP16 deployment. +run, at every M in 1..8, over one K3MoeHeadWorkspace. The head is sharded over the group (TP W: 3584 / W latent + +896 / W router rows and 2 x 6144 / W shared rows per rank; W = 4 on one GB200 tray, the model's TP16 shapes with 16 +processes); the routed experts are one rank of experts TP4 x EP4 (224 local experts, intermediate 768), as in the TP16 +deployment. front : against the unfused chain (the head GEMV in fp32 torch -> the gather -> trtllm::kimi_k3_noaux_tc_mxfp8_quant; shared: cuBLAS gate_up -> trtllm::situ_and_mul): top-16 ids per token (a mismatch only at a reference diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py index a111e9d20ff0..983d688eb35a 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py @@ -424,7 +424,7 @@ def run(): ctx.comm.Barrier() wrapped = run() after = int(flags[0].item()) - same = all(_same(a, b) for fo, wo in zip(fresh, wrapped) for a, b in zip(fo, wo)) + same = all(_same(a, b) for f_out, w_out in zip(fresh, wrapped) for a, b in zip(f_out, w_out)) row = dict(op="k3_sandwich_wrap", case="int32_wrap", M=8, eq_fresh=same, crossed=after < 0, calls=len(calls)) row["ok"] = _all_ranks(ctx, same and row["crossed"]) return [row] @@ -469,7 +469,7 @@ def run(): flags[0] = 2**31 - 4 + (count & 1) torch.cuda.synchronize() wrapped, wrapped_kept = run() - same = all(_same(a, b) for fo, wo in zip(fresh, wrapped) for a, b in zip(fo, wo)) + same = all(_same(a, b) for f_out, w_out in zip(fresh, wrapped) for a, b in zip(f_out, w_out)) kept = all(fresh_kept) and all(wrapped_kept) row = dict(op="k3_sandwich_tail_fold", case="count_wrap", M=8, eq_fresh=same, other_words_kept=kept, count_after=int(flags[0].item())) # fmt: skip diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py index 487fc265ce20..4e5af9f0f4ae 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py @@ -544,9 +544,9 @@ def check_graph_capture_and_replay() -> None: buffers are empty, WS_B's epoch advanced once per head_flags call, replayed or eager, with every ready word at it, and both states armed. """ - statics = [torch.zeros(8, HIDDEN, dtype=torch.bfloat16, device="cuda") for _ in range(LAYERS)] + static_bufs = [torch.zeros(8, HIDDEN, dtype=torch.bfloat16, device="cuda") for _ in range(LAYERS)] shells = [ - Call(0, 0, layer=layer, kind=KINDS[layer], x=statics[layer]) for layer in range(LAYERS) + Call(0, 0, layer=layer, kind=KINDS[layer], x=static_bufs[layer]) for layer in range(LAYERS) ] def step(): @@ -564,7 +564,7 @@ def step(): alone_eager = alone_results(eagers, WS_A, "eager between replays") epoch0 = epoch_and_words(WS_B)[0] - for static, c in zip(statics, warm): + for static, c in zip(static_bufs, warm): static.copy_(c.x) R.barrier() step() # every first call of this step eager @@ -577,7 +577,7 @@ def step(): R.barrier() flag_calls = 1 # the warm-up step's for r in range(REPLAYS): - for static, c in zip(statics, reps[r]): + for static, c in zip(static_bufs, reps[r]): static.copy_(c.x) R.barrier() graph.replay() diff --git a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py index a8ecfd1fcedd..6dee0eb91c1f 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py @@ -502,10 +502,10 @@ def check_graph_capture_and_replay() -> None: AR(6003, 32, LATENT, path="two"), AG(6004, 32, b, f), ] - statics = [c.static() for c in calls] + static_bufs = [c.static() for c in calls] def step(): - return [c.run(WS_B, **s) for c, s in zip(calls, statics)] + return [c.run(WS_B, **s) for c, s in zip(calls, static_bufs)] outs = step() # every call once eagerly, outside capture for i, (c, got) in enumerate(zip(calls, outs)): @@ -526,7 +526,7 @@ def step(): lambda seed: AG(seed, 16, *SPLITS[-1]), ) for rep in range(8): - for i, (c, s) in enumerate(zip(calls, statics)): + for i, (c, s) in enumerate(zip(calls, static_bufs)): c.refill(c.fresh(7000 + 100 * rep + i), s) R.barrier() graph.replay() diff --git a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_fusion_allreduce_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_fusion_allreduce_op_matrix.py index 75057e65e335..a72cd25c6ae8 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_fusion_allreduce_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_fusion_allreduce_op_matrix.py @@ -561,10 +561,10 @@ def check_graph_capture_and_replay() -> None: AR(6004, 32, H_MODEL, path="two"), AR(6005, 16, H_MODEL, residual=True, path="two"), ] - statics = [c.static() for c in calls] + static_bufs = [c.static() for c in calls] def step(): - return [c.run(WS_B, **s) for c, s in zip(calls, statics)] + return [c.run(WS_B, **s) for c, s in zip(calls, static_bufs)] outs = step() # every call once eagerly, outside capture for i, (c, got) in enumerate(zip(calls, outs)): @@ -585,7 +585,7 @@ def step(): lambda seed: AR(seed, 64, H_LATENT, path="one"), ) for rep in range(8): - for i, (c, s) in enumerate(zip(calls, statics)): + for i, (c, s) in enumerate(zip(calls, static_bufs)): c.refill(c.fresh(7000 + 100 * rep + i), s) R.barrier() graph.replay() From 22629dba797bc660d0ebbaeaf711bcfcc4b331ea Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 19:42:28 -0700 Subject: [PATCH 065/161] [None][chore] Kimi K3 collectives: format with the repo's pre-commit hooks The files this PR adds or changes, put through main's pre-commit (ruff format at 100 columns, ruff's import order). Every reformatted Python file parses to the same AST as before (ast.dump compared for all 49 files), so behaviour is unchanged. The second pre-commit run passes clean. Signed-off-by: Vasanth Sabavat --- .../catalog/comm/mnnvl_allgather_split.py | 8 +- .../catalog/comm/mnnvl_fusion_allreduce.py | 4 +- .../modeling_v2/catalog/moe/k3_moe.py | 8 +- .../modeling_v2/catalog/moe/k3_moe_front.py | 5 +- .../k3_fused_moe/latent_op.py | 8 +- .../cute_dsl_kernels/k3_fused_moe/op.py | 40 +++++-- .../k3_sandwich/k3_sandwich_kernel.py | 10 +- .../_torch/cute_dsl_kernels/k3_sandwich/op.py | 22 +++- .../kimi_k3/test_k3_fused_moe.py | 27 ++++- .../kimi_k3/test_k3_mnnvl_comm.py | 14 ++- .../kimi_k3/test_k3_moe_front.py | 11 +- .../kimi_k3/test_k3_moe_wide.py | 5 +- .../kimi_k3/test_k3_route_quant.py | 31 +++++- .../kimi_k3/test_k3_sandwich.py | 105 ++++++++++++++---- .../comm/_k3_moe_front_op_matrix.py | 4 +- .../modeling_v2/comm/_k3_sandwich_common.py | 76 ++++++++++--- .../comm/_k3_sandwich_oproj_op_matrix.py | 50 +++++++-- .../comm/_k3_sandwich_plain_op_matrix.py | 27 ++++- .../comm/_k3_sandwich_tail_op_matrix.py | 26 ++++- 19 files changed, 381 insertions(+), 100 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.py index 6a4ac034f113..8a9719f417ba 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.py @@ -14,7 +14,9 @@ __all__ = ["MnnvlWorkspace", "mnnvl_allgather_split", "required_buffer_bytes"] -def required_buffer_bytes(num_tokens: int, bf16_columns: int, fp32_columns: int, world_size: int) -> int: +def required_buffer_bytes( + num_tokens: int, bf16_columns: int, fp32_columns: int, world_size: int +) -> int: """Bytes of one Lamport buffer a call occupies: every rank's rows, the bf16 columns at 2 bytes and the fp32 columns at 4.""" return num_tokens * world_size * (bf16_columns * 2 + fp32_columns * 4) @@ -31,7 +33,9 @@ def mnnvl_allgather_split( Takes one turn of the workspace's Lamport rotation, as an all-reduce on it does: every rank of the group makes the same MNNVL calls on it in the same order.""" num_tokens, columns = input.shape - need = required_buffer_bytes(num_tokens, bf16_columns, columns - bf16_columns, workspace.world_size) + need = required_buffer_bytes( + num_tokens, bf16_columns, columns - bf16_columns, workspace.world_size + ) if need > workspace.buffer_bytes: raise ValueError( f"mnnvl_allgather_split: the call needs {need} bytes per Lamport buffer, the workspace has " diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.py index 2bde4b87db50..9c48175014d3 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.py @@ -46,7 +46,9 @@ def mnnvl_fusion_allreduce( raise ValueError("mnnvl_fusion_allreduce: residual, norm_weight and eps go together") hidden = input.shape[-1] num_tokens = input.numel() // hidden - need = required_buffer_bytes(num_tokens, hidden, workspace.world_size, input.dtype, one_shot_max_bytes) + need = required_buffer_bytes( + num_tokens, hidden, workspace.world_size, input.dtype, one_shot_max_bytes + ) if need > workspace.buffer_bytes: raise ValueError( f"mnnvl_fusion_allreduce: the call needs {need} bytes per Lamport buffer, the workspace has " diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.py index 31b5daeb99c8..f2190722aac7 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.py @@ -19,7 +19,9 @@ K3MoeWideState, is_supported, ) -from tensorrt_llm._torch.cute_dsl_kernels.k3_route_quant import op as _k3_route_quant_op # noqa: F401 +from tensorrt_llm._torch.cute_dsl_kernels.k3_route_quant import ( + op as _k3_route_quant_op, # noqa: F401 +) __all__ = [ "K3MoeHeadWorkspace", @@ -46,7 +48,9 @@ def k3_moe( ``router_logits`` (fp32 ``[M, 896]``) and ``latent`` (bf16 ``[M, 3584]``), then ``k3_moe`` on ``layer``'s experts (global ids ``[local_expert_offset, local_expert_offset + num_local)``). Writes ``layer``'s state's scratch (left armed) and ``layer``'s counters (left zero).""" - return layer(latent, router_logits, e_score_correction_bias, local_expert_offset, routed_scaling_factor) + return layer( + latent, router_logits, e_score_correction_bias, local_expert_offset, routed_scaling_factor + ) def k3_moe_fused_front( diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.py index c88abf6d0982..0d238c29e28f 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.py @@ -9,7 +9,10 @@ # Importing front_op registers trtllm::k3_moe_front. front_weight packs the front's one weight at load time; # weight_supported reads metadata only. -from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.front_op import front_weight, weight_supported +from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.front_op import ( + front_weight, + weight_supported, +) from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.op import K3MoeHeadWorkspace __all__ = ["K3MoeHeadWorkspace", "front_weight", "k3_moe_front", "weight_supported"] diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py index b43918d9dc51..7b1c025ef5da 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py @@ -91,7 +91,9 @@ def create(cls, mapping, fabric_handle: Optional[bool] = None) -> "K3LatentExcha ) words = _kernel().buffer_words(mapping.tp_size) - use_fabric_handle = mapping.is_multi_node() if fabric_handle is None else bool(fabric_handle) + use_fabric_handle = ( + mapping.is_multi_node() if fabric_handle is None else bool(fabric_handle) + ) comm = _get_mnnvl_workspace_comm(mapping) error: Optional[Exception] = None exchange = None @@ -115,7 +117,9 @@ def create(cls, mapping, fabric_handle: Optional[bool] = None) -> "K3LatentExcha error = exc # Also the barrier that keeps any rank from pushing into a peer's buffer before the peer has emptied it. if not _mnnvl_workspace_all_succeeded(comm, error is None): - raise RuntimeError("K3LatentExchange: allocation failed on at least one rank") from error + raise RuntimeError( + "K3LatentExchange: allocation failed on at least one rank" + ) from error return exchange def push_args(self): diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py index 3fc6d6c760fe..0e3b58cbe7cd 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py @@ -162,7 +162,9 @@ def create(cls, mapping, fabric_handle: Optional[bool] = None) -> "K3MoeHeadWork from . import k3_route_quant_ag as layout words = layout.workspace_words(mapping.tp_size) - use_fabric_handle = mapping.is_multi_node() if fabric_handle is None else bool(fabric_handle) + use_fabric_handle = ( + mapping.is_multi_node() if fabric_handle is None else bool(fabric_handle) + ) comm = _get_mnnvl_workspace_comm(mapping) error: Optional[Exception] = None workspace = None @@ -188,7 +190,9 @@ def create(cls, mapping, fabric_handle: Optional[bool] = None) -> "K3MoeHeadWork error = exc # Also the barrier that keeps any rank from pushing into a peer's buffer before the peer has emptied it. if not _mnnvl_workspace_all_succeeded(comm, error is None): - raise RuntimeError("K3MoeHeadWorkspace: allocation failed on at least one rank") from error + raise RuntimeError( + "K3MoeHeadWorkspace: allocation failed on at least one rank" + ) from error return workspace @@ -212,7 +216,9 @@ def __init__( config: Optional[dict] = None, ): if torch.cuda.is_current_stream_capturing(): - raise RuntimeError("K3MoeState allocates its scratch: build it outside CUDA-graph capture") + raise RuntimeError( + "K3MoeState allocates its scratch: build it outside CUDA-graph capture" + ) # One persistent CTA per SM (config "num_ctas" caps it, e.g. for a grid-size A/B). num_ctas = torch.cuda.get_device_properties(device).multi_processor_count cfg = { @@ -225,7 +231,9 @@ def __init__( cfg.update(config or {}) self.mod = mod = _kernel_module(cfg) if mod.FUSED_AR or mod.FOLD or mod.LAT_SLAB or mod.WIDE: - raise ValueError("K3MoeState is the M <= 8 build without the fused all-reduce, the fold or the slab") + raise ValueError( + "K3MoeState is the M <= 8 build without the fused all-reduce, the fold or the slab" + ) self.device = device self.i_tp = i_tp self.num_local = num_local @@ -259,7 +267,9 @@ def __init__( # The route+quant kernel triggers k3_moe's launch right after its own grid dependency: # k3_moe waits for the whole route+quant grid before reading its outputs. self.route_kwargs = {"early_trigger": True} if mod.USE_PDL else {} - from ..k3_route_quant import op as _k3_route_quant_op # noqa: F401 (registers trtllm::k3_route_quant) + from ..k3_route_quant import ( + op as _k3_route_quant_op, # noqa: F401 (registers trtllm::k3_route_quant) + ) self.compiled = None @@ -287,7 +297,9 @@ def __init__( w2_weight_scale: torch.Tensor, ): if torch.cuda.is_current_stream_capturing(): - raise RuntimeError("K3MoeLayer allocates its counters: build it outside CUDA-graph capture") + raise RuntimeError( + "K3MoeLayer allocates its counters: build it outside CUDA-graph capture" + ) ok, why = is_supported( w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale, state.num_local ) @@ -324,7 +336,9 @@ def __call__( ``[local_expert_offset, local_expert_offset + num_local)``.""" st = self.state if st.head_flags: - raise ValueError("a head_flags build takes the front's ready words: call K3MoeLayer.front") + raise ValueError( + "a head_flags build takes the front's ready words: call K3MoeLayer.front" + ) _check_tokens(hidden_states) ids, weights, x_fp8, x_sf = torch.ops.trtllm.k3_route_quant( router_logits.contiguous(), e_score_correction_bias, hidden_states.contiguous(), @@ -361,7 +375,9 @@ def front( flag_in = None if st.head_flags: flag_in = (_view(head.ready.view(-1), 16, 0), _view(head.flags.view(-1), 16, 0)) - y = self._launch(ids, weights, x_fp8, x_sf, local_expert_offset, routed_scaling_factor, flag_in) + y = self._launch( + ids, weights, x_fp8, x_sf, local_expert_offset, routed_scaling_factor, flag_in + ) return y, shared def _launch(self, ids, weights, x_fp8, x_sf, local_offset, scale, flag_in=None) -> torch.Tensor: @@ -427,7 +443,9 @@ class K3MoeWideState: def __init__(self, device: torch.device, i_tp: int, num_local: int, use_pdl: bool = True): if torch.cuda.is_current_stream_capturing(): - raise RuntimeError("K3MoeWideState allocates its scratch: build it outside CUDA-graph capture") + raise RuntimeError( + "K3MoeWideState allocates its scratch: build it outside CUDA-graph capture" + ) num_ctas = torch.cuda.get_device_properties(device).multi_processor_count config = { "i_tp": i_tp, @@ -508,7 +526,9 @@ def __init__( w2_weight_scale: torch.Tensor, ): if torch.cuda.is_current_stream_capturing(): - raise RuntimeError("K3MoeWideLayer allocates its counters: build it outside CUDA-graph capture") + raise RuntimeError( + "K3MoeWideLayer allocates its counters: build it outside CUDA-graph capture" + ) ok, why = is_supported( w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale, state.num_local ) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/k3_sandwich_kernel.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/k3_sandwich_kernel.py index fb4488da2b36..2af8a6c5d67a 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/k3_sandwich_kernel.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/k3_sandwich_kernel.py @@ -145,9 +145,7 @@ LAT_SLICE_VECS = LAT_SLICE // 8 # its 16-byte vectors per token row LAT_WORDS = 3584 // 2 # int32 words of one latent row LAT_VECS = 3584 // 8 # 16-byte vectors of one latent row -LAT_FLAG_WORDS = ( - 64 # int32: [0] the tail's call count mod 6, [LAT_SCALES + 8 b + t] buffer b of the latent scale slab -) +LAT_FLAG_WORDS = 64 # int32: [0] the tail's call count mod 6, [LAT_SCALES + 8 b + t] buffer b of the latent scale slab LAT_SCALES = 32 LAT_SCALE_BUFS = 3 SCALE_SENTINEL = -1 # 0xFFFFFFFF: an unwritten scale (a computed rsqrt is never this NaN pattern) @@ -1884,7 +1882,11 @@ def sandwich_role( # the scale buffer. Every CTA of this call has read this one (CTA 0's phase 2 needed every CTA's push). if (bx == cutlass.Int32(0)) & (tid == cutlass.Int32(0)): count_period = cutlass.Int32(2 * LAT_SCALE_BUFS) - lat_flags.store((n_calls % count_period + cutlass.Int32(1)) % count_period, idx=0, is_volatile=True) + lat_flags.store( + (n_calls % count_period + cutlass.Int32(1)) % count_period, + idx=0, + is_volatile=True, + ) @cute.kernel diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py index b902e2cbc68f..350865704cc9 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py @@ -235,11 +235,18 @@ def create(cls, mapping, fabric_handle: Optional[bool] = None) -> "K3SandwichLat from . import k3_sandwich_kernel as kernel def arm(flags: torch.Tensor) -> None: - scales = slice(kernel.LAT_SCALES, kernel.LAT_SCALES + kernel.LAT_SCALE_BUFS * MAX_TOKENS) + scales = slice( + kernel.LAT_SCALES, kernel.LAT_SCALES + kernel.LAT_SCALE_BUFS * MAX_TOKENS + ) flags[scales] = kernel.SCALE_SENTINEL return _create_buffer( - cls, mapping, kernel.lat_buffer_words(mapping.tp_size), kernel.LAT_FLAG_WORDS, fabric_handle, arm + cls, + mapping, + kernel.lat_buffer_words(mapping.tp_size), + kernel.LAT_FLAG_WORDS, + fabric_handle, + arm, ) @@ -393,7 +400,16 @@ def supports_tail( @torch.library.custom_op( "trtllm::k3_sandwich_tail", - mutates_args=("ws_uc", "ws_mc", "ws_flags", "x_slab", "lat_uc", "lat_flags", "tap", "updated_out"), + mutates_args=( + "ws_uc", + "ws_mc", + "ws_flags", + "x_slab", + "lat_uc", + "lat_flags", + "tap", + "updated_out", + ), ) def k3_sandwich_tail( latent: torch.Tensor, diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py index cd015657b644..281e93dc9f1a 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py @@ -41,7 +41,10 @@ def _is_sm100() -> bool: H, TOP_K, NUM_EXPERTS, SV = 3584, 16, 896, 32 I_TP, E_LOCAL, MOE_TP, TP_RANK, EP_RANK = 768, 224, 4, 1, 1 # one rank of experts TP4 x EP4 OFFSET = EP_RANK * E_LOCAL -GATE_CAP, LINEAR_CAP = 4.0, 25.0 # the SiTU caps (activation_situ_beta, activation_situ_linear_beta) +GATE_CAP, LINEAR_CAP = ( + 4.0, + 25.0, +) # the SiTU caps (activation_situ_beta, activation_situ_linear_beta) RSF = 2.827 ULP = 2.0**-8 E4M3_MAX = 448.0 @@ -64,7 +67,9 @@ def _rand_mxfp4(rows, k, gen): dot product lands near std 3.""" codes = torch.randint(0, 16, (rows, k), dtype=torch.uint8, device="cuda", generator=gen) base = 127 + round(0.5 * math.log2(0.01057 / k)) - exps = torch.randint(base, base + 6, (rows, k // SV), dtype=torch.uint8, device="cuda", generator=gen) + exps = torch.randint( + base, base + 6, (rows, k // SV), dtype=torch.uint8, device="cuda", generator=gen + ) return (codes[:, 0::2] | (codes[:, 1::2] << 4)).contiguous(), exps @@ -121,7 +126,9 @@ def _deq_w(packed, sf): def _deq_x(x_fp8, x_sf): rows, k = x_fp8.shape - return x_fp8.float() * torch.exp2(x_sf.reshape(rows, k // SV).float() - 127.0).repeat_interleave(SV, dim=1) + return x_fp8.float() * torch.exp2( + x_sf.reshape(rows, k // SV).float() - 127.0 + ).repeat_interleave(SV, dim=1) def _requant(act): @@ -152,7 +159,9 @@ def _reference(raw, x_deq, ids, weights): gate = (xe @ _deq_w(raw["gate"][e], raw["gate_s"][e]).double().t()).float() act = GATE_CAP * torch.tanh(gate / GATE_CAP) * torch.sigmoid(gate) act = act * (LINEAR_CAP * torch.tanh(up / LINEAR_CAP)) - y = (_requant(act).double() @ _deq_w(raw["down"][e], raw["down_s"][e]).double().t()).float() + y = ( + _requant(act).double() @ _deq_w(raw["down"][e], raw["down_s"][e]).double().t() + ).float() out.index_add_(0, tok, y * weights[tok, slot].float().unsqueeze(1)) return out finally: @@ -297,8 +306,14 @@ def test_k3_fused_moe_mixed_m_sequence(): _ops() proc, _, bias = _experts() logits8, x8 = _tokens("random") - alone = {m: _fused(proc, bias, x8[:m].contiguous(), logits8[:m].contiguous()) for m in (1, 2, 5, 7, 8)} - seq = [(m, _fused(proc, bias, x8[:m].contiguous(), logits8[:m].contiguous())) for m in (8, 1, 5, 2, 8, 7)] + alone = { + m: _fused(proc, bias, x8[:m].contiguous(), logits8[:m].contiguous()) + for m in (1, 2, 5, 7, 8) + } + seq = [ + (m, _fused(proc, bias, x8[:m].contiguous(), logits8[:m].contiguous())) + for m in (8, 1, 5, 2, 8, 7) + ] assert all(torch.equal(_bits(y), _bits(alone[m])) for m, y in seq) assert _scratch_rearmed() diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mnnvl_comm.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mnnvl_comm.py index aaa750f960b5..e4e7c2746ddd 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mnnvl_comm.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mnnvl_comm.py @@ -238,12 +238,16 @@ def check_allgather(ctx): exact = _same(got_b, want_b) and _same(got_f, want_f) interleaved = True partial = torch.randn(m, H, device="cuda").bfloat16() - for i in range(6): # two turns of the three-buffer Lamport rotation, mixed with the other collectives + for i in range( + 6 + ): # two turns of the three-buffer Lamport rotation, mixed with the other collectives _allreduce(ctx, partial) if i % 2: block = torch.randn(2, m, H, device="cuda").bfloat16() ones = torch.ones(H, device="cuda").bfloat16() - ctx.mnnvl.allreduce_attn_res_rmsnorm(partial, None, block, ones, ones, ones, EPS, EPS) + ctx.mnnvl.allreduce_attn_res_rmsnorm( + partial, None, block, ones, ones, ones, EPS, EPS + ) b, f = ctx.mnnvl.allgather_split(mine, bf16_cols) interleaved &= _same(b, want_b) and _same(f, want_f) bad = mine.clone() @@ -259,9 +263,9 @@ def check_allgather(ctx): rows_as_m64=_same(got_b, b64[:m]) and _same(got_f, f64[:m]), control=not (_same(bad_b, want_b) and _same(bad_f, want_f)), ) - row["ok"] = _all_ranks(ctx, row["exact"] and interleaved and row["rows_as_m64"]) and _all_ranks( - ctx, row["control"] - ) + row["ok"] = _all_ranks( + ctx, row["exact"] and interleaved and row["rows_as_m64"] + ) and _all_ranks(ctx, row["control"]) results.append(row) return results diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py index da975bfcbc96..dec3f625264c 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py @@ -487,7 +487,9 @@ def set_epoch(epoch): t0 = time.perf_counter() y, _ = fused(x) torch.cuda.synchronize() - ms = (time.perf_counter() - t0) * 1e3 # ~20 ms: the hold ran out, no epoch moved during it + ms = ( + time.perf_counter() - t0 + ) * 1e3 # ~20 ms: the hold ran out, no epoch moved during it after, words = _quiet_check(ctx, lambda: (int(flags[2].item()), ready[:16].tolist())) row.update( epoch_advanced=after == want, off_epoch=[(i, w) for i, w in enumerate(words) if w != want], @@ -495,7 +497,12 @@ def set_epoch(epoch): y_zero=bool((y == 0).all()) if ctx.rank == idle else True, call_ms=[round(v, 2) for v in ctx.comm.allgather(ms)], ) # fmt: skip - good = row["epoch_advanced"] and not row["off_epoch"] and row["buffers_empty"] and row["y_zero"] + good = ( + row["epoch_advanced"] + and not row["off_epoch"] + and row["buffers_empty"] + and row["y_zero"] + ) row["rank"], row["good"] = ctx.rank, bool(good) row["ok"] = _all_ranks(ctx, good) results.append(row) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_wide.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_wide.py index 56f753e5fb2b..da28514533c8 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_wide.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_wide.py @@ -445,7 +445,10 @@ def run(layer, m): graph = torch.cuda.CUDAGraph() with torch.cuda.stream(stream): with torch.cuda.graph(graph, stream=stream): - outs = [_wide_call(layer, bias, static_x[:m], static_logits[:m], routed) for layer, m in calls] + outs = [ + _wide_call(layer, bias, static_x[:m], static_logits[:m], routed) + for layer, m in calls + ] static_x.copy_(x64) static_logits.copy_(logits64) graph.replay() diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_route_quant.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_route_quant.py index 72f646f0ed89..463322629760 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_route_quant.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_route_quant.py @@ -46,7 +46,11 @@ def _ops(): def _bits(t: torch.Tensor) -> torch.Tensor: - view = {torch.bfloat16: torch.int16, torch.float32: torch.int32, torch.float8_e4m3fn: torch.uint8} + view = { + torch.bfloat16: torch.int16, + torch.float32: torch.int32, + torch.float8_e4m3fn: torch.uint8, + } return t.contiguous().view(view.get(t.dtype, t.dtype)) @@ -66,7 +70,12 @@ def _unfused(scores, bias, hidden): weights, ids = ops.noaux_tc_op(scores, bias, 1, 1, K, SCALE) quantized, scales = ops.mxfp8_quantize(hidden, False, alignment=256) m = scores.shape[0] - return ids.int(), weights.to(torch.bfloat16), quantized.view(torch.float8_e4m3fn), scales.view(m, -1) + return ( + ids.int(), + weights.to(torch.bfloat16), + quantized.view(torch.float8_e4m3fn), + scales.view(m, -1), + ) @functools.lru_cache(maxsize=None) @@ -114,7 +123,12 @@ def _edge_cases(): b_tied = b.clone() b_tied[100:140] = 0.25 # 40 equal keys compete for the top 16: ties go to the lower id yield "40_tied_keys", s, b_tied, h - yield "all_equal_logits", torch.full((m, E), 0.3, device="cuda"), torch.zeros(E, device="cuda"), h + yield ( + "all_equal_logits", + torch.full((m, E), 0.3, device="cuda"), + torch.zeros(E, device="cuda"), + h, + ) yield "huge_logits", torch.randn(m, E, generator=gen, device="cuda") * 40.0, b, h s = torch.randn(m, E, generator=gen, device="cuda") s[:, 3::32] += 20.0 # every winner in one selection lane (the exact fallback) @@ -129,7 +143,12 @@ def _edge_cases(): h2[3, ::7] = torch.tensor(1e-39).bfloat16() h2[4, 5] = torch.tensor(-3e38).bfloat16() yield "zero_large_denormal_rows", torch.randn(m, E, generator=gen, device="cuda"), b, h2 - yield "bias_parameter", torch.randn(m, E, generator=gen, device="cuda"), torch.nn.Parameter(b.clone()), h + yield ( + "bias_parameter", + torch.randn(m, E, generator=gen, device="cuda"), + torch.nn.Parameter(b.clone()), + h, + ) EDGE_CASES = [ @@ -148,7 +167,9 @@ def test_k3_route_quant_edge_cases(case): """The PyTorch sort is reported, not asserted: its sigmoid may tie or split keys the kernels' sigmoid does not.""" name, s, b, h = next(c for c in _edge_cases() if c[0] == case) for m in (1, 3, 8): - _, vs_cpp, vs_unfused, _, rerun, early = _check(name, s[:m].contiguous(), b, h[:m].contiguous()) + _, vs_cpp, vs_unfused, _, rerun, early = _check( + name, s[:m].contiguous(), b, h[:m].contiguous() + ) assert vs_cpp and vs_unfused and rerun and early diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py index 983d688eb35a..50da0aeee64d 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py @@ -144,7 +144,9 @@ def _plain(ctx, x, w, residual, norm_w, swiglu=False): def _attn_res_ar(ctx, partial, prefix, block, res_w, rms_w, out_w): """The unfused post-projection step: the MNNVL one-shot all-reduce with the attention-residual epilogue.""" - return ctx.mnnvl.allreduce_attn_res_rmsnorm(partial, prefix, block, res_w, rms_w, out_w, EPS, EPS) + return ctx.mnnvl.allreduce_attn_res_rmsnorm( + partial, prefix, block, res_w, rms_w, out_w, EPS, EPS + ) def _tail_partial(latent, act, w, lo): @@ -163,7 +165,9 @@ def _residual_rms_ar(ctx, partial, residual, norm_w): params = AllReduceParams(fusion_op=AllReduceFusionOp.RESIDUAL_RMS_NORM, residual=residual, norm_weight=norm_w, eps=EPS) # fmt: skip - out = ctx.mnnvl(partial, params, one_shot_max_bytes=partial.numel() * ctx.world * partial.element_size()) + out = ctx.mnnvl( + partial, params, one_shot_max_bytes=partial.numel() * ctx.world * partial.element_size() + ) return out[0], out[1] @@ -180,7 +184,15 @@ def _oproj_inputs(ctx, snapshots, seed): w = _rand((H, K_O), seed + 1000 * r + 2, 0.03) prefix8 = _rand((8, H), seed + 3) block8 = _rand((snapshots, 8, H), seed + 4) - return core8, w, prefix8, block8, _rand((H,), seed + 5, 0.05), _norm_w(seed + 6), _norm_w(seed + 7) + return ( + core8, + w, + prefix8, + block8, + _rand((H,), seed + 5, 0.05), + _norm_w(seed + 6), + _norm_w(seed + 7), + ) def _tail_inputs(ctx, snapshots, seed): @@ -191,8 +203,17 @@ def _tail_inputs(ctx, snapshots, seed): w[:, WIDTH:PAD] = 0 prefix8 = _rand((8, H), seed + 14) block8 = _rand((snapshots, 8, H), seed + 15) - return latent8, act8, w, r * WIDTH, prefix8, block8, _rand((H,), seed + 16, 0.05), _norm_w(seed + 17), _norm_w( - seed + 18) + return ( + latent8, + act8, + w, + r * WIDTH, + prefix8, + block8, + _rand((H,), seed + 16, 0.05), + _norm_w(seed + 17), + _norm_w(seed + 18), + ) def _first(m, prefix8, block8, with_prefix): @@ -211,7 +232,9 @@ def check_oproj(ctx): results = [] for snapshots in SNAPSHOTS: for with_prefix in (True, False): - core8, w, prefix8, block8, res_w, rms_w, out_w = _oproj_inputs(ctx, snapshots, 100 * snapshots) + core8, w, prefix8, block8, res_w, rms_w, out_w = _oproj_inputs( + ctx, snapshots, 100 * snapshots + ) pre8, _ = _first(8, prefix8, block8, with_prefix) n8, u8 = _oproj(ctx, core8, w, pre8, block8, res_w, rms_w, out_w) for m in M_ALL: @@ -221,14 +244,18 @@ def check_oproj(ctx): want_n, want_u = _attn_res_ar(ctx, _row8(lambda x: F.linear(x, w), core, m), pre, block, res_w, rms_w, out_w) # fmt: skip again = [_oproj(ctx, core, w, pre, block, res_w, rms_w, out_w) for _ in range(2)] - bad_n, bad_u = _oproj(ctx, _perturbed(ctx, core), w, pre, block, res_w, rms_w, out_w) + bad_n, bad_u = _oproj( + ctx, _perturbed(ctx, core), w, pre, block, res_w, rms_w, out_w + ) row = dict( op="k3_sandwich_oproj", case=f"S{snapshots}_{'prefix' if with_prefix else 'noprefix'}", M=m, eq_unfused=_same(n, want_n) and _same(u, want_u), rel_updated=_rel(u, want_u), rel_normed=_rel(n, want_n), det=all(_same(a, n) and _same(b, u) for a, b in again), rows_as_m8=_same(n, n8[:m]) and _same(u, u8[:m]), control=not _same(bad_u, u), ) # fmt: skip - row["ok"] = _all_ranks(ctx, row["eq_unfused"] and row["det"] and row["rows_as_m8"] and row["control"]) + row["ok"] = _all_ranks( + ctx, row["eq_unfused"] and row["det"] and row["rows_as_m8"] and row["control"] + ) results.append(row) return results @@ -246,7 +273,9 @@ def check_tail(ctx): results = [] for snapshots in SNAPSHOTS: for with_prefix in (True, False): - latent8, act8, w, lo, prefix8, block8, res_w, rms_w, out_w = _tail_inputs(ctx, snapshots, 50 + snapshots) + latent8, act8, w, lo, prefix8, block8, res_w, rms_w, out_w = _tail_inputs( + ctx, snapshots, 50 + snapshots + ) pre8, _ = _first(8, prefix8, block8, with_prefix) n8, u8 = _tail(ctx, latent8, act8, w, lo, pre8, block8, res_w, rms_w, out_w) for m in M_ALL: @@ -255,8 +284,13 @@ def check_tail(ctx): n, u = _tail(ctx, latent, act, w, lo, pre, block, res_w, rms_w, out_w) part = _tail_partial(latent, act, w, lo) want_n, want_u = _attn_res_ar(ctx, part, pre, block, res_w, rms_w, out_w) - again = [_tail(ctx, latent, act, w, lo, pre, block, res_w, rms_w, out_w) for _ in range(2)] - bad_n, bad_u = _tail(ctx, latent, _perturbed(ctx, act), w, lo, pre, block, res_w, rms_w, out_w) + again = [ + _tail(ctx, latent, act, w, lo, pre, block, res_w, rms_w, out_w) + for _ in range(2) + ] + bad_n, bad_u = _tail( + ctx, latent, _perturbed(ctx, act), w, lo, pre, block, res_w, rms_w, out_w + ) row = dict( op="k3_sandwich_tail", case=f"S{snapshots}_{'prefix' if with_prefix else 'noprefix'}", M=m, rel_updated=_rel(u, want_u), rel_normed=_rel(n, want_n), @@ -311,7 +345,9 @@ def _check_plain(ctx, swiglu): results = [] name = "k3_sandwich_plain_swiglu" if swiglu else "k3_sandwich_plain" for seed in (0, 1): - x8, w, res8, norm_w = _plain_inputs(ctx, 2 * DOWN_K if swiglu else PLAIN_K, 300 + 10 * seed + int(swiglu)) + x8, w, res8, norm_w = _plain_inputs( + ctx, 2 * DOWN_K if swiglu else PLAIN_K, 300 + 10 * seed + int(swiglu) + ) def gemv(x): if swiglu: @@ -324,14 +360,18 @@ def gemv(x): n, u = _plain(ctx, x, w, res, norm_w, swiglu) want_n, want_u = _residual_rms_ar(ctx, gemv(x), res, norm_w) again = [_plain(ctx, x, w, res, norm_w, swiglu) for _ in range(2)] - bad_n, bad_u = _plain(ctx, _perturbed(ctx, x, DOWN_K if swiglu else 0), w, res, norm_w, swiglu) + bad_n, bad_u = _plain( + ctx, _perturbed(ctx, x, DOWN_K if swiglu else 0), w, res, norm_w, swiglu + ) row = dict( op=name, case=f"seed{seed}", M=m, eq_unfused=_same(n, want_n) and _same(u, want_u), rel_updated=_rel(u, want_u), rel_normed=_rel(n, want_n), det=all(_same(a, n) and _same(b, u) for a, b in again), rows_as_m8=_same(n, n8[:m]) and _same(u, u8[:m]), control=not _same(bad_u, u), ) # fmt: skip - row["ok"] = _all_ranks(ctx, row["eq_unfused"] and row["det"] and row["rows_as_m8"] and row["control"]) + row["ok"] = _all_ranks( + ctx, row["eq_unfused"] and row["det"] and row["rows_as_m8"] and row["control"] + ) results.append(row) return results @@ -353,7 +393,15 @@ def check_replay(ctx): core8, w_o, prefix8, block8, res_w, rms_w, out_w = _oproj_inputs(ctx, 3, 700) latent8, act8, w_t, lo, t_prefix8, t_block8, t_res, t_rms, t_out = _tail_inputs(ctx, 3, 710) x8, w_p, res8, norm_w = _plain_inputs(ctx, PLAIN_K, 720) - a_in = [core8[:m].clone(), w_o, prefix8[:m].clone(), block8[:, :m].clone(), res_w, rms_w, out_w] + a_in = [ + core8[:m].clone(), + w_o, + prefix8[:m].clone(), + block8[:, :m].clone(), + res_w, + rms_w, + out_w, + ] b_in = [latent8[:m].clone(), act8[:m].clone(), w_t, lo, t_prefix8[:m].clone(), t_block8[:, :m].clone(), t_res, t_rms, t_out] # fmt: skip c_in = [x8[:m].clone(), w_p, res8[:m].clone(), norm_w] @@ -387,7 +435,15 @@ def check_replay(ctx): want = [_oproj(ctx, *a_in), _tail(ctx, *b_in), _plain(ctx, *c_in)] torch.cuda.synchronize() same = all(_same(g, x) for go, wo in zip(got, want) for g, x in zip(go, wo)) - results.append(dict(op="k3_sandwich_replay", case=f"replay{it}", M=m, eq_eager=same, ok=_all_ranks(ctx, same))) + results.append( + dict( + op="k3_sandwich_replay", + case=f"replay{it}", + M=m, + eq_eager=same, + ok=_all_ranks(ctx, same), + ) + ) del graphs return results @@ -425,7 +481,14 @@ def run(): wrapped = run() after = int(flags[0].item()) same = all(_same(a, b) for f_out, w_out in zip(fresh, wrapped) for a, b in zip(f_out, w_out)) - row = dict(op="k3_sandwich_wrap", case="int32_wrap", M=8, eq_fresh=same, crossed=after < 0, calls=len(calls)) + row = dict( + op="k3_sandwich_wrap", + case="int32_wrap", + M=8, + eq_fresh=same, + crossed=after < 0, + calls=len(calls), + ) row["ok"] = _all_ranks(ctx, same and row["crossed"]) return [row] @@ -443,7 +506,9 @@ def check_fold_wrap(ctx): lanes = ex.mc.view(2, 8, ctx.world, LATENT // 2) latent8, act8, w, lo, prefix8, block8, res_w, rms_w, out_w = _tail_inputs(ctx, 3, 830) part8 = _rand((8, LATENT), 840 + 1000 * ctx.rank, 0.3) - part8[part8 == 0] = 0.0 # pushes never send -0.0: with +0.0 beside it, that word would read as empty + part8[part8 == 0] = ( + 0.0 # pushes never send -0.0: with +0.0 beside it, that word would read as empty + ) slab = slice(kernel.LAT_SCALES, kernel.LAT_SCALES + kernel.LAT_SCALE_BUFS * 8) def run(): @@ -460,7 +525,9 @@ def run(): lat_uc=ex.uc, lat_flags=flags) # fmt: skip torch.cuda.synchronize() outs.append([t.clone() for t in out]) - others.append(torch.equal(before, torch.cat([flags[1 : slab.start], flags[slab.stop :]]))) + others.append( + torch.equal(before, torch.cat([flags[1 : slab.start], flags[slab.stop :]])) + ) return outs, others fresh, fresh_kept = run() diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py index 4e5af9f0f4ae..3abb5cc1a9b0 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py @@ -544,7 +544,9 @@ def check_graph_capture_and_replay() -> None: buffers are empty, WS_B's epoch advanced once per head_flags call, replayed or eager, with every ready word at it, and both states armed. """ - static_bufs = [torch.zeros(8, HIDDEN, dtype=torch.bfloat16, device="cuda") for _ in range(LAYERS)] + static_bufs = [ + torch.zeros(8, HIDDEN, dtype=torch.bfloat16, device="cuda") for _ in range(LAYERS) + ] shells = [ Call(0, 0, layer=layer, kind=KINDS[layer], x=static_bufs[layer]) for layer in range(LAYERS) ] diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_common.py b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_common.py index 5f3dca684c54..d93f600557c7 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_common.py +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_common.py @@ -46,7 +46,9 @@ LAT_EPS = 1e-6 EPS = 1e-6 TOL = 2e-2 # normed and the tap: max |err| / max |ref| -TOL_TAIL = 8e-3 # the tail's updated: max |err| / max |ref|, about one bf16 ulp of its largest elements +TOL_TAIL = ( + 8e-3 # the tail's updated: max |err| / max |ref|, about one bf16 ulp of its largest elements +) GATES = (0.0, 32.0, 64.0) # swiglu gate values on which silu is exact in fp32 (see PlainCall) WEIGHT_SETS = 3 TOKENS = (1, 2, 3, 4, 5, 6, 7, 8) @@ -156,7 +158,9 @@ def attn_res_mixture(updated, block, res_w, rms_w) -> torch.Tensor: return (probs[..., None] * v).sum(dim=0).bfloat16() -def _within(got: torch.Tensor, want: torch.Tensor, tol: float, key: str, where: str, what: str) -> None: +def _within( + got: torch.Tensor, want: torch.Tensor, tol: float, key: str, where: str, what: str +) -> None: err = ls.rel_err(got, want) STATS[key] = max(STATS[key], err) assert err <= tol, f"{where}: {what} rel err {err:.3e} > {tol}" @@ -185,7 +189,9 @@ def refill(self, fresh: "Call", carry: bool) -> None: def verify(self, got, where: str) -> None: normed, updated = got want_normed, want_updated = self.ref() - assert torch.equal(updated, want_updated), f"{where}: updated differs from the exact reference" + assert torch.equal(updated, want_updated), ( + f"{where}: updated differs from the exact reference" + ) _within(normed, want_normed, TOL, "normed", where, "normed") assert R.same_on_ranks(normed, updated), f"{where}: ranks disagree" @@ -211,7 +217,9 @@ def __init__(self, seed, tokens, snapshots, prefix=True, weights=0): self.block, self.res_w, self.rms_w, self.out_w = _attn_res_inputs(g, snapshots, tokens) def ref(self): - updated = reduce_ref([bf16_partial(c, w) for c, w in zip(self.cores, self.weights)], self.carry) + updated = reduce_ref( + [bf16_partial(c, w) for c, w in zip(self.cores, self.weights)], self.carry + ) normed = ls.residual_update_ref( updated, self.block, self.res_w, self.rms_w, RMS_EPS, self.out_w, OUT_EPS ) @@ -237,7 +245,17 @@ class TailCall(Call): [2 H, 3 H) of a NaN-filled [M, 5 H] capture buffer) or "updated" (``updated`` into its columns [4 H, 5 H)); ``updated_out``: store ``updated`` into row 1 of a NaN-filled [3, M, H] bank.""" - def __init__(self, seed, tokens, snapshots, prefix=True, weights=0, shift=None, tap=None, updated_out=False): + def __init__( + self, + seed, + tokens, + snapshots, + prefix=True, + weights=0, + shift=None, + tap=None, + updated_out=False, + ): g = _gen(seed) self.tokens = tokens self.latent = ls.exact_bf16(g, (tokens, LATENT), -16, 17, 1 / 8) @@ -252,11 +270,15 @@ def __init__(self, seed, tokens, snapshots, prefix=True, weights=0, shift=None, def _set_options(self, tap, updated_out) -> None: self.tap_kind, self.cap, self.tap = tap, None, None if tap is not None: - self.cap = torch.full((self.tokens, 5 * H), float("nan"), dtype=torch.bfloat16, device="cuda") + self.cap = torch.full( + (self.tokens, 5 * H), float("nan"), dtype=torch.bfloat16, device="cuda" + ) self.tap = self.cap[:, self._tap_col() * H : (self._tap_col() + 1) * H] self.bank, self.updated_out = None, None if updated_out: - self.bank = torch.full((3, self.tokens, H), float("nan"), dtype=torch.bfloat16, device="cuda") + self.bank = torch.full( + (3, self.tokens, H), float("nan"), dtype=torch.bfloat16, device="cuda" + ) self.updated_out = self.bank[1] def _tap_col(self) -> int: @@ -309,17 +331,25 @@ def verify(self, got, where): _within(normed, want_normed, TOL, "normed", where, "normed") outputs = [normed, updated] if self.updated_out is not None: - assert updated.data_ptr() == self.updated_out.data_ptr(), f"{where}: updated is not updated_out" - assert bool(torch.isnan(self.bank[0::2].float()).all()), f"{where}: another bank row was written" + assert updated.data_ptr() == self.updated_out.data_ptr(), ( + f"{where}: updated is not updated_out" + ) + assert bool(torch.isnan(self.bank[0::2].float()).all()), ( + f"{where}: another bank row was written" + ) if self.tap is not None: col = self._tap_col() rest = torch.cat([self.cap[:, : col * H], self.cap[:, (col + 1) * H :]], dim=1) - assert bool(torch.isnan(rest.float()).all()), f"{where}: the capture buffer was written outside the tap" + assert bool(torch.isnan(rest.float()).all()), ( + f"{where}: the capture buffer was written outside the tap" + ) if self.tap_kind == "mix": mix = attn_res_mixture(want_updated, self.block, self.res_w, self.rms_w) _within(self.tap, mix, TOL, "tap", where, "tapped mixture") else: - assert torch.equal(bits(self.tap), bits(updated)), f"{where}: the tap is not updated" + assert torch.equal(bits(self.tap), bits(updated)), ( + f"{where}: the tap is not updated" + ) outputs.append(self.tap) assert R.same_on_ranks(*outputs), f"{where}: ranks disagree" @@ -347,7 +377,13 @@ def __init__(self, seed, tokens, swiglu=False, residual=True, weights=0): self.swiglu = swiglu if swiglu: self.xs = [ - torch.cat([_gates(g, (tokens, DOWN_K)), ls.exact_bf16(g, (tokens, DOWN_K), -2, 3, 1 / 64)], dim=1) + torch.cat( + [ + _gates(g, (tokens, DOWN_K)), + ls.exact_bf16(g, (tokens, DOWN_K), -2, 3, 1 / 64), + ], + dim=1, + ) for _ in range(R.world) ] self.weights = WEIGHTS["down"][weights] @@ -367,12 +403,16 @@ def ref(self): partials = [bf16_partial(self.operand(x), w) for x, w in zip(self.xs, self.weights)] updated = reduce_ref(partials, self.carry) u = updated.float() - normed = (u * (u.square().mean(dim=-1, keepdim=True) + EPS).rsqrt() * self.norm_w.float()).bfloat16() + normed = ( + u * (u.square().mean(dim=-1, keepdim=True) + EPS).rsqrt() * self.norm_w.float() + ).bfloat16() return normed, updated def run(self, ws): r = R.rank - return OPS["plain"](self.xs[r], self.weights[r], self.carry, self.norm_w, EPS, ws, swiglu=self.swiglu) + return OPS["plain"]( + self.xs[r], self.weights[r], self.carry, self.norm_w, EPS, ws, swiglu=self.swiglu + ) def refill(self, fresh, carry): for mine, new in zip(self.xs, fresh.xs): @@ -385,7 +425,9 @@ def _label(i: int, call: Call) -> str: return f"call {i} ({type(call).__name__} M {call.tokens})" -def run_sequence(seq, ws, verify: bool = True, late: Optional[random.Random] = None, where: str = "sequence"): +def run_sequence( + seq, ws, verify: bool = True, late: Optional[random.Random] = None, where: str = "sequence" +): """Run ``seq`` -- (chain, call) pairs -- in order on ``ws``. A call's carry is the ``updated`` of the previous call of its chain (a chain's first call keeps its own), as the model chains its residual streams. With ``late`` (a random.Random seeded alike on every rank) every call starts after a barrier, a random rank 5 ms late; with @@ -448,7 +490,9 @@ def attempt() -> bool: every = attempt() R.barrier() alone = attempt() if R.rank == R.world - 1 else True - assert R.all_true(every and alone), f"create under capture: every rank raised {every}, one rank alone {alone}" + assert R.all_true(every and alone), ( + f"create under capture: every rank raised {every}, one rank alone {alone}" + ) next_call.verify(next_call.run(ws), "after the refused creates") diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_oproj_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_oproj_op_matrix.py index 333ce767fd3b..cd9ae009f226 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_oproj_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_oproj_op_matrix.py @@ -84,8 +84,16 @@ def check_unsupported_shape_raises_on_every_rank() -> None: def _step(seed, tokens): """One decode step: ``LAYERS`` chained calls, each layer's prefix the previous layer's ``updated``.""" return [ - ("target", cm.OprojCall(seed + 1 + layer, tokens, layer % 9, prefix=True if layer == 0 else None, - weights=layer % cm.WEIGHT_SETS)) # fmt: skip + ( + "target", + cm.OprojCall( + seed + 1 + layer, + tokens, + layer % 9, + prefix=True if layer == 0 else None, + weights=layer % cm.WEIGHT_SETS, + ), + ) # fmt: skip for layer in range(LAYERS) ] @@ -111,7 +119,9 @@ def check_graph_capture_and_replay() -> None: cm.capture_and_replay( WS_B, lambda seed: _step(seed, 8), - lambda rep: [cm.OprojCall(7000 + rep, (3, 1, 6, 5)[rep % 4], rep % 9, weights=rep % cm.WEIGHT_SETS)], + lambda rep: [ + cm.OprojCall(7000 + rep, (3, 1, 6, 5)[rep % 4], rep % 9, weights=rep % cm.WEIGHT_SETS) + ], "captured step", ) @@ -127,9 +137,18 @@ def _shared_step(seed, t_tokens, d_tokens): s = seed + 10 * layer seq.append(("target", cm.OprojCall(s + 1, t_tokens, (4 * layer) % 9, prefix=True if layer == 0 else None, weights=layer))) # fmt: skip - seq.append(("target", cm.TailCall(s + 2, t_tokens, (4 * layer + 1) % 9, prefix=None, updated_out=layer == 1))) + seq.append( + ( + "target", + cm.TailCall( + s + 2, t_tokens, (4 * layer + 1) % 9, prefix=None, updated_out=layer == 1 + ), + ) + ) if layer != 1: - seq.append(("drafter", cm.PlainCall(s + 3, d_tokens, residual=True if layer == 0 else None))) + seq.append( + ("drafter", cm.PlainCall(s + 3, d_tokens, residual=True if layer == 0 else None)) + ) seq.append(("drafter", cm.PlainCall(s + 4, d_tokens, swiglu=True, residual=None))) return seq @@ -138,12 +157,22 @@ def check_shared_sequence() -> None: """The three sandwich ops on one workspace, as the model runs them (``_shared_step``, 10 calls): steps at (target M, drafter M) = (8, 3), (2, 8), (7, 7), a random rank late at every call, every call against its reference. The three wrappers export one workspace type.""" - from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm import k3_sandwich_plain, k3_sandwich_tail + from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm import ( + k3_sandwich_plain, + k3_sandwich_tail, + ) - assert k3_sandwich_tail.K3SandwichWorkspace is WORKSPACE is k3_sandwich_plain.K3SandwichWorkspace + assert ( + k3_sandwich_tail.K3SandwichWorkspace is WORKSPACE is k3_sandwich_plain.K3SandwichWorkspace + ) late = random.Random(11) for i, (t, d) in enumerate(SHARED_STEPS): - cm.run_sequence(_shared_step(9000 + 100 * i, t, d), WS_A, late=late, where=f"shared step {i} M {t} / {d}") + cm.run_sequence( + _shared_step(9000 + 100 * i, t, d), + WS_A, + late=late, + where=f"shared step {i} M {t} / {d}", + ) def check_shared_sequence_captured() -> None: @@ -169,7 +198,10 @@ def check_wrong_call_order_is_detected() -> None: or hangs, but every rank's two results are wrong (more than half the ``updated`` elements differ); then a plain call is correct again.""" cm.swapped_pair_is_wrong( - WS_A, cm.OprojCall(8000, 8, 3), cm.OprojCall(8001, 8, 3, weights=1), cm.OprojCall(8002, 8, 3) + WS_A, + cm.OprojCall(8000, 8, 3), + cm.OprojCall(8001, 8, 3, weights=1), + cm.OprojCall(8002, 8, 3), ) diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_plain_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_plain_op_matrix.py index c9adbff4760f..c0def69078f6 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_plain_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_plain_op_matrix.py @@ -77,7 +77,10 @@ def check_unsupported_calls_raise_on_every_rank() -> None: ValueError on every rank before any launch; the next call is correct.""" big = cm.PlainCall(2000, 9) narrow = cm.PlainCall(2001, 4) - narrow.swiglu, narrow.xs = True, [torch.cat([x, x], dim=1) for x in narrow.xs] # [4, 768] on a K 384 slice + narrow.swiglu, narrow.xs = ( + True, + [torch.cat([x, x], dim=1) for x in narrow.xs], + ) # [4, 768] on a K 384 slice wide = cm.PlainCall(2002, 4) wide.xs = [torch.zeros(4, 512, dtype=torch.bfloat16, device="cuda") for _ in wide.xs] cm.unsupported_raises( @@ -98,7 +101,12 @@ def _step(seed, tokens): for layer in range(DRAFTER_LAYERS): s = seed + 2 * layer w = layer % cm.WEIGHT_SETS - seq.append(("drafter", cm.PlainCall(s + 1, tokens, residual=True if layer == 0 else None, weights=w))) + seq.append( + ( + "drafter", + cm.PlainCall(s + 1, tokens, residual=True if layer == 0 else None, weights=w), + ) + ) seq.append(("drafter", cm.PlainCall(s + 2, tokens, swiglu=True, residual=None, weights=w))) return seq @@ -115,7 +123,9 @@ def check_two_workspaces_interleaved() -> None: cm.interleaved( WS_A, WS_B, - lambda i: cm.PlainCall(4000 + i, (3, 8, 1, 8, 5)[i % 5], swiglu=i % 3 == 1, weights=i % cm.WEIGHT_SETS), + lambda i: cm.PlainCall( + 4000 + i, (3, 8, 1, 8, 5)[i % 5], swiglu=i % 3 == 1, weights=i % cm.WEIGHT_SETS + ), ) @@ -125,8 +135,11 @@ def check_graph_capture_and_replay() -> None: cm.capture_and_replay( WS_B, lambda seed: _step(seed, 8), - lambda rep: [cm.PlainCall(7000 + rep, (3, 1, 6, 5)[rep % 4], swiglu=rep % 2 == 1, - weights=rep % cm.WEIGHT_SETS)], # fmt: skip + lambda rep: [ + cm.PlainCall( + 7000 + rep, (3, 1, 6, 5)[rep % 4], swiglu=rep % 2 == 1, weights=rep % cm.WEIGHT_SETS + ) + ], # fmt: skip "captured step", ) @@ -135,7 +148,9 @@ def check_wrong_call_order_is_detected() -> None: """Negative control: rank 0 swaps two same-shaped calls on one workspace. Every call returns and nothing raises or hangs, but every rank's two results are wrong (more than half the ``updated`` elements differ); then a plain call is correct again.""" - cm.swapped_pair_is_wrong(WS_A, cm.PlainCall(8000, 8), cm.PlainCall(8001, 8, weights=1), cm.PlainCall(8002, 8)) + cm.swapped_pair_is_wrong( + WS_A, cm.PlainCall(8000, 8), cm.PlainCall(8001, 8, weights=1), cm.PlainCall(8002, 8) + ) CHECKS = [ diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_tail_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_tail_op_matrix.py index 718747f2823e..1dba6c7713e7 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_tail_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_tail_op_matrix.py @@ -70,7 +70,10 @@ def check_single_calls() -> None: for with_prefix in (True, False): call = cm.TailCall(1000 + 37 * t + s, t, s, prefix=with_prefix or None, weights=t % cm.WEIGHT_SETS, shift=i) # fmt: skip - call.verify(call.run(WS_A), f"M {t} snapshots {s} prefix {with_prefix} lo {call.los[R.rank]}") + call.verify( + call.run(WS_A), + f"M {t} snapshots {s} prefix {with_prefix} lo {call.los[R.rank]}", + ) i += 1 @@ -101,7 +104,9 @@ def check_unsupported_calls_raise_on_every_rank() -> None: big = cm.TailCall(2000, 9, 1) misaligned = cm.TailCall(2001, 4, 1) store = torch.zeros(4 * cm.LATENT + 8, dtype=torch.bfloat16, device="cuda") - shifted = store[1 : 1 + 4 * cm.LATENT].view(4, cm.LATENT) # contiguous, 2 bytes past an aligned address + shifted = store[1 : 1 + 4 * cm.LATENT].view( + 4, cm.LATENT + ) # contiguous, 2 bytes past an aligned address shifted.copy_(misaligned.latent) misaligned.latent = shifted strided = cm.TailCall(2002, 4, 1) @@ -123,8 +128,17 @@ def _step(seed, tokens, options=None): maps a layer to its tap / updated_out.""" options = options or {} return [ - ("target", cm.TailCall(seed + 1 + layer, tokens, layer % 9, prefix=True if layer == 0 else None, - weights=layer % cm.WEIGHT_SETS, **options.get(layer, {}))) # fmt: skip + ( + "target", + cm.TailCall( + seed + 1 + layer, + tokens, + layer % 9, + prefix=True if layer == 0 else None, + weights=layer % cm.WEIGHT_SETS, + **options.get(layer, {}), + ), + ) # fmt: skip for layer in range(LAYERS) ] @@ -151,7 +165,9 @@ def check_graph_capture_and_replay() -> None: cm.capture_and_replay( WS_B, lambda seed: _step(seed, 8, CAPTURED_OPTIONS), - lambda rep: [cm.TailCall(7000 + rep, (3, 1, 6, 5)[rep % 4], rep % 9, weights=rep % cm.WEIGHT_SETS)], + lambda rep: [ + cm.TailCall(7000 + rep, (3, 1, 6, 5)[rep % 4], rep % 9, weights=rep % cm.WEIGHT_SETS) + ], "captured step", ) From 1de6f8395ffa027d25e65f5160a8ceefd551e93a Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 20:19:51 -0700 Subject: [PATCH 066/161] [None][fix] MNNVL: buffer_flags declared mutable; the split all-gather takes its world size - trtllm::mnnvl_fusion_allreduce's schema declares buffer_flags mutable (Tensor(b!)): every call advances the Lamport rotation's flag words. - trtllm::mnnvl_allgather_split takes world_size, which must equal the workspace's number of ranks (checked). The output widths are world_size * bf16_columns and world_size * (columns - bf16_columns), so the op can now register a fake implementation, and it does. Signed-off-by: Vasanth Sabavat --- cpp/tensorrt_llm/thop/allreduceOp.cpp | 19 +++++++++++-------- .../catalog/comm/mnnvl_allgather_split.py | 1 + .../_torch/custom_ops/cpp_custom_ops.py | 9 +++++++++ tensorrt_llm/_torch/distributed/ops.py | 1 + 4 files changed, 22 insertions(+), 8 deletions(-) diff --git a/cpp/tensorrt_llm/thop/allreduceOp.cpp b/cpp/tensorrt_llm/thop/allreduceOp.cpp index 9ef1ed39b46e..92197cc83f0a 100644 --- a/cpp/tensorrt_llm/thop/allreduceOp.cpp +++ b/cpp/tensorrt_llm/thop/allreduceOp.cpp @@ -2300,8 +2300,8 @@ namespace // The all-gather's checks, outputs and params. tensorrt_llm::kernels::mnnvl::AllGatherSplitParams makeAllGatherSplitParams(torch::Tensor const& input, - int64_t bf16_columns, torch::Tensor& comm_buffer, torch::Tensor& buffer_flags, torch::Tensor& bf16Out, - torch::Tensor& fp32Out) + int64_t bf16_columns, int64_t world_size, torch::Tensor& comm_buffer, torch::Tensor& buffer_flags, + torch::Tensor& bf16Out, torch::Tensor& fp32Out) { namespace mnnvl = tensorrt_llm::kernels::mnnvl; auto* mcast_mem = tensorrt_llm::common::findMcastDevMemBuffer(comm_buffer.data_ptr()); @@ -2317,6 +2317,8 @@ tensorrt_llm::kernels::mnnvl::AllGatherSplitParams makeAllGatherSplitParams(torc TORCH_CHECK(bf16_columns >= 0 && fp32Columns >= 0 && bf16_columns % 8 == 0 && fp32Columns % 4 == 0, "[mnnvlAllGatherSplit] needs bf16_columns a multiple of 8 and the remaining columns a multiple of 4"); int64_t const nRanks = mcast_mem->getWorldSize(); + TORCH_CHECK(world_size == nRanks, "[mnnvlAllGatherSplit] world_size ", world_size, " is not the workspace's ", + nRanks, " ranks"); TORCH_CHECK(mnnvl::mnnvlAllGatherSplitFootprint(numTokens, bf16_columns, fp32Columns, nRanks) <= comm_buffer.size(-1) * comm_buffer.element_size(), "[mnnvlAllGatherSplit] the exchange does not fit in one Lamport buffer"); @@ -2342,12 +2344,13 @@ tensorrt_llm::kernels::mnnvl::AllGatherSplitParams makeAllGatherSplitParams(torc } // namespace -std::vector mnnvlAllGatherSplit( - torch::Tensor const& input, int64_t bf16_columns, torch::Tensor& comm_buffer, torch::Tensor& buffer_flags) +std::vector mnnvlAllGatherSplit(torch::Tensor const& input, int64_t bf16_columns, int64_t world_size, + torch::Tensor& comm_buffer, torch::Tensor& buffer_flags) { torch::Tensor bf16Out; torch::Tensor fp32Out; - auto const params = makeAllGatherSplitParams(input, bf16_columns, comm_buffer, buffer_flags, bf16Out, fp32Out); + auto const params + = makeAllGatherSplitParams(input, bf16_columns, world_size, comm_buffer, buffer_flags, bf16Out, fp32Out); tensorrt_llm::kernels::mnnvl::mnnvlAllGatherSplitOp(params); return {bf16Out, fp32Out}; } @@ -2450,7 +2453,7 @@ TORCH_LIBRARY_FRAGMENT(trtllm, m) { m.def( "mnnvl_fusion_allreduce(Tensor input, Tensor? gamma, Tensor? residual, " - "float? epsilon, Tensor(a!) comm_buffer, Tensor buffer_flags, bool rmsnorm_fusion, " + "float? epsilon, Tensor(a!) comm_buffer, Tensor(b!) buffer_flags, bool rmsnorm_fusion, " "Tensor? scale=None, int fusion_op=0, int one_shot_max_bytes=1048576) -> " "Tensor[]"); m.def( @@ -2458,8 +2461,8 @@ TORCH_LIBRARY_FRAGMENT(trtllm, m) "Tensor rms_weight, Tensor output_rms_weight, float rms_eps, float output_rms_eps, Tensor(a!) comm_buffer, " "Tensor(b!) buffer_flags) -> Tensor[]"); m.def( - "mnnvl_allgather_split(Tensor input, int bf16_columns, Tensor(a!) comm_buffer, Tensor(b!) buffer_flags) " - "-> Tensor[]"); + "mnnvl_allgather_split(Tensor input, int bf16_columns, int world_size, Tensor(a!) comm_buffer, " + "Tensor(b!) buffer_flags) -> Tensor[]"); m.def( "allreduce(" "Tensor input," diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.py index 8a9719f417ba..7a0b701ddbe5 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.py @@ -44,6 +44,7 @@ def mnnvl_allgather_split( bf16_out, fp32_out = torch.ops.trtllm.mnnvl_allgather_split( input, bf16_columns, + workspace.world_size, workspace.comm_buffer(torch.bfloat16), workspace.buffer_flags, ) diff --git a/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py b/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py index 020cf6658bcb..ba6a25655281 100644 --- a/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py +++ b/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py @@ -127,6 +127,15 @@ def _(input, prefix_sum, block_residual, res_weight, rms_weight, buffer_flags): return [torch.empty_like(input), torch.empty_like(input)] + @torch.library.register_fake("trtllm::mnnvl_allgather_split") + def _(input, bf16_columns, world_size, comm_buffer, buffer_flags): + num_tokens, columns = input.shape + return [ + input.new_empty((num_tokens, world_size * bf16_columns), + dtype=torch.bfloat16), + input.new_empty((num_tokens, world_size * (columns - bf16_columns))), + ] + # MNNVL Allreduce @torch.library.register_fake("trtllm::mnnvl_fusion_allreduce") def _(input, diff --git a/tensorrt_llm/_torch/distributed/ops.py b/tensorrt_llm/_torch/distributed/ops.py index 1f7a6a3a6065..aed6c515f142 100644 --- a/tensorrt_llm/_torch/distributed/ops.py +++ b/tensorrt_llm/_torch/distributed/ops.py @@ -1075,6 +1075,7 @@ def allgather_split(self, input: torch.Tensor, bf16_out, fp32_out = torch.ops.trtllm.mnnvl_allgather_split( input, bf16_columns, + self.mapping.tp_size, workspace["uc_buffer"].view(self.dtype).view(3, -1), workspace["buffer_flags"], ) From 823538a31f0575596bb4f5f3670ea7e0959cbf83 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 20:20:12 -0700 Subject: [PATCH 067/161] [None][fix] Kimi K3 collective state: the ranks agree before allocating K3SandwichWorkspace, K3SandwichLatentExchange, K3MoeHeadWorkspace and K3LatentExchange are created with the failure model of MnnvlWorkspace.create: - before allocating, the ranks agree that each of them can (not capturing a CUDA graph, the buffer within its device's free memory). If one cannot, every rank raises RuntimeError and none enters the collective allocation; - a failure that returns from the allocation is agreed the same way; - a rank that fails inside the allocation's handle exchange can still leave its peers waiting there, as the docstrings state. The head workspace and the latent exchange share one helper, k3_fused_moe.op.create_mcast_state. Signed-off-by: Vasanth Sabavat --- .../k3_fused_moe/latent_op.py | 42 ++------- .../cute_dsl_kernels/k3_fused_moe/op.py | 94 +++++++++++-------- .../_torch/cute_dsl_kernels/k3_sandwich/op.py | 30 ++++-- 3 files changed, 88 insertions(+), 78 deletions(-) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py index 7b1c025ef5da..7b4f23e6b689 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py @@ -77,50 +77,24 @@ class K3LatentExchange: @classmethod def create(cls, mapping, fabric_handle: Optional[bool] = None) -> "K3LatentExchange": """Allocate and arm an exchange for ``mapping``'s TP group. Collective: every rank of the group calls it at the - same point, eagerly (not under CUDA-graph capture); it returns on every rank or raises on every rank. + same point, eagerly (not under CUDA-graph capture); the failure model is ``op.create_mcast_state``'s. ``fabric_handle``: share the memory by fabric handle (required across nodes) rather than POSIX file descriptor; default ``mapping.is_multi_node()``.""" - if torch.cuda.is_current_stream_capturing(): - raise RuntimeError( - "K3LatentExchange.create is collective and allocates: call it outside CUDA-graph capture" - ) - from tensorrt_llm._torch.distributed.ops import ( - _get_mnnvl_workspace_comm, - _make_mnnvl_mcast_buffer, - _mnnvl_workspace_all_succeeded, - ) + from .op import create_mcast_state - words = _kernel().buffer_words(mapping.tp_size) - use_fabric_handle = ( - mapping.is_multi_node() if fabric_handle is None else bool(fabric_handle) - ) - comm = _get_mnnvl_workspace_comm(mapping) - error: Optional[Exception] = None - exchange = None - try: - handle = _make_mnnvl_mcast_buffer(comm, words * 4, mapping, use_fabric_handle) - uc = handle.get_uc_buffer(mapping.tp_rank, (words,), torch.int32, 0) - mc = handle.get_mc_buffer((words,), torch.int32, 0) - uc.fill_(EMPTY_WORD) - flags = torch.zeros(4, dtype=torch.int32, device=uc.device) - torch.cuda.synchronize() - exchange = cls( + def build(uc, mc, handle, comm): + return cls( uc=uc, mc=mc, - flags=flags, + flags=torch.zeros(4, dtype=torch.int32, device=uc.device), rank=mapping.tp_rank, world_size=mapping.tp_size, handle=handle, comm=comm, ) - except Exception as exc: # noqa: BLE001 -- reported to every rank below, then re-raised - error = exc - # Also the barrier that keeps any rank from pushing into a peer's buffer before the peer has emptied it. - if not _mnnvl_workspace_all_succeeded(comm, error is None): - raise RuntimeError( - "K3LatentExchange: allocation failed on at least one rank" - ) from error - return exchange + + words = _kernel().buffer_words(mapping.tp_size) + return create_mcast_state("K3LatentExchange", mapping, words, fabric_handle, build) def push_args(self): """``(ar_uc, ar_mc, ar_flags, ar_rank)`` of the push-only producers.""" diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py index 0e3b58cbe7cd..9cbfe2fe7262 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py @@ -117,6 +117,55 @@ def is_supported(w3_w1_weight: torch.Tensor, w3_w1_weight_scale: torch.Tensor, w return True, "" +def create_mcast_state(name: str, mapping, words: int, fabric_handle: Optional[bool], build): + """``build(uc, mc, handle, comm)`` over a new multicast buffer of ``words`` int32 per rank of ``mapping``'s TP + group, every word empty (``uc``: this rank's words; ``mc``: the same words through the multicast mapping). + Collective and eager: every rank of the group calls it at the same point. + + Failure model (as ``MnnvlWorkspace.create``): before allocating, the ranks agree that each of them can (not + capturing a CUDA graph, the buffer within its device's free memory); if one cannot, every rank raises + ``RuntimeError`` and none allocates. A failure that returns from the allocation or from ``build`` is agreed the + same way. A rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange: + that failure is not turned into an error on the other ranks.""" + from tensorrt_llm._torch.distributed.ops import ( + _get_mnnvl_workspace_comm, + _make_mnnvl_mcast_buffer, + _mnnvl_device_index, + _mnnvl_workspace_all_succeeded, + ) + + use_fabric_handle = mapping.is_multi_node() if fabric_handle is None else bool(fabric_handle) + comm = _get_mnnvl_workspace_comm(mapping) + # Every condition one rank alone can fail is checked before the allocation, and the ranks agree on it: a rank + # failing inside the allocation would leave its peers in the handle exchange. + problem: Optional[str] = None + if torch.cuda.is_current_stream_capturing(): + problem = "it is collective and allocates: call it outside CUDA-graph capture" + else: + free_bytes, _ = torch.cuda.mem_get_info(_mnnvl_device_index(mapping)) + if free_bytes < words * 4: + problem = f"its {words * 4} bytes exceed the {free_bytes} free on this rank's device" + if not _mnnvl_workspace_all_succeeded(comm, problem is None): + raise RuntimeError( + f"{name}.create: not every rank can allocate ({problem or 'another rank cannot'})" + ) + error: Optional[Exception] = None + state = None + try: + handle = _make_mnnvl_mcast_buffer(comm, words * 4, mapping, use_fabric_handle) + uc = handle.get_uc_buffer(mapping.tp_rank, (words,), torch.int32, 0) + mc = handle.get_mc_buffer((words,), torch.int32, 0) + uc.fill_(EMPTY_WORD) + state = build(uc, mc, handle, comm) + torch.cuda.synchronize() + except Exception as exc: # noqa: BLE001 -- reported to every rank below, then re-raised + error = exc + # Also the barrier that keeps any rank from pushing into a peer's buffer before the peer has emptied it. + if not _mnnvl_workspace_all_succeeded(comm, error is None): + raise RuntimeError(f"{name}: allocation failed on at least one rank") from error + return state + + @dataclass(eq=False) class K3MoeHeadWorkspace: """One TP group's MoE head all-gather buffers, read and written by ``trtllm::k3_moe_front`` (alone, or as the @@ -146,54 +195,25 @@ class K3MoeHeadWorkspace: @classmethod def create(cls, mapping, fabric_handle: Optional[bool] = None) -> "K3MoeHeadWorkspace": """Allocate and arm a workspace for ``mapping``'s TP group. Collective: every rank of the group calls it at the - same point, eagerly (not under CUDA-graph capture); it returns on every rank or raises on every rank. + same point, eagerly (not under CUDA-graph capture); the failure model is :func:`create_mcast_state`'s. ``fabric_handle``: share the memory by fabric handle (required across nodes) rather than POSIX file descriptor; default ``mapping.is_multi_node()``.""" - if torch.cuda.is_current_stream_capturing(): - raise RuntimeError( - "K3MoeHeadWorkspace.create is collective and allocates: call it outside CUDA-graph capture" - ) - from tensorrt_llm._torch.distributed.ops import ( - _get_mnnvl_workspace_comm, - _make_mnnvl_mcast_buffer, - _mnnvl_workspace_all_succeeded, - ) - from . import k3_route_quant_ag as layout - words = layout.workspace_words(mapping.tp_size) - use_fabric_handle = ( - mapping.is_multi_node() if fabric_handle is None else bool(fabric_handle) - ) - comm = _get_mnnvl_workspace_comm(mapping) - error: Optional[Exception] = None - workspace = None - try: - handle = _make_mnnvl_mcast_buffer(comm, words * 4, mapping, use_fabric_handle) - uc = handle.get_uc_buffer(mapping.tp_rank, (words,), torch.int32, 0) - mc = handle.get_mc_buffer((words,), torch.int32, 0) - uc.fill_(EMPTY_WORD) - flags = torch.zeros(4, dtype=torch.int32, device=uc.device) - ready = torch.zeros(32, dtype=torch.int32, device=uc.device) - torch.cuda.synchronize() - workspace = cls( + def build(uc, mc, handle, comm): + return cls( uc=uc, mc=mc, - flags=flags, - ready=ready, + flags=torch.zeros(4, dtype=torch.int32, device=uc.device), + ready=torch.zeros(32, dtype=torch.int32, device=uc.device), rank=mapping.tp_rank, world_size=mapping.tp_size, handle=handle, comm=comm, ) - except Exception as exc: # noqa: BLE001 -- reported to every rank below, then re-raised - error = exc - # Also the barrier that keeps any rank from pushing into a peer's buffer before the peer has emptied it. - if not _mnnvl_workspace_all_succeeded(comm, error is None): - raise RuntimeError( - "K3MoeHeadWorkspace: allocation failed on at least one rank" - ) from error - return workspace + + words = layout.workspace_words(mapping.tp_size) + return create_mcast_state("K3MoeHeadWorkspace", mapping, words, fabric_handle, build) class K3MoeState: diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py index 350865704cc9..c93bd9ff2458 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py @@ -131,20 +131,36 @@ def _slab_args(x_slab: Optional[torch.Tensor], slab_buf: int, fallback: torch.Te def _create_buffer(cls, mapping, words: int, flag_words: int, fabric_handle: Optional[bool], arm_flags: Optional[Callable[[torch.Tensor], None]] = None): # fmt: skip """A ``cls`` over a new multicast buffer of ``words`` int32 per rank of ``mapping``'s TP group, every word empty, - and ``flag_words`` int32 flags, zero (then ``arm_flags``). Collective and eager: every rank of the group calls it at - the same point, outside CUDA-graph capture; it returns on every rank or raises on every rank.""" - if torch.cuda.is_current_stream_capturing(): - raise RuntimeError( - f"{cls.__name__}.create is collective and allocates: call it outside CUDA-graph capture" - ) + and ``flag_words`` int32 flags, zero (then ``arm_flags``). Collective and eager: every rank of the group calls it + at the same point. + + Failure model (as ``MnnvlWorkspace.create``): before allocating, the ranks agree that each of them can (not + capturing a CUDA graph, the buffer within its device's free memory); if one cannot, every rank raises + ``RuntimeError`` and none allocates. A failure that returns from the allocation is agreed the same way. A rank + that fails inside the allocation's handle exchange can leave its peers waiting in that exchange: that failure is + not turned into an error on the other ranks.""" from tensorrt_llm._torch.distributed.ops import ( _get_mnnvl_workspace_comm, _make_mnnvl_mcast_buffer, + _mnnvl_device_index, _mnnvl_workspace_all_succeeded, ) use_fabric_handle = mapping.is_multi_node() if fabric_handle is None else bool(fabric_handle) comm = _get_mnnvl_workspace_comm(mapping) + # Every condition one rank alone can fail is checked before the allocation, and the ranks agree on it: a rank + # failing inside the allocation would leave its peers in the handle exchange. + problem: Optional[str] = None + if torch.cuda.is_current_stream_capturing(): + problem = "it is collective and allocates: call it outside CUDA-graph capture" + else: + free_bytes, _ = torch.cuda.mem_get_info(_mnnvl_device_index(mapping)) + if free_bytes < words * 4: + problem = f"its {words * 4} bytes exceed the {free_bytes} free on this rank's device" + if not _mnnvl_workspace_all_succeeded(comm, problem is None): + raise RuntimeError( + f"{cls.__name__}.create: not every rank can allocate ({problem or 'another rank cannot'})" + ) error: Optional[Exception] = None state = None try: @@ -197,7 +213,7 @@ class K3SandwichWorkspace: @classmethod def create(cls, mapping, fabric_handle: Optional[bool] = None) -> "K3SandwichWorkspace": """Allocate and arm a workspace for ``mapping``'s TP group. Collective: every rank of the group calls it at the - same point, eagerly (not under CUDA-graph capture); it returns on every rank or raises on every rank. + same point, eagerly (not under CUDA-graph capture); the failure model is ``_create_buffer``'s. ``fabric_handle``: share the memory by fabric handle (required across nodes) rather than POSIX file descriptor; default ``mapping.is_multi_node()``.""" from . import k3_sandwich_kernel as kernel From bd4a4078d993570e014270a27909fe52e9770ffb Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 20:20:24 -0700 Subject: [PATCH 068/161] [None][refactor] Kimi K3 MoE: trtllm::k3_moe, a torch op over the caller's state The routed-experts kernel launches through a torch op again, so its writes are declared, as for every other K3 op: - trtllm::k3_moe takes the routing and MXFP8 latent (k3_route_quant's or k3_moe_front's outputs), the layer's weight buffers, the state's slab and FC2 partial rows, the layer's counters and, for a head_flags build, the head workspace's ready words and flags; - mutates_args names the slab, the partials, the counters, the ready words, the flags and the optional out buffer. - One op serves both builds (up to 8 tokens and the wide build up to 64), each compiled with TVM-FFI on its first call per device. K3MoeState and K3MoeWideState own the scratch; K3MoeLayer owns a layer's weights and counters and calls the op. - The routing is no longer folded into the call (K3MoeLayer() / .front()): the caller runs trtllm::k3_route_quant or trtllm::k3_moe_front first. Signed-off-by: Vasanth Sabavat --- .../cute_dsl_kernels/k3_fused_moe/op.py | 645 ++++++++---------- 1 file changed, 280 insertions(+), 365 deletions(-) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py index 9cbfe2fe7262..85c8ff3f3e59 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py @@ -12,30 +12,24 @@ # 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. -"""Kimi K3 routed experts for decode: the persistent CuTe DSL kernel ``k3_moe`` (``k3_moe_kernel.py``). - -For M <= 8 tokens (:class:`K3MoeState`, :class:`K3MoeLayer`), two kernels on the current stream, no host -synchronization: - -1. the routing and the MXFP8 input quantization: ``trtllm::k3_route_quant`` from the router logits and the latent - (the routing the TRTLLM-Gen path uses under separated routing: sigmoid, top-16 of sigmoid + bias, unbiased scores - renormalized times the routed scaling factor, ties to the lower id; the CuTe DSL form of - ``trtllm::kimi_k3_noaux_tc_mxfp8_quant`` with the same outputs bit for bit), or ``trtllm::k3_moe_front`` from the - MoE input (:meth:`K3MoeLayer.front`: head GEMV, head all-gather, routing, MXFP8 latent, shared gate_up + SiTU); -2. ``k3_moe``, launched as a programmatic dependent of the first: this rank's (expert, token) groups in its prologue, - then FC1 + SiTU + FC2 with the routing-weighted, deterministic combine. - -The result is this rank's routed partial ``[M, 3584]`` bf16, the tensor the TRTLLM-Gen W4A8_MXFP4_MXFP8 op returns, -so the routed-latent all-reduce and the latent-up tail are unchanged. Weights are the TRTLLM-Gen buffers, read in -place. - -Steps of up to 64 tokens use :class:`K3MoeWideState` (the m_max 64 build of ``k3_moe``, launched after -``trtllm::k3_route_quant``). - -The caller owns all state: the scratch a state's layers share (the intermediate slab, left armed by every call, and -the FC2 partial rows), each layer's counters (left zero by every call), and the head all-gather buffers of the front -(:class:`K3MoeHeadWorkspace`, collective over the TP group). Build them before CUDA-graph capture; each kernel compiles -on its first call, which must also come before capture. +"""Kimi K3 routed experts for decode: ``trtllm::k3_moe``, the persistent CuTe DSL kernel ``k3_moe`` +(``k3_moe_kernel.py``): this rank's (expert, token) groups in its prologue, then FC1 + SiTU + FC2 with the +routing-weighted, deterministic combine, over the TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers read in place. + +Its inputs are the routing and the MXFP8 latent of ``trtllm::k3_route_quant`` (the routing the TRTLLM-Gen path uses +under separated routing: sigmoid, top-16 of sigmoid + bias, unbiased scores renormalized times the routed scaling +factor, ties to the lower id; the CuTe DSL form of ``trtllm::kimi_k3_noaux_tc_mxfp8_quant`` with the same outputs bit +for bit) or of ``trtllm::k3_moe_front`` (head GEMV, head all-gather, routing, MXFP8 latent, shared gate_up + SiTU), +and it launches as their programmatic dependent. The result is this rank's routed partial ``[M, 3584]`` bf16, the +tensor the TRTLLM-Gen W4A8_MXFP4_MXFP8 op returns, so the routed-latent all-reduce and the latent-up tail are +unchanged. + +Two builds: up to 8 tokens (:class:`K3MoeState`, optionally acquiring the front's outputs through the head +workspace's ready words) and up to 64 (:class:`K3MoeWideState`). The caller owns all state: the scratch a state's +layers share (the intermediate slab, left armed by every call, and the FC2 partial rows), each layer's counters +(:class:`K3MoeLayer`, left zero by every call), and the head all-gather buffers of the front +(:class:`K3MoeHeadWorkspace`, collective over the TP group). Build them before CUDA-graph capture; each build +compiles on its first call, which must also come before capture. """ from __future__ import annotations @@ -216,82 +210,77 @@ def build(uc, mc, handle, comm): return create_mcast_state("K3MoeHeadWorkspace", mapping, words, fabric_handle, build) -class K3MoeState: - """``k3_moe`` for 1..8 decode tokens on one device: its build and the scratch its layers share, i.e. the FC1 -> - FC2 intermediate slab (armed between calls: FP8 -0.0 values, E8M0 NaN scale words) and the FC2 partial rows. Build - it eagerly before CUDA-graph capture and keep it with the model; every layer takes its own counters from - :meth:`layer`. The layers of one state run in one stream order (they share the scratch). The kernel compiles on the - first call, which must therefore come before capture. +WIDE_MAX_TOKENS = 64 + +# The kernel's tensor arguments after the 18 it always reads: the fused all-reduce's buffers (3), the fold's inputs +# (5), the head flags build's ready words and head flags (2), the latent slab (1). The builds this module compiles +# read only the ready words and head flags (head_flags), so the others are given a stand-in they never touch. +_OPTIONAL_ARGS = 11 +# (alignment, leading dim) of every tensor argument, in the kernel's order. +_ALIGNS = [16, 16, 16, 16, 16, 4, 16, 16, 16, 16, 16, 16, 16, 16, 16, 4, 4, 4] + [16] * _OPTIONAL_ARGS +_LEADING = [0, 0, 0, 1, 2, 2, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 0] + [0] * _OPTIONAL_ARGS +_compiled: Dict[tuple, object] = {} + + +def _part_rows(mod) -> int: + """Rows of the FC2 partial buffer: the M <= 8 build's slices fit in its groups' rows, the wide build's are + PART_ROWS.""" + return mod.PART_ROWS if mod.WIDE else mod.G_CAP * _TOKEN_SLOTS + + +def _config(i_tp: int, num_ctas: int, num_local: int, m_max: int, use_pdl: bool, head_flags: bool) -> dict: + """The kernel options of one build (trace-time constants).""" + return { + "i_tp": i_tp, + "num_ctas": num_ctas, + "num_local": num_local, + "m_max": m_max, + "pdl": int(use_pdl), + "head_flags": int(head_flags), + "lat_slab": 0, + } + - ``head_flags``: the build in which k3_moe acquires the front's routing and MXFP8 rows through the head workspace's - ready words instead of waiting for the front's grid (:meth:`K3MoeLayer.front` only). ``config`` overrides kernel - options (tests and A/B runs: ``pdl``, ``num_ctas``, ...); anything it leaves out takes the kernel's default.""" +class _K3MoeScratch: + """The scratch of one ``k3_moe`` build on one device, shared by the layers that run on it: the FC1 -> FC2 + intermediate slab (armed between calls: FP8 -0.0 values, E8M0 NaN scale words) and the FC2 partial rows.""" def __init__( self, device: torch.device, i_tp: int, num_local: int, - head_flags: bool = False, - config: Optional[dict] = None, + m_max: int, + use_pdl: bool, + head_flags: bool, + num_ctas: Optional[int], ): if torch.cuda.is_current_stream_capturing(): raise RuntimeError( - "K3MoeState allocates its scratch: build it outside CUDA-graph capture" - ) - # One persistent CTA per SM (config "num_ctas" caps it, e.g. for a grid-size A/B). - num_ctas = torch.cuda.get_device_properties(device).multi_processor_count - cfg = { - "i_tp": i_tp, - "num_ctas": num_ctas, - "num_local": num_local, - "head_flags": int(head_flags), - "lat_slab": 0, - } - cfg.update(config or {}) - self.mod = mod = _kernel_module(cfg) - if mod.FUSED_AR or mod.FOLD or mod.LAT_SLAB or mod.WIDE: - raise ValueError( - "K3MoeState is the M <= 8 build without the fused all-reduce, the fold or the slab" + f"{type(self).__name__} allocates its scratch: build it outside CUDA-graph capture" ) + # One persistent CTA per SM (num_ctas caps it, e.g. for a grid-size A/B). + if num_ctas is None: + num_ctas = torch.cuda.get_device_properties(device).multi_processor_count + self.config = _config(i_tp, num_ctas, num_local, m_max, use_pdl, head_flags) + self.mod = mod = _kernel_module(self.config) self.device = device self.i_tp = i_tp self.num_local = num_local - self.head_flags = mod.HEAD_FLAGS + self.m_max = m_max + self.num_ctas = num_ctas + self.use_pdl = bool(use_pdl) + self.head_flags = bool(head_flags) g_cap = mod.G_CAP kw = dict(device=device) - # Lamport slab, armed: FP8 -0.0 values; E8M0 NaN in bytes 0..3 of each 16-byte scale - # group. Every call leaves the groups it used armed again. + # Lamport slab, armed: FP8 -0.0 values; E8M0 NaN in bytes 0..3 of each 16-byte scale group. Every call + # leaves the groups it used armed again. self.c = torch.full((g_cap, _TOKEN_SLOTS, i_tp), -128, dtype=torch.int8, **kw) self.cs = torch.zeros(g_cap, _TOKEN_SLOTS, mod.SF_STRIDE0, dtype=torch.int8, **kw) self.cs.view(g_cap, _TOKEN_SLOTS, mod.K2_TILES, mod.SFB_GROUP_BYTES)[..., :4] = -1 - # FC2 partial rows. A call writes the rows of its M tokens; the combine also loads the rows past M (their sums - # are dropped), so the buffer starts zeroed and those loads never read unwritten memory. - self.part = torch.zeros(g_cap * _TOKEN_SLOTS, HIDDEN_SIZE, dtype=torch.float32, **kw) - self.c_t = _view(self.c, 16, 2) - self.cs_t = _view(self.cs, 4, 2) - self.c_words_t = _view(self.c.view(-1).view(torch.int32), 16, 0) - self.cs_words_t = _view(self.cs.view(-1).view(torch.int32), 16, 0) - self.b2_t = _view(self.c.permute(2, 1, 0), 16, 0, mod.b_dtype) - sfb2 = ( - self.cs.view(torch.uint8) - .reshape(g_cap, _TOKEN_SLOTS, mod.K2_TILES, mod.SFB_GROUP_BYTES) - .permute(3, 2, 1, 0) - ) - self.sfb2_t = _view(sfb2, 16, 0, mod.sf_dtype) - self.part_t = _view(self.part, 16, 1) - # Stand-in for the buffers of the options this build does not have (fused all-reduce, fold inputs, latent - # slab, and the ready words without head_flags). - self.unused = torch.zeros(4, dtype=torch.int32, **kw) - self.unused_t = _view(self.unused, 16, 0) - # The route+quant kernel triggers k3_moe's launch right after its own grid dependency: - # k3_moe waits for the whole route+quant grid before reading its outputs. - self.route_kwargs = {"early_trigger": True} if mod.USE_PDL else {} - from ..k3_route_quant import ( - op as _k3_route_quant_op, # noqa: F401 (registers trtllm::k3_route_quant) - ) - - self.compiled = None + # FC2 partial rows. A call writes the rows of its tokens; the M <= 8 build's combine also loads the rows past + # M (their sums are dropped), so the buffer starts zeroed and those loads never read unwritten memory. + self.part = torch.zeros(_part_rows(mod), HIDDEN_SIZE, dtype=torch.float32, **kw) def layer( self, @@ -303,243 +292,55 @@ def layer( """A layer's handle: its experts' TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers (read in place) and its counters.""" return K3MoeLayer(self, w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale) + @property + def compiled(self) -> bool: + """Whether this build has been compiled on this device (by any state's first call).""" + return _compile_key(self.device, self.config) in _compiled -class K3MoeLayer: - """One MoE layer on a :class:`K3MoeState`: its weights as the kernel reads them and its counters (int32, zero - between calls; every call leaves them zero).""" - def __init__( - self, - state: K3MoeState, - w3_w1_weight: torch.Tensor, - w3_w1_weight_scale: torch.Tensor, - w2_weight: torch.Tensor, - w2_weight_scale: torch.Tensor, - ): - if torch.cuda.is_current_stream_capturing(): - raise RuntimeError( - "K3MoeLayer allocates its counters: build it outside CUDA-graph capture" - ) - ok, why = is_supported( - w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale, state.num_local - ) - if not ok or w3_w1_weight.shape[1] != 2 * state.i_tp: - raise ValueError(f"k3_moe layer: {why or 'intermediate size differs from the state'}") - mod = state.mod - e, two_i, _ = w3_w1_weight.shape - i_tp = two_i // 2 - sfa1 = w3_w1_weight_scale.view(e, two_i // 128, HIDDEN_SIZE // 128, 512).permute(3, 2, 1, 0) - sfa2 = w2_weight_scale.view(e, HIDDEN_SIZE // 128, i_tp // 128, 512).permute(3, 2, 1, 0) - self.state = state - self.counters = torch.zeros(mod.NUM_STATE, dtype=torch.int32, device=state.device) - self.weights = ( - _view(w3_w1_weight.view(torch.int8).permute(2, 1, 0), 16, 0), - _view(sfa1, 16, 0, mod.sf_dtype), - _view(w2_weight.view(torch.int8).permute(2, 1, 0), 16, 0), - _view(sfa2, 16, 0, mod.sf_dtype), - ) - self.counters_t = _view(self.counters, 4, 0) +class K3MoeState(_K3MoeScratch): + """``k3_moe`` for 1..8 decode tokens on one device: the build's scratch, shared by its layers. Build it eagerly + before CUDA-graph capture and keep it with the model; every layer takes its own counters from :meth:`layer`. The + layers of one state run in one stream order (they share the scratch). The kernel compiles on its first call for + the build, which must therefore come before capture. - def __call__( - self, - hidden_states: torch.Tensor, - router_logits: torch.Tensor, - e_score_correction_bias: torch.Tensor, - local_expert_offset: int, - routed_scaling_factor: float, - ) -> torch.Tensor: - """This rank's routed partial ``[M, 3584]`` bf16 for M <= 8 decode tokens: ``trtllm::k3_route_quant`` (the - routing and the MXFP8 latent), then ``k3_moe`` launched as its programmatic dependent. + ``head_flags``: the build in which k3_moe acquires the MoE front's routing and MXFP8 rows through the head + workspace's ready words instead of waiting for the front's grid (``trtllm::k3_moe_front`` with ``ag_ready``, then + ``trtllm::k3_moe`` with the workspace's ``ready`` and ``flags``). ``use_pdl``: launch k3_moe as a programmatic + dependent of its producer (default ``TRTLLM_ENABLE_PDL``). ``num_ctas``: the persistent grid (default one CTA per + SM).""" - ``hidden_states``: bf16 ``[M, 3584]`` latent; ``router_logits``: fp32 ``[M, 896]``; - ``e_score_correction_bias``: fp32 ``[896]``; the layer's experts hold global ids - ``[local_expert_offset, local_expert_offset + num_local)``.""" - st = self.state - if st.head_flags: - raise ValueError( - "a head_flags build takes the front's ready words: call K3MoeLayer.front" - ) - _check_tokens(hidden_states) - ids, weights, x_fp8, x_sf = torch.ops.trtllm.k3_route_quant( - router_logits.contiguous(), e_score_correction_bias, hidden_states.contiguous(), - float(routed_scaling_factor), **st.route_kwargs, - ) # fmt: skip - return self._launch(ids, weights, x_fp8, x_sf, local_expert_offset, routed_scaling_factor) - - def front( + def __init__( self, - x: torch.Tensor, - w_front: torch.Tensor, - e_score_correction_bias: torch.Tensor, - local_expert_offset: int, - routed_scaling_factor: float, - shared_cols: int, - gate_cap: float, - linear_cap: float, - head: K3MoeHeadWorkspace, - ) -> Tuple[torch.Tensor, torch.Tensor]: - """``trtllm::k3_moe_front`` (head GEMV, head all-gather over ``head``, routing, MXFP8 latent, shared gate_up + - SiTU) then ``k3_moe`` for the MoE input ``x`` bf16 ``[M, 7168]`` (the same on every rank). Returns ``(y [M, - 3584] bf16, shared activation [M, shared_cols] bf16)``. A ``head_flags`` build acquires the front's routing and - MXFP8 rows through ``head.ready`` instead of waiting for the front's grid, whose shared tiles may still be - running.""" - from . import front_op # noqa: F401 (registers trtllm::k3_moe_front) - - st = self.state - _check_tokens(x) - ready = head.ready if st.head_flags else None - ids, weights, x_fp8, x_sf, shared = torch.ops.trtllm.k3_moe_front( - x.contiguous(), w_front, e_score_correction_bias, float(routed_scaling_factor), shared_cols, gate_cap, - linear_cap, head.uc, head.mc, head.flags, head.rank, head.world_size, ag_ready=ready, - ) # fmt: skip - flag_in = None - if st.head_flags: - flag_in = (_view(head.ready.view(-1), 16, 0), _view(head.flags.view(-1), 16, 0)) - y = self._launch( - ids, weights, x_fp8, x_sf, local_expert_offset, routed_scaling_factor, flag_in - ) - return y, shared - - def _launch(self, ids, weights, x_fp8, x_sf, local_offset, scale, flag_in=None) -> torch.Tensor: - import cuda.bindings.driver as cuda_driver - import cutlass.cute as cute - - st = self.state - mod = st.mod - num_tokens = ids.shape[0] - a1, sfa1, a2, sfa2 = self.weights - y = torch.empty(num_tokens, HIDDEN_SIZE, dtype=torch.bfloat16, device=x_fp8.device) - stream = cuda_driver.CUstream(torch.cuda.current_stream().cuda_stream) - b1 = _view(x_fp8.view(torch.uint8).permute(1, 0), 16, 0, mod.b_dtype) - sfb1 = _view( - x_sf.view(torch.uint8).view(num_tokens, HIDDEN_SIZE // _SF_VEC), 16, 1, mod.sf_dtype - ) - u = st.unused_t - args = [ - a1, b1, sfa1, sfb1, st.c_t, st.cs_t, st.c_words_t, st.cs_words_t, a2, st.b2_t, sfa2, st.sfb2_t, - _view(y, 16, 1), _view(y.view(torch.int32), 16, 1), st.part_t, _view(ids, 4, 1), _view(weights, 4, 1), - self.counters_t, - u, u, u, # the fused all-reduce's buffers - u, u, u, u, u, # the fold's inputs - *(flag_in or (u, u)), # the ready words and the head flags (head_flags) - u, # the latent slab - ] # fmt: skip - scalars = (num_tokens, local_offset, st.num_local, 0, float(scale), 0, 0) - if st.compiled is None: - if torch.cuda.is_current_stream_capturing(): - raise RuntimeError( - "k3_moe compiles on its first call: call it once before CUDA-graph capture" - ) - st.compiled = cute.compile(mod.k3_moe, *args, *scalars, stream) - st.compiled(*args, *scalars, stream) - return y - - -def _check_tokens(x: torch.Tensor) -> None: - if not 0 < x.shape[0] <= MAX_TOKENS: - raise ValueError(f"k3_moe handles 1 to {MAX_TOKENS} tokens, got {x.shape[0]}") - - -# --------------------------------------------------------------------------- -# Steps of up to 64 tokens (e.g. R x 8 speculative verify tokens): the m_max 64 build of -# ``k3_moe``, one launch per call, after ``trtllm::k3_route_quant``. Its state belongs to the -# caller: nothing here is cached per process beyond the kernel module of the configuration. - -WIDE_MAX_TOKENS = 64 - + device: torch.device, + i_tp: int, + num_local: int, + head_flags: bool = False, + use_pdl: Optional[bool] = None, + num_ctas: Optional[int] = None, + ): + if use_pdl is None: + use_pdl = os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" + super().__init__(device, i_tp, num_local, MAX_TOKENS, use_pdl, head_flags, num_ctas) -class K3MoeWideState: - """``k3_moe`` for 1..64 tokens on one device: the compiled kernel and the scratch its layers - share, i.e. the FC1 -> FC2 intermediate slab (armed between calls: FP8 -0.0 values, E8M0 NaN - scale words) and the FC2 slice partials, all sized for 64 tokens and kept at fixed addresses. - Build it eagerly before CUDA-graph capture and keep it with the model; every layer takes its - own counters from :meth:`layer`. The layers of one state run in one stream order (they share - the scratch). The kernel is compiled (TVM-FFI, explicit stream) by the first call, which must - therefore come before capture. - ``use_pdl``: launch ``k3_moe`` as a programmatic dependent; its producer must then be - ``trtllm::k3_route_quant`` with ``early_trigger=True`` (or any kernel whose outputs ``k3_moe`` - may read once that grid has completed).""" +class K3MoeWideState(_K3MoeScratch): + """``k3_moe`` for 1..64 tokens on one device (the m_max 64 build): its scratch, sized for 64 tokens and shared by + its layers, as :class:`K3MoeState`. ``use_pdl``: launch ``k3_moe`` as a programmatic dependent; its producer must + then be ``trtllm::k3_route_quant`` with ``early_trigger=True`` (or any kernel whose outputs ``k3_moe`` may read + once that grid has completed).""" def __init__(self, device: torch.device, i_tp: int, num_local: int, use_pdl: bool = True): - if torch.cuda.is_current_stream_capturing(): - raise RuntimeError( - "K3MoeWideState allocates its scratch: build it outside CUDA-graph capture" - ) - num_ctas = torch.cuda.get_device_properties(device).multi_processor_count - config = { - "i_tp": i_tp, - "num_ctas": num_ctas, - "num_local": num_local, - "m_max": WIDE_MAX_TOKENS, - "pdl": int(use_pdl), - } - self.mod = mod = _kernel_module(config) - self.device = device - self.i_tp = i_tp - self.num_local = num_local - g_cap = mod.G_CAP - kw = dict(device=device) - self.c = torch.full((g_cap, _TOKEN_SLOTS, i_tp), -128, dtype=torch.int8, **kw) - self.cs = torch.zeros(g_cap, _TOKEN_SLOTS, mod.SF_STRIDE0, dtype=torch.int8, **kw) - self.cs.view(g_cap, _TOKEN_SLOTS, mod.K2_TILES, mod.SFB_GROUP_BYTES)[..., :4] = -1 - self.part = torch.empty(mod.PART_ROWS, HIDDEN_SIZE, dtype=torch.float32, **kw) - # Stand-in for the buffers of the options this build does not have (fused all-reduce, - # fold, head flags, latent slab). - self.unused = torch.zeros(4, dtype=torch.int32, **kw) - # The kernel's views of the scratch, in its argument order (FP8 / E8M0 data as bytes). - sfb2 = ( - self.cs.view(torch.uint8) - .reshape(g_cap, _TOKEN_SLOTS, mod.K2_TILES, mod.SFB_GROUP_BYTES) - .permute(3, 2, 1, 0) - ) - self.scratch = ( - self.c, - self.cs, - self.c.view(-1).view(torch.int32), - self.cs.view(-1).view(torch.int32), - self.c.view(torch.uint8).permute(2, 1, 0), - sfb2, - self.part, - ) - self.compiled = None - - def layer( - self, - w3_w1_weight: torch.Tensor, - w3_w1_weight_scale: torch.Tensor, - w2_weight: torch.Tensor, - w2_weight_scale: torch.Tensor, - ) -> "K3MoeWideLayer": - """A layer's handle: its experts' TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers (read in place) and - its counters.""" - return K3MoeWideLayer(self, w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale) - - def _compile(self, args, scalars): - """TVM-FFI build for these torch arguments' types and layouts (any M up to 64).""" - import cutlass.cute as cute - - # As the M <= 8 op's views: (alignment, leading dim) per tensor argument; the 11 stand-ins last. - aligns = [16, 16, 16, 16, 16, 4, 16, 16, 16, 16, 16, 16, 16, 16, 16, 4, 4, 4] + [16] * 11 - leading = [0, 0, 0, 1, 2, 2, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 0] + [0] * 11 - assert len(args) == len(aligns) - signature = [_view(t, a, d) for t, a, d in zip(args, aligns, leading)] - return cute.compile( - self.mod.k3_moe, - *signature, - *scalars, - cute.runtime.make_fake_stream(), - options="--enable-tvm-ffi", - ) + super().__init__(device, i_tp, num_local, WIDE_MAX_TOKENS, use_pdl, False, None) -class K3MoeWideLayer: - """One MoE layer on a :class:`K3MoeWideState`: its weights as the kernel reads them and its +class K3MoeLayer: + """One MoE layer on a :class:`K3MoeState` or :class:`K3MoeWideState`: its experts' weight buffers and its counters (int32, zero between calls; every call leaves them zero).""" def __init__( self, - state: K3MoeWideState, + state: _K3MoeScratch, w3_w1_weight: torch.Tensor, w3_w1_weight_scale: torch.Tensor, w2_weight: torch.Tensor, @@ -547,25 +348,16 @@ def __init__( ): if torch.cuda.is_current_stream_capturing(): raise RuntimeError( - "K3MoeWideLayer allocates its counters: build it outside CUDA-graph capture" + "K3MoeLayer allocates its counters: build it outside CUDA-graph capture" ) ok, why = is_supported( w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale, state.num_local ) if not ok or w3_w1_weight.shape[1] != 2 * state.i_tp: - raise ValueError( - f"k3_moe wide layer: {why or 'intermediate size differs from the state'}" - ) - e, two_i, _ = w3_w1_weight.shape - i_tp = two_i // 2 + raise ValueError(f"k3_moe layer: {why or 'intermediate size differs from the state'}") self.state = state + self.weights = (w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale) self.counters = torch.zeros(state.mod.NUM_STATE, dtype=torch.int32, device=state.device) - self.weights = ( - w3_w1_weight.view(torch.int8).permute(2, 1, 0), - w3_w1_weight_scale.view(e, two_i // 128, HIDDEN_SIZE // 128, 512).permute(3, 2, 1, 0), - w2_weight.view(torch.int8).permute(2, 1, 0), - w2_weight_scale.view(e, HIDDEN_SIZE // 128, i_tp // 128, 512).permute(3, 2, 1, 0), - ) def __call__( self, @@ -574,58 +366,181 @@ def __call__( topk_ids: torch.Tensor, topk_weights: torch.Tensor, local_expert_offset: int, + head_ready: Optional[torch.Tensor] = None, + head_flags: Optional[torch.Tensor] = None, out: Optional[torch.Tensor] = None, ) -> torch.Tensor: - """This rank's routed partial ``[M, 3584]`` bf16 for the outputs of - ``trtllm::k3_route_quant``: ``x_fp8`` float8_e4m3fn ``[M, 3584]``, ``x_sf`` its E8M0 - scales (``M * 112`` bytes), ``topk_ids`` int32 ``[M, 16]`` global expert ids, - ``topk_weights`` bf16 ``[M, 16]``; 1 <= M <= 64. The layer's experts hold global ids - ``[local_expert_offset, local_expert_offset + num_local)``. ``out``: bf16, contiguous, at - least ``[M, 3584]``; its first M rows are the result (a fresh tensor without it). Writes - the state's slab (left armed) and partials and this layer's counters (left zero).""" + """``trtllm::k3_moe`` on this layer's experts, counters and its state's scratch: see the op.""" st = self.state - num_tokens = topk_ids.shape[0] - if not 0 < num_tokens <= WIDE_MAX_TOKENS: - raise ValueError(f"k3_moe wide: M must be in [1, {WIDE_MAX_TOKENS}], got {num_tokens}") + y = torch.ops.trtllm.k3_moe( + x_fp8, x_sf, topk_ids, topk_weights, *self.weights, st.c, st.cs, st.part, self.counters, + local_expert_offset, st.num_local, st.num_ctas, st.m_max, st.use_pdl, head_ready, head_flags, out, + ) # fmt: skip + return y if out is None else out[: topk_ids.shape[0]] + + +def _compile_key(device: torch.device, config: dict) -> tuple: + index = device.index if device.index is not None else torch.cuda.current_device() + return (index, tuple(sorted(config.items()))) + + +def _compile(mod, args, scalars): + """TVM-FFI build of ``mod``'s k3_moe for these torch arguments' types and layouts (any M up to the build's).""" + import cutlass.cute as cute + + assert len(args) == len(_ALIGNS) + signature = [_view(t, a, d) for t, a, d in zip(args, _ALIGNS, _LEADING)] + return cute.compile( + mod.k3_moe, + *signature, + *scalars, + cute.runtime.make_fake_stream(), + options="--enable-tvm-ffi", + ) + + +@torch.library.custom_op( + "trtllm::k3_moe", + mutates_args=("c", "cs", "part", "counters", "head_ready", "head_flags", "out"), +) +def k3_moe( + x_fp8: torch.Tensor, + x_sf: torch.Tensor, + topk_ids: torch.Tensor, + topk_weights: torch.Tensor, + w3_w1_weight: torch.Tensor, + w3_w1_weight_scale: torch.Tensor, + w2_weight: torch.Tensor, + w2_weight_scale: torch.Tensor, + c: torch.Tensor, + cs: torch.Tensor, + part: torch.Tensor, + counters: torch.Tensor, + local_expert_offset: int, + num_local: int, + num_ctas: int, + m_max: int, + use_pdl: bool, + head_ready: Optional[torch.Tensor] = None, + head_flags: Optional[torch.Tensor] = None, + out: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """This rank's routed partial ``[M, 3584]`` bf16 from the persistent ``k3_moe`` kernel: FC1 + SiTU + FC2 with the + routing-weighted, deterministic combine over this rank's experts. + + ``x_fp8`` float8_e4m3fn ``[M, 3584]``, ``x_sf`` its E8M0 scales (``M * 112`` bytes), ``topk_ids`` int32 + ``[M, 16]`` global expert ids and ``topk_weights`` bf16 ``[M, 16]``: the outputs of ``trtllm::k3_route_quant`` + (or ``trtllm::k3_moe_front``). The weights are the TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers of ``num_local`` experts, + which hold global ids ``[local_expert_offset, local_expert_offset + num_local)``, read in place. ``c``, ``cs``, + ``part``: the scratch of a :class:`K3MoeState` (``m_max`` 8) or :class:`K3MoeWideState` (``m_max`` 64) of this + build (``num_ctas``, ``use_pdl``); every call leaves the slab armed. ``counters``: the layer's, left zero. + ``head_ready`` / ``head_flags``: a ``head_flags`` build's ready words and head flags (a ``K3MoeHeadWorkspace``'s + ``ready`` and ``flags``), whose epoch the call advances. ``out``: bf16, contiguous, at least ``[M, 3584]``; the + call writes its first M rows and returns an empty ``[0, 3584]`` instead of a new tensor. 1 <= M <= ``m_max``.""" + num_tokens = topk_ids.shape[0] + head = head_ready is not None + if head != (head_flags is not None): + raise ValueError("k3_moe: head_ready and head_flags go together") + if m_max not in (MAX_TOKENS, WIDE_MAX_TOKENS) or (head and m_max != MAX_TOKENS): + raise ValueError( + f"k3_moe: m_max is {MAX_TOKENS} (head flags possible) or {WIDE_MAX_TOKENS}, got {m_max}" + ) + if not 0 < num_tokens <= m_max: + raise ValueError(f"k3_moe: M must be in [1, {m_max}], got {num_tokens}") + if ( + topk_ids.dtype != torch.int32 + or tuple(topk_ids.shape) != (num_tokens, TOP_K) + or topk_weights.dtype != torch.bfloat16 + or tuple(topk_weights.shape) != (num_tokens, TOP_K) + or x_fp8.dtype != torch.float8_e4m3fn + or tuple(x_fp8.shape) != (num_tokens, HIDDEN_SIZE) + or x_sf.numel() != num_tokens * (HIDDEN_SIZE // _SF_VEC) + or not (topk_ids.is_contiguous() and topk_weights.is_contiguous()) + or not (x_fp8.is_contiguous() and x_sf.is_contiguous()) + ): + raise ValueError("k3_moe: expects trtllm::k3_route_quant's outputs for M tokens") + ok, why = is_supported(w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale, num_local) + if not ok: + raise ValueError(f"k3_moe: {why}") + e, two_i, _ = w3_w1_weight.shape + i_tp = two_i // 2 + config = _config(i_tp, num_ctas, num_local, m_max, use_pdl, head) + mod = _kernel_module(config) + g_cap = mod.G_CAP + if ( + c.dtype != torch.int8 + or tuple(c.shape) != (g_cap, _TOKEN_SLOTS, i_tp) + or cs.dtype != torch.int8 + or tuple(cs.shape) != (g_cap, _TOKEN_SLOTS, mod.SF_STRIDE0) + or part.dtype != torch.float32 + or tuple(part.shape) != (_part_rows(mod), HIDDEN_SIZE) + or counters.dtype != torch.int32 + or tuple(counters.shape) != (mod.NUM_STATE,) + or not all(t.is_contiguous() for t in (c, cs, part, counters)) + ): + raise ValueError("k3_moe: c / cs / part / counters are not this build's scratch and layer counters") + if head and ( + head_ready.dtype != torch.int32 + or head_ready.numel() < 2 * MAX_TOKENS + or head_flags.dtype != torch.int32 + or head_flags.numel() < 3 + ): + raise ValueError("k3_moe: head_ready / head_flags must be a K3MoeHeadWorkspace's ready and flags") + if out is None: + y = torch.empty(num_tokens, HIDDEN_SIZE, dtype=torch.bfloat16, device=x_fp8.device) + else: if ( - topk_ids.dtype != torch.int32 - or tuple(topk_ids.shape) != (num_tokens, TOP_K) - or topk_weights.dtype != torch.bfloat16 - or tuple(topk_weights.shape) != (num_tokens, TOP_K) - or x_fp8.dtype != torch.float8_e4m3fn - or tuple(x_fp8.shape) != (num_tokens, HIDDEN_SIZE) - or x_sf.numel() != num_tokens * (HIDDEN_SIZE // _SF_VEC) - or not (topk_ids.is_contiguous() and topk_weights.is_contiguous()) - or not (x_fp8.is_contiguous() and x_sf.is_contiguous()) + out.dtype != torch.bfloat16 + or out.dim() != 2 + or out.shape[0] < num_tokens + or out.shape[1] != HIDDEN_SIZE + or not out.is_contiguous() ): - raise ValueError("k3_moe wide: expects trtllm::k3_route_quant's outputs for M tokens") - if out is None: - y = torch.empty(num_tokens, HIDDEN_SIZE, dtype=torch.bfloat16, device=x_fp8.device) - else: - if ( - out.dtype != torch.bfloat16 - or out.dim() != 2 - or out.shape[0] < num_tokens - or out.shape[1] != HIDDEN_SIZE - or not out.is_contiguous() - ): - raise ValueError("k3_moe wide: out must be contiguous bf16 [>= M, 3584]") - y = out[:num_tokens] - a1, sfa1, a2, sfa2 = self.weights - c, cs, c_words, cs_words, b2, sfb2, part = st.scratch - u = st.unused - args = ( - a1, x_fp8.view(torch.uint8).permute(1, 0), sfa1, - x_sf.view(torch.uint8).view(num_tokens, HIDDEN_SIZE // _SF_VEC), c, cs, c_words, cs_words, a2, b2, - sfa2, sfb2, y, y.view(torch.int32), part, topk_ids, topk_weights, self.counters, - u, u, u, u, u, u, u, u, u, u, u, - ) # fmt: skip - scalars = (num_tokens, local_expert_offset, st.num_local, 0, 1.0, 0, 0) - if st.compiled is None: - if torch.cuda.is_current_stream_capturing(): - raise RuntimeError( - "k3_moe wide compiles on its first call: call it once before CUDA-graph capture" - ) - st.compiled = st._compile(args, scalars) - st.compiled(*args, *scalars, torch.cuda.current_stream().cuda_stream) - return y + raise ValueError("k3_moe: out must be contiguous bf16 [>= M, 3584]") + y = out[:num_tokens] + sfb2 = ( + cs.view(torch.uint8) + .reshape(g_cap, _TOKEN_SLOTS, mod.K2_TILES, mod.SFB_GROUP_BYTES) + .permute(3, 2, 1, 0) + ) + # Stand-in for the options this build does not have (fused all-reduce, fold, latent slab, and the head flags + # without head_flags): the kernel never touches it. + unused = counters + flag_args = (head_ready.view(-1), head_flags.view(-1)) if head else (unused, unused) + args = ( + w3_w1_weight.view(torch.int8).permute(2, 1, 0), + x_fp8.view(torch.uint8).permute(1, 0), + w3_w1_weight_scale.view(e, two_i // 128, HIDDEN_SIZE // 128, 512).permute(3, 2, 1, 0), + x_sf.view(torch.uint8).view(num_tokens, HIDDEN_SIZE // _SF_VEC), + c, cs, c.view(-1).view(torch.int32), cs.view(-1).view(torch.int32), + w2_weight.view(torch.int8).permute(2, 1, 0), + c.view(torch.uint8).permute(2, 1, 0), + w2_weight_scale.view(e, HIDDEN_SIZE // 128, i_tp // 128, 512).permute(3, 2, 1, 0), + sfb2, y, y.view(torch.int32), part, topk_ids, topk_weights, counters, + unused, unused, unused, # the fused all-reduce's buffers + unused, unused, unused, unused, unused, # the fold's inputs + *flag_args, # the ready words and the head flags + unused, # the latent slab + ) # fmt: skip + # (tokens, local offset, local experts, all-reduce rank, routed scaling factor (fold only), slab buffer, re-arm 0) + scalars = (num_tokens, local_expert_offset, num_local, 0, 1.0, 0, 0) + key = _compile_key(x_fp8.device, config) + fn = _compiled.get(key) + if fn is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "trtllm::k3_moe compiles on its first call for each build: call it once before CUDA-graph capture" + ) + fn = _compiled[key] = _compile(mod, args, scalars) + fn(*args, *scalars, torch.cuda.current_stream(x_fp8.device).cuda_stream) + if out is not None: + return y.new_empty((0, HIDDEN_SIZE)) + return y + + +@k3_moe.register_fake +def _(x_fp8, x_sf, topk_ids, topk_weights, w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale, c, cs, + part, counters, local_expert_offset, num_local, num_ctas, m_max, use_pdl, head_ready=None, head_flags=None, + out=None): # fmt: skip + rows = 0 if out is not None else topk_ids.shape[0] + return x_fp8.new_empty((rows, HIDDEN_SIZE), dtype=torch.bfloat16) From 26a3dbc62068b34490ea95152efa367b0e0122f4 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 20:47:47 -0700 Subject: [PATCH 069/161] [None][test] MNNVL catalog matrices: calls queued without host sync; create() refused on every rank Both MNNVL matrices enqueue three rounds of 52 calls (the three MNNVL ops, both all-reduce paths, Kimi K3's step sizes, a layer's inputs taken from the previous layer's outputs on the device) with no host synchronization between calls while a random rank starts late and another pauses, then check every result and the flag words. The create check follows MnnvlWorkspace.create's failure model: one rank capturing while the others call it eagerly, then every rank capturing; every rank raises and the workspaces in use are untouched; a create right after returns an armed workspace. The attention- residual calls go through the entry's wrapper, and the all-gather matrix checks that a direct call with the wrong world size raises on every rank. The contracts state the failure model, the order-not-timing invariant, the one-shot kernel's early trigger for every caller, and the new all-gather schema. Signed-off-by: Vasanth Sabavat --- .../catalog/comm/mnnvl_allgather_split.md | 72 +++++--- .../catalog/comm/mnnvl_fusion_allreduce.md | 66 ++++--- .../comm/_mnnvl_allgather_split_op_matrix.py | 164 +++++++++++++++--- .../comm/_mnnvl_fusion_allreduce_op_matrix.py | 154 +++++++++++++--- 4 files changed, 363 insertions(+), 93 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md index 576e54f7d960..0e49452aa4e9 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md @@ -26,6 +26,13 @@ rounded), exact midpoints between two bf16 values in every fourth of them (ties fourth column of each part, and in the fp32 part infinities, a NaN, signed denormals and the largest float, which arrive unchanged. The result is bitwise the same on every rank (certified, every call of the test). +The op also takes `world_size`, the number of ranks it sizes the outputs for (`world_size x B` and `world_size x F` +columns). It must equal the workspace's rank count, which the op checks: another value raises `RuntimeError` on every +rank before the workspace is touched (certified, calling the op directly). The wrapper passes `workspace.world_size`, +so its own signature has no such argument. With the output widths given by an argument, the op has a fake (shape-only) +implementation (`tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py`), so fake-tensor tracing such as `torch.compile` +can run it. + Kimi K3's use: the row-sharded MoE head of a wide decode step (9 to 64 tokens), and of a decode step of at most 8 tokens where the fused MoE front kernel does not run. Rank `r`'s GEMV gives fp32 `[T, 3584/W + 896/W]`: the latent down projection's columns `[r * 3584/W, (r+1) * 3584/W)`, then the router logits of @@ -37,9 +44,9 @@ Fusion boundary. Inside: the bf16 rounding of the leading columns and the exchan `k3_fused_moe/k3_route_quant_ag.py` fuses this exchange with them and states it is bit for bit this op followed by `k3_route_quant`). -The kernel releases its programmatic dependents as soon as it starts; its outputs are complete only when its grid is, -so a kernel launched as its programmatic dependent must wait for the grid before reading them (the kernel's -statement). +The kernel releases its programmatic dependents as soon as its own grid-dependency wait returns; its outputs are +complete only when its grid is, so a kernel launched as its programmatic dependent must wait for the grid before +reading them (the kernel's statement). ## Signature @@ -53,6 +60,14 @@ def mnnvl_allgather_split( def required_buffer_bytes(num_tokens: int, bf16_columns: int, fp32_columns: int, world_size: int) -> int ``` +The wrapper makes one op call, passing `workspace.world_size`, `workspace.comm_buffer(torch.bfloat16)` and +`workspace.buffer_flags` for the op's last three arguments: + +``` +mnnvl_allgather_split(Tensor input, int bf16_columns, int world_size, Tensor(a!) comm_buffer, + Tensor(b!) buffer_flags) -> Tensor[] +``` + ### Certified arguments | Argument | Shape | Dtype | Layout | Device | @@ -62,8 +77,9 @@ def required_buffer_bytes(num_tokens: int, bf16_columns: int, fp32_columns: int, | `workspace` | an `MnnvlWorkspace` of this rank's TP group whose buffers hold the call (see *State*) | — | — | — | | returns | `(bf16_out, fp32_out)`: `[T, W x B]` and `[T, W x F]` | bf16, fp32 | contiguous, newly allocated | = `input.device` | -`input` is read only. `required_buffer_bytes` = `T x W x (2B + 4F)`, the bytes the call writes into one Lamport -buffer (certified equal to what the call records, every call of the split grid). +`input` is read only. The op writes the workspace's buffers and flag words; its schema declares both mutable. +`required_buffer_bytes` = `T x W x (2B + 4F)`, the bytes the call writes into one Lamport buffer (certified equal to +what the call records, every call of the split grid). ## State @@ -81,10 +97,21 @@ buffer; the flags do not move and the next call is correct; the op has the same two-shot `[64, 7168]` all-reduce of its shared-workspace sequence). **Who creates it, and when.** The target, in `post_load_weights`, with -`MnnvlWorkspace.create(mapping, buffer_bytes, fabric_handle=None)`: collective over the TP group, eager, every word -and flag armed before any rank returns (see `mnnvl_allreduce_attn_res.md`). It refuses CUDA-graph capture: certified -with every rank capturing, each raising `RuntimeError`. The check runs before any communication, so a rank that is -not capturing while its peers are would go on into the communicator split and wait for them (code). +`MnnvlWorkspace.create(mapping, buffer_bytes, fabric_handle=None)` (see `mnnvl_allreduce_attn_res.md`): + +- collective over the TP group: every rank calls it at the same point; +- failure model: + - before allocating, the ranks agree that each of them can (not capturing, a valid `buffer_bytes`, the three + buffers within that rank's free device memory). If one cannot, every rank raises `RuntimeError` and none + allocates (certified: one rank inside a CUDA-graph capture while the others are not, and then every rank + capturing; each time every rank raises, the capturing ranks' message naming the capture, and the workspaces in + use are untouched; a create right after, eager on every rank, returns an armed workspace whose first call is + correct); + - a failure that returns from the allocation is agreed the same way; + - a rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange; this + is not turned into an error on the other ranks; +- eager: it allocates and exchanges handles, so it refuses to run under CUDA-graph capture (every rank raises); +- it arms every buffer word and the flags before any rank returns. **Which ops may share one object.** Every MNNVL op of the group takes the same `comm_buffer` / `buffer_flags`: `comm/mnnvl_allreduce_attn_res`, `comm/mnnvl_fusion_allreduce` on either path, and this entry. Their calls form one @@ -101,7 +128,14 @@ ranks issuing B-then-A against A-then-B deadlocked; this op waits for its peers **Call-order invariant.** Every rank of the group makes the same sequence of calls on one workspace — the same number, the `k`-th with the same op, `T`, `B` and `F` — across layers and decode steps, eager calls and graph replays -alike; and on one stream the same order of calls across workspaces. +alike; and on one stream the same order of calls across workspaces. The invariant is on the order of calls, not on +their timing: a rank may enqueue any number of calls ahead of its peers, since each call waits on the device for its +peers' words of that call only, and every rank's earliest pending call can always complete. Certified: three rounds +of 52 calls on one workspace — the three MNNVL ops (an all-gather in every layer), both all-reduce paths, `T` = 8, 2, +64, 16, 1, 32, 7, two layers each, the second layer's prefix sum and residual taken from the first layer's outputs on +the device — enqueued by every rank with no host synchronization between calls, a random rank 20 ms late before it +starts enqueueing and another pausing 20 ms halfway; then every result against the reference and the flags against +the model. **What a later launch reads.** `buffer_flags`, which every call leaves as: current = its own buffer plus one, mod 3; dirty = its own buffer; bytes per buffer unchanged; dirty stage count 1; bytes to clear `(T x W x (2B + 4F), 0, 0, @@ -139,9 +173,11 @@ depend on it. ## Preconditions -- `input` fp32, contiguous, 2-D, 16-byte aligned; `B` a multiple of 8 and `F` a multiple of 4; `T` at least 1. - Otherwise the op raises `RuntimeError` on every rank before it touches the workspace (certified: `B` = 12, `F` = 2, - a bf16 input, `T` = 0; the flags do not move and the next call is correct). +- `input` fp32, contiguous, 2-D, 16-byte aligned; `B` a multiple of 8 and `F` a multiple of 4; `T` at least 1; the + op's `world_size` equal to the workspace's rank count (the wrapper passes it). Otherwise the op raises + `RuntimeError` on every rank before it touches the workspace (certified: `B` = 12, `F` = 2, a bf16 input, `T` = 0, + and, calling the op directly, `world_size` one more than the workspace's; the flags do not move and the next call + is correct). - `required_buffer_bytes(T, B, F, W) <= workspace.buffer_bytes` (*State*). - Every rank calls with the same `T`, `B` and `F`; the call order is the *State* invariant. - `workspace` was created before any capture. Calls may be captured: certified with a captured step of five calls on @@ -158,16 +194,12 @@ depend on it. and exact; this op's outputs are compared bit for bit. The all-reduce and attention-residual calls of its sequences are checked as in their own matrices (sums bit for bit, normed outputs within a tolerance). - State and test design: a typed state object built by an explicit, collective, eager `create()`; a test that drives - call sequences on real state (layers x steps, capture + replay, two objects interleaved) plus a negative control; - every written buffer named in the schema (the op falls short there, see the gaps below); the matrix takes + call sequences on real state (layers x steps, calls queued without host synchronization, capture + replay, two + objects interleaved) plus a negative control; every written buffer named in the schema; the matrix takes `--world-size` and `--launcher` (`mpirun` on one node, `srun` across trays) and CI runs it at 4 ranks on one GB200 tray; one `MnnvlWorkspace` shared by every MNNVL entry of the TP group. The 16-rank receipt is pending. - Not exercised: `B` = 0 or `F` = 0 (the op accepts both), denormal and non-finite values in the bf16 columns, an - accepted call of more than 64 tokens. -- Gaps (the op is unchanged by this entry): the schema marks `comm_buffer` mutable `(a!)` but not `buffer_flags`, - which every call advances; the op has no `register_fake` - (`tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py` registers one for the other two MNNVL ops), so fake-tensor - tracing, e.g. `torch.compile`, cannot run it. + accepted call of more than 64 tokens, the fake implementation. - In the model today the call is `MNNVLAllReduce.allgather_split(input, bf16_columns)` on `MNNVLAllReduce`'s workspace (a dict keyed by `Mapping`, grown to the call's footprint by the first eager call that needs more). This entry takes the explicit object instead, sized at construction. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md index 75dafff24a0c..279d61d57fa6 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md @@ -34,6 +34,8 @@ call leaves in the workspace's flags shows (*State*). - One-shot: one kernel. Each rank writes its rows into every rank's buffer through the multicast mapping, waits for all `W` rows of every token in its own copy and sums them; the residual add and the RMSNorm run in the same kernel. + It releases its programmatic dependents as soon as its own grid-dependency wait returns, before it waits for its + peers (see *Notes* for this change). - Two-shot: each rank writes token `t`'s row into rank `t mod W`'s buffer, which sums the `W` rows of its tokens and writes the bf16 sums into every rank's buffer through the multicast mapping. Plain, the same kernel then waits for every token's sum and copies it out; fused, a second kernel does that and adds the residual and normalizes. @@ -84,9 +86,10 @@ def required_buffer_bytes( | `eps` | `None`, or a float with `residual`; 1e-5 certified | Python float | — | — | | returns | plain: the sum `[T, H]`; fused: `(normed, updated)`, each `[T, H]` | bf16 | contiguous, newly allocated | = `input.device` | -`input`, `residual` and `norm_weight` are read only. `required_buffer_bytes` is the space one Lamport buffer must -have for the call: `T x H x W x 2` one-shot, `2 x ceil(T / W) x W x H x 2` two-shot (two stages); certified equal to -the space the call's stages take, every call of the shape grid. +`input`, `residual` and `norm_weight` are read only. The op writes the workspace's buffers and flag words; its +schema declares both mutable (`comm_buffer` `Tensor(a!)`, `buffer_flags` `Tensor(b!)`). `required_buffer_bytes` is +the space one Lamport buffer must have for the call: `T x H x W x 2` one-shot, `2 x ceil(T / W) x W x H x 2` two-shot +(two stages); certified equal to the space the call's stages take, every call of the shape grid. ## State @@ -108,10 +111,21 @@ calls of this op (arithmetic, not a test). The test's buffer is the one-shot foo `W` = 4. **Who creates it, and when.** The target, in `post_load_weights`, with -`MnnvlWorkspace.create(mapping, buffer_bytes, fabric_handle=None)`: collective over the TP group, eager, every word -and flag armed before any rank returns (see `mnnvl_allreduce_attn_res.md`). It refuses CUDA-graph capture: certified -with every rank capturing, each raising `RuntimeError`. The check runs before any communication, so a rank that is -not capturing while its peers are would go on into the communicator split and wait for them (code). +`MnnvlWorkspace.create(mapping, buffer_bytes, fabric_handle=None)` (see `mnnvl_allreduce_attn_res.md`): + +- collective over the TP group: every rank calls it at the same point; +- failure model: + - before allocating, the ranks agree that each of them can (not capturing, a valid `buffer_bytes`, the three + buffers within that rank's free device memory). If one cannot, every rank raises `RuntimeError` and none + allocates (certified: one rank inside a CUDA-graph capture while the others are not, and then every rank + capturing; each time every rank raises, the capturing ranks' message naming the capture, and the workspaces in + use are untouched; a create right after, eager on every rank, returns an armed workspace whose first call is + correct); + - a failure that returns from the allocation is agreed the same way; + - a rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange; this + is not turned into an error on the other ranks; +- eager: it allocates and exchanges handles, so it refuses to run under CUDA-graph capture (every rank raises); +- it arms every buffer word and the flags before any rank returns. **Which ops may share one object.** Every MNNVL op of the group takes the same `comm_buffer` / `buffer_flags`: `comm/mnnvl_allreduce_attn_res`, this entry on either path, and `comm/mnnvl_allgather_split`. Their calls form one @@ -130,7 +144,13 @@ the same way). number, the `k`-th with the same op, `T`, `H`, fusion and path — across layers and decode steps, eager calls and graph replays alike; and on one stream the same order of calls across workspaces. The path is part of the sequence: the two paths write and wait for different words, so the ranks' `one_shot_max_bytes` must pick the same path (they do -when they pass the same value). +when they pass the same value). The invariant is on the order of calls, not on their timing: a rank may enqueue any +number of calls ahead of its peers, since each call waits on the device for its peers' words of that call only, and +every rank's earliest pending call can always complete. Certified: three rounds of 52 calls on one workspace — the +three MNNVL ops, both paths, `T` = 8, 2, 64, 16, 1, 32, 7, two layers each, the second layer's prefix sum and residual +taken from the first layer's outputs on the device — enqueued by every rank with no host synchronization between +calls, a random rank 20 ms late before it starts enqueueing and another pausing 20 ms halfway; then every result +against the reference and the flags against the model. **What a later launch reads.** `buffer_flags`, which every call leaves as: current = its own buffer plus one, mod 3; dirty = its own buffer; bytes per buffer unchanged; dirty stage count 1 (one-shot) or 2 (two-shot); bytes to clear @@ -170,9 +190,9 @@ None besides `workspace` and `one_shot_max_bytes`, both explicit. The op keeps n (precompiled kernels). It finds the multicast mapping by looking `comm_buffer`'s address up in a process registry of multicast buffers, which the workspace's handle keeps registered. `TRTLLM_ENABLE_PDL` (read once per process, default on at SM 90 and newer) launches the kernels as programmatic dependents: the one-shot kernel releases its own -dependents as it starts, the two-shot one after its scatter; consumers of the outputs wait for the grid (the kernels' -statement); results do not depend on it. The launch shape (CTAs per token, cluster size) follows `T`, `H` and the -device's SM count, and in the fused form it sets the order in which the squares are summed. +dependents right after its grid-dependency wait, the two-shot one after its scatter; consumers of the outputs wait +for the grid (the kernels' statement); results do not depend on it. The launch shape (CTAs per token, cluster size) +follows `T`, `H` and the device's SM count, and in the fused form it sets the order in which the squares are summed. ## Preconditions @@ -202,11 +222,18 @@ device's SM count, and in the fused form it sets the order in which the squares `updated` are compared bit for bit, `normed` against the fp32 RMSNorm of `updated` within 1e-2 of its largest magnitude (the bf16 squares cost at most 2^-9 on `rcp`, the bf16 output 2^-8). - State and test design: a typed state object built by an explicit, collective, eager `create()`; a test that drives - call sequences on real state (layers x steps, capture + replay, two objects interleaved) plus a negative control; - every written buffer named in the schema (the op falls short there, see the gaps below); the matrix takes + call sequences on real state (layers x steps, calls queued without host synchronization, capture + replay, two + objects interleaved) plus a negative control; every written buffer named in the schema; the matrix takes `--world-size` and `--launcher` (`mpirun` on one node, `srun` across trays) and CI runs it at 4 ranks on one GB200 tray; one caller-owned `MnnvlWorkspace` shared by every MNNVL entry of the TP group, and `one_shot_max_bytes` per call. +- The one-shot kernel differs from main's for every caller, main's `MNNVLAllReduce` included. It releases its + programmatic dependents right after its own grid-dependency wait, where main's released them after the reduction, + so a dependent kernel (a GEMV) can launch and stream its weights while this kernel waits for its peers; a + dependent still reads the output and the flags only after its own grid wait. Its Lamport reduction is the shared + `reduceOneshotLamport` of the attention-residual one-shot kernel: the code is moved, not changed, and the reduction + order is the same. Results unchanged (main's MNNVL test bodies pass on this build; the kernel's SASS changes). + EVIDENCE: - At `W` = 16 the one-shot kernel adds the ranks in two chunks of 8, a branch a 4-rank run never reaches. The 16-rank receipt is pending. - Kimi K3's calls (its decode path, not this test): the model sets every `MNNVLAllReduce` of the target, its @@ -217,13 +244,12 @@ device's SM count, and in the fused form it sets the order in which the squares `k3_sandwich_plain` does not take the call. At 4 MiB a `[T, 7168]` call goes one-shot up to `T` = 18 at `W` = 16 (73 at `W` = 4) and a `[T, 3584]` one up to 36 (146); at 1 MiB `[T, 7168]` up to 4 (18) and `[T, 3584]` up to 9 (36). -- Gaps (the op is unchanged by this entry): the schema marks `comm_buffer` mutable `(a!)` but not `buffer_flags`, - which every call advances; the op does not check that the call fits `comm_buffer` (the - attention-residual and all-gather ops do), so a direct op call over one buffer writes past it — the wrapper's - `required_buffer_bytes` check is the guard; `MnnvlWorkspace.create` accepts any multiple of 16 bytes, but the - two-shot broadcast stage starts at `buffer_bytes / 2` and is accessed in 16-byte vectors, so a two-shot call needs - `buffer_bytes` to be a multiple of 32 (code; every buffer in the test is); the schema's default - `one_shot_max_bytes=1048576` applies to a direct op call (the wrapper always passes one). +- Gaps: the op does not check that the call fits `comm_buffer` (the attention-residual and all-gather ops do), so a + direct op call over one buffer writes past it — the wrapper's `required_buffer_bytes` check is the guard; + `MnnvlWorkspace.create` accepts any multiple of 16 bytes, but the two-shot broadcast stage starts at + `buffer_bytes / 2` and is accessed in 16-byte vectors, so a two-shot call needs `buffer_bytes` to be a multiple of + 32 (code; every buffer in the test is); the schema's default `one_shot_max_bytes=1048576` applies to a direct op + call (the wrapper always passes one). - In the model today the workspace is `MNNVLAllReduce`'s (a dict keyed by `Mapping`, grown on demand by the first eager call that needs more, in 8 MiB steps) and the one-shot ceiling is a module attribute with a per-call override. This entry takes both explicitly: the workspace sized at construction, the ceiling per call. diff --git a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py index 6dee0eb91c1f..12babfbd03fc 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py @@ -6,10 +6,11 @@ ``comm/mnnvl_fusion_allreduce`` on either path and ``comm/mnnvl_allreduce_attn_res``). So beyond single calls (every certified split at every certified token count against a bit-exact reference) this drives call *sequences*: decode steps whose token count dips and grows back with a random rank late at every call; two workspaces interleaved; the -three MNNVL ops interleaved on one workspace; CUDA-graph capture and replay mixed with eager calls; and a negative -control in which one rank swaps two calls and every rank gets a wrong answer without an error. After every eager call -the workspace's ``buffer_flags`` are compared with this file's model of the rotation (``Rotation``): one turn per call -whatever the op, one stage, the bytes the call wrote. +three MNNVL ops interleaved on one workspace; the same mix queued with no host synchronization between calls while +ranks run ahead of one another; CUDA-graph capture and replay mixed with eager calls; ``MnnvlWorkspace.create``'s +failure model; and a negative control in which one rank swaps two calls and every rank gets a wrong answer without +an error. After every eager call the workspace's ``buffer_flags`` are compared with this file's model of the rotation +(``Rotation``): one turn per call whatever the op, one stage, the bytes the call wrote. CUDA_VISIBLE_DEVICES=0,1,2,3 python _mnnvl_allgather_split_op_matrix.py [--world-size 4] srun -N 4 --ntasks-per-node 4 --mpi=pmix python _mnnvl_allgather_split_op_matrix.py --launcher srun \ @@ -64,16 +65,21 @@ LAYERS = 6 DIP_STEPS = (8, 8, 8, 2, 7, 8, 1, 1, 64, 3, 32, 8, 16, 1, 64, 8) SHARED_STEPS = (8, 2, 16, 64, 1, 7, 32, 8, 3, 16) +QUEUED_STEPS = (8, 2, 64, 16, 1, 32, 7) # one round of the queued check: 52 calls +QUEUED_ROUNDS = 3 +QUEUE_LATE_S = 0.02 INTERLEAVED_TOKENS = (3, 8, 1, 64, 5, 16, 2) R = None allgather_split = None required_buffer_bytes = None fusion_allreduce = None +allreduce_attn_res = None MnnvlWorkspace = None BUFFER_BYTES = None WS_A = None WS_B = None +WS_C = None # created by the failure-model check after the refused creates ROT = {} STATS = {"normed_err": 0.0, "attn_res_err": 0.0} @@ -250,6 +256,11 @@ def fresh(self, seed): seed, self.tokens, self.hidden, residual=True if self.fused else None, path=self.path ) + def chain_from(self, updated: torch.Tensor) -> None: + """Take an earlier fused call's ``updated`` as this fused call's residual (as a layer's residual stream).""" + assert self.fused + self.residual = updated + def run(self, ws, x=None, residual=None): x = self.inputs[R.rank] if x is None else x if not self.fused: @@ -302,8 +313,8 @@ def refill(self, fresh, static) -> None: class AttnRes: - """One comm/mnnvl_allreduce_attn_res call (G1's entry, called through its op on the workspace's comm_buffer and - buffer_flags, as that entry's wrapper does) and its reference. ``prefix``: a tensor (chained) or True (drawn).""" + """One comm/mnnvl_allreduce_attn_res call (the entry's wrapper, on the same workspace) and its reference. + ``prefix``: a tensor (chained) or True (drawn).""" def __init__(self, seed, tokens, snapshots, prefix=True): g = torch.Generator(device="cuda").manual_seed(seed) @@ -317,8 +328,12 @@ def __init__(self, seed, tokens, snapshots, prefix=True): self.rms_w = (1.0 + 0.1 * torch.randn(H_MODEL, generator=g, device="cuda")).bfloat16() self.out_w = (1.0 + 0.1 * torch.randn(H_MODEL, generator=g, device="cuda")).bfloat16() + def chain_from(self, updated: torch.Tensor) -> None: + """Take an earlier call's ``updated`` as this call's prefix sum (as the next layer does).""" + self.prefix = updated + def run(self, ws): - normed, updated = torch.ops.trtllm.mnnvl_allreduce_attn_res( + return allreduce_attn_res( self.inputs[R.rank], self.prefix, self.block, @@ -327,10 +342,8 @@ def run(self, ws): self.out_w, EPS, EPS, - ws.comm_buffer(torch.bfloat16), - ws.buffer_flags, + ws, ) - return normed, updated def ref(self): updated = (sum(x.float() for x in self.inputs) + self.prefix.float()).bfloat16() @@ -420,7 +433,8 @@ def check_single_calls() -> None: def check_unsupported_calls_raise_on_every_rank() -> None: """Refused on every rank before the workspace is touched (its flags do not move), and the next call is correct: more rows than one Lamport buffer holds (the wrapper's ValueError); bf16 columns not a multiple of 8, remaining - columns not a multiple of 4, a bf16 input, no rows (the op's RuntimeError).""" + columns not a multiple of 4, a bf16 input, no rows, and -- through the op itself, as the wrapper always passes + the workspace's -- a world_size other than the workspace's rank count (the op's RuntimeError).""" b, f = k3_split() t = BUFFER_BYTES // (R.world * (2 * b + 4 * f)) + 1 over = AG(2000, t, b, f) @@ -431,6 +445,13 @@ def check_unsupported_calls_raise_on_every_rank() -> None: expect_refusal(lambda: allgather_split(narrow, 8, WS_A), RuntimeError, "2 fp32 columns") expect_refusal(lambda: allgather_split(rows.bfloat16(), 8, WS_A), RuntimeError, "bf16 rows") expect_refusal(lambda: allgather_split(rows[:0], 8, WS_A), RuntimeError, "no rows") + expect_refusal( + lambda: torch.ops.trtllm.mnnvl_allgather_split( + rows, 8, R.world + 1, WS_A.comm_buffer(torch.bfloat16), WS_A.buffer_flags + ), + RuntimeError, + f"world_size {R.world + 1}", + ) call_and_check(AG(2001, 8, b, f), WS_A, "after the refused calls") @@ -488,6 +509,60 @@ def check_one_workspace_three_ops() -> None: residual = call_late(fused, WS_A, f"{where} fused", late)[1] +def queued_plan(seed: int): + """One round of the queued check, in issue order: per step of QUEUED_STEPS, two layers of Kimi K3's calls -- a + step of at most 16 tokens: the attention-residual all-reduce, the head all-gather, the latent all-reduce and the + fused all-reduce, the two all-reduces on opposite paths, alternating from layer to layer; a wide step: the wide + all-reduce [T, 7168], the all-gather and the latent all-reduce at Kimi K3's ceiling. The second layer's prefix sum + and residual are the first layer's outputs. Returns ``[(call, index of the call it chains from, or None)]``.""" + b, f = k3_split() + plan = [] + for i, t in enumerate(QUEUED_STEPS): + decode = t <= ATTN_RES_MAX_TOKENS + prev_attn = prev_fused = None + for layer in range(2): + s = seed + 100 * i + 10 * layer + one, two = ("one", "two") if (i + layer) % 2 == 0 else ("two", "one") + if decode: + plan.append((AttnRes(s, t, (0, 2, 5)[(i + layer) % 3]), prev_attn)) + prev_attn = len(plan) - 1 + else: + plan.append((AR(s, t, H_MODEL), None)) + plan.append((AG(s + 1, t, b, f), None)) + plan.append((AR(s + 2, t, LATENT, path=one if decode else "k3"), None)) + if decode: + plan.append((AR(s + 3, t, H_MODEL, residual=True, path=two), prev_fused)) + prev_fused = len(plan) - 1 + return plan + + +def check_queued_calls_without_host_sync() -> None: + """Ranks running ahead of one another: in each round every rank enqueues the same 52 calls on one workspace (the + three ops, both all-reduce paths, Kimi K3's step sizes dipping and growing back, a layer's prefix sum and residual + taken from the previous layer's outputs on the device) with no host synchronization between them -- no barrier, + no .item(), no synchronize -- a random rank sleeping 20 ms before it starts enqueueing and another pausing 20 ms + halfway, so the other ranks' queues run deep while it is idle. Then every result is checked, and the flags. This + cannot deadlock: a call waits only for its peers' words of the same call, every rank enqueues every call, and a + stream runs its calls in order, so the ranks' earliest pending call always completes.""" + late = random.Random(13) + for rnd in range(QUEUED_ROUNDS): + plan = queued_plan(60_000 + 1000 * rnd) + sleeper, pauser = late.randrange(R.world), late.randrange(R.world) + R.barrier() + R.late(sleeper, QUEUE_LATE_S) + outs = [] + for k, (call, src) in enumerate(plan): + if k == len(plan) // 2: + R.late(pauser, QUEUE_LATE_S) + if src is not None: + call.chain_from(outs[src][1]) + outs.append(call.run(WS_A)) + torch.cuda.synchronize() + for k, ((call, _), got) in enumerate(zip(plan, outs)): + call.verify(got, f"queued round {rnd} call {k}") + rot(WS_A).advance(plan[-1][0].record(), f"queued round {rnd}", calls=len(plan)) + + def check_graph_capture_and_replay() -> None: """A captured step of five calls on WS_B, the MoE layers of a decode step: the head all-gather (T 8), the routed-latent all-reduce one-shot, the head all-gather again, a [32, 3584] all-reduce sent two-shot and the head @@ -537,20 +612,50 @@ def step(): del graph -def check_create_under_capture_raises() -> None: - """MnnvlWorkspace.create allocates and exchanges handles, so it refuses CUDA-graph capture: with every rank - capturing, it raises RuntimeError on every rank (before any communication), and the next call on a workspace in - use is correct.""" - graph, stream = torch.cuda.CUDAGraph(), torch.cuda.Stream() - raised = False - with torch.cuda.graph(graph, stream=stream): - try: +def try_create(capturing: bool): + """Every rank calls MnnvlWorkspace.create; this one inside a CUDA-graph capture if ``capturing``. Returns the + RuntimeError's message, or None if it returned a workspace.""" + try: + if not capturing: MnnvlWorkspace.create(R.mapping, BUFFER_BYTES, fabric_handle=R.fabric) - except RuntimeError as exc: - raised = "before capture" in str(exc) - del graph - assert R.all_true(raised), "create() under capture did not raise on every rank" - call_and_check(AG(8500, 8, *k3_split()), WS_A, "after the refused create") + return None + graph, stream = torch.cuda.CUDAGraph(), torch.cuda.Stream() + try: + with torch.cuda.graph(graph, stream=stream): + MnnvlWorkspace.create(R.mapping, BUFFER_BYTES, fabric_handle=R.fabric) + finally: + del graph + return None + except RuntimeError as exc: + return str(exc) + + +def check_create_refuses_on_every_rank() -> None: + """MnnvlWorkspace.create's failure model: before allocating, the ranks agree that each of them can. With one + random rank inside a CUDA-graph capture and the others not, and then with every rank capturing, every rank raises + RuntimeError and none allocates; the capturing ranks' message names the capture. A create right after, eager on + every rank, returns an armed workspace whose first call is correct, and the workspaces in use are untouched.""" + global WS_C + capturer = random.Random(17).randrange(R.world) + b, f = k3_split() + for every in (False, True): + capturing = every or R.rank == capturer + message = try_create(capturing) + refused = message is not None and "not every rank can allocate" in message + named = not capturing or (refused and "before capture" in message) + assert R.all_true(refused and named), ( + f"{'every rank' if every else f'rank {capturer}'} capturing: not refused on every rank ({message})" + ) + rot(WS_A).unchanged("after the refused create") + WS_C = MnnvlWorkspace.create(R.mapping, BUFFER_BYTES, fabric_handle=R.fabric) + ROT[id(WS_C)] = Rotation(WS_C) + armed = WS_C.lamport.view(torch.int32) + assert bool((armed == torch.tensor(INT32_MIN, dtype=torch.int32, device="cuda")).all()), ( + "the workspace created after the refusals: every word -0.0" + ) + rot(WS_C).unchanged("created after the refusals") + call_and_check(AG(8500, 8, b, f), WS_C, "first call on the new workspace") + call_and_check(AG(8501, 8, b, f), WS_A, "after the refusals") def check_wrong_call_order_is_detected() -> None: @@ -588,21 +693,25 @@ def check_wrong_call_order_is_detected() -> None: check_dip_and_regrow_sequence, check_two_workspaces_interleaved, check_one_workspace_three_ops, + check_queued_calls_without_host_sync, check_graph_capture_and_replay, - check_create_under_capture_raises, + check_create_refuses_on_every_rank, # Stays last: it deliberately disagrees on call order. check_wrong_call_order_is_detected, ] def _run_one_rank(args) -> int: - global R, allgather_split, required_buffer_bytes, fusion_allreduce, MnnvlWorkspace - global BUFFER_BYTES, WS_A, WS_B + global R, allgather_split, required_buffer_bytes, fusion_allreduce, allreduce_attn_res + global MnnvlWorkspace, BUFFER_BYTES, WS_A, WS_B R = ls.Rank(args) assert R.world in (2, 4, 8, 16), f"K3's shapes shard over 2, 4, 8 or 16 ranks, not {R.world}" from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm import ( mnnvl_allgather_split as module, ) + from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm import ( + mnnvl_allreduce_attn_res as attn_res_module, + ) from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm import ( mnnvl_fusion_allreduce as reduce_module, ) @@ -610,6 +719,7 @@ def _run_one_rank(args) -> int: allgather_split = module.mnnvl_allgather_split required_buffer_bytes = module.required_buffer_bytes fusion_allreduce = reduce_module.mnnvl_fusion_allreduce + allreduce_attn_res = attn_res_module.mnnvl_allreduce_attn_res MnnvlWorkspace = module.MnnvlWorkspace BUFFER_BYTES = buffer_bytes(R.world) with torch.inference_mode(): diff --git a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_fusion_allreduce_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_fusion_allreduce_op_matrix.py index a72cd25c6ae8..14114f8c348c 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_fusion_allreduce_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_fusion_allreduce_op_matrix.py @@ -7,10 +7,12 @@ So beyond single calls (every certified shape sent one-shot and two-shot back to back, the path chosen per call by ``one_shot_max_bytes`` at the exact boundary) this drives call *sequences*: decode steps whose token count dips and grows back with Kimi K3's one-shot ceilings flipping the path inside the sequence, a random rank late at every call; -two workspaces interleaved; the three MNNVL ops interleaved on one workspace; CUDA-graph capture and replay mixed with -eager calls; and a negative control in which one rank swaps two calls and every rank gets a wrong answer without an -error. After every eager call the workspace's ``buffer_flags`` are compared with this file's model of the rotation -(``Rotation``): one turn per call whatever the op, the path the call took, the bytes it wrote. +two workspaces interleaved; the three MNNVL ops interleaved on one workspace; the same mix queued with no host +synchronization between calls while ranks run ahead of one another; CUDA-graph capture and replay mixed with eager +calls; ``MnnvlWorkspace.create``'s failure model; and a negative control in which one rank swaps two calls and every +rank gets a wrong answer without an error. After every eager call the workspace's ``buffer_flags`` are compared with +this file's model of the rotation (``Rotation``): one turn per call whatever the op, the path the call took, the +bytes it wrote. CUDA_VISIBLE_DEVICES=0,1,2,3 python _mnnvl_fusion_allreduce_op_matrix.py [--world-size 4] srun -N 4 --ntasks-per-node 4 --mpi=pmix python _mnnvl_fusion_allreduce_op_matrix.py --launcher srun \ @@ -63,6 +65,9 @@ LAYERS = 8 DIP_STEPS = (8, 8, 8, 2, 7, 8, 1, 1, 64, 3, 32, 8, 16, 1, 64, 8) SHARED_STEPS = (8, 2, 16, 64, 1, 7, 32, 8, 3, 16) +QUEUED_STEPS = (8, 2, 64, 16, 1, 32, 7) # one round of the queued check: 52 calls +QUEUED_ROUNDS = 3 +QUEUE_LATE_S = 0.02 INTERLEAVED = ( # (T, H, fused, path) (3, H_LATENT, False, "one"), (8, H_MODEL, True, "two"), @@ -77,10 +82,12 @@ fusion_allreduce = None required_buffer_bytes = None allgather_split = None +allreduce_attn_res = None MnnvlWorkspace = None BUFFER_BYTES = None WS_A = None WS_B = None +WS_C = None # created by the failure-model check after the refused creates ROT = {} STATS = {"normed_err": 0.0, "normed_paths": 0.0, "attn_res_err": 0.0} @@ -189,6 +196,11 @@ def fresh(self, seed): seed, self.tokens, self.hidden, residual=True if self.fused else None, path=self.path ) + def chain_from(self, updated: torch.Tensor) -> None: + """Take an earlier fused call's ``updated`` as this fused call's residual (as a layer's residual stream).""" + assert self.fused + self.residual = updated + def run(self, ws, x=None, residual=None): x = self.inputs[R.rank] if x is None else x if not self.fused: @@ -304,8 +316,8 @@ def refill(self, fresh, static) -> None: class AttnRes: - """One comm/mnnvl_allreduce_attn_res call (G1's entry, called through its op on the workspace's comm_buffer and - buffer_flags, as that entry's wrapper does) and its reference. ``prefix``: a tensor (chained) or True (drawn).""" + """One comm/mnnvl_allreduce_attn_res call (the entry's wrapper, on the same workspace) and its reference. + ``prefix``: a tensor (chained) or True (drawn).""" def __init__(self, seed, tokens, snapshots, prefix=True): g = torch.Generator(device="cuda").manual_seed(seed) @@ -322,8 +334,12 @@ def __init__(self, seed, tokens, snapshots, prefix=True): def fresh(self, seed): return AttnRes(seed, self.tokens, self.snapshots) + def chain_from(self, updated: torch.Tensor) -> None: + """Take an earlier call's ``updated`` as this call's prefix sum (as the next layer does).""" + self.prefix = updated + def run(self, ws, x=None, prefix=None, block=None): - normed, updated = torch.ops.trtllm.mnnvl_allreduce_attn_res( + return allreduce_attn_res( self.inputs[R.rank] if x is None else x, self.prefix if prefix is None else prefix, self.block if block is None else block, @@ -332,10 +348,8 @@ def run(self, ws, x=None, prefix=None, block=None): self.out_w, EPS, EPS, - ws.comm_buffer(torch.bfloat16), - ws.buffer_flags, + ws, ) - return normed, updated def ref(self): updated = (sum(x.float() for x in self.inputs) + self.prefix.float()).bfloat16() @@ -545,6 +559,60 @@ def check_one_workspace_three_ops() -> None: residual = call_late(fused, WS_A, f"{where} fused", late)[1] +def queued_plan(seed: int): + """One round of the queued check, in issue order: per step of QUEUED_STEPS, two layers of Kimi K3's calls -- a + step of at most 16 tokens: the attention-residual all-reduce, the head all-gather, the latent all-reduce and the + fused all-reduce, the two all-reduces on opposite paths, alternating from layer to layer; a wide step: the wide + all-reduce [T, 7168], the all-gather and the latent all-reduce at Kimi K3's ceiling. The second layer's prefix sum + and residual are the first layer's outputs. Returns ``[(call, index of the call it chains from, or None)]``.""" + b, f = k3_split() + plan = [] + for i, t in enumerate(QUEUED_STEPS): + decode = t <= ATTN_RES_MAX_TOKENS + prev_attn = prev_fused = None + for layer in range(2): + s = seed + 100 * i + 10 * layer + one, two = ("one", "two") if (i + layer) % 2 == 0 else ("two", "one") + if decode: + plan.append((AttnRes(s, t, (0, 2, 5)[(i + layer) % 3]), prev_attn)) + prev_attn = len(plan) - 1 + else: + plan.append((AR(s, t, H_MODEL), None)) + plan.append((AG(s + 1, t, b, f), None)) + plan.append((AR(s + 2, t, H_LATENT, path=one if decode else "k3"), None)) + if decode: + plan.append((AR(s + 3, t, H_MODEL, residual=True, path=two), prev_fused)) + prev_fused = len(plan) - 1 + return plan + + +def check_queued_calls_without_host_sync() -> None: + """Ranks running ahead of one another: in each round every rank enqueues the same 52 calls on one workspace (the + three ops, both all-reduce paths, Kimi K3's step sizes dipping and growing back, a layer's prefix sum and residual + taken from the previous layer's outputs on the device) with no host synchronization between them -- no barrier, + no .item(), no synchronize -- a random rank sleeping 20 ms before it starts enqueueing and another pausing 20 ms + halfway, so the other ranks' queues run deep while it is idle. Then every result is checked, and the flags. This + cannot deadlock: a call waits only for its peers' words of the same call, every rank enqueues every call, and a + stream runs its calls in order, so the ranks' earliest pending call always completes.""" + late = random.Random(13) + for rnd in range(QUEUED_ROUNDS): + plan = queued_plan(60_000 + 1000 * rnd) + sleeper, pauser = late.randrange(R.world), late.randrange(R.world) + R.barrier() + R.late(sleeper, QUEUE_LATE_S) + outs = [] + for k, (call, src) in enumerate(plan): + if k == len(plan) // 2: + R.late(pauser, QUEUE_LATE_S) + if src is not None: + call.chain_from(outs[src][1]) + outs.append(call.run(WS_A)) + torch.cuda.synchronize() + for k, ((call, _), got) in enumerate(zip(plan, outs)): + call.verify(got, f"queued round {rnd} call {k}") + rot(WS_A).advance(plan[-1][0].record(), f"queued round {rnd}", calls=len(plan)) + + def check_graph_capture_and_replay() -> None: """A captured step of six calls on WS_B: the attention-residual all-reduce, the plain and fused all-reduces one-shot (T 8, Kimi K3's ceiling), the head all-gather, a plain [32, 7168] and a fused [16, 7168] sent two-shot @@ -596,20 +664,49 @@ def step(): del graph -def check_create_under_capture_raises() -> None: - """MnnvlWorkspace.create allocates and exchanges handles, so it refuses CUDA-graph capture: with every rank - capturing, it raises RuntimeError on every rank (before any communication), and the next call on a workspace in - use is correct.""" - graph, stream = torch.cuda.CUDAGraph(), torch.cuda.Stream() - raised = False - with torch.cuda.graph(graph, stream=stream): - try: +def try_create(capturing: bool): + """Every rank calls MnnvlWorkspace.create; this one inside a CUDA-graph capture if ``capturing``. Returns the + RuntimeError's message, or None if it returned a workspace.""" + try: + if not capturing: MnnvlWorkspace.create(R.mapping, BUFFER_BYTES, fabric_handle=R.fabric) - except RuntimeError as exc: - raised = "before capture" in str(exc) - del graph - assert R.all_true(raised), "create() under capture did not raise on every rank" - call_and_check(AR(8500, 8, H_MODEL, residual=True), WS_A, "after the refused create") + return None + graph, stream = torch.cuda.CUDAGraph(), torch.cuda.Stream() + try: + with torch.cuda.graph(graph, stream=stream): + MnnvlWorkspace.create(R.mapping, BUFFER_BYTES, fabric_handle=R.fabric) + finally: + del graph + return None + except RuntimeError as exc: + return str(exc) + + +def check_create_refuses_on_every_rank() -> None: + """MnnvlWorkspace.create's failure model: before allocating, the ranks agree that each of them can. With one + random rank inside a CUDA-graph capture and the others not, and then with every rank capturing, every rank raises + RuntimeError and none allocates; the capturing ranks' message names the capture. A create right after, eager on + every rank, returns an armed workspace whose first call is correct, and the workspaces in use are untouched.""" + global WS_C + capturer = random.Random(17).randrange(R.world) + for every in (False, True): + capturing = every or R.rank == capturer + message = try_create(capturing) + refused = message is not None and "not every rank can allocate" in message + named = not capturing or (refused and "before capture" in message) + assert R.all_true(refused and named), ( + f"{'every rank' if every else f'rank {capturer}'} capturing: not refused on every rank ({message})" + ) + rot(WS_A).unchanged("after the refused create") + WS_C = MnnvlWorkspace.create(R.mapping, BUFFER_BYTES, fabric_handle=R.fabric) + ROT[id(WS_C)] = Rotation(WS_C) + armed = WS_C.lamport.view(torch.int32) + assert bool((armed == torch.tensor(INT32_MIN, dtype=torch.int32, device="cuda")).all()), ( + "the workspace created after the refusals: every word -0.0" + ) + rot(WS_C).unchanged("created after the refusals") + call_and_check(AR(8500, 8, H_MODEL, residual=True), WS_C, "first call on the new workspace") + call_and_check(AR(8501, 8, H_MODEL, residual=True, path="two"), WS_A, "after the refusals") def check_wrong_call_order_is_detected() -> None: @@ -654,21 +751,25 @@ def check_wrong_call_order_is_detected() -> None: check_dip_and_regrow_sequence, check_two_workspaces_interleaved, check_one_workspace_three_ops, + check_queued_calls_without_host_sync, check_graph_capture_and_replay, - check_create_under_capture_raises, + check_create_refuses_on_every_rank, # Stays last: it deliberately disagrees on call order. check_wrong_call_order_is_detected, ] def _run_one_rank(args) -> int: - global R, fusion_allreduce, required_buffer_bytes, allgather_split, MnnvlWorkspace - global BUFFER_BYTES, WS_A, WS_B + global R, fusion_allreduce, required_buffer_bytes, allgather_split, allreduce_attn_res + global MnnvlWorkspace, BUFFER_BYTES, WS_A, WS_B R = ls.Rank(args) assert R.world in (2, 4, 8, 16), f"K3's shapes shard over 2, 4, 8 or 16 ranks, not {R.world}" from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm import ( mnnvl_allgather_split as gather_module, ) + from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm import ( + mnnvl_allreduce_attn_res as attn_res_module, + ) from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm import ( mnnvl_fusion_allreduce as module, ) @@ -676,6 +777,7 @@ def _run_one_rank(args) -> int: fusion_allreduce = module.mnnvl_fusion_allreduce required_buffer_bytes = module.required_buffer_bytes allgather_split = gather_module.mnnvl_allgather_split + allreduce_attn_res = attn_res_module.mnnvl_allreduce_attn_res MnnvlWorkspace = module.MnnvlWorkspace BUFFER_BYTES = buffer_bytes(R.world) with torch.inference_mode(): From e22949b8d7b301a0ce935701cae4d93181defa04 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 20:47:55 -0700 Subject: [PATCH 070/161] [None][test] Kimi K3 sandwich and latent exchange matrices: create() refused on every rank; ranks a step apart The workspaces' create() lets the ranks agree before allocating, so the capture checks now cover every rank capturing and one rank capturing while its peers call create() eagerly: every rank raises, none allocates (the sandwich matrices count the multicast allocations), no stream is left capturing, and the next call is correct. The old case of one rank calling create() alone is gone: the collective now waits for its peers there. The latent exchange matrix adds ranks running a whole step apart (a rank's GPU runs at most one call ahead of its peers, however far its host enqueues) and checks the swapped-calls control's results bit for bit. The contracts state the failure model. Signed-off-by: Vasanth Sabavat --- .../catalog/comm/k3_latent_reduce.md | 35 ++++-- .../catalog/comm/k3_sandwich_oproj.md | 19 ++- .../catalog/comm/k3_sandwich_plain.md | 24 +++- .../catalog/comm/k3_sandwich_tail.md | 21 +++- .../comm/_k3_latent_reduce_op_matrix.py | 115 ++++++++++++++---- .../modeling_v2/comm/_k3_sandwich_common.py | 83 +++++++++---- .../comm/_k3_sandwich_oproj_op_matrix.py | 8 +- .../comm/_k3_sandwich_plain_op_matrix.py | 5 +- .../comm/_k3_sandwich_tail_op_matrix.py | 5 +- 9 files changed, 242 insertions(+), 73 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md index 89663c308c43..4112de1901b7 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md @@ -81,12 +81,19 @@ of up to 8 tokens fits. **Who creates it, and when.** The target, in `post_load_weights`, with `K3LatentExchange.create(mapping, fabric_handle=None)`: -- collective over `mapping`'s TP group: every rank calls it at the same point; it returns on every rank or raises on - every rank (each rank's success is agreed before anyone proceeds); -- eager: it allocates and exchanges handles, so under CUDA-graph capture it raises `RuntimeError` before any - collective step (certified on every rank at once, and on one rank alone while its peers do not call it); -- it empties every word and zeroes `flags`, synchronizes, and returns only once every rank has done so (the success - agreement is the barrier), so no producer can push into a buffer before its rank has armed it; armed and sized on +- collective over `mapping`'s TP group: every rank calls it at the same point. Every rank first joins the split of + the TP group's communicator, so a rank that calls it while its peers do not waits for them there; +- failure model (`k3_fused_moe.op.create_mcast_state`, the same as `MnnvlWorkspace.create`'s): + - before allocating, the ranks agree that each of them can (not capturing, the buffer within that rank's free + device memory). If one cannot, every rank raises `RuntimeError` and none allocates (certified: every rank + capturing, and one rank capturing while its peers call it eagerly at the same point; every rank raises at that + agreement, the capturing rank's message naming the capture, and the next call is correct); + - a failure that returns from the allocation is agreed the same way; + - a rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange; this + is not turned into an error on the other ranks; +- eager: it allocates and exchanges handles, so it refuses to run under CUDA-graph capture (every rank raises); +- it empties every word and zeroes `flags`, and returns only once every rank has done so (the agreement after the + allocation is the barrier), so no producer can push into a buffer before its rank has armed it; armed and sized on every rank right after `create` is certified; - `fabric_handle`: share the memory by fabric handle (required across nodes) or POSIX file descriptor; default `mapping.is_multi_node()`. No environment variable is read. @@ -117,6 +124,19 @@ it calls `griddepcontrol.wait`, or it launches without programmatic dependent la the count before that reduce has advanced it and write into the half the reduce is still reading (the op's statement, `latent_op.py`). +**How far ranks run apart.** Only the reduce waits, so a rank's host may enqueue any number of calls ahead of its +peers, but its GPU runs at most one call ahead of the slowest peer: its push of a call lands only after its reduce of +the previous call has ended, which waited for every rank's push of that call, and each of those landed after that +rank's reduce of the call before. So a push lands only in a half that every rank has finished reading. Certified +with the test's pushes, which are copies issued after the previous reduce on the same stream: one rank enqueues a +whole step of 12 calls (`M` dipping and growing back) and waits until its first push has completed, while its peers +are held at a host barrier. Before they issue anything, every peer finds that push in its buffer and the other half +empty (the rank's second push is held behind its first reduce). Then the converse, every rank but one a whole step +ahead of it. Every result is correct and the exchange ends clean. The random late rank of the sequence below adds +per-call skew on top. Not certified here: a producer launched as a programmatic dependent starts before the previous +reduce has ended and stores only after its grid-dependency wait; that ordering, and the condition above that it +relies on, are not exercised by the copies. + **What a later launch reads.** Before its grid-dependency wait the op reads `flags[0]` (its half) and counts each CTA into `flags[2]`. After the wait it reads rows `0..M-1` of every rank's slot of that half, polling until none of the words is empty. At its very end CTA 0 waits until all `M x CTAs` CTAs have counted in, then sets `flags[2]` back @@ -144,7 +164,8 @@ push lands while the others' reduces poll), every call against the reference. - Swapped calls. Rank 0 makes two same-shaped calls in swapped order (it pushes its partial of the second call first). Nothing raises and nothing hangs, since the counts still agree, but every rank's two results are wrong: - each reduce sums rank 0's partial of the other call, and more than half of the elements differ on every rank. The + the `k`-th reduce sums every rank's `k`-th push, so each pairs rank 0's partial of one call with the others' of + the other (bit for bit), and more than half of the elements differ from the intended call's on every rank. The exchange is clean afterwards and a plain call right after is correct. - A token count that differs from the push. Every rank pushes 8 rows and rank 0 reduces 4: its 4 rows are right, but rows 4-7 of that half stay full in its buffer. Two calls later every rank pushes 4 rows into that half and rank diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md index 34216e3e2f02..dc0bee244c37 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md @@ -88,12 +88,19 @@ every word empty, every counter zero. **Who creates it, and when.** The target, in `post_load_weights`, with `K3SandwichWorkspace.create(mapping, fabric_handle=None)`: -- collective over `mapping`'s TP group: every rank calls it at the same point; it returns on every rank or raises on - every rank (each rank's success is agreed before any returns, which is also the barrier that keeps a peer from - pushing into a buffer its owner has not emptied yet); -- eager: it allocates and exchanges handles, so under CUDA-graph capture it raises `RuntimeError` before it enters - any collective — certified on every rank at once, and on one rank alone while its peers do not call it; -- every word emptied and every counter zeroed before any rank returns; +- collective over `mapping`'s TP group: every rank calls it at the same point; +- failure model: + - before allocating, the ranks agree that each of them can (not capturing, the buffer within that rank's free + device memory). If one cannot, every rank raises `RuntimeError` and none allocates (certified: one rank + capturing while its peers call it eagerly, every rank raises, the capturing rank naming the capture and its + peers another rank; no rank reaches the allocation, and the next call is correct); + - a failure that returns from the allocation is agreed the same way; + - a rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange; this + is not turned into an error on the other ranks; +- eager: it allocates and exchanges handles, so it refuses to run under CUDA-graph capture (certified: every rank + capturing, every rank raises; no stream is left capturing); +- it empties every word and zeroes every counter, and returns only once every rank has (the agreement after the + allocation), so no peer can push into a buffer its owner has not emptied yet; - `fabric_handle`: share the memory by fabric handle (required across nodes) or POSIX file descriptor; default `mapping.is_multi_node()`. No environment variable is read. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md index 4dcd66c12552..5049373fd590 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md @@ -84,11 +84,25 @@ CTAs, with the handle that owns the memory and the communicator. This op's parti target's sandwiches'. Certified after `create`: sized for the group, every word empty, every counter zero. **Who creates it, and when.** The target, in `post_load_weights`, with -`K3SandwichWorkspace.create(mapping, fabric_handle=None)`: collective over the TP group (it returns on every rank or -raises on every rank), eager, every word emptied and every counter zeroed before any rank returns. Under CUDA-graph -capture it raises `RuntimeError` before it enters any collective: certified on every rank at once, and on one rank -alone while its peers do not call it. Details in `k3_sandwich_oproj.md`. The drafter does not create its own: it -takes the target's. +`K3SandwichWorkspace.create(mapping, fabric_handle=None)`: + +- collective over `mapping`'s TP group: every rank calls it at the same point; +- failure model: + - before allocating, the ranks agree that each of them can (not capturing, the buffer within that rank's free + device memory). If one cannot, every rank raises `RuntimeError` and none allocates (certified: one rank + capturing while its peers call it eagerly, every rank raises, the capturing rank naming the capture and its + peers another rank; no rank reaches the allocation, and the next call is correct); + - a failure that returns from the allocation is agreed the same way; + - a rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange; this + is not turned into an error on the other ranks; +- eager: it allocates and exchanges handles, so it refuses to run under CUDA-graph capture (certified: every rank + capturing, every rank raises; no stream is left capturing); +- it empties every word and zeroes every counter, and returns only once every rank has (the agreement after the + allocation), so no peer can push into a buffer its owner has not emptied yet; +- `fabric_handle`: share the memory by fabric handle (required across nodes) or POSIX file descriptor; default + `mapping.is_multi_node()`. No environment variable is read. + +The drafter does not create its own workspace: its calls take the target's. **Which ops may share one object.** The three sandwich entries — `comm/k3_sandwich_oproj`, `comm/k3_sandwich_tail` and this one — take the same object, and in the model the drafter's calls run on the target's workspace of their TP diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md index a8eccdfb6fba..1c2aea006675 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md @@ -116,10 +116,23 @@ CTAs, with the handle that owns the memory and the communicator. This op's parti `k3_sandwich_oproj`'s. Certified after `create`: sized for the group, every word empty, every counter zero. **Who creates it, and when.** The target, in `post_load_weights`, with -`K3SandwichWorkspace.create(mapping, fabric_handle=None)`: collective over the TP group (it returns on every rank or -raises on every rank), eager, every word emptied and every counter zeroed before any rank returns. Under CUDA-graph -capture it raises `RuntimeError` before it enters any collective: certified on every rank at once, and on one rank -alone while its peers do not call it. Details in `k3_sandwich_oproj.md`. +`K3SandwichWorkspace.create(mapping, fabric_handle=None)`: + +- collective over `mapping`'s TP group: every rank calls it at the same point; +- failure model: + - before allocating, the ranks agree that each of them can (not capturing, the buffer within that rank's free + device memory). If one cannot, every rank raises `RuntimeError` and none allocates (certified: one rank + capturing while its peers call it eagerly, every rank raises, the capturing rank naming the capture and its + peers another rank; no rank reaches the allocation, and the next call is correct); + - a failure that returns from the allocation is agreed the same way; + - a rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange; this + is not turned into an error on the other ranks; +- eager: it allocates and exchanges handles, so it refuses to run under CUDA-graph capture (certified: every rank + capturing, every rank raises; no stream is left capturing); +- it empties every word and zeroes every counter, and returns only once every rank has (the agreement after the + allocation), so no peer can push into a buffer its owner has not emptied yet; +- `fabric_handle`: share the memory by fabric handle (required across nodes) or POSIX file descriptor; default + `mapping.is_multi_node()`. No environment variable is read. **Which ops may share one object.** The three sandwich entries — `comm/k3_sandwich_oproj`, this one and `comm/k3_sandwich_plain` — take the same object, and in the model the target's and the drafter's calls run on one diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_latent_reduce_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_k3_latent_reduce_op_matrix.py index 228606e27021..21c4e9527de3 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_k3_latent_reduce_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_latent_reduce_op_matrix.py @@ -6,9 +6,9 @@ partial into the exchange, and the op, on every rank, sums the ranks' rows. Its correctness depends on state that outlives a call (the call count in ``flags[0]``, whose parity picks the half every push and reduce use, and the words a reduce empties for the push two calls later), so beyond single calls this drives call *sequences*: layers x steps -with the token count dipping and growing back and a random rank late, two exchanges interleaved, CUDA-graph capture -and replay mixed with eager calls, the count across its int32 wrap, and two negative controls in which ranks break -the call order and get a wrong answer without an error. +with the token count dipping and growing back and a random rank late, ranks a whole step apart, two exchanges +interleaved, CUDA-graph capture and replay mixed with eager calls, the count across its int32 wrap, and two negative +controls in which ranks break the call order and get a wrong answer without an error. CUDA_VISIBLE_DEVICES=0,1,2,3 python _k3_latent_reduce_op_matrix.py [--world-size 4] srun -N 4 --ntasks-per-node 4 --mpi=pmix python _k3_latent_reduce_op_matrix.py --launcher srun --world-size 16 @@ -55,6 +55,8 @@ TOKENS = (1, 2, 3, 4, 5, 6, 7, 8) LAYERS = 12 # even: a captured step keeps the halves' parity (check_graph_capture_and_replay) DIP_STEPS = (8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8) +# One step of LAYERS calls whose M dips and grows back (check_ranks_run_a_step_apart). +RUN_AHEAD_TOKENS = (8, 3, 8, 1, 5, 8, 2, 8, 7, 4, 8, 6) GRAPH_TOKENS = (8, 3) # one captured step per batch size, as an engine keeps one graph per size EAGER_BETWEEN = (5, 1) # an even number of eager calls between two replays REPLAYS = 4 @@ -246,26 +248,45 @@ def check_exchange_is_armed_and_sized() -> None: def check_capture_refusals() -> None: - """Under CUDA-graph capture ``K3LatentExchange.create`` raises RuntimeError on every rank at once, and on one rank - alone while its peers do not call it: before any collective step, where that rank would wait for its peers. The - op's first call, which would compile the kernel, raises on every rank before it launches anything. The exchange - is untouched. Runs before every eager call of the op: the compile cache must still be cold.""" + """``K3LatentExchange.create`` is collective: every rank joins the TP group's communicator, and before allocating + the ranks agree that each of them can. With every rank capturing a CUDA graph, and with one rank capturing while + its peers call it eagerly at the same point, every rank raises RuntimeError, and a capturing rank's message names + the capture. Every rank raises at that agreement ("not every rank can allocate"; a failure after allocating reads + "allocation failed"), so nothing is allocated. The op's first call, which would compile the kernel, is refused per + rank: under capture it raises on every rank before it launches anything. The exchange is untouched and the next + call is correct. Runs before every eager call of the op: the compile cache must still be cold.""" # Imported by the op's first call; imported here so that nothing is imported inside the capture. import cutlass.cute.runtime # noqa: F401 from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import k3_latent_reduce # noqa: F401 - def refused() -> bool: - message = raised_under_capture(lambda: create_exchange(R.mapping, fabric_handle=R.fabric)) - return "outside CUDA-graph capture" in message + def create() -> None: + create_exchange(R.mapping, fabric_handle=R.fabric) - every = refused() + def refusal(capturing: bool) -> str: + """The message of the RuntimeError ``create`` raised on this rank ('' if it raised none).""" + if capturing: + return raised_under_capture(create) + try: + create() + except RuntimeError as exc: + return str(exc) + return "" + + for case, capturing in (("every rank", True), ("one rank", R.rank == R.world - 1)): + R.barrier() + message = refusal(capturing) + refused = "not every rank can allocate" in message + named = "outside CUDA-graph capture" in message or not capturing + assert R.all_true(refused and named), ( + f"{case} capturing: rank {R.rank} (capturing {capturing}) got {message!r}" + ) R.barrier() - alone = refused() if R.rank == R.world - 1 else True - assert R.all_true(every and alone), f"create: every rank {every}, one rank alone {alone}" first = raised_under_capture(lambda: entry(MAX_TOKENS, EX_A.state)) assert R.all_true("outside CUDA-graph capture first" in first), f"first call: {first!r}" assert_clean(EX_A, "after the refused calls") + call = Call(900, MAX_TOKENS) + verify(call, call.run(EX_A), "after the refused calls") def check_single_calls() -> None: @@ -343,6 +364,49 @@ def check_dip_and_regrow_sequence() -> None: assert_clean(EX_A, f"after step {i} (M {t})") +def check_ranks_run_a_step_apart() -> None: + """Calls queued across ranks. Rank 0 enqueues a whole step (LAYERS push + reduce pairs, M dipping and growing + back) and waits until its first push has completed, while its peers wait at a host barrier; only then do they + start. Before issuing anything, each peer finds rank 0's first push in its buffer and nothing of its second (the + other half is empty): a push is issued after its rank's previous reduce, which waits for every rank's push of its + call, so however far a rank's host runs ahead, its GPU runs at most one call ahead. Then the same with every rank + but the last a step ahead of the last. Every result is correct and the exchange ends clean.""" + ex = EX_A + sizes = RUN_AHEAD_TOKENS + for phase, ahead in enumerate((range(1), range(R.world - 1))): + calls = [Call(90_000 + 100 * phase + layer, t) for layer, t in enumerate(sizes)] + half = ex.calls & 1 + R.barrier() + outs = [] + seen = (True, True, True) + if R.rank in ahead: + pushed = torch.cuda.Event() + for layer, call in enumerate(calls): + ex.push(call.parts[R.rank]) + if layer == 0: + pushed.record() + outs.append(ex.reduce(call.tokens)) + # The event, not the stream: this rank's first reduce waits for the peers' pushes. + pushed.synchronize() + R.comm.Barrier() + else: + R.comm.Barrier() + rows = ex.state.uc.view(2, MAX_TOKENS, R.world, ROW_WORDS) + seen = ( + all(bool((rows[half, : sizes[0], a] != EMPTY_WORD).all()) for a in ahead), + bool((rows[half, :, R.rank] == EMPTY_WORD).all()), + bool((rows[half ^ 1] == EMPTY_WORD).all()), + ) + assert R.all_true(all(seen)), ( + f"phase {phase}: rank {R.rank} saw (first pushes landed, own slot empty, other half empty) = {seen}" + ) + if R.rank not in ahead: + outs = [call.run(ex) for call in calls] + for layer, (call, got) in enumerate(zip(calls, outs)): + verify(call, got, f"phase {phase} (ranks {list(ahead)} ahead), layer {layer}") + assert_clean(ex, f"after phase {phase}") + + def check_two_exchanges_interleaved() -> None: """Two exchanges are two counts and two buffers: calls alternate between them in an irregular pattern (A A B A B B ...), so the halves they use differ from call to call, a random rank late before each; every call is correct and @@ -426,18 +490,22 @@ def check_count_parity_across_the_int32_wrap() -> None: def check_swapped_calls_are_wrong() -> None: """Negative control: rank 0 makes two same-shaped calls in swapped order (it pushes its partial of the second call first). Every push is still followed by one reduce of its token count, so the counts agree and nothing raises or - hangs, but every rank's two results are wrong: each reduce sums rank 0's partial of the other call. The exchange - is clean afterwards and a plain call is correct.""" + hangs, but every rank's two results are wrong: each reduce sums rank 0's partial of the other call (exactly the + position-paired sums, bit for bit), and more than half of each result differs from the intended call's. The + exchange is clean afterwards and a plain call is correct.""" first, second = Call(70_000, MAX_TOKENS), Call(70_001, MAX_TOKENS) + order = (second, first) if R.rank == 0 else (first, second) R.barrier() - if R.rank == 0: - got_second, got_first = second.run(EX_A), first.run(EX_A) - else: - got_first, got_second = first.run(EX_A), second.run(EX_A) - torch.cuda.synchronize() - wrong = [differing(got, call.ref()) for call, got in ((first, got_first), (second, got_second))] - assert R.all_true(min(wrong) > 0.5), ( - f"rank {R.rank}: the swap went unnoticed, wrong shares {wrong}" + got = [call.run(EX_A) for call in order] + # The k-th reduce sums every rank's k-th push: rank 0's partial of one call with the others' of the other. + paired = [ + ordered_sum([second.parts[0]] + first.parts[1:], RANK_CHUNK), + ordered_sum([first.parts[0]] + second.parts[1:], RANK_CHUNK), + ] + exact = all(same_bits(g, p) for g, p in zip(got, paired)) + wrong = [differing(g, call.ref()) for g, call in zip(got, order)] + assert R.all_true(exact and min(wrong) > 0.5), ( + f"rank {R.rank}: position-paired sums {exact}, shares differing from the intended calls {wrong}" ) assert_clean(EX_A, "after the swapped pair") call = Call(70_002, MAX_TOKENS) @@ -492,6 +560,7 @@ def check_token_count_mismatch_returns_stale_rows() -> None: check_bit_identical_to_mnnvl_oneshot, check_unsupported_token_counts_raise_on_every_rank, check_dip_and_regrow_sequence, + check_ranks_run_a_step_apart, check_two_exchanges_interleaved, check_graph_capture_and_replay, check_count_parity_across_the_int32_wrap, diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_common.py b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_common.py index d93f600557c7..1fae85271a29 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_common.py +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_common.py @@ -3,8 +3,8 @@ """Shared pieces of the Kimi K3 sandwich matrices (``_k3_sandwich_oproj_op_matrix.py``, ``_k3_sandwich_tail_op_matrix.py``, ``_k3_sandwich_plain_op_matrix.py``): one class per op holding one call's arguments on every rank and its native-torch reference, a driver for call sequences (chained the way the model chains -its residual streams; eager with a random rank late, or captured and replayed), and the A1 checks the three matrices -run the same way on a ``K3SandwichWorkspace``. +its residual streams; eager with a random rank late, or captured and replayed), and the state checks the three +matrices run the same way on a ``K3SandwichWorkspace``. Started by file path like ``_lockstep`` (this tree is not a package). Importing it pulls in torch only; ``bind`` (called by the rank body) imports the catalog wrappers. @@ -57,15 +57,21 @@ REPLAYS = 8 WEIGHT_SEEDS = {"oproj": 900, "tail": 910, "plain": 920, "down": 930} +CAPTURE_REFUSED = "outside CUDA-graph capture" # create()'s reason on a rank that is capturing +PEER_REFUSED = "another rank cannot" # its reason on the other ranks + R = None # this rank (a _lockstep.Rank), set by bind() OPS: Dict[str, Callable] = {} WEIGHTS: Dict[str, List[List[torch.Tensor]]] = {} STATS = {"normed": 0.0, "tail_updated": 0.0, "tap": 0.0} +# The multicast allocations create() has reached on this rank, counted from bind() on. +ALLOCATIONS = [0] def bind(rank: ls.Rank, weight_sets: Dict[str, int]) -> None: - """Bind this rank and the three catalog wrappers, and draw ``weight_sets[kind]`` weight sets of each kind ("oproj", - "tail", "plain", "down"), every rank's slice, from fixed seeds.""" + """Bind this rank and the three catalog wrappers, draw ``weight_sets[kind]`` weight sets of each kind ("oproj", + "tail", "plain", "down"), every rank's slice, from fixed seeds, and count from here on the multicast allocations + ``create`` reaches (``ALLOCATIONS``), so that a refused ``create`` can be shown to allocate nothing.""" global R from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm import ( k3_sandwich_oproj, @@ -82,6 +88,21 @@ def bind(rank: ls.Rank, weight_sets: Dict[str, int]) -> None: for s in range(sets): g = _gen(WEIGHT_SEEDS[kind] + s) WEIGHTS[kind].append([_weight(kind, g) for _ in range(R.world)]) + _count_allocations() + + +def _count_allocations() -> None: + """Wrap the multicast allocation ``K3SandwichWorkspace.create`` makes (``_make_mnnvl_mcast_buffer``, which it looks + up at every call) with a counter; the allocation itself is unchanged.""" + from tensorrt_llm._torch.distributed import ops + + allocate = ops._make_mnnvl_mcast_buffer + + def counted(*args, **kwargs): + ALLOCATIONS[0] += 1 + return allocate(*args, **kwargs) + + ops._make_mnnvl_mcast_buffer = counted def report() -> None: @@ -471,27 +492,47 @@ def armed_and_sized(ws) -> None: assert int(ws.flags.abs().sum()) == 0, "every counter 0" -def create_refuses_capture(workspace_type, ws, next_call: Call) -> None: - """``create`` is collective and allocates: under CUDA-graph capture it raises RuntimeError on every rank at once, - and on one rank alone while its peers do not call it -- before any collective, since a rank inside one would wait - for its peers there. The workspace in use is untouched: the next call is correct.""" - - def attempt() -> bool: - stream = torch.cuda.Stream() - stream.wait_stream(torch.cuda.current_stream()) - graph = torch.cuda.CUDAGraph() - try: +def _create(workspace_type, capture: bool) -> str: + """This rank's ``create``, under CUDA-graph capture on a side stream or eagerly; returns the message of the + RuntimeError it raised ("" if it returned), once neither stream is left capturing.""" + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + graph = torch.cuda.CUDAGraph() + message = "" + try: + if capture: with torch.cuda.graph(graph, stream=stream): workspace_type.create(R.mapping, fabric_handle=R.fabric) - except RuntimeError as exc: - return "outside CUDA-graph capture" in str(exc) - return False + else: + workspace_type.create(R.mapping, fabric_handle=R.fabric) + except RuntimeError as exc: + message = str(exc) + with torch.cuda.stream(stream): + side_capturing = torch.cuda.is_current_stream_capturing() + assert not (torch.cuda.is_current_stream_capturing() or side_capturing), ( + "a stream was left capturing" + ) + return message + - every = attempt() +def create_refuses_capture(workspace_type, ws, next_call: Call) -> None: + """``create`` is collective and the ranks agree before it allocates, so a rank under CUDA-graph capture makes every + rank raise RuntimeError. (a) Every rank captures: every rank raises, naming the capture. (b) The last rank captures + while its peers call ``create`` eagerly at the same point: every rank raises, the capturing rank naming the + capture, its peers saying another rank cannot. No rank reaches the allocation in either case (the counted + multicast allocation, which the workspaces in use went through), no stream is left capturing, and the workspace + in use is untouched: the next call is correct.""" + assert ALLOCATIONS[0] > 0, "the counter did not see the workspaces in use being created" + before = ALLOCATIONS[0] + every = _create(workspace_type, capture=True) + every_ok = CAPTURE_REFUSED in every R.barrier() - alone = attempt() if R.rank == R.world - 1 else True - assert R.all_true(every and alone), ( - f"create under capture: every rank raised {every}, one rank alone {alone}" + capturing = R.world - 1 + one = _create(workspace_type, capture=R.rank == capturing) + one_ok = (CAPTURE_REFUSED if R.rank == capturing else PEER_REFUSED) in one + allocated = ALLOCATIONS[0] - before + assert R.all_true(every_ok and one_ok and allocated == 0), ( + f"every rank capturing: {every!r}; rank {capturing} capturing: {one!r}; allocations reached {allocated}" ) next_call.verify(next_call.run(ws), "after the refused creates") diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_oproj_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_oproj_op_matrix.py index cd9ae009f226..1230862ee4fa 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_oproj_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_oproj_op_matrix.py @@ -8,7 +8,8 @@ replay mixed with eager calls, and a negative control in which one rank swaps two calls and every rank gets a wrong answer without an error. The workspace is also shared the way the model shares it: one decode step of target layers (this op, then ``k3_sandwich_tail``) with the drafter's ``k3_sandwich_plain`` calls interleaved, call by call, then -that step captured and replayed between eager calls of the three ops. And ``create`` must refuse CUDA-graph capture. +that step captured and replayed between eager calls of the three ops. And ``create`` must agree across the +ranks before it allocates: when any rank is capturing a CUDA graph, every rank raises and none allocates. CUDA_VISIBLE_DEVICES=0,1,2,3 python _k3_sandwich_oproj_op_matrix.py [--world-size 4] srun -n 16 --mpi=pmix python _k3_sandwich_oproj_op_matrix.py --launcher srun --world-size 16 @@ -56,8 +57,9 @@ def check_workspace_is_armed_and_sized() -> None: def check_create_refuses_capture() -> None: - """``create`` under CUDA-graph capture raises on every rank at once and on one rank alone (before any - collective); the workspace in use is untouched.""" + """``create`` under CUDA-graph capture: every rank capturing, every rank raises; one rank capturing while its + peers call it eagerly, every rank raises too. No rank allocates, no stream is left capturing, and the workspace + in use is untouched.""" cm.create_refuses_capture(WORKSPACE, WS_A, cm.OprojCall(600, 8, 2)) diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_plain_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_plain_op_matrix.py index c0def69078f6..ff3d41250b01 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_plain_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_plain_op_matrix.py @@ -52,8 +52,9 @@ def check_workspace_is_armed_and_sized() -> None: def check_create_refuses_capture() -> None: - """``create`` under CUDA-graph capture raises on every rank at once and on one rank alone (before any - collective); the workspace in use is untouched.""" + """``create`` under CUDA-graph capture: every rank capturing, every rank raises; one rank capturing while its + peers call it eagerly, every rank raises too. No rank allocates, no stream is left capturing, and the workspace + in use is untouched.""" cm.create_refuses_capture(WORKSPACE, WS_A, cm.PlainCall(600, 8)) diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_tail_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_tail_op_matrix.py index 1dba6c7713e7..074f6b22886e 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_tail_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_tail_op_matrix.py @@ -56,8 +56,9 @@ def check_workspace_is_armed_and_sized() -> None: def check_create_refuses_capture() -> None: - """``create`` under CUDA-graph capture raises on every rank at once and on one rank alone (before any - collective); the workspace in use is untouched.""" + """``create`` under CUDA-graph capture: every rank capturing, every rank raises; one rank capturing while its + peers call it eagerly, every rank raises too. No rank allocates, no stream is left capturing, and the workspace + in use is untouched.""" cm.create_refuses_capture(WORKSPACE, WS_A, cm.TailCall(600, 8, 2)) From 344d3409c51cd383cb6d3a5240d968bb1c000216 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 20:58:16 -0700 Subject: [PATCH 071/161] [None][feat] modeling_v2 catalog: moe/k3_moe over trtllm::k3_moe; moe/k3_route_quant; the front's publish flag moe/k3_moe is now one call: trtllm::k3_moe with the layer's weights and counters and its state's scratch, for a K3MoeState (up to 8 tokens) or a K3MoeWideState (up to 64); head= passes the TP group's K3MoeHeadWorkspace for, and only for, a head_flags state (K3MoeLayer checks the same). The routing and quantization before it are their own stateless entry, moe/k3_route_quant (trtllm::k3_route_quant, bit-identical to trtllm::kimi_k3_noaux_tc_mxfp8_quant). moe/k3_moe_front takes publish=True to release the ready words a head_flags k3_moe acquires, so the head_flags pair runs on the entries alone. The k3_moe contract states what the kernel touches before its grid-dependency wait and the precondition that follows. The kernel tests and the front matrix call the new API; the index rows and the B200 list follow, and op.py's k3_moe code is formatted with the repo's hooks. Signed-off-by: Vasanth Sabavat --- .../modeling_v2/catalog/index.yaml | 8 +- .../modeling_v2/catalog/moe/k3_moe.md | 291 +++++++------- .../modeling_v2/catalog/moe/k3_moe.py | 97 ++--- .../modeling_v2/catalog/moe/k3_moe_front.md | 117 +++--- .../modeling_v2/catalog/moe/k3_moe_front.py | 7 +- .../modeling_v2/catalog/moe/k3_route_quant.md | 107 +++++ .../modeling_v2/catalog/moe/k3_route_quant.py | 39 ++ .../cute_dsl_kernels/k3_fused_moe/op.py | 33 +- .../test_lists/test-db/l0_b200.yml | 3 +- .../kimi_k3/test_k3_fused_moe.py | 15 +- .../kimi_k3/test_k3_moe_front.py | 40 +- .../kimi_k3/test_k3_moe_wide.py | 11 +- .../comm/_k3_moe_front_op_matrix.py | 149 ++++--- .../moe/test_modeling_v2_k3_moe.py | 367 +++++++++--------- ...test_modeling_v2_k3_moe_front_op_matrix.py | 10 +- .../moe/test_modeling_v2_k3_route_quant.py | 262 +++++++++++++ 16 files changed, 1004 insertions(+), 552 deletions(-) create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_route_quant.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_route_quant.py create mode 100644 tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_route_quant.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml index 2d628ad02393..6ee1d5eac663 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml @@ -254,6 +254,10 @@ entries: impl: torch.ops.trtllm.fp4_block_scale_moe_runner summary: "Full NVFP4-weight / NVFP4-activation MoE layer (trtllm-gen W4A4): caller-supplied top-k ids/weights + grouped FC1 GEMM over pre-shuffled block-scaled [up; gate] expert weights + SwiGLU + NVFP4 requantization of that activation + grouped FC2 GEMM + routing-weighted fp32 combine, into a fresh or caller-provided [tokens, hidden] bf16 buffer; the checkpoint's per-tensor global scales enter as three per-expert fp32 scalars, hidden must be a multiple of 256, expert-parallel window selected by local_expert_offset/local_num_experts, do_finalize=False returns the per-slot rows plus their permuted-row map instead of the combine" + - path: moe/k3_route_quant.py + impl: torch.ops.trtllm.k3_route_quant + summary: "Kimi K3's top-16 routing (sigmoid + bias-corrected selection, the unbiased weights renormalized and scaled) and the MXFP8 quantization of the routed latent for 1-64 decode tokens in one CuTe DSL kernel, bit-identical to trtllm::kimi_k3_noaux_tc_mxfp8_quant; stateless (a compile cache only); early_trigger lets a programmatic-dependent k3_moe launch before the outputs are written" + # ─── quantization ────────────────────────────────────────────── - path: quantization/mxfp8_quantize.py impl: torch.ops.trtllm.mxfp8_quantize @@ -310,5 +314,5 @@ entries: summary: "Kimi K3's MoE front at decode size in one kernel: the row-sharded MoE head GEMV, its all-gather over a caller-owned K3MoeHeadWorkspace (two alternating Lamport buffers rotated by every front call), the top-16 routing and MXFP8 latent as trtllm::k3_route_quant returns them for the gathered head, and the shared experts' gate_up + SiTU-and-mul; optionally publishing per-token ready words that a head_flags build of k3_moe acquires" - path: moe/k3_moe.py - impl: tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.op.K3MoeLayer - summary: "Kimi K3's routed experts at decode size: this rank's routed partial from the persistent CuTe DSL kernel k3_moe (FC1 + SiTU + FC2 with the routing-weighted combine over the TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers, read in place) after k3_route_quant or the MoE front, on caller-owned per-rank state: a K3MoeState and one K3MoeLayer per MoE layer for up to 8 tokens, a K3MoeWideState and one K3MoeWideLayer per layer for up to 64; a state's layers share its scratch (left armed by every call) and run in one stream order, each layer's counters are left zero" + impl: torch.ops.trtllm.k3_moe + summary: "Kimi K3's routed experts at decode size: this rank's routed partial from the persistent CuTe DSL kernel k3_moe (FC1 + SiTU + FC2 with the routing-weighted combine over the TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers, read in place) on the routing and MXFP8 latent of k3_route_quant or the MoE front, over caller-owned per-rank state: a K3MoeState (up to 8 tokens; its head_flags build acquires a publishing front's ready words on the group's K3MoeHeadWorkspace) or a K3MoeWideState (up to 64), and one K3MoeLayer per MoE layer; a state's layers share its scratch (left armed by every call) and run in one stream order, each layer's counters are left zero" diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md index 78ff45874b0f..590d921bf862 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md @@ -5,38 +5,21 @@ receipts: # k3_moe -**Wraps** three calls of the caller-owned layer objects in `cute_dsl_kernels/k3_fused_moe/op.py`, one each. The -expert kernel `k3_moe` is a CuTe DSL kernel launched through its compiled function, not a torch op. - -| Function | Runs | Tokens `M` | -|---|---|---| -| `k3_moe` | `K3MoeLayer.__call__`: `torch.ops.trtllm.k3_route_quant`, then `k3_moe` as its programmatic dependent | 1-8 | -| `k3_moe_fused_front` | `K3MoeLayer.front`: `torch.ops.trtllm.k3_moe_front` (entry `moe/k3_moe_front`), then `k3_moe` | 1-8 | -| `k3_moe_wide` | `torch.ops.trtllm.k3_route_quant(early_trigger=True)`, then `K3MoeWideLayer.__call__`: the m_max 64 build of `k3_moe` | 1-64 | +**Wraps** `torch.ops.trtllm.k3_moe` (one call), on caller-owned state: a `K3MoeState` (up to 8 tokens) or +`K3MoeWideState` (up to 64) and one `K3MoeLayer` per MoE layer; for a head_flags state, also the TP group's +`K3MoeHeadWorkspace`. ## Semantics Kimi K3's routed experts at decode size: this rank's routed partial, i.e. for each token the sum over its top-16 -experts that this rank holds of the expert's output times its routing weight. Two launches on the current stream, no -host synchronization (the op module's statement). - -Routing and MXFP8 quantization of the latent, `k3_route_quant` (inside `k3_moe` and `k3_moe_wide`): the CuTe DSL -form of `trtllm::kimi_k3_noaux_tc_mxfp8_quant`, whose four outputs it returns bit for bit (certified at `M` 1-8, 16, -33 and 64, with and without the early dependent trigger). Per token, as the kernel module states it: +experts that this rank holds of the expert's output times its routing weight, from the persistent CuTe DSL kernel +`k3_moe`. Its inputs are the routing and MXFP8 latent of `moe/k3_route_quant` or `moe/k3_moe_front`: `topk_ids`, +`topk_weights`, and `x_fp8` with its UE8M0 scales `x_sf` (`x[t] = x_fp8[t] * 2^(x_sf[t] - 127)` per 32 columns). Per +token `t`: ``` -s = 0.5 * tanh(0.5 * router_logits) + 0.5 # fp32, [896] -ids = the 16 experts with the largest s + e_score_correction_bias, descending, ties to the lower id -weights = bf16(s[ids] * routed_scaling_factor / (sum of the 16 s[ids] + 1e-20)) # in fp64 -xq, xs = MXFP8(latent): e4m3 codes, one UE8M0 scale per 32 columns, scale 2^ceil(log2(amax / 448)) -``` - -The experts, `k3_moe`, per token `t`: - -``` -y[t] = bf16( sum over the slots j with offset <= ids[t, j] < offset + num_local of - weights[t, j] * FC2_e(q8(SiTU(FC1_e(x_t)))) ), e = ids[t, j] - offset -x_t = xq[t] * 2^(xs[t] - 127) # the dequantized latent row, [3584] +y[t] = bf16( sum over the slots j with offset <= topk_ids[t, j] < offset + num_local of + topk_weights[t, j] * FC2_e(q8(SiTU(FC1_e(x[t])))) ), e = topk_ids[t, j] - offset FC1_e(x) = (x @ W_up[e]^T, x @ W_gate[e]^T) # [i_tp] each, fp32 accumulation SiTU = 4 tanh(gate / 4) sigmoid(gate) * 25 tanh(up / 25) # caps 4.0 / 25.0 q8 = MXFP8 per 32 intermediate columns, scale 2^ceil(log2(amax / 448)), e4m3 round to nearest even @@ -46,13 +29,14 @@ FC2_e(a) = a @ W_down[e]^T # [3584], fp32 `W_up[e]`, `W_gate[e]` (`[i_tp, 3584]`) and `W_down[e]` (`[3584, i_tp]`) are the values of the layer's MXFP4 expert `e` (E2M1 codes times `2^(E8M0 - 127)` per 32 K elements), read in place from the TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers (*Preconditions*). `offset` is `local_expert_offset`, `num_local` the state's. The SiTU caps are constants of the -`k3_moe` build (Kimi K3's `activation_situ_beta` 4.0 and `activation_situ_linear_beta` 25.0), not arguments; the -`gate_cap` / `linear_cap` arguments of `k3_moe_fused_front` are the shared activation's (the front's). +kernel (Kimi K3's `activation_situ_beta` 4.0 and `activation_situ_linear_beta` 25.0), not arguments. The combine adds +each token's expert terms in fp32 in an order fixed by the call's grouping. -Numerics, certified for `k3_moe` at every `M` 1-8 and for `k3_moe_wide` at `M` 1, 2, 7, 8, 9, 16, 33, 40 and 64, in -three routing cases each: random logits and no local expert for both; 16 local experts per token, none shared (128 -groups at `M` 8, the M <= 8 build's group capacity) for `k3_moe`; 100 local experts with 9 of 64 tokens and 124 with -one (324 groups at `M` 64, the wide build's capacity) for `k3_moe_wide`: +Numerics, certified with `moe/k3_route_quant` as the producer, for the M <= 8 build at every `M` 1-8 and for the wide +build at `M` 1, 2, 7, 8, 9, 16, 33, 40 and 64, in three routing cases each: random logits and no local expert for +both; 16 local experts per token, none shared (128 groups at `M` 8, the M <= 8 build's group capacity), for the M <= 8 +build; 100 local experts with 9 of 64 tokens and 124 with one (324 groups at `M` 64, the wide build's capacity), for +the wide build: - within the op-catalog gates (8 bf16 ulp of the token row's largest magnitude per element, 4 ulp relative RMS) of an fp64 reference of the formula above over the dequantized experts, and of the stock path @@ -60,88 +44,65 @@ one (324 groups at `M` 64, the wide build's capacity) for `k3_moe_wide`: - run-to-run bit identical; - a token's row may round differently when other tokens share the call: FC2 adds a token's expert terms in slices whose bounds follow the call's group count. At most 1 bf16 ulp: each `M`'s rows against the same rows of the - 8-token call, and `k3_moe_wide`'s rows at `M` <= 8 against `k3_moe`'s; + 8-token call, and the wide build's rows at `M` <= 8 against the M <= 8 build's; - a token with no expert on this rank gets a zero row. -`k3_moe_fused_front` returns `(y, shared)`: `y` is `k3_moe` applied to the front's routing and MXFP8 latent, `shared` -the front's shared activation (`moe/k3_moe_front`). Certified in that entry's 4-rank matrix: `y` within the op-catalog -gates of the stock runner on the front's own routing and latent; `shared` bit for bit the front's; on a head_flags -state, `y` and `shared` bit for bit the plain state's. +With `moe/k3_moe_front` as the producer (certified in that entry's 4-rank matrix): `y` within the op-catalog gates of +the stock runner on the front's own routing and latent, and on a head_flags state bit for bit the plain state's. -Fusion boundary. Inside: the routing and the MXFP8 latent (`k3_moe`, `k3_moe_wide`) or the whole MoE front -(`k3_moe_fused_front`); the grouping of (expert, token) pairs; FC1, SiTU, the MXFP8 intermediate, FC2; the -routing-weighted combine. Outside: the router and latent-down GEMMs that produce `router_logits` and `latent` -(`k3_moe`, `k3_moe_wide`); the sum of the routed partials over the ranks that hold the other experts and -intermediate slices (the routed-latent all-reduce); the latent-up projection; the shared experts' down projection -(and, except in `k3_moe_fused_front`, their gate_up and activation); the residual. +Fusion boundary. Inside: the grouping of (expert, token) pairs, FC1, SiTU, the MXFP8 intermediate, FC2, the +routing-weighted combine. Outside: the routing and the latent's MXFP8 quantization (`moe/k3_route_quant` or +`moe/k3_moe_front`); the sum of the routed partials over the ranks that hold the other experts and intermediate +slices (the routed-latent all-reduce); the latent-up projection; the shared experts; the residual. ## Signature ```python def k3_moe( - latent: torch.Tensor, - router_logits: torch.Tensor, - e_score_correction_bias: torch.Tensor, + x_fp8: torch.Tensor, + x_sf: torch.Tensor, + topk_ids: torch.Tensor, + topk_weights: torch.Tensor, local_expert_offset: int, - routed_scaling_factor: float, layer: K3MoeLayer, -) -> torch.Tensor - -def k3_moe_fused_front( - x: torch.Tensor, - w_front: torch.Tensor, - e_score_correction_bias: torch.Tensor, - local_expert_offset: int, - routed_scaling_factor: float, - shared_cols: int, - gate_cap: float, - linear_cap: float, - head: K3MoeHeadWorkspace, - layer: K3MoeLayer, -) -> Tuple[torch.Tensor, torch.Tensor] - -def k3_moe_wide( - latent: torch.Tensor, - router_logits: torch.Tensor, - e_score_correction_bias: torch.Tensor, - local_expert_offset: int, - routed_scaling_factor: float, - layer: K3MoeWideLayer, + head: Optional[K3MoeHeadWorkspace] = None, out: Optional[torch.Tensor] = None, ) -> torch.Tensor ``` -The wrapper module re-exports the state types (`K3MoeState`, `K3MoeLayer`, `K3MoeWideState`, `K3MoeWideLayer`, -`K3MoeHeadWorkspace`) and `is_supported`. +The entry passes the layer's weight buffers and counters and its state's scratch and build options to the op: +`trtllm::k3_moe(x_fp8, x_sf, topk_ids, topk_weights, w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale, +c, cs, part, counters, local_expert_offset, num_local, num_ctas, m_max, use_pdl, head_ready=None, head_flags=None, +out=None)`, with `mutates_args = (c, cs, part, counters, head_ready, head_flags, out)`: every buffer the kernel +writes. The wrapper module re-exports `K3MoeState`, `K3MoeWideState`, `K3MoeLayer`, `K3MoeHeadWorkspace` and +`is_supported`. ### Certified arguments | Argument | Shape | Dtype | Layout | Device | |---|---|---|---|---| -| `latent` | `[M, 3584]`; `M` 1-8 (`k3_moe`), 1-64 (`k3_moe_wide`) | bf16 | contiguous | CUDA, the state's device | -| `router_logits` | `[M, 896]` | fp32 | as `latent` | CUDA | -| `e_score_correction_bias` | `[896]` | fp32 | contiguous | CUDA | +| `x_fp8` | `[M, 3584]`; `M` 1-8 on a `K3MoeState`, 1-64 on a `K3MoeWideState` | float8_e4m3fn | contiguous | CUDA, the state's device | +| `x_sf` | `[M, 112]` (`M * 112` bytes, linear: one byte per 32 columns) | uint8 (UE8M0) | contiguous | CUDA | +| `topk_ids` | `[M, 16]`, global expert ids | int32 | contiguous | CUDA | +| `topk_weights` | `[M, 16]` | bf16 | contiguous | CUDA | | `local_expert_offset` | scalar: the global id of the layer's expert 0 (224 certified: experts `[224, 448)`) | Python int | — | — | -| `routed_scaling_factor` | scalar (2.827 certified) | Python float | — | — | -| `layer` | a `K3MoeLayer` of a plain `K3MoeState` (`k3_moe`); of a plain or head_flags one (`k3_moe_fused_front`); a `K3MoeWideLayer` (`k3_moe_wide`) | — | — | — | -| `out` (`k3_moe_wide`) | `None`, or `[>= M, 3584]`: the result is `out[:M]` (the same storage), rows past `M` untouched (certified at `M` 9 and 64) | bf16 | contiguous | CUDA | -| `x`, `w_front`, `shared_cols`, `gate_cap`, `linear_cap`, `head` (`k3_moe_fused_front`) | as `moe/k3_moe_front`; `head` that entry's `K3MoeHeadWorkspace` | — | — | — | -| returns | `y [M, 3584]` (`k3_moe_fused_front`: `(y, shared [M, shared_cols])`) | bf16 | contiguous, newly allocated (or `out[:M]`) | the inputs' device | +| `layer` | a `K3MoeLayer` of a plain or head_flags `K3MoeState`, or of a `K3MoeWideState` | — | — | — | +| `head` | `None`; for and only for a head_flags state's layers, the TP group's `K3MoeHeadWorkspace` (certified in `moe/k3_moe_front`'s matrix) | — | — | — | +| `out` | `None`, or `[>= M, 3584]`: the call writes `out[:M]` and returns an empty `[0, 3584]`; rows past `M` untouched (certified at `M` 3, 8 and 9, 64 on the two builds) | bf16 | contiguous | CUDA | +| returns | `y [M, 3584]`, or `[0, 3584]` with `out` | bf16 | contiguous, newly allocated | the inputs' device | -The state objects' certified construction: `K3MoeState(device, 768, 224)` (`head_flags` False or True, `config` None), -`K3MoeWideState(device, 768, 224)` (`use_pdl` True), and `state.layer(w3_w1_weight, w3_w1_weight_scale, w2_weight, -w2_weight_scale)` over this rank's 224 experts in the TRTLLM-Gen W4A8_MXFP4_MXFP8 layout at intermediate 768: -`[224, 1536, 1792]`, `[224, 1536, 112]`, `[224, 3584, 384]`, `[224, 3584, 24]`, all uint8. +The four inputs are `moe/k3_route_quant`'s outputs for `M` tokens (certified) or `moe/k3_moe_front`'s (certified in +its matrix). State construction certified: `K3MoeState(device, 768, 224)` (`head_flags` False or True; `use_pdl` and +`num_ctas` at their defaults), `K3MoeWideState(device, 768, 224)` (`use_pdl` True), and `state.layer(w3_w1_weight, +w3_w1_weight_scale, w2_weight, w2_weight_scale)` over this rank's 224 experts in the TRTLLM-Gen W4A8_MXFP4_MXFP8 +layout at intermediate 768: `[224, 1536, 1792]`, `[224, 1536, 112]`, `[224, 3584, 384]`, `[224, 3584, 24]`, uint8. ## State **Objects.** Per rank and per device, owned by the caller (`cute_dsl_kernels/k3_fused_moe/op.py`, re-exported by -`catalog/moe/k3_moe.py`): - -- `K3MoeState` (`M` <= 8) and one `K3MoeLayer` per MoE layer from `state.layer(...)`; -- `K3MoeWideState` (`M` <= 64) and one `K3MoeWideLayer` per MoE layer; -- `k3_moe_fused_front` also takes the TP group's `K3MoeHeadWorkspace` (collective; its *State* is in - `moe/k3_moe_front`). +`catalog/moe/k3_moe.py`): a build's scratch, `K3MoeState` (`m_max` 8) or `K3MoeWideState` (`m_max` 64), and one +`K3MoeLayer` per MoE layer from `state.layer(...)`. A head_flags `K3MoeState`'s calls also read and write the TP +group's `K3MoeHeadWorkspace` (collective; its *State* is in `moe/k3_moe_front`). **Contents and size.** At the certified layout (`i_tp` 768, `num_local` 224), certified right after construction: @@ -150,49 +111,47 @@ w2_weight_scale)` over this rank's 224 experts in the TRTLLM-Gen W4A8_MXFP4_MXFP | `K3MoeState` | `c`: the FC1 -> FC2 intermediate slab | int8 `[G, 8, i_tp]`, `G = min(num_local, 128)` = 128 | 786,432 | armed: every byte 0x80 (FP8 -0.0) | | | `cs`: its E8M0 scales | int8 `[G, 8, i_tp / 8]` | 98,304 | armed: bytes 0-3 of every 16-byte group 0xFF (E8M0 NaN), bytes 4-15 zero | | | `part`: the FC2 partial rows | fp32 `[8 G, 3584]` | 14,680,064 | zero when built; then what the last call left | -| `K3MoeLayer` | `counters` | int32 `[32 + 2 G]` = `[288]` | 1,152 | zero | | `K3MoeWideState` | `c`, `cs` | int8 `[G, 8, i_tp]`, `[G, 8, i_tp / 8]`, `G = 224 + (1024 - 224) / 8` = 324 | 1,990,656 + 248,832 | armed, as above | -| | `part` | fp32 `[1024, 3584]` | 14,680,064 | not initialized | -| `K3MoeWideLayer` | `counters` | int32 `[32 + 2 G]` = `[680]` | 2,720 | zero | +| | `part` | fp32 `[1024, 3584]` | 14,680,064 | zero when built; then what the last call left | +| `K3MoeLayer` | `counters` | int32 `[32 + 2 G]`: `[288]` on a `K3MoeState`, `[680]` on a `K3MoeWideState` | 1,152 / 2,720 | zero | `G` is the build's group capacity (`group_capacity` in the kernel module): an (expert, up to 8 tokens) group per local -expert a call routes to, at most `min(num_local, 8 x 16)` for `M` <= 8; the wide build gives an expert with `t` tokens -`ceil(t / 8)` groups. Each state also holds `mod` (its configuration's kernel module, from a process-wide cache keyed -by the configuration), `compiled` (the compiled `k3_moe`, built by the state's first call) and, for the M <= 8 build, -`head_flags`. A layer holds views of its four weight buffers (no copy) and its counters; it does not hold +expert a call routes to, at most `min(num_local, 8 x 16)` for `M` <= 8; the wide build gives an expert with `t` +tokens `ceil(t / 8)` groups. A state also holds its build's options (`m_max`, `num_ctas`: one CTA per SM by default, +`use_pdl`, `head_flags`; certified), its kernel module and `compiled` (whether its build is in the compile cache). A +layer holds its four weight tensors themselves (no copy; certified) and its counters; it does not hold `local_expert_offset`. **Who creates it, and when.** The target, in `post_load_weights`, after the expert weights are final (a layer reads the buffers it was built over; see *What a wrong order does*): `K3MoeState(device, i_tp, num_local, -head_flags=False)` once per device and `state.layer(w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale)` -once per MoE layer; likewise `K3MoeWideState(device, i_tp, num_local)` and its layers. Not collective. Eager: the -four constructors (`K3MoeState()`, `K3MoeState.layer()`, `K3MoeWideState()`, `K3MoeWideState.layer()`) raise -`RuntimeError` under CUDA-graph capture (certified). The kernel compiles on the first call of each -state object (seconds), which must be eager: under capture that call raises `RuntimeError` before the `k3_moe` launch -and leaves the state uncompiled (certified for both builds). `config` (kernel options for tests and A/B runs) stays -`None`. No environment variable selects anything here except PDL (*Metadata consumed*). +head_flags=False, use_pdl=None, num_ctas=None)` and `K3MoeWideState(device, i_tp, num_local, use_pdl=True)` once per +device, `state.layer(w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale)` once per MoE layer. Not +collective. Eager: the four constructors raise `RuntimeError` under CUDA-graph capture (certified). Each build +compiles on its first call on a device (seconds; one process-wide cache keyed by device and build options), which must +be eager: under capture that call raises `RuntimeError` before the `k3_moe` launch and the build stays uncompiled +(certified for both builds, with the cache cold). **Which ops may share one object.** All layers of a state share its slab and partial rows; each layer has its own counters. Layers of one state may be built over the same weight buffers (two counter sets) or over different ones -(certified, both). On a plain `K3MoeState`, `k3_moe` and `k3_moe_fused_front` calls may be mixed (one build). A -head_flags `K3MoeState` serves `k3_moe_fused_front` only: `k3_moe` on its layers raises `ValueError` before any -launch (certified). `K3MoeWideState` serves `k3_moe_wide` only. Separate states share nothing: with calls on two -`K3MoeState`s and a `K3MoeWideState` interleaved in an irregular pattern (26 calls, back to back), each returns the -bits of the same call made alone (certified). +(certified, both). Separate states share nothing: with calls on two `K3MoeState`s and a `K3MoeWideState` interleaved +in an irregular pattern (26 calls, back to back), each returns the bits of the same call made alone (certified). The +entry passes `head` for, and only for, a head_flags state's layers: both mismatches raise `ValueError` before any +launch (certified). A head_flags state's calls pair with the front's on one `K3MoeHeadWorkspace`: the front call +before each must publish the workspace's ready words (the `moe/k3_moe_front` entry with `publish=True`), and each +publishing front call must be followed by exactly one such `k3_moe` call on that workspace (*Preconditions*). **Call-order invariant.** The calls on all layers of one state run one after the other in one stream order. Every call needs the slab armed and its layer's counters at zero, which only the end of the previous call on the state -guarantees; and (the kernel's statement) a call issues its first tile claim on its layer's counters before its grid -dependency wait, relying on the previous call of that layer having completed before the producer kernel launched this -one, which holds on one stream. Calls on one state from two streams at once are not exercised: nothing in the state -keeps two concurrent calls' groups apart. The state is per rank: there is no cross-rank order, except through the -head workspace for `k3_moe_fused_front` (`moe/k3_moe_front`). +guarantees, and a call claims its first tile on its layer's counters before its grid-dependency wait (the kernel's +statement; *Preconditions*). Calls on one state from two streams at once are not exercised: nothing in the state keeps +two concurrent calls' groups apart. The state is per rank: no cross-rank order, except through the head workspace on +a head_flags state (`moe/k3_moe_front`). **What a later launch reads.** The slab armed: FC2 treats a group's intermediate as written only once its FP8 -0.0 and E8M0 NaN sentinels are gone. The layer's counters at zero: the tile-queue cursor, the FC2 m-tile arrivals, the per-group FC1 and FC2 counts. The partial rows: the M <= 8 build's combine also loads the rows of the token slots past `M`, whose sums it drops; they are zero in a new state, so those loads never read unwritten memory (the op module's -statement). With head_flags: the head workspace's epoch `flags[2]` and its ready words (below). +statement). On a head_flags state: the workspace's epoch `flags[2]` and its ready words (below). **How it is re-armed.** By `k3_moe` itself, on every call: the last FC2 task that reads a group's intermediate puts its sentinels back, and each counter is reset by its last user (the cursor by the grid's last claim). Nothing is @@ -200,18 +159,17 @@ cleared between calls and nothing records a call's size, so a call after a small larger call left. Certified: the slab armed and every counter zero after every single call and every sequence of the test; decode steps of three layers at `M` 8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8 and of two wide layers at `M` 64, 64, 16, 1, 40, 64, 8, 23, 64, 9, 64, new inputs every call, back to back, each call the bits of the same call made alone -(itself within the gates of the stock path); a captured step (three `k3_moe` and one `k3_moe_wide` call) replayed 6 -times with rewritten inputs and eager calls of other `M` on the same states between replays, every replayed and -eager call the bits of the same call alone. - -With head_flags (`k3_moe_fused_front`), the ready-word handoff. The front releases `ready[t]` (token `t`'s ids and -weights) and `ready[8 + t]` (its MXFP8 row) as `epoch + 1`, with `epoch = head.flags[2]`; the head_flags `k3_moe` -reads the epoch, acquires those words for its `M` tokens instead of waiting for the front's grid, waits for the grid -only before its FC2 phase (which writes the output, memory it does not own), and at its last tile claim writes -`epoch + 1` into the ready words of the tokens past `M` and into `flags[2]` (the kernel's statement). So after every -head_flags call `flags[2]` and all 16 words hold one value, and the next call waits for a value no word holds, across -the int32 wraps too. Certified in `moe/k3_moe_front`'s matrix: before every head_flags call made alone no polled word -already holds `epoch + 1`, and after it the epoch advanced by one with all 16 words at it, across -1 -> 0 and +(itself within the gates of the stock path); a captured step (three M <= 8 calls and one wide call, each routed by +`k3_route_quant` inside the step) replayed 6 times with rewritten inputs and eager calls of other `M` on the same +states between replays, every replayed and eager call the bits of the same call alone. + +On a head_flags state, the ready-word handoff (the kernel's statement). The front releases `ready[t]` (token `t`'s +ids and weights) and `ready[8 + t]` (its MXFP8 row) as `epoch + 1`, with `epoch = head.flags[2]`; `k3_moe` reads the +epoch, acquires those words for its `M` tokens instead of waiting for the front's grid, and at its last tile claim +writes `epoch + 1` into the ready words of the tokens past `M` and into `flags[2]`. So after every head_flags call +`flags[2]` and all 16 words hold one value, and the next call waits for a value no word holds, across the int32 +wraps too. Certified in `moe/k3_moe_front`'s matrix: before every head_flags call made alone no polled word already +holds `epoch + 1`, and after it the epoch advanced by one with all 16 words at it, across -1 -> 0 and 2^31 - 1 -> -2^31 too; after every back-to-back sequence and every replayed step, the epoch advanced by the number of head_flags calls with all 16 words at it. @@ -222,57 +180,74 @@ returning the old experts' partial bit for bit, outside the op-catalog gates of silently stale. Copying the new weights into the old buffers in place is seen by the next call (the new experts' partial, bit for bit), and layers built over the new tensors are correct. So build the layers once the weights are final, and reload weights in place. Not exercised: one state on two streams at once (a race, no deterministic -control); a head_flags call whose polled ready words already hold `epoch + 1`, e.g. after a caller resets -`flags[2]` (`k3_moe` would read the front's output buffers before the front writes them; the matrix checks this -precondition before every head_flags call instead). +control); breaking the head_flags pairing. Without a publishing front before it, a head_flags `k3_moe` waits for +ready words nobody releases; after a publishing front not followed by one, the epoch does not advance, and the next +head_flags call finds its words already at `epoch + 1` and can read the routing before the front writes it (the +matrix checks that precondition before every head_flags call made alone). ## Metadata consumed Besides the state objects (explicit arguments): -- `TRTLLM_ENABLE_PDL` (default on), read by `k3_route_quant` on every call (part of its compile key), by - `k3_moe_front` on every call, and by the M <= 8 build's kernel module when a configuration is first loaded (the - wide build takes `use_pdl` from its constructor). It changes scheduling, not results (the ops' statement); the - tests run with the default. -- Process-wide caches, result-neutral: the kernel modules keyed by configuration (`op._modules`), `k3_route_quant`'s - compiled kernels keyed by (early trigger, PDL), `k3_moe_front`'s keyed by its configuration. The compiled `k3_moe` - is per state object (`state.compiled`): every new state compiles on its first call. +- The launch's programmatic dependent launch is the state's `use_pdl`: a `K3MoeState` built with `use_pdl=None` (the + default) reads `TRTLLM_ENABLE_PDL` (default on) at construction; a `K3MoeWideState` defaults to `True`. The op + itself reads no environment. PDL changes scheduling, not results (the op's statement); the tests run with the + default. +- Process-wide caches, result-neutral: the compiled kernels keyed by (device, build options) (`op._compiled`) and the + kernel modules keyed by build options (`op._modules`). A state's `compiled` reads the former. ## Preconditions -- sm_100, with the CuTe DSL package (`is_supported`): the layer constructor raises `ValueError` otherwise. +- sm_100, with the CuTe DSL package (`is_supported`): `state.layer()` and the op raise `ValueError` otherwise. - The weights are the four contiguous uint8 buffers W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod writes (rows `[up; gate]` interleaved and shuffled in 32-row blocks, scales block-interleaved 128 x 4): `w3_w1_weight [E, 2 i_tp, 1792]`, `w3_w1_weight_scale [E, 2 i_tp, 112]`, `w2_weight [E, 3584, i_tp / 2]`, `w2_weight_scale [E, 3584, i_tp / 32]`, with `E` the state's `num_local` (at most 896) and `i_tp` the state's, a multiple of 128. Anything else raises `ValueError` in `state.layer()`. -- `M` 1-8 (`k3_moe`, `k3_moe_fused_front`) or 1-64 (`k3_moe_wide`): `M` 0 and 9, and 0 and 65, raise `ValueError` - before any launch, every state's slab, partial rows and counters keep their bits, and the next call returns the - bits of the same call made before (certified). +- The four inputs as *Certified arguments* (the op checks dtypes, shapes and contiguity: `ValueError`, certified for + int64 ids); `M` 1-8 (`K3MoeState`) or 1-64 (`K3MoeWideState`): `M` 0 and 9, and 0 and 65, raise `ValueError` before + any launch. Certified: every state's slab, partial rows and counters keep their bits through the refused calls, and + the next call returns the bits of the same call made before. `out` with fewer than `M` rows raises `ValueError` + (certified). - `local_expert_offset` is the global id of the expert the layer's buffers start with. The layer does not record it: another value computes other global ids with these weights, without an error. -- `k3_moe_wide`'s `latent`, `router_logits` and `e_score_correction_bias`, and `k3_moe`'s bias, contiguous: - `k3_route_quant` raises `ValueError` otherwise (its check). `k3_moe` makes its `latent` and `router_logits` - contiguous itself; only contiguous inputs are certified. - Inputs on the state's device, with that device current (not checked). -- Calls may be captured once the state's first call has run eagerly: certified with the captured step above, and in - `moe/k3_moe_front`'s matrix for `k3_moe_fused_front` on both states. +- On a head_flags state: `head` is the TP group's `K3MoeHeadWorkspace` whose ready words the front call just before + published (`moe/k3_moe_front` with `publish=True`), and every publishing front call is followed by exactly one such + call (*State*). +- Before its grid-dependency wait (with `use_pdl`, `k3_moe` launches as a programmatic dependent of the kernel before + it), the kernel touches, from its code: + - M <= 8 and wide builds without head_flags: only the layer's `counters` word 0, the tile-queue cursor, one atomic + add per CTA. It reads its inputs, the weights and the scratch, and writes `part` and the output, only after the + wait. So the kernel before it must not write the layer's counters (no producer does), and the previous call on + the same layer must have ended by the time that kernel lets `k3_moe` launch. On one stream that holds when that + kernel triggers its dependents only after its own grid-dependency wait, as `k3_route_quant` (its early trigger + comes right after that wait) and `k3_moe_front` (its role CTAs trigger after theirs) do, and every kernel since + the previous call waited for its predecessor before finishing. + - head_flags build: it does not wait for the front's grid until its FC2 phase. Before that it reads + `head.flags[2]`, acquires `head.ready[t]` and `head.ready[8 + t]` for its tokens and then reads the routing and + the MXFP8 rows they release, streams the weights, writes the slab and the layer's counters, and at its last claim + writes `head.ready` past `M` and `head.flags[2]`. It waits for the front's grid before writing `part` and the + output, and before it exits. So the front must write each token's routing and MXFP8 row before releasing their + ready words (it fences, then releases), and must not write the weights, the scratch, the counters or + `head.flags[2]` (it does not). + - Either build lets its own dependents launch early (plain builds right after the wait, the head_flags build at + launch): a consumer of `y` must wait for `k3_moe`'s grid (`griddepcontrol.wait`, or plain stream order) before + reading it. +- Calls may be captured once the build's first call has run eagerly: certified with the captured step above, and in + `moe/k3_moe_front`'s matrix on both K3MoeStates. ## Notes -- Certified path: one GB200 GPU (sm_100) for `k3_moe` and `k3_moe_wide`, test - `tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe.py`, at the Kimi K3 TP16 deployment's routed-expert - rank layout (experts TP4 x EP4: 224 local experts of 896, intermediate 768), with random checkpoint-format MXFP4 - experts put through TRT-LLM's own loader. `k3_moe_fused_front`: 4 ranks of one GB200 tray in +- Certified path: one GB200 GPU (sm_100), test `tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe.py`, + at the Kimi K3 TP16 deployment's routed-expert rank layout (experts TP4 x EP4: 224 local experts of 896, + intermediate 768), with random checkpoint-format MXFP4 experts put through TRT-LLM's own loader, routed by + `moe/k3_route_quant`. The head_flags build and the front as producer: 4 ranks of one GB200 tray in `tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py` (entry point `moe/test_modeling_v2_k3_moe_front_op_matrix.py`); its 16-rank receipt is pending with `moe/k3_moe_front`'s. - References: an fp64 reference over the dequantized experts (from the checkpoint-format tensors) and the stock path. - The kernel tests (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py`, `test_k3_moe_wide.py`, - `test_k3_route_quant.py`) remain the exhaustive numerics; this entry's test copies their references. -- Gaps (the kernels are unchanged by this entry): - - The `k3_moe` launch is not a torch op, so nothing declares what it writes: the state's slab and partial rows, the - layer's counters, and with head_flags the head workspace's `flags[2]` and ready words (a torch op would name them - in `mutates_args`). `trtllm::k3_route_quant` writes only its new outputs (`mutates_args=()`). - - A layer does not record `local_expert_offset`, and nothing ties a state to a stream or checks the inputs' device. - - The compiled kernel is per state object, not per configuration: a second state of the same configuration - compiles again. + The kernel tests (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py`, `test_k3_moe_wide.py`) + remain the exhaustive numerics; this entry's test copies their references. +- Gaps (the op is unchanged by this entry): a layer does not record `local_expert_offset`; nothing ties a state to a + stream or checks the inputs' device; the head_flags pairing with the front is the caller's (the entry checks only + that `head` matches the state); the compile cache is a module-level dict (result-neutral). diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.py index f2190722aac7..96ace8337365 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.py @@ -1,94 +1,61 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""Kimi K3's routed experts at decode size: this rank's routed partial from the persistent CuTe DSL kernel ``k3_moe`` -(FC1 + SiTU + FC2 with the routing-weighted combine over the TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers, read in place), -after the routing and MXFP8 quantization of its producer, on caller-owned state: a :class:`K3MoeState` and one -:class:`K3MoeLayer` per MoE layer for up to 8 tokens, a :class:`K3MoeWideState` and one :class:`K3MoeWideLayer` per -layer for up to 64.""" +"""Kimi K3's routed experts at decode size: ``trtllm::k3_moe``, this rank's routed partial from the persistent CuTe +DSL kernel (FC1 + SiTU + FC2 with the routing-weighted combine over the TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers, read in +place), on caller-owned state: a :class:`K3MoeState` (up to 8 tokens) or :class:`K3MoeWideState` (up to 64) and one +:class:`K3MoeLayer` per MoE layer. Its inputs are the outputs of ``moe/k3_route_quant`` or ``moe/k3_moe_front``.""" -from typing import Optional, Tuple +from typing import Optional import torch -# The state types; is_supported reads metadata only. Importing k3_route_quant's op registers trtllm::k3_route_quant. +# The state types (they launch nothing per call); importing the op module registers trtllm::k3_moe. is_supported reads +# metadata only. from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.op import ( K3MoeHeadWorkspace, K3MoeLayer, K3MoeState, - K3MoeWideLayer, K3MoeWideState, is_supported, ) -from tensorrt_llm._torch.cute_dsl_kernels.k3_route_quant import ( - op as _k3_route_quant_op, # noqa: F401 -) __all__ = [ "K3MoeHeadWorkspace", "K3MoeLayer", "K3MoeState", - "K3MoeWideLayer", "K3MoeWideState", "is_supported", "k3_moe", - "k3_moe_fused_front", - "k3_moe_wide", ] def k3_moe( - latent: torch.Tensor, - router_logits: torch.Tensor, - e_score_correction_bias: torch.Tensor, - local_expert_offset: int, - routed_scaling_factor: float, - layer: K3MoeLayer, -) -> torch.Tensor: - """This rank's routed partial ``[M, 3584]`` bf16 for ``M <= 8`` tokens: ``trtllm::k3_route_quant`` of - ``router_logits`` (fp32 ``[M, 896]``) and ``latent`` (bf16 ``[M, 3584]``), then ``k3_moe`` on ``layer``'s experts - (global ids ``[local_expert_offset, local_expert_offset + num_local)``). Writes ``layer``'s state's scratch (left - armed) and ``layer``'s counters (left zero).""" - return layer( - latent, router_logits, e_score_correction_bias, local_expert_offset, routed_scaling_factor - ) - - -def k3_moe_fused_front( - x: torch.Tensor, - w_front: torch.Tensor, - e_score_correction_bias: torch.Tensor, + x_fp8: torch.Tensor, + x_sf: torch.Tensor, + topk_ids: torch.Tensor, + topk_weights: torch.Tensor, local_expert_offset: int, - routed_scaling_factor: float, - shared_cols: int, - gate_cap: float, - linear_cap: float, - head: K3MoeHeadWorkspace, layer: K3MoeLayer, -) -> Tuple[torch.Tensor, torch.Tensor]: - """``(routed partial [M, 3584] bf16, shared activation [M, shared_cols] bf16)`` for the MoE input ``x`` (bf16 - ``[M <= 8, 7168]``): ``trtllm::k3_moe_front`` over ``head`` (see ``moe/k3_moe_front``), then ``k3_moe`` on - ``layer``'s experts. Advances ``head`` by one front call; writes ``layer``'s scratch and counters as - :func:`k3_moe`.""" - return layer.front( - x, w_front, e_score_correction_bias, local_expert_offset, routed_scaling_factor, shared_cols, gate_cap, - linear_cap, head, - ) # fmt: skip - - -def k3_moe_wide( - latent: torch.Tensor, - router_logits: torch.Tensor, - e_score_correction_bias: torch.Tensor, - local_expert_offset: int, - routed_scaling_factor: float, - layer: K3MoeWideLayer, + head: Optional[K3MoeHeadWorkspace] = None, out: Optional[torch.Tensor] = None, ) -> torch.Tensor: - """This rank's routed partial ``[M, 3584]`` bf16 for ``1 <= M <= 64`` tokens: ``trtllm::k3_route_quant`` (its - dependents launched early, as ``k3_moe``'s PDL producer), then the m_max 64 build of ``k3_moe`` on ``layer``'s - experts. ``out``: bf16, contiguous, at least ``[M, 3584]``; its first M rows are the result (a new tensor without - it). Writes ``layer``'s state's scratch (left armed) and ``layer``'s counters (left zero).""" - ids, weights, x_fp8, x_sf = torch.ops.trtllm.k3_route_quant( - router_logits, e_score_correction_bias, latent, routed_scaling_factor, early_trigger=True - ) - return layer(x_fp8, x_sf, ids, weights, local_expert_offset, out=out) + """This rank's routed partial bf16 ``[M, 3584]`` for the routing and MXFP8 latent of ``k3_route_quant`` or + ``k3_moe_front`` (``x_fp8`` float8_e4m3fn ``[M, 3584]``, ``x_sf`` its UE8M0 scales ``[M, 112]``, ``topk_ids`` + int32 and ``topk_weights`` bf16 ``[M, 16]``), over ``layer``'s experts (global ids + ``[local_expert_offset, local_expert_offset + num_local)``); M <= 8 on a K3MoeState, <= 64 on a K3MoeWideState. + + ``head``: the TP group's K3MoeHeadWorkspace, for and only for a ``head_flags`` state's layers: the call then + acquires the front's outputs through the workspace's ready words and advances its epoch, so the front call before + it must have published them (``moe/k3_moe_front`` with ``publish=True``). ``out``: bf16, contiguous, at + least ``[M, 3584]``; the call writes its first M rows and returns an empty ``[0, 3584]`` tensor. Writes the state's + slab (left armed) and partial rows, and the layer's counters (left zero).""" + state = layer.state + if (head is not None) != state.head_flags: + raise ValueError( + "k3_moe: head is given for, and only for, the layers of a head_flags K3MoeState" + ) + return torch.ops.trtllm.k3_moe( + x_fp8, x_sf, topk_ids, topk_weights, *layer.weights, state.c, state.cs, state.part, layer.counters, + local_expert_offset, state.num_local, state.num_ctas, state.m_max, state.use_pdl, + None if head is None else head.ready, None if head is None else head.flags, out, + ) # fmt: skip diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md index 7f4e669ce4d1..ac00b649759c 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md @@ -26,11 +26,11 @@ shared = bf16(gate_cap * tanh(g / gate_cap) * sigmoid(g) * linear_cap * tanh(u ``` `head_weight_r` is rank `r`'s `[WL + WE, K]` head slice and `gate`, `up` this rank's `[shared_cols, K]` shared rows, -all inside the ranks' `w_front`. `k3_route_quant` is the routing and quantization of `moe/k3_moe` (top-16 of the -sigmoid plus the bias, ties to the lower id, weights renormalized times `routed_scaling_factor`; MXFP8 with one UE8M0 -scale per 32 columns): the front selects with all warps of a CTA but returns `top16_warp`'s experts, order and weight -bits, and quantizes with the same device code (the kernel's statement). Every head and shared value is an fp32 sum -over `K` split across the 8 CTAs of a cluster, the 8 partials added in cluster-rank order from +0.0 (the kernel's +all inside the ranks' `w_front`. `k3_route_quant` is the routing and quantization of `moe/k3_route_quant` (top-16 of +the sigmoid plus the bias, ties to the lower id, weights renormalized times `routed_scaling_factor`; MXFP8 with one +UE8M0 scale per 32 columns): the front selects with all warps of a CTA but returns `top16_warp`'s experts, order and +weight bits, and quantizes with the same device code (the kernel's statement). Every head and shared value is an fp32 +sum over `K` split across the 8 CTAs of a cluster, the 8 partials added in cluster-rank order from +0.0 (the kernel's statement): deterministic, but not the summation order of another GEMM. Certified at every `M` 1-8 and for every front call the matrix makes alone, with payloads whose head and shared sums @@ -48,9 +48,9 @@ another GEMM's in the last bit; the kernel test (`test_k3_moe_front.py`) bounds the 16th and 17th selection keys are within 1e-4, more than 99.9 % of the MXFP8 codes and scales equal. Fusion boundary. Inside: the head GEMV, the head all-gather, the routing, the MXFP8 latent, the shared gate_up and -its SiTU. Outside: the producer of the MoE input `x`; the routed experts (`moe/k3_moe`'s `k3_moe_fused_front` runs -`k3_moe` on this op's outputs within the same call); the shared experts' down projection and its all-reduce; the -routed partials' all-reduce and the latent-up projection. +its SiTU. Outside: the producer of the MoE input `x`; the routed experts (`moe/k3_moe`, called on this op's routing +and MXFP8 latent right after it); the shared experts' down projection and its all-reduce; the routed partials' +all-reduce and the latent-up projection. ## Signature @@ -64,6 +64,7 @@ def k3_moe_front( gate_cap: float, linear_cap: float, workspace: K3MoeHeadWorkspace, + publish: bool = False, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor] ``` @@ -83,11 +84,12 @@ The wrapper module also exports the load-time helpers `front_weight(head_weight, | `shared_cols` | 384 (Kimi K3 TP16's per-rank width: two shared experts of 3072 over 16 ranks) | Python int | — | — | | `gate_cap`, `linear_cap` | 4.0, 25.0 (Kimi K3's SiTU caps) | Python float | — | — | | `workspace` | the `K3MoeHeadWorkspace` of this rank's TP group, `W` = 4 certified (see *State*) | — | — | — | +| `publish` | `True`: also release the workspace's per-token ready words for the head_flags `moe/k3_moe` call that follows (*State*); `False` (default): leave them and the epoch untouched | Python bool | — | — | | returns | `topk_ids [M, 16]` int32, `topk_weights [M, 16]` bf16, `quantized [M, 3584]` float8_e4m3fn, `scales [M, 112]` uint8 (linear: one byte per 32 columns), `shared [M, shared_cols]` bf16 | — | contiguous, newly allocated | `x.device` | -Inert (not exposed by the wrapper): `ring` (4, the weight ring's stages) and `ag_ready` (`None`: the standalone front -publishes no ready words; `K3MoeLayer.front` on a head_flags state passes the workspace's `ready`, see *State*). -`gate_cap` and `linear_cap` are compile-time constants of the kernel: each pair compiles once (*Metadata consumed*). +Inert (not exposed by the wrapper): `ring` (4, the weight ring's stages). `publish` passes the workspace's ready +words as the op's `ag_ready`. `gate_cap` and `linear_cap` are compile-time constants of the kernel: each pair +compiles once (*Metadata consumed*). ## State @@ -108,29 +110,38 @@ at `W` 4, 8 and 16 (certified at 4): advanced only by a head_flags build of `k3_moe` (`moe/k3_moe`); `[3]` the sign-ins of the CTAs that read `[0]` in the current call. - `ready`, int32 `[32]`: `[t]` token `t`'s routing and `[8 + t]` its MXFP8 row, released as `epoch + 1` by a - publishing front (`K3MoeLayer.front` on a head_flags state); `[16, 32)` unused. + publishing front (this entry with `publish=True`); `[16, 32)` unused. - `rank`, `world_size`; `handle`, the `McastGPUBuffer` that owns the memory (the workspace is valid while this object lives); `comm`, the TP-group communicator the handles were exchanged over. The size depends on `W` only, not on `M`: every call fits. **Who creates it, and when.** The target, in `post_load_weights`, with `K3MoeHeadWorkspace.create(mapping, -fabric_handle=None)`: collective over `mapping`'s TP group (every rank calls it at the same point; it returns on every -rank or raises on every rank, the agreement also being the barrier that keeps any rank from pushing into a peer's -buffer before the peer has emptied it); eager: under CUDA-graph capture it raises `RuntimeError` before entering the -collective (certified, on every rank). It empties every word and zeroes `flags` and `ready` (certified, both of the -matrix's workspaces). `fabric_handle`: share the memory by fabric handle (required across nodes) or POSIX file -descriptor; default `mapping.is_multi_node()`. No environment variable is read. - -**Which ops may share one object.** Every MoE front call of the TP group: this entry and `moe/k3_moe`'s -`k3_moe_fused_front`, on a plain or a head_flags `K3MoeState`. They form one sequence on the workspace: mixed in one -step (the fused front on each state, then the front alone) they all stay correct (certified, below). Separate from the -MNNVL all-reduce workspace and the sandwich workspace: a front call advances neither. Two workspaces are two -independent rotations and two independent epochs: 20 calls alternating between two workspaces in an irregular -pattern, kinds mixed, each return the bits of the same call made alone, and each workspace's epoch advances by its -own head_flags calls only (certified). They are not independent orders: each call spins until its peers' pushes of -the same call arrive and the calls of one stream run one after the other, so ranks that order calls on two -workspaces differently on one stream would deadlock (the kernel's design; not exercised). +fabric_handle=None)`: collective over `mapping`'s TP group and eager, every rank of the group calling it at the same +point. Under MPI each call first splits the group's communicator off the session's, a collective of every rank of the +session (`_get_mnnvl_workspace_comm`). Before allocating, the ranks agree that each of them can: not capturing a CUDA +graph, the buffer within its device's free memory. If one cannot, every rank raises `RuntimeError` ("not every rank +can allocate") and none allocates. A failure returned by the allocation is agreed the same way, and that second +agreement is also the barrier that keeps any rank from pushing into a peer's buffer before the peer has emptied it; a +rank that fails inside the allocation's handle exchange can leave its peers waiting there (the op module's +statements). Certified: with every rank capturing, and with the last rank capturing while its peers call it eagerly, +every rank raises `RuntimeError` (the capturing ranks' messages naming the capture) and the existing workspace keeps +its bits. So even a refusal needs every rank to call `create()`: a rank calling it alone waits for its peers. It +empties every word and zeroes `flags` and `ready` (certified, both of the matrix's workspaces). `fabric_handle`: share +the memory by fabric handle (required across nodes) or POSIX file descriptor; default `mapping.is_multi_node()`. No +environment variable is read. + +**Which ops may share one object.** Every MoE front call of the TP group: this entry, plain or publishing +(`publish=True`, the front a head_flags `k3_moe` pairs with). The head_flags calls of `moe/k3_moe` (`head=workspace`) +read and advance its epoch `flags[2]` and its ready words (the kernel's statement). They form one sequence on the +workspace: mixed in one step (the front then `k3_moe` on a plain state, the publishing front then `k3_moe` on a +head_flags state, the front alone) they all stay correct (certified, below). Separate from the MNNVL all-reduce +workspace and the sandwich workspace: a front call advances neither. Two workspaces are two independent rotations and +two independent epochs: 20 calls alternating between two workspaces in an irregular pattern, kinds mixed, each return +the bits of the same call made alone, and each workspace's epoch advances by its own head_flags calls only +(certified). They are not independent orders: each call spins until its peers' pushes of the same call arrive and the +calls of one stream run one after the other, so ranks that order calls on two workspaces differently on one stream +would deadlock (the kernel's design; not exercised). **Call-order invariant.** Every rank of the group makes the same sequence of front calls on one workspace (the same number of calls, the `k`-th with the same `M`), eager calls and graph replays alike, with the same `x`; and on one @@ -139,7 +150,8 @@ pushes its latent slice and router partials into that buffer on every rank, poll pushes of this call are there, and empties what it read. Every CTA that reads `flags[0]` signs in on `flags[3]` right after the read; the CTA that flips `flags[0]` for the next call waits for all of them and zeroes `flags[3]` (the kernel's statement; certified: `flags[0]` flips once per call, and `flags[3]` is zero after each single call and -each sequence). +each sequence). Each publishing front call is followed by exactly one head_flags `k3_moe` call on the same workspace +(`moe/k3_moe`, *Preconditions*). **What a later launch reads.** `flags[0]`, and its buffer's words, which must be empty except for this call's pushes; `flags[3]` at zero. A publishing front also reads the epoch `flags[2]`, before it lets its dependent launch (the @@ -150,10 +162,10 @@ words of the same call's tokens. The next push into a buffer comes from a call t after this call has ended on this rank (the kernel's statement, from the stream order and the grid-dependency waits). So no separate clear and no record of the previous call's size is needed, and a call after a smaller one finds no word of an older, larger one. Certified: every word of this rank's buffers empty and `flags[1]`, `flags[3]` zero -after each single call and after each sequence below; decode steps of three layers (the fused front on the plain -state, the fused front on the head_flags state, the front alone) at `M` 8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8 with new -inputs every call, back to back, a random rank 5 ms late before every call: every call the bits of the same call made -alone (itself checked against the references). +after each single call and after each sequence below; decode steps of three layers (the front then `k3_moe` on the +plain state, the publishing front then `k3_moe` on the head_flags state, the front alone) at `M` 8, 8, 8, 2, 7, 8, 1, +1, 8, 3, 8 with new inputs every call, back to back, a random rank 5 ms late before every call: every call the bits of +the same call made alone (itself checked against the references). The ready words are re-armed by the head_flags `k3_moe` (`moe/k3_moe`, *State*): at its last tile claim it writes `epoch + 1` into the words of the tokens past its `M` and into `flags[2]`, so after every head_flags call `flags[2]` @@ -163,9 +175,9 @@ at it; after every back-to-back sequence and every replayed step, the epoch adva calls with all 16 words at it; across -1 -> 0 (from zeroed words and epoch 0, calls at `M` 1, 1, then the epoch preset to -2 and calls at `M` 1, 8, 3, 8: the `M` 8 call at epoch -1 waits for 0, the value a word past an earlier call's tokens would still hold without that re-arm) and across 2^31 - 1 -> -2^31 (epoch preset to 2^31 - 2, calls at -`M` 8, 1, 8), each call's outputs the plain build's bits. The standalone front leaves `flags[2]` and the ready words -untouched (certified). Nothing else may write them: a caller that resets `flags[2]` while the words hold -`epoch + 1` would let the next head_flags `k3_moe` read the front's outputs before the front writes them (not +`M` 8, 1, 8), each call's outputs the plain build's bits. A front call without `publish` leaves `flags[2]` and +the ready words untouched (certified). Nothing else may write them: a caller that resets `flags[2]` while the words +hold `epoch + 1` would let the next head_flags `k3_moe` read the front's outputs before the front writes them (not exercised). **What a wrong order does.** Certified at `W` = 4 (the negative control): rank 0 issues two same-shaped front calls @@ -190,9 +202,16 @@ Besides `workspace` (an explicit argument): - `TRTLLM_ENABLE_PDL` (default on), read on every call and part of the key; it changes scheduling, not results (the op's statement). - The device's cluster capacity (`max_clusters`, cached per device), which sizes the grid and picks the head's - geometry: 128-row tiles, or one round of 64-row half-tiles when they fit next to the shared tiles - (`half_geometry`; the op's statement: TP16). At `W` 4 the head (1120 rows per rank, 18 half-tiles) does not fit in - one round, so the 4-rank run uses 128-row tiles. + geometry (the kernel module's `geometry` and `half_geometry`): one round of 64-row half-tiles, every head k-tile of + a CTA on chip before the grid-dependency wait, when they fit beside the shared tiles in the clusters left next to + the two role clusters; else 128-row tiles. A GB200 holds 15 clusters of 8 CTAs (the kernel's statement), so with + 384 shared columns (6 tiles) the half-tile path runs at `W` 16 only: 280 head rows per rank make 5 half-tiles, 11 + GEMV clusters of the 13 left. At `W` 8 (560 rows, 9 half-tiles) and `W` 4 (1120 rows, 18) they do not fit, and the + head runs as 5 and 9 tiles of 128 rows. The CPU test + `tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front_geometry.py` checks this choice and that + `front_weight`'s rows are the rows each plan reads. So this entry's 4-rank matrix runs 128-row tiles only; the + half-tile path's 16-rank record is the kernel test `test_k3_moe_front.py` run as 16 processes (the model's TP16 + shapes), not this matrix. ## Preconditions @@ -206,26 +225,26 @@ Besides `workspace` (an explicit argument): invariant. Nothing checks that `x` agrees across ranks: a rank with another `x` mixes its slice into every rank's result (the negative control shows that effect). - `workspace` was created, and the kernel compiled (one eager call per configuration), before any capture. Calls may - be captured: certified with a captured step of three calls at `M` 8 (the fused front on the plain and on the - head_flags state, the front alone) replayed 6 times with rewritten inputs, an eager call of another `M` on the same - workspace between replays, every replayed and eager call the bits of the same call made alone, the epoch advanced - once per head_flags call, replayed or eager. + be captured: certified with a captured step of three calls at `M` 8 (the front then `k3_moe` on the plain state, + the publishing front then `k3_moe` on the head_flags state, the front alone) replayed 6 times with rewritten + inputs, an eager call of another `M` on the same workspace between replays, every replayed and eager call the bits + of the same call made alone, the epoch advanced once per head_flags call, replayed or eager. ## Notes - Certified path: 4 ranks of one GB200 tray (sm_100), one rank per GPU, the head sharded over those 4 ranks (1120 rows per rank, 128-row tiles), the shared activation at TP16's per-rank width (384). Kimi K3 TP16 shards the head - over 16 ranks on four trays (280 rows per rank), which runs the half-tile geometry and which only a 16-rank run - reaches; the matrix takes `--world-size` and `--launcher`, and that receipt is pending. + over 16 ranks on four trays (280 rows per rank, the half-tile geometry: *Metadata consumed*); the matrix takes + `--world-size` and `--launcher`, and its 16-rank receipt is pending. - Test: `tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py` (rank body), collected by `tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_front_op_matrix.py`. It also certifies - `moe/k3_moe`'s `k3_moe_fused_front` cells. + `moe/k3_moe` behind the front, on a plain and on a head_flags `K3MoeState`. - Reference: native torch for the head (every rank's slice in fp64, exact for the payloads, then fp32; the latent columns rounded to bf16) and the shared gate_up (fp64, exact, rounded to bf16) with SiTU in fp32; the stock `trtllm::kimi_k3_noaux_tc_mxfp8_quant` for the routing and quantization of the gathered head. Every rank draws every - rank's head slice from one seed, so each holds the whole reference. The routed experts of the fused cells are - random checkpoint-format MXFP4 (224 per rank, rank `r` at global ids `[224 (r % 4), 224 (r % 4) + 224)`), the - reference for `y` the stock TRTLLM-Gen runner. The kernel test + rank's head slice from one seed, so each holds the whole reference. The routed experts behind the front are random + checkpoint-format MXFP4 (224 per rank, rank `r` at global ids `[224 (r % 4), 224 (r % 4) + 224)`), the reference + for `y` the stock TRTLLM-Gen runner. The kernel test (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py`) covers Gaussian payloads and races the publishing front's epoch read against `k3_moe`'s epoch advance (`check_publish_order`); neither is repeated here. - `mutates_args` names every buffer the op writes (`ag_uc`, `ag_mc`, `ag_flags`, `ag_ready`). Gaps (the op is diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.py index 0d238c29e28f..05bf3f55e37b 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.py @@ -27,12 +27,16 @@ def k3_moe_front( gate_cap: float, linear_cap: float, workspace: K3MoeHeadWorkspace, + publish: bool = False, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: """Return ``(topk_ids, topk_weights, quantized, scales, shared)`` for the MoE input ``x`` (bf16 ``[M <= 8, 7168]``, the same on every rank): the top-16 routing of the gathered router logits with ``e_score_correction_bias`` and the MXFP8 latent with its UE8M0 scales, as ``trtllm::k3_route_quant`` returns them for the gathered head, and the shared experts' activation (bf16 ``[M, shared_cols]``). ``w_front`` from :func:`front_weight`. Advances ``workspace`` by - one call: every rank of the group makes the same front calls on it in the same order.""" + one call: every rank of the group makes the same front calls on it in the same order. + + ``publish``: also release the workspace's per-token ready words, which the next ``moe/k3_moe`` call on a head_flags + state (``head=workspace``) acquires; every publishing call is followed by exactly one such call.""" return torch.ops.trtllm.k3_moe_front( x, w_front, @@ -46,4 +50,5 @@ def k3_moe_front( workspace.flags, workspace.rank, workspace.world_size, + ag_ready=workspace.ready if publish else None, ) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_route_quant.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_route_quant.md new file mode 100644 index 000000000000..3d3f5f1851d8 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_route_quant.md @@ -0,0 +1,107 @@ +--- +receipts: + sm_100: {status: pending, tests: 18} +--- + +# k3_route_quant + +**Wraps** `torch.ops.trtllm.k3_route_quant` (one call). + +## Semantics + +Kimi K3's top-16 routing and the MXFP8 quantization of the routed latent, for a decode batch of up to 64 tokens, in +one kernel: the CuTe DSL form of `trtllm::kimi_k3_noaux_tc_mxfp8_quant`, whose four outputs it returns bit for bit +(certified, below). Per token `t`, as the kernel module states it: + +``` +s = 0.5 * tanhf(0.5 * router_logits[t]) + 0.5 # fp32, [896] +ids = the 16 experts with the largest s + e_score_correction_bias, descending, ties to the lower id +weights = bf16(s[ids] * routed_scaling_factor / (sum of the 16 s[ids] + 1e-20)) + # the division in fp64; the sum in fp32, in the order of a 16-lane xor-butterfly warp reduction +x_fp8, x_sf = MXFP8(latent[t]): one UE8M0 scale per 32 columns, rounded up from amax / 448 + (the recipe of cvt_warp_fp16_to_mxfp8), and the e4m3 codes of the scaled values +``` + +It returns `(topk_ids, topk_weights, quantized, scales)`: the global expert ids (int32 `[M, 16]`, in selection +order), their routing weights (bf16 `[M, 16]`), the latent's e4m3 codes (`[M, 3584]`) and its scales (uint8 `[M, +112]`, linear: byte `t * 112 + b` scales columns `[32 b, 32 b + 32)` of token `t`). The bias enters the selection +only; the weights are the unbiased sigmoids, renormalized over the 16 and scaled. The kernel selects differently from +the C++ one (one warp per token, each lane sorting its 28 keys, 16 rounds of warp max and min reductions; the +kernel's statement); the results are the same bits. + +Certified, bit for bit against `trtllm::kimi_k3_noaux_tc_mxfp8_quant`, with and without the early trigger: at `M` +1-8, 16, 33 and 64 with random logits, and at `M` 1, 3 and 8 for 40 tied selection keys (ties to the lower id), equal +logits, huge logits (saturated sigmoids) and zero, large and denormal latent rows. Run-to-run bit identical, and each +`M`'s outputs the bits of the same rows of the 64-token call. + +Fusion boundary. Inside: the sigmoid, the bias-corrected top-16 selection, the renormalized weights, the MXFP8 +quantization of the latent. Outside: the router GEMM that produces the logits and the latent-down GEMM that produces +the latent (`moe/k3_moe_front` fuses both, the head all-gather and this op's device code into one kernel); the routed +experts (`moe/k3_moe`, which consumes these outputs). + +## Signature + +```python +def k3_route_quant( + router_logits: torch.Tensor, + e_score_correction_bias: torch.Tensor, + latent: torch.Tensor, + routed_scaling_factor: float, + early_trigger: bool = False, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor] +``` + +The op's schema has the same arguments (`scores`, `bias`, `hidden_states`, `routed_scaling_factor`, +`early_trigger`); `mutates_args=()`: it writes only its new outputs. + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `router_logits` | `[M, 896]`, `M` 1-64 | fp32 | contiguous | CUDA | +| `e_score_correction_bias` | `[896]` | fp32 | contiguous | CUDA, the same device | +| `latent` | `[M, 3584]`, the same `M` | bf16 | contiguous | CUDA, the same device | +| `routed_scaling_factor` | scalar (2.827 certified) | Python float | — | — | +| `early_trigger` | `False`, `True` | Python bool | — | — | +| returns | `topk_ids [M, 16]` int32, `topk_weights [M, 16]` bf16, `quantized [M, 3584]` float8_e4m3fn, `scales [M, 112]` uint8 | — | contiguous, newly allocated | `router_logits.device` | + +`early_trigger` lets the next kernel on the stream, launched as a programmatic dependent, start as soon as every CTA +of this one has passed its own grid-dependency wait, before any output is written; without it, each CTA lets it start +once it has written its outputs, behind a GPU-scope fence (the kernel's code). A dependent launched early must wait +for this grid (`griddepcontrol.wait`) before it reads the outputs, as `moe/k3_moe` does: a caller sets it when the +next kernel is `k3_moe` launched as a programmatic dependent (a state with `use_pdl`). The outputs are the same bits +either way (certified). + +## Metadata consumed + +Stateless: no state object, nothing kept from one call to the next. Process state: + +- A cache of compiled kernels keyed by (early trigger, PDL). The first call of a key compiles (seconds) and must be + eager: under capture it raises `RuntimeError` ("must run once outside CUDA-graph capture first") instead of + compiling, certified for both early-trigger builds with the cache cold. Result-neutral. +- `TRTLLM_ENABLE_PDL` (default on), read on every call and part of the key: launch the kernel as a programmatic + dependent of the kernel before it. It changes scheduling, not results (the op's statement). + +## Preconditions + +- CUDA tensors on one device; fp32 logits and bias, bf16 latent; all contiguous; `router_logits` `[M, 896]`, the bias + 896 elements, `latent` `[M, 3584]` with the same `M`; 1 <= `M` <= 64. Anything else raises `ValueError` before any + launch (the op's check). Certified: `M` 0 and 65, a strided view of the logits, bf16 logits, an fp16 latent, 895 + biases and a latent with another `M` raise `ValueError`, and the next call is correct. +- Before its own grid-dependency wait the kernel reads only `e_score_correction_bias`, a weight (the kernel's code): + the kernel before it must not be writing the bias when this one launches early. +- The first call of each build eagerly (*Metadata consumed*). Calls may be captured: certified with a captured + sequence at `M` 8, 3 and 64 (early trigger on, off, on) replayed 4 times with rewritten inputs, an eager call of + another `M` between replays, every replayed and eager call the bits of the same call made alone. +- Certified on sm_100 only; the stock op it is compared with requires an SM 10.x device. + +## Notes + +- Certified path: one GB200 GPU (sm_100), test + `tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_route_quant.py`. Reference: the stock + `trtllm::kimi_k3_noaux_tc_mxfp8_quant`, bit for bit. The kernel test + (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_route_quant.py`) also compares with the unfused chain + (`trtllm::noaux_tc_op`, then `trtllm::mxfp8_quantize`) and reports a stable PyTorch sort of sigmoid + bias. +- Users: `moe/k3_moe` consumes the four outputs directly; `moe/k3_moe_front` runs this op's device code on the + gathered head (and so returns the same bits for the same logits and latent). +- The compile cache is a module-level dict (result-neutral, above). diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_route_quant.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_route_quant.py new file mode 100644 index 000000000000..618e2ec846a2 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_route_quant.py @@ -0,0 +1,39 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3's top-16 routing and MXFP8 latent quantization at decode size, ``trtllm::k3_route_quant``: the CuTe DSL +form of ``trtllm::kimi_k3_noaux_tc_mxfp8_quant``, the same four outputs bit for bit, for up to 64 tokens.""" + +from typing import Tuple + +import torch + +# Importing the op module registers trtllm::k3_route_quant; CuTe DSL is imported on its first call. +from tensorrt_llm._torch.cute_dsl_kernels.k3_route_quant import ( + op as _k3_route_quant_op, # noqa: F401 +) + +__all__ = ["k3_route_quant"] + + +def k3_route_quant( + router_logits: torch.Tensor, + e_score_correction_bias: torch.Tensor, + latent: torch.Tensor, + routed_scaling_factor: float, + early_trigger: bool = False, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """``(topk_ids, topk_weights, quantized, scales)`` for ``router_logits`` (fp32 ``[M, 896]``) and ``latent`` (bf16 + ``[M, 3584]``), 1 <= M <= 64: the 16 experts with the largest sigmoid + ``e_score_correction_bias`` (int32 + ``[M, 16]``), their unbiased sigmoids renormalized times ``routed_scaling_factor`` (bf16 ``[M, 16]``), and the + latent as MXFP8 (float8_e4m3fn ``[M, 3584]``, one UE8M0 scale per 32 columns, uint8 ``[M, 112]``). + + ``early_trigger``: let the next kernel launch (as a programmatic dependent) as soon as every CTA has passed its own + grid-dependency wait, before the outputs are written; for a dependent that waits for this whole grid before + reading them, as ``trtllm::k3_moe`` does.""" + return torch.ops.trtllm.k3_route_quant( + router_logits, + e_score_correction_bias, + latent, + routed_scaling_factor, + early_trigger=early_trigger, + ) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py index 85c8ff3f3e59..8f3a118ca86b 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py @@ -162,10 +162,10 @@ def create_mcast_state(name: str, mapping, words: int, fabric_handle: Optional[b @dataclass(eq=False) class K3MoeHeadWorkspace: - """One TP group's MoE head all-gather buffers, read and written by ``trtllm::k3_moe_front`` (alone, or as the - producer of :meth:`K3MoeLayer.front`): two alternating Lamport buffers of every rank's head slice per token behind - one multicast mapping, then the front's router partials; the flag words that rotate them; and the per-token ready - words a ``head_flags`` build of k3_moe acquires. Every front call on it takes the next buffer, so all of a group's + """One TP group's MoE head all-gather buffers, read and written by ``trtllm::k3_moe_front``: two alternating Lamport + buffers of every rank's head slice per token behind one multicast mapping, then the front's router partials; the + flag words that rotate them; and the per-token ready words a publishing front releases and the ``head_flags`` + build of ``trtllm::k3_moe`` after it acquires. Every front call on it takes the next buffer, so all of a group's ranks make the same front calls on it in the same order. Pass ``uc``, ``mc``, ``flags``, ``rank`` and ``world_size`` as the front's ``ag_uc``, ``ag_mc``, ``ag_flags``, ``ag_rank`` and ``ag_world``, and ``ready`` as its ``ag_ready``. Separate from the MNNVL all-reduce workspace.""" @@ -217,8 +217,10 @@ def build(uc, mc, handle, comm): # read only the ready words and head flags (head_flags), so the others are given a stand-in they never touch. _OPTIONAL_ARGS = 11 # (alignment, leading dim) of every tensor argument, in the kernel's order. -_ALIGNS = [16, 16, 16, 16, 16, 4, 16, 16, 16, 16, 16, 16, 16, 16, 16, 4, 4, 4] + [16] * _OPTIONAL_ARGS -_LEADING = [0, 0, 0, 1, 2, 2, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 0] + [0] * _OPTIONAL_ARGS +_ALIGNS = [16, 16, 16, 16, 16, 4, 16, 16, 16, 16, 16, 16, 16, 16, 16, 4, 4, 4] +_ALIGNS += [16] * _OPTIONAL_ARGS +_LEADING = [0, 0, 0, 1, 2, 2, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 0] +_LEADING += [0] * _OPTIONAL_ARGS _compiled: Dict[tuple, object] = {} @@ -228,7 +230,9 @@ def _part_rows(mod) -> int: return mod.PART_ROWS if mod.WIDE else mod.G_CAP * _TOKEN_SLOTS -def _config(i_tp: int, num_ctas: int, num_local: int, m_max: int, use_pdl: bool, head_flags: bool) -> dict: +def _config( + i_tp: int, num_ctas: int, num_local: int, m_max: int, use_pdl: bool, head_flags: bool +) -> dict: """The kernel options of one build (trace-time constants).""" return { "i_tp": i_tp, @@ -370,8 +374,13 @@ def __call__( head_flags: Optional[torch.Tensor] = None, out: Optional[torch.Tensor] = None, ) -> torch.Tensor: - """``trtllm::k3_moe`` on this layer's experts, counters and its state's scratch: see the op.""" + """``trtllm::k3_moe`` on this layer's experts, counters and its state's scratch: see the op. ``head_ready`` and + ``head_flags`` are given for, and only for, the layers of a ``head_flags`` state.""" st = self.state + if (head_ready is not None) != st.head_flags: + raise ValueError( + "K3MoeLayer: head_ready / head_flags are given for, and only for, a head_flags state's layers" + ) y = torch.ops.trtllm.k3_moe( x_fp8, x_sf, topk_ids, topk_weights, *self.weights, st.c, st.cs, st.part, self.counters, local_expert_offset, st.num_local, st.num_ctas, st.m_max, st.use_pdl, head_ready, head_flags, out, @@ -478,14 +487,18 @@ def k3_moe( or tuple(counters.shape) != (mod.NUM_STATE,) or not all(t.is_contiguous() for t in (c, cs, part, counters)) ): - raise ValueError("k3_moe: c / cs / part / counters are not this build's scratch and layer counters") + raise ValueError( + "k3_moe: c / cs / part / counters are not this build's scratch and layer counters" + ) if head and ( head_ready.dtype != torch.int32 or head_ready.numel() < 2 * MAX_TOKENS or head_flags.dtype != torch.int32 or head_flags.numel() < 3 ): - raise ValueError("k3_moe: head_ready / head_flags must be a K3MoeHeadWorkspace's ready and flags") + raise ValueError( + "k3_moe: head_ready / head_flags must be a K3MoeHeadWorkspace's ready and flags" + ) if out is None: y = torch.empty(num_tokens, HIDDEN_SIZE, dtype=torch.bfloat16, device=x_fp8.device) else: diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index c5d5bcc35668..4c7f1e397e61 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -318,11 +318,12 @@ l0_b200: - kv_cache/test_kv_cache_iteration_stats.py::TestKvCacheIterationStats::test_field_completeness # ------------- Prefix-aware scheduling E2E tests --------------- - kv_cache/test_prefix_aware_scheduling.py::TestServePrefixAwareScheduling::test_multi_round_qa_shared_prefix_smoke - # ------------- Kimi K3 decode MoE kernels and their catalog entry (sm_100) --------------- + # ------------- Kimi K3 decode MoE kernels and their catalog entries (sm_100) --------------- - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_wide.py - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_route_quant.py - unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe.py + - unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_route_quant.py # ------------- Visual Gen tests --------------- - unittest/_torch/cute_dsl_kernels/test_nvfp4_conv3d.py - unittest/_torch/visual_gen/test_media_decode.py diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py index 281e93dc9f1a..47a70551a3b9 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py @@ -12,8 +12,9 @@ # 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. -"""K3MoeLayer (trtllm::k3_route_quant + the persistent k3_moe kernel, M <= 8) on one GPU, at the Kimi K3 TP16 -deployment's routed-expert rank layout (experts TP4 x EP4: 224 local experts, intermediate 768 per rank), at every M +"""trtllm::k3_route_quant, then trtllm::k3_moe through a K3MoeLayer (the persistent k3_moe kernel, M <= 8) on one GPU, +at the Kimi K3 TP16 deployment's routed-expert rank layout (experts TP4 x EP4: 224 local experts, intermediate 768 +per rank), at every M in 1..8 with random routing, with 0, 4 and 16 of each token's experts local, and with 16 local experts per token none shared (16 M groups: the kernel's group capacity at M = 8): against the stock path (trtllm::kimi_k3_noaux_tc_mxfp8_quant then the TRTLLM-Gen W4A8_MXFP4_MXFP8 MoE runner with those ids) and an fp32 @@ -54,6 +55,7 @@ def _is_sm100() -> bool: def _ops(): import tensorrt_llm # noqa: F401 from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op # noqa: F401 + from tensorrt_llm._torch.cute_dsl_kernels.k3_route_quant import op as _rq # noqa: F401 return torch.ops.trtllm @@ -207,7 +209,10 @@ def _layer(): def _fused(proc, bias, x, logits): - return _layer()(x, logits, bias, OFFSET, RSF) + """trtllm::k3_route_quant (its dependents launched early: k3_moe is its programmatic dependent), then + trtllm::k3_moe on _layer().""" + ids, w, x_fp8, x_sf = _ops().k3_route_quant(logits, bias, x, RSF, early_trigger=True) + return _layer()(x_fp8, x_sf, ids, w, OFFSET) def _scratch_rearmed(): @@ -350,8 +355,8 @@ def test_token_limit(): def test_collective_workspaces_refuse_graph_capture(): """The head all-gather's buffers (the front's) are created collectively (an MNNVL multicast allocation over the TP - group): creating them under CUDA-graph capture raises instead of entering the collective. The per-rank state and - a layer's counters refuse capture too (they allocate).""" + group): under CUDA-graph capture create() raises (the group's ranks agree first that each can allocate; a group of + one here) instead of allocating. The per-rank state and a layer's counters refuse capture too (they allocate).""" from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op from tensorrt_llm.mapping import Mapping diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py index dec3f625264c..f3e046826e1f 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py @@ -12,12 +12,13 @@ # 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. -"""trtllm::k3_moe_front and K3MoeLayer.front (the Kimi K3 MoE front: sharded head GEMV, head all-gather, top-16 -routing, MXFP8 latent, shared gate_up + SiTU; then k3_moe on its grid), one process per GPU over the TP group of this -run, at every M in 1..8, over one K3MoeHeadWorkspace. The head is sharded over the group (TP W: 3584 / W latent + -896 / W router rows and 2 x 6144 / W shared rows per rank; W = 4 on one GB200 tray, the model's TP16 shapes with 16 -processes); the routed experts are one rank of experts TP4 x EP4 (224 local experts, intermediate 768), as in the TP16 -deployment. +"""trtllm::k3_moe_front (the Kimi K3 MoE front: sharded head GEMV, head all-gather, top-16 routing, MXFP8 latent, +shared gate_up + SiTU), alone and followed by trtllm::k3_moe through a K3MoeLayer (the plain build waits for the +front's grid; the head_flags build acquires the ready words the front publishes with ag_ready), one process per GPU +over the TP group of this run, at every M in 1..8, over one K3MoeHeadWorkspace. The head is sharded over the group +(TP W: 3584 / W latent + 896 / W router rows and 2 x 6144 / W shared rows per rank; W = 4 on one GB200 tray, the +model's TP16 shapes with 16 processes); the routed experts are one rank of experts TP4 x EP4 (224 local experts, +intermediate 768), as in the TP16 deployment. front : against the unfused chain (the head GEMV in fp32 torch -> the gather -> trtllm::kimi_k3_noaux_tc_mxfp8_quant; shared: cuBLAS gate_up -> trtllm::situ_and_mul): top-16 ids per token (a mismatch only at a reference @@ -169,13 +170,14 @@ def _context(with_experts): return ctx -def _layer(ctx, head_flags=False, config=None): - """This rank's experts as a layer of a new K3MoeState: the plain build, or the head_flags build.""" +def _layer(ctx, head_flags=False, num_ctas=None): + """This rank's experts as a layer of a new K3MoeState: the plain build, or the head_flags build; ``num_ctas`` caps + k3_moe's persistent grid (default one CTA per SM).""" from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op as moe_op p = ctx.experts state = moe_op.K3MoeState(torch.device("cuda", torch.cuda.current_device()), I_TP, E_LOCAL, head_flags=head_flags, - config=config) # fmt: skip + num_ctas=num_ctas) # fmt: skip return state.layer(p["w31"], p["w31s"], p["w2"], p["w2s"]) @@ -266,10 +268,18 @@ def check_front(ctx): def _fused(ctx, x, layer=None, bias=None): - """K3MoeLayer.front on ``layer`` (default: the plain build's ``ctx.layer``).""" + """trtllm::k3_moe_front, then trtllm::k3_moe on ``layer`` (default: the plain build's ``ctx.layer``); returns + (y, shared). A head_flags layer's front publishes the ready words (``ag_ready``) and its k3_moe acquires them.""" layer = layer or ctx.layer - return layer.front(x, ctx.front, ctx.bias if bias is None else bias, ctx.offset, RSF, ctx.inter, GATE_CAP, - LINEAR_CAP, ctx.ws) # fmt: skip + head = layer.state.head_flags + ids, w, q, s, shared = torch.ops.trtllm.k3_moe_front( + x, ctx.front, ctx.bias if bias is None else bias, RSF, ctx.inter, GATE_CAP, LINEAR_CAP, *ctx.ag, ctx.world, + ag_ready=ctx.ws.ready if head else None, + ) # fmt: skip + if not head: + return layer(q, s, ids, w, ctx.offset), shared + y = layer(q, s, ids, w, ctx.offset, head_ready=ctx.ws.ready, head_flags=ctx.ws.flags) + return y, shared def _runner(ctx, ids, w, q, s): @@ -349,8 +359,8 @@ def _i32(v: int) -> int: def check_head_flags(ctx): - """K3MoeLayer.front with the ready-word handoff (``ag_ready``: k3_moe built with head_flags acquires the front's - ready words, ready[t] / ready[8 + t] = the head epoch flags[2] + 1 for token t, instead of waiting for its grid) + """The front, then k3_moe with the ready-word handoff (``ag_ready``: k3_moe built with head_flags acquires the + front's ready words, ready[t] / ready[8 + t] = the head epoch flags[2] + 1 for token t, instead of its grid) across the epoch's int32 wrap. From a new workspace's state (epoch 0, ready words 0), two calls at M 1, then the epoch preset to -2, then calls at M 1, 8, 3, 8: the M 8 call at epoch -1 waits for 0, the value of the words that no call has published. Per call: no word the call polls already holds its epoch + 1 (such a word would let k3_moe @@ -457,7 +467,7 @@ def check_publish_order(ctx, num_ctas=None): ctas = num_ctas or ctas // 2 saved_kernel, saved_compiled = front_op._kernel, dict(front_op._compiled) tmp_dir = tempfile.mkdtemp(prefix="k3_moe_front_") - held_layer = _layer(ctx, head_flags=True, config={"num_ctas": ctas}) + held_layer = _layer(ctx, head_flags=True, num_ctas=ctas) def fused(x): return _fused(ctx, x, held_layer, bias) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_wide.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_wide.py index da28514533c8..af6c48a8b09e 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_wide.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_wide.py @@ -12,9 +12,9 @@ # 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. -"""k3_moe for steps of up to 64 tokens (K3MoeWideState: trtllm::k3_route_quant, then the m_max 64 build of k3_moe) -on one GPU at the Kimi K3 TP16 deployment's routed-expert rank layout (experts TP4 x EP4: 224 local experts, -intermediate 768 per rank), at every M in 1..64, for these routings: +"""k3_moe for steps of up to 64 tokens (trtllm::k3_route_quant, then trtllm::k3_moe on a K3MoeWideState's layer: the +m_max 64 build) on one GPU at the Kimi K3 TP16 deployment's routed-expert rank layout (experts TP4 x EP4: 224 local +experts, intermediate 768 per rank), at every M in 1..64, for these routings: - random router logits; - 16 local experts per token (1024 local pairs at M = 64); - disjoint: 3 local experts per token, no two tokens sharing one; @@ -26,7 +26,7 @@ Checks: against the stock path (trtllm::kimi_k3_noaux_tc_mxfp8_quant, then the TRTLLM-Gen W4A8_MXFP4_MXFP8 MoE runner with those ids) and an fp64 reference over the dequantized MXFP4 experts (op-catalog gates: 8 ulp of the row max per element, 4 ulp relative RMS); run-to-run identical bits; the slab armed and the layer's counters zero after every -call; at M <= 8, within one bf16 ulp of K3MoeLayer (the decode build; bit-identity reported). Then: calls +call; at M <= 8, within one bf16 ulp of a K3MoeState's layer (the decode build; bit-identity reported). Then: calls of two layers at mixed M on one stream and replayed from a CUDA graph give each call's bits alone; 0 and 65 tokens are refused. Weights are random checkpoint-format MXFP4 experts put through TRT-LLM's own loader.""" @@ -367,7 +367,8 @@ def _check(case, m): assert groups == state.mod.G_CAP == 324 decode = "" if m <= 8: - y8 = _decode_layer()(x, logits, bias, OFFSET, RSF) + ids8, w8, x_fp8_8, x_sf_8 = _ops().k3_route_quant(logits, bias, x, RSF, early_trigger=True) + y8 = _decode_layer()(x_fp8_8, x_sf_8, ids8, w8, OFFSET) ulp_dec = _max_ulp(y, y8) decode = f" max_ulp_vs_decode={ulp_dec} bits_as_decode={torch.equal(_bits(y), _bits(y8))}" assert ulp_dec <= 1 diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py index 3abb5cc1a9b0..d636d0bd91cb 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py @@ -1,7 +1,9 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""GPU certification matrix for the ``moe/k3_moe_front`` catalog entry and its ``K3MoeHeadWorkspace``, with the -fused-front cells of ``moe/k3_moe`` (``k3_moe_fused_front`` on a plain and on a head_flags ``K3MoeState``). +"""GPU certification matrix for the ``moe/k3_moe_front`` catalog entry and its ``K3MoeHeadWorkspace``, with the cells +of ``moe/k3_moe`` behind the front: ``trtllm::k3_moe_front``, then the ``k3_moe`` entry on a plain and on a head_flags +``K3MoeState`` (the head_flags pair: the front with ``publish=True`` releasing the workspace's ready words, k3_moe +acquiring them). The front's correctness depends on state that outlives a call: the head workspace's two alternating Lamport buffers (flags[0] says which one a call uses; every call flips it), the readers' re-arm of every word they read, and, with a @@ -32,9 +34,9 @@ weights, MXFP8 codes and scales must equal it bit for bit (the front routes with k3_route_quant's selection, weights and quantization: the kernel's statement), and be bit for bit the same on every rank. The shared activation: this rank's gate_up in fp64 (exact) rounded to bf16, SiTU in fp32, within 2e-2 of its largest magnitude (the kernel's tanh -and sigmoid are fast approximations). k3_moe_fused_front's routed partial against the stock TRTLLM-Gen +and sigmoid are fast approximations). The routed partial of the front + k3_moe pair against the stock TRTLLM-Gen W4A8_MXFP4_MXFP8 runner on the front's own routing and latent (op-catalog gates: 8 bf16 ulp of the row's max per -element, 4 ulp relative RMS); the head_flags build bit for bit against the plain one. Call sequences compare every +element, 4 ulp relative RMS); the head_flags pair bit for bit against the plain one. Call sequences compare every call bit for bit with the same call made alone (itself checked against the references first). """ @@ -72,7 +74,8 @@ SHARED_TOL = 2e-2 TOKENS = (1, 2, 3, 4, 5, 6, 7, 8) LAYERS = 3 -# Layer l of a step: k3_moe_fused_front on the plain state, on the head_flags state, and k3_moe_front alone. +# Layer l of a step: the front then k3_moe on the plain state, the publishing front then k3_moe on the head_flags state, +# and the front alone. KINDS = ("fused", "flags", "front") DIP_STEPS = (8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8) REPLAYS = 6 @@ -134,8 +137,9 @@ def ordered(x): class Call: """One call: the MoE input ``x`` (the same on every rank), the layer whose weights it uses, and its kind. - Kinds: "front" (k3_moe_front), "fused" (k3_moe_fused_front on the plain K3MoeState's layer), "flags" - (k3_moe_fused_front on the head_flags K3MoeState's layer). + Kinds: "front" (the k3_moe_front entry, five outputs), "fused" (the entry, then the k3_moe entry on the plain + K3MoeState's layer: ``(y, shared)``), "flags" (the publishing front, then the k3_moe entry on the head_flags + K3MoeState's layer with ``head=ws``: ``(y, shared)``). """ def __init__(self, seed: int, tokens: int, layer: int = 0, kind: str = "front", x=None): @@ -154,12 +158,16 @@ def as_kind(self, kind: str) -> "Call": def run(self, ws, x=None): x = self.x if x is None else x lw = LAYER_WEIGHTS[self.layer] + if self.kind == "flags": + ids, w, q, s, shared = OPS.front( + x, lw.front, BIAS, RSF, SHARED_COLS, GATE_CAP, LINEAR_CAP, ws, publish=True + ) + return OPS.k3_moe(q, s, ids, w, OFFSET, FLAGS.layers[self.layer], head=ws), shared + out = OPS.front(x, lw.front, BIAS, RSF, SHARED_COLS, GATE_CAP, LINEAR_CAP, ws) if self.kind == "front": - return OPS.front(x, lw.front, BIAS, RSF, SHARED_COLS, GATE_CAP, LINEAR_CAP, ws) - layer = (PLAIN if self.kind == "fused" else FLAGS).layers[self.layer] - return OPS.fused( - x, lw.front, BIAS, OFFSET, RSF, SHARED_COLS, GATE_CAP, LINEAR_CAP, ws, layer - ) + return out + ids, w, q, s, shared = out + return OPS.k3_moe(q, s, ids, w, OFFSET, PLAIN.layers[self.layer]), shared # ── references ──────────────────────────────────────────────────────────── @@ -235,7 +243,7 @@ def compare(y, ref): def verify_fused(got, call: Call, ws, where: str) -> None: - """k3_moe_fused_front's (y, shared) for ``call``. + """The (y, shared) of a front + k3_moe pair for ``call``. The front alone on the same input (made here, on ``ws``, by every rank) gives the routing and MXFP8 latent k3_moe consumed: y within the op-catalog gates of the stock runner on them (all zeros on a rank no token routes to), and @@ -405,12 +413,12 @@ def check_front_single_calls() -> None: ) -def check_fused_front_single_calls() -> None: - """k3_moe_fused_front at every M 1-8 on the plain K3MoeState and on the head_flags one. +def check_front_and_k3_moe_single_calls() -> None: + """The front then the k3_moe entry at every M 1-8, on the plain K3MoeState and on the head_flags one. Layer 0's weights, the first M tokens of one batch. The plain y within the op-catalog gates of the stock runner on the front's own routing and latent (all zeros on a rank no token routes to), the shared activation the front's - bits; the head_flags build's y and shared the plain build's bits, its handoff checked before and after + bits; the head_flags pair's y and shared the plain pair's bits, its handoff checked before and after (flags_call_checked); the plain y's rows within one bf16 ulp of the same rows of the 8-token call (k3_moe's slice FC2 groups a token's expert terms by the step's group count; bit-identity counted) and the shared rows their bits; run-to-run identical bits; afterwards both states armed with every counter zero and the head buffers empty. @@ -477,8 +485,9 @@ def preset(epoch, zero_words=False): def check_dip_and_regrow_sequence() -> None: """Decode steps of three layers on one workspace, the token count dipping and growing back, a random rank late. - Each step: k3_moe_fused_front on the plain state (layer 0's weights), on the head_flags state (layer 1's) and the - front alone (layer 2's), at M 8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8 with new inputs every call, the layers back to back + Each step: the front then k3_moe on the plain state (layer 0's weights), the publishing front then k3_moe on the + head_flags state (layer 1's), and the front alone (layer 2's), at M 8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8 with new inputs + every call, the layers back to back with a random rank 5 ms late before every call. Every call returns the bits of the same call made alone (each checked first). Afterwards the head buffers are empty, the epoch advanced once per head_flags call with every ready word at it, and both states armed: a call after a smaller one reads nothing an older, larger call left. @@ -538,8 +547,8 @@ def check_two_workspaces_interleaved() -> None: def check_graph_capture_and_replay() -> None: """A captured step replayed with rewritten inputs, eager calls of other token counts between replays. - The step: the three layers at M 8 on WS_B (k3_moe_fused_front plain and head_flags, the front alone), captured - once and replayed 6 times with new inputs copied into its static buffers; between replays an eager call of + The step: the three layers at M 8 on WS_B (the front + k3_moe pairs, plain and head_flags, and the front alone), + captured once and replayed 6 times with new inputs copied into its static buffers; between replays an eager call of another M on WS_B. Every replayed and eager call returns the bits of the same call alone; afterwards the head buffers are empty, WS_B's epoch advanced once per head_flags call, replayed or eager, with every ready word at it, and both states armed. @@ -603,8 +612,9 @@ def step(): def check_unsupported_shape_raises_on_every_rank() -> None: """M 0 and 9 raise ValueError on every rank before touching anything, and the next call is correct. - The front alone and k3_moe_fused_front (plain and head_flags) at M 0 and 9: the head workspace's words, flags and - ready words keep their bits; the next call returns the bits of the same call made alone. + The front alone and the two front + k3_moe pairs (plain, and head_flags with the publishing front) at M 0 and 9: + the front refuses them before touching the workspace (its words, flags and ready words keep their bits); the next + call returns the bits of the same call made alone. """ nxt = Call(7000, 8, layer=0, kind="fused") want = alone_results([nxt], WS_A, "before the unsupported calls")[0] @@ -632,41 +642,72 @@ def check_unsupported_shape_raises_on_every_rank() -> None: ) -def check_create_and_first_compile_refuse_capture() -> None: - """Under CUDA-graph capture on every rank, create() and a not yet compiled front configuration raise. +def raised_under_capture(fn) -> str: + """Run ``fn`` under CUDA-graph capture; the message of the RuntimeError it raised, or '' if it raised none.""" + graph, stream = torch.cuda.CUDAGraph(), torch.cuda.Stream() + resting = torch.cuda.current_stream() + stream.wait_stream(resting) + message = "" + try: + with torch.cuda.graph(graph, stream=stream): + try: + fn() + except RuntimeError as exc: + message = str(exc) + finally: + # A capture that fails when it ends leaves its own stream current; put the resting one back. + torch.cuda.set_stream(resting) + del graph + return message + - K3MoeHeadWorkspace.create raises RuntimeError before entering the collective (no rank waits for another), and a - front call of a configuration not yet compiled (other SiTU caps) raises RuntimeError before any launch. The - workspace keeps its bits and the next call returns the bits of the same call made alone. +def raised_eagerly(fn) -> str: + """Run ``fn``; the message of the RuntimeError it raised, or '' if it raised none.""" + try: + fn() + except RuntimeError as exc: + return str(exc) + return "" + + +def check_create_and_first_compile_refuse_capture() -> None: + """Under CUDA-graph capture, create() is refused on every rank, and a not yet compiled front configuration raises. + + K3MoeHeadWorkspace.create has the ranks agree before allocating that each of them can (not capturing, enough free + memory): with every rank capturing, and with the last rank capturing while its peers call it eagerly, every rank + raises RuntimeError at that agreement ("not every rank can allocate", a capturing rank's message naming the + capture), so none allocates. Every rank must call it: a rank calling it alone would wait for its peers. A front + call of a configuration not yet compiled (other SiTU caps) raises RuntimeError under capture before any launch. + The workspace keeps its bits and the next call returns the bits of the same call made alone. """ nxt = Call(7500, 4, layer=1, kind="front") want = alone_results([nxt], WS_A, "before the capture refusals")[0] before = workspace_snapshot(WS_A) - messages = [] - stream = torch.cuda.Stream() - stream.wait_stream(torch.cuda.current_stream()) - graph = torch.cuda.CUDAGraph() - with torch.cuda.graph(graph, stream=stream): - try: - OPS.workspace.create(R.mapping, fabric_handle=R.fabric) - except RuntimeError as exc: - messages.append(str(exc)) - try: - OPS.front( - nxt.x, - LAYER_WEIGHTS[1].front, - BIAS, - RSF, - SHARED_COLS, - GATE_CAP + 1.0, - LINEAR_CAP, - WS_A, - ) - except RuntimeError as exc: - messages.append(str(exc)) - del graph - refused = len(messages) == 2 and all("outside CUDA-graph capture" in msg for msg in messages) - assert R.all_true(refused), f"under capture: {messages}" + + def create(): + OPS.workspace.create(R.mapping, fabric_handle=R.fabric) + + every = raised_under_capture(create) + R.barrier() + capturing = R.world - 1 + one = raised_under_capture(create) if R.rank == capturing else raised_eagerly(create) + R.barrier() + first = raised_under_capture( + lambda: OPS.front( + nxt.x, LAYER_WEIGHTS[1].front, BIAS, RSF, SHARED_COLS, GATE_CAP + 1.0, LINEAR_CAP, WS_A + ) + ) + refused = ( + "not every rank can allocate" in every + and "outside CUDA-graph capture" in every + and "not every rank can allocate" in one + and (R.rank != capturing or "outside CUDA-graph capture" in one) + and "outside CUDA-graph capture" in first + ) + assert R.all_true(refused), ( + f"rank {R.rank}: every rank capturing {every!r}; rank {capturing} capturing {one!r}; " + f"uncompiled front {first!r}" + ) after = workspace_snapshot(WS_A) assert all(same(a, b) for a, b in zip(before, after)), ( "a refused call touched the head workspace" @@ -742,7 +783,7 @@ def check_wrong_call_order_is_detected() -> None: CHECKS = [ check_workspace_is_armed_and_sized, check_front_single_calls, - check_fused_front_single_calls, + check_front_and_k3_moe_single_calls, check_head_flags_epoch_wraps, check_dip_and_regrow_sequence, check_two_workspaces_interleaved, @@ -836,7 +877,7 @@ def _run_one_rank(args) -> int: OPS = SimpleNamespace( front=front.k3_moe_front, - fused=moe.k3_moe_fused_front, + k3_moe=moe.k3_moe, workspace=front.K3MoeHeadWorkspace, K3MoeState=moe.K3MoeState, ) diff --git a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe.py b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe.py index b4f8f8743980..4e85a0f2696f 100644 --- a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe.py +++ b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe.py @@ -1,11 +1,12 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""GPU test for the moe/k3_moe catalog entry: k3_moe and k3_moe_wide on their caller-owned state. +"""GPU test for the moe/k3_moe catalog entry: trtllm::k3_moe on its caller-owned state, both builds. One GPU (sm_100), one rank of the Kimi K3 TP16 deployment's routed experts: experts TP4 x EP4, so 224 of the 896 experts are local (global ids [224, 448)) with intermediate 768 per rank. The experts are random checkpoint-format MXFP4 tensors put through TRT-LLM's own W4A8_MXFP4_MXFP8 TRTLLM-Gen loader, which writes the buffers that both k3_moe -and the stock runner read. +and the stock runner read. Every call routes with moe/k3_route_quant first (its early trigger on when the state +launches k3_moe as a programmatic dependent), as a target does. References: the stock path (trtllm::kimi_k3_noaux_tc_mxfp8_quant, then the TRTLLM-Gen W4A8_MXFP4_MXFP8 MoE runner on those ids) and an fp64 reference over the dequantized MXFP4 experts (SiTU, the MXFP8 intermediate with the round-up @@ -17,10 +18,11 @@ first (compiling) call eagerly where it needs one, so it also runs alone (-k). The kernel tests under tests/unittest/_torch/cute_dsl_kernels/kimi_k3/ remain the exhaustive numerics; this file copies what it needs. -k3_moe_fused_front needs the TP group's head workspace: its cells run in moe/k3_moe_front's 4-rank matrix, -tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py. +The head_flags build needs the TP group's head workspace and the MoE front: its cells run in moe/k3_moe_front's +4-rank matrix, tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py. """ +import contextlib import functools import math from types import SimpleNamespace @@ -34,8 +36,8 @@ K3MoeWideState, is_supported, k3_moe, - k3_moe_wide, ) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.k3_route_quant import k3_route_quant H, TOP_K, NUM_EXPERTS, SV = 3584, 16, 896, 32 I_TP, E_LOCAL, MOE_TP, TP_RANK, EP_RANK = 768, 224, 4, 1, 1 # one rank of experts TP4 x EP4 @@ -50,7 +52,6 @@ DECODE_CASES = ("random", "16_local_disjoint", "none_local") WIDE_CASES = ("random", "group_cap", "none_local") WIDE_M = (1, 2, 7, 8, 9, 16, 33, 40, 64) -ROUTE_M = (1, 2, 3, 4, 5, 6, 7, 8, 16, 33, 64) DIP_STEPS = (8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8) WIDE_DIP_STEPS = (64, 64, 16, 1, 40, 64, 8, 23, 64, 9, 64) REPLAYS = 6 @@ -224,10 +225,13 @@ def _draw(seed: int, rows: int): return logits, torch.randn(rows, H, generator=gen, device=DEV).bfloat16() -def _zeros(rows: int): +def _routed_zeros(rows: int): + """k3_moe's inputs for ``rows`` tokens, all zero: (x_fp8, x_sf, topk_ids, topk_weights).""" return ( - torch.zeros(rows, NUM_EXPERTS, device=DEV), - torch.zeros(rows, H, dtype=torch.bfloat16, device=DEV), + torch.zeros(rows, H, dtype=torch.float8_e4m3fn, device=DEV), + torch.zeros(rows, H // SV, dtype=torch.uint8, device=DEV), + torch.zeros(rows, TOP_K, dtype=torch.int32, device=DEV), + torch.zeros(rows, TOP_K, dtype=torch.bfloat16, device=DEV), ) @@ -262,12 +266,16 @@ def _experts_of(i: int): return (_experts()[0], _rolled(), _experts()[0])[i] -def _decode_call(layer, logits, x): - return k3_moe(x, logits, _bias(), OFFSET, RSF, layer) +def _route(layer, logits, x): + """moe/k3_route_quant for a call on ``layer``: its dependents launched early when the layer's state launches + k3_moe as a programmatic dependent.""" + return k3_route_quant(logits, _bias(), x, RSF, early_trigger=layer.state.use_pdl) -def _wide_call(layer, logits, x, out=None): - return k3_moe_wide(x, logits, _bias(), OFFSET, RSF, layer, out=out) +def _call(layer, logits, x, out=None): + """moe/k3_route_quant, then the moe/k3_moe entry on ``layer`` (either build).""" + ids, weights, x_fp8, x_sf = _route(layer, logits, x) + return k3_moe(x_fp8, x_sf, ids, weights, OFFSET, layer, out=out) def _armed(state, layers) -> bool: @@ -292,6 +300,21 @@ def _snapshot(state, layers): ] +@contextlib.contextmanager +def _cold_k3_moe_cache(): + """trtllm::k3_moe's compile cache emptied for the duration (and restored after): the next call of every build is + a first call.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op + + saved = dict(op._compiled) + op._compiled.clear() + try: + yield + finally: + op._compiled.clear() + op._compiled.update(saved) + + # ── references ──────────────────────────────────────────────────────────── @@ -429,10 +452,10 @@ def _check_vs_stock(y, experts, logits, x, where: str) -> None: assert c["ok"], f"{where}: against the stock path {c}" -def _alone(fn, layer, logits, x, experts, where: str): +def _alone(layer, logits, x, experts, where: str): """One call made alone (synchronized before and after), checked against the stock path.""" torch.cuda.synchronize() - y = fn(layer, logits, x) + y = _call(layer, logits, x) torch.cuda.synchronize() _check_vs_stock(y, experts, logits, x, where) return y @@ -444,55 +467,47 @@ def _alone(fn, layer, logits, x, experts, where: str): def test_state_armed_and_sized(): """New state objects are armed and sized as the contract states for 224 local experts and intermediate 768. - K3MoeState: slab [128, 8, 768] FP8 codes and [128, 8, 96] scale bytes armed, FC2 partial rows fp32 [1024, 3584] - zero, not compiled; a layer's counters int32 [288] zero. K3MoeWideState: slab [324, 8, 768] and [324, 8, 96] - armed, partials fp32 [1024, 3584]; a layer's counters int32 [680] zero. The loader's buffers fit the kernels. + K3MoeState: the M <= 8 build (m_max 8, one CTA per SM), slab [128, 8, 768] FP8 codes and [128, 8, 96] scale bytes + armed, FC2 partial rows fp32 [1024, 3584] zero; a layer's counters int32 [288] zero and its weights the tensors + it was built over. K3MoeWideState: m_max 64, slab [324, 8, 768] and [324, 8, 96] armed, partials fp32 + [1024, 3584] zero; a layer's counters int32 [680] zero. The loader's buffers fit the kernels. """ proc, _, _ = _experts() ok, why = is_supported(*_weights(proc), E_LOCAL) assert ok, why + sms = torch.cuda.get_device_properties(_device()).multi_processor_count state = K3MoeState(_device(), I_TP, E_LOCAL) layer = state.layer(*_weights(proc)) g = min(E_LOCAL, DECODE_MAX * TOP_K) assert state.mod.G_CAP == g == 128 + assert state.m_max == DECODE_MAX and state.num_ctas == sms and not state.head_flags + assert state.num_local == E_LOCAL and state.i_tp == I_TP assert state.c.dtype == state.cs.dtype == torch.int8 assert tuple(state.c.shape) == (g, 8, I_TP) and tuple(state.cs.shape) == (g, 8, I_TP // 8) assert state.part.dtype == torch.float32 and tuple(state.part.shape) == (8 * g, H) assert bool((state.part == 0).all()) assert layer.counters.dtype == torch.int32 and tuple(layer.counters.shape) == (32 + 2 * g,) - assert _armed(state, [layer]) and not state.head_flags and state.compiled is None - assert K3MoeState(_device(), I_TP, E_LOCAL, head_flags=True).head_flags + assert all(a is b for a, b in zip(layer.weights, _weights(proc))) + assert _armed(state, [layer]) + flags_state = K3MoeState(_device(), I_TP, E_LOCAL, head_flags=True) + assert flags_state.head_flags and flags_state.m_max == DECODE_MAX wide = K3MoeWideState(_device(), I_TP, E_LOCAL) wide_layer = wide.layer(*_weights(proc)) gw = E_LOCAL + (WIDE_MAX * TOP_K - E_LOCAL) // 8 assert wide.mod.G_CAP == gw == 324 + assert wide.m_max == WIDE_MAX and wide.num_ctas == sms and wide.use_pdl and not wide.head_flags assert tuple(wide.c.shape) == (gw, 8, I_TP) and tuple(wide.cs.shape) == (gw, 8, I_TP // 8) assert wide.part.dtype == torch.float32 and tuple(wide.part.shape) == (16 * WIDE_MAX, H) + assert bool((wide.part == 0).all()) assert tuple(wide_layer.counters.shape) == (32 + 2 * gw,) - assert _armed(wide, [wide_layer]) and wide.compiled is None - - -@pytest.mark.parametrize("m", ROUTE_M) -def test_k3_route_quant_is_the_stock_routing(m): - """k3_route_quant (inside k3_moe and k3_moe_wide) returns kimi_k3_noaux_tc_mxfp8_quant's four outputs, bit for bit. - - Top-16 ids, routing weights, MXFP8 codes and scales, with and without the early dependent trigger. - """ - logits64, x64 = _tokens("random", WIDE_MAX) - logits, x = logits64[:m].contiguous(), x64[:m].contiguous() - want = torch.ops.trtllm.kimi_k3_noaux_tc_mxfp8_quant(logits, _bias(), x, RSF) - for early in (False, True): - got = torch.ops.trtllm.k3_route_quant(logits, _bias(), x, RSF, early_trigger=early) - same = [_same(a, b) for a, b in zip(got, want)] - print(f"OPCHECK op=k3_route_quant M={m} early_trigger={early} same(ids,w,q,sf)={same}") - assert all(same), f"M {m} early_trigger {early}: {same}" + assert _armed(wide, [wide_layer]) @pytest.mark.parametrize("m", range(1, DECODE_MAX + 1)) @pytest.mark.parametrize("case", DECODE_CASES) def test_k3_moe_single_call(case, m): - """One k3_moe call of M tokens on layer A against the fp64 reference and the stock path. + """One call of M tokens on the K3MoeState's layer A against the fp64 reference and the stock path. Also: run-to-run identical bits; each row within one bf16 ulp of the same row of the 8-token call (the slice FC2 adds a token's expert terms in slices whose bounds follow the step's group count; bit-identity is reported); the @@ -502,7 +517,7 @@ def test_k3_moe_single_call(case, m): state, layers = _decode() logits8, x8 = _tokens(case, DECODE_MAX) logits, x = logits8[:m].contiguous(), x8[:m].contiguous() - y = _decode_call(layers[0], logits, x) + y = _call(layers[0], logits, x) torch.cuda.synchronize() rearmed = _armed(state, layers) assert ( @@ -511,8 +526,8 @@ def test_k3_moe_single_call(case, m): and y.is_contiguous() and y.device == x.device ) - det = all(_same(_decode_call(layers[0], logits, x), y) for _ in range(2)) - y8 = _decode_call(layers[0], logits8, x8) + det = all(_same(_call(layers[0], logits, x), y) for _ in range(2)) + y8 = _call(layers[0], logits8, x8) ulp_m8 = _max_ulp(y, y8[:m]) y_stock, ids, w, x_fp8, x_sf = _stock(proc, bias, x, logits) local = _local_pairs(ids) @@ -524,7 +539,7 @@ def test_k3_moe_single_call(case, m): if local == 0: zeros = bool((y == 0).all()) print( - f"OPCHECK op=k3_moe case={case} M={m} zeros={zeros} det={det} " + f"OPCHECK op=k3_moe build=m8 case={case} M={m} zeros={zeros} det={det} " f"rows_as_m8={_same(y, y8[:m])} scratch_rearmed={rearmed}" ) assert zeros and det and ulp_m8 == 0 and rearmed @@ -532,7 +547,7 @@ def test_k3_moe_single_call(case, m): ref = _reference(raw, _deq_x(x_fp8, x_sf), ids, w).bfloat16() c_ref, c_stock, c_stock_ref = _compare(y, ref), _compare(y, y_stock), _compare(y_stock, ref) print( - f"OPCHECK op=k3_moe case={case} M={m} local_pairs={local} groups={groups} " + f"OPCHECK op=k3_moe build=m8 case={case} M={m} local_pairs={local} groups={groups} " f"vs_ref_elt_ulp={c_ref['elt_ulp']:.2f} vs_ref_rms_ulp={c_ref['rms_ulp']:.2f} " f"vs_stock_elt_ulp={c_stock['elt_ulp']:.2f} vs_stock_rms_ulp={c_stock['rms_ulp']:.2f} " f"stock_vs_ref_elt_ulp={c_stock_ref['elt_ulp']:.2f} det={det} " @@ -545,36 +560,37 @@ def test_k3_moe_single_call(case, m): @pytest.mark.parametrize("m", WIDE_M) @pytest.mark.parametrize("case", WIDE_CASES) def test_k3_moe_wide_single_call(case, m): - """One k3_moe_wide call of M tokens on its layer A against the fp64 reference and the stock path. + """One call of M tokens on the K3MoeWideState's layer A against the fp64 reference and the stock path. Also: run-to-run identical bits; the slab armed and every counter zero after the call; at M <= 8 within one bf16 - ulp of k3_moe on the same experts (bit-identity reported); at group_cap M 64 the group count at the capacity, 324. + ulp of the M <= 8 build on the same experts (bit-identity reported); at group_cap M 64 the group count at the + capacity, 324. """ proc, raw, bias = _experts() state, layers = _wide() _, decode_layers = _decode() logits64, x64 = _tokens(case, WIDE_MAX) logits, x = logits64[:m].contiguous(), x64[:m].contiguous() - y = _wide_call(layers[0], logits, x) + y = _call(layers[0], logits, x) torch.cuda.synchronize() rearmed = _armed(state, layers) assert y.shape == (m, H) and y.dtype == torch.bfloat16 and y.is_contiguous() - det = all(_same(_wide_call(layers[0], logits, x), y) for _ in range(2)) + det = all(_same(_call(layers[0], logits, x), y) for _ in range(2)) y_stock, ids, w, x_fp8, x_sf = _stock(proc, bias, x, logits) groups = _groups_wide(ids) if case == "group_cap" and m == WIDE_MAX: assert groups == state.mod.G_CAP == 324 decode = "" if m <= DECODE_MAX: - y_dec = _decode_call(decode_layers[0], logits, x) + y_dec = _call(decode_layers[0], logits, x) ulp_dec = _max_ulp(y, y_dec) - decode = f" max_ulp_vs_k3_moe={ulp_dec} bits_as_k3_moe={_same(y, y_dec)}" + decode = f" max_ulp_vs_m8_build={ulp_dec} bits_as_m8_build={_same(y, y_dec)}" assert ulp_dec <= 1 local = _local_pairs(ids) if local == 0: zeros = bool((y == 0).all()) print( - f"OPCHECK op=k3_moe_wide case={case} M={m} zeros={zeros} det={det} " + f"OPCHECK op=k3_moe build=m64 case={case} M={m} zeros={zeros} det={det} " f"scratch_rearmed={rearmed}{decode}" ) assert zeros and det and rearmed @@ -582,7 +598,7 @@ def test_k3_moe_wide_single_call(case, m): ref = _reference(raw, _deq_x(x_fp8, x_sf), ids, w).bfloat16() c_ref, c_stock, c_stock_ref = _compare(y, ref), _compare(y, y_stock), _compare(y_stock, ref) print( - f"OPCHECK op=k3_moe_wide case={case} M={m} local_pairs={local} groups={groups} " + f"OPCHECK op=k3_moe build=m64 case={case} M={m} local_pairs={local} groups={groups} " f"vs_ref_elt_ulp={c_ref['elt_ulp']:.2f} vs_ref_rms_ulp={c_ref['rms_ulp']:.2f} " f"vs_stock_elt_ulp={c_stock['elt_ulp']:.2f} vs_stock_rms_ulp={c_stock['rms_ulp']:.2f} " f"stock_vs_ref_elt_ulp={c_stock_ref['elt_ulp']:.2f} det={det} scratch_rearmed={rearmed}{decode}" @@ -591,74 +607,62 @@ def test_k3_moe_wide_single_call(case, m): assert det and rearmed -@pytest.mark.parametrize("m", [9, WIDE_MAX]) -def test_k3_moe_wide_out_buffer(m): - """k3_moe_wide with ``out``: the result is out[:M] (the same storage), the bits of the call without ``out``. +@pytest.mark.parametrize("build,m", [("m8", 3), ("m8", 8), ("m64", 9), ("m64", 64)]) +def test_k3_moe_out_buffer(build, m): + """``out``: the call writes out[:M], the bits of the call without ``out``, and returns an empty [0, 3584] tensor. The rows of ``out`` past M keep their bits. """ - _, layers = _wide() + layer = (_decode() if build == "m8" else _wide())[1][0] logits, x = _draw(900 + m, m) - want = _wide_call(layers[0], logits, x) + want = _call(layer, logits, x) out = torch.full((WIDE_MAX + 3, H), -3.0, dtype=torch.bfloat16, device=DEV) tail = out[m:].clone() - got = _wide_call(layers[0], logits, x, out=out) + got = _call(layer, logits, x, out=out) torch.cuda.synchronize() - assert got.data_ptr() == out.data_ptr() and got.shape == (m, H) - assert _same(got, want) and _same(out[m:], tail) + assert got.dtype == torch.bfloat16 and tuple(got.shape) == (0, H) + assert _same(out[:m], want) and _same(out[m:], tail) def test_layers_by_steps_dip_and_regrow(): """Decode steps of several layers on one state, the token count dipping and growing back, back to back. - Steps of three k3_moe calls (layers A, B, C of one state) at M 8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8, then steps of two - k3_moe_wide calls (layers A, B of the wide state) at M 64, 64, 16, 1, 40, 64, 8, 23, 64, 9, 64, with new inputs - every call. Each call is first made alone (synchronized, checked against the stock path); in the sequence, back to - back on one stream, each returns the bits of its call alone, and afterwards both slabs are armed and every counter - is zero: a call after a smaller one reads nothing an older, larger call left. + Steps of three calls (layers A, B, C of the K3MoeState) at M 8, 8, 8, 2, 7, 8, 1, 1, 8, 3, 8, then steps of two + calls (layers A, B of the K3MoeWideState) at M 64, 64, 16, 1, 40, 64, 8, 23, 64, 9, 64, with new inputs every call. + Each call is first made alone (synchronized, checked against the stock path); in the sequence, back to back on one + stream, each returns the bits of its call alone, and afterwards both slabs are armed and every counter is zero: a + call after a smaller one reads nothing an older, larger call left. """ - state, layers = _decode() - calls = [(i, *_draw(1000 + 10 * s + i, m)) for s, m in enumerate(DIP_STEPS) for i in range(3)] - alone = [ - _alone( - _decode_call, layers[i], lg, x, _experts_of(i), f"k3_moe layer {i} M {x.shape[0]} alone" - ) - for i, lg, x in calls - ] - seq = [_decode_call(layers[i], lg, x) for i, lg, x in calls] - torch.cuda.synchronize() - bad = [k for k, (y, a) in enumerate(zip(seq, alone)) if not _same(y, a)] - assert not bad, f"k3_moe calls {bad} of the sequence differ from the same calls alone" - assert _armed(state, layers) - - wide, wide_layers = _wide() - wide_calls = [ - (i, *_draw(2000 + 10 * s + i, m)) for s, m in enumerate(WIDE_DIP_STEPS) for i in range(2) - ] - alone = [ - _alone( - _wide_call, - wide_layers[i], - lg, - x, - _experts_of(i), - f"k3_moe_wide layer {i} M {x.shape[0]} alone", + for state, layers, steps, seed in ( + (*_decode(), DIP_STEPS, 1000), + (*_wide(), WIDE_DIP_STEPS, 2000), + ): + calls = [ + (i, *_draw(seed + 10 * s + i, m)) + for s, m in enumerate(steps) + for i in range(len(layers)) + ] + alone = [ + _alone( + layers[i], lg, x, _experts_of(i), f"m{state.m_max} layer {i} M {x.shape[0]} alone" + ) + for i, lg, x in calls + ] + seq = [_call(layers[i], lg, x) for i, lg, x in calls] + torch.cuda.synchronize() + bad = [k for k, (y, a) in enumerate(zip(seq, alone)) if not _same(y, a)] + assert not bad, ( + f"m{state.m_max}: calls {bad} of the sequence differ from the same calls alone" ) - for i, lg, x in wide_calls - ] - seq = [_wide_call(wide_layers[i], lg, x) for i, lg, x in wide_calls] - torch.cuda.synchronize() - bad = [k for k, (y, a) in enumerate(zip(seq, alone)) if not _same(y, a)] - assert not bad, f"k3_moe_wide calls {bad} of the sequence differ from the same calls alone" - assert _armed(wide, wide_layers) + assert _armed(state, layers) def test_two_states_interleaved(): - """Two K3MoeStates (each its own scratch) and the wide state, calls interleaved in an irregular pattern. + """Two K3MoeStates (each its own scratch) and the K3MoeWideState, calls interleaved in an irregular pattern. - A A B A B B A A A B, twice, with a k3_moe_wide call after every third, M varying, back to back on one stream: every - call returns the bits of the same call alone (checked against the stock path), and every slab is armed and every - counter zero afterwards. + A A B A B B A A A B, twice, with a call on the wide state after every third, M varying, back to back on one + stream: every call returns the bits of the same call alone (checked against the stock path), and every slab is + armed and every counter zero afterwards. """ state_a, layers_a = _decode() state_b, layer_b = _decode_b() @@ -667,17 +671,17 @@ def test_two_states_interleaved(): for i, which in enumerate("AABABBAAAB" * 2): m = (3, 8, 1, 8, 5)[i % 5] if which == "A": - plan.append((_decode_call, layers_a[i % 3], _experts_of(i % 3), *_draw(3000 + i, m))) + plan.append((layers_a[i % 3], _experts_of(i % 3), *_draw(3000 + i, m))) else: - plan.append((_decode_call, layer_b, _experts()[0], *_draw(3000 + i, m))) + plan.append((layer_b, _experts()[0], *_draw(3000 + i, m))) if i % 3 == 2: mw = (40, 64, 9)[(i // 3) % 3] - plan.append((_wide_call, wide_layers[i % 2], _experts_of(i % 2), *_draw(3500 + i, mw))) + plan.append((wide_layers[i % 2], _experts_of(i % 2), *_draw(3500 + i, mw))) alone = [ - _alone(fn, layer, lg, x, ex, f"interleaved call {k} alone") - for k, (fn, layer, ex, lg, x) in enumerate(plan) + _alone(layer, lg, x, ex, f"interleaved call {k} alone") + for k, (layer, ex, lg, x) in enumerate(plan) ] - seq = [fn(layer, lg, x) for fn, layer, _, lg, x in plan] + seq = [_call(layer, lg, x) for layer, _, lg, x in plan] torch.cuda.synchronize() bad = [k for k, (y, a) in enumerate(zip(seq, alone)) if not _same(y, a)] assert not bad, f"interleaved calls {bad} differ from the same calls alone" @@ -687,19 +691,20 @@ def test_two_states_interleaved(): def test_graph_capture_and_replay(): """A captured step replayed with rewritten inputs, eager calls of other token counts between replays. - The step: k3_moe on layers A, B, C at M 8, then k3_moe_wide on its layer A at M 64, captured once and replayed six - times with new inputs copied into its static buffers; between replays an eager k3_moe call (M 3, 1, 6, 5, 2, 7) - and an eager k3_moe_wide call (M 23, 9, 40, 1, 64, 16) on the same states. Every replayed and eager call returns - the bits of the same call alone (checked against the stock path), and the slabs are armed and every counter zero - afterwards. + The step: layers A, B, C of the K3MoeState at M 8, then layer A of the K3MoeWideState at M 64 (each routed by + k3_route_quant inside the step), captured once and replayed six times with new inputs copied into its static + buffers; between replays an eager call of another M on each state (M 3, 1, 6, 5, 2, 7 and 23, 9, 40, 1, 64, 16). + Every replayed and eager call returns the bits of the same call alone (checked against the stock path), and the + slabs are armed and every counter zero afterwards. """ state, layers = _decode() wide, wide_layers = _wide() + step_layers = [layers[0], layers[1], layers[2], wide_layers[0]] + step_experts = [_experts_of(0), _experts_of(1), _experts_of(2), _experts_of(0)] static = [_draw(4000 + i, DECODE_MAX) for i in range(3)] + [_draw(4003, WIDE_MAX)] def step(): - outs = [_decode_call(layers[i], *static[i]) for i in range(3)] - return outs + [_wide_call(wide_layers[0], *static[3])] + return [_call(layer, lg, x) for layer, (lg, x) in zip(step_layers, static)] step() # every first call eager: k3_route_quant and both k3_moe builds compile here if nothing has yet torch.cuda.synchronize() @@ -709,37 +714,21 @@ def step(): with torch.cuda.graph(graph, stream=stream): outs = step() for rep in range(REPLAYS): - inputs = [_draw(5000 + 10 * rep + i, DECODE_MAX) for i in range(3)] + [ - _draw(5003 + 10 * rep, WIDE_MAX) - ] + inputs = [_draw(5000 + 10 * rep + i, DECODE_MAX) for i in range(3)] + inputs.append(_draw(5003 + 10 * rep, WIDE_MAX)) alone = [ - _alone( - _decode_call, layers[i], *inputs[i], _experts_of(i), f"replay {rep} layer {i} alone" - ) - for i in range(3) + _alone(layer, lg, x, ex, f"replay {rep} call {i} alone") + for i, (layer, ex, (lg, x)) in enumerate(zip(step_layers, step_experts, inputs)) ] - alone.append( - _alone( - _wide_call, wide_layers[0], *inputs[3], _experts_of(0), f"replay {rep} wide alone" - ) - ) - eager = ( - _draw(6000 + rep, (3, 1, 6, 5, 2, 7)[rep]), - _draw(6100 + rep, (23, 9, 40, 1, 64, 16)[rep]), - ) - eager_layers = (layers[rep % 3], wide_layers[rep % 2]) - eager_alone = ( - _alone( - _decode_call, eager_layers[0], *eager[0], _experts_of(rep % 3), f"eager {rep} alone" - ), - _alone( - _wide_call, - eager_layers[1], - *eager[1], + eager = [ + (layers[rep % 3], _experts_of(rep % 3), *_draw(6000 + rep, (3, 1, 6, 5, 2, 7)[rep])), + ( + wide_layers[rep % 2], _experts_of(rep % 2), - f"eager wide {rep} alone", + *_draw(6100 + rep, (23, 9, 40, 1, 64, 16)[rep]), ), - ) + ] + eager_alone = [_alone(layer, lg, x, ex, f"eager {rep} alone") for layer, ex, lg, x in eager] for (lg, x), (new_lg, new_x) in zip(static, inputs): lg.copy_(new_lg) x.copy_(new_x) @@ -747,7 +736,7 @@ def step(): torch.cuda.synchronize() bad = [i for i, (y, a) in enumerate(zip(outs, alone)) if not _same(y, a)] assert not bad, f"replay {rep}: calls {bad} differ from the same calls alone" - got = (_decode_call(eager_layers[0], *eager[0]), _wide_call(eager_layers[1], *eager[1])) + got = [_call(layer, lg, x) for layer, _, lg, x in eager] torch.cuda.synchronize() assert all(_same(g, a) for g, a in zip(got, eager_alone)), f"eager calls after replay {rep}" del graph @@ -757,43 +746,56 @@ def step(): def test_unsupported_calls_refused_before_launch(): """Unsupported calls raise ValueError before any launch, and the next call is correct. - k3_moe at M 0 and 9, k3_moe_wide at M 0 and 65, and k3_moe on a head_flags state's layer (that build takes the - front's ready words: k3_moe_fused_front only). The slabs, the FC2 partial rows and every counter keep their bits; - the next calls return the bits of the same calls made before. + M 0 and 9 on the K3MoeState, M 0 and 65 on the K3MoeWideState; a head_flags state's layer without ``head`` and a + plain state's layer with one; int64 ids; an ``out`` with too few rows. The slabs, the FC2 partial rows and every + counter keep their bits; the next calls return the bits of the same calls made before. """ state, layers = _decode() wide, wide_layers = _wide() lg8, x8 = _draw(7000, DECODE_MAX) lg40, x40 = _draw(7001, 40) - want = _decode_call(layers[0], lg8, x8) - want_wide = _wide_call(wide_layers[0], lg40, x40) + want = _call(layers[0], lg8, x8) + want_wide = _call(wide_layers[0], lg40, x40) flags_state = K3MoeState(_device(), I_TP, E_LOCAL, head_flags=True) flags_layer = flags_state.layer(*_weights(_experts()[0])) + ids, weights, x_fp8, x_sf = _route(layers[0], lg8, x8) + # Never touched: the entry refuses a head for a plain state's layer before reading it. + head = SimpleNamespace( + ready=torch.zeros(32, dtype=torch.int32, device=DEV), + flags=torch.zeros(4, dtype=torch.int32, device=DEV), + ) torch.cuda.synchronize() objects = ((state, layers), (wide, wide_layers), (flags_state, [flags_layer])) before = [_snapshot(st, lyrs) for st, lyrs in objects] for m in (0, DECODE_MAX + 1): with pytest.raises(ValueError): - _decode_call(layers[0], *_zeros(m)) + k3_moe(*_routed_zeros(m), OFFSET, layers[0]) for m in (0, WIDE_MAX + 1): with pytest.raises(ValueError): - _wide_call(wide_layers[0], *_zeros(m)) - with pytest.raises(ValueError, match="head_flags"): - _decode_call(flags_layer, lg8, x8) + k3_moe(*_routed_zeros(m), OFFSET, wide_layers[0]) + with pytest.raises(ValueError, match="head"): + k3_moe(x_fp8, x_sf, ids, weights, OFFSET, flags_layer) + with pytest.raises(ValueError, match="head"): + k3_moe(x_fp8, x_sf, ids, weights, OFFSET, layers[0], head=head) + with pytest.raises(ValueError): + k3_moe(x_fp8, x_sf, ids.long(), weights, OFFSET, layers[0]) + with pytest.raises(ValueError): + short = torch.empty(DECODE_MAX - 1, H, dtype=torch.bfloat16, device=DEV) + k3_moe(x_fp8, x_sf, ids, weights, OFFSET, layers[0], out=short) torch.cuda.synchronize() after = [_snapshot(st, lyrs) for st, lyrs in objects] assert all(_same(a, b) for snap_a, snap_b in zip(before, after) for a, b in zip(snap_a, snap_b)) - assert flags_state.compiled is None - assert _same(_decode_call(layers[0], lg8, x8), want) - assert _same(_wide_call(wide_layers[0], lg40, x40), want_wide) + assert _same(_call(layers[0], lg8, x8), want) + assert _same(_call(wide_layers[0], lg40, x40), want_wide) def test_construction_and_first_call_refuse_capture(): - """Under CUDA-graph capture the states refuse to allocate, and a new state refuses its first (compiling) call. + """Under CUDA-graph capture the states refuse to allocate, and a build refuses its first (compiling) call. - K3MoeState(), K3MoeState.layer(), K3MoeWideState() and K3MoeWideState.layer() raise RuntimeError; the first call - on a new K3MoeState and on a new K3MoeWideState raises RuntimeError before the k3_moe launch (k3_route_quant, - compiled eagerly first, is captured and discarded with the graph); neither state is compiled afterwards. + K3MoeState(), K3MoeState.layer(), K3MoeWideState() and K3MoeWideState.layer() raise RuntimeError. With + trtllm::k3_moe's compile cache cold, the first call of either build raises RuntimeError before the k3_moe launch + (k3_route_quant, compiled eagerly first, is captured and discarded with the graph), and the build stays + uncompiled. """ proc, _, bias = _experts() fresh = K3MoeState(_device(), I_TP, E_LOCAL) @@ -802,28 +804,29 @@ def test_construction_and_first_call_refuse_capture(): fresh_wide_layer = fresh_wide.layer(*_weights(proc)) lg4, x4 = _draw(7100, 4) lg16, x16 = _draw(7101, 16) - # Both k3_route_quant builds compiled, whichever PDL setting the layers use. + # Both k3_route_quant builds compiled, whichever PDL setting the states use. for early in (False, True): - torch.ops.trtllm.k3_route_quant(lg4, bias, x4, RSF, early_trigger=early) + k3_route_quant(lg4, bias, x4, RSF, early_trigger=early) torch.cuda.synchronize() stream = torch.cuda.Stream() stream.wait_stream(torch.cuda.current_stream()) graph = torch.cuda.CUDAGraph() - with torch.cuda.graph(graph, stream=stream): - with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): - K3MoeState(_device(), I_TP, E_LOCAL) - with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): - fresh.layer(*_weights(proc)) - with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): - K3MoeWideState(_device(), I_TP, E_LOCAL) - with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): - fresh_wide.layer(*_weights(proc)) - with pytest.raises(RuntimeError, match="compiles on its first call"): - _decode_call(fresh_layer, lg4, x4) - with pytest.raises(RuntimeError, match="compiles on its first call"): - _wide_call(fresh_wide_layer, lg16, x16) + with _cold_k3_moe_cache(): + with torch.cuda.graph(graph, stream=stream): + with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): + K3MoeState(_device(), I_TP, E_LOCAL) + with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): + fresh.layer(*_weights(proc)) + with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): + K3MoeWideState(_device(), I_TP, E_LOCAL) + with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): + fresh_wide.layer(*_weights(proc)) + with pytest.raises(RuntimeError, match="compiles on its first call"): + _call(fresh_layer, lg4, x4) + with pytest.raises(RuntimeError, match="compiles on its first call"): + _call(fresh_wide_layer, lg16, x16) + assert not fresh.compiled and not fresh_wide.compiled del graph - assert fresh.compiled is None and fresh_wide.compiled is None def test_negative_control_weights_rebound_after_layer(): @@ -843,14 +846,14 @@ def test_negative_control_weights_rebound_after_layer(): stale_wide = wide.layer(*_weights(e1)) lg8, x8 = _tokens("random", DECODE_MAX) lg40, x40 = _draw(8000, 40) - y_e1 = _decode_call(stale, lg8, x8) - yw_e1 = _wide_call(stale_wide, lg40, x40) + y_e1 = _call(stale, lg8, x8) + yw_e1 = _call(stale_wide, lg40, x40) torch.cuda.synchronize() # The reload: the caller's weights are now these new tensors; E1's buffers stay as they were. e2 = _rolled() - y_stale = _decode_call(stale, lg8, x8) - yw_stale = _wide_call(stale_wide, lg40, x40) + y_stale = _call(stale, lg8, x8) + yw_stale = _call(stale_wide, lg40, x40) stock_e2 = _stock(e2, bias, x8, lg8)[0] stock_wide_e2 = _stock(e2, bias, x40, lg40)[0] c, cw = _compare(y_stale, stock_e2), _compare(yw_stale, stock_wide_e2) @@ -864,10 +867,10 @@ def test_negative_control_weights_rebound_after_layer(): fresh = state.layer(*_weights(e2)) fresh_wide = wide.layer(*_weights(e2)) - y_e2 = _decode_call(fresh, lg8, x8) - yw_e2 = _wide_call(fresh_wide, lg40, x40) + y_e2 = _call(fresh, lg8, x8) + yw_e2 = _call(fresh_wide, lg40, x40) assert _compare(y_e2, stock_e2)["ok"] and _compare(yw_e2, stock_wide_e2)["ok"] for name, t in e1.items(): t.copy_(e2[name]) # in place: the buffers the stale layers read now hold E2 - assert _same(_decode_call(stale, lg8, x8), y_e2) - assert _same(_wide_call(stale_wide, lg40, x40), yw_e2) + assert _same(_call(stale, lg8, x8), y_e2) + assert _same(_call(stale_wide, lg40, x40), yw_e2) diff --git a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_front_op_matrix.py b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_front_op_matrix.py index ea93c75c4902..d19d722ff2f2 100644 --- a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_front_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_front_op_matrix.py @@ -2,11 +2,11 @@ # SPDX-License-Identifier: Apache-2.0 """Collected entry point for the k3_moe_front op's certification matrix. -The matrix also certifies moe/k3_moe's fused-front cells (k3_moe_fused_front on a plain and a head_flags -K3MoeState). It is ``comm/_k3_moe_front_op_matrix.py``: rank bodies live beside ``_lockstep`` and ``_rank_job`` in -``comm/``, and it is its own W-rank launcher (call sequences over caller-owned workspaces and states, so one job, not -independent cases); see ``_rank_job`` for why that is left intact. This file puts ``comm/`` on the import path to -reach ``_rank_job``. +The matrix also certifies moe/k3_moe's cells behind the front (trtllm::k3_moe_front, then the k3_moe entry on a plain +and on a head_flags K3MoeState). It is ``comm/_k3_moe_front_op_matrix.py``: rank bodies live beside ``_lockstep`` and +``_rank_job`` in ``comm/``, and it is its own W-rank launcher (call sequences over caller-owned workspaces and states, +so one job, not independent cases); see ``_rank_job`` for why that is left intact. This file puts ``comm/`` on the +import path to reach ``_rank_job``. """ import sys diff --git a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_route_quant.py b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_route_quant.py new file mode 100644 index 000000000000..d0b0a1e79470 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_route_quant.py @@ -0,0 +1,262 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the moe/k3_route_quant catalog entry: trtllm::k3_route_quant on one GPU (sm_100). + +The reference is the stock op it replaces, trtllm::kimi_k3_noaux_tc_mxfp8_quant: the four outputs (top-16 ids, +routing weights, MXFP8 codes, UE8M0 scales) must be its bits. Checks, in file order: +- single calls at M 1-8, 16, 33 and 64 with random logits, with and without the early trigger: the stock op's bits, + run-to-run identical, and each M's rows the bits of the same rows of the 64-token call; +- edge cases at M 1, 3 and 8 (40 tied selection keys, equal logits, huge logits, zero / large / denormal latent + rows): the stock op's bits; +- a captured sequence of calls at three M replayed with rewritten inputs, eager calls of other M between replays: + every call the bits of the same call made alone; +- M 0 and 65, a non-contiguous or mistyped input and a mis-sized bias raise ValueError, and the next call is correct; +- the first call of a build under CUDA-graph capture raises RuntimeError (the kernel compiles on its first call). +The op keeps no state between calls: its only process state is the compile cache. +""" + +import contextlib +import functools + +import pytest +import torch + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.k3_route_quant import k3_route_quant + +E, K, H = 896, 16, 3584 +RSF = 2.827 +DEV = "cuda" +M_ALL = (1, 2, 3, 4, 5, 6, 7, 8, 16, 33, 64) +EDGE_CASES = ("40_tied_keys", "all_equal_logits", "huge_logits", "zero_large_denormal_rows") +REPLAYS = 4 + + +def _is_sm100() -> bool: + return torch.cuda.is_available() and torch.cuda.get_device_capability() == (10, 0) + + +pytestmark = pytest.mark.skipif(not _is_sm100(), reason="k3_route_quant needs sm_100") + + +@pytest.fixture(autouse=True) +def _inference_mode(): + """Run every check as the model runs the op, under inference mode.""" + with torch.inference_mode(): + yield + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.uint8) + + +def _same(a: torch.Tensor, b: torch.Tensor) -> bool: + """Same shape, dtype and bits.""" + return a.shape == b.shape and a.dtype == b.dtype and torch.equal(_bits(a), _bits(b)) + + +def _same_outputs(got, want) -> bool: + return len(got) == len(want) and all(_same(a, b) for a, b in zip(got, want)) + + +@functools.lru_cache(maxsize=None) +def _bias() -> torch.Tensor: + gen = torch.Generator(device=DEV).manual_seed(20260929) + return (torch.randn(E, generator=gen, device=DEV) * 0.1).float() + + +def _draw(seed: int, m: int): + """Router logits fp32 [m, 896] (std 2.5) and the latent bf16 [m, 3584] from a seed.""" + gen = torch.Generator(device=DEV).manual_seed(seed) + logits = (torch.randn(m, E, generator=gen, device=DEV) * 2.5).float() + return logits, (torch.randn(m, H, generator=gen, device=DEV) * 0.7).bfloat16() + + +def _stock(logits, latent): + return torch.ops.trtllm.kimi_k3_noaux_tc_mxfp8_quant(logits, _bias(), latent, RSF) + + +def _edge_case(case: str): + """8 tokens of an edge case: router logits, bias and latent.""" + gen = torch.Generator(device=DEV).manual_seed(20260930) + m = 8 + bias = _bias().clone() + latent = (torch.randn(m, H, generator=gen, device=DEV) * 0.7).bfloat16() + logits = torch.randn(m, E, generator=gen, device=DEV) * 2.5 + if case == "40_tied_keys": + # 40 equal keys compete for the top 16: ties go to the lower expert id. + logits[:, 100:140] = 3.0 + bias[100:140] = 0.25 + elif case == "all_equal_logits": + logits = torch.full((m, E), 0.3, device=DEV) + bias = torch.zeros(E, device=DEV) + elif case == "huge_logits": + logits = torch.randn(m, E, generator=gen, device=DEV) * 40.0 + elif case == "zero_large_denormal_rows": + latent[0] = 0 + latent[1, :64] = 0 + latent[2] = latent[2] * 3e4 + latent[3, ::7] = torch.tensor(1e-39).bfloat16() + latent[4, 5] = torch.tensor(-3e38).bfloat16() + else: + raise ValueError(case) + return logits.float().contiguous(), bias, latent + + +@contextlib.contextmanager +def _cold_compile_cache(): + """The op's compile cache emptied for the duration (and restored after), so that the next call is a first call.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_route_quant import op + + saved = dict(op._compiled) + op._compiled.clear() + try: + yield + finally: + op._compiled.clear() + op._compiled.update(saved) + + +def _raised_under_capture(fn) -> str: + """Run ``fn`` under CUDA-graph capture; the message of the RuntimeError it raised, or '' if it raised none.""" + graph, stream = torch.cuda.CUDAGraph(), torch.cuda.Stream() + resting = torch.cuda.current_stream() + stream.wait_stream(resting) + message = "" + try: + with torch.cuda.graph(graph, stream=stream): + try: + fn() + except RuntimeError as exc: + message = str(exc) + finally: + torch.cuda.set_stream(resting) + del graph + return message + + +@pytest.mark.parametrize("m", M_ALL) +def test_single_call(m): + """The stock op's four outputs bit for bit, with and without the early trigger. + + Shapes and dtypes as the contract states; run-to-run identical; the rows of each M the bits of the same rows of + the 64-token call. + """ + logits64, latent64 = _draw(1, 64) + logits, latent = logits64[:m].contiguous(), latent64[:m].contiguous() + want = _stock(logits, latent) + full = k3_route_quant(logits64, _bias(), latent64, RSF) + for early in (False, True): + got = k3_route_quant(logits, _bias(), latent, RSF, early_trigger=early) + ids, weights, quantized, scales = got + assert ids.dtype == torch.int32 and tuple(ids.shape) == (m, K) + assert weights.dtype == torch.bfloat16 and tuple(weights.shape) == (m, K) + assert quantized.dtype == torch.float8_e4m3fn and tuple(quantized.shape) == (m, H) + assert scales.dtype == torch.uint8 and tuple(scales.shape) == (m, H // 32) + assert all(t.is_contiguous() and t.device == logits.device for t in got) + same = [_same(a, b) for a, b in zip(got, want)] + rerun = _same_outputs( + k3_route_quant(logits, _bias(), latent, RSF, early_trigger=early), got + ) + rows = all(_same(a, b[:m]) for a, b in zip(got, full)) + print( + f"OPCHECK op=k3_route_quant M={m} early_trigger={early} " + f"stock_bits(ids,w,q,sf)={same} det={rerun} rows_as_m64={rows}" + ) + assert all(same) and rerun and rows + + +@pytest.mark.parametrize("case", EDGE_CASES) +def test_edge_case(case): + """Ties, saturation and zero / large / denormal latent rows at M 1, 3 and 8: the stock op's bits.""" + logits8, bias, latent8 = _edge_case(case) + for m in (1, 3, 8): + logits, latent = logits8[:m].contiguous(), latent8[:m].contiguous() + want = torch.ops.trtllm.kimi_k3_noaux_tc_mxfp8_quant(logits, bias, latent, RSF) + for early in (False, True): + got = k3_route_quant(logits, bias, latent, RSF, early_trigger=early) + same = [_same(a, b) for a, b in zip(got, want)] + print(f"OPCHECK op=k3_route_quant case={case} M={m} early_trigger={early} same={same}") + assert all(same), f"{case} M {m} early_trigger {early}: {same}" + + +def test_graph_capture_and_replay(): + """A captured sequence (M 8 with the early trigger, M 3 without, M 64 with) replayed with rewritten inputs. + + Between replays an eager call of another M (5, 1, 40, 2). Every replayed and eager call returns the bits of the + same call made alone. + """ + plan = ((8, True), (3, False), (64, True)) + static = [_draw(100 + i, m) for i, (m, _) in enumerate(plan)] + + def step(): + return [ + k3_route_quant(lg, _bias(), x, RSF, early_trigger=early) + for (lg, x), (_, early) in zip(static, plan) + ] + + step() # every build compiled eagerly + torch.cuda.synchronize() + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + outs = step() + for rep in range(REPLAYS): + inputs = [_draw(200 + 10 * rep + i, m) for i, (m, _) in enumerate(plan)] + alone = [ + k3_route_quant(lg, _bias(), x, RSF, early_trigger=early) + for (lg, x), (_, early) in zip(inputs, plan) + ] + eager_logits, eager_latent = _draw(300 + rep, (5, 1, 40, 2)[rep]) + eager_alone = k3_route_quant(eager_logits, _bias(), eager_latent, RSF) + torch.cuda.synchronize() + for (lg, x), (new_lg, new_x) in zip(static, inputs): + lg.copy_(new_lg) + x.copy_(new_x) + graph.replay() + torch.cuda.synchronize() + bad = [i for i, (got, want) in enumerate(zip(outs, alone)) if not _same_outputs(got, want)] + assert not bad, f"replay {rep}: calls {bad} differ from the same calls alone" + eager = k3_route_quant(eager_logits, _bias(), eager_latent, RSF) + assert _same_outputs(eager, eager_alone), f"eager call after replay {rep}" + del graph + + +def test_unsupported_calls_refused(): + """M 0 and 65, a strided or mistyped input and a mis-sized bias raise ValueError; the next call is correct.""" + logits, latent = _draw(400, 4) + want = k3_route_quant(logits, _bias(), latent, RSF) + for m in (0, 65): + bad_logits = torch.zeros(m, E, device=DEV) + bad_latent = torch.zeros(m, H, dtype=torch.bfloat16, device=DEV) + with pytest.raises(ValueError): + k3_route_quant(bad_logits, _bias(), bad_latent, RSF) + wide = torch.zeros(4, 2 * E, device=DEV) + wide[:, :E] = logits + with pytest.raises(ValueError): + k3_route_quant(wide[:, :E], _bias(), latent, RSF) # a strided view of the logits + with pytest.raises(ValueError): + k3_route_quant(logits.bfloat16(), _bias(), latent, RSF) # bf16 logits + with pytest.raises(ValueError): + k3_route_quant(logits, _bias(), latent.half(), RSF) # fp16 latent + with pytest.raises(ValueError): + k3_route_quant(logits, _bias()[:-1].contiguous(), latent, RSF) # 895 biases + with pytest.raises(ValueError): + k3_route_quant(logits, _bias(), latent[:3].contiguous(), RSF) # M differs + assert _same_outputs(k3_route_quant(logits, _bias(), latent, RSF), want) + + +def test_first_call_of_a_build_refuses_capture(): + """The kernel compiles on the first call of each build (early trigger, PDL): under CUDA-graph capture that call + raises RuntimeError instead of compiling into the capture, for both early-trigger builds.""" + logits, latent = _draw(500, 2) + want = k3_route_quant(logits, _bias(), latent, RSF) + for early in (False, True): + with _cold_compile_cache(): + message = _raised_under_capture( + lambda early=early: k3_route_quant( + logits, _bias(), latent, RSF, early_trigger=early + ) + ) + assert "outside CUDA-graph capture" in message, f"early_trigger {early}: {message!r}" + assert _same_outputs(k3_route_quant(logits, _bias(), latent, RSF), want) From dcb4d88b130d67b498adbf470d0297df8e5793ef Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 20:58:31 -0700 Subject: [PATCH 072/161] [None][test] Kimi K3 MoE front: the head geometry each TP size runs, without a GPU trtllm::k3_moe_front runs its head in 64-row half-tiles only when they fit beside the shared tiles in the device's free clusters, which at Kimi K3's shapes is W 16 alone; a 4-rank run never reaches that path. The new test checks the choice at W 4, 8 and 16, the cluster boundary, and that front_weight's rows are the rows each plan reads. It imports the CuTe DSL kernel modules, so it is listed with the B200 kernel tests. Signed-off-by: Vasanth Sabavat --- .../test_lists/test-db/l0_b200.yml | 1 + .../kimi_k3/test_k3_moe_front_geometry.py | 88 +++++++++++++++++++ 2 files changed, 89 insertions(+) create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front_geometry.py diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index 4c7f1e397e61..0f2d55dacd77 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -322,6 +322,7 @@ l0_b200: - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_wide.py - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_route_quant.py + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front_geometry.py - unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe.py - unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_route_quant.py # ------------- Visual Gen tests --------------- diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front_geometry.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front_geometry.py new file mode 100644 index 000000000000..2e64e1f20f2b --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front_geometry.py @@ -0,0 +1,88 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""The Kimi K3 MoE front's geometry choice and the weight rows each plan reads, without a GPU. + +trtllm::k3_moe_front runs the head in 64-row half-tiles, one round with every head k-tile of a CTA on chip before +the grid wait (``half_geometry``), only when those half-tiles fit beside the shared tiles in the GEMV clusters the +device leaves next to the two role clusters; otherwise in 128-row tiles (``geometry``). With 384 shared columns +(Kimi K3 TP16's per-rank width) and a GB200's 15 clusters of 8 CTAs, that is W = 16 alone: 280 head rows per rank +are 5 half-tiles, + 6 shared tiles = 11 GEMV clusters of the 13 left. At W = 4 (1120 rows, 18 half-tiles) and W = 8 +(560 rows, 9 half-tiles) they do not fit. A 4-rank run therefore never reaches the half-tile path; this checks the +selection, and that ``front_weight``'s rows are the rows each plan's weight descriptor covers. +""" + +import pytest +import torch + +kernel = pytest.importorskip("tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.k3_moe_front") +front_op = pytest.importorskip("tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.front_op") + +K = 7168 # the MoE input's width +SHARED_COLS = 384 +RING = 4 +# Clusters of 8 CTAs, one per SM, that a GB200 holds at once (the kernel module's statement). +GB200_CLUSTERS = 15 + + +def _plan(world: int, max_clusters: int = GB200_CLUSTERS): + """(head tiles of the plan, shared tiles, GEMV clusters, half-tile head) as trtllm::k3_moe_front picks them.""" + half = kernel.half_geometry(world, SHARED_COLS, max_clusters, K, RING) + if half is not None: + return (*half, True) + return (*kernel.geometry(world, SHARED_COLS, max_clusters), False) + + +def test_half_tile_head_only_at_w16(): + """W 16 runs the head as 5 half-tiles in one round of 11 GEMV clusters; W 4 and W 8 run 128-row tiles.""" + assert kernel.half_geometry(16, SHARED_COLS, GB200_CLUSTERS, K, RING) == (5, 6, 11) + assert _plan(16) == (5, 6, 11, True) + assert kernel.half_geometry(8, SHARED_COLS, GB200_CLUSTERS, K, RING) is None + assert _plan(8) == (5, 6, 11, False) + assert kernel.half_geometry(4, SHARED_COLS, GB200_CLUSTERS, K, RING) is None + assert _plan(4) == (9, 6, 13, False) + for world in (4, 8, 16): + assert kernel.supports(world, SHARED_COLS, GB200_CLUSTERS, K, RING), world + + +def test_half_tile_boundary(): + """At W 16 the half-tile plan needs 11 GEMV clusters beside the 2 role clusters: 13 clusters fit it, 12 do not.""" + assert _plan(16, max_clusters=13) == (5, 6, 11, True) + assert _plan(16, max_clusters=12)[-1] is False + + +@pytest.mark.parametrize("world", [4, 8, 16]) +def test_front_weight_rows_match_the_plan(world): + """``front_weight``'s rows are the rows the plan's weight descriptor covers, and the shared rows start where the + plan's first shared tile reads; the half-tiles stay inside the zero-padded head rows.""" + head_rows = kernel.head_rows(world) + head = torch.zeros(head_rows, K, dtype=torch.bfloat16) + gate_up = torch.zeros(2 * SHARED_COLS, K, dtype=torch.bfloat16) + rows = front_op.front_weight(head, gate_up).shape[0] + padded_head = kernel.head_tiles(world) * kernel.CTA_M + assert rows == padded_head + 2 * SHARED_COLS + n_head, n_shared, _, half = _plan(world) + assert kernel._weight_rows(world, n_head, n_head + n_shared, half) == rows + my_tiles = K // kernel.CTA_K // kernel.SPLIT + shift, _ = kernel._half_plan(half, world, n_head, RING, my_tiles) + assert n_head * kernel.CTA_M + shift == padded_head # the first shared tile's first weight row + if half: + assert head_rows <= n_head * kernel.HALF_M <= padded_head + + +def test_front_weight_packs_the_shared_rows(): + """The head rows first, zero-padded to whole 128-row tiles; then every 32 rows hold 16 gate rows and the 16 up rows + of the same columns.""" + world, k = 16, 64 + head_rows = kernel.head_rows(world) + head = torch.arange(head_rows, dtype=torch.float32)[:, None].expand(head_rows, k).contiguous() + gate_up = (10_000 + torch.arange(2 * SHARED_COLS, dtype=torch.float32))[:, None].expand(-1, k) + w = front_op.front_weight(head, gate_up.contiguous()) + padded_head = kernel.head_tiles(world) * kernel.CTA_M + assert torch.equal(w[:head_rows], head) + assert bool((w[head_rows:padded_head] == 0).all()) + shared = w[padded_head:, 0] + for block in range(SHARED_COLS // 16): + gate = 10_000 + torch.arange(16 * block, 16 * block + 16, dtype=torch.float32) + up = gate + SHARED_COLS + assert torch.equal(shared[32 * block : 32 * block + 16], gate) + assert torch.equal(shared[32 * block + 16 : 32 * block + 32], up) From d0c5e7586830a31da59bfa8af53728b7c22a6107 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 20:58:38 -0700 Subject: [PATCH 073/161] [None][chore] MNNVL split all-gather fake: format with the repo's hooks Signed-off-by: Vasanth Sabavat --- tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py b/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py index ba6a25655281..95460ad6d213 100644 --- a/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py +++ b/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py @@ -133,7 +133,8 @@ def _(input, bf16_columns, world_size, comm_buffer, buffer_flags): return [ input.new_empty((num_tokens, world_size * bf16_columns), dtype=torch.bfloat16), - input.new_empty((num_tokens, world_size * (columns - bf16_columns))), + input.new_empty( + (num_tokens, world_size * (columns - bf16_columns))), ] # MNNVL Allreduce From f8a29e3e0748c2431f2de1aee3ece9b5ba2caba0 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 20:59:00 -0700 Subject: [PATCH 074/161] [None][test] Kimi K3 MoE front test: a latent code may differ only at a reference rounding tie The test's head reference is a torch fp32 matmul, which sums in another order than the front's split-K GEMV. Where the reference's bf16 latent sits on an E4M3 rounding midpoint, the front's code is the adjacent value: one or a few codes per M from M 2, and at the top of the range one step is 16 block-scale units, which the old 1-unit gate rejected (from M 5 on, every rank). A code now differs from the reference's only where that explains it: the same block scale, adjacent E4M3 codes, and the reference latent within two bf16 ulps of their midpoint. The 99.9 % equal-code gate stays. Signed-off-by: Vasanth Sabavat --- .../kimi_k3/test_k3_moe_front.py | 32 ++++++++++++++++--- 1 file changed, 27 insertions(+), 5 deletions(-) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py index f3e046826e1f..f789e9808487 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py @@ -23,8 +23,10 @@ trtllm::kimi_k3_noaux_tc_mxfp8_quant; shared: cuBLAS gate_up -> trtllm::situ_and_mul): top-16 ids per token (a mismatch only at a reference 16th / 17th key margin below 1e-4: the split-K head sums in another order), routing weights, MXFP8 codes and - scales (> 99.9 % equal, dequantized within one block-scale unit), the shared activation within 2e-2; the same - routing and latent bits on every rank; the head buffers empty and the buffer index flipped after each call; + scales (> 99.9 % equal; for the same reason a code may differ only at a reference rounding tie: the + adjacent E4M3 value under the same block scale, the reference latent within two bf16 ulps of the two + values' midpoint), the shared activation within 2e-2; the same routing and latent bits on every rank; the + head buffers empty and the buffer index flipped after each call; fused : y against the TRTLLM-Gen W4A8_MXFP4_MXFP8 runner on the front's own routing and MXFP8 latent (op-catalog gates), the shared activation the front's bits, k3_moe's scratch re-armed; head_flags : fused with the ready-word handoff (k3_moe built with head_flags) across the head epoch's int32 wrap: @@ -210,7 +212,25 @@ def _reference(ctx, x): shared = torch.ops.trtllm.situ_and_mul(F.linear(x, ctx.gate_up), GATE_CAP, LINEAR_CAP) key = torch.sigmoid(logits) + ctx.bias top = key.sort(dim=1, descending=True).values - return ids, w, q, s, shared, top[:, TOP_K - 1] - top[:, TOP_K] + return ids, w, q, s, shared, top[:, TOP_K - 1] - top[:, TOP_K], latent + + +def _latent_mismatch_not_near_tie(q, s, r_q, r_s, r_latent) -> int: + """MXFP8 latent codes that differ from the reference's anywhere but at a reference rounding tie. The front sums + the head in its own split-K order, so its bf16 latent can sit one bf16 ulp from the reference's; where that ulp + straddles an E4M3 rounding midpoint the code moves to the adjacent value. A difference is allowed only there: + the same block scale, adjacent codes, and the reference latent within two bf16 ulps of their midpoint.""" + qb, rb = _bits(q).int(), _bits(r_q).int() + mismatch = qb != rb + same_scale = (_bits(s) == _bits(r_s)).repeat_interleave(SV, dim=1) + # E4M3 codes of one sign are ordered by magnitude: adjacent values differ by one in the 7 magnitude bits. + adjacent = ((qb & 0x80) == (rb & 0x80)) & (((qb & 0x7F) - (rb & 0x7F)).abs() == 1) + scale = torch.pow(2.0, r_s.float() - 127.0).repeat_interleave(SV, dim=1) + mid = (q.float() + r_q.float()) * 0.5 * scale + lat = r_latent.float() + bf16_ulp = torch.pow(2.0, torch.floor(torch.log2(lat.abs().clamp_min(2.0**-126))) - 7) + near_tie = (lat - mid).abs() <= 2 * bf16_ulp + return int((mismatch & ~(same_scale & adjacent & near_tie)).sum()) def _dequant(q, s): @@ -232,7 +252,7 @@ def check_front(ctx): torch.cuda.synchronize() flag1 = int(ctx.ws.flags[0].item()) empty = _quiet_check(ctx, lambda: _buffers_empty(ctx)) - r_ids, r_w, r_q, r_s, r_shared, margin = _reference(ctx, x) + r_ids, r_w, r_q, r_s, r_shared, margin, r_latent = _reference(ctx, x) # The selected experts per token (their order inside the top 16 may differ at near-equal keys) and each # selected expert's weight. sets_equal = (ids.sort(dim=1).values == r_ids.sort(dim=1).values).all(dim=1) @@ -251,6 +271,7 @@ def check_front(ctx): weight_max_err=(dense[same_tok] - r_dense[same_tok]).abs().max().item() if same_tok.numel() else 0.0, codes_equal=(_bits(q) == _bits(r_q)).float().mean().item(), scales_equal=(s == r_s).float().mean().item(), latent_err_scale_units=((dq - rdq).abs() / torch.maximum(scale, rscale)).max().item(), + latent_mismatch_not_near_tie=_latent_mismatch_not_near_tie(q, s, r_q, r_s, r_latent), shared_rel=((shared.float() - r_shared.float()).abs().max() / r_shared.float().abs().max()).item(), det=all(all(_same(a, b) for a, b in zip(r, out)) for r in again), rows_as_m8=all(_same(a, b[:m]) for a, b in zip(out, out8)), @@ -258,7 +279,8 @@ def check_front(ctx): buffers_empty=empty and _quiet_check(ctx, lambda: _buffers_empty(ctx)), flag_flipped=flag1 != flag0, ) # fmt: skip good = (row["mismatch_not_near_tie"] == 0 and row["weight_max_err"] <= 0.01 * RSF and row["codes_equal"] > 0.999 - and row["scales_equal"] > 0.999 and row["latent_err_scale_units"] <= 1.0 and row["shared_rel"] <= 2e-2 + and row["scales_equal"] > 0.999 and row["latent_mismatch_not_near_tie"] == 0 + and row["shared_rel"] <= 2e-2 and row["det"] and row["rows_as_m8"] and row["ranks_agree"] and row["buffers_empty"] and row["flag_flipped"]) # fmt: skip row["rank"], row["good"] = ctx.rank, bool(good) From 8c0431e33c94b88146440de968cbbc19c79b77e4 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 21:00:29 -0700 Subject: [PATCH 075/161] [None][fix] Kimi K3 collectives: review nits K3LatentExchange.create refuses a TP size other than 4, 8 or 16 before any collective step. mnnvl_fusion_allreduce refuses a two-shot call on a workspace whose buffer size is not a multiple of 32 (the two-shot broadcast stage starts at half the buffer and moves 16-byte vectors). The k3_fused_moe package docstring names its three ops, and two matrix docstrings drop an internal run path. Signed-off-by: Vasanth Sabavat --- .../modeling_v2/catalog/comm/k3_latent_reduce.md | 4 ++-- .../modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md | 6 +++--- .../modeling_v2/catalog/comm/mnnvl_fusion_allreduce.py | 9 +++++++++ .../_torch/cute_dsl_kernels/k3_fused_moe/__init__.py | 7 ++++--- .../_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py | 4 ++++ .../modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py | 3 +-- .../comm/_mnnvl_fusion_allreduce_op_matrix.py | 3 +-- 7 files changed, 24 insertions(+), 12 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md index 4112de1901b7..24ddae59fa0c 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md @@ -98,8 +98,8 @@ of up to 8 tokens fits. - `fabric_handle`: share the memory by fabric handle (required across nodes) or POSIX file descriptor; default `mapping.is_multi_node()`. No environment variable is read. -`create` does not check `W`: it builds an exchange for any TP size, and the op then raises `ValueError` at every -call unless `W` is 4, 8 or 16 (the op's code). +`create` raises `ValueError` for a TP size other than 4, 8 or 16, on every rank and before any collective step (the +size is the same on every rank of the group; code). **Which ops may share one object.** One call's producers and this op. The producers are the push-only builds of the routed experts (`k3_moe_m1` / `k3_moe_m2` push, `trtllm::k3_fused_moe_push`, `trtllm::k3_fused_moe_front_push`; diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md index 279d61d57fa6..edffc4e9d2d6 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md @@ -247,9 +247,9 @@ follows `T`, `H` and the device's SM count, and in the fused form it sets the or - Gaps: the op does not check that the call fits `comm_buffer` (the attention-residual and all-gather ops do), so a direct op call over one buffer writes past it — the wrapper's `required_buffer_bytes` check is the guard; `MnnvlWorkspace.create` accepts any multiple of 16 bytes, but the two-shot broadcast stage starts at - `buffer_bytes / 2` and is accessed in 16-byte vectors, so a two-shot call needs `buffer_bytes` to be a multiple of - 32 (code; every buffer in the test is); the schema's default `one_shot_max_bytes=1048576` applies to a direct op - call (the wrapper always passes one). + `buffer_bytes / 2` and is accessed in 16-byte vectors, so the wrapper raises `ValueError` before a two-shot call + on a workspace whose `buffer_bytes` is not a multiple of 32 (code; every buffer in the test is a multiple); the + schema's default `one_shot_max_bytes=1048576` applies to a direct op call (the wrapper always passes one). - In the model today the workspace is `MNNVLAllReduce`'s (a dict keyed by `Mapping`, grown on demand by the first eager call that needs more, in 8 MiB steps) and the one-shot ceiling is a module attribute with a per-call override. This entry takes both explicitly: the workspace sized at construction, the ceiling per call. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.py index 9c48175014d3..20a6e5c5a4e4 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.py @@ -54,6 +54,15 @@ def mnnvl_fusion_allreduce( f"mnnvl_fusion_allreduce: the call needs {need} bytes per Lamport buffer, the workspace has " f"{workspace.buffer_bytes}" ) + two_shot = ( + num_tokens * hidden * workspace.world_size * input.element_size() > one_shot_max_bytes + ) + if two_shot and workspace.buffer_bytes % 32: + # The two-shot broadcast stage starts at buffer_bytes / 2 and is accessed in 16-byte vectors. + raise ValueError( + "mnnvl_fusion_allreduce: a two-shot call needs the workspace's buffer_bytes to be a multiple of 32, " + f"not {workspace.buffer_bytes}" + ) fusion_op = AllReduceFusionOp.RESIDUAL_RMS_NORM if fused else AllReduceFusionOp.NONE outputs = torch.ops.trtllm.mnnvl_fusion_allreduce( input, diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/__init__.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/__init__.py index c20fcf7c5dfa..57beb0dee6f5 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/__init__.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/__init__.py @@ -12,8 +12,9 @@ # 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. -"""Kimi K3 fused routed-expert decode path (``trtllm::k3_fused_moe``). +"""Kimi K3 MoE decode kernels: the routed experts (``trtllm::k3_moe``, :mod:`.op`), the MoE front +(``trtllm::k3_moe_front``, :mod:`.front_op`) and the latent reduce (``trtllm::k3_latent_reduce``, :mod:`.latent_op`). -Importing :mod:`.op` registers the torch op; nothing is imported eagerly here so that -the CuTe DSL dependency stays optional for every other model. +Importing a module registers its torch op; nothing is imported eagerly here so that the CuTe DSL dependency stays +optional for every other model. """ diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py index 7b4f23e6b689..5fee5cbb5c4c 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/latent_op.py @@ -93,6 +93,10 @@ def build(uc, mc, handle, comm): comm=comm, ) + if mapping.tp_size not in (4, 8, 16): + raise ValueError( + f"K3LatentExchange.create: TP size {mapping.tp_size}; the exchange supports 4, 8 and 16" + ) words = _kernel().buffer_words(mapping.tp_size) return create_mcast_state("K3LatentExchange", mapping, words, fabric_handle, build) diff --git a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py index 12babfbd03fc..0814b7cb9554 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py @@ -471,8 +471,7 @@ def check_two_workspaces_interleaved() -> None: """Two workspaces are two rotations: all-gathers of mixed token counts and splits alternate between them in an irregular pattern (A A B A B B ...), every call is correct and each workspace's flags move with its own calls only. The pattern is the same on every rank: calls on one stream are serialized and each waits for its peers, so - ranks issuing calls on two workspaces in different orders deadlock (measured for mnnvl_allreduce_attn_res, - runs/drafter/u4-mnnvl-srun-2).""" + ranks issuing calls on two workspaces in different orders deadlock (seen with mnnvl_allreduce_attn_res).""" pattern = "AABABBAAAB" * 2 for i, which in enumerate(pattern): t = INTERLEAVED_TOKENS[i % len(INTERLEAVED_TOKENS)] diff --git a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_fusion_allreduce_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_fusion_allreduce_op_matrix.py index 14114f8c348c..58ac651c2d65 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_fusion_allreduce_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_fusion_allreduce_op_matrix.py @@ -521,8 +521,7 @@ def check_two_workspaces_interleaved() -> None: """Two workspaces are two rotations: calls of mixed shapes and paths alternate between them in an irregular pattern (A A B A B B ...), every call is correct and each workspace's flags move with its own calls only. The pattern is the same on every rank: calls on one stream are serialized and each waits for its peers, so ranks - issuing calls on two workspaces in different orders deadlock (measured for mnnvl_allreduce_attn_res, - runs/drafter/u4-mnnvl-srun-2).""" + issuing calls on two workspaces in different orders deadlock (seen with mnnvl_allreduce_attn_res).""" pattern = "AABABBAAAB" * 2 for i, which in enumerate(pattern): t, hidden, fused, path = INTERLEAVED[i % len(INTERLEAVED)] From 9fbc51f81426e66e3fafbc3d71b806c454e50efa Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 21:06:12 -0700 Subject: [PATCH 076/161] [None][fix] Kimi K3 collective state: create() frees its communicator split when it raises As MnnvlWorkspace.create (G1): the workspaces' create() splits an MPI communicator for each call, and on the two failures every rank agrees on (a rank cannot allocate; the allocation failed on a rank) it raised and kept the split. Every rank now frees it on both paths; under Ray the communicator is c10d's TP group and is only dropped. The sandwich, latent exchange and MoE front matrices also assert that every rank freed the communicators its refused creates split. Signed-off-by: Vasanth Sabavat --- .../catalog/comm/k3_latent_reduce.md | 9 ++-- .../catalog/comm/k3_sandwich_oproj.md | 9 ++-- .../catalog/comm/k3_sandwich_plain.md | 9 ++-- .../catalog/comm/k3_sandwich_tail.md | 9 ++-- .../modeling_v2/catalog/moe/k3_moe_front.md | 18 ++++---- .../cute_dsl_kernels/k3_fused_moe/op.py | 15 +++++-- .../_torch/cute_dsl_kernels/k3_sandwich/op.py | 15 +++++-- .../comm/_k3_latent_reduce_op_matrix.py | 44 +++++++++++++------ .../comm/_k3_moe_front_op_matrix.py | 28 +++++++++--- .../modeling_v2/comm/_k3_sandwich_common.py | 31 +++++++++---- 10 files changed, 128 insertions(+), 59 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md index 24ddae59fa0c..9f67e9d0cdc1 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md @@ -85,10 +85,11 @@ of up to 8 tokens fits. the TP group's communicator, so a rank that calls it while its peers do not waits for them there; - failure model (`k3_fused_moe.op.create_mcast_state`, the same as `MnnvlWorkspace.create`'s): - before allocating, the ranks agree that each of them can (not capturing, the buffer within that rank's free - device memory). If one cannot, every rank raises `RuntimeError` and none allocates (certified: every rank - capturing, and one rank capturing while its peers call it eagerly at the same point; every rank raises at that - agreement, the capturing rank's message naming the capture, and the next call is correct); - - a failure that returns from the allocation is agreed the same way; + device memory). If one cannot, every rank raises `RuntimeError`, none allocates, and under MPI each frees the + communicator it split for the call (certified: every rank capturing, and one rank capturing while its peers + call it eagerly at the same point; every rank raises at that agreement and frees its split, the capturing + rank's message naming the capture, and the next call is correct); + - a failure that returns from the allocation is agreed and handled the same way; - a rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange; this is not turned into an error on the other ranks; - eager: it allocates and exchanges handles, so it refuses to run under CUDA-graph capture (every rank raises); diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md index dc0bee244c37..b1310543bd46 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md @@ -91,10 +91,11 @@ every word empty, every counter zero. - collective over `mapping`'s TP group: every rank calls it at the same point; - failure model: - before allocating, the ranks agree that each of them can (not capturing, the buffer within that rank's free - device memory). If one cannot, every rank raises `RuntimeError` and none allocates (certified: one rank - capturing while its peers call it eagerly, every rank raises, the capturing rank naming the capture and its - peers another rank; no rank reaches the allocation, and the next call is correct); - - a failure that returns from the allocation is agreed the same way; + device memory). If one cannot, every rank raises `RuntimeError`, none allocates, and under MPI each frees the + communicator it split for the call (certified: one rank capturing while its peers call it eagerly, every rank + raises, the capturing rank naming the capture and its peers another rank; no rank reaches the allocation, + every rank frees its split, and the next call is correct); + - a failure that returns from the allocation is agreed and handled the same way; - a rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange; this is not turned into an error on the other ranks; - eager: it allocates and exchanges handles, so it refuses to run under CUDA-graph capture (certified: every rank diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md index 5049373fd590..3b50dd1c8fd9 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md @@ -89,10 +89,11 @@ target's sandwiches'. Certified after `create`: sized for the group, every word - collective over `mapping`'s TP group: every rank calls it at the same point; - failure model: - before allocating, the ranks agree that each of them can (not capturing, the buffer within that rank's free - device memory). If one cannot, every rank raises `RuntimeError` and none allocates (certified: one rank - capturing while its peers call it eagerly, every rank raises, the capturing rank naming the capture and its - peers another rank; no rank reaches the allocation, and the next call is correct); - - a failure that returns from the allocation is agreed the same way; + device memory). If one cannot, every rank raises `RuntimeError`, none allocates, and under MPI each frees the + communicator it split for the call (certified: one rank capturing while its peers call it eagerly, every rank + raises, the capturing rank naming the capture and its peers another rank; no rank reaches the allocation, + every rank frees its split, and the next call is correct); + - a failure that returns from the allocation is agreed and handled the same way; - a rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange; this is not turned into an error on the other ranks; - eager: it allocates and exchanges handles, so it refuses to run under CUDA-graph capture (certified: every rank diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md index 1c2aea006675..6c069b684d37 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md @@ -121,10 +121,11 @@ CTAs, with the handle that owns the memory and the communicator. This op's parti - collective over `mapping`'s TP group: every rank calls it at the same point; - failure model: - before allocating, the ranks agree that each of them can (not capturing, the buffer within that rank's free - device memory). If one cannot, every rank raises `RuntimeError` and none allocates (certified: one rank - capturing while its peers call it eagerly, every rank raises, the capturing rank naming the capture and its - peers another rank; no rank reaches the allocation, and the next call is correct); - - a failure that returns from the allocation is agreed the same way; + device memory). If one cannot, every rank raises `RuntimeError`, none allocates, and under MPI each frees the + communicator it split for the call (certified: one rank capturing while its peers call it eagerly, every rank + raises, the capturing rank naming the capture and its peers another rank; no rank reaches the allocation, + every rank frees its split, and the next call is correct); + - a failure that returns from the allocation is agreed and handled the same way; - a rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange; this is not turned into an error on the other ranks; - eager: it allocates and exchanges handles, so it refuses to run under CUDA-graph capture (certified: every rank diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md index ac00b649759c..f7608ddea53d 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md @@ -121,15 +121,15 @@ fabric_handle=None)`: collective over `mapping`'s TP group and eager, every rank point. Under MPI each call first splits the group's communicator off the session's, a collective of every rank of the session (`_get_mnnvl_workspace_comm`). Before allocating, the ranks agree that each of them can: not capturing a CUDA graph, the buffer within its device's free memory. If one cannot, every rank raises `RuntimeError` ("not every rank -can allocate") and none allocates. A failure returned by the allocation is agreed the same way, and that second -agreement is also the barrier that keeps any rank from pushing into a peer's buffer before the peer has emptied it; a -rank that fails inside the allocation's handle exchange can leave its peers waiting there (the op module's -statements). Certified: with every rank capturing, and with the last rank capturing while its peers call it eagerly, -every rank raises `RuntimeError` (the capturing ranks' messages naming the capture) and the existing workspace keeps -its bits. So even a refusal needs every rank to call `create()`: a rank calling it alone waits for its peers. It -empties every word and zeroes `flags` and `ready` (certified, both of the matrix's workspaces). `fabric_handle`: share -the memory by fabric handle (required across nodes) or POSIX file descriptor; default `mapping.is_multi_node()`. No -environment variable is read. +can allocate"), none allocates, and under MPI each frees the communicator it split for the call. A failure returned by +the allocation is agreed and handled the same way, and that second agreement is also the barrier that keeps any rank +from pushing into a peer's buffer before the peer has emptied it; a rank that fails inside the allocation's handle +exchange can leave its peers waiting there (the op module's statements). Certified: with every rank capturing, and +with the last rank capturing while its peers call it eagerly, every rank raises `RuntimeError` (the capturing ranks' +messages naming the capture) and frees its split, and the existing workspace keeps its bits. So even a refusal needs +every rank to call `create()`: a rank calling it alone waits for its peers. It empties every word and zeroes `flags` +and `ready` (certified, both of the matrix's workspaces). `fabric_handle`: share the memory by fabric handle (required +across nodes) or POSIX file descriptor; default `mapping.is_multi_node()`. No environment variable is read. **Which ops may share one object.** Every MoE front call of the TP group: this entry, plain or publishing (`publish=True`, the front a head_flags `k3_moe` pairs with). The head_flags calls of `moe/k3_moe` (`head=workspace`) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py index 8f3a118ca86b..bc51c9e7a4d3 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py @@ -118,15 +118,17 @@ def create_mcast_state(name: str, mapping, words: int, fabric_handle: Optional[b Failure model (as ``MnnvlWorkspace.create``): before allocating, the ranks agree that each of them can (not capturing a CUDA graph, the buffer within its device's free memory); if one cannot, every rank raises - ``RuntimeError`` and none allocates. A failure that returns from the allocation or from ``build`` is agreed the - same way. A rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange: - that failure is not turned into an error on the other ranks.""" + ``RuntimeError``, none allocates, and under MPI each frees the communicator it split for the call. A failure that + returns from the allocation or from ``build`` is agreed and handled the same way. A rank that fails inside the + allocation's handle exchange can leave its peers waiting in that exchange: that failure is not turned into an + error on the other ranks.""" from tensorrt_llm._torch.distributed.ops import ( _get_mnnvl_workspace_comm, _make_mnnvl_mcast_buffer, _mnnvl_device_index, _mnnvl_workspace_all_succeeded, ) + from tensorrt_llm._utils import mpi_disabled use_fabric_handle = mapping.is_multi_node() if fabric_handle is None else bool(fabric_handle) comm = _get_mnnvl_workspace_comm(mapping) @@ -140,6 +142,9 @@ def create_mcast_state(name: str, mapping, words: int, fabric_handle: Optional[b if free_bytes < words * 4: problem = f"its {words * 4} bytes exceed the {free_bytes} free on this rank's device" if not _mnnvl_workspace_all_succeeded(comm, problem is None): + # Every rank takes this path: free the MPI communicator split above (a ProcessGroup is c10d's). + if not mpi_disabled(): + comm.Free() raise RuntimeError( f"{name}.create: not every rank can allocate ({problem or 'another rank cannot'})" ) @@ -156,6 +161,10 @@ def create_mcast_state(name: str, mapping, words: int, fabric_handle: Optional[b error = exc # Also the barrier that keeps any rank from pushing into a peer's buffer before the peer has emptied it. if not _mnnvl_workspace_all_succeeded(comm, error is None): + # Every rank takes this path too. A handle only borrows the communicator and makes no MPI call when + # destroyed, so it may outlive the communicator. + if not mpi_disabled(): + comm.Free() raise RuntimeError(f"{name}: allocation failed on at least one rank") from error return state diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py index c93bd9ff2458..00091954c490 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py @@ -136,15 +136,17 @@ def _create_buffer(cls, mapping, words: int, flag_words: int, fabric_handle: Opt Failure model (as ``MnnvlWorkspace.create``): before allocating, the ranks agree that each of them can (not capturing a CUDA graph, the buffer within its device's free memory); if one cannot, every rank raises - ``RuntimeError`` and none allocates. A failure that returns from the allocation is agreed the same way. A rank - that fails inside the allocation's handle exchange can leave its peers waiting in that exchange: that failure is - not turned into an error on the other ranks.""" + ``RuntimeError``, none allocates, and under MPI each frees the communicator it split for the call. A failure + that returns from the allocation is agreed and handled the same way. A rank that fails inside the allocation's + handle exchange can leave its peers waiting in that exchange: that failure is not turned into an error on the + other ranks.""" from tensorrt_llm._torch.distributed.ops import ( _get_mnnvl_workspace_comm, _make_mnnvl_mcast_buffer, _mnnvl_device_index, _mnnvl_workspace_all_succeeded, ) + from tensorrt_llm._utils import mpi_disabled use_fabric_handle = mapping.is_multi_node() if fabric_handle is None else bool(fabric_handle) comm = _get_mnnvl_workspace_comm(mapping) @@ -158,6 +160,9 @@ def _create_buffer(cls, mapping, words: int, flag_words: int, fabric_handle: Opt if free_bytes < words * 4: problem = f"its {words * 4} bytes exceed the {free_bytes} free on this rank's device" if not _mnnvl_workspace_all_succeeded(comm, problem is None): + # Every rank takes this path: free the MPI communicator split above (a ProcessGroup is c10d's). + if not mpi_disabled(): + comm.Free() raise RuntimeError( f"{cls.__name__}.create: not every rank can allocate ({problem or 'another rank cannot'})" ) @@ -185,6 +190,10 @@ def _create_buffer(cls, mapping, words: int, flag_words: int, fabric_handle: Opt error = exc # Also the barrier that keeps any rank from pushing into a peer's buffer before the peer has emptied it. if not _mnnvl_workspace_all_succeeded(comm, error is None): + # Every rank takes this path too. A handle only borrows the communicator and makes no MPI call when + # destroyed, so it may outlive the communicator. + if not mpi_disabled(): + comm.Free() raise RuntimeError(f"{cls.__name__}: allocation failed on at least one rank") from error return state diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_latent_reduce_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_k3_latent_reduce_op_matrix.py index 21c4e9527de3..83ca05f580fe 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_k3_latent_reduce_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_latent_reduce_op_matrix.py @@ -249,12 +249,13 @@ def check_exchange_is_armed_and_sized() -> None: def check_capture_refusals() -> None: """``K3LatentExchange.create`` is collective: every rank joins the TP group's communicator, and before allocating - the ranks agree that each of them can. With every rank capturing a CUDA graph, and with one rank capturing while - its peers call it eagerly at the same point, every rank raises RuntimeError, and a capturing rank's message names - the capture. Every rank raises at that agreement ("not every rank can allocate"; a failure after allocating reads - "allocation failed"), so nothing is allocated. The op's first call, which would compile the kernel, is refused per - rank: under capture it raises on every rank before it launches anything. The exchange is untouched and the next - call is correct. Runs before every eager call of the op: the compile cache must still be cold.""" + the ranks agree that each of them can. With every rank capturing a CUDA graph, and with one rank capturing while its + peers call it eagerly at the same point, every rank raises RuntimeError, and a capturing rank's message names the + capture. Every rank raises at that agreement ("not every rank can allocate"; a failure after allocating reads + "allocation failed"), so nothing is allocated, and each frees the communicator it split. The op's first call, which + would compile the kernel, is refused per rank: under capture it raises on every rank before it launches anything. + The exchange is untouched and the next call is correct. Runs before every eager call of the op: the compile cache + must still be cold.""" # Imported by the op's first call; imported here so that nothing is imported inside the capture. import cutlass.cute.runtime # noqa: F401 @@ -273,14 +274,29 @@ def refusal(capturing: bool) -> str: return str(exc) return "" - for case, capturing in (("every rank", True), ("one rank", R.rank == R.world - 1)): - R.barrier() - message = refusal(capturing) - refused = "not every rank can allocate" in message - named = "outside CUDA-graph capture" in message or not capturing - assert R.all_true(refused and named), ( - f"{case} capturing: rank {R.rank} (capturing {capturing}) got {message!r}" - ) + from tensorrt_llm._torch.distributed import ops + + split = ops._get_mnnvl_workspace_comm + comms = [] + + def recording_split(mapping): + comms.append(split(mapping)) + return comms[-1] + + ops._get_mnnvl_workspace_comm = recording_split + try: + for case, capturing in (("every rank", True), ("one rank", R.rank == R.world - 1)): + R.barrier() + message = refusal(capturing) + refused = "not every rank can allocate" in message + named = "outside CUDA-graph capture" in message or not capturing + assert R.all_true(refused and named), ( + f"{case} capturing: rank {R.rank} (capturing {capturing}) got {message!r}" + ) + finally: + ops._get_mnnvl_workspace_comm = split + freed = len(comms) == 2 and all(c == R.MPI.COMM_NULL for c in comms) + assert R.all_true(freed), "a refused create kept the communicator it split" R.barrier() first = raised_under_capture(lambda: entry(MAX_TOKENS, EX_A.state)) assert R.all_true("outside CUDA-graph capture first" in first), f"first call: {first!r}" diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py index d636d0bd91cb..ac5c76ef6f28 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py @@ -676,10 +676,13 @@ def check_create_and_first_compile_refuse_capture() -> None: K3MoeHeadWorkspace.create has the ranks agree before allocating that each of them can (not capturing, enough free memory): with every rank capturing, and with the last rank capturing while its peers call it eagerly, every rank raises RuntimeError at that agreement ("not every rank can allocate", a capturing rank's message naming the - capture), so none allocates. Every rank must call it: a rank calling it alone would wait for its peers. A front - call of a configuration not yet compiled (other SiTU caps) raises RuntimeError under capture before any launch. - The workspace keeps its bits and the next call returns the bits of the same call made alone. + capture), so none allocates, and each frees the communicator it split for the call. Every rank must call it: a + rank calling it alone would wait for its peers. A front call of a configuration not yet compiled (other SiTU caps) + raises RuntimeError under capture before any launch. The workspace keeps its bits and the next call returns the + bits of the same call made alone. """ + from tensorrt_llm._torch.distributed import ops + nxt = Call(7500, 4, layer=1, kind="front") want = alone_results([nxt], WS_A, "before the capture refusals")[0] before = workspace_snapshot(WS_A) @@ -687,10 +690,21 @@ def check_create_and_first_compile_refuse_capture() -> None: def create(): OPS.workspace.create(R.mapping, fabric_handle=R.fabric) - every = raised_under_capture(create) - R.barrier() + split = ops._get_mnnvl_workspace_comm + comms = [] + + def recording_split(mapping): + comms.append(split(mapping)) + return comms[-1] + capturing = R.world - 1 - one = raised_under_capture(create) if R.rank == capturing else raised_eagerly(create) + ops._get_mnnvl_workspace_comm = recording_split + try: + every = raised_under_capture(create) + R.barrier() + one = raised_under_capture(create) if R.rank == capturing else raised_eagerly(create) + finally: + ops._get_mnnvl_workspace_comm = split R.barrier() first = raised_under_capture( lambda: OPS.front( @@ -708,6 +722,8 @@ def create(): f"rank {R.rank}: every rank capturing {every!r}; rank {capturing} capturing {one!r}; " f"uncompiled front {first!r}" ) + freed = len(comms) == 2 and all(c == R.MPI.COMM_NULL for c in comms) + assert R.all_true(freed), "a refused create kept the communicator it split" after = workspace_snapshot(WS_A) assert all(same(a, b) for a, b in zip(before, after)), ( "a refused call touched the head workspace" diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_common.py b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_common.py index 1fae85271a29..d59e6773d128 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_common.py +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_common.py @@ -520,19 +520,34 @@ def create_refuses_capture(workspace_type, ws, next_call: Call) -> None: rank raise RuntimeError. (a) Every rank captures: every rank raises, naming the capture. (b) The last rank captures while its peers call ``create`` eagerly at the same point: every rank raises, the capturing rank naming the capture, its peers saying another rank cannot. No rank reaches the allocation in either case (the counted - multicast allocation, which the workspaces in use went through), no stream is left capturing, and the workspace - in use is untouched: the next call is correct.""" + multicast allocation, which the workspaces in use went through), each rank frees the communicator it split, no + stream is left capturing, and the workspace in use is untouched: the next call is correct.""" + from tensorrt_llm._torch.distributed import ops + assert ALLOCATIONS[0] > 0, "the counter did not see the workspaces in use being created" before = ALLOCATIONS[0] - every = _create(workspace_type, capture=True) - every_ok = CAPTURE_REFUSED in every - R.barrier() + split = ops._get_mnnvl_workspace_comm + comms = [] + + def recording_split(mapping): + comms.append(split(mapping)) + return comms[-1] + capturing = R.world - 1 - one = _create(workspace_type, capture=R.rank == capturing) + ops._get_mnnvl_workspace_comm = recording_split + try: + every = _create(workspace_type, capture=True) + R.barrier() + one = _create(workspace_type, capture=R.rank == capturing) + finally: + ops._get_mnnvl_workspace_comm = split + every_ok = CAPTURE_REFUSED in every one_ok = (CAPTURE_REFUSED if R.rank == capturing else PEER_REFUSED) in one allocated = ALLOCATIONS[0] - before - assert R.all_true(every_ok and one_ok and allocated == 0), ( - f"every rank capturing: {every!r}; rank {capturing} capturing: {one!r}; allocations reached {allocated}" + freed = len(comms) == 2 and all(c == R.MPI.COMM_NULL for c in comms) + assert R.all_true(every_ok and one_ok and allocated == 0 and freed), ( + f"every rank capturing: {every!r}; rank {capturing} capturing: {one!r}; allocations reached {allocated}; " + f"communicator splits freed {freed}" ) next_call.verify(next_call.run(ws), "after the refused creates") From 4ca4bfe436255c2b9caf265ebd243b2028965525 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 18:12:40 -0700 Subject: [PATCH 077/161] [None][feat] Kimi K3 DSpark drafter attention (trtllm::k3_drafter_attn, k3_drafter_attn_qknorm) A CuTe DSL kernel for the drafter's block attention, on SM 100. Each of up to 8 requests brings a draft block of up to 8 tokens. Each block attends densely to its request's paged context and to its own K / V, which are read from qkv and not written to the cache. GQA uses 6 query heads per KV head, head dim 64, HND pages of 64. - trtllm::k3_drafter_attn takes normalized and roped q / k. - trtllm::k3_drafter_attn_qknorm applies the per-head q / k RMSNorm and the NeoX RoPE in the kernel, with fused_qk_norm_rope's arithmetic. The kernel compiles on its first call for its head counts, page stride and whether there is more than one request. That first call must be outside CUDA-graph capture. Lengths, page tables and qkv are read on the device, so captured calls replay with them rewritten. test_k3_drafter_attn.py checks every split the engine schedules, with mixed context lengths: - against an fp64 reference; - against the DFlash TRTLLM path the kernel replaces (append_paged_kv_cache and the trtllm-gen context kernel, non-causal); - no NaN, and the pool left untouched; - reruns bit-identical, and requests isolated; - CUDA-graph replays with rewritten inputs. It runs in l0_b200.yml. Signed-off-by: Vasanth Sabavat --- .../cute_dsl_kernels/k3_drafter/__init__.py | 19 + .../k3_drafter/k3_drafter_attn_kernel.py | 969 ++++++++++++++++++ .../_torch/cute_dsl_kernels/k3_drafter/op.py | 349 +++++++ .../test_lists/test-db/l0_b200.yml | 2 + .../kimi_k3/test_k3_drafter_attn.py | 423 ++++++++ 5 files changed, 1762 insertions(+) create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_drafter/__init__.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_drafter/k3_drafter_attn_kernel.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_drafter/op.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_drafter_attn.py diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_drafter/__init__.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_drafter/__init__.py new file mode 100644 index 000000000000..051716c5f638 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_drafter/__init__.py @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Kimi K3 DSpark drafter CTM kernels (``trtllm::k3_drafter_attn``). + +Importing :mod:`.op` registers the torch ops; nothing is imported eagerly here so that +the CuTe DSL dependency stays optional for every other model. +""" diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_drafter/k3_drafter_attn_kernel.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_drafter/k3_drafter_attn_kernel.py new file mode 100644 index 000000000000..e70869719b68 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_drafter/k3_drafter_attn_kernel.py @@ -0,0 +1,969 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +# ============================================================================= +# Kimi K3 DSpark drafter block attention -- CTM (prims/cute) kernel: R <= 8 requests, a block of T <= 8 draft tokens +# each, GQA groups of 6 query heads per KV head, head dim 64 +# ============================================================================= +# +# o[t, h] = softmax_k( scale * q[t, h] . K[k] ) @ V[k] for every k < L = ctx + T (dense: no causal mask; the +# block attends to itself bidirectionally) +# q = qkv[t, 64 h : 64 h + 64] (after the q/k RMSNorm + RoPE), 48 rows per group r = 8 h + t (head-major) +# K, V = rows k < ctx: the drafter's paged cache (HND pages of 64 rows x 64, K and V blocks per head); +# rows ctx + t: the block's own k, v, straight from qkv (they are not appended to the cache) +# o = [R T, heads * 64] bf16 (the o_proj input) +# Request r owns rows r T .. r T + T - 1 of qkv, positions and o, length ctx = ctx_len[r] and page-table row r +# (page_table[r * table_stride ...]); t above is the token within its request. +# +# One cluster of 16 CTAs per (KV head, request): grid (16 kv_heads, R). CTA c takes the 128-row tiles c, c + 16, ... +# (two pages each) of its request. One request compiles without any request indexing, so that build is the +# single-sequence kernel; the indexing sits on the step's critical path. Per tile: +# S = Q K^T A = Q [64 rows (48 real), 64] K-major, B = the tile's K [128 rows, 64] K-major; M 64, N 128, into TMEM +# softmax each row's 128 scores over the 4 lanes of a quad (tcgen05.ld 16x256b), exp2 against the running max, +# P = bf16(p) into a 128B-swizzled K-major tile, fenced to the async proxy +# O = P V A = P [64, 128], B = V [128 rows, 64] read MN-major; M 64, N 64 +# fold the CTA's running (m, l, O) per row lives in the softmax threads' registers (32 fp32 each) +# Rows >= ctx of a tile are rewritten in shared memory after its TMA lands: the block's k / v from qkv (loaded while +# the TMA is in flight), and V zeros past L (a masked score gives p = 0, and 0 x a stale non-finite V would be NaN). +# Merge: every CTA (with or without tiles) stages its (m, l, O) of the 48 rows and st.async's row r's 64 + 2 floats +# into slot [cta] of CTA r / 3's mailbox, completing that CTA's barrier by bytes (the byte count is fixed, so it is +# armed at init). CTA c then merges its 3 rows over the 16 slots in fp32 (flash-decoding combine) and stores bf16. +# Nothing is read before griddepcontrol.wait: the context lengths, the page table and qkv are this step's. +# norm_rope: qkv holds the raw projection; the kernel applies the per-head q / k RMSNorm and the (NeoX) RoPE itself, +# with fused_qk_norm_rope's arithmetic (a warp per head row, two elements per lane): warps 4-7 build the 48 Q rows in +# shared memory (128 arrivals on q_full instead of the TMA), warp 1 the block's K rows. +# +# Warps: 0 TMA, 1 block rows / tail zeros, 2 TMEM allocation + MMA, 4-7 softmax / fold / push / merge, 3 idle. +# ============================================================================= +"""CTM kernel: Kimi K3 DSpark drafter block attention over the paged drafter cache, split-KV over a 16-CTA cluster +per request.""" + +from __future__ import annotations + +import math + +import cuda.bindings.driver as cuda_driver +import cutlass +import cutlass.cute as cute +import cutlass.experimental.cuda as cuda +from cutlass._mlir.dialects import llvm as _llvm +from cutlass._mlir.extras import types as _T +from cutlass.cutlass_dsl import dsl_user_op +from cutlass.experimental import primitives as prims + +HEADS = 6 # query heads per KV head (GQA group) +MAX_TOKENS = 8 +ROWS = HEADS * MAX_TOKENS # 48 query rows per group +D = 64 # head dim (one 128-byte swizzle row) +PAGE = 64 +TILE = 128 # KV rows per tile = 2 pages +PAGES_PER_TILE = TILE // PAGE +CLUSTER = 16 +OWN = ROWS // CLUSTER # 3 output rows merged per CTA +MMA_M = 64 +MMA_K = 16 +THREADS = 256 +EPI_THREADS = 128 +TMEM_COLS = 256 +TMEM_O = TILE # O's 64 columns after S's 128 +ELEM_BYTES = 2 +LOG2E = 1.4426950408889634 +NEG_INF = float("-inf") + +# Shared-memory tiles, bf16 elements. +Q_ELEMS = MMA_M * D # the M = 64 MMA reads 16 rows past the 48 of the group +KV_ELEMS = TILE * D # 16 KB: rows [0, 64) page 0, [64, 128) page 1 +PAGE_ELEMS = PAGE * D +P_HALF_ELEMS = MMA_M * 64 # P [64 rows][128 kv] as two 128B-swizzled [64][64] halves +P_ELEMS = 2 * P_HALF_ELEMS + +# Descriptor offsets in 16-byte units. +LEADING = 16 +SBO = 8 * D * ELEM_BYTES # 1024 B between 8-row swizzle atoms +STEP_K = (MMA_K * ELEM_BYTES) >> 4 # K-major: 32 B per K16 step +STEP_MN = (2 * SBO) >> 4 # MN-major: 16 rows per K16 step +P_HALF_U = (P_HALF_ELEMS * ELEM_BYTES) >> 4 + +# Mailbox of the 3 rows a CTA merges: [16 sources][3 rows][64] fp32 and [16][3][2] (m, l). +MAIL_O = CLUSTER * OWN * D +MAIL_ML = CLUSTER * OWN * 2 +MAIL_BYTES = (MAIL_O + MAIL_ML) * 4 + +io_dtype = cutlass.BFloat16 + + +@dsl_user_op +def _mapa_u32(smem_ptr, peer, *, loc=None, ip=None): + """The shared::cluster address of this CTA's shared-memory location in cluster CTA ``peer``.""" + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [smem_ptr.toint(loc=loc, ip=ip).ir_value(), cutlass.Int32(peer).ir_value(loc=loc, ip=ip)], + "mapa.shared::cluster.u32 $0, $1, $2;", "=r,r,r", has_side_effects=False, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _try_wait_cluster(mbar, parity, *, loc=None, ip=None): + """mbarrier.try_wait.parity.acquire.cluster on a barrier of this CTA (``mbar``, its shared-memory pointer): as + ``_test_wait_cluster``, with try_wait's bounded suspend.""" + done = cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [mbar.toint(loc=loc, ip=ip).ir_value(), cutlass.Int32(parity).ir_value(loc=loc, ip=ip)], + "{\n\t.reg .pred p;\n\tmbarrier.try_wait.parity.acquire.cluster.shared::cta.b64 p, [$1], $2;\n\t" + "selp.u32 $0, 1, 0, p;\n\t}", "=r,r,r,~{memory}", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + return done != cutlass.Int32(0) + + +@dsl_user_op +def _st_async_v2(dst, a, b, mbar, *, loc=None, ip=None): + """st.async of two fp32 to a shared::cluster address, completing ``mbar`` (shared::cluster) by 8 bytes.""" + _llvm.inline_asm( + None, + [cutlass.Int32(dst).ir_value(loc=loc, ip=ip), cutlass.Float32(a).ir_value(loc=loc, ip=ip), + cutlass.Float32(b).ir_value(loc=loc, ip=ip), cutlass.Int32(mbar).ir_value(loc=loc, ip=ip)], + "st.async.shared::cluster.mbarrier::complete_tx::bytes.v2.f32 [$0], {$1, $2}, [$3];", "r,f,f,r", + has_side_effects=True, is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _st_async_v4(dst, a, b, c, d, mbar, *, loc=None, ip=None): + """st.async of four fp32 to a shared::cluster address, completing ``mbar`` (shared::cluster) by 16 bytes.""" + _llvm.inline_asm( + None, + [cutlass.Int32(dst).ir_value(loc=loc, ip=ip), cutlass.Float32(a).ir_value(loc=loc, ip=ip), + cutlass.Float32(b).ir_value(loc=loc, ip=ip), cutlass.Float32(c).ir_value(loc=loc, ip=ip), + cutlass.Float32(d).ir_value(loc=loc, ip=ip), cutlass.Int32(mbar).ir_value(loc=loc, ip=ip)], + "st.async.shared::cluster.mbarrier::complete_tx::bytes.v4.f32 [$0], {$1, $2, $3, $4}, [$5];", + "r,f,f,f,f,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _approx(op, x, *, loc=None, ip=None): + """One MUFU op (rsqrt / ex2 / lg2 / sin / cos .approx.f32), as CUDA's rsqrtf / exp2f / __log2f / __sincosf.""" + return cutlass.Float32( + _llvm.inline_asm( + _T.f32(), [cutlass.Float32(x).ir_value(loc=loc, ip=ip)], f"{op}.approx.f32 $0, $1;", "=f,f", + has_side_effects=False, is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +def _norm_rope_pair(x0, x1, w0, w1, lane, pos_f, eps, rope_c): + """fused_qk_norm_rope on one head row held by a warp, lane l with elements (2 l, 2 l + 1): RMSNorm over the 64 + elements (butterfly sum), x * (rrms * w), then NeoX RoPE against the partner lane l ^ 16. Returns fp32.""" + ss = cutlass.Float32(0.0) + x0 * x0 + ss = ss + x1 * x1 + for mask in (16, 8, 4, 2, 1): + ss = ss + cute.arch.shuffle_sync_bfly(ss, offset=mask) + r = _approx("rsqrt", ss * cutlass.Float32(1.0 / D) + eps) + y0 = x0 * (r * w0) + y1 = x1 * (r * w1) + p0 = cute.arch.shuffle_sync_bfly(y0, offset=16) + p1 = cute.arch.shuffle_sync_bfly(y1, offset=16) + first_half = lane < cutlass.Int32(16) + p0 = cutlass.Float32(cutlass.select_(first_half, -p0, p0)) + p1 = cutlass.Float32(cutlass.select_(first_half, -p1, p1)) + hd0 = cutlass.Float32((cutlass.Int32(2) * lane) & cutlass.Int32(31)) + hd1 = cutlass.Float32((cutlass.Int32(2) * lane + cutlass.Int32(1)) & cutlass.Int32(31)) + th0 = pos_f * _approx("ex2", hd0 * rope_c) + th1 = pos_f * _approx("ex2", hd1 * rope_c) + o0 = y0 * _approx("cos", th0) + p0 * _approx("sin", th0) + o1 = y1 * _approx("cos", th1) + p1 * _approx("sin", th1) + return o0, o1 + + +def _plus(base, x): + """``base + x``; ``x`` itself when ``base`` is the single-request build's Python 0 (no op in the IR, so that build + is the single-request kernel instruction for instruction).""" + if type(base) is int and base == 0: + return x + return base + x + + +def _swz(row, chunk): + """Element index of 16-byte vector `chunk` (< 8) of `row` in a 128B-swizzled [rows][64] bf16 tile.""" + return row * cutlass.Int32(D) + ((chunk ^ (row % cutlass.Int32(8))) * cutlass.Int32(8)) + + +@cute.kernel +def k3_drafter_attn_kernel( + tma_q: cutlass.GridConstant[ + cuda.TensorMap + ], # qkv's q columns as (64 cols, token, head): box 64 x 8 x 6 + tma_kv: cutlass.GridConstant[ + cuda.TensorMap + ], # the layer's cache as (64 cols, 64 rows, K/V x head, page) + qkv: cutlass.Array, # [R T, (heads + 2 kv_heads) * 64] bf16, after the q/k RMSNorm + RoPE + page_table: cutlass.Array, # int32, row r at r * table_stride: >= ceil((ctx_len[r] + T) / 64) pages of request r + ctx_len: cutlass.Array, # int32 [R]: cached rows before each request's block + out: cutlass.Array, # [R T, heads * 64] bf16 + q_w: cutlass.Array, # norm_rope: [64] bf16 q_norm weight + k_w: cutlass.Array, # norm_rope: [64] bf16 k_norm weight + positions: cutlass.Array, # norm_rope: int32 or int64 [R T] RoPE positions of the blocks' tokens + num_tokens: cutlass.Int32, # T: tokens per request + scale_log2: cutlass.Float32, # softmax scale * log2(e) + norm_eps: cutlass.Float32, + rope_base: cutlass.Float32, + table_stride: cutlass.Int32, # elements between page-table rows (last: the other parameters keep their offsets) + total_heads: cutlass.Constexpr[int], # query heads of the rank: one cluster per 6 + kv_heads: cutlass.Constexpr[int], + norm_rope: cutlass.Constexpr[bool], + multi_request: cutlass.Constexpr[bool], # R > 1 (one request: request 0, no request indexing) +): + tx, _, _ = cute.arch.thread_idx() + warp_id = cute.arch.warp_idx() + cta = cute.arch.block_idx_in_cluster() + if cutlass.const_expr(multi_request): + bid, req, _ = cute.arch.block_idx() + tok0 = req * num_tokens # the request's first row of qkv, positions and out + table_row = req * table_stride + else: + # One request: no request arithmetic at all (Python zeros, see _plus). + bid, _, _ = cute.arch.block_idx() + req = tok0 = table_row = 0 + g_kv = bid // cutlass.Int32(CLUSTER) # KV head; query heads 6 g .. 6 g + 5 + ptr_q = tma_q.get_ptr() + ptr_kv = tma_kv.get_ptr() + qkv_cols = (total_heads + 2 * kv_heads) * D + + smem_q = cutlass.Array(io_dtype, Q_ELEMS, space=cutlass.AddressSpace.smem, alignment=1024) + smem_k = cutlass.Array(io_dtype, KV_ELEMS, space=cutlass.AddressSpace.smem, alignment=1024) + smem_v = cutlass.Array(io_dtype, KV_ELEMS, space=cutlass.AddressSpace.smem, alignment=1024) + smem_p = cutlass.Array(io_dtype, P_ELEMS, space=cutlass.AddressSpace.smem, alignment=1024) + # This CTA's final (m, l, O) of the 48 rows, staged for the push: O [48][64] fp32, (m, l) [48][2]. + stage_o = cutlass.Array( + cutlass.Float32, ROWS * D, space=cutlass.AddressSpace.smem, alignment=16 + ) + stage_ml = cutlass.Array( + cutlass.Float32, ROWS * 2, space=cutlass.AddressSpace.smem, alignment=16 + ) + mail_o = cutlass.Array(cutlass.Float32, MAIL_O, space=cutlass.AddressSpace.smem, alignment=16) + mail_ml = cutlass.Array(cutlass.Float32, MAIL_ML, space=cutlass.AddressSpace.smem, alignment=16) + q_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + kv_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + # The tile with rows >= ctx lands on its own barrier (used once), so warp 1 waits for exactly that load. + kv_fix_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + kv_fix = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + s_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + p_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + o_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + o_drained = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + mail_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + tmem_ptr_i32 = cutlass.Array(cutlass.Int32, 1, space=cutlass.AddressSpace.smem) + + if warp_id == 0: + prims.prefetch_tensormap(ptr_q) + prims.prefetch_tensormap(ptr_kv) + if prims.elect_sync(): + prims.mbarrier_init(q_full, EPI_THREADS if norm_rope else 1) + prims.mbarrier_init(kv_full, 1) + prims.mbarrier_init(kv_fix_full, 1) + prims.mbarrier_init(kv_fix, 32) + prims.mbarrier_init(s_full, 1) + prims.mbarrier_init(p_full, EPI_THREADS) + prims.mbarrier_init(o_full, 1) + prims.mbarrier_init(o_drained, EPI_THREADS) + # Every CTA of the cluster pushes all 48 rows: the byte count does not depend on the length. + prims.mbarrier_init(mail_full, 1) + prims.mbarrier_arrive_expect_tx(mail_full, MAIL_BYTES) + if warp_id == 2: + prims.tcgen05_alloc(tmem_ptr_i32, TMEM_COLS) + prims.tcgen05_relinquish_alloc_permit() + prims.fence_mbarrier_init() + prims.barrier_cluster_arrive_relaxed() + prims.barrier_cluster_wait() + prims.barrier_cta_sync(0) + tmem_base = tmem_ptr_i32.load() + + # Everything below reads this step's data (length, pages, qkv): after the grid dependency. + prims.griddepcontrol(prims.GridDepAction.WAIT) + ctx = ctx_len.load(idx=req) + kv_len = ctx + num_tokens + n_tiles = (kv_len + cutlass.Int32(TILE - 1)) // cutlass.Int32(TILE) + n_pages = (kv_len + cutlass.Int32(PAGE - 1)) // cutlass.Int32(PAGE) + my_count = cutlass.Int32(0) + if cta < n_tiles: + my_count = (n_tiles - cta + cutlass.Int32(CLUSTER - 1)) // cutlass.Int32(CLUSTER) + # The one tile of this CTA (if any) that holds rows >= ctx: the block's rows and the zeros past L. + fix_tile = ctx // cutlass.Int32(TILE) + fix_k = cutlass.Int32(-1) + if (fix_tile % cutlass.Int32(CLUSTER)) == cta: + fix_k = fix_tile // cutlass.Int32(CLUSTER) + # A block that straddles a tile boundary puts its tail in the next tile (another CTA). + fix_tile2 = (kv_len - cutlass.Int32(1)) // cutlass.Int32(TILE) + if fix_tile2 != fix_tile: + if (fix_tile2 % cutlass.Int32(CLUSTER)) == cta: + fix_k = fix_tile2 // cutlass.Int32(CLUSTER) + + if warp_id == 0: + # ===================================================================== + # TMA: Q once; each tile's two pages of K and V (a page past L is not + # loaded: warp 1 zeroes its rows). + # ===================================================================== + if prims.elect_sync(): + if cutlass.const_expr(not norm_rope): + if my_count > cutlass.Int32(0): + # The box's 8 token rows start at the request's first; rows >= T are another request's (or + # zeros past the last row) and only feed output rows that are not stored. + prims.mbarrier_arrive_expect_tx(q_full, ROWS * D * ELEM_BYTES) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_q, ptr_q, (cutlass.Int32(0), cutlass.Int32(tok0), g_kv * cutlass.Int32(HEADS)), q_full, + ) # fmt: skip + for k in range(my_count): + tile = cta + k * cutlass.Int32(CLUSTER) + if k > 0: + # K / V are free once the previous tile's P V is done. + while not cute.arch.mbarrier_try_wait( + o_full.data_ptr(), (k - cutlass.Int32(1)) & cutlass.Int32(1) + ): + pass + n_here = cutlass.Int32(0) + for p in cutlass.range_constexpr(PAGES_PER_TILE): + if tile * cutlass.Int32(PAGES_PER_TILE) + cutlass.Int32(p) < n_pages: + n_here = n_here + cutlass.Int32(1) + if k == fix_k: + prims.mbarrier_arrive_expect_tx( + kv_fix_full, n_here * cutlass.Int32(2 * PAGE_ELEMS * ELEM_BYTES) + ) + for p in cutlass.range_constexpr(PAGES_PER_TILE): + page_idx = tile * cutlass.Int32(PAGES_PER_TILE) + cutlass.Int32(p) + if page_idx < n_pages: + page = page_table.load(idx=_plus(table_row, page_idx)) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_k.subview(p * PAGE_ELEMS), ptr_kv, + (cutlass.Int32(0), cutlass.Int32(0), g_kv, page), kv_fix_full, + ) # fmt: skip + prims.cp_async_bulk_tensor_shared_cta_global( + smem_v.subview(p * PAGE_ELEMS), ptr_kv, + (cutlass.Int32(0), cutlass.Int32(0), cutlass.Int32(kv_heads) + g_kv, page), kv_fix_full, + ) # fmt: skip + else: + prims.mbarrier_arrive_expect_tx( + kv_full, n_here * cutlass.Int32(2 * PAGE_ELEMS * ELEM_BYTES) + ) + for p in cutlass.range_constexpr(PAGES_PER_TILE): + page_idx = tile * cutlass.Int32(PAGES_PER_TILE) + cutlass.Int32(p) + if page_idx < n_pages: + page = page_table.load(idx=_plus(table_row, page_idx)) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_k.subview(p * PAGE_ELEMS), ptr_kv, + (cutlass.Int32(0), cutlass.Int32(0), g_kv, page), kv_full, + ) # fmt: skip + prims.cp_async_bulk_tensor_shared_cta_global( + smem_v.subview(p * PAGE_ELEMS), ptr_kv, + (cutlass.Int32(0), cutlass.Int32(0), cutlass.Int32(kv_heads) + g_kv, page), kv_full, + ) # fmt: skip + # The o_proj after this kernel may launch and stream its weights; it waits for this grid before reading o. + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + elif warp_id == 1: + # ===================================================================== + # Rows >= ctx of the tile that holds them: the block's k / v from qkv + # (loaded while the tile's TMA is in flight), zeros for V past L; then + # release the MMA warp. + # ===================================================================== + lane = tx % 32 + if fix_k >= cutlass.Int32(0): + tile = cta + fix_k * cutlass.Int32(CLUSTER) + row0 = tile * cutlass.Int32(TILE) + # Lane l: block token l / 8 (+ 4 on the second pass), 16-byte vector l % 8 of its k and v. + c = lane % cutlass.Int32(8) + kvecs = [] + vvecs = [] + for j in cutlass.range_constexpr(MAX_TOKENS // 4): + t = lane // cutlass.Int32(8) + cutlass.Int32(4 * j) + t_c = _plus( + tok0, cutlass.Int32(cutlass.select_(t < num_tokens, t, cutlass.Int32(0))) + ) + kvecs.append( + qkv.load( + idx=t_c * cutlass.Int32(qkv_cols) + + (cutlass.Int32(total_heads) + g_kv) * cutlass.Int32(D) + + c * cutlass.Int32(8), + vector_size=8, + alignment=16, + ) + ) + vvecs.append( + qkv.load( + idx=t_c * cutlass.Int32(qkv_cols) + + (cutlass.Int32(total_heads + kv_heads) + g_kv) * cutlass.Int32(D) + + c * cutlass.Int32(8), + vector_size=8, + alignment=16, + ) + ) + if cutlass.const_expr(norm_rope): + # K rows: token t's raw k (lane: elements 2 lane, 2 lane + 1), normalized and roped like Q. + kw0 = cutlass.Float32(k_w.load(idx=cutlass.Int32(2) * lane)) + kw1 = cutlass.Float32(k_w.load(idx=cutlass.Int32(2) * lane + cutlass.Int32(1))) + rope_c = (cutlass.Float32(-2.0) * _approx("lg2", rope_base)) * cutlass.Float32( + 1.0 / D + ) + kraw = [] + kpos = [] + for t in cutlass.range_constexpr(MAX_TOKENS): + t_c = _plus( + tok0, + cutlass.Int32( + cutlass.select_( + cutlass.Int32(t) < num_tokens, cutlass.Int32(t), cutlass.Int32(0) + ) + ), + ) + kraw.append( + qkv.load( + idx=t_c * cutlass.Int32(qkv_cols) + + (cutlass.Int32(total_heads) + g_kv) * cutlass.Int32(D) + + cutlass.Int32(2) * lane, + vector_size=2, + alignment=4, + ) + ) + kpos.append(cutlass.Int32(positions.load(idx=t_c)).to(cutlass.Float32)) + krot = [] + for t in cutlass.range_constexpr(MAX_TOKENS): + o0, o1 = _norm_rope_pair( + cutlass.Float32(kraw[t][0]), + cutlass.Float32(kraw[t][1]), + kw0, + kw1, + lane, + kpos[t], + norm_eps, + rope_c, + ) + krot.append( + cutlass.Vector.from_elements((o0.to(io_dtype), o1.to(io_dtype)), io_dtype) + ) + while not cute.arch.mbarrier_try_wait(kv_fix_full.data_ptr(), 0): + pass + for j in cutlass.range_constexpr(MAX_TOKENS // 4): + t = lane // cutlass.Int32(8) + cutlass.Int32(4 * j) + r = ctx + t - row0 # the block row's row in this tile + if (t < num_tokens) & (r >= cutlass.Int32(0)) & (r < cutlass.Int32(TILE)): + if cutlass.const_expr(not norm_rope): + smem_k.store(kvecs[j], idx=_swz(r, c), vector_size=8, alignment=16) + smem_v.store(vvecs[j], idx=_swz(r, c), vector_size=8, alignment=16) + if cutlass.const_expr(norm_rope): + for t in cutlass.range_constexpr(MAX_TOKENS): + r = ctx + cutlass.Int32(t) - row0 + if ( + (cutlass.Int32(t) < num_tokens) + & (r >= cutlass.Int32(0)) + & (r < cutlass.Int32(TILE)) + ): + smem_k.store( + krot[t], + idx=_swz(r, lane // cutlass.Int32(4)) + + (cutlass.Int32(2) * lane) % cutlass.Int32(8), + vector_size=2, + alignment=4, + ) + # V rows >= L of the tile are zero: a masked score gives p = 0, and 0 x a stale non-finite V is NaN. + # (K rows >= L only feed masked scores.) + first_zero = kv_len - row0 + zeros = [] + for e in cutlass.range_constexpr(8): + zeros.append(cutlass.BFloat16(0.0)) + zero_v = cutlass.Vector.from_elements(tuple(zeros), io_dtype) + for j in cutlass.range_constexpr(TILE * 8 // 32): + vi = lane + cutlass.Int32(32 * j) + r = vi // cutlass.Int32(8) + if r >= first_zero: + smem_v.store( + zero_v, idx=_swz(r, vi % cutlass.Int32(8)), vector_size=8, alignment=16 + ) + prims.fence_proxy(prims.Proxy.ASYNC_SHARED, space=prims.SharedSpace.shared_cta) + prims.mbarrier_arrive(kv_fix) + elif warp_id == 2: + # ===================================================================== + # MMA: S = Q K^T (N 128), then O = P V (N 64), per tile. + # ===================================================================== + idesc_s = prims.Tcgen05InstrDesc.build( + c_dtype=cutlass.Float32, a_dtype=io_dtype, b_dtype=io_dtype, n_dim=TILE, m_dim=MMA_M + ) + idesc_o = prims.Tcgen05InstrDesc.build( + c_dtype=cutlass.Float32, + a_dtype=io_dtype, + b_dtype=io_dtype, + n_dim=D, + m_dim=MMA_M, + b_major=1, + ) + swz = prims.Tcgen05SmemSwizzle.SWIZZLE_128B + desc_q = prims.Tcgen05SmemDesc.build( + start_address=smem_q, leading_byte_offset=LEADING, stride_byte_offset=SBO, layout=swz + ) + desc_k = prims.Tcgen05SmemDesc.build( + start_address=smem_k, leading_byte_offset=LEADING, stride_byte_offset=SBO, layout=swz + ) + desc_p = prims.Tcgen05SmemDesc.build( + start_address=smem_p, leading_byte_offset=LEADING, stride_byte_offset=SBO, layout=swz + ) + desc_v = prims.Tcgen05SmemDesc.build( + start_address=smem_v, + leading_byte_offset=KV_ELEMS * ELEM_BYTES, + stride_byte_offset=SBO, + layout=swz, + ) + tmem_s = cutlass.inttoptr(tmem_base, 6, cutlass.Int32) + tmem_o = cutlass.inttoptr(tmem_base + cutlass.Int32(TMEM_O), 6, cutlass.Int32) + if my_count > cutlass.Int32(0): + while not cute.arch.mbarrier_try_wait(q_full.data_ptr(), 0): + pass + for k in range(my_count): + phase = k & cutlass.Int32(1) + if k == fix_k: + # Warp 1 saw the load land and rewrote the rows >= ctx. + while not cute.arch.mbarrier_try_wait(kv_fix.data_ptr(), 0): + pass + else: + # kv_full completes once per tile other than the fix tile. + seen_fix = cutlass.Int32( + cutlass.select_( + (fix_k >= cutlass.Int32(0)) & (fix_k < k), + cutlass.Int32(1), + cutlass.Int32(0), + ) + ) + while not cute.arch.mbarrier_try_wait( + kv_full.data_ptr(), (k - seen_fix) & cutlass.Int32(1) + ): + pass + if k > 0: + # S's columns are free once the softmax warps have read the previous S (they wrote its P). + while not cute.arch.mbarrier_try_wait(p_full.data_ptr(), phase ^ cutlass.Int32(1)): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + for kk in cutlass.range_constexpr(D // MMA_K): + if prims.elect_sync(): + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, prims.CTAGroup.CTA_1, tmem_s, + desc_q + kk * STEP_K, desc_k + kk * STEP_K, idesc_s, kk != 0, + ) # fmt: skip + if prims.elect_sync(): + prims.tcgen05_commit(s_full) + while not cute.arch.mbarrier_try_wait(p_full.data_ptr(), phase): + pass + if k > 0: + # O's columns are free once the softmax warps have folded the previous O. + while not cute.arch.mbarrier_try_wait( + o_drained.data_ptr(), phase ^ cutlass.Int32(1) + ): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + for kk in cutlass.range_constexpr(TILE // MMA_K): + if prims.elect_sync(): + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, prims.CTAGroup.CTA_1, tmem_o, + desc_p + ((kk // (64 // MMA_K)) * P_HALF_U + (kk % (64 // MMA_K)) * STEP_K), + desc_v + kk * STEP_MN, idesc_o, kk != 0, + ) # fmt: skip + if prims.elect_sync(): + prims.tcgen05_commit(o_full) + elif warp_id >= 4: + # ===================================================================== + # Softmax, fold, stage, push; then merge this CTA's 3 rows. + # 16x256b: lane l of warp w holds rows 16 w + l / 4 and + 8, columns + # 8 g + 2 (l % 4) + {0, 1} of every 8-column group g. + # ===================================================================== + tid = tx - cutlass.Int32(EPI_THREADS) + lane = tx % 32 + w = warp_id - 4 + quad = lane % cutlass.Int32(4) + r0 = w * cutlass.Int32(16) + lane // cutlass.Int32(4) + r1 = r0 + cutlass.Int32(8) + if cutlass.const_expr(norm_rope): + if my_count > cutlass.Int32(0): + # Q rows 8 h + t: warp w takes (h, t) pairs w, w + 4, ...; lane l elements 2 l, 2 l + 1 of the row. + qw0 = cutlass.Float32(q_w.load(idx=cutlass.Int32(2) * lane)) + qw1 = cutlass.Float32(q_w.load(idx=cutlass.Int32(2) * lane + cutlass.Int32(1))) + rope_c = (cutlass.Float32(-2.0) * _approx("lg2", rope_base)) * cutlass.Float32( + 1.0 / D + ) + qraw = [] + qpos = [] + for i in cutlass.range_constexpr(ROWS // 4): + row = w + cutlass.Int32(4 * i) + h = row // cutlass.Int32(MAX_TOKENS) + t = row % cutlass.Int32(MAX_TOKENS) + t_c = _plus( + tok0, cutlass.Int32(cutlass.select_(t < num_tokens, t, cutlass.Int32(0))) + ) + qraw.append( + qkv.load( + idx=t_c * cutlass.Int32(qkv_cols) + + (g_kv * cutlass.Int32(HEADS) + h) * cutlass.Int32(D) + + cutlass.Int32(2) * lane, + vector_size=2, + alignment=4, + ) + ) + qpos.append(cutlass.Int32(positions.load(idx=t_c)).to(cutlass.Float32)) + for i in cutlass.range_constexpr(ROWS // 4): + row = w + cutlass.Int32(4 * i) + o0, o1 = _norm_rope_pair( + cutlass.Float32(qraw[i][0]), + cutlass.Float32(qraw[i][1]), + qw0, + qw1, + lane, + qpos[i], + norm_eps, + rope_c, + ) + smem_q.store( + cutlass.Vector.from_elements((o0.to(io_dtype), o1.to(io_dtype)), io_dtype), + idx=_swz(row, lane // cutlass.Int32(4)) + (cutlass.Int32(2) * lane) % cutlass.Int32(8), + vector_size=2, alignment=4, + ) # fmt: skip + prims.fence_proxy(prims.Proxy.ASYNC_SHARED, space=prims.SharedSpace.shared_cta) + prims.mbarrier_arrive(q_full) + m_run0 = cutlass.Float32(NEG_INF) + m_run1 = cutlass.Float32(NEG_INF) + l_run0 = cutlass.Float32(0.0) + l_run1 = cutlass.Float32(0.0) + o_run = [cutlass.Float32(0.0)] * (2 * D // 4) # rows r0, r1 x 16 columns: [4 g + 2 row + e] + for k in range(my_count): + phase = k & cutlass.Int32(1) + tile = cta + k * cutlass.Int32(CLUSTER) + kv0 = tile * cutlass.Int32(TILE) + while not cute.arch.mbarrier_try_wait(s_full.data_ptr(), phase): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + s = prims.tcgen05_ld( + "16x256b", cutlass.inttoptr(tmem_base, 6, cutlass.Float32), num=TILE // 8 + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + m0 = cutlass.Float32(NEG_INF) + m1 = cutlass.Float32(NEG_INF) + v0 = [] + v1 = [] + for g in cutlass.range_constexpr(TILE // 8): + for e in cutlass.range_constexpr(2): + col = kv0 + cutlass.Int32(8 * g + e) + cutlass.Int32(2) * quad + a = cutlass.Float32( + cutlass.select_( + col < kv_len, + cutlass.Float32(s[4 * g + e]) * scale_log2, + cutlass.Float32(NEG_INF), + ) + ) + b = cutlass.Float32( + cutlass.select_( + col < kv_len, + cutlass.Float32(s[4 * g + 2 + e]) * scale_log2, + cutlass.Float32(NEG_INF), + ) + ) + v0.append(a) + v1.append(b) + m0 = cute.arch.fmax(m0, a) + m1 = cute.arch.fmax(m1, b) + for offset in (1, 2): + m0 = cute.arch.fmax(m0, cute.arch.shuffle_sync_bfly(m0, offset=offset)) + m1 = cute.arch.fmax(m1, cute.arch.shuffle_sync_bfly(m1, offset=offset)) + m_new0 = cute.arch.fmax(m_run0, m0) + m_new1 = cute.arch.fmax(m_run1, m1) + base0 = cutlass.Float32( + cutlass.select_(m_new0 == cutlass.Float32(NEG_INF), cutlass.Float32(0.0), m_new0) + ) + base1 = cutlass.Float32( + cutlass.select_(m_new1 == cutlass.Float32(NEG_INF), cutlass.Float32(0.0), m_new1) + ) + l0 = cutlass.Float32(0.0) + l1 = cutlass.Float32(0.0) + for g in cutlass.range_constexpr(TILE // 8): + pv0 = [] + pv1 = [] + for e in cutlass.range_constexpr(2): + pa = cute.math.exp2(v0[2 * g + e] - base0, fastmath=True) + pb = cute.math.exp2(v1[2 * g + e] - base1, fastmath=True) + l0 = l0 + pa + l1 = l1 + pb + pv0.append(pa) + pv1.append(pb) + # P [row][kv] bf16, K-major 128B swizzle: half g // 8 (= page), chunk g % 8, elements 2 quad, +1. + half = g // 8 + chunk = g % 8 + pair0 = cutlass.Vector.from_elements( + (pv0[0].to(io_dtype), pv0[1].to(io_dtype)), io_dtype + ) + pair1 = cutlass.Vector.from_elements( + (pv1[0].to(io_dtype), pv1[1].to(io_dtype)), io_dtype + ) + off0 = ( + cutlass.Int32(half * P_HALF_ELEMS) + + _swz(r0, cutlass.Int32(chunk)) + + cutlass.Int32(2) * quad + ) + off1 = ( + cutlass.Int32(half * P_HALF_ELEMS) + + _swz(r1, cutlass.Int32(chunk)) + + cutlass.Int32(2) * quad + ) + smem_p.store(pair0, idx=off0, vector_size=2, alignment=4) + smem_p.store(pair1, idx=off1, vector_size=2, alignment=4) + for offset in (1, 2): + l0 = l0 + cute.arch.shuffle_sync_bfly(l0, offset=offset) + l1 = l1 + cute.arch.shuffle_sync_bfly(l1, offset=offset) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + prims.fence_proxy(prims.Proxy.ASYNC_SHARED, space=prims.SharedSpace.shared_cta) + prims.mbarrier_arrive(p_full) + # The running (m, l, O) against the new max: alpha = exp2(m_run - m_new) (0 while nothing was seen). + alpha0 = cute.math.exp2(m_run0 - base0, fastmath=True) + alpha1 = cute.math.exp2(m_run1 - base1, fastmath=True) + l_run0 = l_run0 * alpha0 + l0 + l_run1 = l_run1 * alpha1 + l1 + m_run0 = m_new0 + m_run1 = m_new1 + while not cute.arch.mbarrier_try_wait(o_full.data_ptr(), phase): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + o = prims.tcgen05_ld( + "16x256b", + cutlass.inttoptr(tmem_base + cutlass.Int32(TMEM_O), 6, cutlass.Float32), + num=D // 8, + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + prims.mbarrier_arrive(o_drained) + for g in cutlass.range_constexpr(D // 8): + for e in cutlass.range_constexpr(2): + o_run[4 * g + e] = o_run[4 * g + e] * alpha0 + cutlass.Float32(o[4 * g + e]) + o_run[4 * g + 2 + e] = o_run[4 * g + 2 + e] * alpha1 + cutlass.Float32( + o[4 * g + 2 + e] + ) + # Stage the 48 rows (rows 48-63 of the M = 64 accumulator are padding). + for g in cutlass.range_constexpr(D // 8): + col = cutlass.Int32(8 * g) + cutlass.Int32(2) * quad + if r0 < cutlass.Int32(ROWS): + stage_o.store( + cutlass.Vector.from_elements((o_run[4 * g], o_run[4 * g + 1]), cutlass.Float32), + idx=r0 * cutlass.Int32(D) + col, vector_size=2, alignment=8, + ) # fmt: skip + if r1 < cutlass.Int32(ROWS): + stage_o.store( + cutlass.Vector.from_elements((o_run[4 * g + 2], o_run[4 * g + 3]), cutlass.Float32), + idx=r1 * cutlass.Int32(D) + col, vector_size=2, alignment=8, + ) # fmt: skip + if quad == cutlass.Int32(0): + if r0 < cutlass.Int32(ROWS): + stage_ml.store( + cutlass.Vector.from_elements((m_run0, l_run0), cutlass.Float32), + idx=r0 * cutlass.Int32(2), + vector_size=2, + alignment=8, + ) + if r1 < cutlass.Int32(ROWS): + stage_ml.store( + cutlass.Vector.from_elements((m_run1, l_run1), cutlass.Float32), + idx=r1 * cutlass.Int32(2), + vector_size=2, + alignment=8, + ) + prims.barrier_cta_sync(1, thread_count=EPI_THREADS) + # Push: 48 rows x 16 vectors of 4 fp32 into slot [cta] of the owner's mailbox (row r -> CTA r / 3, row r % 3). + for i in cutlass.range_constexpr(ROWS * (D // 4) // EPI_THREADS): + vi = tid + cutlass.Int32(i * EPI_THREADS) + row = vi // cutlass.Int32(D // 4) + c4 = vi % cutlass.Int32(D // 4) + own = row // cutlass.Int32(OWN) + vals = stage_o.load( + idx=row * cutlass.Int32(D) + c4 * cutlass.Int32(4), vector_size=4, alignment=16 + ) + _st_async_v4( + _mapa_u32( + mail_o.data_ptr( + (cta * cutlass.Int32(OWN) + row % cutlass.Int32(OWN)) * cutlass.Int32(D) + + c4 * cutlass.Int32(4) + ), + own, + ), + vals[0], + vals[1], + vals[2], + vals[3], + _mapa_u32(mail_full.data_ptr(), own), + ) + if tid < cutlass.Int32(ROWS): + own = tid // cutlass.Int32(OWN) + ml = stage_ml.load(idx=tid * cutlass.Int32(2), vector_size=2, alignment=8) + _st_async_v2( + _mapa_u32( + mail_ml.data_ptr( + (cta * cutlass.Int32(OWN) + tid % cutlass.Int32(OWN)) * cutlass.Int32(2) + ), + own, + ), + ml[0], + ml[1], + _mapa_u32(mail_full.data_ptr(), own), + ) + # ===================================================================== + # Merge rows 3 cta .. 3 cta + 2 over the 16 slots (fp32), store bf16. + # ===================================================================== + while not _try_wait_cluster(mail_full.data_ptr(), 0): + pass + if tid < cutlass.Int32(OWN * D // 2): + j = tid // cutlass.Int32(D // 2) + c2 = (tid % cutlass.Int32(D // 2)) * cutlass.Int32(2) + row = cta * cutlass.Int32(OWN) + j + ms = [] + for s_ in cutlass.range_constexpr(CLUSTER): + ms.append(mail_ml.load(idx=(cutlass.Int32(s_ * OWN) + j) * cutlass.Int32(2))) + mx = ms[0] + for s_ in cutlass.range_constexpr(1, CLUSTER): + mx = cute.arch.fmax(mx, ms[s_]) + den = cutlass.Float32(0.0) + acc0 = cutlass.Float32(0.0) + acc1 = cutlass.Float32(0.0) + for s_ in cutlass.range_constexpr(CLUSTER): + wgt = cutlass.Float32( + cutlass.select_( + ms[s_] == cutlass.Float32(NEG_INF), + cutlass.Float32(0.0), + cute.math.exp2(ms[s_] - mx, fastmath=True), + ) + ) + den = den + wgt * mail_ml.load( + idx=(cutlass.Int32(s_ * OWN) + j) * cutlass.Int32(2) + cutlass.Int32(1) + ) + ov = mail_o.load( + idx=(cutlass.Int32(s_ * OWN) + j) * cutlass.Int32(D) + c2, + vector_size=2, + alignment=8, + ) + acc0 = acc0 + wgt * ov[0] + acc1 = acc1 + wgt * ov[1] + inv = cutlass.Float32(1.0) / den + h_loc = row // cutlass.Int32(MAX_TOKENS) + t_out = row % cutlass.Int32(MAX_TOKENS) + if t_out < num_tokens: + out.store( + cutlass.Vector.from_elements( + ((acc0 * inv).to(io_dtype), (acc1 * inv).to(io_dtype)), io_dtype + ), + idx=( + _plus(tok0, t_out) * cutlass.Int32(total_heads) + + g_kv * cutlass.Int32(HEADS) + + h_loc + ) + * cutlass.Int32(D) + + c2, + vector_size=2, + alignment=4, + ) + # Every TMEM reader has waited for its loads; no peer writes this CTA's mailbox after its barrier completed. + prims.barrier_cta_sync(0) + if warp_id == 2: + prims.tcgen05_dealloc(cutlass.inttoptr(tmem_base, 6, cutlass.Int32), TMEM_COLS) + + +def _q_map(qkv, num_rows, total_heads, kv_heads): + """qkv's query columns as (64 columns, token, head): one call at (token r T, head 6 g) lands the group's + [6 heads][8 tokens][64] of request r (row 8 h + t; rows past the last token zero).""" + qkv_cols = (total_heads + 2 * kv_heads) * D + return cuda.create_tensor_map_tiled( + global_address=qkv.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[D, num_rows, total_heads], + global_strides=[(qkv_cols * ELEM_BYTES) // 16, (D * ELEM_BYTES) // 16], + box_dims=[D, MAX_TOKENS, HEADS], + swizzle=cuda.TensorMapSwizzle.s128b, + ) + + +def _kv_map(cache, total_pages, kv_heads, page_stride): + """The layer's HND cache [pages, 2, kv_heads, 64, 64] (page stride `page_stride` elements) as (64 columns, 64 rows, + K / V x head, page): one call lands one head's K or V page.""" + return cuda.create_tensor_map_tiled( + global_address=cache.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[D, PAGE, 2 * kv_heads, total_pages], + global_strides=[ + (D * ELEM_BYTES) // 16, + (PAGE * D * ELEM_BYTES) // 16, + (page_stride * ELEM_BYTES) // 16, + ], + box_dims=[D, PAGE, 1, 1], + swizzle=cuda.TensorMapSwizzle.s128b, + ) + + +@cute.jit +def k3_drafter_attn( + qkv: cute.Tensor, # [R T * (heads + 2 kv_heads) * 64] bf16 + cache: cute.Tensor, # the layer's cache, pointer carrier (first page) + page_table: cute.Tensor, # int32: R rows of pages, table_stride apart + ctx_len: cute.Tensor, # int32 [R] + out: cute.Tensor, # [R T * heads * 64] bf16 + q_w: cute.Tensor, # norm_rope: [64] bf16 + k_w: cute.Tensor, # norm_rope: [64] bf16 + positions: cute.Tensor, # norm_rope: int32 or int64 [R T] + num_tokens: cutlass.Int32, # T: tokens per request + num_requests: cutlass.Int32, # R + table_stride: cutlass.Int32, + scale_log2: cutlass.Float32, + total_pages: cutlass.Int32, + norm_eps: cutlass.Float32, + rope_base: cutlass.Float32, + total_heads: cutlass.Constexpr[int], + kv_heads: cutlass.Constexpr[int], + page_stride: cutlass.Constexpr[int], + norm_rope: cutlass.Constexpr[bool], + multi_request: cutlass.Constexpr[bool], + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + tma_q = _q_map(qkv, num_tokens * num_requests, total_heads, kv_heads) + tma_kv = _kv_map(cache, total_pages, kv_heads, page_stride) + k3_drafter_attn_kernel( + tma_q, + tma_kv, + qkv, + page_table, + ctx_len, + out, + q_w, + k_w, + positions, + num_tokens, + scale_log2, + norm_eps, + rope_base, + table_stride, + total_heads, + kv_heads, + norm_rope, + multi_request, + ).launch( # fmt: skip + grid=(CLUSTER * kv_heads, num_requests, 1), + block=(THREADS, 1, 1), + cluster=(CLUSTER, 1, 1), + stream=stream, + use_pdl=use_pdl, + ) + + +def softmax_scale_log2(head_dim: int = D) -> float: + return LOG2E / math.sqrt(head_dim) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_drafter/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_drafter/op.py new file mode 100644 index 000000000000..3f5f2522190a --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_drafter/op.py @@ -0,0 +1,349 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Torch ops of the Kimi K3 DSpark drafter CTM kernels. + +``trtllm::k3_drafter_attn``: the draft blocks of R <= 8 requests (T <= 8 tokens each, M = R T rows of ``qkv``), each +attending densely to its request's pages of the drafter's paged cache plus its own K / V (taken from ``qkv``, not +appended to the cache), GQA groups of 6 query heads per KV head, head dim 64, pages of 64 rows (HND). + +Requests: R = ``ctx_len.numel()`` (int32 [R], the cached rows before each block) and T = M / R. ``page_table`` is +int32 with dense rows: [>= R, width] (any row stride, e.g. rows of the pool's block table) or, for R = 1, one row +[width]. Request r reads rows r T .. r T + T - 1 of ``qkv`` (and ``positions``), page-table row r and +``ctx_len[r]``, and writes rows r T .. of ``out``. + +Compiled on the first call for its head counts, cache page stride and whether R > 1 (not for R or T otherwise), +which must happen outside CUDA-graph capture. +""" + +from __future__ import annotations + +import os +import threading +from typing import Dict + +import torch + +MAX_TOKENS = 8 # per request +MAX_REQUESTS = 8 +HEADS_PER_KV = 6 +HEAD_DIM = 64 +PAGE = 64 + +_lock = threading.Lock() +_compiled: Dict[tuple, object] = {} +_arg_dummies: Dict[tuple, torch.Tensor] = {} + + +def _dummy(device: torch.device, dtype: torch.dtype) -> torch.Tensor: + """A tensor argument the build does not touch (the norm / RoPE inputs when the kernel does not apply them).""" + key = (device.index, dtype) + t = _arg_dummies.get(key) + if t is None: + t = _arg_dummies[key] = torch.zeros(MAX_TOKENS, dtype=dtype, device=device) + return t + + +def _arg(t: torch.Tensor, align: int = 16): + from cutlass.cute.runtime import from_dlpack + + return from_dlpack(t.detach(), assumed_align=align).mark_layout_dynamic(leading_dim=t.dim() - 1) + + +def _use_pdl() -> bool: + return os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" + + +def supports_attn( + qkv: torch.Tensor, cache: torch.Tensor, num_heads: int, num_kv_heads: int, num_requests: int = 1 +) -> bool: + """Whether ``k3_drafter_attn`` takes the call: ``qkv`` [M = R T, (heads + 2 kv_heads) * 64] dense bf16 of + R = ``num_requests`` <= 8 blocks of T <= 8 tokens, 6 query heads per KV head, and an HND cache + [pages, 2, kv_heads, 64, 64] bf16 whose pages are dense (any page stride).""" + if not (qkv.is_cuda and cache.is_cuda and qkv.dtype == cache.dtype == torch.bfloat16): + return False + if num_kv_heads < 1 or num_heads != HEADS_PER_KV * num_kv_heads: + return False + if not 0 < num_requests <= MAX_REQUESTS or qkv.dim() != 2 or not qkv.is_contiguous(): + return False + if qkv.shape[0] % num_requests or not 0 < qkv.shape[0] // num_requests <= MAX_TOKENS: + return False + if qkv.shape[1] != (num_heads + 2 * num_kv_heads) * HEAD_DIM: + return False + if cache.dim() != 5 or tuple(cache.shape[1:]) != (2, num_kv_heads, PAGE, HEAD_DIM): + return False + inner = (num_kv_heads * PAGE * HEAD_DIM, PAGE * HEAD_DIM, HEAD_DIM, 1) + return ( + tuple(cache.stride()[1:]) == inner + and cache.stride(0) % 8 == 0 + and cache.data_ptr() % 16 == 0 + ) + + +def _launch_attn( + qkv, cache, page_table, ctx_len, num_heads, num_kv_heads, out, norm=None +) -> torch.Tensor: + """``norm``: None (qkv already normalized and roped) or dict(q_w, k_w, positions, eps, base) (qkv raw).""" + import cuda.bindings.driver as cuda_driver + + num_requests = ctx_len.numel() + if not supports_attn(qkv, cache, num_heads, num_kv_heads, num_requests): + raise ValueError( + f"k3_drafter_attn: unsupported call qkv {tuple(qkv.shape)} {qkv.dtype} for {num_requests} requests, " + f"cache {tuple(cache.shape)} {tuple(cache.stride())} {cache.dtype}, heads {num_heads} / {num_kv_heads}" + ) + from . import k3_drafter_attn_kernel as kernel + + num_rows = qkv.shape[0] + num_tokens = num_rows // num_requests + if not ( + out.is_contiguous() + and out.dtype == torch.bfloat16 + and out.numel() == num_rows * num_heads * HEAD_DIM + ): + raise ValueError( + f"k3_drafter_attn: output {tuple(out.shape)} {out.dtype} is not a dense [M, heads * 64] bf16" + ) + if page_table.dtype != torch.int32 or ctx_len.dtype != torch.int32: + raise ValueError( + f"k3_drafter_attn: page table {page_table.dtype} / lengths {ctx_len.dtype} not int32" + ) + # Request r's pages start at element r * table_stride of the table (a view over the R rows, no copy). + if ( + page_table.dim() == 2 + and page_table.shape[0] >= num_requests + and (page_table.shape[1] == 1 or page_table.stride(1) == 1) + ): + table_stride = page_table.stride(0) + table = page_table.as_strided( + ((num_requests - 1) * table_stride + page_table.shape[1],), (1,) + ) + elif page_table.dim() == 1 and num_requests == 1: + table = page_table.reshape(-1) + table_stride = table.numel() + else: + raise ValueError( + f"k3_drafter_attn: page table {tuple(page_table.shape)} {tuple(page_table.stride())} is not " + f"[>= {num_requests}, width] with dense rows (or one dense row for one request)" + ) + norm_rope = norm is not None + q_w = k_w = _dummy(qkv.device, torch.bfloat16) + positions = _dummy(qkv.device, torch.int32) + norm_eps, rope_base = 0.0, 1.0 + if norm_rope: + q_w, k_w, positions = norm["q_w"], norm["k_w"], norm["positions"] + norm_eps, rope_base = float(norm["eps"]), float(norm["base"]) + if not ( + q_w.dtype == k_w.dtype == torch.bfloat16 + and q_w.numel() == k_w.numel() == HEAD_DIM + and q_w.is_contiguous() + and k_w.is_contiguous() + and positions.dtype in (torch.int32, torch.int64) + and positions.numel() >= num_rows + ): + raise ValueError( + "k3_drafter_attn: q/k norm weights must be bf16 [64], positions int32 / int64 [>= M]" + ) + # The cache's first page carries the pointer; the tensor map gets the page count and stride separately. + base = cache.as_strided((HEAD_DIM,), (1,)) + args = ( + _arg(qkv.view(-1)), + _arg(base), + _arg(table, align=4), # a view at any row of a block table; read one int32 at a time + _arg(ctx_len.reshape(-1)), + _arg(out.view(-1)), + _arg(q_w.reshape(-1)), + _arg(k_w.reshape(-1)), + _arg(positions.reshape(-1)), + ) + stream = cuda_driver.CUstream(torch.cuda.current_stream(qkv.device).cuda_stream) + use_pdl = _use_pdl() + page_stride = cache.stride(0) + total_pages = cache.shape[0] + scale_log2 = kernel.softmax_scale_log2(HEAD_DIM) + multi_request = num_requests > 1 + key = ( + "k3_drafter_attn", + num_heads, + num_kv_heads, + page_stride, + norm_rope, + positions.dtype, + multi_request, + use_pdl, + ) + fn = _compiled.get(key) + if fn is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "trtllm::k3_drafter_attn must run once outside CUDA-graph capture first (it compiles)." + ) + import cutlass.cute as cute + + with _lock: + fn = _compiled.get(key) + if fn is None: + fn = _compiled[key] = cute.compile( + kernel.k3_drafter_attn, *args, num_tokens, num_requests, table_stride, scale_log2, total_pages, + norm_eps, rope_base, num_heads, num_kv_heads, page_stride, norm_rope, multi_request, use_pdl, + stream, + ) # fmt: skip + fn(*args, num_tokens, num_requests, table_stride, scale_log2, total_pages, norm_eps, rope_base, + stream) # fmt: skip + return out + + +@torch.library.custom_op("trtllm::k3_drafter_attn", mutates_args=("out",)) +def k3_drafter_attn( + qkv: torch.Tensor, + cache: torch.Tensor, + page_table: torch.Tensor, + ctx_len: torch.Tensor, + num_heads: int, + num_kv_heads: int, + out: torch.Tensor, +) -> None: + """``out`` [M, heads * 64] = softmax(q K^T / 8) V per query head of each of the R = ``ctx_len.numel()`` requests, + dense over its ``ctx_len[r]`` cached rows of ``cache`` (pages: row r of ``page_table``) and its block's own + T = M / R rows of k / v in ``qkv`` (q, k after RMSNorm + RoPE). See the module docstring for the layouts.""" + _launch_attn(qkv, cache, page_table, ctx_len, num_heads, num_kv_heads, out) + + +@k3_drafter_attn.register_fake +def _(qkv, cache, page_table, ctx_len, num_heads, num_kv_heads, out): + return None + + +@torch.library.custom_op("trtllm::k3_drafter_attn_qknorm", mutates_args=("out",)) +def k3_drafter_attn_qknorm( + qkv: torch.Tensor, + q_w: torch.Tensor, + k_w: torch.Tensor, + positions: torch.Tensor, + eps: float, + rope_base: float, + cache: torch.Tensor, + page_table: torch.Tensor, + ctx_len: torch.Tensor, + num_heads: int, + num_kv_heads: int, + out: torch.Tensor, +) -> None: + """``k3_drafter_attn`` on the raw projection output: the per-head q / k RMSNorm (``q_w``, ``k_w``, ``eps``) and the + NeoX RoPE (``rope_base``, ``positions`` [M], one per row of ``qkv``) are applied in the kernel with + fused_qk_norm_rope's arithmetic.""" + norm = dict(q_w=q_w, k_w=k_w, positions=positions, eps=eps, base=rope_base) + _launch_attn(qkv, cache, page_table, ctx_len, num_heads, num_kv_heads, out, norm=norm) + + +@k3_drafter_attn_qknorm.register_fake +def _( + qkv, + q_w, + k_w, + positions, + eps, + rope_base, + cache, + page_table, + ctx_len, + num_heads, + num_kv_heads, + out, +): + return None + + +def _per_request(qkv: torch.Tensor, page_table: torch.Tensor, num_requests: int): + """(rows of ``qkv``, page-table row) of each request.""" + tables = page_table.unsqueeze(0) if page_table.dim() == 1 else page_table + t = qkv.shape[0] // num_requests + return [(qkv[r * t : (r + 1) * t], tables[r]) for r in range(num_requests)] + + +def reference(qkv: torch.Tensor, cache: torch.Tensor, page_table: torch.Tensor, ctx, num_heads: int, + num_kv_heads: int) -> torch.Tensor: # fmt: skip + """torch reference of ``k3_drafter_attn`` with host-known lengths ``ctx`` (an int for one request, else one per + request), [M, heads * 64] fp32 (computed in float64: no TF32 even where cuBLAS is told to use it).""" + ctxs = [ctx] if isinstance(ctx, int) else [int(c) for c in ctx] + return torch.cat([ + _reference_one(rows, cache, table, c, num_heads, num_kv_heads) + for (rows, table), c in zip(_per_request(qkv, page_table, len(ctxs)), ctxs) + ]).float() # fmt: skip + + +def _reference_one(qkv, cache, page_table, ctx, num_heads, num_kv_heads): + m = qkv.shape[0] + q = qkv[:, : num_heads * HEAD_DIM].double().view(m, num_heads, HEAD_DIM) + k_blk = ( + qkv[:, num_heads * HEAD_DIM : (num_heads + num_kv_heads) * HEAD_DIM] + .double() + .view(m, num_kv_heads, HEAD_DIM) + ) + v_blk = qkv[:, (num_heads + num_kv_heads) * HEAD_DIM :].double().view(m, num_kv_heads, HEAD_DIM) + n_pages = (ctx + PAGE - 1) // PAGE + pages = page_table[:n_pages].long() + k_ctx = ( + cache[pages, 0].double().permute(1, 0, 2, 3).reshape(num_kv_heads, -1, HEAD_DIM)[:, :ctx] + ) + v_ctx = ( + cache[pages, 1].double().permute(1, 0, 2, 3).reshape(num_kv_heads, -1, HEAD_DIM)[:, :ctx] + ) + k_all = torch.cat([k_ctx, k_blk.permute(1, 0, 2)], dim=1) # [kv, L, 64] + v_all = torch.cat([v_ctx, v_blk.permute(1, 0, 2)], dim=1) + group = num_heads // num_kv_heads + k_h = k_all.repeat_interleave(group, dim=0) # [heads, L, 64] + v_h = v_all.repeat_interleave(group, dim=0) + s = torch.einsum("thd,hld->thl", q, k_h) / (HEAD_DIM**0.5) + p = torch.softmax(s, dim=-1) + return torch.einsum("thl,hld->thd", p, v_h).reshape(m, num_heads * HEAD_DIM) + + +def reference_masked(qkv: torch.Tensor, cache: torch.Tensor, page_table: torch.Tensor, ctx_len: torch.Tensor, + num_heads: int, num_kv_heads: int) -> torch.Tensor: # fmt: skip + """torch reference of ``k3_drafter_attn`` with the lengths on the device (CUDA-graph safe): every page of a + request's page-table row is gathered and its rows >= ``ctx_len[r]`` are masked (their V zeroed), [M, heads * 64] + fp32 (computed in float64).""" + lens = ctx_len.reshape(-1) + return torch.cat([ + _reference_masked_one(rows, cache, table, lens[r : r + 1], num_heads, num_kv_heads) + for r, (rows, table) in enumerate(_per_request(qkv, page_table, lens.numel())) + ]).float() # fmt: skip + + +def _reference_masked_one(qkv, cache, page_table, ctx_len, num_heads, num_kv_heads): + m = qkv.shape[0] + pages = page_table.reshape(-1).long().clamp(min=0, max=cache.shape[0] - 1) + n_rows = pages.numel() * PAGE + k_ctx = cache[pages, 0].double().permute(1, 0, 2, 3).reshape(num_kv_heads, n_rows, HEAD_DIM) + v_ctx = cache[pages, 1].double().permute(1, 0, 2, 3).reshape(num_kv_heads, n_rows, HEAD_DIM) + valid_ctx = torch.arange(n_rows, device=qkv.device) < ctx_len.long() + k_blk = ( + qkv[:, num_heads * HEAD_DIM : (num_heads + num_kv_heads) * HEAD_DIM] + .double() + .view(m, num_kv_heads, HEAD_DIM) + ) + v_blk = qkv[:, (num_heads + num_kv_heads) * HEAD_DIM :].double().view(m, num_kv_heads, HEAD_DIM) + k_all = torch.cat([k_ctx, k_blk.permute(1, 0, 2)], dim=1) + v_all = torch.cat([v_ctx, v_blk.permute(1, 0, 2)], dim=1) + valid = torch.cat([valid_ctx, torch.ones(m, dtype=torch.bool, device=qkv.device)]) + v_all = torch.where(valid.view(1, -1, 1), v_all, torch.zeros_like(v_all)) + group = num_heads // num_kv_heads + q = qkv[:, : num_heads * HEAD_DIM].double().view(m, num_heads, HEAD_DIM) + s = torch.einsum("thd,hld->thl", q, k_all.repeat_interleave(group, dim=0)) / (HEAD_DIM**0.5) + s = s.masked_fill(~valid.view(1, 1, -1), float("-inf")) + p = torch.softmax(s, dim=-1) + return torch.einsum("thl,hld->thd", p, v_all.repeat_interleave(group, dim=0)).reshape( + m, num_heads * HEAD_DIM + ) diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index d5faa587114b..df09a270ce4a 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -245,6 +245,8 @@ l0_b200: # decode-step inputs (SM 100 only; they skip on every other list). - unittest/_torch/cute_dsl_kernels/test_spec_step_copies.py - unittest/_torch/executor/test_pytorch_model_engine.py -k "test_staged_spec_decode_graph_step" + # Kimi K3 DSpark drafter attention (CuTe DSL, SM 100). + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_drafter_attn.py - unittest/_torch/thop/parallel TIMEOUT (90) - unittest/_torch/visual_gen/kernels/parallel - unittest/_torch/thop/serial diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_drafter_attn.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_drafter_attn.py new file mode 100644 index 000000000000..0230d4da41ec --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_drafter_attn.py @@ -0,0 +1,423 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``trtllm::k3_drafter_attn`` / ``k3_drafter_attn_qknorm`` (the Kimi K3 DSpark drafter's block attention) for R +requests of T tokens, at the in-model TP16 group (6 query heads, 1 KV head; TP4 24 / 4 for a few splits), head dim 64, +HND pages of 64. + +Every split R x T (R T <= 8, and R x 8 for R <= 8) with mixed per-request context lengths (blocks crossing a page, a +128-row tile and the cluster's 16-tile round, no context, several tiles per CTA), page tables as strided row views +and as dense rows, against an fp32 reference and the model's production path (flashinfer append_paged_kv_cache + +trtllm-gen batch_context_with_kv_cache, non-causal), plus: no NaN (the pool's unused rows are NaN), the pool left +untouched, reruns bit-identical, requests isolated (a change in one request's context or block changes only its +rows), CUDA-graph replays with rewritten inputs. + +Error table: ``python3 test_k3_drafter_attn.py report``. +""" + +import sys + +import pytest +import torch + +D = 64 +PAGE = 64 +EPS = 1e-5 +THETA = 10000.0 +TOL = 1e-2 # max |err| / max |ref|; bf16 P and output roundings are ~4e-3 + +# Every split the engine schedules for the drafter: R requests x T tokens with R T <= 8, and DSpark's R x 8. +SPLITS = [(1, 8), (2, 4), (4, 2), (8, 1), (1, 1), (2, 1), (3, 1), (4, 1), (5, 1)] + [ + (r, 8) for r in range(2, 9) +] +# Context lengths, cycled over the requests: blocks crossing a page (60, 121), a 128-row tile (124, 127), the +# cluster's 16-tile round (2040, 2041); starting a page / tile (0, 64, 128, 1024, 2048); several tiles per CTA. +LENGTHS = [60, 2040, 5, 124, 1000, 0, 2041, 64, 3000, 127, 1024, 121, 128, 4000, 2048, 1500] + + +def _sm100() -> bool: + return torch.cuda.is_available() and torch.cuda.get_device_capability()[0] == 10 + + +pytestmark = pytest.mark.skipif(not _sm100(), reason="needs SM100 (tcgen05, TMA, clusters)") + + +def _op(): + import tensorrt_llm # noqa: F401 (registers the trtllm ops) + from tensorrt_llm._torch.cute_dsl_kernels.k3_drafter import op + + return op + + +def bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +class Case: + """R requests of T tokens: an HND pool [pages, 2, kv, 64, 64] (NaN except each request's context rows, optionally + with a page stride larger than a page), distinct random pages per request, page-table rows, qkv and lengths.""" + + def __init__( + self, gen, heads, kv, num_requests, tokens, lengths, page_pad=0, strided_table=True + ): + self.heads, self.kv, self.r, self.t = heads, kv, num_requests, tokens + self.lengths = list(lengths) + need = [(c + tokens + PAGE - 1) // PAGE for c in self.lengths] + self.width = max(need) + 2 + n_pool = sum(need) + 2 * num_requests + 8 + inner = 2 * kv * PAGE * D + flat = torch.full( + (n_pool * (inner + page_pad),), float("nan"), dtype=torch.bfloat16, device="cuda" + ) + self.pool = flat.as_strided( + (n_pool, 2, kv, PAGE, D), (inner + page_pad, kv * PAGE * D, PAGE * D, D, 1) + ) + perm = torch.randperm( + n_pool, generator=torch.Generator().manual_seed(int(sum(need)) + num_requests) + ) + spare = perm[sum(need) :] + rows, used = [], 0 + for r, n in enumerate(need): + own = perm[used : used + n] + used += n + # Entries past the request's pages: spare (NaN) pages, never read by the kernel. + rows.append( + torch.cat([own, spare[(torch.arange(self.width - n) + 2 * r) % spare.numel()]]) + ) + dense = torch.stack(rows).to(torch.int32).cuda() + if strided_table: + # Rows 1 .. R of a wider table: a row stride larger than the width and a nonzero offset. + full = torch.full( + (num_requests + 2, self.width + 3), -1, dtype=torch.int32, device="cuda" + ) + full[1 : num_requests + 1, : self.width] = dense + self.table = full[1 : num_requests + 1, : self.width] + else: + self.table = dense + for r, c in enumerate(self.lengths): + for i in range((c + PAGE - 1) // PAGE): + n_rows = min(PAGE, c - i * PAGE) + p = int(self.table[r, i]) + self.pool[p, :, :, :n_rows] = ( + torch.randn(2, kv, n_rows, D, generator=gen, device="cuda") * 0.5 + ).bfloat16() + m = num_requests * tokens + self.qkv = ( + torch.randn(m, (heads + 2 * kv) * D, generator=gen, device="cuda") * 0.5 + ).bfloat16() + self.ctx_len = torch.tensor(self.lengths, dtype=torch.int32, device="cuda") + self.positions = ( + self.ctx_len.view(-1, 1) + torch.arange(tokens, device="cuda", dtype=torch.int32) + ).reshape(-1) + + def run(self, qkv=None, pool=None, table=None): + out = torch.empty(self.r * self.t, self.heads * D, dtype=torch.bfloat16, device="cuda") + torch.ops.trtllm.k3_drafter_attn( + self.qkv if qkv is None else qkv, self.pool if pool is None else pool, + self.table if table is None else table, self.ctx_len, self.heads, self.kv, out, + ) # fmt: skip + return out + + def rows(self, r): + return slice(r * self.t, (r + 1) * self.t) + + +def production(case: Case) -> torch.Tensor: + """The model's DFlash TRTLLM path for a non-causal layer (modeling_dflash): append the blocks' K / V into the pool + (a copy, NaN rows zeroed: the context kernel reads the last page's rows past L), then the batched context kernel.""" + from tensorrt_llm._torch.speculative.dflash_attention import get_dflash_trtllm_gen_ops + + ops = get_dflash_trtllm_gen_ops() + h, kv, b, t = case.heads, case.kv, case.r, case.t + pool = torch.nan_to_num(case.pool.clone(), nan=0.0) + table = case.table.contiguous() + seq_after = case.ctx_len + t + ws = ops.get_workspace_size(dtype=torch.bfloat16, num_tokens=b * t, num_gen_tokens=b * t, num_heads=h, + num_kv_heads=kv, head_size=D, max_num_requests=b, rotary_embedding_dim=0, + fp8_context_fmha=False) # fmt: skip + workspace = torch.empty(ws, dtype=torch.uint8, device="cuda") + sm = torch.cuda.get_device_properties(0).multi_processor_count + counters = torch.zeros( + ops.get_multi_ctas_kv_counter_size(h, b, sm), dtype=torch.uint8, device="cuda" + ) + qkv = case.qkv + ops.append_paged_kv_cache( + append_key=qkv[:, h * D : (h + kv) * D].reshape(-1, kv, D), + append_value=qkv[:, (h + kv) * D :].reshape(-1, kv, D), + batch_indices=torch.arange(b, dtype=torch.int32, device="cuda").repeat_interleave(t), + positions=case.positions, paged_kv_cache=pool, kv_indices=table.flatten(), + kv_indptr=torch.arange(0, (b + 1) * case.width, case.width, dtype=torch.int32, device="cuda"), + kv_last_page_len=((seq_after - 1) % PAGE) + 1, kv_layout="HND", + ) # fmt: skip + out = torch.empty(b * t, h * D, dtype=torch.bfloat16, device="cuda") + ops.batch_context_with_kv_cache( + query=qkv[:, : h * D].reshape(-1, h, D), kv_cache=(pool[:, 0], pool[:, 1]), workspace_buffer=workspace, + block_tables=table, seq_lens=seq_after, max_q_len=t, max_kv_len=case.width * PAGE, bmm1_scale=D**-0.5, + bmm2_scale=1.0, batch_size=b, + cum_seq_lens_q=torch.arange(0, (b + 1) * t, t, dtype=torch.int32, device="cuda"), + cum_seq_lens_kv=torch.cat([torch.zeros(1, dtype=torch.int32, device="cuda"), + seq_after.cumsum(0, dtype=torch.int32)]), + window_left=-1, out=out.view(-1, h, D), sinks=None, enable_pdl=False, kv_layout="HND", kv_cache_sf=None, + uses_shared_paged_kv_idx=True, causal=False, multi_ctas_kv_counter_buffer=counters, + ) # fmt: skip + return out + + +def rel_err(a: torch.Tensor, ref: torch.Tensor) -> float: + return (a.float() - ref.float()).abs().max().item() / max(ref.float().abs().max().item(), 1e-6) + + +def case_lengths(num_requests: int, offset: int): + return [LENGTHS[(offset + 5 * r) % len(LENGTHS)] for r in range(num_requests)] + + +def measure_split(num_requests, tokens, offset, heads=6, kv=1): + """One split's checks: errors per request against the fp32 reference and the production path, NaN, pool + untouched, rerun, isolation, the device-length reference.""" + op = _op() + gen = torch.Generator(device="cuda").manual_seed( + 1000 * num_requests + 10 * tokens + offset + heads + ) + case = Case(gen, heads, kv, num_requests, tokens, case_lengths(num_requests, offset), + page_pad=0 if offset % 2 == 0 else 3 * 2 * kv * PAGE * D, strided_table=offset != 3) # fmt: skip + before = case.pool.clone() + out = case.run() + torch.cuda.synchronize() + res = dict(split=f"{num_requests}x{tokens}", heads=f"{heads}/{kv}", lengths=case.lengths, + table="strided" if offset != 3 else "dense", page_pad=offset % 2 == 1) # fmt: skip + res["untouched"] = torch.equal(bits(case.pool), bits(before)) + res["nan"] = bool(torch.isnan(out.float()).any()) + ref = op.reference(case.qkv, case.pool, case.table, case.lengths, heads, kv) + prod = production(case) + res["err"] = [rel_err(out[case.rows(r)], ref[case.rows(r)]) for r in range(num_requests)] + res["err_prod"] = [rel_err(prod[case.rows(r)], ref[case.rows(r)]) for r in range(num_requests)] + res["err_vs_prod"] = [ + rel_err(out[case.rows(r)], prod[case.rows(r)]) for r in range(num_requests) + ] + res["abs_err"] = (out.float() - ref).abs().max().item() + res["bits_vs_prod"] = int((bits(out) != bits(prod)).sum()) + # The device-length reference (what a CUDA-graph check would use) agrees with the host-length one. + masked = op.reference_masked( + case.qkv, torch.nan_to_num(case.pool, nan=0.0), case.table, case.ctx_len, heads, kv + ) + res["masked_ref"] = rel_err(masked, ref) + res["rerun"] = torch.equal(bits(case.run()), bits(out)) + # Isolation: request 0's context row and the last request's block V change only their own rows (a block's V + # always reaches its outputs; its K does not when the block is a request's only key). + iso = True + if case.lengths[0] > 0: + pool_c = case.pool.clone() + row = case.lengths[0] // 2 + pool_c[int(case.table[0, row // PAGE]), 0, 0, row % PAGE] += 1.0 + out_c = case.run(pool=pool_c) + iso &= not torch.equal(bits(out_c[case.rows(0)]), bits(out[case.rows(0)])) + iso &= torch.equal(bits(out_c[tokens:]), bits(out[tokens:])) + qkv_c = case.qkv.clone() + qkv_c[num_requests * tokens - 1, (heads + kv) * D + 5] += 1.0 + out_c = case.run(qkv=qkv_c) + last = case.rows(num_requests - 1) + iso &= not torch.equal(bits(out_c[last]), bits(out[last])) + iso &= torch.equal(bits(out_c[: last.start]), bits(out[: last.start])) + res["isolated"] = iso + res["ok"] = (res["untouched"] and not res["nan"] and max(res["err"]) <= TOL and max(res["err_vs_prod"]) <= TOL + and res["masked_ref"] <= 1e-5 and res["rerun"] and iso) # fmt: skip + return res + + +@pytest.mark.parametrize("num_requests,tokens", SPLITS, ids=[f"{r}x{t}" for r, t in SPLITS]) +@pytest.mark.parametrize("offset", [0, 3, 7]) +def test_split(num_requests, tokens, offset): + with torch.inference_mode(): + res = measure_split(num_requests, tokens, offset) + assert res["ok"], res + + +TP4_SPLITS = [(1, 8), (2, 4), (8, 1), (4, 8), (8, 8)] + + +@pytest.mark.parametrize("num_requests,tokens", TP4_SPLITS, ids=[f"{r}x{t}" for r, t in TP4_SPLITS]) +def test_split_tp4(num_requests, tokens): + """TP4's group (24 query heads, 4 KV heads: four clusters per request).""" + with torch.inference_mode(): + res = measure_split(num_requests, tokens, 1, heads=24, kv=4) + assert res["ok"], res + + +def qk_norm_rope(qkv, heads, kv, q_w, k_w, positions): + """The model's fused_qk_norm_rope, in place (plain NeoX RoPE).""" + torch.ops.trtllm.fused_qk_norm_rope( + qkv, + heads, + kv, + kv, + D, + D, + EPS, + q_w, + k_w, + THETA, + True, + positions, + 1.0, + 0.0, + 0.0, + 1.0, + True, + False, + False, + 0, + 0, + ) + return qkv + + +@pytest.mark.parametrize("num_requests,tokens", SPLITS, ids=[f"{r}x{t}" for r, t in SPLITS]) +def test_split_qknorm(num_requests, tokens): + """k3_drafter_attn_qknorm (q/k RMSNorm + RoPE in the kernel, raw qkv) against fused_qk_norm_rope + k3_drafter_attn + and the fp32 reference; its input untouched; int64 positions give the same bits.""" + op = _op() + gen = torch.Generator(device="cuda").manual_seed(5000 + 10 * num_requests + tokens) + with torch.inference_mode(): + case = Case(gen, 6, 1, num_requests, tokens, case_lengths(num_requests, 2)) + raw = (case.qkv.float() * 4.0).bfloat16() + q_w = (1.0 + 0.2 * torch.randn(D, generator=gen, device="cuda")).bfloat16() + k_w = (1.0 + 0.2 * torch.randn(D, generator=gen, device="cuda")).bfloat16() + normed = qk_norm_rope(raw.clone(), 6, 1, q_w, k_w, case.positions) + out_a = case.run(qkv=normed) + raw_before = raw.clone() + out_b = torch.empty_like(out_a) + torch.ops.trtllm.k3_drafter_attn_qknorm(raw, q_w, k_w, case.positions, EPS, THETA, case.pool, case.table, + case.ctx_len, 6, 1, out_b) # fmt: skip + out_c = torch.empty_like(out_a) + torch.ops.trtllm.k3_drafter_attn_qknorm(raw, q_w, k_w, case.positions.long(), EPS, THETA, case.pool, + case.table, case.ctx_len, 6, 1, out_c) # fmt: skip + assert torch.equal(bits(raw), bits(raw_before)), "the kernel wrote its input" + assert torch.equal(bits(out_c), bits(out_b)), "int64 positions differ from int32" + ref = op.reference(normed, case.pool, case.table, case.lengths, 6, 1) + for r in range(num_requests): + rows = case.rows(r) + assert rel_err(out_b[rows], ref[rows]) <= TOL, f"request {r}" + assert rel_err(out_b[rows], out_a[rows]) <= TOL, ( + f"request {r} vs norm/rope + k3_drafter_attn" + ) + + +@pytest.mark.parametrize("num_requests,tokens", [(4, 8), (8, 1), (2, 4)]) +def test_graph_replay(num_requests, tokens): + """One captured call of each op, replayed with the lengths, page tables and qkv rewritten in place.""" + op = _op() + gen = torch.Generator(device="cuda").manual_seed(31 + num_requests) + with torch.inference_mode(): + lengths = [3000] * num_requests + case = Case(gen, 6, 1, num_requests, tokens, lengths, strided_table=False) + q_w = (1.0 + 0.2 * torch.randn(D, generator=gen, device="cuda")).bfloat16() + k_w = (1.0 + 0.2 * torch.randn(D, generator=gen, device="cuda")).bfloat16() + out = torch.empty(num_requests * tokens, 6 * D, dtype=torch.bfloat16, device="cuda") + out_n = torch.empty_like(out) + + def body(): + torch.ops.trtllm.k3_drafter_attn( + case.qkv, case.pool, case.table, case.ctx_len, 6, 1, out + ) + torch.ops.trtllm.k3_drafter_attn_qknorm(case.qkv, q_w, k_w, case.positions, EPS, THETA, case.pool, + case.table, case.ctx_len, 6, 1, out_n) # fmt: skip + + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(stream): + body() + torch.cuda.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + body() + torch.cuda.synchronize() + for rep in range(6): + new = [min(c, 3000 - tokens) for c in case_lengths(num_requests, rep)] + case.ctx_len.copy_(torch.tensor(new, dtype=torch.int32)) + case.positions.copy_( + (case.ctx_len.view(-1, 1) + torch.arange(tokens, device="cuda")).reshape(-1) + ) + # Rotate the page-table rows (each request reads another's pages) and redraw qkv. + case.table.copy_(torch.roll(case.table, shifts=rep + 1, dims=0)) + case.qkv.copy_( + (torch.randn(case.qkv.shape, generator=gen, device="cuda") * 0.5).bfloat16() + ) + graph.replay() + torch.cuda.synchronize() + want = torch.empty_like(out) + torch.ops.trtllm.k3_drafter_attn( + case.qkv, case.pool, case.table, case.ctx_len, 6, 1, want + ) + want_n = torch.empty_like(out) + torch.ops.trtllm.k3_drafter_attn_qknorm(case.qkv, q_w, k_w, case.positions, EPS, THETA, case.pool, + case.table, case.ctx_len, 6, 1, want_n) # fmt: skip + assert torch.equal(bits(out), bits(want)), f"replay {rep}" + assert torch.equal(bits(out_n), bits(want_n)), f"replay {rep} (qknorm)" + # Every request has 3000 rows of context in its first pages, so any length <= 3000 reads finite rows. + ref = op.reference(case.qkv, case.pool, case.table, new, 6, 1) + assert rel_err(out, ref) <= TOL, f"replay {rep}" + + +def test_unsupported(): + op = _op() + with torch.inference_mode(): + case = Case(torch.Generator(device="cuda").manual_seed(3), 6, 1, 2, 8, [100, 200]) + assert op.supports_attn(case.qkv, case.pool, 6, 1, 2) + assert op.supports_attn(case.qkv[:8], case.pool, 6, 1) + assert not op.supports_attn(case.qkv, case.pool, 6, 1) # 16 rows of one request + assert not op.supports_attn(case.qkv, case.pool, 6, 1, 3) # 16 rows over 3 requests + assert not op.supports_attn(case.qkv.repeat(5, 1), case.pool, 6, 1, 10) # 10 requests + out = torch.empty(16, 6 * D, dtype=torch.bfloat16, device="cuda") + with pytest.raises(ValueError): # one page-table row for two requests + torch.ops.trtllm.k3_drafter_attn( + case.qkv, case.pool, case.table[0], case.ctx_len, 6, 1, out + ) + with pytest.raises(ValueError): # columns not dense + torch.ops.trtllm.k3_drafter_attn(case.qkv, case.pool, case.table.t().contiguous().t(), case.ctx_len, 6, + 1, out) # fmt: skip + + +def report() -> int: + """The split checks as a markdown table (max over requests and the length offsets).""" + print(f"{torch.cuda.get_device_name()}") + print( + "| heads/kv | split | lengths (offset 0) | CTM vs fp32 | prod vs fp32 | CTM vs prod | max abs | bits != prod " + "| NaN | pool untouched | rerun | isolated | result |" + ) + print("| :-- | :-- | :-- | --: | --: | --: | --: | --: | :-- | :-- | :-- | :-- | :-- |") + ok_all = True + with torch.inference_mode(): + for heads, kv, splits in ((6, 1, SPLITS), (24, 4, TP4_SPLITS)): + for r, t in splits: + rows = [ + measure_split(r, t, off, heads, kv) + for off in ((0, 3, 7) if heads == 6 else (1,)) + ] + ok = all(x["ok"] for x in rows) + ok_all &= ok + print(f"| {heads}/{kv} | {r}x{t} | {rows[0]['lengths']} | {max(max(x['err']) for x in rows):.2e} | " + f"{max(max(x['err_prod']) for x in rows):.2e} | {max(max(x['err_vs_prod']) for x in rows):.2e} | " + f"{max(x['abs_err'] for x in rows):.2e} | {sum(x['bits_vs_prod'] for x in rows)}/" + f"{len(rows) * r * t * heads * D} | {any(x['nan'] for x in rows)} | " + f"{all(x['untouched'] for x in rows)} | {all(x['rerun'] for x in rows)} | " + f"{all(x['isolated'] for x in rows)} | {'PASS' if ok else 'FAIL'} |", flush=True) # fmt: skip + print("ALL PASS" if ok_all else "FAIL") + return 0 if ok_all else 1 + + +if __name__ == "__main__": + if len(sys.argv) > 1 and sys.argv[1] == "report": + sys.exit(report()) + else: + sys.exit(pytest.main([__file__, "-q", "-p", "no:cacheprovider", *sys.argv[1:]])) From a6a8cb73a306656af764ed532e985bd576c21473 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 18:13:04 -0700 Subject: [PATCH 078/161] [None][feat] modeling_v2 catalog: attention/k3_drafter_attn, k3_drafter_attn_qknorm; sm_100 cells for fused_qk_norm_rope - attention/k3_drafter_attn and attention/k3_drafter_attn_qknorm: contracts, wrappers (argument for argument with the op schemas) and GPU tests on sm_100. The tests run against fp64 references written in the test: - both head layouts Kimi K3 runs (6 / 1, 24 / 4); - eight request x token splits; - context lengths 0 to 2041; - page tables as strided or dense rows, padded page strides; - int32 and int64 positions; - CUDA-graph replay with rewritten inputs; - the out-of-contract calls that must raise. The tests skip off SM 10.0. - attention/fused_qk_norm_rope gains the drafter's cell (head_dim 64, 6 query heads per KV head, 1-64 tokens, eps 1e-5, base 10000, NeoX), which skips off SM 10.0, and an sm_100 receipt for the whole test file. The sm_103 receipt stays, noted as pre-dating that cell. - index.yaml and README record the three sm_100 entries and the rule for reused entries. - The three test files run in l0_b200.yml. Signed-off-by: Vasanth Sabavat --- .../_experimental/modeling_v2/README.md | 5 +- .../catalog/attention/fused_qk_norm_rope.md | 5 + .../catalog/attention/k3_drafter_attn.md | 81 ++++++ .../catalog/attention/k3_drafter_attn.py | 22 ++ .../attention/k3_drafter_attn_qknorm.md | 74 ++++++ .../attention/k3_drafter_attn_qknorm.py | 40 +++ .../modeling_v2/catalog/index.yaml | 21 +- .../test_lists/test-db/l0_b200.yml | 6 +- .../test_modeling_v2_fused_qk_norm_rope.py | 16 ++ .../test_modeling_v2_k3_drafter_attn.py | 234 ++++++++++++++++++ ...test_modeling_v2_k3_drafter_attn_qknorm.py | 158 ++++++++++++ 11 files changed, 654 insertions(+), 8 deletions(-) create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn_qknorm.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn_qknorm.py create mode 100644 tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn.py create mode 100644 tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn_qknorm.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/README.md b/tensorrt_llm/_torch/_experimental/modeling_v2/README.md index cf0738128514..e7ca121547e4 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/README.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/README.md @@ -187,8 +187,9 @@ Perf is measured, never gated. ## Status of every record in this tree **The catalog is certified: 19 entries on sm_103, and -`comm/mnnvl_allreduce_attn_res` on sm_100 (GB200), where its first caller -runs. The targets construct but have never executed.** +`comm/mnnvl_allreduce_attn_res`, `attention/k3_drafter_attn` and +`attention/k3_drafter_attn_qknorm` on sm_100 (GB200), where their first +caller runs. The targets construct but have never executed.** Two things voided every receipt in the move: each catalog test file was rewritten, and the targets moved from sm_100 (B200) to sm_103 (GB300), where diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/fused_qk_norm_rope.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/fused_qk_norm_rope.md index 33a26545df88..01a11bb23d82 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/fused_qk_norm_rope.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/fused_qk_norm_rope.md @@ -1,6 +1,7 @@ --- receipts: sm_103: {status: passed, tests: 9} + sm_100: {status: passed, tests: 10} --- # fused_qk_norm_rope @@ -156,3 +157,7 @@ purely per-token). - TRT-LLM's own caller derives `factor/low/high/attention_factor` from a YaRN config and uses `rotary_dim = head_dim * partial_rotary_factor`; this entry exposes them raw. +- Kimi K3's DSpark drafter calls it with `head_dim` 64, 6 query heads per KV head (6 / 1 at TP16, 24 / 4 at TP4), + 1-64 tokens, `eps` 1e-5, `base` 10000, NeoX: `test_bf16_neox_kimi_k3_drafter`, which runs on sm_100 only. The + sm_100 receipt covers the whole file (10 tests). The sm_103 receipt pre-dates that cell and covers the other 9; + CI's `l0_b300` run re-certifies them, since the K3 cell skips off SM 10.0. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn.md new file mode 100644 index 000000000000..8f816d76aca1 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn.md @@ -0,0 +1,81 @@ +--- +receipts: + sm_100: {status: passed, tests: 20} +--- + +# k3_drafter_attn + +**Wraps** `torch.ops.trtllm.k3_drafter_attn` (one call), a CuTe DSL kernel +(`tensorrt_llm/_torch/cute_dsl_kernels/k3_drafter/`). + +## Semantics + +The block attention of Kimi K3's DSpark drafter. `R = ctx_len.numel()` requests each bring a draft block of +`T = M / R` tokens; the `M` rows of `qkv` are the requests' blocks in order. Per request `r`, row `i` of its block +and query head `h` (groups of `num_heads / num_kv_heads` = 6 query heads per KV head `g(h)`), head dim 64: + +``` +K_r, V_r = [rows 0 .. ctx_len[r] - 1 of request r's cache pages] ++ [the block's own k, v rows of qkv] +out[r T + i, h] = softmax(q[r T + i, h] . K_r[g(h)]^T / 8) V_r[g(h)] +``` + +Within a block every row attends to every row (non-causal). The block's own `k` / `v` are read from `qkv`, not +from the cache, and are not written into it. `out` is written; `qkv`, `cache`, `page_table` and `ctx_len` are read +only (the cache certified bit for bit). Each request's rows depend only on its own context and block (the op's own +test checks this). + +Fusion boundary. Inside: the paged gather of each request's context `K` / `V`, the block's own `K` / `V`, the +softmax over both, and the product with `V`. Outside: storing the block's `K` / `V` in the cache, if the caller keeps +them, and the q / k norm and RoPE (`attention/k3_drafter_attn_qknorm` fuses those in). + +## Signature + +```python +def k3_drafter_attn( + qkv: torch.Tensor, + cache: torch.Tensor, + page_table: torch.Tensor, + ctx_len: torch.Tensor, + num_heads: int, + num_kv_heads: int, + out: torch.Tensor, +) -> None +``` + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `qkv` | `[M = R T, (num_heads + 2 num_kv_heads) * 64]`: q heads, then k heads, then v heads; `R` <= 8, `T` <= 8 | bf16 | contiguous | CUDA | +| `cache` | `[pages, 2, num_kv_heads, 64, 64]` (HND pages of 64 rows, `K` then `V`) | bf16 | each page dense; the page stride may exceed a page (a multiple of 8 elements) | CUDA | +| `page_table` | `[>= R, width]`, row `r` = request `r`'s pages, dense rows at any row stride (e.g. a view of a block table); or `[width]` when `R` = 1 | int32 | — | CUDA | +| `ctx_len` | `[R]`, cached rows per request, 0 allowed | int32 | — | CUDA | +| `num_heads`, `num_kv_heads` | 6 / 1 (Kimi K3 at TP16) and 24 / 4 (TP4) certified; `num_heads` = 6 `num_kv_heads` | Python int | — | — | +| `out` | `[M, num_heads * 64]` | bf16 | contiguous | CUDA | + +Certified splits `R x T`: 1x1, 1x8, 2x4, 4x2, 8x1, 3x1, 2x8, 8x8, with context lengths 0 to 2041 (blocks crossing a +page, 64- and 128-row boundaries, 2041 rows spanning the kernel's 16-tile round), at both head layouts. + +## Metadata consumed + +None. The op compiles on its first call for each (`num_heads`, `num_kv_heads`, page stride, `R` > 1, PDL) and keeps +the result in a process-wide cache; that first call must happen outside CUDA-graph capture (it raises +`RuntimeError` there). The lengths, page tables and `qkv` are read on the device, so a captured call replays with +them rewritten in place (certified: 4 replays with new lengths, rotated page-table rows and new `qkv`). + +## Preconditions + +- SM100 (tcgen05, TMA, clusters). +- The shapes and dtypes above; a call outside them raises `ValueError` before any launch (certified: 12 query heads + per KV head, one page-table row for two requests, page-table columns not dense, int64 lengths, an fp32 output, + `T` = 16). +- Page-table entries below `ceil(ctx_len[r] / 64)` name pages holding request `r`'s rows; later entries and the + rows of the last page past `ctx_len[r]` are not read (certified with those rows NaN: the output has no NaN). + +## Notes + +- Certified on GB200 (sm_100) against an fp64 torch reference that gathers the pages, within 1e-2 of the largest + output magnitude (the kernel rounds the probabilities and the output to bf16). Reruns are bit-identical. +- The op's own test (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_drafter_attn.py`) also compares every + split against the DFlash TRTLLM path it replaces (`append_paged_kv_cache` + the trtllm-gen context kernel, + non-causal) and checks request isolation. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn.py new file mode 100644 index 000000000000..fd152cd11001 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn.py @@ -0,0 +1,22 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3 DSpark drafter block attention: each request's draft block attends densely to its paged context and to the +block's own K / V.""" + +import torch + +import tensorrt_llm._torch.cute_dsl_kernels.k3_drafter.op # noqa: F401 — registers torch.ops.trtllm.k3_drafter_attn + + +def k3_drafter_attn( + qkv: torch.Tensor, + cache: torch.Tensor, + page_table: torch.Tensor, + ctx_len: torch.Tensor, + num_heads: int, + num_kv_heads: int, + out: torch.Tensor, +) -> None: + """Write ``out`` [R T, num_heads * 64]: request r's T rows of ``qkv`` attend to its ``ctx_len[r]`` cached rows + (pages: row r of ``page_table``) and to its block's own k / v in ``qkv``. Returns None.""" + torch.ops.trtllm.k3_drafter_attn(qkv, cache, page_table, ctx_len, num_heads, num_kv_heads, out) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn_qknorm.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn_qknorm.md new file mode 100644 index 000000000000..45b60ecba263 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn_qknorm.md @@ -0,0 +1,74 @@ +--- +receipts: + sm_100: {status: passed, tests: 8} +--- + +# k3_drafter_attn_qknorm + +**Wraps** `torch.ops.trtllm.k3_drafter_attn_qknorm` (one call), the CuTe DSL kernel of `attention/k3_drafter_attn` +built with the q / k normalization in front. + +## Semantics + +`attention/k3_drafter_attn` on the raw output of the drafter's QKV projection. Before the attention, the kernel +applies to every q head and every k head of `qkv` (the block's own `k`, not the cached rows): + +``` +x = x / sqrt(mean(x^2) + eps) * w # w = q_w for q heads, k_w for k heads; per head over its 64 dims +x = NeoX RoPE(x, positions[row], rope_base) # rotary dim 64: the halves (x[:32], x[32:]) rotate together +``` + +with `attention/fused_qk_norm_rope`'s arithmetic (`is_neox = True`, full rotary dim, no YaRN), then attends exactly +as `k3_drafter_attn`. The normalized values exist only inside the kernel: `qkv` is read only (certified bit for +bit) and `out` is written. + +Fusion boundary. Inside: the q / k RMSNorm, the RoPE and the whole of `k3_drafter_attn`. Outside: the projection +that produced `qkv`, and storing the roped `k` / `v` in the cache, if the caller keeps them. + +## Signature + +```python +def k3_drafter_attn_qknorm( + qkv: torch.Tensor, + q_w: torch.Tensor, + k_w: torch.Tensor, + positions: torch.Tensor, + eps: float, + rope_base: float, + cache: torch.Tensor, + page_table: torch.Tensor, + ctx_len: torch.Tensor, + num_heads: int, + num_kv_heads: int, + out: torch.Tensor, +) -> None +``` + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `qkv` | as `k3_drafter_attn`, unnormalized | bf16 | contiguous | CUDA | +| `q_w`, `k_w` | `[64]` | bf16 | contiguous | CUDA | +| `positions` | `[>= M]`, one per row of `qkv` | int32 or int64 (bit-identical results, certified) | — | CUDA | +| `eps`, `rope_base` | 1e-5 and 10000.0 certified (Kimi K3's drafter) | Python float | — | — | +| `cache`, `page_table`, `ctx_len`, `out` | as `k3_drafter_attn` | | | | +| `num_heads`, `num_kv_heads` | 6 / 1 certified | Python int | — | — | + +Certified splits `R x T`: 1x1, 1x8, 2x4, 4x2, 8x1, 3x1, 2x8, 8x8, context lengths 0 to 2041. + +## Metadata consumed + +None; the compile cache and its capture rule are `k3_drafter_attn`'s (the normalization is part of the compile +key). A captured call replays with `positions` rewritten in place (certified with `k3_drafter_attn`'s replays). + +## Preconditions + +`k3_drafter_attn`'s, plus: `q_w` / `k_w` bf16 `[64]` contiguous and `positions` int32 / int64 with at least `M` +entries, else `ValueError`. + +## Notes + +- Certified on GB200 (sm_100) against fp64 RMSNorm + NeoX RoPE rounded to bf16 (as `fused_qk_norm_rope` stores + them), then the fp64 attention, within 1e-2 of the largest output magnitude. +- The op's own test also checks it against `fused_qk_norm_rope` followed by `k3_drafter_attn`. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn_qknorm.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn_qknorm.py new file mode 100644 index 000000000000..97bb4911b68b --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn_qknorm.py @@ -0,0 +1,40 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3 DSpark drafter block attention on the raw projection output: per-head q / k RMSNorm and NeoX RoPE inside +the kernel, then k3_drafter_attn.""" + +import torch + +import tensorrt_llm._torch.cute_dsl_kernels.k3_drafter.op # noqa: F401 — registers torch.ops.trtllm.k3_drafter_attn_qknorm + + +def k3_drafter_attn_qknorm( + qkv: torch.Tensor, + q_w: torch.Tensor, + k_w: torch.Tensor, + positions: torch.Tensor, + eps: float, + rope_base: float, + cache: torch.Tensor, + page_table: torch.Tensor, + ctx_len: torch.Tensor, + num_heads: int, + num_kv_heads: int, + out: torch.Tensor, +) -> None: + """``k3_drafter_attn`` after RMSNorm (``q_w`` / ``k_w``, ``eps``) and NeoX RoPE (``rope_base``, one position per + row of ``qkv``) of every q and k head, computed in the kernel; ``qkv`` is not modified. Returns None.""" + torch.ops.trtllm.k3_drafter_attn_qknorm( + qkv, + q_w, + k_w, + positions, + eps, + rope_base, + cache, + page_table, + ctx_len, + num_heads, + num_kv_heads, + out, + ) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml index 456e86438283..20880b0547a4 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml @@ -28,7 +28,7 @@ # limits, e.g. no float8 arithmetic, are noted in the entry docstring). # Entries wrapping trtllm ops carry all three. # -# ── RECEIPT STATUS: 19 entries certified on sm_103, 1 on sm_100 ─────────── +# ── RECEIPT STATUS: 19 entries certified on sm_103, 3 on sm_100 ─────────── # # A receipt says this entry's test passed on a stated GPU architecture. The # architecture is the whole key: it is a real axis -- see the @@ -68,10 +68,13 @@ # exist. Read them as provenance; the frontmatter is the certification. # # An entry first called by a GB200 (sm_100) target is certified there: its -# receipt key is sm_100 (comm/mnnvl_allreduce_attn_res). An entry already -# certified elsewhere gains its sm_100 key in the change that adds its first -# sm_100 caller. A collective's receipt records the world size its matrix ran -# at. +# receipt key is sm_100 (comm/mnnvl_allreduce_attn_res, +# attention/k3_drafter_attn, attention/k3_drafter_attn_qknorm). An entry +# already certified elsewhere gains its sm_100 key in the change that adds its +# first sm_100 caller and keeps the key it had, noting that the new cells +# post-date it (attention/fused_qk_norm_rope). Those cells skip off SM 10.0, +# so CI on the other architecture keeps re-certifying the existing ones. A +# collective's receipt records the world size its matrix ran at. # ────────────────────────────────────────────────────────────────────────── # # STATEFUL ENTRIES. When an op's result depends on state that outlives the @@ -192,6 +195,14 @@ entries: impl: tensorrt_llm.bindings.internal.thop.attention summary: "Full attention core with fully explicit state (pybind binding, approved exception): paged KV-cache append + causal/padding-masked GQA FMHA over a caller-owned pool (bf16, or fp8-e4m3 with per-tensor kv scales) and explicit length/offset tensors, written into a caller buffer, on either context execution path (use_paged_context_fmha selects packed-QKV context FMHA, or paged-KV context FMHA so a context call may run over a cached prefix — KV-cache reuse and chunked prefill), with optional per-query-head attention sinks (one extra softmax-denominator logit, dropped from the output) and an optional per-call sliding window (attention_window_size keys ending at the query's absolute position; a pure mask — the append stays at absolute positions, so cyclic pool reuse is the caller's page mapping); MLA mode runs context prefill (in-kernel RoPE + latent append), no-append context over explicit K/V (latent_cache=None: cached-KV prefixes and chunked partial passes with softmax-stats output), and generation latent-MQA decode as separate calls over a paged latent pool (bf16 or fp8-e4m3)" + - path: attention/k3_drafter_attn.py + impl: torch.ops.trtllm.k3_drafter_attn + summary: "Kimi K3 DSpark drafter block attention (CuTe DSL, sm_100): R <= 8 requests' draft blocks of T <= 8 tokens, each attending densely (non-causal within the block) to its request's paged context (HND pages of 64, rows 0 .. ctx_len[r] - 1, page-table rows at any row stride) and to the block's own K / V taken from qkv, not written to the cache; GQA with 6 query heads per KV head, head dim 64, bf16; compiles on its first call outside capture, then replays under CUDA graphs with lengths, page tables and qkv rewritten in place" + + - path: attention/k3_drafter_attn_qknorm.py + impl: torch.ops.trtllm.k3_drafter_attn_qknorm + summary: "k3_drafter_attn on the raw QKV projection: per-head q / k RMSNorm and NeoX RoPE (fused_qk_norm_rope's arithmetic, positions int32 or int64) applied inside the kernel, qkv left unmodified" + # ─── moe ─────────────────────────────────────────────────────── - path: moe/noaux_tc_op.py impl: torch.ops.trtllm.noaux_tc_op diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index df09a270ce4a..86bf36671f68 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -245,8 +245,12 @@ l0_b200: # decode-step inputs (SM 100 only; they skip on every other list). - unittest/_torch/cute_dsl_kernels/test_spec_step_copies.py - unittest/_torch/executor/test_pytorch_model_engine.py -k "test_staged_spec_decode_graph_step" - # Kimi K3 DSpark drafter attention (CuTe DSL, SM 100). + # Kimi K3 DSpark drafter attention (CuTe DSL, SM 100) and its modeling_v2 catalog entries, plus the sm_100 cells of + # fused_qk_norm_rope, which the drafter calls. - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_drafter_attn.py + - unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn.py + - unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn_qknorm.py + - unittest/_torch/modeling_v2/attention/test_modeling_v2_fused_qk_norm_rope.py - unittest/_torch/thop/parallel TIMEOUT (90) - unittest/_torch/visual_gen/kernels/parallel - unittest/_torch/thop/serial diff --git a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_fused_qk_norm_rope.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_fused_qk_norm_rope.py index b71a64912ef6..5f3402ae200e 100644 --- a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_fused_qk_norm_rope.py +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_fused_qk_norm_rope.py @@ -2,6 +2,7 @@ # SPDX-License-Identifier: Apache-2.0 """GPU test for the fused_qk_norm_rope catalog entry.""" +import pytest import torch from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.fused_qk_norm_rope import ( @@ -151,6 +152,21 @@ def test_bf16_neox_head_dims() -> None: ) # fmt: skip +@pytest.mark.skipif( + torch.cuda.get_device_capability() != (10, 0), reason="the Kimi K3 cell is certified on sm_100" +) +def test_bf16_neox_kimi_k3_drafter() -> None: + torch.manual_seed(9) + # Kimi K3's DSpark drafter: head_dim 64, 6 query heads per KV head (TP16: 6 / 1, TP4: 24 / 4), draft blocks of + # up to 8 requests x 8 tokens + for num_heads_q, num_heads_kv in ((6, 1), (24, 4)): + for num_tokens in (1, 8, 64): + _check( + num_tokens, num_heads_q, num_heads_kv, head_dim=64, rotary_dim=64, eps=1e-5, + base=10000.0, is_neox=True, + ) # fmt: skip + + def test_bf16_neox_prefill() -> None: torch.manual_seed(1) # prefill-like: many tokens, Qwen3-32B-like head layout diff --git a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn.py new file mode 100644 index 000000000000..f7180da1bee5 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn.py @@ -0,0 +1,234 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the k3_drafter_attn catalog entry (sm_100); its graph-replay cell also replays +k3_drafter_attn_qknorm.""" + +import pytest +import torch + +assert torch.cuda.is_available(), "k3_drafter_attn requires a CUDA device" + +if torch.cuda.get_device_capability() != (10, 0): + # The CuTe DSL kernel uses tcgen05, TMA and clusters; it is certified on sm_100 (B200 / GB200) only. + pytest.skip("k3_drafter_attn is certified on sm_100 only", allow_module_level=True) + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.k3_drafter_attn import ( # noqa: E402 + k3_drafter_attn, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.k3_drafter_attn_qknorm import ( # noqa: E402 + k3_drafter_attn_qknorm, +) + +D = 64 +PAGE = 64 +EPS = 1e-5 +THETA = 10000.0 +# max |err| / max |ref|: the kernel rounds P and the output to bf16 (~4e-3 each). +TOL = 1e-2 +# Context lengths, cycled over the requests: none, a block crossing a page, page and 128-row tile boundaries, the +# cluster's 16-tile round, several tiles per CTA. +LENGTHS = (60, 2041, 0, 127, 1000, 64, 128, 5) +SPLITS = ((1, 1), (1, 8), (2, 4), (4, 2), (8, 1), (3, 1), (2, 8), (8, 8)) + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def _rel_err(a: torch.Tensor, ref: torch.Tensor) -> float: + return (a.double() - ref).abs().max().item() / max(ref.abs().max().item(), 1e-6) + + +class _Case: + """R requests of T tokens: an HND pool [pages, 2, kv, 64, 64] that is NaN outside the requests' context rows + (optionally with a page stride larger than a page), distinct random pages per request, page-table rows (a strided + view into a wider table, or dense), qkv and lengths.""" + + def __init__(self, seed, heads, kv, num_requests, tokens, page_pad=0, strided_table=True): + gen = torch.Generator(device="cuda").manual_seed(seed) + self.heads, self.kv, self.r, self.t = heads, kv, num_requests, tokens + self.lengths = [LENGTHS[(seed + 3 * r) % len(LENGTHS)] for r in range(num_requests)] + need = [(c + PAGE - 1) // PAGE for c in self.lengths] + self.width = max(need) + 2 + n_pool = sum(need) + 2 * num_requests + 4 + inner = 2 * kv * PAGE * D + flat = torch.full( + (n_pool * (inner + page_pad),), float("nan"), dtype=torch.bfloat16, device="cuda" + ) + self.pool = flat.as_strided( + (n_pool, 2, kv, PAGE, D), (inner + page_pad, kv * PAGE * D, PAGE * D, D, 1) + ) + perm = torch.randperm(n_pool, generator=torch.Generator().manual_seed(seed)) + rows, used = [], 0 + for n in need: + rows.append(torch.cat([perm[used : used + n], perm[-(self.width - n) :]])) + used += n + dense = torch.stack(rows).to(torch.int32).cuda() + if strided_table: + full = torch.full( + (num_requests + 2, self.width + 3), -1, dtype=torch.int32, device="cuda" + ) + full[1 : num_requests + 1, : self.width] = dense + self.table = full[1 : num_requests + 1, : self.width] + else: + self.table = dense + for r, c in enumerate(self.lengths): + for i in range((c + PAGE - 1) // PAGE): + n_rows = min(PAGE, c - i * PAGE) + p = int(self.table[r, i]) + self.pool[p, :, :, :n_rows] = ( + torch.randn(2, kv, n_rows, D, generator=gen, device="cuda") * 0.5 + ).bfloat16() + m = num_requests * tokens + self.qkv = ( + torch.randn(m, (heads + 2 * kv) * D, generator=gen, device="cuda") * 0.5 + ).bfloat16() + self.ctx_len = torch.tensor(self.lengths, dtype=torch.int32, device="cuda") + self.positions = ( + self.ctx_len.view(-1, 1) + torch.arange(tokens, device="cuda", dtype=torch.int32) + ).reshape(-1) + + def out(self) -> torch.Tensor: + return torch.empty(self.r * self.t, self.heads * D, dtype=torch.bfloat16, device="cuda") + + +def _attention_ref(case: _Case, qkv: torch.Tensor) -> torch.Tensor: + """fp64: each request's q against its cached rows and its block's own k / v, softmax(q k^T / 8) v.""" + h, kv, t = case.heads, case.kv, case.t + outs = [] + for r, c in enumerate(case.lengths): + rows = qkv[r * t : (r + 1) * t].double() + q = rows[:, : h * D].view(t, h, D) + k_blk = rows[:, h * D : (h + kv) * D].view(t, kv, D).permute(1, 0, 2) + v_blk = rows[:, (h + kv) * D :].view(t, kv, D).permute(1, 0, 2) + pages = case.table[r, : (c + PAGE - 1) // PAGE].long() + k_ctx = case.pool[pages, 0].double().permute(1, 0, 2, 3).reshape(kv, -1, D)[:, :c] + v_ctx = case.pool[pages, 1].double().permute(1, 0, 2, 3).reshape(kv, -1, D)[:, :c] + k = torch.cat([k_ctx, k_blk], 1).repeat_interleave(h // kv, 0) + v = torch.cat([v_ctx, v_blk], 1).repeat_interleave(h // kv, 0) + p = torch.softmax(torch.einsum("thd,hld->thl", q, k) / D**0.5, dim=-1) + outs.append(torch.einsum("thl,hld->thd", p, v).reshape(t, h * D)) + return torch.cat(outs) + + +def _qk_norm_rope_ref(case: _Case, qkv: torch.Tensor, q_w, k_w) -> torch.Tensor: + """Per-head RMSNorm of the q and k heads, then NeoX RoPE (base THETA) at each row's position, in fp64 and rounded + to bf16 as fused_qk_norm_rope stores them; v as is.""" + h, kv = case.heads, case.kv + x = qkv.double().view(qkv.shape[0], h + 2 * kv, D).clone() + inv_freq = 1.0 / THETA ** (torch.arange(0, D, 2, dtype=torch.float64, device="cuda") / D) + angle = case.positions.double()[:, None] * inv_freq + cos, sin = angle.cos()[:, None, :], angle.sin()[:, None, :] + for start, count, w in ((0, h, q_w), (h, kv, k_w)): + y = x[:, start : start + count] + y = y * torch.rsqrt(y.pow(2).mean(-1, keepdim=True) + EPS) * w.double() + y1, y2 = y[..., : D // 2], y[..., D // 2 :] + x[:, start : start + count] = torch.cat([y1 * cos - y2 * sin, y1 * sin + y2 * cos], -1) + return x.view(qkv.shape[0], -1).bfloat16() + + +def _check_attn(case: _Case) -> None: + before = case.pool.clone() + out = case.out() + k3_drafter_attn(case.qkv, case.pool, case.table, case.ctx_len, case.heads, case.kv, out) + torch.cuda.synchronize() + assert not torch.isnan(out.float()).any(), "NaN in the output (the pool's unused rows are NaN)" + assert torch.equal(_bits(case.pool), _bits(before)), "the pool was written" + err = _rel_err(out, _attention_ref(case, case.qkv)) + assert err <= TOL, f"rel err {err:.3e}" + again = case.out() + k3_drafter_attn(case.qkv, case.pool, case.table, case.ctx_len, case.heads, case.kv, again) + assert torch.equal(_bits(again), _bits(out)), "a rerun differs" + + +@pytest.mark.parametrize("heads,kv", [(6, 1), (24, 4)], ids=["h6kv1", "h24kv4"]) +@pytest.mark.parametrize("num_requests,tokens", SPLITS, ids=[f"{r}x{t}" for r, t in SPLITS]) +def test_k3_drafter_attn(heads, kv, num_requests, tokens) -> None: + with torch.inference_mode(): + _check_attn(_Case(100 * num_requests + tokens + heads, heads, kv, num_requests, tokens)) + + +@pytest.mark.parametrize( + "page_pad,strided_table", [(0, False), (3 * 2 * PAGE * D, True)], ids=["dense", "padded"] +) +def test_k3_drafter_attn_table_forms(page_pad, strided_table) -> None: + """Dense page-table rows; pages with a stride larger than a page; one request's 1-D page-table row.""" + with torch.inference_mode(): + case = _Case(7, 6, 1, 4, 8, page_pad=page_pad, strided_table=strided_table) + _check_attn(case) + one = _Case(8, 6, 1, 1, 8, page_pad=page_pad) + out_2d, out_1d = one.out(), one.out() + k3_drafter_attn(one.qkv, one.pool, one.table, one.ctx_len, 6, 1, out_2d) + k3_drafter_attn(one.qkv, one.pool, one.table[0], one.ctx_len, 6, 1, out_1d) + assert torch.equal(_bits(out_1d), _bits(out_2d)) + + +def test_k3_drafter_attn_graph_replay() -> None: + """Both entries captured once and replayed with the lengths, page-table rows and qkv rewritten in place.""" + with torch.inference_mode(): + case = _Case(41, 6, 1, 4, 8, strided_table=False) + case.lengths = [2041] * 4 # every request has 2041 cached rows in its first pages + case.ctx_len.fill_(2041) + for r in range(4): + for i in range(32): + case.pool[int(case.table[r, i]), :, :, :] = torch.nan_to_num( + case.pool[int(case.table[r, i])], nan=0.25 + ) + gen = torch.Generator(device="cuda").manual_seed(42) + q_w = (1.0 + 0.2 * torch.randn(D, generator=gen, device="cuda")).bfloat16() + k_w = (1.0 + 0.2 * torch.randn(D, generator=gen, device="cuda")).bfloat16() + out, out_n = case.out(), case.out() + + def body(): + k3_drafter_attn(case.qkv, case.pool, case.table, case.ctx_len, 6, 1, out) + k3_drafter_attn_qknorm(case.qkv, q_w, k_w, case.positions, EPS, THETA, case.pool, case.table, case.ctx_len, + 6, 1, out_n) # fmt: skip + + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(stream): + body() # the first call of each entry compiles, outside the capture + torch.cuda.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + body() + torch.cuda.synchronize() + for rep in range(4): + lengths = [(500 * (rep + 1) + 37 * r) % 2000 for r in range(4)] + case.lengths = lengths + case.ctx_len.copy_(torch.tensor(lengths, dtype=torch.int32)) + case.positions.copy_( + (case.ctx_len.view(-1, 1) + torch.arange(8, device="cuda")).reshape(-1) + ) + case.table.copy_(torch.roll(case.table, shifts=1, dims=0)) + case.qkv.copy_( + (torch.randn(case.qkv.shape, generator=gen, device="cuda") * 0.5).bfloat16() + ) + graph.replay() + torch.cuda.synchronize() + assert _rel_err(out, _attention_ref(case, case.qkv)) <= TOL, f"replay {rep}" + want_n = _attention_ref(case, _qk_norm_rope_ref(case, case.qkv, q_w, k_w)) + assert _rel_err(out_n, want_n) <= TOL, f"replay {rep} (qknorm)" + + +def test_k3_drafter_attn_rejects_out_of_contract() -> None: + with torch.inference_mode(): + case = _Case(3, 6, 1, 2, 8) + out = case.out() + bad = [ + dict(num_heads=12), # not 6 query heads per KV head + dict(page_table=case.table[0]), # one page-table row for two requests + dict(page_table=case.table.t().contiguous().t()), # page-table columns not dense + dict(ctx_len=case.ctx_len.long()), # int64 lengths + dict(out=out.float()), # fp32 output + ] + for change in bad: + args = dict(qkv=case.qkv, cache=case.pool, page_table=case.table, ctx_len=case.ctx_len, num_heads=6, + num_kv_heads=1, out=out) # fmt: skip + args.update(change) + with pytest.raises(ValueError): + k3_drafter_attn(**args) + many = _Case(4, 6, 1, 1, 8) # 16 rows of one request: T > 8 + with pytest.raises(ValueError): + k3_drafter_attn(many.qkv.repeat(2, 1), many.pool, many.table, many.ctx_len, 6, 1, + torch.empty(16, 6 * D, dtype=torch.bfloat16, device="cuda")) # fmt: skip diff --git a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn_qknorm.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn_qknorm.py new file mode 100644 index 000000000000..3f08befe6d89 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn_qknorm.py @@ -0,0 +1,158 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the k3_drafter_attn_qknorm catalog entry (sm_100).""" + +import pytest +import torch + +assert torch.cuda.is_available(), "k3_drafter_attn requires a CUDA device" + +if torch.cuda.get_device_capability() != (10, 0): + # The CuTe DSL kernel uses tcgen05, TMA and clusters; it is certified on sm_100 (B200 / GB200) only. + pytest.skip("k3_drafter_attn is certified on sm_100 only", allow_module_level=True) + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.k3_drafter_attn_qknorm import ( # noqa: E402 + k3_drafter_attn_qknorm, +) + +D = 64 +PAGE = 64 +EPS = 1e-5 +THETA = 10000.0 +# max |err| / max |ref|: the kernel rounds P and the output to bf16 (~4e-3 each). +TOL = 1e-2 +# Context lengths, cycled over the requests: none, a block crossing a page, page and 128-row tile boundaries, the +# cluster's 16-tile round, several tiles per CTA. +LENGTHS = (60, 2041, 0, 127, 1000, 64, 128, 5) +SPLITS = ((1, 1), (1, 8), (2, 4), (4, 2), (8, 1), (3, 1), (2, 8), (8, 8)) + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def _rel_err(a: torch.Tensor, ref: torch.Tensor) -> float: + return (a.double() - ref).abs().max().item() / max(ref.abs().max().item(), 1e-6) + + +class _Case: + """R requests of T tokens: an HND pool [pages, 2, kv, 64, 64] that is NaN outside the requests' context rows + (optionally with a page stride larger than a page), distinct random pages per request, page-table rows (a strided + view into a wider table, or dense), qkv and lengths.""" + + def __init__(self, seed, heads, kv, num_requests, tokens, page_pad=0, strided_table=True): + gen = torch.Generator(device="cuda").manual_seed(seed) + self.heads, self.kv, self.r, self.t = heads, kv, num_requests, tokens + self.lengths = [LENGTHS[(seed + 3 * r) % len(LENGTHS)] for r in range(num_requests)] + need = [(c + PAGE - 1) // PAGE for c in self.lengths] + self.width = max(need) + 2 + n_pool = sum(need) + 2 * num_requests + 4 + inner = 2 * kv * PAGE * D + flat = torch.full( + (n_pool * (inner + page_pad),), float("nan"), dtype=torch.bfloat16, device="cuda" + ) + self.pool = flat.as_strided( + (n_pool, 2, kv, PAGE, D), (inner + page_pad, kv * PAGE * D, PAGE * D, D, 1) + ) + perm = torch.randperm(n_pool, generator=torch.Generator().manual_seed(seed)) + rows, used = [], 0 + for n in need: + rows.append(torch.cat([perm[used : used + n], perm[-(self.width - n) :]])) + used += n + dense = torch.stack(rows).to(torch.int32).cuda() + if strided_table: + full = torch.full( + (num_requests + 2, self.width + 3), -1, dtype=torch.int32, device="cuda" + ) + full[1 : num_requests + 1, : self.width] = dense + self.table = full[1 : num_requests + 1, : self.width] + else: + self.table = dense + for r, c in enumerate(self.lengths): + for i in range((c + PAGE - 1) // PAGE): + n_rows = min(PAGE, c - i * PAGE) + p = int(self.table[r, i]) + self.pool[p, :, :, :n_rows] = ( + torch.randn(2, kv, n_rows, D, generator=gen, device="cuda") * 0.5 + ).bfloat16() + m = num_requests * tokens + self.qkv = ( + torch.randn(m, (heads + 2 * kv) * D, generator=gen, device="cuda") * 0.5 + ).bfloat16() + self.ctx_len = torch.tensor(self.lengths, dtype=torch.int32, device="cuda") + self.positions = ( + self.ctx_len.view(-1, 1) + torch.arange(tokens, device="cuda", dtype=torch.int32) + ).reshape(-1) + + def out(self) -> torch.Tensor: + return torch.empty(self.r * self.t, self.heads * D, dtype=torch.bfloat16, device="cuda") + + +def _attention_ref(case: _Case, qkv: torch.Tensor) -> torch.Tensor: + """fp64: each request's q against its cached rows and its block's own k / v, softmax(q k^T / 8) v.""" + h, kv, t = case.heads, case.kv, case.t + outs = [] + for r, c in enumerate(case.lengths): + rows = qkv[r * t : (r + 1) * t].double() + q = rows[:, : h * D].view(t, h, D) + k_blk = rows[:, h * D : (h + kv) * D].view(t, kv, D).permute(1, 0, 2) + v_blk = rows[:, (h + kv) * D :].view(t, kv, D).permute(1, 0, 2) + pages = case.table[r, : (c + PAGE - 1) // PAGE].long() + k_ctx = case.pool[pages, 0].double().permute(1, 0, 2, 3).reshape(kv, -1, D)[:, :c] + v_ctx = case.pool[pages, 1].double().permute(1, 0, 2, 3).reshape(kv, -1, D)[:, :c] + k = torch.cat([k_ctx, k_blk], 1).repeat_interleave(h // kv, 0) + v = torch.cat([v_ctx, v_blk], 1).repeat_interleave(h // kv, 0) + p = torch.softmax(torch.einsum("thd,hld->thl", q, k) / D**0.5, dim=-1) + outs.append(torch.einsum("thl,hld->thd", p, v).reshape(t, h * D)) + return torch.cat(outs) + + +def _qk_norm_rope_ref(case: _Case, qkv: torch.Tensor, q_w, k_w) -> torch.Tensor: + """Per-head RMSNorm of the q and k heads, then NeoX RoPE (base THETA) at each row's position, in fp64 and rounded + to bf16 as fused_qk_norm_rope stores them; v as is.""" + h, kv = case.heads, case.kv + x = qkv.double().view(qkv.shape[0], h + 2 * kv, D).clone() + inv_freq = 1.0 / THETA ** (torch.arange(0, D, 2, dtype=torch.float64, device="cuda") / D) + angle = case.positions.double()[:, None] * inv_freq + cos, sin = angle.cos()[:, None, :], angle.sin()[:, None, :] + for start, count, w in ((0, h, q_w), (h, kv, k_w)): + y = x[:, start : start + count] + y = y * torch.rsqrt(y.pow(2).mean(-1, keepdim=True) + EPS) * w.double() + y1, y2 = y[..., : D // 2], y[..., D // 2 :] + x[:, start : start + count] = torch.cat([y1 * cos - y2 * sin, y1 * sin + y2 * cos], -1) + return x.view(qkv.shape[0], -1).bfloat16() + + +@pytest.mark.parametrize("num_requests,tokens", SPLITS, ids=[f"{r}x{t}" for r, t in SPLITS]) +def test_k3_drafter_attn_qknorm(num_requests, tokens) -> None: + """The raw projection output, normed and roped in the kernel, against the fp64 norm + RoPE + attention; ``qkv`` + left as it was; int64 positions give the same bits as int32.""" + with torch.inference_mode(): + case = _Case(300 + 10 * num_requests + tokens, 6, 1, num_requests, tokens) + gen = torch.Generator(device="cuda").manual_seed(num_requests * 9 + tokens) + raw = (case.qkv.float() * 4.0).bfloat16() + q_w = (1.0 + 0.2 * torch.randn(D, generator=gen, device="cuda")).bfloat16() + k_w = (1.0 + 0.2 * torch.randn(D, generator=gen, device="cuda")).bfloat16() + raw_before = raw.clone() + out, out64 = case.out(), case.out() + k3_drafter_attn_qknorm( + raw, + q_w, + k_w, + case.positions, + EPS, + THETA, + case.pool, + case.table, + case.ctx_len, + 6, + 1, + out, + ) + k3_drafter_attn_qknorm(raw, q_w, k_w, case.positions.long(), EPS, THETA, case.pool, case.table, case.ctx_len, + 6, 1, out64) # fmt: skip + assert torch.equal(_bits(raw), _bits(raw_before)), "qkv was written" + assert torch.equal(_bits(out64), _bits(out)), "int64 positions differ from int32" + want = _attention_ref(case, _qk_norm_rope_ref(case, raw, q_w, k_w)) + err = _rel_err(out, want) + assert err <= TOL, f"rel err {err:.3e}" From 90638fd88ae2d3dd18441ddac2ea7e362b9977cc Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 21:21:20 -0700 Subject: [PATCH 079/161] [None][fix] Kimi K3 collective state: create() involves its TP group's ranks only Every Kimi K3 state's create() made its communicator by splitting the session communicator, a collective over every rank of the session: with pipeline or context parallelism every rank of the job had to call create() together. Both create helpers now take it from _get_mnnvl_tp_group_comm, as MnnvlWorkspace.create does: a communicator of the TP group's ranks alone (MPI_Comm_create_group; the TP ProcessGroup under Ray), freed on the agreed failures. The latent exchange matrix checks the helper with the job split into two TP groups, one group making its communicator while the other does not call; the contracts say which ranks call create(). Also: the MoE front's compile cache is keyed by device, and the k3_moe contract states that the op picks its head_flags build from the head tensors it is given, which the entry and K3MoeLayer tie to the state. Signed-off-by: Vasanth Sabavat --- .../catalog/comm/k3_latent_reduce.md | 10 ++-- .../catalog/comm/k3_sandwich_oproj.md | 10 ++-- .../catalog/comm/k3_sandwich_plain.md | 10 ++-- .../catalog/comm/k3_sandwich_tail.md | 10 ++-- .../modeling_v2/catalog/moe/k3_moe.md | 8 +-- .../modeling_v2/catalog/moe/k3_moe_front.md | 36 +++++++------ .../cute_dsl_kernels/k3_fused_moe/front_op.py | 1 + .../cute_dsl_kernels/k3_fused_moe/op.py | 11 ++-- .../_torch/cute_dsl_kernels/k3_sandwich/op.py | 13 ++--- .../comm/_k3_latent_reduce_op_matrix.py | 54 +++++++++++++++---- .../comm/_k3_moe_front_op_matrix.py | 14 ++--- .../modeling_v2/comm/_k3_sandwich_common.py | 14 ++--- 12 files changed, 123 insertions(+), 68 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md index 9f67e9d0cdc1..cce0226627f0 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md @@ -81,13 +81,15 @@ of up to 8 tokens fits. **Who creates it, and when.** The target, in `post_load_weights`, with `K3LatentExchange.create(mapping, fabric_handle=None)`: -- collective over `mapping`'s TP group: every rank calls it at the same point. Every rank first joins the split of - the TP group's communicator, so a rank that calls it while its peers do not waits for them there; +- collective over `mapping`'s TP group only: every rank of the group, and no other rank of the session, calls it + at the same point. Under MPI its communicator is made from the group's ranks alone (`MPI_Comm_create_group`; + certified: with the job split into two TP groups of `W / 2`, one group makes its communicator while the other's + ranks do not call); a rank that calls it while its group's peers do not waits for them; - failure model (`k3_fused_moe.op.create_mcast_state`, the same as `MnnvlWorkspace.create`'s): - before allocating, the ranks agree that each of them can (not capturing, the buffer within that rank's free device memory). If one cannot, every rank raises `RuntimeError`, none allocates, and under MPI each frees the - communicator it split for the call (certified: every rank capturing, and one rank capturing while its peers - call it eagerly at the same point; every rank raises at that agreement and frees its split, the capturing + communicator made for the call (certified: every rank capturing, and one rank capturing while its peers + call it eagerly at the same point; every rank raises at that agreement and frees that communicator, the capturing rank's message naming the capture, and the next call is correct); - a failure that returns from the allocation is agreed and handled the same way; - a rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange; this diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md index b1310543bd46..b858953e3a9a 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md @@ -88,13 +88,17 @@ every word empty, every counter zero. **Who creates it, and when.** The target, in `post_load_weights`, with `K3SandwichWorkspace.create(mapping, fabric_handle=None)`: -- collective over `mapping`'s TP group: every rank calls it at the same point; +- collective over `mapping`'s TP group only: every rank of the group, and no other rank of the session, calls it + at the same point. Under MPI its communicator is made from the group's ranks alone (`MPI_Comm_create_group`; + certified in `comm/k3_latent_reduce`'s matrix: with the job split into two TP groups of `W / 2`, one group + makes its communicator while the other's ranks do not call); a rank that calls it while its group's peers do + not waits for them; - failure model: - before allocating, the ranks agree that each of them can (not capturing, the buffer within that rank's free device memory). If one cannot, every rank raises `RuntimeError`, none allocates, and under MPI each frees the - communicator it split for the call (certified: one rank capturing while its peers call it eagerly, every rank + communicator made for the call (certified: one rank capturing while its peers call it eagerly, every rank raises, the capturing rank naming the capture and its peers another rank; no rank reaches the allocation, - every rank frees its split, and the next call is correct); + every rank frees that communicator, and the next call is correct); - a failure that returns from the allocation is agreed and handled the same way; - a rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange; this is not turned into an error on the other ranks; diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md index 3b50dd1c8fd9..29f0b09cc826 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md @@ -86,13 +86,17 @@ target's sandwiches'. Certified after `create`: sized for the group, every word **Who creates it, and when.** The target, in `post_load_weights`, with `K3SandwichWorkspace.create(mapping, fabric_handle=None)`: -- collective over `mapping`'s TP group: every rank calls it at the same point; +- collective over `mapping`'s TP group only: every rank of the group, and no other rank of the session, calls it + at the same point. Under MPI its communicator is made from the group's ranks alone (`MPI_Comm_create_group`; + certified in `comm/k3_latent_reduce`'s matrix: with the job split into two TP groups of `W / 2`, one group + makes its communicator while the other's ranks do not call); a rank that calls it while its group's peers do + not waits for them; - failure model: - before allocating, the ranks agree that each of them can (not capturing, the buffer within that rank's free device memory). If one cannot, every rank raises `RuntimeError`, none allocates, and under MPI each frees the - communicator it split for the call (certified: one rank capturing while its peers call it eagerly, every rank + communicator made for the call (certified: one rank capturing while its peers call it eagerly, every rank raises, the capturing rank naming the capture and its peers another rank; no rank reaches the allocation, - every rank frees its split, and the next call is correct); + every rank frees that communicator, and the next call is correct); - a failure that returns from the allocation is agreed and handled the same way; - a rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange; this is not turned into an error on the other ranks; diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md index 6c069b684d37..068c1532da24 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md @@ -118,13 +118,17 @@ CTAs, with the handle that owns the memory and the communicator. This op's parti **Who creates it, and when.** The target, in `post_load_weights`, with `K3SandwichWorkspace.create(mapping, fabric_handle=None)`: -- collective over `mapping`'s TP group: every rank calls it at the same point; +- collective over `mapping`'s TP group only: every rank of the group, and no other rank of the session, calls it + at the same point. Under MPI its communicator is made from the group's ranks alone (`MPI_Comm_create_group`; + certified in `comm/k3_latent_reduce`'s matrix: with the job split into two TP groups of `W / 2`, one group + makes its communicator while the other's ranks do not call); a rank that calls it while its group's peers do + not waits for them; - failure model: - before allocating, the ranks agree that each of them can (not capturing, the buffer within that rank's free device memory). If one cannot, every rank raises `RuntimeError`, none allocates, and under MPI each frees the - communicator it split for the call (certified: one rank capturing while its peers call it eagerly, every rank + communicator made for the call (certified: one rank capturing while its peers call it eagerly, every rank raises, the capturing rank naming the capture and its peers another rank; no rank reaches the allocation, - every rank frees its split, and the next call is correct); + every rank frees that communicator, and the next call is correct); - a failure that returns from the allocation is agreed and handled the same way; - a rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange; this is not turned into an error on the other ranks; diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md index 590d921bf862..e4f139541269 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md @@ -136,9 +136,11 @@ counters. Layers of one state may be built over the same weight buffers (two cou (certified, both). Separate states share nothing: with calls on two `K3MoeState`s and a `K3MoeWideState` interleaved in an irregular pattern (26 calls, back to back), each returns the bits of the same call made alone (certified). The entry passes `head` for, and only for, a head_flags state's layers: both mismatches raise `ValueError` before any -launch (certified). A head_flags state's calls pair with the front's on one `K3MoeHeadWorkspace`: the front call -before each must publish the workspace's ready words (the `moe/k3_moe_front` entry with `publish=True`), and each -publishing front call must be followed by exactly one such `k3_moe` call on that workspace (*Preconditions*). +launch (certified). The op itself picks the head_flags build by whether `head_ready` / `head_flags` are passed; this +check, and the same one in `K3MoeLayer`, is what ties the build to the state. A head_flags state's calls pair with the +front's on one `K3MoeHeadWorkspace`: the front call before each must publish the workspace's ready words (the +`moe/k3_moe_front` entry with `publish=True`), and each publishing front call must be followed by exactly one such +`k3_moe` call on that workspace (*Preconditions*). **Call-order invariant.** The calls on all layers of one state run one after the other in one stream order. Every call needs the slab armed and its layer's counters at zero, which only the end of the previous call on the state diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md index f7608ddea53d..fdb5ad7ef751 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md @@ -117,19 +117,21 @@ at `W` 4, 8 and 16 (certified at 4): The size depends on `W` only, not on `M`: every call fits. **Who creates it, and when.** The target, in `post_load_weights`, with `K3MoeHeadWorkspace.create(mapping, -fabric_handle=None)`: collective over `mapping`'s TP group and eager, every rank of the group calling it at the same -point. Under MPI each call first splits the group's communicator off the session's, a collective of every rank of the -session (`_get_mnnvl_workspace_comm`). Before allocating, the ranks agree that each of them can: not capturing a CUDA -graph, the buffer within its device's free memory. If one cannot, every rank raises `RuntimeError` ("not every rank -can allocate"), none allocates, and under MPI each frees the communicator it split for the call. A failure returned by -the allocation is agreed and handled the same way, and that second agreement is also the barrier that keeps any rank -from pushing into a peer's buffer before the peer has emptied it; a rank that fails inside the allocation's handle -exchange can leave its peers waiting there (the op module's statements). Certified: with every rank capturing, and -with the last rank capturing while its peers call it eagerly, every rank raises `RuntimeError` (the capturing ranks' -messages naming the capture) and frees its split, and the existing workspace keeps its bits. So even a refusal needs -every rank to call `create()`: a rank calling it alone waits for its peers. It empties every word and zeroes `flags` -and `ready` (certified, both of the matrix's workspaces). `fabric_handle`: share the memory by fabric handle (required -across nodes) or POSIX file descriptor; default `mapping.is_multi_node()`. No environment variable is read. +fabric_handle=None)`: collective over `mapping`'s TP group only and eager: every rank of the group, and no other rank +of the session, calls it at the same point. Under MPI its communicator is made from the group's ranks alone +(`MPI_Comm_create_group`; certified in `comm/k3_latent_reduce`'s matrix, where one of two TP groups makes its +communicator while the other's ranks do not call). Before allocating, the ranks agree that each of them can: not +capturing a CUDA graph, the buffer within its device's free memory. If one cannot, every rank raises `RuntimeError` +("not every rank can allocate"), none allocates, and under MPI each frees the communicator made for the call. A +failure returned by the allocation is agreed and handled the same way, and that second agreement is also the barrier +that keeps any rank from pushing into a peer's buffer before the peer has emptied it; a rank that fails inside the +allocation's handle exchange can leave its peers waiting there (the op module's statements). Certified: with every +rank capturing, and with the last rank capturing while its peers call it eagerly, every rank raises `RuntimeError` +(the capturing ranks' messages naming the capture) and frees that communicator, and the existing workspace keeps its +bits. So even a refusal needs every rank of the group to call `create()`: a rank calling it alone waits for its peers. +It empties every word and zeroes `flags` and `ready` (certified, both of the matrix's workspaces). `fabric_handle`: +share the memory by fabric handle (required across nodes) or POSIX file descriptor; default `mapping.is_multi_node()`. +No environment variable is read. **Which ops may share one object.** Every MoE front call of the TP group: this entry, plain or publishing (`publish=True`, the front a head_flags `k3_moe` pairs with). The head_flags calls of `moe/k3_moe` (`head=workspace`) @@ -195,10 +197,10 @@ never push. Besides `workspace` (an explicit argument): -- A process-wide cache of compiled kernels keyed by (`W`, `shared_cols`, tiles, clusters, `K`, ring, `gate_cap`, - `linear_cap`, publish, PDL, half-tile head). The first call of a key compiles (seconds) and must be eager: under - capture it raises `RuntimeError` ("must run once per configuration outside CUDA-graph capture first") before any - launch, the workspace untouched (certified with other SiTU caps). The cache is result-neutral. +- A process-wide cache of compiled kernels keyed by (device, `W`, `shared_cols`, tiles, clusters, `K`, ring, + `gate_cap`, `linear_cap`, publish, PDL, half-tile head). The first call of a key compiles (seconds) and must be + eager: under capture it raises `RuntimeError` ("must run once per configuration outside CUDA-graph capture first") + before any launch, the workspace untouched (certified with other SiTU caps). The cache is result-neutral. - `TRTLLM_ENABLE_PDL` (default on), read on every call and part of the key; it changes scheduling, not results (the op's statement). - The device's cluster capacity (`max_clusters`, cached per device), which sizes the grid and picks the head's diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/front_op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/front_op.py index dcb359241cf3..53c50262d952 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/front_op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/front_op.py @@ -192,6 +192,7 @@ def arg(t): use_pdl = os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" stream = cuda_driver.CUstream(torch.cuda.current_stream(device).cuda_stream) key = ( + device.index if device.index is not None else torch.cuda.current_device(), ag_world, shared_cols, ht + st, diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py index bc51c9e7a4d3..69f70d200ba9 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py @@ -114,16 +114,17 @@ def is_supported(w3_w1_weight: torch.Tensor, w3_w1_weight_scale: torch.Tensor, w def create_mcast_state(name: str, mapping, words: int, fabric_handle: Optional[bool], build): """``build(uc, mc, handle, comm)`` over a new multicast buffer of ``words`` int32 per rank of ``mapping``'s TP group, every word empty (``uc``: this rank's words; ``mc``: the same words through the multicast mapping). - Collective and eager: every rank of the group calls it at the same point. + Collective over the TP group only, and eager: every rank of the group, and no other rank, calls it at the same + point (its communicator is made from the group's ranks alone). Failure model (as ``MnnvlWorkspace.create``): before allocating, the ranks agree that each of them can (not capturing a CUDA graph, the buffer within its device's free memory); if one cannot, every rank raises - ``RuntimeError``, none allocates, and under MPI each frees the communicator it split for the call. A failure that + ``RuntimeError``, none allocates, and under MPI each frees the communicator made for the call. A failure that returns from the allocation or from ``build`` is agreed and handled the same way. A rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange: that failure is not turned into an error on the other ranks.""" from tensorrt_llm._torch.distributed.ops import ( - _get_mnnvl_workspace_comm, + _get_mnnvl_tp_group_comm, _make_mnnvl_mcast_buffer, _mnnvl_device_index, _mnnvl_workspace_all_succeeded, @@ -131,7 +132,7 @@ def create_mcast_state(name: str, mapping, words: int, fabric_handle: Optional[b from tensorrt_llm._utils import mpi_disabled use_fabric_handle = mapping.is_multi_node() if fabric_handle is None else bool(fabric_handle) - comm = _get_mnnvl_workspace_comm(mapping) + comm = _get_mnnvl_tp_group_comm(mapping) # Every condition one rank alone can fail is checked before the allocation, and the ranks agree on it: a rank # failing inside the allocation would leave its peers in the handle exchange. problem: Optional[str] = None @@ -142,7 +143,7 @@ def create_mcast_state(name: str, mapping, words: int, fabric_handle: Optional[b if free_bytes < words * 4: problem = f"its {words * 4} bytes exceed the {free_bytes} free on this rank's device" if not _mnnvl_workspace_all_succeeded(comm, problem is None): - # Every rank takes this path: free the MPI communicator split above (a ProcessGroup is c10d's). + # Every rank takes this path: free the MPI communicator made above (a ProcessGroup is c10d's). if not mpi_disabled(): comm.Free() raise RuntimeError( diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py index 00091954c490..a467d0db8c16 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py @@ -131,17 +131,18 @@ def _slab_args(x_slab: Optional[torch.Tensor], slab_buf: int, fallback: torch.Te def _create_buffer(cls, mapping, words: int, flag_words: int, fabric_handle: Optional[bool], arm_flags: Optional[Callable[[torch.Tensor], None]] = None): # fmt: skip """A ``cls`` over a new multicast buffer of ``words`` int32 per rank of ``mapping``'s TP group, every word empty, - and ``flag_words`` int32 flags, zero (then ``arm_flags``). Collective and eager: every rank of the group calls it - at the same point. + and ``flag_words`` int32 flags, zero (then ``arm_flags``). Collective over the TP group only, and eager: every + rank of the group, and no other rank, calls it at the same point (its communicator is made from the group's + ranks alone). Failure model (as ``MnnvlWorkspace.create``): before allocating, the ranks agree that each of them can (not capturing a CUDA graph, the buffer within its device's free memory); if one cannot, every rank raises - ``RuntimeError``, none allocates, and under MPI each frees the communicator it split for the call. A failure + ``RuntimeError``, none allocates, and under MPI each frees the communicator made for the call. A failure that returns from the allocation is agreed and handled the same way. A rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange: that failure is not turned into an error on the other ranks.""" from tensorrt_llm._torch.distributed.ops import ( - _get_mnnvl_workspace_comm, + _get_mnnvl_tp_group_comm, _make_mnnvl_mcast_buffer, _mnnvl_device_index, _mnnvl_workspace_all_succeeded, @@ -149,7 +150,7 @@ def _create_buffer(cls, mapping, words: int, flag_words: int, fabric_handle: Opt from tensorrt_llm._utils import mpi_disabled use_fabric_handle = mapping.is_multi_node() if fabric_handle is None else bool(fabric_handle) - comm = _get_mnnvl_workspace_comm(mapping) + comm = _get_mnnvl_tp_group_comm(mapping) # Every condition one rank alone can fail is checked before the allocation, and the ranks agree on it: a rank # failing inside the allocation would leave its peers in the handle exchange. problem: Optional[str] = None @@ -160,7 +161,7 @@ def _create_buffer(cls, mapping, words: int, flag_words: int, fabric_handle: Opt if free_bytes < words * 4: problem = f"its {words * 4} bytes exceed the {free_bytes} free on this rank's device" if not _mnnvl_workspace_all_succeeded(comm, problem is None): - # Every rank takes this path: free the MPI communicator split above (a ProcessGroup is c10d's). + # Every rank takes this path: free the MPI communicator made above (a ProcessGroup is c10d's). if not mpi_disabled(): comm.Free() raise RuntimeError( diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_latent_reduce_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_k3_latent_reduce_op_matrix.py index 83ca05f580fe..785711e8e883 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_k3_latent_reduce_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_latent_reduce_op_matrix.py @@ -247,15 +247,48 @@ def check_exchange_is_armed_and_sized() -> None: assert EX_A.state.flags.data_ptr() != EX_B.state.flags.data_ptr(), "two exchanges, two counts" +def check_create_needs_only_the_tp_group() -> None: + """Under MPI a create() makes its communicator from its TP group's ranks alone (the helper every Kimi K3 state's + create() uses). With the job split into two TP groups of W / 2 (pipeline parallel 2), the first group's ranks make + theirs while the second group's ranks do not call at all; each communicator holds exactly its group, in TP-rank + order. The exchanges in use also hold their TP group's communicator.""" + from tensorrt_llm._torch.distributed import ops + from tensorrt_llm.mapping import Mapping + + for ex in (EX_A, EX_B): + assert ex.state.comm.Get_size() == R.world and ex.state.comm.Get_rank() == R.rank + half = Mapping( + world_size=R.world, + rank=R.rank, + gpus_per_node=R.mapping.gpus_per_node, + tp_size=R.world // 2, + pp_size=2, + ) + R.barrier() + ok = True + if half.pp_rank == 0: + comm = ops._get_mnnvl_tp_group_comm(half) + ok = ( + comm.Get_size() == half.tp_size + and comm.Get_rank() == half.tp_rank + and comm.allgather(R.rank) == list(half.tp_group) + ) + comm.Free() + R.comm.Barrier() + assert R.all_true(ok), ( + f"rank {R.rank}: the first TP group's communicator is not exactly that group" + ) + + def check_capture_refusals() -> None: """``K3LatentExchange.create`` is collective: every rank joins the TP group's communicator, and before allocating the ranks agree that each of them can. With every rank capturing a CUDA graph, and with one rank capturing while its peers call it eagerly at the same point, every rank raises RuntimeError, and a capturing rank's message names the capture. Every rank raises at that agreement ("not every rank can allocate"; a failure after allocating reads - "allocation failed"), so nothing is allocated, and each frees the communicator it split. The op's first call, which - would compile the kernel, is refused per rank: under capture it raises on every rank before it launches anything. - The exchange is untouched and the next call is correct. Runs before every eager call of the op: the compile cache - must still be cold.""" + "allocation failed"), so nothing is allocated, and each frees the communicator made for it. The op's first call, + which would compile the kernel, is refused per rank: under capture it raises on every rank before it launches + anything. The exchange is untouched and the next call is correct. Runs before every eager call of the op: the + compile cache must still be cold.""" # Imported by the op's first call; imported here so that nothing is imported inside the capture. import cutlass.cute.runtime # noqa: F401 @@ -276,14 +309,14 @@ def refusal(capturing: bool) -> str: from tensorrt_llm._torch.distributed import ops - split = ops._get_mnnvl_workspace_comm + make = ops._get_mnnvl_tp_group_comm comms = [] - def recording_split(mapping): - comms.append(split(mapping)) + def recording_make(mapping): + comms.append(make(mapping)) return comms[-1] - ops._get_mnnvl_workspace_comm = recording_split + ops._get_mnnvl_tp_group_comm = recording_make try: for case, capturing in (("every rank", True), ("one rank", R.rank == R.world - 1)): R.barrier() @@ -294,9 +327,9 @@ def recording_split(mapping): f"{case} capturing: rank {R.rank} (capturing {capturing}) got {message!r}" ) finally: - ops._get_mnnvl_workspace_comm = split + ops._get_mnnvl_tp_group_comm = make freed = len(comms) == 2 and all(c == R.MPI.COMM_NULL for c in comms) - assert R.all_true(freed), "a refused create kept the communicator it split" + assert R.all_true(freed), "a refused create kept the communicator made for it" R.barrier() first = raised_under_capture(lambda: entry(MAX_TOKENS, EX_A.state)) assert R.all_true("outside CUDA-graph capture first" in first), f"first call: {first!r}" @@ -570,6 +603,7 @@ def check_token_count_mismatch_returns_stale_rows() -> None: CHECKS = [ check_exchange_is_armed_and_sized, + check_create_needs_only_the_tp_group, # Before every eager call of the op: it needs a cold compile cache. check_capture_refusals, check_single_calls, diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py index ac5c76ef6f28..b6de2bc7fb96 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py @@ -676,7 +676,7 @@ def check_create_and_first_compile_refuse_capture() -> None: K3MoeHeadWorkspace.create has the ranks agree before allocating that each of them can (not capturing, enough free memory): with every rank capturing, and with the last rank capturing while its peers call it eagerly, every rank raises RuntimeError at that agreement ("not every rank can allocate", a capturing rank's message naming the - capture), so none allocates, and each frees the communicator it split for the call. Every rank must call it: a + capture), so none allocates, and each frees the communicator made for the call. Every rank must call it: a rank calling it alone would wait for its peers. A front call of a configuration not yet compiled (other SiTU caps) raises RuntimeError under capture before any launch. The workspace keeps its bits and the next call returns the bits of the same call made alone. @@ -690,21 +690,21 @@ def check_create_and_first_compile_refuse_capture() -> None: def create(): OPS.workspace.create(R.mapping, fabric_handle=R.fabric) - split = ops._get_mnnvl_workspace_comm + make = ops._get_mnnvl_tp_group_comm comms = [] - def recording_split(mapping): - comms.append(split(mapping)) + def recording_make(mapping): + comms.append(make(mapping)) return comms[-1] capturing = R.world - 1 - ops._get_mnnvl_workspace_comm = recording_split + ops._get_mnnvl_tp_group_comm = recording_make try: every = raised_under_capture(create) R.barrier() one = raised_under_capture(create) if R.rank == capturing else raised_eagerly(create) finally: - ops._get_mnnvl_workspace_comm = split + ops._get_mnnvl_tp_group_comm = make R.barrier() first = raised_under_capture( lambda: OPS.front( @@ -723,7 +723,7 @@ def recording_split(mapping): f"uncompiled front {first!r}" ) freed = len(comms) == 2 and all(c == R.MPI.COMM_NULL for c in comms) - assert R.all_true(freed), "a refused create kept the communicator it split" + assert R.all_true(freed), "a refused create kept the communicator made for it" after = workspace_snapshot(WS_A) assert all(same(a, b) for a, b in zip(before, after)), ( "a refused call touched the head workspace" diff --git a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_common.py b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_common.py index d59e6773d128..af13c070ecf7 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_common.py +++ b/tests/unittest/_torch/modeling_v2/comm/_k3_sandwich_common.py @@ -520,34 +520,34 @@ def create_refuses_capture(workspace_type, ws, next_call: Call) -> None: rank raise RuntimeError. (a) Every rank captures: every rank raises, naming the capture. (b) The last rank captures while its peers call ``create`` eagerly at the same point: every rank raises, the capturing rank naming the capture, its peers saying another rank cannot. No rank reaches the allocation in either case (the counted - multicast allocation, which the workspaces in use went through), each rank frees the communicator it split, no + multicast allocation, which the workspaces in use went through), each rank frees the communicator made for it, no stream is left capturing, and the workspace in use is untouched: the next call is correct.""" from tensorrt_llm._torch.distributed import ops assert ALLOCATIONS[0] > 0, "the counter did not see the workspaces in use being created" before = ALLOCATIONS[0] - split = ops._get_mnnvl_workspace_comm + make = ops._get_mnnvl_tp_group_comm comms = [] - def recording_split(mapping): - comms.append(split(mapping)) + def recording_make(mapping): + comms.append(make(mapping)) return comms[-1] capturing = R.world - 1 - ops._get_mnnvl_workspace_comm = recording_split + ops._get_mnnvl_tp_group_comm = recording_make try: every = _create(workspace_type, capture=True) R.barrier() one = _create(workspace_type, capture=R.rank == capturing) finally: - ops._get_mnnvl_workspace_comm = split + ops._get_mnnvl_tp_group_comm = make every_ok = CAPTURE_REFUSED in every one_ok = (CAPTURE_REFUSED if R.rank == capturing else PEER_REFUSED) in one allocated = ALLOCATIONS[0] - before freed = len(comms) == 2 and all(c == R.MPI.COMM_NULL for c in comms) assert R.all_true(every_ok and one_ok and allocated == 0 and freed), ( f"every rank capturing: {every!r}; rank {capturing} capturing: {one!r}; allocations reached {allocated}; " - f"communicator splits freed {freed}" + f"communicators freed {freed}" ) next_call.verify(next_call.run(ws), "after the refused creates") From 90158d95bc60f6d5d04c7e4092f18d79d633e92c Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 21:22:06 -0700 Subject: [PATCH 080/161] [None][doc] MNNVL catalog contracts: MnnvlWorkspace.create involves the TP group's ranks only As the workspace's create() now does: every rank of the TP group calls it, ranks outside the group take no part, and on a refusal each rank frees the communicator it made for the call. Signed-off-by: Vasanth Sabavat --- .../catalog/comm/mnnvl_allgather_split.md | 17 +++++++++-------- .../catalog/comm/mnnvl_fusion_allreduce.md | 17 +++++++++-------- 2 files changed, 18 insertions(+), 16 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md index 0e49452aa4e9..b09a7c2ceb78 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md @@ -99,15 +99,16 @@ two-shot `[64, 7168]` all-reduce of its shared-workspace sequence). **Who creates it, and when.** The target, in `post_load_weights`, with `MnnvlWorkspace.create(mapping, buffer_bytes, fabric_handle=None)` (see `mnnvl_allreduce_attn_res.md`): -- collective over the TP group: every rank calls it at the same point; +- collective over `mapping`'s TP group only: every rank of the group calls it at the same point, and ranks outside + the group take no part (certified in `mnnvl_allreduce_attn_res`'s matrix); - failure model: - - before allocating, the ranks agree that each of them can (not capturing, a valid `buffer_bytes`, the three - buffers within that rank's free device memory). If one cannot, every rank raises `RuntimeError` and none - allocates (certified: one rank inside a CUDA-graph capture while the others are not, and then every rank - capturing; each time every rank raises, the capturing ranks' message naming the capture, and the workspaces in - use are untouched; a create right after, eager on every rank, returns an armed workspace whose first call is - correct); - - a failure that returns from the allocation is agreed the same way; + - before allocating, the ranks agree that each of them can (not capturing, a valid `buffer_bytes`, the three buffers + within that rank's free device memory). If one cannot, every rank raises `RuntimeError`, none allocates, and under + MPI each frees the communicator it made for the call (certified: one rank inside a CUDA-graph capture while the + others are not, and then every rank capturing; each time every rank raises, the capturing ranks' message naming + the capture, and the workspaces in use are untouched; a create right after, eager on every rank, returns an armed + workspace whose first call is correct); + - a failure that returns from the allocation is agreed and handled the same way; - a rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange; this is not turned into an error on the other ranks; - eager: it allocates and exchanges handles, so it refuses to run under CUDA-graph capture (every rank raises); diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md index edffc4e9d2d6..a50d3fcb2fd8 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md @@ -113,15 +113,16 @@ calls of this op (arithmetic, not a test). The test's buffer is the one-shot foo **Who creates it, and when.** The target, in `post_load_weights`, with `MnnvlWorkspace.create(mapping, buffer_bytes, fabric_handle=None)` (see `mnnvl_allreduce_attn_res.md`): -- collective over the TP group: every rank calls it at the same point; +- collective over `mapping`'s TP group only: every rank of the group calls it at the same point, and ranks outside + the group take no part (certified in `mnnvl_allreduce_attn_res`'s matrix); - failure model: - - before allocating, the ranks agree that each of them can (not capturing, a valid `buffer_bytes`, the three - buffers within that rank's free device memory). If one cannot, every rank raises `RuntimeError` and none - allocates (certified: one rank inside a CUDA-graph capture while the others are not, and then every rank - capturing; each time every rank raises, the capturing ranks' message naming the capture, and the workspaces in - use are untouched; a create right after, eager on every rank, returns an armed workspace whose first call is - correct); - - a failure that returns from the allocation is agreed the same way; + - before allocating, the ranks agree that each of them can (not capturing, a valid `buffer_bytes`, the three buffers + within that rank's free device memory). If one cannot, every rank raises `RuntimeError`, none allocates, and under + MPI each frees the communicator it made for the call (certified: one rank inside a CUDA-graph capture while the + others are not, and then every rank capturing; each time every rank raises, the capturing ranks' message naming + the capture, and the workspaces in use are untouched; a create right after, eager on every rank, returns an armed + workspace whose first call is correct); + - a failure that returns from the allocation is agreed and handled the same way; - a rank that fails inside the allocation's handle exchange can leave its peers waiting in that exchange; this is not turned into an error on the other ranks; - eager: it allocates and exchanges handles, so it refuses to run under CUDA-graph capture (every rank raises); From 068241fa82d5aefd95311bedd7777df8a16821e5 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 21:22:35 -0700 Subject: [PATCH 081/161] [None][feat] modeling_v2 Kimi K3 target: attention outputs for fused post-attention steps The attention modules expose what a fused post-attention step needs: - will_run_decode_branch(attn_metadata, step), on K3DecodeKDA, K3DecodeMLA and KimiMLARuntime: whether forward runs the step on the decode branch. KDA: every step decode_step classifies, outside a breakable CUDA graph. MLA: a decode step whose paged cache the decode kernels read. MLA's KV writes allow one run per step, so a caller asks before the call. - reduce_output=False, on K3DecodeKDA and KimiMLARuntime: o_proj's TP partial, without the all-reduce. - project_output=False, on K3DecodeKDA, K3DecodeMLA and KimiMLARuntime: the input of o_proj, which is KDA's post-o_norm core or MLA's gated attention output. The flags raise ValueError where the decode branch does not run, except KimiMLARuntime's reduce_output, which holds on every step. The default call path is unchanged. Signed-off-by: Vasanth Sabavat --- .../modeling.py | 108 +++++++++++++----- 1 file changed, 81 insertions(+), 27 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index dc8a27202f56..6003fd0882c7 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -1249,15 +1249,27 @@ def __init__( aux_stream_dict=aux_stream_dict, ) + def will_run_decode_branch( + self, attn_metadata: AttentionMetadata, step: Optional[DecodeStep] + ) -> bool: + """Whether the mixer runs ``step`` on the decode kernels (``K3DecodeMLA.will_run_decode_branch``).""" + return self.mixer.will_run_decode_branch(attn_metadata, step) + def forward( self, hidden_states: torch.Tensor, attn_metadata: AttentionMetadata, step: Optional[DecodeStep] = None, + reduce_output: bool = True, + project_output: bool = True, ) -> torch.Tensor: + """``reduce_output=False`` returns ``o_proj``'s TP partial (no all-reduce); ``project_output=False`` the + mixer's gated attention output before ``o_proj``, only where ``will_run_decode_branch`` holds.""" # MLA.forward takes position_ids first; K3 is NoPE, so pass None. - out = self.mixer(None, hidden_states, attn_metadata, step=step) - if self._o_allreduce is not None: + out = self.mixer( + None, hidden_states, attn_metadata, step=step, project_output=project_output + ) + if project_output and reduce_output and self._o_allreduce is not None: # Head-sharded TP: sum the row-sharded o_proj partials across # the head-shard group. out = self._o_allreduce(out) @@ -1735,16 +1747,32 @@ def finalize_decode_weights(self) -> None: self.k3_proj_weight = fused self._qkvg_proj_weight, self._bfa_proj_weight = fused[:rows], fused[rows:] + def will_run_decode_branch( + self, attn_metadata: AttentionMetadata, step: Optional[DecodeStep] + ) -> bool: + """Whether ``forward`` runs ``step`` on the decode branch: a step ``decode_step`` classifies, outside a + breakable CUDA graph. Only there does it return the unreduced ``o_proj`` output or the core.""" + return step is not None and not is_in_breakable_cuda_graph() + def forward( self, hidden_states: torch.Tensor, attn_metadata: AttentionMetadata, step: Optional[DecodeStep] = None, + reduce_output: bool = True, + project_output: bool = True, ) -> torch.Tensor: - """The built-in forward on a step ``decode_step`` does not classify, and under a breakable CUDA graph. On the - others: the plain decode on ``ssm/k3_kda_decode_attn`` on a decode step of one token per request, else the - built-in dispatch; then ``o_proj`` on its decode GEMV site.""" - if step is None or is_in_breakable_cuda_graph(): + """The built-in forward where ``will_run_decode_branch`` does not hold. On the decode branch: the plain decode + on ``ssm/k3_kda_decode_attn`` on a decode step of one token per request, else the built-in dispatch; then + ``o_proj`` on its decode GEMV site and the TP all-reduce. + + On the decode branch only: ``reduce_output=False`` returns ``o_proj``'s TP partial (no all-reduce), and + ``project_output=False`` the post-o_norm core ``[N, H * 128]`` (no ``o_proj``).""" + if not self.will_run_decode_branch(attn_metadata, step): + if not (reduce_output and project_output): + raise ValueError( + "reduce_output / project_output need a step will_run_decode_branch takes" + ) return super().forward(hidden_states, attn_metadata) if ( step.decode @@ -1755,19 +1783,24 @@ def forward( core = self._k3_decode(hidden_states[: step.num_tokens], attn_metadata) else: core = self._forward_impl(hidden_states, attn_metadata) - return self._k3_project_output(core) - - def _k3_project_output(self, core: torch.Tensor) -> torch.Tensor: - """``o_proj`` on the ``o_proj`` decode GEMV site where it takes the rows (else the module), then the TP - all-reduce.""" + if not project_output: + return core.reshape(-1, self.proj_size) + return self._k3_project_output(core, reduce_output) + + def _k3_project_output(self, core: torch.Tensor, reduce_output: bool = True) -> torch.Tensor: + """``o_proj`` on the ``o_proj`` decode GEMV site where it takes the rows (else the module), then, with + ``reduce_output``, the TP all-reduce.""" + core2d = core.reshape(-1, self.proj_size) out = None if self.decode_gemvs is not None: - out = self.decode_gemvs.project( - "o_proj", core.reshape(-1, self.proj_size), self.o_proj.weight - ) + out = self.decode_gemvs.project("o_proj", core2d, self.o_proj.weight) if out is None: - return self._project_output(core) - return out if self._o_allreduce is None else self._o_allreduce(out) + if reduce_output: + return self._project_output(core) + out = self.o_proj(core2d) + if reduce_output and self._o_allreduce is not None: + out = self._o_allreduce(out) + return out def _k3_decode(self, x: torch.Tensor, attn_metadata: AttentionMetadata) -> torch.Tensor: """``ssm/k3_kda_decode_attn``: the core output ``[R, H, 128]`` of one token of each of the step's R requests; @@ -1958,6 +1991,29 @@ def _k3_layout_gaps(self) -> list: ) return [why for ok, why in checks if not ok] + def will_run_decode_branch( + self, attn_metadata: AttentionMetadata, step: Optional[DecodeStep] + ) -> bool: + """Whether ``forward`` runs ``step`` on the decode kernels (see ``_k3_step_view``). Only there does it return + the gated attention output before ``o_proj``; the KV writes forbid running the attention twice, so a caller + that needs it asks first.""" + return self._k3_step_view(attn_metadata, step) is not None + + def _k3_step_view( + self, attn_metadata: AttentionMetadata, step: Optional[DecodeStep] + ) -> Optional[dict]: + """The paged-cache view the decode kernels read on ``step``: a decode step, this layer's ``[W_a; W_g]`` and + workspace built, outside a breakable CUDA graph, and a cache ``k3_mla_decode_view`` takes. Else None.""" + if not ( + step is not None + and step.decode + and self.k3_ag_weight is not None + and self.k3_workspace is not None + and not is_in_breakable_cuda_graph() + ): + return None + return self._k3_decode_view(attn_metadata, step.num_tokens) + def forward( self, position_ids: Optional[torch.Tensor], @@ -1966,19 +2022,15 @@ def forward( all_reduce_params=None, latent_cache_gen: Optional[torch.Tensor] = None, step: Optional[DecodeStep] = None, + project_output: bool = True, ) -> torch.Tensor: - """The built-in forward, except on a decode step whose cache the decode kernels read.""" - view = None - if ( - step is not None - and step.decode - and latent_cache_gen is None - and self.k3_ag_weight is not None - and self.k3_workspace is not None - and not is_in_breakable_cuda_graph() - ): - view = self._k3_decode_view(attn_metadata, step.num_tokens) + """The built-in forward, except on a decode step whose cache the decode kernels read + (``will_run_decode_branch``). There only, ``project_output=False`` returns the gated attention output + ``[M, H * 128]``, the input of ``o_proj``.""" + view = None if latent_cache_gen is not None else self._k3_step_view(attn_metadata, step) if view is None: + if not project_output: + raise ValueError("project_output=False needs a step will_run_decode_branch takes") return super().forward( position_ids, hidden_states, attn_metadata, all_reduce_params, latent_cache_gen ) @@ -2018,6 +2070,8 @@ def forward( gate=ag, gate_col0=rows, ) + if not project_output: + return attn_output out = None if gemvs is None else gemvs.project("o_proj", attn_output, self.o_proj.weight) if out is None: out = self._project_output( From 10f4a6ed3272368fd11fe6f5c7a6ca7064d79525 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 21:57:37 -0700 Subject: [PATCH 082/161] [None][feat] modeling_v2 Kimi K3 tp16_moetp16ep1: tp16_moetp4ep4's modules, copied Targets share no files, so the tp16_moetp16ep1 target ships a copy of the tp16_moetp4ep4 target's modeling.py, weights.py and decode_gemv.py. The copy differs only inside blocks marked "# >>> route B: " ... "# <<< route B": * the module docstrings: experts split 16 x 1, no speculative decoding; * the construction checks: topology (16, 16, 1, 16, 1, no attention DP), the expert split set explicitly (left unset, Kimi K3 runs the experts expert-parallel over the 16 ranks), and no speculative decoding config (the message names the 4 x 4 split for DSpark, DFlash and SA); * the registered name and class; * the speculative path: no SA / DFlash / DSpark admission, no drafter LM head hand-off, no per-token KDA verify states (kda_token_states) and no KDA verify kernels (k3_kda_attn, k3_kda_verify, trtllm::kda_mtp_decode), with REQUIRED_TRTLLM_OPS and UNCERTIFIED_GENERIC_CALLS to match. decode_gemv.py is byte-identical. The causal LM stays the stock one-engine shell the built-in model builds on; without a speculative config it builds no drafter or worker, and the text model's hidden-state taps never fire. test_modeling_v2_kimi_k3_drift.py reads both targets' files, imports neither, and fails on any difference outside the blocks; a control checks that it flags changes outside blocks only. test_modeling_v2_kimi_k3_construction.py runs both targets' construction checks on the host. Signed-off-by: Vasanth Sabavat --- .../decode_gemv.py | 411 +++ .../modeling.py | 2199 ++++++++++++++++- .../weights.py | 2 + .../test_modeling_v2_kimi_k3_construction.py | 72 + .../test_modeling_v2_kimi_k3_drift.py | 135 + 5 files changed, 2753 insertions(+), 66 deletions(-) create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_gemv.py create mode 100644 tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_construction.py create mode 100644 tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drift.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_gemv.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_gemv.py new file mode 100644 index 000000000000..b384010b65dd --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_gemv.py @@ -0,0 +1,411 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""The decode path's GEMVs, LM head and embedding on the catalog's single-GPU Kimi K3 entries. + +* **Per-site GEMVs** (`K3DecodeGemvs.project`): a projection of a decode step runs on the kernel measured fastest + at its call site's weight shape (`SITES`): at most `MAX_ROWS` rows on `gemm/k3_decode_gemv`, + `gemm/k3_ctm_gemv_wide` or `gemm/k3_ctm_gemv_long`, and, where the site lists it, up to `WIDE_ROWS` rows on + `gemm/k3_ctm_gemv_wide`. The sites are the decode kernels' fused projections (MLA's [W_a; W_g] with the gate rows + through a sigmoid, KDA's [q | k | v | g | f_a | b]), the attention output projection, and the built-in MLA path's + q_a / kv_a, q_b and gate projections. +* **LM head** (`K3LogitsProcessor`): at most `MAX_ROWS` rows of this rank's vocabulary shard on + `gemm/k3_head_gemv` over the target's `K3HeadGemvWorkspace`, then the shards gathered (`comm/allgather`) as the + stock head gathers them. It is the shell's logits processor, so the speculative worker's target logits and the + drafter's logits on the same head take it too. +* **Embedding** (`K3DecodeGemvs.embed_norm`): a decode step's embedding rows written into the attention-residual + bank's slot 0 (layer 0's first snapshot) and layer 0's input RMSNorm applied, in one `norm/k3_embed_norm` launch. +* **Dense MLP** (`K3DecodeGemvs.dense_mlp`): layer 0's MLP at most `MAX_ROWS` rows, split over the whole TP group: + gate_up on `gemm/k3_ctm_gemv_long`, `activation/k3_situ_mul`, down on `gemm/k3_ctm_gemv_long`; the caller then + runs the down projection's all-reduce. + +Each returns None where its kernel does not take the call, and the caller then runs the generic path's module. + +A kernel compiles on its first call for a shape, which must not happen under CUDA-graph capture. +`K3DecodeGemvs.create`, run once the weights are final, runs every site's kernel and the head once, eagerly. The +embedding kernel compiles per token count, on the eager warm-up step before each capture. Under capture, a call +whose kernel has not run eagerly is refused. + +The head's workspace serves every `k3_head_gemv` call of its weight shape, so those calls must be ordered on one +stream: the logits are computed on the model's stream (see the `gemm/k3_head_gemv` contract's State section). +""" + +from __future__ import annotations + +import math +from dataclasses import dataclass +from typing import Dict, Iterable, Optional, Set + +import torch +from torch import nn + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.activation.k3_situ_mul import k3_situ_mul +from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.allgather import allgather +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv_long import ( + k3_ctm_gemv_long, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv_wide import ( + k3_ctm_gemv_wide, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_decode_gemv import k3_decode_gemv +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_head_gemv import ( + K3HeadGemvWorkspace, + k3_head_gemv, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.norm.k3_embed_norm import k3_embed_norm +from tensorrt_llm._torch._experimental.modeling_v2.catalog.torch.concat import concat +from tensorrt_llm._torch._experimental.modeling_v2.catalog.torch.split import split + +# The kernels' support predicates: metadata reads only. +from tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv import op as _ctm_op +from tensorrt_llm._torch.cute_dsl_kernels.k3_decode_gemv import op as _decode_op +from tensorrt_llm._torch.cute_dsl_kernels.k3_embed import op as _embed_op +from tensorrt_llm._torch.cute_dsl_kernels.k3_head_gemv import op as _head_op +from tensorrt_llm._torch.flashinfer_utils import IS_FLASHINFER_AVAILABLE +from tensorrt_llm._utils import mpi_disabled + +# The row limit of the decode GEMV and head kernels (one token tile), and of k3_ctm_gemv_wide (a decode step of 8 +# requests of 8 tokens). +MAX_ROWS = 8 +WIDE_ROWS = 64 + + +@dataclass(frozen=True) +class Site: + """A call site's weight shape (this target's per-rank shapes) and its kernels: ``small`` at 1..MAX_ROWS rows + ("decode", "wide" or "long"), and k3_ctm_gemv_wide at MAX_ROWS+1..WIDE_ROWS rows where ``wide``. Output columns + from ``sig_col0`` on are stored through a sigmoid. ``split`` / ``ring`` / ``push``: k3_ctm_gemv_long's CTAs per + 128-row weight tile, weight-ring stages, and whether the partial sums are pushed to each row's owner.""" + + n: int + k: int + small: str + wide: bool = False + sig_col0: int = -1 + split: int = 0 + ring: int = 0 + push: bool = False + + +SITES: Dict[str, Site] = { + # MLA's [W_a; W_g] on a decode step: [q_a 1536 | kv_a 512 | k_pe 64] then the output gate (6 heads x 128), + # the gate rows through a sigmoid. + "mla_ag": Site(2880, 7168, "long", wide=True, sig_col0=2112, split=6, ring=6, push=True), + # KDA's [q | k | v | g | f_a | b] on a decode step (6 heads x 128 each, then 128 and 6, padded to 3208 rows). + "kda_proj": Site(3208, 7168, "long", wide=True, split=5, ring=6, push=True), + # The attention output projection (row parallel; 6 heads x 128 in). + "o_proj": Site(7168, 768, "decode", wide=True), + # The built-in MLA path's projections: kv_a_proj_with_mqa, q_b_proj and the output gate. + "kv_a": Site(2112, 7168, "decode"), + "q_b": Site(1152, 1536, "wide"), + "g_proj": Site(768, 7168, "wide"), + # Layer 0's dense MLP split over the 16-way TP group: gate_up [gate 2112 | up 2112] and down. + "dense_gate_up": Site(4224, 7168, "long", split=4, ring=5), + "dense_down": Site(7168, 2112, "long", split=2, ring=6), +} + + +def _wide_tile(rows: int) -> int: + from tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv import k3_ctm_gemv_kernel + + return k3_ctm_gemv_kernel.wide_tile(rows) + + +def _run( + spec: Site, kernel: str, x2d: torch.Tensor, weight: torch.Tensor +) -> Optional[torch.Tensor]: + """``spec``'s ``kernel`` on dense rows ``x2d``, or None where it does not take them.""" + if kernel == "decode": + if spec.sig_col0 >= 0 or not _decode_op.supports(x2d, weight): + return None + return k3_decode_gemv(x2d, weight) + if kernel == "wide": + if not _ctm_op.supports_wide(x2d, weight, spec.sig_col0, False): + return None + return k3_ctm_gemv_wide(x2d, weight, sig_col0=spec.sig_col0) + # One wave of the GPU's SMs: beyond it the long GEMV loses to the others. + sms = torch.cuda.get_device_properties(x2d.device).multi_processor_count + if math.ceil(spec.n / 128) * spec.split > sms or not _ctm_op.supports_long( + x2d, weight, spec.split, spec.ring + ): + return None + return k3_ctm_gemv_long( + x2d, + weight, + sig_col0=spec.sig_col0, + split=spec.split, + ring=spec.ring, + trigger_early=True, + push=spec.push, + ) + + +def _capturing() -> bool: + return not torch.compiler.is_compiling() and torch.cuda.is_current_stream_capturing() + + +def _dense_rows(x2d: torch.Tensor) -> torch.Tensor: + """``x2d`` itself, or a dense copy: the kernels' TMA descriptors need dense, 16-byte-aligned rows.""" + if x2d.stride() != (x2d.shape[1], 1) or x2d.data_ptr() % 16: + return x2d.clone(memory_format=torch.contiguous_format) + return x2d + + +def _head_takes_module(lm_head: nn.Module) -> bool: + """Whether ``lm_head(rows)`` is a plain vocabulary-parallel GEMM of a bf16 weight whose shards it gathers along + the vocabulary in rank order, with nothing else applied: the stock head this path reproduces.""" + weight = getattr(lm_head, "weight", None) + mapping = getattr(lm_head, "mapping", None) + return ( + isinstance(weight, torch.Tensor) + and weight.dim() == 2 + and weight.dtype == torch.bfloat16 + and weight.is_contiguous() + and getattr(getattr(lm_head, "tp_mode", None), "name", None) == "COLUMN" + and getattr(lm_head, "gather_output", False) + and getattr(lm_head, "gather_output_sizes", None) is None + and getattr(lm_head, "padding_size", None) == 0 + and getattr(lm_head, "bias", None) is None + and not getattr(lm_head, "has_any_quant", True) + and mapping is not None + and not mapping.enable_attention_dp + ) + + +def _plain_rmsnorm(norm: nn.Module) -> bool: + """Whether ``norm`` is the stock RMSNorm that runs flashinfer's kernel, which ``k3_embed_norm`` reproduces bit for + bit. An unknown module fails closed.""" + weight = getattr(norm, "weight", None) + return ( + IS_FLASHINFER_AVAILABLE + and isinstance(weight, torch.Tensor) + and weight.dtype == torch.bfloat16 + and not getattr(norm, "use_gemma", True) + and not getattr(norm, "is_nvfp4", True) + and not getattr(norm, "use_cuda_tile", True) + and not getattr(norm, "return_hp_output", True) + and getattr(norm, "nvfp4_scale", None) is None + and hasattr(norm, "variance_epsilon") + ) + + +class K3DecodeGemvs: + """The decode GEMVs' state for one target: the LM head's `K3HeadGemvWorkspace`, and the calls whose kernels ran + eagerly (compiled), which are the only ones a CUDA-graph capture may take. Built by `create` once the weights are + final; owned by the target.""" + + def __init__(self, head_workspace: Optional[K3HeadGemvWorkspace] = None) -> None: + self.head_workspace = head_workspace + self._ran: Set[tuple] = set() + + @classmethod + def create( + cls, + lm_head: Optional[nn.Module] = None, + sites: Iterable[str] = tuple(SITES), + device: Optional[torch.device] = None, + ) -> "K3DecodeGemvs": + """The state for ``lm_head`` and ``sites`` on ``device`` (default: the head's, else the current one). Eager: + it allocates the head's workspace and runs every kernel of every site once (each wide row class once) on + zero rows of a zero weight of the site's shape, so they compile here and not under a capture. A site or + head whose kernel does not take its shape keeps the generic path.""" + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "K3DecodeGemvs.create allocates and compiles: run it before CUDA-graph capture" + ) + head_weight = getattr(lm_head, "weight", None) if lm_head is not None else None + if device is None: + device = ( + head_weight.device + if isinstance(head_weight, torch.Tensor) + else torch.device("cuda", torch.cuda.current_device()) + ) + state = cls() + for site in sites: + spec = SITES[site] + weight = torch.zeros(spec.n, spec.k, dtype=torch.bfloat16, device=device) + for rows in (1, 16, 32, 64) if spec.wide else (1,): + state._project(site, weight.new_zeros(rows, spec.k), weight, warm=True) + del weight + if "dense_gate_up" in sites: + gu = torch.zeros(1, SITES["dense_gate_up"].n, dtype=torch.bfloat16, device=device) + if _ctm_op.supports_situ_mul(gu): + for linear_beta in (None, 1.0): + k3_situ_mul(gu, 1.0, linear_beta) + state._ran.add(("situ_mul", linear_beta is not None)) + if lm_head is not None and _head_takes_module(lm_head): + x = head_weight.new_zeros(1, head_weight.shape[1]) + if _head_op.supports(x, head_weight): + workspace = K3HeadGemvWorkspace.create( + head_weight.shape[0], head_weight.shape[1], head_weight.device + ) + k3_head_gemv(x, head_weight, workspace) + state.head_workspace = workspace + state._ran.add(("lm_head",)) + torch.cuda.synchronize(device) + return state + + def project(self, site: str, x: torch.Tensor, weight: torch.Tensor) -> Optional[torch.Tensor]: + """``x @ weight.T`` (bf16 ``[..., N]``, the site's sigmoid columns through the sigmoid) for ``site``'s weight + on its decode kernel, or None where none takes the call: more rows than the site's kernels take, another + shape or dtype, or, under capture, a kernel that has not run eagerly. The caller then runs its GEMM.""" + return self._project(site, x, weight, warm=False) + + def _project( + self, site: str, x: torch.Tensor, weight: torch.Tensor, warm: bool + ) -> Optional[torch.Tensor]: + spec = SITES[site] + if ( + weight.dtype != torch.bfloat16 + or tuple(weight.shape) != (spec.n, spec.k) + or not weight.is_contiguous() + or x.dtype != torch.bfloat16 + or x.dim() < 1 + or x.shape[-1] != spec.k + ): + return None + rows = x.numel() // spec.k + if 0 < rows <= MAX_ROWS: + kernel = spec.small + elif spec.wide and MAX_ROWS < rows <= WIDE_ROWS: + kernel = "wide" + else: + return None + key = (site, kernel, _wide_tile(rows) if kernel == "wide" else 0) + capturing = _capturing() + if capturing and not warm and key not in self._ran: + return None + y = _run(spec, kernel, _dense_rows(x.reshape(rows, spec.k)), weight) + if y is None: + return None + if not capturing: + self._ran.add(key) + return y.view(*x.shape[:-1], spec.n) + + def lm_head_logits(self, rows: torch.Tensor, lm_head: nn.Module) -> Optional[torch.Tensor]: + """``lm_head(rows)``, the gathered bf16 logits ``[M, vocab]``, with this rank's shard on + ``gemm/k3_head_gemv``; None where it does not take the call (more than `MAX_ROWS` rows, another head, ...).""" + workspace = self.head_workspace + if workspace is None or not _head_takes_module(lm_head): + return None + weight = lm_head.weight + if ( + tuple(weight.shape) != (workspace.n_out, workspace.k_in) + or weight.device != workspace.partials.device + or rows.dim() != 2 + or rows.dtype != torch.bfloat16 + or not 0 < rows.shape[0] <= MAX_ROWS + or rows.shape[1] != workspace.k_in + ): + return None + if _capturing() and ("lm_head",) not in self._ran: + return None + group = lm_head.mapping.tp_group + if len(group) > 1 and mpi_disabled(): + return None + x = _dense_rows(rows) + if not _head_op.supports(x, weight): + return None + local = k3_head_gemv(x, weight, workspace) + if len(group) == 1: + return local + gathered = allgather(local, None, group) + return concat(list(split(gathered, rows.shape[0], dim=0)), dim=-1) + + def dense_mlp( + self, + x: torch.Tensor, + gate_up_weight: torch.Tensor, + down_weight: torch.Tensor, + beta: float, + linear_beta: Optional[float], + ) -> Optional[torch.Tensor]: + """Layer 0's dense MLP of at most `MAX_ROWS` rows, before its down projection's all-reduce: gate_up on + ``gemm/k3_ctm_gemv_long``, ``SituAndMul(beta, linear_beta)`` on ``activation/k3_situ_mul``, down on + ``gemm/k3_ctm_gemv_long``. None, with nothing launched, where a kernel does not take the call; the caller then + runs the module.""" + gate_up, down = SITES["dense_gate_up"], SITES["dense_down"] + rows = x.shape[0] if x.dim() == 2 else 0 + if ( + not 0 < rows <= MAX_ROWS + or linear_beta == 0.0 + or tuple(gate_up_weight.shape) != (gate_up.n, gate_up.k) + or tuple(down_weight.shape) != (down.n, down.k) + or gate_up.n != 2 * down.k + ): + return None + if ( + _capturing() + and not { + ("dense_gate_up", "long", 0), + ("dense_down", "long", 0), + ("situ_mul", linear_beta is not None), + } + <= self._ran + ): + return None + gu = self.project("dense_gate_up", x, gate_up_weight) + if gu is None or not _ctm_op.supports_situ_mul(gu): + return None + return self.project("dense_down", k3_situ_mul(gu, beta, linear_beta), down_weight) + + def embed_norm( + self, + input_ids: torch.Tensor, + table: torch.Tensor, + norm: nn.Module, + bank: torch.Tensor, + ) -> Optional[torch.Tensor]: + """Layer 0's normed input for the step's tokens, with their embedding rows written into ``bank[0]``, in one + ``norm/k3_embed_norm`` launch: bit-identical to the embedding followed by ``norm``. None where it does not + apply (another norm, a token count or table the kernel does not take, or, under capture, a token count whose + kernel has not run eagerly).""" + if not _plain_rmsnorm(norm): + return None + ids = input_ids.reshape(-1) + if not _embed_op.supports_norm(ids, table, norm.weight, bank[0]): + return None + key = ("embed_norm", ids.numel(), ids.dtype) + capturing = _capturing() + if capturing and key not in self._ran: + return None + normed = k3_embed_norm(ids, table, norm.weight, norm.variance_epsilon, bank[0]) + if not capturing: + self._ran.add(key) + return normed + + +class _K3Head: + """``lm_head(rows)``, with the rows on ``k3_head_gemv`` where it takes them.""" + + __slots__ = ("_gemvs", "_lm_head") + + def __init__(self, gemvs: K3DecodeGemvs, lm_head: nn.Module) -> None: + self._gemvs = gemvs + self._lm_head = lm_head + + def __call__(self, rows: torch.Tensor) -> torch.Tensor: + logits = self._gemvs.lm_head_logits(rows, self._lm_head) + return self._lm_head(rows) if logits is None else logits + + +class K3LogitsProcessor(nn.Module): + """The shell's logits processor, with its LM head call on ``gemm/k3_head_gemv`` where ``gemvs`` takes the rows. + + It wraps the stock processor instead of repeating it: the row selection and the fp32 conversion stay the stock + processor's. ``gemvs`` is set once the weights are final; until then every call is the stock one. + """ + + def __init__(self, stock: nn.Module) -> None: + super().__init__() + self.stock = stock + self.gemvs: Optional[K3DecodeGemvs] = None + + def forward( + self, + hidden_states: torch.Tensor, + lm_head: nn.Module, + attn_metadata, + return_context_logits: bool = False, + ) -> torch.Tensor: + head = lm_head if self.gemvs is None else _K3Head(self.gemvs, lm_head) + return self.stock.forward(hidden_states, head, attn_metadata, return_context_logits) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py index 3c86aabb0b0e..4c40d44b28d3 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py @@ -1,5 +1,6 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 +# >>> route B: this target's layout, and decoding without speculation """ModelingV2 target: Kimi K3 (MXFP4) / sm_100 / tp16 attention, routed experts moe_tp 16 x moe_ep 1. Kimi K3's language model: 93 layers, hidden 7168. Every fourth layer (3, 7, ..., 91) is MLA attention with 96 @@ -11,54 +12,129 @@ `tp16_moetp16ep1` is `tensor_parallel_size: 16` with `moe_tensor_parallel_size: 16` and `moe_expert_parallel_size: 1`, both set explicitly, no attention data parallelism, all 16 GPUs in one NVLink domain. Attention is head-split (6 MLA query heads per rank), and every rank holds all 896 experts at a sixteenth of their -width (192 of 3072 intermediate values, zero-padded to 256 by the loader). With the expert split left unset the -built-in model runs the experts expert-parallel over the 16 ranks instead, and no target serves that layout. - -**Each step takes one of two paths, chosen on the host from the step's shape** (`_step_path`): - -* **The fused decode path**: pure decode steps of at most 8 tokens, on the K3 decode kernels' catalog entries. The - routed experts of one and two tokens run as `moe/k3_moe_m1` and `moe/k3_moe_m2`, and of 3-8 tokens as k3_moe, each - pushing its partial into the latent exchange that `trtllm::k3_latent_reduce` sums. The state those kernels share - (MNNVL workspace, sandwich and MoE Lamport buffers, the latent exchange, the engines' workspaces, KDA / MLA scratch) - lives in typed objects this target creates in `post_load_weights`, before any graph capture. Until those entries - are wired `_fused_decode` stays None, and every step takes the generic path. -* **The generic path**: prefill, mixed steps, and decode steps above those bounds, on the built-in Kimi K3 text - model, whose modules and ops have no catalog entries yet. `UNCERTIFIED_GENERIC_CALLS` names them. - -**What this target asserts rather than adapts**: SM 10.0; the topology above, with its expert split explicit; no -speculative decoding; bf16 weights and a bf16 KV pool; -tokens_per_block 64 (the MLA generation kernels K3's 96 heads reach exist only at 64); the V2 hybrid KV / state -manager, which holds the KDA states, with block reuse off; an all-reduce strategy of AUTO or MNNVL. The -construction-time ones fail in `__init__`, the per-engine ones on the first forward, each naming the setting. +intermediate width (192 of 3072 values, zero-padded to 256 when they load). With the expert split left unset the +engine resolves the same sizes, but Kimi K3 then runs the experts expert-parallel over the 16 ranks, a layout no +target serves. This is Kimi K3's layout for decoding without speculation; `tp16_moetp4ep4` serves speculative +decoding. + +**This module is `tp16_moetp4ep4`'s, copied** (targets share no files) and changed only inside its marked route B +blocks; test_modeling_v2_kimi_k3_drift.py checks that the two modules match everywhere else. + +**Each step is classified once, on the host, from its shape** (`decode_step`, a `DecodeStep` or None), and the +classification decides which kernels each module runs: + +* **small**: at most 8 tokens (`DECODE_MAX_TOKENS`, one token tile), context requests included. The token-count + kernels take it: decode GEMVs, the MoE front and routed experts, the sandwiches, the embedding and residual + epilogues. +* **decode**: a pure decode step of R <= 8 generation requests and no context request; without speculation each + request has one token. The request-aware kernels take it: MLA attention and its KV store. + +Every other step (prefill, mixed steps, decode steps above those bounds) runs the **generic path**: this target's +text model (`KimiLinearModel` below: decoder layers, attention residuals, the MLA / KDA / MoE runtimes), computed +exactly as the built-in Kimi K3 text model computes it, on stock modules and ops that have no catalog entries yet. +`UNCERTIFIED_GENERIC_CALLS` names them. + +The text model hands each step's classification to its attention modules, which run a **decode step** on the K3 +decode kernels' catalog entries: + +* KDA (`K3DecodeKDA`): one token per request, the fused input projection and the plain decode in one + `ssm/k3_kda_decode_attn` launch. +* MLA (`K3DecodeMLA`): `attention/k3_mla_qkv` (the query path and the step's latent KV rows into the paged cache), + then `attention/k3_mla_attn_vb_out` (the attention, v_b and the output gate in one launch). + +The projections around them run on the decode GEMV sites of `decode_gemv.py` (the [W_a; W_g] projection, `o_proj` +on every classified step), as do the LM head, the embedding and layer 0's dense MLP. The state those kernels share +(the KDA projection's Lamport buffers, the MLA attention workspace, the decode GEMVs' state) lives in typed objects +this target creates in `post_load_weights`, before any graph capture. The MoE front and routed experts, the +sandwiches and the residual epilogues come with their own entries; until then they run the generic path on every +step. + +**What this target asserts rather than adapts**: SM 10.0; the topology above, with the expert split set explicitly; +no speculative decoding; bf16 weights and a bf16 KV pool; tokens_per_block 64 (the MLA generation kernels K3's 96 +heads reach exist only at 64); the V2 hybrid KV / state manager, which holds the KDA states, with block reuse off and +fp32 recurrent states; an all-reduce strategy of AUTO or MNNVL. The construction-time ones fail in `__init__`, the +per-engine ones on the first forward, each naming the setting. A layer the decode kernels do not take fails the +weight load. **Text only.** The checkpoint is the vision-language wrapper. This target builds and loads no vision tower (its weights are a predicted non-load, `weights.py`), and a step carrying multimodal input raises. -**No speculative decoding.** This layout is the low-latency route for decoding without a drafter; DSpark runs on -the `tp16_moetp4ep4` target. A configuration with a speculative decoding config stops at construction. +**No speculative decoding.** A speculative decoding config fails construction, naming `tp16_moetp4ep4`'s expert +split. The causal LM is still the stock one-engine shell the built-in model builds on, which without that config +builds no drafter and no worker. This target drops `tp16_moetp4ep4`'s KDA verify kernels and the per-token verify +states they keep; the text model's speculative hidden-state taps stay as `tp16_moetp4ep4` has them and never fire. """ +# <<< route B + +from __future__ import annotations import copy -from typing import Any, Literal, Optional +import math +import os +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any, Callable, Dict, Literal, NamedTuple, Optional, Tuple import torch +from torch import nn + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.k3_mla_attn_vb_out import ( + k3_mla_attn_vb_out, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.k3_mla_attn_workspace import ( + K3MlaAttnWorkspace, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.k3_mla_qkv import k3_mla_qkv -from tensorrt_llm._torch.attention.backends.interface import AttentionMetadata +# >>> route B: no KDA verify kernels, and no trtllm::kda_mtp_decode (the built-in KDA verify's) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.ssm.k3_kda_buffers import K3KdaBuffers +from tensorrt_llm._torch._experimental.modeling_v2.catalog.ssm.k3_kda_decode_attn import ( + k3_kda_decode_attn, +) +from tensorrt_llm._torch.attention.backends import AttentionMetadata +from tensorrt_llm._torch.attention.backends.fmha.cute_dsl_mla import k3_mla_decode_view + +# <<< route B +from tensorrt_llm._torch.distributed import AllReduce from tensorrt_llm._torch.model_config import ModelConfig from tensorrt_llm._torch.models.modeling_kimi_linear import KimiLinearForCausalLM -from tensorrt_llm._torch.models.modeling_utils import register_auto_model +from tensorrt_llm._torch.models.modeling_speculative import SpecDecOneEngineForCausalLM +from tensorrt_llm._torch.models.modeling_utils import DecoderModel, register_auto_model +from tensorrt_llm._torch.modules.gated_mlp import GatedMLP +from tensorrt_llm._torch.modules.kimi_k3_mla import KimiK3MLAAttention +from tensorrt_llm._torch.modules.kimi_kda import KimiKDALinearAttention +from tensorrt_llm._torch.modules.multi_stream_utils import maybe_execute_in_parallel +from tensorrt_llm._torch.modules.rms_norm import RMSNorm +from tensorrt_llm._torch.modules.situ import SituAndMul +from tensorrt_llm._torch.moe.fused_moe import ( + ConfigurableMoE, + SiTuActivation, + TRTLLMGenFusedMoE, + create_moe, +) +from tensorrt_llm._torch.moe.fused_moe.interface import MoESchedulerKind +from tensorrt_llm._torch.moe.fused_moe.routing import DeepSeekV3MoeRoutingMethod +from tensorrt_llm._torch.pyexecutor.breakable_cuda_graph import is_in_breakable_cuda_graph from tensorrt_llm._torch.pyexecutor.kv_cache.mamba_cache_manager import MambaHybridCacheManagerV2 +from tensorrt_llm._torch.utils import AuxStreamType from tensorrt_llm.functional import AllReduceStrategy +from tensorrt_llm.logger import logger +from tensorrt_llm.mapping import Mapping +from tensorrt_llm.models.modeling_utils import QuantAlgo, QuantConfig +from . import decode_gemv as _decode_gemv from . import weights as _weights +if TYPE_CHECKING: + from transformers import PretrainedConfig + + # The GPU architecture this target IS. Routing will not send another one here, but a direct instantiation could, # and the certification is per arch. _SM = (10, 0) #: Every trtllm op this target reaches for, in its forward and in the weight load. Declared here, asserted in -#: tests/unittest/_torch/modeling_v2. Today these are the K3-specific ops of the generic path (attention residuals, -#: KDA, the router and fused-A GEMMs); the fused decode path adds its own. +#: tests/unittest/_torch/modeling_v2: the K3-specific ops of the generic path (attention residuals, KDA, the router +#: and fused-A GEMMs), then the decode kernels'. REQUIRED_TRTLLM_OPS = ( "attn_res_fwd", "attn_res_rmsnorm_fwd", @@ -66,9 +142,19 @@ "attn_res_add_rmsnorm_persistent_fwd", "kda_prefill", "kda_decode", - "kda_mtp_decode", + # >>> route B: no KDA verify (kda_mtp_decode, k3_kda_attn, k3_kda_verify) "dsv3_router_gemm_op", "dsv3_fused_a_gemm_op", + "k3_kda_decode_attn", + # <<< route B + "k3_mla_qkv", + "k3_mla_attn_vb_out", + # The decode path's GEMVs, LM head and embedding (decode_gemv.py). + "k3_decode_gemv", + "k3_ctm_gemv_wide", + "k3_head_gemv", + "k3_embed_norm", + "allgather", ) #: The engine surface the first forward checks before this target relies on it: per object, the attributes read. @@ -76,23 +162,46 @@ REQUIRED_ENGINE_FIELDS = { "attn_metadata": ( "num_contexts", + "num_generations", "num_seqs", "num_tokens", + "seq_lens", "tokens_per_block", "kv_cache_manager", + "mamba_metadata", ), - "kv_cache_manager": ("enable_block_reuse",), + "kv_cache_manager": ("enable_block_reuse", "mamba_layer_cache"), } -#: Calls the generic path makes outside the catalog, declared so they are not consumed silently. A call leaves this -#: list when a catalog entry replaces it. +#: Stock code the generic path runs outside the catalog, declared so it is not consumed silently: every +#: tensorrt_llm import of this module that computes (test_modeling_v2_claims.py checks the list both ways). An entry +#: leaves the list when a catalog entry replaces it. UNCERTIFIED_GENERIC_CALLS = ( + # The checkpoint load and the engine hooks this target inherits, and the causal LM around the text model. "tensorrt_llm._torch.models.modeling_kimi_linear.KimiLinearForCausalLM", + "tensorrt_llm._torch.models.modeling_speculative.SpecDecOneEngineForCausalLM", + "tensorrt_llm._torch.models.modeling_utils.DecoderModel", + # The text model's stock modules. + "tensorrt_llm._torch.modules.kimi_kda.KimiKDALinearAttention", + # >>> route B: no trtllm::kda_mtp_decode registration (the built-in KDA verify's) + # <<< route B + "tensorrt_llm._torch.modules.kimi_k3_mla.KimiK3MLAAttention", + "tensorrt_llm._torch.moe.fused_moe.create_moe", + "tensorrt_llm._torch.moe.fused_moe.ConfigurableMoE", + "tensorrt_llm._torch.moe.fused_moe.TRTLLMGenFusedMoE", + "tensorrt_llm._torch.moe.fused_moe.routing.DeepSeekV3MoeRoutingMethod", + "tensorrt_llm._torch.modules.gated_mlp.GatedMLP", + "tensorrt_llm._torch.modules.situ.SituAndMul", + "tensorrt_llm._torch.modules.rms_norm.RMSNorm", + "tensorrt_llm._torch.distributed.AllReduce", + "tensorrt_llm._torch.modules.multi_stream_utils.maybe_execute_in_parallel", ) -# The fused decode path's token bound per step: one token per request (the K3 decode kernels are built for up to 8 -# rows). -_FUSED_MAX_TOKENS = 8 +# The K3 decode kernels' bounds: the token-count kernels take one tile of DECODE_MAX_TOKENS rows; the request-aware +# kernels take MAX_REQUESTS generation requests of at most MAX_TOKENS_PER_REQUEST tokens (1 + 7 drafts with DSpark). +DECODE_MAX_TOKENS = 8 +MAX_REQUESTS = 8 +MAX_TOKENS_PER_REQUEST = 8 # The MLA generation kernels for K3's 96 query heads exist only at a 64-token page (the built-in model's own # get_model_defaults sets it for the same reason). @@ -101,6 +210,1853 @@ _LANG_PREFIX = "language_model." +# ---------------------------------------------------------------------------------------------------------------------- +# The text model: Kimi K3's decoder (93 layers: KDA / MLA attention, attention residuals, the dense layer-0 MLP and +# the latent MoE), its generic path. Each step's classification goes to the attention modules (`K3DecodeKDA`, +# `K3DecodeMLA` below). +# ---------------------------------------------------------------------------------------------------------------------- + +# A/B escape hatch: restore nn.Linear for the K3 latent MoE projections +# instead of the min-latency fused GEMM op (read once at import). +_K3_DISABLE_MIN_LATENCY_LATENT_PROJ = ( + os.environ.get("TLLM_K3_DISABLE_MIN_LATENCY_LATENT_PROJ", "0") == "1" +) + + +# Identity-RoPE table positions for the MLA backends. K3 is NoPE (the table +# holds cos=1/sin=0), but the chunked-context path indexes the table by +# absolute position, so it must cover max_position_embeddings (~512MB per +# backend for the 1M-position checkpoint); a smaller table is read out of +# bounds. KIMI_K3_MLA_MAX_POSITIONS overrides the size for short-context +# deployments. +_KIMI_K3_MLA_MAX_POSITIONS_ENV = "KIMI_K3_MLA_MAX_POSITIONS" + + +class KimiK3MoEGate(nn.Module): + """Kimi K3 gate weights and routing method for ``ConfigurableMoE``.""" + + def __init__( + self, + config: Any, + *, + logits_gemm_dtype: torch.dtype | None = None, + device: torch.device | None = None, + ) -> None: + super().__init__() + self.config = config + self.top_k = config.num_experts_per_token + self.num_experts = config.num_experts + self.routed_scaling_factor = config.routed_scaling_factor + self.moe_router_activation_func = config.moe_router_activation_func + self.num_expert_group = getattr(config, "num_expert_group", 1) + self.topk_group = getattr(config, "topk_group", 1) + self.moe_renormalize = config.moe_renormalize + self.gating_dim = config.hidden_size + + assert self.moe_router_activation_func in ("sigmoid", "softmax"), ( + "K3 MoE gate supports sigmoid or softmax scoring only" + ) + + # The checkpoint stores the gate weight in bf16. Storing it in bf16 + # permits the single bf16xbf16 router GEMM while retaining fp32 output. + weight_dtype = logits_gemm_dtype or torch.float32 + self.weight = nn.Parameter( + torch.empty((self.num_experts, self.gating_dim), dtype=weight_dtype, device=device) + ) + self.e_score_correction_bias = nn.Parameter( + torch.empty(self.num_experts, dtype=torch.float32, device=device) + ) + + def compute_logits(self, hidden_states: torch.Tensor) -> torch.Tensor: + """Compute fp32 routing logits shaped ``[num_tokens, num_experts]``.""" + hidden_2d = hidden_states.reshape(-1, self.gating_dim) + if self.weight.dtype == torch.bfloat16 and hidden_2d.dtype == torch.bfloat16: + return torch.ops.trtllm.dsv3_router_gemm_op( + hidden_2d.contiguous(), + self.weight.t(), + bias=None, + out_dtype=torch.float32, + ) + return torch.nn.functional.linear( + hidden_2d.type(torch.float32), + self.weight.type(torch.float32), + None, + ) + + @property + def routing_method(self) -> DeepSeekV3MoeRoutingMethod: + """Return the shared DeepSeek-V3 router used by ``ConfigurableMoE``.""" + if self.moe_router_activation_func != "sigmoid": + raise ValueError("Kimi K3 ConfigurableMoE routing requires sigmoid scores.") + if not self.moe_renormalize: + raise ValueError( + "Kimi K3 ConfigurableMoE routing requires top-k weight renormalization." + ) + return DeepSeekV3MoeRoutingMethod( + top_k=self.top_k, + n_group=self.num_expert_group, + topk_group=self.topk_group, + routed_scaling_factor=self.routed_scaling_factor, + callable_e_score_correction_bias=lambda: self.e_score_correction_bias, + is_fused=True, + ) + + +class KimiK3RMSNorm(nn.Module): + """RMSNorm matching the Kimi checkpoint implementation's rounding.""" + + def __init__( + self, + hidden_size: int, + eps: float = 1e-6, + dtype: torch.dtype = torch.float32, + device: Optional[torch.device] = None, + ) -> None: + super().__init__() + self.weight = nn.Parameter(torch.ones(hidden_size, dtype=dtype, device=device)) + self.eps = eps + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + input_dtype = hidden_states.dtype + hidden_states_float = hidden_states.to(torch.float32) + variance = hidden_states_float.pow(2).mean(-1, keepdim=True) + hidden_states_float = hidden_states_float * torch.rsqrt(variance + self.eps) + return self.weight * hidden_states_float.to(input_dtype) + + +def _resolve_kimi_situ_betas(cfg: Any) -> tuple[float, float]: + """Return the finite SiTu betas required by the routed-expert kernels.""" + config_situ_beta = getattr(cfg, "activation_situ_beta", None) + situ_beta = 1.0 if config_situ_beta is None else config_situ_beta + situ_linear_beta = getattr(cfg, "activation_situ_linear_beta", None) + if situ_linear_beta is None: + raise ValueError( + "Kimi K3 routed SiTu experts require activation_situ_linear_beta; " + "None means an identity linear branch that the fused kernels cannot represent." + ) + if situ_beta <= 0 or situ_linear_beta <= 0: + raise ValueError( + f"Kimi K3 SiTu betas must be positive; got {situ_beta} and {situ_linear_beta}." + ) + return float(situ_beta), float(situ_linear_beta) + + +def _get_text_config(pretrained_config: "PretrainedConfig"): + """Return the Kimi text config, unwrapping a composite kimi_k3 config.""" + if getattr(pretrained_config, "model_type", None) == "kimi_k3" or ( + not hasattr(pretrained_config, "linear_attn_config") + and hasattr(pretrained_config, "text_config") + ): + return pretrained_config.text_config + return pretrained_config + + +def _is_kda_layer(cfg, layer_idx: int) -> bool: + return (layer_idx + 1) in cfg.linear_attn_config["kda_layers"] + + +def _is_mla_layer(cfg, layer_idx: int) -> bool: + return (layer_idx + 1) in cfg.linear_attn_config["full_attn_layers"] + + +KIMI_K3_AUX_ATTN_RES_STREAM_ENV = "KIMI_K3_AUX_ATTN_RES_STREAM" + + +_AUX_ATTN_RES_STREAM_ENABLED = os.environ.get(KIMI_K3_AUX_ATTN_RES_STREAM_ENV, "1") == "1" + + +KIMI_K3_FUSED_ATTN_RES_ENV = "KIMI_K3_FUSED_ATTN_RES" + + +_FUSED_ATTN_RES_ENABLED = os.environ.get(KIMI_K3_FUSED_ATTN_RES_ENV, "1") == "1" + + +KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS_ENV = "KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS" + + +KIMI_K3_ATTN_RES_TOPOLOGY_ENV = "KIMI_K3_ATTN_RES_TOPOLOGY" + + +_ATTN_RES_TOPOLOGIES = ("per_token", "persistent", "split") + + +# Turning the port on should not require knowing which of the three topologies +# is the right one: "1" selects the measured policy (``split``), so enabling the +# feature and choosing the policy are one action through one variable. +_ATTN_RES_TOPOLOGY_ON = "1" + + +def _read_attn_res_topology() -> str: + """Default ``per_token``: the persistent kernel is opt-in. + + ``1`` is the only accepted on-value and resolves to ``split``. The named + topologies stay available for measurement: ``persistent`` uses the + persistent kernel at every shape it implements, ``per_token`` at none. + """ + raw = os.environ.get(KIMI_K3_ATTN_RES_TOPOLOGY_ENV, "per_token") + if raw == _ATTN_RES_TOPOLOGY_ON: + return "split" + if raw not in _ATTN_RES_TOPOLOGIES: + # Loudly, for the same reason as the token ceiling below: a mistyped A/B + # arm that silently fell back to the default would measure one side + # twice and report no difference. + raise ValueError( + f"{KIMI_K3_ATTN_RES_TOPOLOGY_ENV} must be one of " + f"{_ATTN_RES_TOPOLOGIES} or {_ATTN_RES_TOPOLOGY_ON!r} " + f"(which means 'split'), got {raw!r}" + ) + return raw + + +def _read_fused_attn_res_max_tokens() -> int: + """Resolved after the topology, because its default follows it. + + With the persistent port off -- the default -- the ceiling is 1, the + pre-existing gate: the fused epilogue is taken at the single-token decode + shape and nowhere else. Enabling the port raises it to 32, the top of the + measured range, so that "off" keeps meaning "unchanged". + """ + default = "1" if _ATTN_RES_TOPOLOGY == "per_token" else "32" + raw = os.environ.get(KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS_ENV, default) + try: + value = int(raw) + except ValueError: + # Failing loudly matters more than usual here: a mistyped A/B arm that + # silently fell back to the default would measure the candidate twice + # and report no difference. + raise ValueError( + f"{KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS_ENV} must be a positive integer, got {raw!r}" + ) from None + if value < 1: + raise ValueError(f"{KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS_ENV} must be >= 1, got {value}") + return value + + +_ATTN_RES_TOPOLOGY = _read_attn_res_topology() + + +_FUSED_ATTN_RES_MAX_TOKENS = _read_fused_attn_res_max_tokens() + + +def _persistent_attn_res_applicable(M: int, H: int, N: int) -> bool: + """Shape gate for the persistent kernel: H == 7168 and 2 <= N <= 9. + + No token ceiling: the persistent grid is sized by the SM count, not by the + token count, so prefill is the case it exists for. + """ + del M # deliberately unused; see above + return H == 7168 and 2 <= N <= 9 + + +def _use_persistent_attn_res(M: int, H: int, N: int) -> bool: + """Pick between the two fused kernels for this call site. + + ``persistent`` takes the persistent kernel at every shape it implements; + ``split`` takes it only above ``KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS`` tokens, + which stands in for the prefill/decode boundary. Shapes the persistent + kernel does not implement fall through to the caller's existing gate and + land on the unfused path. + """ + if not _persistent_attn_res_applicable(M, H, N): + return False + if _ATTN_RES_TOPOLOGY == "persistent": + return True + return _ATTN_RES_TOPOLOGY == "split" and M > _FUSED_ATTN_RES_MAX_TOKENS + + +def _apply_attn_res_fused( + prefix_sum: torch.Tensor, block_residual: torch.Tensor, proj: nn.Linear, norm: KimiK3RMSNorm +) -> Optional[torch.Tensor]: + """Fused attn_res via the in-tree ``trtllm::attn_res_fwd`` op. + + Returns ``None`` when the call falls outside the fused kernel's + contract (dtype/shape/arch) so the caller can use the exact fp32 reference + instead. ``block_residual`` is kept in the kernel-native ``[K, M, H]`` + layout. Candidate order matches the reference: snapshots first, the + running prefix sum last. + """ + if ( + prefix_sum.dtype is not torch.bfloat16 + or not prefix_sum.is_cuda + or not block_residual.is_cuda + ): + return None + M, H = prefix_sum.shape + K = int(block_residual.shape[0]) + if K + 1 > 12 or M > 16384 or not (4096 <= H <= 8192 and H % 1024 == 0): + return None + try: + attn_res_op = torch.ops.trtllm.attn_res_fwd + except (AttributeError, RuntimeError): + return None + layer_kernel = prefix_sum.reshape(M, 1, H).contiguous() + block_kernel = block_residual.reshape(K, M, 1, H).contiguous() + output, _rsigma, _probs, _logits = attn_res_op( + layer_kernel, + block_kernel, + proj.weight.reshape(-1).to(torch.bfloat16).contiguous(), + norm.weight.to(torch.bfloat16).contiguous(), + float(norm.eps), + ) + return output.reshape(M, H) + + +def _rms_norm_eps(norm: nn.Module) -> float: + if hasattr(norm, "eps"): + return float(norm.eps) + return float(norm.variance_epsilon) + + +def _note_attn_res_fusion(site: str, fused: bool, M: int, H: int, N: int) -> None: + """Report whether the fused path was actually reached, once per shape. + + ``_FUSED_ATTN_RES_ENABLED`` only says the feature is switched on, not that + the shape gate let the call through, and a rejected call looks exactly like + a disabled one in the logs. Emitted at debug level: it is once per distinct + shape, not once per process, so it is a diagnostic rather than a summary. + """ + logger.debug_once( + f"Kimi K3 attn-res fusion [{site}]: " + f"{'FUSED' if fused else 'fallback'} (M={M}, H={H}, N={N})", + key=f"kimi_k3_attn_res_fusion_{site}_{fused}_{M}_{H}_{N}", + ) + + +def _apply_attn_res_rmsnorm_fused( + prefix_sum: torch.Tensor, + block_residual: torch.Tensor, + proj: nn.Linear, + norm: KimiK3RMSNorm, + output_norm: nn.Module, +) -> Optional[torch.Tensor]: + """Fuse attention-residual mixing with its immediately following norm.""" + if ( + prefix_sum.dtype is not torch.bfloat16 + or not prefix_sum.is_cuda + or not block_residual.is_cuda + ): + return None + M, H = prefix_sum.shape + K = int(block_residual.shape[0]) + N = K + 1 + # The fused path is taken for M <= _FUSED_ATTN_RES_MAX_TOKENS, H == 7168 and + # N <= 12, which is the measured window; larger token counts have not been + # measured and fall back to the unfused add + attn_res_fwd + RMSNorm path. + if _use_persistent_attn_res(M, H, N): + try: + persistent_op = torch.ops.trtllm.attn_res_add_rmsnorm_persistent_fwd + except (AttributeError, RuntimeError): + return None + _, output = persistent_op( + prefix_sum.reshape(M, 1, H).contiguous(), + None, + block_residual.reshape(K, M, 1, H).contiguous(), + proj.weight.reshape(-1).to(torch.bfloat16).contiguous(), + norm.weight.to(torch.bfloat16).contiguous(), + output_norm.weight.to(torch.bfloat16).contiguous(), + float(norm.eps), + _rms_norm_eps(output_norm), + ) + _note_attn_res_fusion("attn_res+norm/persistent", True, M, H, N) + return output.reshape(M, H) + + if M > _FUSED_ATTN_RES_MAX_TOKENS or H != 7168 or N > 12: + _note_attn_res_fusion("attn_res+norm", False, M, H, N) + return None + try: + attn_res_rmsnorm_op = torch.ops.trtllm.attn_res_rmsnorm_fwd + except (AttributeError, RuntimeError): + return None + layer_kernel = prefix_sum.reshape(M, 1, H).contiguous() + block_kernel = block_residual.reshape(K, M, 1, H).contiguous() + output = attn_res_rmsnorm_op( + layer_kernel, + block_kernel, + proj.weight.reshape(-1).to(torch.bfloat16).contiguous(), + norm.weight.to(torch.bfloat16).contiguous(), + output_norm.weight.to(torch.bfloat16).contiguous(), + float(norm.eps), + _rms_norm_eps(output_norm), + ) + _note_attn_res_fusion("attn_res+norm", True, M, H, N) + return output.reshape(M, H) + + +def _apply_attn_res_add_rmsnorm_fused( + prefix_sum: torch.Tensor, + addend: torch.Tensor, + block_residual: torch.Tensor, + proj: nn.Linear, + norm: KimiK3RMSNorm, + output_norm: nn.Module, +) -> Optional[Tuple[torch.Tensor, torch.Tensor]]: + """Fuse ``prefix_sum + addend``, attention-residual, and trailing norm. + + The production residual add produces a BF16 tensor that remains live across + the following MLP. The kernel therefore returns that materialized, + BF16-rounded prefix sum alongside the normalized attention-residual output, + while avoiding a separate add launch and a re-read of the intermediate by + attention-residual selection. + """ + if ( + prefix_sum.dtype is not torch.bfloat16 + or addend.dtype is not torch.bfloat16 + or not prefix_sum.is_cuda + or not addend.is_cuda + or not block_residual.is_cuda + or prefix_sum.shape != addend.shape + ): + return None + M, H = prefix_sum.shape + K = int(block_residual.shape[0]) + N = K + 1 + # Same measured window as _apply_attn_res_rmsnorm_fused above. + if _use_persistent_attn_res(M, H, N): + try: + persistent_op = torch.ops.trtllm.attn_res_add_rmsnorm_persistent_fwd + except (AttributeError, RuntimeError): + return None + updated_prefix_sum, output = persistent_op( + prefix_sum.reshape(M, 1, H).contiguous(), + addend.reshape(M, 1, H).contiguous(), + block_residual.reshape(K, M, 1, H).contiguous(), + proj.weight.reshape(-1).to(torch.bfloat16).contiguous(), + norm.weight.to(torch.bfloat16).contiguous(), + output_norm.weight.to(torch.bfloat16).contiguous(), + float(norm.eps), + _rms_norm_eps(output_norm), + ) + _note_attn_res_fusion("add+attn_res+norm/persistent", True, M, H, N) + return updated_prefix_sum.reshape(M, H), output.reshape(M, H) + + if M > _FUSED_ATTN_RES_MAX_TOKENS or H != 7168 or N > 12: + _note_attn_res_fusion("add+attn_res+norm", False, M, H, N) + return None + try: + attn_res_add_rmsnorm_op = torch.ops.trtllm.attn_res_add_rmsnorm_fwd + except (AttributeError, RuntimeError): + return None + layer_kernel = prefix_sum.reshape(M, 1, H).contiguous() + addend_kernel = addend.reshape(M, 1, H).contiguous() + block_kernel = block_residual.reshape(K, M, 1, H).contiguous() + updated_prefix_sum, output = attn_res_add_rmsnorm_op( + layer_kernel, + addend_kernel, + block_kernel, + proj.weight.reshape(-1).to(torch.bfloat16).contiguous(), + norm.weight.to(torch.bfloat16).contiguous(), + output_norm.weight.to(torch.bfloat16).contiguous(), + float(norm.eps), + _rms_norm_eps(output_norm), + ) + _note_attn_res_fusion("add+attn_res+norm", True, M, H, N) + return updated_prefix_sum.reshape(M, H), output.reshape(M, H) + + +def _apply_attn_res( + prefix_sum: torch.Tensor, block_residual: torch.Tensor, proj: nn.Linear, norm: KimiK3RMSNorm +) -> torch.Tensor: + """Exact port of HF ``modeling_kimi._apply_attn_res`` (fp32 math). + + prefix_sum: ``[num_tokens, hidden_size]`` + block_residual: ``[num_snapshots, num_tokens, hidden_size]`` + + Unless ``KIMI_K3_FUSED_ATTN_RES=0``, inputs fitting the fused kernel's + contract dispatch directly to the in-tree ``trtllm::attn_res_fwd`` op. + Only the fallback boundary restores the HF ``[M, K, H]`` layout. + """ + if _FUSED_ATTN_RES_ENABLED: + fused = _apply_attn_res_fused(prefix_sum, block_residual, proj, norm) + if fused is not None: + return fused + block_residual_hf = block_residual.transpose(0, 1) + v = torch.cat((block_residual_hf, prefix_sum.unsqueeze(1)), dim=1) + v_float = v.float() + variance = v_float.pow(2).mean(-1, keepdim=True) + k = v_float * torch.rsqrt(variance + norm.eps) + score_weight = norm.weight.float() * proj.weight.squeeze(0).float() + scores = (k * score_weight).sum(-1) + probs = scores.softmax(-1).unsqueeze(1) + hidden_states = torch.matmul(probs, v_float).squeeze(1) + return hidden_states.to(v.dtype) + + +def _apply_attn_res_and_rmsnorm( + prefix_sum: torch.Tensor, + block_residual: torch.Tensor, + proj: nn.Linear, + norm: KimiK3RMSNorm, + output_norm: nn.Module, +) -> torch.Tensor: + """Apply attention-residual selection and the next RMSNorm.""" + if _FUSED_ATTN_RES_ENABLED: + fused = _apply_attn_res_rmsnorm_fused(prefix_sum, block_residual, proj, norm, output_norm) + if fused is not None: + return fused + return output_norm(_apply_attn_res(prefix_sum, block_residual, proj, norm)) + + +def _apply_attn_res_add_and_rmsnorm( + prefix_sum: torch.Tensor, + addend: torch.Tensor, + block_residual: torch.Tensor, + proj: nn.Linear, + norm: KimiK3RMSNorm, + output_norm: nn.Module, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Add an attention output to the running residual, then select and norm.""" + if _FUSED_ATTN_RES_ENABLED: + fused = _apply_attn_res_add_rmsnorm_fused( + prefix_sum, addend, block_residual, proj, norm, output_norm + ) + if fused is not None: + return fused + updated_prefix_sum = prefix_sum + addend + return updated_prefix_sum, _apply_attn_res_and_rmsnorm( + updated_prefix_sum, block_residual, proj, norm, output_norm + ) + + +# Routed-expert key spellings that ModelOpt emits for Kimi K3. The NVFP4 +# checkpoint (``nvidia/Kimi-K3-NVFP4``) lists every prefix x module-name +# combination in ``quantized_layers``, so a lookup over this product finds it +# without needing the MiniMax-M3-style prefix normalization in ``ModelConfig``. +_K3_ROUTED_EXPERT_KEY_PREFIXES = ("language_model.model.", "model.", "") + + +_K3_ROUTED_EXPERT_KEY_SUFFIXES = ("block_sparse_moe.experts", "mlp.experts") + + +# The subset of the above that can be a real module path. ``exclude_modules`` +# matches with wildcards and walks ancestor prefixes, so an empty prefix would +# widen what matches instead of just missing, as it does in the dict lookup. +_K3_ROUTED_EXPERT_MODULE_PREFIXES = ("language_model.model.", "model.") + + +# Routed-expert quantization used when the checkpoint declares nothing per +# layer. The original ``moonshotai/Kimi-K3`` ships a compressed-tensors +# ``mxfp4-pack-quantized`` config with no ModelOpt per-layer entries, and that +# checkpoint is what this default has always served. +_K3_DEFAULT_ROUTED_QUANT_ALGO = QuantAlgo.W4A8_MXFP4_MXFP8 + + +def _load_packed_mxfp4_expert(backend, base, expert_idx, local_slot_id, get_tensor) -> None: + backend.quant_method.load_packed_mxfp4_expert( + backend, + global_expert_id=expert_idx, + local_slot_id=local_slot_id, + w1_weight=get_tensor(f"{base}.{expert_idx}.w1.weight_packed"), + w1_weight_scale=get_tensor(f"{base}.{expert_idx}.w1.weight_scale"), + w2_weight=get_tensor(f"{base}.{expert_idx}.w2.weight_packed"), + w2_weight_scale=get_tensor(f"{base}.{expert_idx}.w2.weight_scale"), + w3_weight=get_tensor(f"{base}.{expert_idx}.w3.weight_packed"), + w3_weight_scale=get_tensor(f"{base}.{expert_idx}.w3.weight_scale"), + ) + + +def _load_nvfp4_expert(backend, base, expert_idx, local_slot_id, get_tensor) -> None: + backend.quant_method.load_streaming_nvfp4_expert( + backend, + global_expert_id=expert_idx, + local_slot_id=local_slot_id, + **{ + f"{w}_{kind}": get_tensor(f"{base}.{expert_idx}.{w}.{kind}") + for w in ("w1", "w2", "w3") + for kind in ("weight", "weight_scale", "weight_scale_2", "input_scale") + }, + ) + + +class _K3ExpertCkptSpec(NamedTuple): + """How one routed-expert quantization is spelled and loaded.""" + + # Per-``w{1,2,3}`` checkpoint tensor suffixes this layout stores. + kinds: Tuple[str, ...] + loader: Callable[..., None] + # Set of filled slots the loader maintains, checked after the load. + loaded_slots_attr: str + # NVFP4 defers cat/pad/interleave and the alpha computation to + # ``process_weights_after_loading``; the MXFP4 loaders write through. + needs_layer_finalize: bool + + +_K3_EXPERT_CKPT_SPECS = { + QuantAlgo.W4A8_MXFP4_MXFP8: _K3ExpertCkptSpec( + kinds=("weight_packed", "weight_scale"), + loader=_load_packed_mxfp4_expert, + loaded_slots_attr="_packed_mxfp4_loaded_slots", + needs_layer_finalize=False, + ), + QuantAlgo.NVFP4: _K3ExpertCkptSpec( + kinds=("weight", "weight_scale", "weight_scale_2", "input_scale"), + loader=_load_nvfp4_expert, + loaded_slots_attr="_streamed_expert_slots", + needs_layer_finalize=True, + ), +} + + +def _k3_expert_ckpt_spec(quant_algo: Optional[QuantAlgo]) -> _K3ExpertCkptSpec: + spec = _K3_EXPERT_CKPT_SPECS.get(quant_algo) + if spec is None: + raise NotImplementedError( + f"Kimi K3 routed experts are quantized as {quant_algo}, for which " + "no per-expert checkpoint layout is known. Supported: " + f"{sorted(a.name for a in _K3_EXPERT_CKPT_SPECS)}." + ) + return spec + + +class KimiK3MoERuntime(nn.Module): + """Kimi K3 latent MoE block backed by ConfigurableMoE.""" + + def __init__( + self, + model_config: ModelConfig, + cfg, + layer_idx: int, + aux_stream_dict: Dict[AuxStreamType, torch.cuda.Stream], + ): + """Build the routed experts and the shared expert for one MoE layer. + + ``cfg`` is the raw ``PretrainedConfig`` rather than anything derived: + the SiTU soft-caps and the routed-expert geometry are Kimi K3 fields + that ``ModelConfig`` does not carry. + + ``aux_stream_dict`` is shared across every layer of the model, so the + streams reached through it are borrowed and must not be synchronized + or reassigned here. + """ + super().__init__() + self.layer_idx = layer_idx + self.hidden_size = cfg.hidden_size + self.num_experts = cfg.num_experts + self.top_k = cfg.num_experts_per_token + self.moe_hidden_size = cfg.routed_expert_hidden_size + # ValueError (not assert): these guard unsupported checkpoint + # configurations and must stay active under ``python -O``. + if self.moe_hidden_size is None: + raise ValueError("Kimi K3 runtime expects the latent MoE (routed_expert_hidden_size)") + if not getattr(cfg, "latent_moe_use_norm", False): + raise ValueError("Kimi K3 runtime expects latent_moe_use_norm=True") + + situ_beta, situ_linear_beta = _resolve_kimi_situ_betas(cfg) + dtype = torch.bfloat16 + + # Routing scores stay fp32; with attention-DP off the gate GEMM runs + # bf16xbf16 with fp32 accumulate/output (checkpoint stores the gate + # weight in bf16; saves a per-layer input cast + fp32 splitK pair on + # the bs1 decode path). Under attention-DP the legacy upcast-to-fp32 + # GEMM is kept: the bf16-input min-latency GEMM's different reduction + # order flips borderline top-16 picks (GSM8K 96.7 -> 96.1/96.4, + # 3-run bisect on 62b20dd868), and the bs1-latency win is irrelevant + # at DEP batch sizes. KIMI_K3_ROUTER_BF16=1/0 forces either path. + _router_bf16_env = os.environ.get("KIMI_K3_ROUTER_BF16") + _router_bf16 = ( + _router_bf16_env == "1" + if _router_bf16_env is not None + else not model_config.mapping.enable_attention_dp + ) + self.gate = KimiK3MoEGate(cfg, logits_gemm_dtype=torch.bfloat16 if _router_bf16 else None) + + routed_moe_model_config = self._routed_moe_model_config(model_config) + routed_quant_config = self._resolve_routed_quant_config(model_config, layer_idx) + # Resolved here so ``load_weights`` reads the checkpoint layout off the + # module instead of re-deriving it at each of its three call sites. + self.expert_ckpt_spec = _k3_expert_ckpt_spec(routed_quant_config.quant_algo) + routed_moe_kwargs = dict( + routing_method=self.gate.routing_method, + num_experts=self.num_experts, + hidden_size=self.moe_hidden_size, + intermediate_size=cfg.moe_intermediate_size, + dtype=dtype, + # Kimi owns the latent reduction so it can order that collective + # after the shared expert's auxiliary-stream reduction. + reduce_results=False, + model_config=routed_moe_model_config, + override_quant_config=routed_quant_config, + layer_idx=layer_idx, + aux_stream_dict=aux_stream_dict, + # Let CommunicationFactory select the best available strategy. + communication_method=None, + activation=SiTuActivation( + gate_softcap=situ_beta, + linear_softcap=situ_linear_beta, + ), + # A request that silently degraded to CUTLASS would be benchmarked + # as if it were the backend that was asked for, and the decline is + # easy to trigger: MegaMoE has its own token / top-k limits and is + # EP-only, and CuteDSL declines on activation shape, SM version and + # the CuTe DSL dependency. As measured once: a CUTEDSL request + # was turned down on every one of the 92 MoE layers, on all 16 + # ranks, and still produced correct text and a zero exit -- the + # only trace was a warning line per layer. Fail in the resolver + # instead, which reports the rejection trail. + # + # CUTLASS is absent on purpose: it is the fallback target, so + # "degraded to CUTLASS" is not a thing that can happen to it. + allow_backend_degradation=routed_moe_model_config.moe_backend + not in ("MEGAMOE_DEEPGEMM", "MEGAMOE_CUTEDSL", "CUTEDSL"), + ) + self._check_trtllm_situ_quant( + routed_moe_model_config.moe_backend, routed_quant_config.quant_algo + ) + + self.routed_experts = create_moe(**routed_moe_kwargs) + if not isinstance(self.routed_experts, ConfigurableMoE): + raise RuntimeError( + "Kimi K3 requires ConfigurableMoE; ENABLE_CONFIGURABLE_MOE must not be disabled." + ) + if self.routed_experts.layer_load_balancer is not None: + raise NotImplementedError( + "Kimi K3 packed-checkpoint streaming does not yet support " + "dynamic EPLB or replicated expert slots." + ) + local_expert_ids = list(self.routed_experts.backend.initial_local_expert_ids) + if local_expert_ids != list( + range(local_expert_ids[0], local_expert_ids[0] + len(local_expert_ids)) + ): + raise NotImplementedError( + "Kimi K3 packed-checkpoint streaming currently requires a " + "contiguous static expert partition." + ) + self.local_expert_ids = tuple(local_expert_ids) + self.experts_per_rank = len(local_expert_ids) + self.expert_lo = local_expert_ids[0] + self.expert_hi = self.expert_lo + self.experts_per_rank + + shared_intermediate = cfg.moe_intermediate_size * cfg.num_shared_experts + attention_dp = model_config.mapping.enable_attention_dp + shared_model_config = copy.copy(model_config) + shared_model_config.quant_config = QuantConfig() + # Under attention DP each rank owns different tokens, so the shared + # expert is replicated (TP size 1) and must not reduce across ranks. + # Direct MoE-TP leaves both branches as partials for one concatenated + # all-reduce. + use_shared_tp = not attention_dp and model_config.mapping.tp_size > 1 + self._reduce_routed_output = ( + use_shared_tp + and self.routed_experts.backend.scheduler_kind != MoESchedulerKind.FUSED_COMM + ) + if self._reduce_routed_output and self.routed_experts.all_reduce is None: + raise RuntimeError( + "Kimi K3 direct MoE tensor parallelism requires the " + "ConfigurableMoE all-reduce even when reduce_results=False." + ) + self.shared_experts = GatedMLP( + hidden_size=cfg.hidden_size, + intermediate_size=shared_intermediate, + bias=False, + activation=SituAndMul( + beta=situ_beta, + linear_beta=situ_linear_beta, + use_fused_activation=True, + ), + dtype=dtype, + config=shared_model_config, + overridden_tp_size=1 if attention_dp else None, + reduce_output=use_shared_tp, + layer_idx=layer_idx, + is_shared_expert=True, + ) + # Side stream (+ fork/join events) for overlapping shared-expert + # compute with the routed chain. Only engaged when multi-stream is + # active (CUDA graphs on); otherwise both run in order on the default + # stream. + self.shared_expert_stream = aux_stream_dict[AuxStreamType.MoeShared] + self.moe_main_event = torch.cuda.Event() + self.moe_shared_event = torch.cuda.Event() + self.routed_expert_down_proj = nn.Linear( + cfg.hidden_size, self.moe_hidden_size, bias=False, dtype=dtype + ) + self.routed_expert_up_proj = nn.Linear( + self.moe_hidden_size, cfg.hidden_size, bias=False, dtype=dtype + ) + # Stock fused RMSNorm (flashinfer kernel; the no-flashinfer + # fallback is the same fp32-variance eager math as KimiK3RMSNorm). + self.routed_expert_norm = RMSNorm( + hidden_size=self.moe_hidden_size, eps=cfg.rms_norm_eps, dtype=dtype + ) + + @staticmethod + def _routed_projection(hidden_states: torch.Tensor, projection: nn.Module) -> torch.Tensor: + if _K3_DISABLE_MIN_LATENCY_LATENT_PROJ or not isinstance(projection, nn.Linear): + return projection(hidden_states) + return torch.ops.trtllm.dsv3_fused_a_gemm_op( + hidden_states, projection.weight.t(), None, None + ) + + @staticmethod + def _select_moe_tp_ep(mapping: Mapping) -> Tuple[int, int]: + """Resolve the routed-expert ``(moe_tp, moe_ep)`` split. + + Precedence: + + 1. Explicit ``moe_tensor_parallel_size`` / ``moe_expert_parallel_size`` + from the user config. Detected via + ``mapping.moe_tp_ep_user_specified`` so the auto-resolved mapping + default (``moe_tp=tp_size, moe_ep=1``) is NOT mistaken for a TP + request. + 2. Default: EP-only (``moe_tp=1, moe_ep=tp_size``), the historical + K3 layout. + """ + tp_size = mapping.tp_size + if getattr(mapping, "moe_tp_ep_user_specified", False): + return mapping.moe_tp_size, mapping.moe_ep_size + return 1, tp_size + + @staticmethod + def _resolve_routed_quant_config(model_config: ModelConfig, layer_idx: int) -> QuantConfig: + """Routed-expert quantization for ``layer_idx``, taken from the checkpoint. + + ``nvidia/Kimi-K3-NVFP4`` declares the routed experts per layer as + ``NVFP4`` with ``group_size=16``; the original ``moonshotai/Kimi-K3`` + declares nothing per layer and keeps the historical + ``W4A8_MXFP4_MXFP8`` default. Reading the checkpoint instead of + hardcoding is what lets one code path serve both. + + An exclusion outranks the per-layer entry and the default below: + ``create_weights`` treats an override as authoritative over anything + ``__post_init__`` wrote, so this return value stands in for both + quantization passes and exclusion is the one that runs second. It is + matched as a pattern, so it is asked only about real module names. + """ + quant_config = model_config.quant_config + if quant_config is not None and any( + quant_config.is_module_excluded_from_quantization( + f"{prefix}layers.{layer_idx}.{suffix}" + ) + for prefix in _K3_ROUTED_EXPERT_MODULE_PREFIXES + for suffix in _K3_ROUTED_EXPERT_KEY_SUFFIXES + ): + logger.debug( + "Kimi K3 layer %d routed experts: excluded from quantization, " + "keeping them unquantized", + layer_idx, + ) + return QuantConfig(kv_cache_quant_algo=quant_config.kv_cache_quant_algo) + + per_layer = getattr(model_config, "quant_config_dict", None) + if per_layer: + for prefix in _K3_ROUTED_EXPERT_KEY_PREFIXES: + for suffix in _K3_ROUTED_EXPERT_KEY_SUFFIXES: + cfg = per_layer.get(f"{prefix}layers.{layer_idx}.{suffix}") + if cfg is not None and cfg.quant_algo is not None: + # Logged once per layer: the routed-expert format decides + # which MoE backends can serve this checkpoint at all. + logger.debug( + "Kimi K3 layer %d routed experts: %s (group_size=%s) " + "from the checkpoint", + layer_idx, + cfg.quant_algo, + cfg.group_size, + ) + return cfg + logger.debug( + "Kimi K3 layer %d routed experts: no per-layer quant config in the " + "checkpoint, defaulting to %s", + layer_idx, + _K3_DEFAULT_ROUTED_QUANT_ALGO, + ) + return QuantConfig(quant_algo=_K3_DEFAULT_ROUTED_QUANT_ALGO) + + @staticmethod + def _check_trtllm_situ_quant(moe_backend: str, quant_algo: Optional[QuantAlgo]) -> None: + """Reject a routed-expert format trtllm-gen has no fused SiTu cubin for. + + trtllm-gen has fused SiTu FC1 cubins for two input formats and no + standalone SiTu activation kernel, so anything else has to die here + rather than in a cubin lookup deep inside the runner. Checked against + the resolved backend, not the K3 architecture branch, because the + generic FP8_BLOCK_SCALES fallback in ``resolve_moe_backend`` can also + land on TRTLLM. + + The admitted set is read off the backend rather than restated here, + because restating it is what broke. This guard was written in #17865 + when MXFP4 was the only fused SiTu drop; #17940 then added the NVFP4 + (group-16 ``Bmm_E2m1_E2m1E2m1_..._siTuGlu_*``) cubins and updated + ``TRTLLMGenFusedMoE``'s set without touching this copy. For the week + in between, an NVFP4 K3 checkpoint could not start at all -- and not + only when TRTLLM was asked for by name, because + ``ModelConfig.resolve_moe_backend`` sends every K3 architecture to + TRTLLM, so the default AUTO configuration hit this raise too. The unit + tests did not catch it: they call ``create_moe`` directly and never + reach this guard, so the kernel path stayed green while the model path + was closed. + + A staticmethod, not an inline block, so that the invariant is + reachable from a test without constructing the whole runtime. + """ + situ_supported = TRTLLMGenFusedMoE.situ_supported_quant_algos() + if moe_backend != "TRTLLM" or quant_algo in situ_supported: + return + supported = ", ".join(sorted(algo.name for algo in situ_supported)) + raise ValueError( + f"Kimi K3 routed experts are quantized as {quant_algo}, which the " + "TRTLLM (trtllm-gen) MoE backend cannot serve: fused SiTu cubins " + f"exist only for {supported}. Set moe_config.backend to CUTLASS " + "or MEGAMOE_CUTEDSL." + ) + + @staticmethod + def _routed_moe_model_config(model_config: ModelConfig) -> ModelConfig: + """Build a private routed-expert mapping without mutating the shared + config. Default split is EP-only; see ``_select_moe_tp_ep``.""" + # Every backend here declares ``ActivationType.SiTu`` in its + # ``activation_support``; the list is not a preference order. CUTEDSL + # joined once its act-fusion kernel grew the SiTU epilogue. + supported_backends = { + "CUTLASS", + "TRTLLM", + "CUTEDSL", + "MEGAMOE_DEEPGEMM", + "MEGAMOE_CUTEDSL", + } + if model_config.moe_backend not in supported_backends: + raise ValueError( + "Kimi K3 SiTU routed experts only support the CUTLASS, TRTLLM, " + "CUTEDSL, MEGAMOE_DEEPGEMM, and MEGAMOE_CUTEDSL backends; " + f"got {model_config.moe_backend!r}." + ) + if model_config.moe_load_balancer is not None: + raise NotImplementedError( + "Kimi K3 packed-checkpoint streaming does not yet support " + "EPLB or replicated expert slots." + ) + mapping = model_config.mapping + if getattr(mapping, "_dwdp_size", 0) > 1: + raise NotImplementedError("Kimi K3 packed-checkpoint streaming does not support DWDP.") + + moe_tp, moe_ep = KimiK3MoERuntime._select_moe_tp_ep(mapping) + if moe_tp < 1 or moe_ep < 1 or moe_tp * moe_ep != mapping.tp_size: + raise ValueError( + f"Kimi K3 routed MoE split moe_tp={moe_tp} x moe_ep={moe_ep} " + f"must multiply to tp_size={mapping.tp_size}." + ) + if moe_tp > 1 and mapping.enable_attention_dp: + raise NotImplementedError( + "Kimi K3 MoE tensor parallelism requires " + "enable_attention_dp=false (the attention-DP dispatch/combine " + "path is validated for EP-only splits)." + ) + logger.info_once( + f"Kimi K3 routed MoE parallelism: moe_tp={moe_tp}, " + f"moe_ep={moe_ep} (tp_size={mapping.tp_size})", + key="kimi_k3_moe_tp_ep_split", + ) + + mapping_dict = mapping.to_dict() + mapping_dict["moe_cluster_size"] = 1 + mapping_dict["moe_tp_size"] = moe_tp + mapping_dict["moe_ep_size"] = moe_ep + routed_mapping = Mapping.from_dict(mapping_dict) + + routed_model_config = copy.copy(model_config) + routed_model_config._frozen = False + routed_model_config.extra_attrs = copy.copy(model_config.extra_attrs) + routed_model_config.mapping = routed_mapping + routed_model_config.moe_backend = model_config.moe_backend + # MegaMoE uses this value as global DP SymmBuffer capacity, then + # divides it by EP size for the per-rank allocation. Other backends + # keep the user-configured value as their MoE chunking bound. + # Preserve an explicitly larger capacity. + if routed_model_config.moe_backend in { + "MEGAMOE_DEEPGEMM", + "MEGAMOE_CUTEDSL", + }: + default_moe_max_num_tokens = routed_model_config.max_num_tokens * routed_mapping.dp_size + configured_moe_max_num_tokens = int(routed_model_config.moe_max_num_tokens or 0) + if configured_moe_max_num_tokens < default_moe_max_num_tokens: + logger.info_once( + "Kimi K3 MegaMoE raises moe_max_num_tokens from " + f"{configured_moe_max_num_tokens} to {default_moe_max_num_tokens} " + "because the global DP SymmBuffer requires capacity for " + "max_num_tokens * dp_size.", + key=( + "kimi_k3_megamoe_capacity_override_" + f"{configured_moe_max_num_tokens}_{default_moe_max_num_tokens}" + ), + ) + routed_model_config.moe_max_num_tokens = max( + configured_moe_max_num_tokens, + default_moe_max_num_tokens, + ) + routed_model_config._frozen = True + return routed_model_config + + def forward(self, hidden_states: torch.Tensor, all_rank_num_tokens=None) -> torch.Tensor: + """``hidden_states``: ``[num_tokens, hidden_size]`` bf16.""" + identity = hidden_states + router_logits = self.gate.compute_logits(hidden_states) + moe_all_reduce = self.routed_experts.all_reduce if self._reduce_routed_output else None + + def _routed_output(): + # Latent down/up projections via the min-latency fused GEMM op: + # at <=16 tokens (decode graphs) it runs a single pipelined + # bf16 kernel per projection instead of cuBLAS's split-K GEMV + + # splitKreduce pair (~17+3.6us -> ~8us for 7168->3584 at M=1); + # for larger token counts the op falls back to cuBLAS internally. + # TLLM_K3_DISABLE_MIN_LATENCY_LATENT_PROJ=1 restores nn.Linear + # (A/B escape hatch). When the FP8 weight-read conversion has + # replaced the projection module, call it directly: its weight is + # an e4m3 buffer the bf16 dsv3 op must not read, and its forward + # is already a single fused GEMM (fp8_swap_ab_gemm). + routed_in = self._routed_projection(hidden_states, self.routed_expert_down_proj) + y = self.routed_experts( + routed_in, + router_logits, + all_rank_num_tokens=all_rank_num_tokens, + ) + if self._reduce_routed_output: + return y + # Communication-backed paths return a complete routed result. + y = self.routed_expert_norm(y) + return self._routed_projection(y, self.routed_expert_up_proj) + + # Shared experts depend only on the block input, so overlap their GEMMs + # with the routed dispatch/expert/combine chain. Multi-stream engages + # only under CUDA graphs; otherwise both branches run in order on the + # default stream. The shared GatedMLP includes its output all-reduce on + # the auxiliary stream. The join below must precede the routed + # all-reduce: concurrent collectives on different streams can corrupt + # SYMM_MEM all-reduce state. + routed_out, shared_out = maybe_execute_in_parallel( + _routed_output, + lambda: self.shared_experts(identity), + self.moe_main_event, + self.moe_shared_event, + self.shared_expert_stream, + disable_on_compile=True, + ) + if self._reduce_routed_output: + routed_latent = moe_all_reduce(routed_out) + routed_latent = self.routed_expert_norm(routed_latent) + routed_out = self._routed_projection(routed_latent, self.routed_expert_up_proj) + return routed_out + shared_out + + +def resolve_attention_quant_config( + config: ModelConfig | None, layer_idx: int, projection: str +) -> QuantConfig: + """Resolve a checkpoint projection, including mixed-precision exclusions.""" + if config is None: + return QuantConfig() + global_config = config.quant_config or QuantConfig() + names = [ + f"{prefix}layers.{layer_idx}.self_attn.{projection}" + for prefix in ("language_model.model.", "model.", "") + ] + if any(global_config.is_module_excluded_from_quantization(name) for name in names): + return QuantConfig(kv_cache_quant_algo=global_config.kv_cache_quant_algo) + declarations = config.quant_config_dict or {} + matches = [declarations[name] for name in names if name in declarations] + if matches: + selected = matches[0] + if any(match.quant_algo != selected.quant_algo for match in matches[1:]): + raise ValueError(f"Conflicting Kimi K3 quantization aliases for {names[0]}") + elif global_config.quant_algo == QuantAlgo.MIXED_PRECISION: + selected = QuantConfig() + else: + selected = global_config + if selected.quant_algo not in (None, QuantAlgo.FP8_BLOCK_SCALES): + raise ValueError( + f"Kimi K3 attention projection {names[0]} has unsupported checkpoint " + f"quantization {selected.quant_algo}" + ) + if selected.quant_algo == QuantAlgo.FP8_BLOCK_SCALES and selected.group_size not in (None, 128): + raise ValueError(f"Kimi K3 attention requires 128x128 FP8 blocks for {names[0]}") + result = copy.copy(selected) + result.kv_cache_quant_algo = global_config.kv_cache_quant_algo + return result + + +class KimiMLARuntime(nn.Module): + """Wraps K3 MLA and applies its external TP output reduction.""" + + def __init__( + self, + cfg: "PretrainedConfig", + layer_idx: int, + model_config: ModelConfig, + aux_stream_dict: Dict[AuxStreamType, torch.cuda.Stream], + mapping_with_cp: Optional[Mapping] = None, + ) -> None: + super().__init__() + + max_positions = int( + os.environ.get( + _KIMI_K3_MLA_MAX_POSITIONS_ENV, + cfg.max_position_embeddings, + ) + ) + self.layer_idx = layer_idx + # KimiK3MLAAttention owns MLA projection/head sharding. Keep only the + # final output reduction in this wrapper so the output gate remains + # between attention and the row-parallel o_proj. + # Helix: mapping_with_cp (the CP original) activates the base MLA's + # helix machinery; this wrapper's allreduce over the repurposed + # mapping sums the base o_proj's tp*cp partials. + mapping = model_config.mapping + reduce_output = not mapping.enable_attention_dp and mapping.tp_size > 1 + self._o_allreduce = ( + AllReduce( + mapping=mapping, + strategy=model_config.allreduce_strategy, + dtype=torch.bfloat16, + ) + if reduce_output + else None + ) + attention_config = copy.copy(model_config) + attention_config._frozen = False + attention_config.quant_config_dict = { + name: resolve_attention_quant_config(model_config, layer_idx, name) + for name in ( + "q_a_proj", + "kv_a_proj_with_mqa", + "q_b_proj", + "kv_b_proj", + "g_proj", + "o_proj", + ) + } + attention_config.quant_config = QuantConfig( + kv_cache_quant_algo=model_config.quant_config.kv_cache_quant_algo + if model_config.quant_config is not None + else None + ) + attention_config._frozen = model_config._frozen + self.mixer = K3DecodeMLA( + hidden_size=cfg.hidden_size, + num_heads=cfg.num_attention_heads, + q_lora_rank=cfg.q_lora_rank, + kv_lora_rank=cfg.kv_lora_rank, + qk_nope_head_dim=cfg.qk_nope_head_dim, + qk_rope_head_dim=cfg.qk_rope_head_dim, + v_head_dim=cfg.v_head_dim, + rms_norm_eps=cfg.rms_norm_eps, + dtype=torch.bfloat16, + layer_idx=layer_idx, + use_output_gate=cfg.mla_use_output_gate, + max_position_embeddings=max_positions, + model_config=attention_config, + aux_stream_dict=aux_stream_dict, + mapping_with_cp=mapping_with_cp, + ) + + def forward( + self, + hidden_states: torch.Tensor, + attn_metadata: AttentionMetadata, + step: Optional[DecodeStep] = None, + ) -> torch.Tensor: + # MLA.forward takes position_ids first; K3 is NoPE, so pass None. + out = self.mixer(None, hidden_states, attn_metadata, step=step) + if self._o_allreduce is not None: + # Head-sharded TP: sum the row-sharded o_proj partials across + # the head-shard group. + out = self._o_allreduce(out) + return out + + +class KimiLinearDecoderLayer(nn.Module): + def __init__( + self, + model_config: ModelConfig, + cfg, + layer_idx: int, + aux_stream_dict: Dict[AuxStreamType, torch.cuda.Stream], + ): + super().__init__() + self.layer_idx = layer_idx + self.hidden_size = cfg.hidden_size + dtype = torch.bfloat16 + + self.is_kda = _is_kda_layer(cfg, layer_idx) + is_mla = _is_mla_layer(cfg, layer_idx) + if self.is_kda == is_mla: + raise ValueError(f"Kimi K3 layer {layer_idx} must be exactly one of KDA/MLA") + + if self.is_kda: + projection_names = ("q_proj", "k_proj", "v_proj", "g_proj", "o_proj") + attention_config = copy.copy(model_config) + attention_config._frozen = False + attention_config.quant_config_dict = { + name: resolve_attention_quant_config(model_config, layer_idx, name) + for name in projection_names + } + attention_config._frozen = model_config._frozen + self.linear_attn = K3DecodeKDA( + cfg, + layer_idx, + mapping=model_config.mapping, + allreduce_strategy=model_config.allreduce_strategy, + aux_stream=aux_stream_dict[AuxStreamType.Attention], + model_config=attention_config, + ) + else: + self.self_attn = KimiMLARuntime( + cfg, + layer_idx, + model_config=model_config, + aux_stream_dict=aux_stream_dict, + # CP original stashed by _setup_helix_mappings; None outside helix. + mapping_with_cp=getattr(model_config, "_helix_mapping_with_cp", None), + ) + + self.is_moe = ( + cfg.num_experts is not None + and layer_idx >= cfg.first_k_dense_replace + and layer_idx % getattr(cfg, "moe_layer_freq", 1) == 0 + ) + if self.is_moe: + self.block_sparse_moe = KimiK3MoERuntime(model_config, cfg, layer_idx, aux_stream_dict) + else: + situ_beta = getattr(cfg, "activation_situ_beta", None) or 1.0 + situ_linear_beta = getattr(cfg, "activation_situ_linear_beta", None) + attention_dp = model_config.mapping.enable_attention_dp + if attention_dp: + self.mlp_tp_size = 1 + else: + self.mlp_tp_size = math.gcd(cfg.intermediate_size, model_config.mapping.tp_size) + # Over MNNVL (one NVLink domain across the nodes, where a cross-node all-reduce costs what a node's + # does) the MLP stays split over the whole TP group, the per-rank shapes its decode GEMVs take + # (decode_gemv.SITES); otherwise it stays within one node. + spans_nodes = self._mnnvl_allreduce() is not None + if self.mlp_tp_size > model_config.mapping.gpus_per_node and not spans_nodes: + self.mlp_tp_size = math.gcd( + self.mlp_tp_size, model_config.mapping.gpus_per_node + ) + mlp_model_config = copy.copy(model_config) + mlp_model_config.quant_config = QuantConfig() + # K3's dense layer is BF16, so a unit block size gives the same + # subgroup selection as DeepSeek-V3. Attention DP replicates the + # MLP because ranks own different tokens. + self.mlp = GatedMLP( + hidden_size=cfg.hidden_size, + intermediate_size=cfg.intermediate_size, + bias=False, + activation=SituAndMul( + beta=situ_beta, + linear_beta=situ_linear_beta, + use_fused_activation=True, + ), + dtype=dtype, + config=mlp_model_config, + overridden_tp_size=self.mlp_tp_size, + reduce_output=self.mlp_tp_size > 1, + layer_idx=layer_idx, + ) + self._situ = (situ_beta, situ_linear_beta) + # The decode GEMVs' state (decode_gemv.py), set by the target's cache_derived_state. + self.decode_gemvs: Optional[_decode_gemv.K3DecodeGemvs] = None + + # Stock fused RMSNorm for the plain (whole-tensor) norms; numerics + # are drop-in for KimiK3RMSNorm (fp32 variance, weight applied + # after downcast, use_gemma=False). + self.input_layernorm = RMSNorm( + hidden_size=cfg.hidden_size, eps=cfg.rms_norm_eps, dtype=dtype + ) + self.post_attention_layernorm = RMSNorm( + hidden_size=cfg.hidden_size, eps=cfg.rms_norm_eps, dtype=dtype + ) + + # Attention residual scheme (always on for K3). The res norms stay + # KimiK3RMSNorm: they are consumed field-wise (.weight/.eps) by + # _apply_attn_res and the fused attn_res op, never called as + # modules. + self.attn_res_block_size = cfg.attn_res_block_size + assert self.attn_res_block_size is not None, ( + "Kimi K3 runtime expects attn_res_block_size to be set" + ) + self.self_attention_res_norm = KimiK3RMSNorm( + cfg.hidden_size, eps=cfg.rms_norm_eps, dtype=dtype + ) + self.mlp_res_norm = KimiK3RMSNorm(cfg.hidden_size, eps=cfg.rms_norm_eps, dtype=dtype) + self.self_attention_res_proj = nn.Linear(cfg.hidden_size, 1, bias=False, dtype=dtype) + self.mlp_res_proj = nn.Linear(cfg.hidden_size, 1, bias=False, dtype=dtype) + + def forward( + self, + hidden_states: torch.Tensor, + block_residual: torch.Tensor, + num_snapshots: int, + attn_metadata: AttentionMetadata, + capture: Optional[Tuple[Any, int]] = None, + step: Optional[DecodeStep] = None, + prenormed: bool = False, + ) -> Tuple[torch.Tensor, int]: + """Port of HF ``KimiDecoderLayer._forward_attn_residual`` (per token). + + ``block_residual`` is a preallocated snapshot bank in kernel-native + ``[K_max, M, H]`` layout. Returns the running prefix sum and the + number of valid bank rows. + + ``capture`` is ``(spec_metadata, layer_id)`` and taps the DSpark aux + stream for the layer BEFORE this one: the aggregated stream for layer j + is by definition what its next consumer sees, so the mixture computed + below already is it. Reading it here beats recomputing it, and is only + possible because K3 asserts pp_size == 1 -- layer j+1 is always local. + PP support would need a recompute at the rank boundary. + + ``step`` is the step's classification (``decode_step``), handed to the + attention module. + + ``prenormed`` (layer 0 on a decode step): ``hidden_states`` already is + this layer's input norm, and the layer's input, the step's embedding, + already is in ``block_residual[0]`` (``K3DecodeGemvs.embed_norm``). + """ + prefix_sum = hidden_states + valid_block_residual = block_residual[:num_snapshots] + + if prenormed: + assert num_snapshots == 0 and self.layer_idx % self.attn_res_block_size == 0 + elif capture is not None: + # The mixture tap needs the PRE-norm value, which the fused + # attn-res + RMSNorm kernel does not expose. Keep the two steps + # split on captured layers only and fuse everywhere else. + if num_snapshots > 0: + hidden_states = _apply_attn_res( + prefix_sum, + valid_block_residual, + self.self_attention_res_proj, + self.self_attention_res_norm, + ) + # A property of the DRAFTER checkpoint, not a knob: a mismatch only lowers + # acceptance, silently. hidden_states is the pre-norm attn_res mixture; + # prefix_only wants the running prefix, already in hand as prefix_sum. + tapped = hidden_states if _AUX_ATTN_RES_STREAM_ENABLED else prefix_sum + capture[0].maybe_capture_hidden_states(capture[1], tapped, None) + hidden_states = self.input_layernorm(hidden_states) + elif num_snapshots > 0: + hidden_states = _apply_attn_res_and_rmsnorm( + prefix_sum, + valid_block_residual, + self.self_attention_res_proj, + self.self_attention_res_norm, + self.input_layernorm, + ) + else: + hidden_states = self.input_layernorm(hidden_states) + + if self.layer_idx % self.attn_res_block_size == 0: + if not prenormed: + block_residual[num_snapshots].copy_(prefix_sum) + num_snapshots += 1 + valid_block_residual = block_residual[:num_snapshots] + prefix_sum = None + if self.is_kda: + hidden_states = self.linear_attn(hidden_states, attn_metadata, step=step) + else: + hidden_states = self.self_attn(hidden_states, attn_metadata, step=step) + + if prefix_sum is None: + prefix_sum = hidden_states + hidden_states = _apply_attn_res_and_rmsnorm( + prefix_sum, + valid_block_residual, + self.mlp_res_proj, + self.mlp_res_norm, + self.post_attention_layernorm, + ) + else: + prefix_sum, hidden_states = _apply_attn_res_add_and_rmsnorm( + prefix_sum, + hidden_states, + valid_block_residual, + self.mlp_res_proj, + self.mlp_res_norm, + self.post_attention_layernorm, + ) + if self.is_moe: + hidden_states = self.block_sparse_moe( + hidden_states, getattr(attn_metadata, "all_rank_num_tokens", None) + ) + else: + hidden_states = self._dense_mlp(hidden_states, step) + + prefix_sum = prefix_sum + hidden_states + return prefix_sum, num_snapshots + + def _dense_mlp(self, hidden_states: torch.Tensor, step: Optional[DecodeStep]) -> torch.Tensor: + """The dense MLP: on a step of at most DECODE_MAX_TOKENS tokens, its GEMVs and activation on the decode + kernels (``K3DecodeGemvs.dense_mlp``), then the down projection's all-reduce; else the module.""" + gemvs = self.decode_gemvs + if gemvs is not None and step is not None and step.small: + out = gemvs.dense_mlp( + hidden_states, + self.mlp.gate_up_proj.weight, + self.mlp.down_proj.weight, + *self._situ, + ) + if out is not None: + return self.mlp.down_proj.all_reduce(out) if self.mlp_tp_size > 1 else out + return self.mlp(hidden_states) + + def _mnnvl_allreduce(self): + """The MNNVL all-reduce of this layer's attention output, or None.""" + attention = self.linear_attn if self.is_kda else self.self_attn + return getattr(getattr(attention, "_o_allreduce", None), "mnnvl_allreduce", None) + + def skip_forward( + self, + hidden_states: torch.Tensor, + block_residual: torch.Tensor, + attn_metadata: AttentionMetadata, + ) -> Tuple[torch.Tensor, torch.Tensor]: + """No-op stand-in for ``forward``, matching ``DecoderLayer.skip_forward``. + + ``modeling_utils.skip_forward()`` only drops a module's weights when it + finds this attribute, so without it the layer-wise benchmarks would + allocate all 93 layers instead of the profiled slice. + """ + return hidden_states, block_residual + + +class KimiLinearModel(DecoderModel): + def __init__(self, model_config: ModelConfig): + super().__init__(model_config) + cfg = _get_text_config(model_config.pretrained_config) + self._text_cfg = cfg + dtype = torch.bfloat16 + + # Attention and MoE phases are sequential, so their branch-overlap + # roles share one stream; MoE-internal overlap roles remain separate. + aux_stream_list = [torch.cuda.Stream() for _ in range(4)] + self.aux_stream_dict = { + AuxStreamType.Attention: aux_stream_list[0], + AuxStreamType.MoeShared: aux_stream_list[0], + AuxStreamType.MoeChunkingOverlap: aux_stream_list[1], + AuxStreamType.MoeBalancer: aux_stream_list[2], + AuxStreamType.MoeOutputMemset: aux_stream_list[3], + } + + self.embed_tokens = nn.Embedding(cfg.vocab_size, cfg.hidden_size, dtype=dtype) + self.layers = nn.ModuleList( + [ + KimiLinearDecoderLayer(model_config, cfg, layer_idx, self.aux_stream_dict) + for layer_idx in range(cfg.num_hidden_layers) + ] + ) + self.norm = RMSNorm(hidden_size=cfg.hidden_size, eps=cfg.rms_norm_eps, dtype=dtype) + + # KimiK3RMSNorm (not RMSNorm): consumed field-wise (.weight/.eps) + # by _apply_attn_res and the fused attn_res op. + self.output_attn_res_norm = KimiK3RMSNorm( + cfg.hidden_size, eps=cfg.rms_norm_eps, dtype=dtype + ) + self.output_attn_res_proj = nn.Linear(cfg.hidden_size, 1, bias=False, dtype=dtype) + self.num_attn_res_snapshots = ( + cfg.num_hidden_layers + cfg.attn_res_block_size - 1 + ) // cfg.attn_res_block_size + # The decode path's GEMVs and embedding (decode_gemv.py), built by the target's cache_derived_state once the + # weights are final. None: every step embeds on the generic path. + self.decode_gemvs: Optional[_decode_gemv.K3DecodeGemvs] = None + + # Which convention the drafter tap is on is not recoverable from the + # served output -- a mismatch only lowers acceptance -- so state it once + # at construction rather than leaving it to be inferred from an AL. + logger.info_once( + "Kimi K3 aux hidden capture: mode=" + f"{'attn_res_stream' if _AUX_ATTN_RES_STREAM_ENABLED else 'prefix_only'} " + f"({KIMI_K3_AUX_ATTN_RES_STREAM_ENV}={int(_AUX_ATTN_RES_STREAM_ENABLED)})", + key="kimi_k3_aux_capture_mode", + ) + + # >>> route B: no per-token KDA verify states (kda_token_states): no speculative decoding + # <<< route B + + def forward( + self, + attn_metadata: AttentionMetadata, + input_ids: Optional[torch.IntTensor] = None, + position_ids: Optional[torch.IntTensor] = None, + inputs_embeds: Optional[torch.FloatTensor] = None, + spec_metadata=None, + **kwargs, + ) -> torch.Tensor: + if (input_ids is None) ^ (inputs_embeds is not None): + raise ValueError("You must specify exactly one of input_ids or inputs_embeds") + + num_tokens = (input_ids if inputs_embeds is None else inputs_embeds).shape[0] + step = decode_step(attn_metadata, num_tokens) + # A decode step embeds and norms for layer 0 in one launch, the embedding written as layer 0's first snapshot. + prenormed = None + if ( + inputs_embeds is None + and step is not None + and self.decode_gemvs is not None + and len(self.layers) > 0 + and self.num_attn_res_snapshots > 0 + ): + table = self.embed_tokens.weight + block_residual = table.new_empty( + self.num_attn_res_snapshots, num_tokens, table.shape[1] + ) + prenormed = self.decode_gemvs.embed_norm( + input_ids, table, self.layers[0].input_layernorm, block_residual + ) + if prenormed is not None: + hidden_states = prenormed + else: + hidden_states = self.embed_tokens(input_ids) if inputs_embeds is None else inputs_embeds + block_residual = hidden_states.new_empty( + self.num_attn_res_snapshots, + hidden_states.shape[0], + hidden_states.shape[1], + ) + num_snapshots = 0 + capture_set = ( + getattr(spec_metadata, "_capture_layer_set", None) + if spec_metadata is not None + else None + ) + for i, layer in enumerate(self.layers): + # DFlash/DSpark hidden-state capture. The drafter is distilled on + # the aggregated stream value -- the pre-norm softmax mixture its + # next consumer sees -- not on the raw prefix sum a layer returns, + # which is SGLang's fallback for models without the + # attention-residual scheme. Capturing the prefix sum costs 4.5pt + # of draft acceptance on K3 + RadixArk DSpark (AR 66.9% -> 71.4%). + # The tap fires inside layer i+1, which computes that tensor + # anyway; see its forward docstring. Ground truth: SGLang + # kimi_k3.py:2697 _dspark_capture_stream, attn_residual.py:313 + # aggregate_stream_torch. + capture = None + if ( + spec_metadata is not None + and i > 0 + and (capture_set is None or self.layers[i - 1].layer_idx in capture_set) + ): + capture = (spec_metadata, self.layers[i - 1].layer_idx) + hidden_states, num_snapshots = layer( + hidden_states, + block_residual, + num_snapshots, + attn_metadata, + capture=capture, + step=step, + prenormed=i == 0 and prenormed is not None, + ) + + # The last layer has no successor, so this one recompute is + # unavoidable -- output-side score weights, matching SGLang's + # layer_idx + 1 >= end_layer branch. Unreachable for K3's capture set + # against 93 layers; kept so a set that does include the final layer + # gets the right tensor rather than the raw prefix sum. + if spec_metadata is not None and len(self.layers) > 0: + last = self.layers[-1] + if capture_set is None or last.layer_idx in capture_set: + tail = ( + _apply_attn_res( + hidden_states, + block_residual[:num_snapshots], + self.output_attn_res_proj, + self.output_attn_res_norm, + ) + if num_snapshots > 0 and _AUX_ATTN_RES_STREAM_ENABLED + else hidden_states + ) + spec_metadata.maybe_capture_hidden_states(last.layer_idx, tail, None) + + return _apply_attn_res_and_rmsnorm( + hidden_states, + block_residual[:num_snapshots], + self.output_attn_res_proj, + self.output_attn_res_norm, + self.norm, + ) + + +# ---------------------------------------------------------------------------------------------------------------------- +# The attention modules: the built-in KDA and MLA modules, with the steps the decode kernels take on those kernels. +# ---------------------------------------------------------------------------------------------------------------------- + + +class K3DecodeKDA(KimiKDALinearAttention): + # >>> route B: no verify kernels + """Kimi K3's KDA attention: the built-in module, with the decode kernels on the steps they take. + + * A decode step of one token per request runs the fused input projection and the plain decode in one + ``ssm/k3_kda_decode_attn`` launch. + * On every step ``decode_step`` classifies, ``o_proj`` runs on the ``o_proj`` decode GEMV site where it takes the + rows, then the module's all-reduce. + + Every other step runs the built-in module. The kernels read one ``[q | k | v | g | f_a | b]`` weight, built at the + checkpoint load from the module's own, and the device's ``K3KdaBuffers`` and decode GEMVs' state, which the + target sets in ``post_load_weights``. + """ + + # <<< route B + + def __init__(self, *args, **kwargs) -> None: + super().__init__(*args, **kwargs) + # The fused [q | k | v | g | f_a | b | pad] projection weight; the six projections' weights and the built-in + # fused [q | k | v | g] and [f_a | b | pad] ones are views of it. + self.k3_proj_weight: Optional[torch.Tensor] = None + # The fused projection's Lamport buffers: one set per device, shared by every KDA layer. + self.k3_buffers: Optional[K3KdaBuffers] = None + # The decode GEMVs' state (decode_gemv.py), shared by the target's layers. + self.decode_gemvs: Optional[_decode_gemv.K3DecodeGemvs] = None + + @property + def takes_k3_kernels(self) -> bool: + """Whether the decode kernels can run this layer: its fused projection weight is built.""" + return self.k3_proj_weight is not None + + def finalize_decode_weights(self) -> None: + """The built-in fused weights, then one ``[q | k | v | g | f_a | b | pad]`` weight of both.""" + super().finalize_decode_weights() + assert ( + self.use_full_rank_gate + and self.gate_lower_bound is not None + and self._qkvg_proj_weight is not None + and self._bfa_proj_weight is not None + and self._qkvg_proj_weight.dtype == self._bfa_proj_weight.dtype == torch.bfloat16 + ), ( + f"Kimi K3 KDA layer {self.layer_idx}: the decode kernels read the bf16 fused projections the built-in " + "module builds on CUDA at head dim 128, with a full-rank output gate and a gate lower bound" + ) + rows = self._qkvg_proj_weight.shape[0] + with torch.no_grad(): + fused = self._merge_projection_weights( + (self.q_proj, self.k_proj, self.v_proj, self.g_proj, self.f_a_proj, self.b_proj), + pad_rows_to=8, + ) + self.k3_proj_weight = fused + self._qkvg_proj_weight, self._bfa_proj_weight = fused[:rows], fused[rows:] + + def forward( + self, + hidden_states: torch.Tensor, + attn_metadata: AttentionMetadata, + step: Optional[DecodeStep] = None, + ) -> torch.Tensor: + """The built-in forward on a step ``decode_step`` does not classify, and under a breakable CUDA graph. On the + others: the plain decode on ``ssm/k3_kda_decode_attn`` on a decode step of one token per request, else the + built-in dispatch; then ``o_proj`` on its decode GEMV site.""" + if step is None or is_in_breakable_cuda_graph(): + return super().forward(hidden_states, attn_metadata) + if ( + step.decode + and step.tokens_per_request == 1 + and self.takes_k3_kernels + and self.k3_buffers is not None + ): + core = self._k3_decode(hidden_states[: step.num_tokens], attn_metadata) + else: + core = self._forward_impl(hidden_states, attn_metadata) + return self._k3_project_output(core) + + def _k3_project_output(self, core: torch.Tensor) -> torch.Tensor: + """``o_proj`` on the ``o_proj`` decode GEMV site where it takes the rows (else the module), then the TP + all-reduce.""" + out = None + if self.decode_gemvs is not None: + out = self.decode_gemvs.project( + "o_proj", core.reshape(-1, self.proj_size), self.o_proj.weight + ) + if out is None: + return self._project_output(core) + return out if self._o_allreduce is None else self._o_allreduce(out) + + def _k3_decode(self, x: torch.Tensor, attn_metadata: AttentionMetadata) -> torch.Tensor: + """``ssm/k3_kda_decode_attn``: the core output ``[R, H, 128]`` of one token of each of the step's R requests; + each slot's conv window and state advance in place.""" + mamba_metadata = attn_metadata.mamba_metadata + slots = getattr(mamba_metadata, "generation_state_indices", None) + if slots is None: + slots = mamba_metadata.state_indices[: x.shape[0]] + layer_cache = attn_metadata.kv_cache_manager.mamba_layer_cache(self.layer_idx) + w_q, w_k, w_v = self._get_mtp_conv_weights() + core = k3_kda_decode_attn( + x.contiguous(), + self.k3_proj_weight, + self.f_b_proj.weight, + w_q, + w_k, + w_v, + self._A_log_f32, + self._dt_bias_f32, + self._onorm_w_f32, + layer_cache.conv, + layer_cache.temporal, + slots, + self.k3_buffers, + float(self.gate_lower_bound), + self.head_k_dim**-0.5, + float(self.o_norm.eps), + ) + # Speculative decoding's replay caches keep their committed conv window in step with the pool's. + self._sync_kda_replay_conv_window(layer_cache, slots, layer_cache.conv) + return core + + # >>> route B: no verify kernels (a verify step needs speculative decoding) + # <<< route B + + +class K3DecodeMLA(KimiK3MLAAttention): + """Kimi K3's MLA attention: the built-in module, with a decode step's attention on the decode kernels. + + A decode step runs ``x [W_a; W_g]^T`` with the gate columns through a sigmoid (the ``mla_ag`` decode GEMV site + where it takes the rows, else one GEMM), then ``attention/k3_mla_qkv`` (the q_a / kv_a RMSNorms, q_b and the k_b + absorption into the fused query, the step's latent rows stored into the paged cache) and + ``attention/k3_mla_attn_vb_out`` (the attention over the paged cache, v_b and the output gate in one launch), then + ``o_proj`` (the ``o_proj`` site where it takes the rows, else the module). Every other step, and a decode step whose + cache the kernels do not read (``k3_mla_decode_view`` says why), runs the built-in module. + + ``[W_a; W_g]`` is built at load from the module's weights, which become views of it. The attention workspace is + the device's ``K3MlaAttnWorkspace``; it and the decode GEMVs' state are set by the target in + ``post_load_weights``. + """ + + def __init__(self, **kwargs) -> None: + super().__init__(**kwargs) + # [W_a; W_g]: the fused q_a / kv_a projection's rows, then the output gate's. + self.k3_ag_weight: Optional[torch.Tensor] = None + # The decode attention's workspace: one per device, shared by every MLA layer. + self.k3_workspace: Optional[K3MlaAttnWorkspace] = None + # The decode GEMVs' state (decode_gemv.py), shared by the target's layers. + self.decode_gemvs: Optional[_decode_gemv.K3DecodeGemvs] = None + + def post_load_weights(self) -> None: + """The built-in post-load, then ``[W_a; W_g]`` (once: CUDA graphs captured since read it).""" + super().post_load_weights() + if self.k3_ag_weight is not None: + return + gaps = self._k3_layout_gaps() + assert not gaps, ( + f"Kimi K3 MLA layer {self.layer_idx}: the decode kernels do not take {'; '.join(gaps)}" + ) + qkv_a, gate = self.kv_a_proj_with_mqa, self.g_proj + rows = qkv_a.weight.shape[0] + with torch.no_grad(): + fused = torch.cat([qkv_a.weight, gate.weight]) + qkv_a.weight = nn.Parameter(fused[:rows], requires_grad=False) + gate.weight = nn.Parameter(fused[rows:], requires_grad=False) + self.k3_ag_weight = fused + + def _k3_layout_gaps(self) -> list: + """What of this layer the decode kernels do not take (empty when they take all of it).""" + if not (self.use_output_gate and self.fuse_qkv_a_proj and not self.is_lite): + return ["a layer without the output gate or the fused q_a / kv_a projection"] + linears = (self.kv_a_proj_with_mqa, self.g_proj, self.q_b_proj, self.o_proj) + checks = ( + (not self.mapping.has_cp_helix(), "helix context parallelism"), + (not self.apply_rotary_emb and not self.llama_4_scaling, "RoPE or llama-4 scaling"), + (self.sparse_attn_hooks is None, "sparse attention"), + ( + self.kv_cache_dtype != "fp8_ds_mla" + and not getattr(self.mqa, "has_fp8_kv_cache", False) + and not getattr(self.mqa, "has_fp4_kv_cache", False), + "a quantized KV cache", + ), + ( + all(m.weight.dtype == torch.bfloat16 and m.bias is None for m in linears) + and self.k_b_proj_trans.dtype == self.v_b_proj.dtype == torch.bfloat16, + "projections other than bf16 and unbiased", + ), + ( + not getattr(self.q_a_layernorm, "is_nvfp4", False) + and not getattr(self.kv_a_layernorm, "use_gemma", False), + "an NVFP4 q_a norm or a Gemma kv_a norm", + ), + ( + self.num_heads_tp % 6 == 0 + and self.kv_lora_rank == 512 + and self.qk_rope_head_dim == 64 + and self.q_lora_rank == 1536, + f"{self.num_heads_tp} heads, latent {self.kv_lora_rank}, rope {self.qk_rope_head_dim}, " + f"q_lora {self.q_lora_rank}", + ), + ) + return [why for ok, why in checks if not ok] + + def forward( + self, + position_ids: Optional[torch.Tensor], + hidden_states: torch.Tensor, + attn_metadata: AttentionMetadata, + all_reduce_params=None, + latent_cache_gen: Optional[torch.Tensor] = None, + step: Optional[DecodeStep] = None, + ) -> torch.Tensor: + """The built-in forward, except on a decode step whose cache the decode kernels read.""" + view = None + if ( + step is not None + and step.decode + and latent_cache_gen is None + and self.k3_ag_weight is not None + and self.k3_workspace is not None + and not is_in_breakable_cuda_graph() + ): + view = self._k3_decode_view(attn_metadata, step.num_tokens) + if view is None: + return super().forward( + position_ids, hidden_states, attn_metadata, all_reduce_params, latent_cache_gen + ) + x = hidden_states[: step.num_tokens].contiguous() + rows = self.kv_a_proj_with_mqa.weight.shape[0] + gemvs = self.decode_gemvs + ag = None if gemvs is None else gemvs.project("mla_ag", x, self.k3_ag_weight) + if ag is None: + ag = torch.nn.functional.linear(x, self.k3_ag_weight) + ag[:, rows:].sigmoid_() + fused_q = k3_mla_qkv( + ag, + self.q_a_layernorm.weight, + float(self.q_a_layernorm.variance_epsilon), + self.q_b_proj.weight, + self.k_b_proj_trans, + self.kv_a_layernorm.weight, + float(self.kv_a_layernorm.variance_epsilon), + view["pool"], + view["row_stride"], + view["page_table"], + view["page_offset"], + view["seq_len"], + ) + attn_output = self.create_output(x, 0) + k3_mla_attn_vb_out( + fused_q, + view["pool"], + view["row_stride"], + view["page_table"], + view["page_offset"], + view["seq_len"], + view["softmax_scale"], + self.v_b_proj, + attn_output, + self.k3_workspace, + gate=ag, + gate_col0=rows, + ) + out = None if gemvs is None else gemvs.project("o_proj", attn_output, self.o_proj.weight) + if out is None: + out = self._project_output( + [attn_output], position_ids, attn_metadata, all_reduce_params + ) + return out + + def _k3_decode_view(self, attn_metadata: AttentionMetadata, num_tokens: int) -> Optional[dict]: + """The paged latent cache as the decode kernels read it this step, or None when they do not read it (the + reason is logged once).""" + view = k3_mla_decode_view(self.mqa, attn_metadata, num_tokens) + if isinstance(view, str): + logger.info_once( + f"Kimi K3 MLA: the built-in path for a decode step the decode kernels do not read ({view})", + key=f"k3_mla_decode_view_{view}", + ) + return None + return view + + +# ---------------------------------------------------------------------------------------------------------------------- +# The target: step classification, the construction checks and the registration shell. +# ---------------------------------------------------------------------------------------------------------------------- + + def _text_model_config(model_config: ModelConfig) -> ModelConfig: """The language model's ModelConfig: the checkpoint's text_config, with quant exclusions renamed to match. @@ -127,6 +2083,67 @@ def _text_model_config(model_config: ModelConfig) -> ModelConfig: return text +@dataclass(frozen=True) +class DecodeStep: + """A step the Kimi K3 decode kernels take: ``num_tokens`` rows and, on a pure decode step, ``num_requests`` + generation requests of ``tokens_per_request`` tokens each (None on a step with context requests).""" + + num_tokens: int + num_requests: Optional[int] = None + tokens_per_request: Optional[int] = None + + @property + def small(self) -> bool: + """Whether the step fits one token tile of the token-count kernels.""" + return self.num_tokens <= DECODE_MAX_TOKENS + + @property + def decode(self) -> bool: + """Whether the step is a pure decode step the request-aware kernels take.""" + return self.num_requests is not None + + @property + def wide(self) -> bool: + """Whether the step is a pure decode step of more than one token tile: its token-count work keeps the decode + layout's MoE head and tail, on M-general ops.""" + return self.decode and not self.small + + +def decode_step(attn_metadata: AttentionMetadata, num_tokens: int) -> Optional[DecodeStep]: + """The step's shape if any Kimi K3 decode kernel takes it, else None (the generic path runs). + + ``num_tokens`` is the step's token count (the rows of the model input). Read on the host from per-step integers + and the host copy of the sequence lengths only. A CUDA graph is captured per decode batch shape, and every input + here is fixed by that shape, so a captured step and its replays are classified alike. + """ + if num_tokens <= 0: + return None + requests = _decode_requests(attn_metadata, num_tokens) + if requests is not None: + return DecodeStep(num_tokens, requests, num_tokens // requests) + if num_tokens <= DECODE_MAX_TOKENS: + return DecodeStep(num_tokens) + return None + + +def _decode_requests(attn_metadata: AttentionMetadata, num_tokens: int) -> Optional[int]: + """R when the step is R <= 8 generation requests of the same T <= 8 tokens and no context request, else None.""" + if attn_metadata.num_contexts != 0: + return None + requests = attn_metadata.num_generations + if not 0 < requests <= MAX_REQUESTS or num_tokens % requests != 0: + return None + tokens = num_tokens // requests + if not 0 < tokens <= MAX_TOKENS_PER_REQUEST: + return None + seq_lens = getattr(attn_metadata, "seq_lens", None) + if seq_lens is not None and seq_lens.device.type == "cpu": + lens = seq_lens[:requests] + if lens.numel() != requests or bool((lens != tokens).any()): + return None + return requests + + def _check_construction(model_config: ModelConfig) -> None: """The settings this target is built for that are fixed before the first step.""" capability = torch.cuda.get_device_capability() @@ -142,19 +2159,22 @@ def _check_construction(model_config: ModelConfig) -> None: mapping.moe_ep_size, mapping.enable_attention_dp, ) + # >>> route B: this target's topology, its expert split set explicitly, and no speculative decoding assert topology == (16, 16, 1, 16, 1, False), ( "the tp16_moetp16ep1 target needs world_size 16, tensor_parallel_size 16, pipeline_parallel_size 1, " "moe_tensor_parallel_size 16, moe_expert_parallel_size 1 and enable_attention_dp false; the engine built " f"(world, tp, pp, moe_tp, moe_ep, attention_dp) = {topology}" ) - assert getattr(mapping, "moe_tp_ep_user_specified", False), ( - "the tp16_moetp16ep1 target needs moe_tensor_parallel_size 16 and moe_expert_parallel_size 1 set explicitly; " - "with the split unset the built-in Kimi K3 model runs the experts expert-parallel over the 16 ranks" + assert mapping.moe_tp_ep_user_specified, ( + "the tp16_moetp16ep1 target needs moe_tensor_parallel_size 16 and moe_expert_parallel_size 1 set " + "explicitly; with the split unset the engine resolves the same sizes, but Kimi K3 then runs the experts " + "expert-parallel over the 16 ranks" ) - assert model_config.spec_config is None, ( - "the tp16_moetp16ep1 target decodes without speculation; DSpark runs on the tp16_moetp4ep4 target " - f"(moe_tensor_parallel_size 4, moe_expert_parallel_size 4); got {type(model_config.spec_config).__name__}" + assert getattr(model_config, "spec_config", None) is None, ( + "the tp16_moetp16ep1 target decodes without speculation; for DSpark, DFlash or SA set " + "moe_tensor_parallel_size 4 and moe_expert_parallel_size 4 (the tp16_moetp4ep4 target)" ) + # <<< route B assert model_config.torch_dtype == torch.bfloat16, ( f"this target computes in bf16; the engine resolved dtype {model_config.torch_dtype}" ) @@ -168,9 +2188,16 @@ def _check_construction(model_config: ModelConfig) -> None: ) +# >>> route B: this target's name @register_auto_model("ModelingV2KimiK3Mxfp4Sm100Tp16Moetp16ep1") class ModelingV2KimiK3Mxfp4Sm100Tp16Moetp16ep1(KimiLinearForCausalLM): - """The registration shell: the built-in Kimi K3 text model as the generic path, behind this target's checks.""" + # <<< route B + """The registration shell: this target's text model (`KimiLinearModel` above) behind its checks. + + It inherits the built-in Kimi K3 causal LM for the checkpoint load and the engine hooks (`load_weights` through + weights.py, the KDA metadata class, the model defaults), which walk the model by its module names; the text model + keeps the built-in one's. + """ @classmethod def get_preferred_kv_cache_manager_version(cls, pretrained_config: Any = None) -> Literal["V2"]: @@ -184,11 +2211,27 @@ def __init__(self, model_config: ModelConfig): "text_config" ) _check_construction(model_config) - super().__init__(_text_model_config(model_config)) + text = _text_model_config(model_config) + # >>> route B: no speculative decoding (_check_construction), so the one-engine shell builds no drafter + # <<< route B + # The inherited loader reads these: this target has neither helix context parallelism nor the fp8 + # weight-read conversion of the shared / latent MLPs. + self._fp8_weight_read_moe_mlp = False + self.mapping_with_cp = None + self._repurposed_tp_mapping = None + cfg = text.pretrained_config + SpecDecOneEngineForCausalLM.__init__( + self, + KimiLinearModel(text), + text, + hidden_size=cfg.hidden_size, + vocab_size=cfg.vocab_size, + ) self._step_checked = False - # The fused decode path and the state its kernels share, built in post_load_weights once the catalog entries - # it calls exist. None: every step takes the generic path. - self._fused_decode = None + # >>> route B: no drafter to hand the LM head to + # The LM head on gemm/k3_head_gemv at decode size, once cache_derived_state has built its state. + self.logits_processor = _decode_gemv.K3LogitsProcessor(self.logits_processor) + # <<< route B # The executor reads generation settings (eos_token_id, ...) off the model config the engine holds, which # must therefore be the text config, as the built-in wrapper leaves it. model_config._frozen = False @@ -198,6 +2241,46 @@ def __init__(self, model_config: ModelConfig): def load_weights(self, weights, *args, **kwargs): _weights.load(self, weights) + def cache_derived_state(self) -> None: + """Build the decode GEMVs' state once the weights are final: the LM head's workspace, and one eager call of + every decode GEMV kernel at its site's shape (decode_gemv.SITES), so none compiles under capture. Built once: + a later call keeps it, since CUDA graphs captured in between hold its workspace.""" + super().cache_derived_state() + if self.model.decode_gemvs is not None: + return + gemvs = _decode_gemv.K3DecodeGemvs.create(self.lm_head) + self.model.decode_gemvs = gemvs + self.logits_processor.gemvs = gemvs + for layer in self.model.layers: + if not layer.is_moe: + layer.decode_gemvs = gemvs + + def post_load_weights(self) -> None: + """The state the decode kernels share, built once per device before any CUDA-graph capture and handed to + every layer of its kind: the KDA projection's Lamport buffers and the MLA decode attention's workspace; and + the decode GEMVs' state (built by ``cache_derived_state``) handed to every attention module.""" + super().post_load_weights() + kda = [layer.linear_attn for layer in self.model.layers if layer.is_kda] + mla = [layer.self_attn.mixer for layer in self.model.layers if not layer.is_kda] + if kda[0].k3_buffers is not None: + return # built by an earlier call; CUDA graphs captured since hold it + device = self.model.embed_tokens.weight.device + buffers = K3KdaBuffers.create(device) + workspace = K3MlaAttnWorkspace.create(device, mla[0].num_heads_tp // 6) + for module in kda: + module.k3_buffers = buffers + for module in mla: + module.k3_workspace = workspace + for module in kda + mla: + module.decode_gemvs = self.model.decode_gemvs + logger.info( + # >>> route B: no verify kernels + "Kimi K3 decode kernels: KDA on k3_kda_decode_attn " + # <<< route B + f"({sum(m.takes_k3_kernels for m in kda)} / {len(kda)} layers take them), MLA on k3_mla_qkv and " + f"k3_mla_attn_vb_out ({len(mla)} layers)" + ) + def _check_step_contract(self, attn_metadata: AttentionMetadata) -> None: """First-forward checks of the engine surface and the per-engine settings.""" objects = { @@ -224,18 +2307,14 @@ def _check_step_contract(self, attn_metadata: AttentionMetadata) -> None: assert not manager.enable_block_reuse, ( "this target runs with kv_cache_config.enable_block_reuse false; the engine enabled it" ) + kda_layer = next(layer.layer_idx for layer in self.model.layers if layer.is_kda) + state_dtype = manager.mamba_layer_cache(kda_layer).temporal.dtype + assert state_dtype == torch.float32, ( + "this target's KDA kernels keep fp32 recurrent states (kv_cache_config.mamba_ssm_cache_dtype float32 or " + f"auto); the engine built a {state_dtype} state pool" + ) self._step_checked = True - def _step_path(self, attn_metadata: AttentionMetadata, spec_metadata) -> str: - """`"fused"` for a pure decode step within the fused path's bounds, else `"generic"`. - - Read on the host from per-step integers only. A CUDA graph is captured per decode batch shape, and every - input here is fixed by that shape, so a captured step and its replays take the same path. - """ - if attn_metadata.num_contexts: - return "generic" - return "fused" if attn_metadata.num_tokens <= _FUSED_MAX_TOKENS else "generic" - def forward( self, attn_metadata: AttentionMetadata, @@ -253,18 +2332,6 @@ def forward( ) if not self._step_checked: self._check_step_contract(attn_metadata) - if ( - self._fused_decode is not None - and self._step_path(attn_metadata, spec_metadata) == "fused" - ): - return self._fused_decode( - attn_metadata=attn_metadata, - input_ids=input_ids, - position_ids=position_ids, - spec_metadata=spec_metadata, - resource_manager=resource_manager, - **kwargs, - ) return super().forward( attn_metadata, input_ids, diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/weights.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/weights.py index e4c24eaaea22..66a1f4f72049 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/weights.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/weights.py @@ -1,5 +1,6 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 +# >>> route B: this target's expert slice """Weight loading: Kimi K3 (MXFP4) / sm_100 / tp16_moetp16ep1. The checkpoint is the vision-language wrapper's. `language_model.*` holds the language model; `vision_tower.*` and @@ -11,6 +12,7 @@ * The vision tower and the projector are a predicted non-load: listed here, never read. * Any other key fails the load, naming it, rather than being dropped. """ +# <<< route B from tensorrt_llm._torch.models.checkpoints.base_weight_loader import ConsumableWeightsDict from tensorrt_llm._torch.models.modeling_kimi_linear import KimiLinearForCausalLM diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_construction.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_construction.py new file mode 100644 index 000000000000..0763b4667356 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_construction.py @@ -0,0 +1,72 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""The Kimi K3 targets' construction checks, host-side: each target requires its own parallel layout, and +``tp16_moetp16ep1`` also an explicit expert split and no speculative decoding.""" + +import types + +import pytest +import torch + +from tensorrt_llm._torch._experimental.modeling_v2.models.kimi_k3_vl.kimi_k3_mxfp4__sm_100__tp16_moetp4ep4 import ( # noqa: E501 + modeling as route_a, +) +from tensorrt_llm._torch._experimental.modeling_v2.models.kimi_k3_vl.kimi_k3_mxfp4__sm_100__tp16_moetp16ep1 import ( # noqa: E501 + modeling as route_b, +) +from tensorrt_llm.functional import AllReduceStrategy + + +def _config(moe_tp, moe_ep, split_set=True, spec_config=None, attention_dp=False): + """The ModelConfig fields the construction checks read.""" + mapping = types.SimpleNamespace( + world_size=16, + tp_size=16, + pp_size=1, + moe_tp_size=moe_tp, + moe_ep_size=moe_ep, + enable_attention_dp=attention_dp, + moe_tp_ep_user_specified=split_set, + ) + return types.SimpleNamespace( + mapping=mapping, + spec_config=spec_config, + torch_dtype=torch.bfloat16, + quant_config=types.SimpleNamespace(quant_algo=None, kv_cache_quant_algo=None), + quant_config_dict=None, + allreduce_strategy=AllReduceStrategy.AUTO, + ) + + +@pytest.fixture(autouse=True) +def sm_100(monkeypatch): + """The targets' architecture, on any host.""" + monkeypatch.setattr(torch.cuda, "get_device_capability", lambda *args, **kwargs: (10, 0)) + + +@pytest.mark.parametrize( + "target,parallel,split,other", + [ + (route_a, "tp16_moetp4ep4", (4, 4), (16, 1)), + (route_b, "tp16_moetp16ep1", (16, 1), (4, 4)), + ], + ids=["tp16_moetp4ep4", "tp16_moetp16ep1"], +) +def test_each_target_requires_its_layout(target, parallel, split, other): + target._check_construction(_config(*split)) + with pytest.raises(AssertionError, match=f"the {parallel} target needs"): + target._check_construction(_config(*other)) + with pytest.raises(AssertionError, match="enable_attention_dp false"): + target._check_construction(_config(*split, attention_dp=True)) + + +def test_tp16_moetp16ep1_requires_the_split_set_explicitly(): + with pytest.raises(AssertionError, match="set explicitly"): + route_b._check_construction(_config(16, 1, split_set=False)) + + +def test_tp16_moetp16ep1_decodes_without_speculation(): + with pytest.raises( + AssertionError, match="moe_tensor_parallel_size 4 and moe_expert_parallel_size 4" + ): + route_b._check_construction(_config(16, 1, spec_config=object())) diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drift.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drift.py new file mode 100644 index 000000000000..716249d68953 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drift.py @@ -0,0 +1,135 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""The Kimi K3 `tp16_moetp16ep1` target is a copy of `tp16_moetp4ep4`'s files, changed only in marked blocks. + +Targets share no files (test_modeling_v2_claims.py), so `tp16_moetp16ep1` (route B) ships its own copy of every +module of `tp16_moetp4ep4` (route A). Each place the copy differs is a block: + + # >>> route B: + + # <<< route B + +This reads both targets' files and imports neither. It fails on any difference outside the blocks, so a change to +route A's modules reaches route B's copy, or becomes one of its blocks, in the same commit. +""" + +from __future__ import annotations + +import difflib +from pathlib import Path +from typing import List, Tuple + +import pytest + +import tensorrt_llm._torch._experimental.modeling_v2 as _modeling_v2 + +# The package, not this file: the targets live under tensorrt_llm/ while this test lives under tests/. +_FAMILY = Path(_modeling_v2.__file__).resolve().parent / "models" / "kimi_k3_vl" +_ROUTE_A = _FAMILY / "kimi_k3_mxfp4__sm_100__tp16_moetp4ep4" +_ROUTE_B = _FAMILY / "kimi_k3_mxfp4__sm_100__tp16_moetp16ep1" + +_BEGIN = "# >>> route B" +_END = "# <<< route B" + +# Each target's package marker holds only its own one-line description. +_NOT_COPIED = {"__init__.py"} + + +def _modules(target: Path) -> List[str]: + return sorted(p.name for p in target.glob("*.py") if p.name not in _NOT_COPIED) + + +def _blocks(lines: List[str]) -> List[Tuple[int, int]]: + """The route B blocks as (begin, end) indices of their marker lines; asserts the markers are well formed.""" + blocks, begin = [], None + for i, line in enumerate(lines): + text = line.strip() + if text.startswith(_BEGIN): + assert begin is None, ( + f"line {i + 1}: a block opens inside the one opened on line {begin + 1}" + ) + assert text[len(_BEGIN) :].startswith(":") and text[len(_BEGIN) + 1 :].strip(), ( + f"line {i + 1}: a block opens with '{_BEGIN}: '" + ) + begin = i + elif text.startswith(_END): + assert text == _END, f"line {i + 1}: a block closes with '{_END}' alone" + assert begin is not None, f"line {i + 1}: a block closes without one open" + blocks.append((begin, i)) + begin = None + assert begin is None, f"line {begin + 1}: a block is never closed" + return blocks + + +def _spans(lines: List[str], blocks: List[Tuple[int, int]]) -> List[Tuple[int, int]]: + """Each block with the blank lines around it, which the formatters add and drop next to a comment.""" + spans = [] + for begin, end in blocks: + while begin > 0 and not lines[begin - 1].strip(): + begin -= 1 + while end + 1 < len(lines) and not lines[end + 1].strip(): + end += 1 + spans.append((begin, end)) + return spans + + +def _inside(spans: List[Tuple[int, int]], j1: int, j2: int) -> bool: + """Whether route B's lines [j1, j2), or the point between lines j1 - 1 and j1 when j1 == j2, are in one block.""" + if j1 == j2: + return any(begin < j1 <= end for begin, end in spans) + return any(begin <= j1 and j2 <= end + 1 for begin, end in spans) + + +def _drift(a: List[str], b: List[str]) -> List[str]: + """Each difference between route A's lines ``a`` and route B's ``b`` outside route B's blocks, as text.""" + assert not _blocks(a) and not any(_BEGIN in line or _END in line for line in a), ( + "route A's files have no route B markers" + ) + blocks = _spans(b, _blocks(b)) + drift = [] + matcher = difflib.SequenceMatcher(None, a, b, autojunk=False) + for tag, i1, i2, j1, j2 in matcher.get_opcodes(): + if tag == "equal" or _inside(blocks, j1, j2): + continue + shown = [f" route A {i + 1}: {a[i]}" for i in range(i1, min(i2, i1 + 5))] + shown += [f" route B {j + 1}: {b[j]}" for j in range(j1, min(j2, j1 + 5))] + drift.append( + f"route A lines {i1 + 1}-{i2}, route B lines {j1 + 1}-{j2}:\n" + "\n".join(shown) + ) + return drift + + +def test_route_b_copies_every_module(): + assert _modules(_ROUTE_A), f"no modules in {_ROUTE_A}" + assert _modules(_ROUTE_B) == _modules(_ROUTE_A) + + +@pytest.mark.parametrize("name", _modules(_ROUTE_A)) +def test_route_b_matches_route_a_outside_its_blocks(name): + a = (_ROUTE_A / name).read_text(encoding="utf-8").splitlines() + b = (_ROUTE_B / name).read_text(encoding="utf-8").splitlines() + drift = _drift(a, b) + assert not drift, ( + f"{name}: route B's copy differs from route A's outside its blocks. Carry route A's " + "change into the copy, or mark route B's lines as a block that says why they differ:\n" + + "\n".join(drift) + ) + + +def test_the_check_sees_drift_only_outside_the_blocks(): + a = ["x = 1", "y = 2", "z = 3"] + changed = ["x = 1", "# >>> route B: why", "y = 4", "# <<< route B", "z = 3"] + dropped = ["x = 1", "", "# >>> route B: why", "# <<< route B", "", "z = 3"] + unmarked = ["x = 1", "# >>> route B: why", "y = 2", "# <<< route B", "z = 4"] + added = ["x = 1", "y = 2", "w = 0", "z = 3"] + spaced = ["x = 1", "", "y = 2", "z = 3"] + assert not _drift(a, changed) + assert not _drift(a, dropped) + assert _drift(a, unmarked) + assert _drift(a, added) + assert _drift(a, spaced) + assert _drift(a[:2], a) + with pytest.raises(AssertionError, match="never closed"): + _drift(a, ["x = 1", "# >>> route B: why", "y = 2", "z = 3"]) + with pytest.raises(AssertionError, match="why it differs"): + _drift(a, ["x = 1", "# >>> route B", "y = 2", "# <<< route B", "z = 3"]) From 03df8ee849fe1e11b93265718e4be75df7bf140d Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 22:03:21 -0700 Subject: [PATCH 083/161] [None][fix] modeling_v2 Kimi K3 target: route the MXFP4 checkpoint as W4A16_MXFP4 KimiK3Config surfaces the checkpoint's compressed-tensors declaration (text_config.quantization_config, mxfp4-pack-quantized), and the model config reads it as W4A16_MXFP4 with no per-layer declarations. The routing criterion and the target's construction assert expected no quantization, so TRTLLM_MODELING_V2=require refused the checkpoint. The routing tests now derive the quantization through KimiK3Config and ModelConfig.load_hf_quant_config from the checkpoint's declaration, and an unquantized checkpoint of the same shape does not match. Signed-off-by: Vasanth Sabavat --- .../modeling.py | 13 +-- .../modeling_v2/models/kimi_k3_vl/routing.py | 16 ++-- .../modeling_v2/test_modeling_v2_routing.py | 91 ++++++++++++++++--- 3 files changed, 93 insertions(+), 27 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index dc8a27202f56..ca8876712688 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -46,8 +46,9 @@ experts, the sandwiches and the residual epilogues come with their own entries; until then they run the generic path on every step. -**What this target asserts rather than adapts**: SM 10.0; the topology above; the MXFP4 checkpoint's quantization (none -the model config reads, so the routed experts keep the MXFP4 default); bf16 weights and a bf16 KV pool; +**What this target asserts rather than adapts**: SM 10.0; the topology above; the MXFP4 checkpoint's quantization +(W4A16_MXFP4 with no per-layer declarations, so the routed experts run the W4A8_MXFP4_MXFP8 default and the excluded +modules stay bf16); bf16 weights and a bf16 KV pool; tokens_per_block 64 (the MLA generation kernels K3's 96 heads reach exist only at 64); the V2 hybrid KV / state manager, which holds the KDA states, with block reuse off and fp32 recurrent states; an all-reduce strategy of AUTO or MNNVL. The construction-time ones fail in `__init__`, the per-engine ones on the first forward, each naming the @@ -2158,10 +2159,10 @@ def _check_construction(model_config: ModelConfig) -> None: f"this target's MLA kernels read a bf16 KV pool; kv_cache_config.dtype resolved to {kv_algo}" ) quant_algo = model_config.quant_config.quant_algo - assert quant_algo is None and not model_config.quant_config_dict, ( - "this target loads the MXFP4 checkpoint, which declares no quantization the model config reads (its routed " - f"experts keep the W4A8_MXFP4_MXFP8 default); the engine read {quant_algo} with " - f"{len(model_config.quant_config_dict or {})} per-layer declarations" + assert quant_algo == QuantAlgo.W4A16_MXFP4 and not model_config.quant_config_dict, ( + "this target loads the MXFP4 checkpoint, whose compressed-tensors config the model config reads as " + "W4A16_MXFP4 with no per-layer declarations (its routed experts run the W4A8_MXFP4_MXFP8 default); the " + f"engine read {quant_algo} with {len(model_config.quant_config_dict or {})} per-layer declarations" ) strategy = model_config.allreduce_strategy assert strategy in (AllReduceStrategy.AUTO, AllReduceStrategy.MNNVL), ( diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/routing.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/routing.py index 048c790a0b84..c406423e441d 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/routing.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/routing.py @@ -16,6 +16,8 @@ from typing import Any, Optional +from tensorrt_llm.quantization.mode import QuantAlgo + from ..._router_index import NULL_TRACE, ModelingV2Context, Trace # The one GPU architecture these targets are written for. sm is part of a @@ -55,13 +57,15 @@ def _field(config: Any, name: str) -> Any: def _mxfp4(quant_config: Any) -> bool: """Whether the checkpoint is the MXFP4 one these targets load. - The MXFP4 checkpoint declares its quantization only inside - ``text_config.quantization_config`` (compressed-tensors), which the model - config does not surface, so it reads as no quantization at all. The NVFP4 - requant ships ``hf_quant_config.json`` and reads as MIXED_PRECISION; its - experts would land in a loader that reads packed MXFP4 tensors. + The MXFP4 checkpoint declares its quantization inside + ``text_config.quantization_config`` (compressed-tensors, + ``mxfp4-pack-quantized``); ``KimiK3Config`` surfaces it and the model + config reads it as W4A16_MXFP4. The NVFP4 requant ships + ``hf_quant_config.json`` and reads as MIXED_PRECISION, and a checkpoint + without quantization reads as none; the experts of either would land in a + loader that reads packed MXFP4 tensors. """ - return quant_config is None or quant_config.quant_algo is None + return quant_config is not None and quant_config.quant_algo == QuantAlgo.W4A16_MXFP4 def _parallel(m) -> Optional[str]: diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py index 9728ba55d985..25fe8e786b66 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py @@ -28,6 +28,7 @@ ModelingV2Mode, modeling_v2_resolve, ) +from tensorrt_llm._torch.configs.kimi_k3 import KimiK3Config from tensorrt_llm._torch.model_config import ModelConfig from tensorrt_llm._torch.models.modeling_auto import AutoModelForCausalLM from tensorrt_llm._torch.models.modeling_utils import ( @@ -35,6 +36,8 @@ get_registered_model_class, ) from tensorrt_llm.mapping import Mapping +from tensorrt_llm.models.modeling_utils import QuantConfig +from tensorrt_llm.quantization.mode import QuantAlgo _SM103 = (10, 3) @@ -260,41 +263,99 @@ def _k3_config(text_as_dict=False, **text_overrides): ) +# The quantization the MXFP4 checkpoint declares in its config.json, inside ``text_config``. +_K3_CHECKPOINT_QUANTIZATION = { + "quant_method": "compressed-tensors", + "format": "mxfp4-pack-quantized", + "config_groups": { + "group_0": { + "format": "mxfp4-pack-quantized", + "input_activations": None, + "output_activations": None, + "targets": ["Linear"], + "weights": { + "group_size": 32, + "num_bits": 4, + "strategy": "group", + "symmetric": True, + "type": "float", + }, + } + }, + "ignore": [ + "re:.*self_attn.*", + "re:.*shared_experts.*", + r"re:.*mlp\.(gate|up|gate_up|down)_proj.*", + "re:.*lm_head.*", + "re:.*vision_tower.*", + "re:.*mm_projector.*", + ], + "kv_cache_scheme": None, +} + + +def _k3_checkpoint_quant_config(): + """What the engine reads for the MXFP4 checkpoint: ``KimiK3Config`` surfaces the text config's + declaration, and ``ModelConfig`` parses it.""" + config = KimiK3Config( + text_config=dict(quantization_config=_K3_CHECKPOINT_QUANTIZATION), + architectures=["KimiK3ForConditionalGeneration"], + ) + quant_config, layer_quant_config = ModelConfig.load_hf_quant_config( + config.quantization_config, "TRTLLM" + ) + assert layer_quant_config is None + return quant_config + + +def _k3_model_config(pretrained_config, quant_config=None, **mapping_kwargs): + """A model config of the MXFP4 checkpoint's quantization (or ``quant_config``).""" + return ModelConfig( + pretrained_config=pretrained_config, + mapping=Mapping(**mapping_kwargs), + quant_config=quant_config if quant_config is not None else _k3_checkpoint_quant_config(), + ) + + @pytest.fixture def _on_sm100(monkeypatch): """Route as if this were a GB200.""" monkeypatch.setattr(torch.cuda, "get_device_capability", lambda *a, **k: _SM100) +def test_kimi_k3_checkpoint_reads_as_w4a16_mxfp4(): + quant_config = _k3_checkpoint_quant_config() + assert quant_config.quant_algo == QuantAlgo.W4A16_MXFP4 + assert quant_config.group_size == 32 + + @pytest.mark.usefixtures("_on_sm100") @pytest.mark.parametrize("mode", ["auto", "require"]) @pytest.mark.parametrize("text_as_dict", [False, True], ids=["text-config", "text-dict"]) def test_kimi_k3_tp16_moetp4ep4_matches(monkeypatch, mode, text_as_dict): _set_mode(monkeypatch, mode) - config = _model_config(_k3_config(text_as_dict=text_as_dict), **_TP16_MOETP4EP4) + config = _k3_model_config(_k3_config(text_as_dict=text_as_dict), **_TP16_MOETP4EP4) assert modeling_v2_resolve(config) == _K3_TARGET @pytest.mark.usefixtures("_on_sm100") def test_kimi_k3_target_registers_and_counts_as_external(): - config = _model_config(_k3_config(), **_TP16_MOETP4EP4) + config = _k3_model_config(_k3_config(), **_TP16_MOETP4EP4) cls = get_registered_model_class(modeling_v2_resolve(config)) assert cls is not None and cls.__name__ == _K3_TARGET assert not _is_builtin_model_class(cls) @pytest.mark.usefixtures("_on_sm100") -def test_kimi_k3_nvfp4_requant_does_not_match(monkeypatch): - """The NVFP4 requant has the MXFP4 checkpoint's shape; its quantization - is what keeps it out of a target whose loader reads packed MXFP4 experts.""" - from tensorrt_llm.models.modeling_utils import QuantConfig - from tensorrt_llm.quantization.mode import QuantAlgo - - config = ModelConfig( - pretrained_config=_k3_config(), - mapping=Mapping(**_TP16_MOETP4EP4), - quant_config=QuantConfig(quant_algo=QuantAlgo.MIXED_PRECISION), - ) +@pytest.mark.parametrize( + "quant_algo", + [QuantAlgo.MIXED_PRECISION, None], + ids=["nvfp4-requant", "unquantized"], +) +def test_kimi_k3_other_quantizations_do_not_match(monkeypatch, quant_algo): + """The NVFP4 requant (MIXED_PRECISION) and an unquantized checkpoint have the MXFP4 checkpoint's shape; the + quantization is what keeps them out of a target whose loader reads packed MXFP4 experts.""" + config = _k3_model_config(_k3_config(), QuantConfig(quant_algo=quant_algo), **_TP16_MOETP4EP4) assert modeling_v2_resolve(config) is None _set_mode(monkeypatch, "require") @@ -317,7 +378,7 @@ def test_kimi_k3_nvfp4_requant_does_not_match(monkeypatch): ], ) def test_kimi_k3_near_misses_do_not_match(monkeypatch, text_overrides, mapping_kwargs, missed): - config = _model_config(_k3_config(**text_overrides), **mapping_kwargs) + config = _k3_model_config(_k3_config(**text_overrides), **mapping_kwargs) assert modeling_v2_resolve(config) is None _set_mode(monkeypatch, "require") @@ -327,7 +388,7 @@ def test_kimi_k3_near_misses_do_not_match(monkeypatch, text_overrides, mapping_k def test_kimi_k3_on_another_sm_does_not_match(monkeypatch): """The autouse fixture routes as a GB300 (sm 10.3); the K3 target is sm 10.0 only.""" - config = _model_config(_k3_config(), **_TP16_MOETP4EP4) + config = _k3_model_config(_k3_config(), **_TP16_MOETP4EP4) assert modeling_v2_resolve(config) is None _set_mode(monkeypatch, "require") From fa24aca986ee57fd3582a29262d100afe18ebcc8 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 22:06:51 -0700 Subject: [PATCH 084/161] [None][chore] modeling_v2 Kimi K3 tp16_moetp16ep1: carry the trim and the W4A16_MXFP4 fix into the copy The tp16_moetp4ep4 target's text model is trimmed to its settings (f01a46b67f), and its construction assert reads the MXFP4 checkpoint as W4A16_MXFP4 (03df8ee849). The copy here takes the same changes. Two of them name the layout: the routed-expert split's docstring (16 x 1, a new route B block) and the module docstring's list of asserted settings (inside its block). The construction test's model config carries the checkpoint's quantization, W4A16_MXFP4. Signed-off-by: Vasanth Sabavat --- .../modeling.py | 179 +++++------------- .../test_modeling_v2_kimi_k3_construction.py | 6 +- 2 files changed, 50 insertions(+), 135 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py index 4c40d44b28d3..638ccc2bb741 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py @@ -50,9 +50,11 @@ step. **What this target asserts rather than adapts**: SM 10.0; the topology above, with the expert split set explicitly; -no speculative decoding; bf16 weights and a bf16 KV pool; tokens_per_block 64 (the MLA generation kernels K3's 96 -heads reach exist only at 64); the V2 hybrid KV / state manager, which holds the KDA states, with block reuse off and -fp32 recurrent states; an all-reduce strategy of AUTO or MNNVL. The construction-time ones fail in `__init__`, the +no speculative decoding; the MXFP4 checkpoint's quantization (W4A16_MXFP4 with no per-layer declarations, so the +routed experts run the W4A8_MXFP4_MXFP8 default and the excluded modules stay bf16); bf16 weights and a bf16 KV pool; +tokens_per_block 64 (the MLA generation kernels K3's 96 heads reach exist only at 64); the V2 hybrid KV / state +manager, which holds the KDA states, with block reuse off and fp32 recurrent states; an all-reduce strategy of AUTO or +MNNVL. The construction-time ones fail in `__init__`, the per-engine ones on the first forward, each naming the setting. A layer the decode kernels do not take fails the weight load. @@ -717,19 +719,12 @@ def _apply_attn_res_add_and_rmsnorm( ) -# Routed-expert key spellings that ModelOpt emits for Kimi K3. The NVFP4 -# checkpoint (``nvidia/Kimi-K3-NVFP4``) lists every prefix x module-name -# combination in ``quantized_layers``, so a lookup over this product finds it -# without needing the MiniMax-M3-style prefix normalization in ``ModelConfig``. -_K3_ROUTED_EXPERT_KEY_PREFIXES = ("language_model.model.", "model.", "") - - _K3_ROUTED_EXPERT_KEY_SUFFIXES = ("block_sparse_moe.experts", "mlp.experts") -# The subset of the above that can be a real module path. ``exclude_modules`` -# matches with wildcards and walks ancestor prefixes, so an empty prefix would -# widen what matches instead of just missing, as it does in the dict lookup. +# The routed experts' module path prefixes. ``exclude_modules`` matches with +# wildcards and walks ancestor prefixes, so an empty prefix would widen what +# matches instead of just missing. _K3_ROUTED_EXPERT_MODULE_PREFIXES = ("language_model.model.", "model.") @@ -754,19 +749,6 @@ def _load_packed_mxfp4_expert(backend, base, expert_idx, local_slot_id, get_tens ) -def _load_nvfp4_expert(backend, base, expert_idx, local_slot_id, get_tensor) -> None: - backend.quant_method.load_streaming_nvfp4_expert( - backend, - global_expert_id=expert_idx, - local_slot_id=local_slot_id, - **{ - f"{w}_{kind}": get_tensor(f"{base}.{expert_idx}.{w}.{kind}") - for w in ("w1", "w2", "w3") - for kind in ("weight", "weight_scale", "weight_scale_2", "input_scale") - }, - ) - - class _K3ExpertCkptSpec(NamedTuple): """How one routed-expert quantization is spelled and loaded.""" @@ -775,8 +757,7 @@ class _K3ExpertCkptSpec(NamedTuple): loader: Callable[..., None] # Set of filled slots the loader maintains, checked after the load. loaded_slots_attr: str - # NVFP4 defers cat/pad/interleave and the alpha computation to - # ``process_weights_after_loading``; the MXFP4 loaders write through. + # Whether the layer is finalized after its experts load (the MXFP4 loaders write through). needs_layer_finalize: bool @@ -787,12 +768,6 @@ class _K3ExpertCkptSpec(NamedTuple): loaded_slots_attr="_packed_mxfp4_loaded_slots", needs_layer_finalize=False, ), - QuantAlgo.NVFP4: _K3ExpertCkptSpec( - kinds=("weight", "weight_scale", "weight_scale_2", "input_scale"), - loader=_load_nvfp4_expert, - loaded_slots_attr="_streamed_expert_slots", - needs_layer_finalize=True, - ), } @@ -843,20 +818,12 @@ def __init__( situ_beta, situ_linear_beta = _resolve_kimi_situ_betas(cfg) dtype = torch.bfloat16 - # Routing scores stay fp32; with attention-DP off the gate GEMM runs - # bf16xbf16 with fp32 accumulate/output (checkpoint stores the gate - # weight in bf16; saves a per-layer input cast + fp32 splitK pair on - # the bs1 decode path). Under attention-DP the legacy upcast-to-fp32 - # GEMM is kept: the bf16-input min-latency GEMM's different reduction - # order flips borderline top-16 picks (GSM8K 96.7 -> 96.1/96.4, - # 3-run bisect on 62b20dd868), and the bs1-latency win is irrelevant - # at DEP batch sizes. KIMI_K3_ROUTER_BF16=1/0 forces either path. + # Routing scores stay fp32; the gate GEMM runs bf16xbf16 with fp32 + # accumulate/output (checkpoint stores the gate weight in bf16; saves a + # per-layer input cast + fp32 splitK pair on the bs1 decode path). + # KIMI_K3_ROUTER_BF16=0 forces the upcast-to-fp32 GEMM. _router_bf16_env = os.environ.get("KIMI_K3_ROUTER_BF16") - _router_bf16 = ( - _router_bf16_env == "1" - if _router_bf16_env is not None - else not model_config.mapping.enable_attention_dp - ) + _router_bf16 = _router_bf16_env == "1" if _router_bf16_env is not None else True self.gate = KimiK3MoEGate(cfg, logits_gemm_dtype=torch.bfloat16 if _router_bf16 else None) routed_moe_model_config = self._routed_moe_model_config(model_config) @@ -926,14 +893,11 @@ def __init__( self.expert_hi = self.expert_lo + self.experts_per_rank shared_intermediate = cfg.moe_intermediate_size * cfg.num_shared_experts - attention_dp = model_config.mapping.enable_attention_dp shared_model_config = copy.copy(model_config) shared_model_config.quant_config = QuantConfig() - # Under attention DP each rank owns different tokens, so the shared - # expert is replicated (TP size 1) and must not reduce across ranks. # Direct MoE-TP leaves both branches as partials for one concatenated # all-reduce. - use_shared_tp = not attention_dp and model_config.mapping.tp_size > 1 + use_shared_tp = model_config.mapping.tp_size > 1 self._reduce_routed_output = ( use_shared_tp and self.routed_experts.backend.scheduler_kind != MoESchedulerKind.FUSED_COMM @@ -954,7 +918,6 @@ def __init__( ), dtype=dtype, config=shared_model_config, - overridden_tp_size=1 if attention_dp else None, reduce_output=use_shared_tp, layer_idx=layer_idx, is_shared_expert=True, @@ -988,38 +951,23 @@ def _routed_projection(hidden_states: torch.Tensor, projection: nn.Module) -> to @staticmethod def _select_moe_tp_ep(mapping: Mapping) -> Tuple[int, int]: - """Resolve the routed-expert ``(moe_tp, moe_ep)`` split. - - Precedence: - - 1. Explicit ``moe_tensor_parallel_size`` / ``moe_expert_parallel_size`` - from the user config. Detected via - ``mapping.moe_tp_ep_user_specified`` so the auto-resolved mapping - default (``moe_tp=tp_size, moe_ep=1``) is NOT mistaken for a TP - request. - 2. Default: EP-only (``moe_tp=1, moe_ep=tp_size``), the historical - K3 layout. - """ - tp_size = mapping.tp_size - if getattr(mapping, "moe_tp_ep_user_specified", False): - return mapping.moe_tp_size, mapping.moe_ep_size - return 1, tp_size + # >>> route B: this target's split + """The routed-expert ``(moe_tp, moe_ep)`` split: the user config's explicit + ``moe_tensor_parallel_size`` / ``moe_expert_parallel_size`` (16 x 1, asserted at + construction).""" + # <<< route B + return mapping.moe_tp_size, mapping.moe_ep_size @staticmethod def _resolve_routed_quant_config(model_config: ModelConfig, layer_idx: int) -> QuantConfig: - """Routed-expert quantization for ``layer_idx``, taken from the checkpoint. - - ``nvidia/Kimi-K3-NVFP4`` declares the routed experts per layer as - ``NVFP4`` with ``group_size=16``; the original ``moonshotai/Kimi-K3`` - declares nothing per layer and keeps the historical - ``W4A8_MXFP4_MXFP8`` default. Reading the checkpoint instead of - hardcoding is what lets one code path serve both. - - An exclusion outranks the per-layer entry and the default below: - ``create_weights`` treats an override as authoritative over anything - ``__post_init__`` wrote, so this return value stands in for both - quantization passes and exclusion is the one that runs second. It is - matched as a pattern, so it is asked only about real module names. + """Routed-expert quantization for ``layer_idx``: the MXFP4 checkpoint declares nothing per layer (asserted at + construction), so the experts keep the ``W4A8_MXFP4_MXFP8`` default. + + An exclusion outranks the default: ``create_weights`` treats an override + as authoritative over anything ``__post_init__`` wrote, so this return + value stands in for both quantization passes and exclusion is the one + that runs second. It is matched as a pattern, so it is asked only about + real module names. """ quant_config = model_config.quant_config if quant_config is not None and any( @@ -1036,22 +984,6 @@ def _resolve_routed_quant_config(model_config: ModelConfig, layer_idx: int) -> Q ) return QuantConfig(kv_cache_quant_algo=quant_config.kv_cache_quant_algo) - per_layer = getattr(model_config, "quant_config_dict", None) - if per_layer: - for prefix in _K3_ROUTED_EXPERT_KEY_PREFIXES: - for suffix in _K3_ROUTED_EXPERT_KEY_SUFFIXES: - cfg = per_layer.get(f"{prefix}layers.{layer_idx}.{suffix}") - if cfg is not None and cfg.quant_algo is not None: - # Logged once per layer: the routed-expert format decides - # which MoE backends can serve this checkpoint at all. - logger.debug( - "Kimi K3 layer %d routed experts: %s (group_size=%s) " - "from the checkpoint", - layer_idx, - cfg.quant_algo, - cfg.group_size, - ) - return cfg logger.debug( "Kimi K3 layer %d routed experts: no per-layer quant config in the " "checkpoint, defaulting to %s", @@ -1101,7 +1033,7 @@ def _check_trtllm_situ_quant(moe_backend: str, quant_algo: Optional[QuantAlgo]) @staticmethod def _routed_moe_model_config(model_config: ModelConfig) -> ModelConfig: """Build a private routed-expert mapping without mutating the shared - config. Default split is EP-only; see ``_select_moe_tp_ep``.""" + config, with the split of ``_select_moe_tp_ep``.""" # Every backend here declares ``ActivationType.SiTu`` in its # ``activation_support``; the list is not a preference order. CUTEDSL # joined once its act-fusion kernel grew the SiTU epilogue. @@ -1124,21 +1056,8 @@ def _routed_moe_model_config(model_config: ModelConfig) -> ModelConfig: "EPLB or replicated expert slots." ) mapping = model_config.mapping - if getattr(mapping, "_dwdp_size", 0) > 1: - raise NotImplementedError("Kimi K3 packed-checkpoint streaming does not support DWDP.") moe_tp, moe_ep = KimiK3MoERuntime._select_moe_tp_ep(mapping) - if moe_tp < 1 or moe_ep < 1 or moe_tp * moe_ep != mapping.tp_size: - raise ValueError( - f"Kimi K3 routed MoE split moe_tp={moe_tp} x moe_ep={moe_ep} " - f"must multiply to tp_size={mapping.tp_size}." - ) - if moe_tp > 1 and mapping.enable_attention_dp: - raise NotImplementedError( - "Kimi K3 MoE tensor parallelism requires " - "enable_attention_dp=false (the attention-DP dispatch/combine " - "path is validated for EP-only splits)." - ) logger.info_once( f"Kimi K3 routed MoE parallelism: moe_tp={moe_tp}, " f"moe_ep={moe_ep} (tp_size={mapping.tp_size})", @@ -1279,7 +1198,6 @@ def __init__( layer_idx: int, model_config: ModelConfig, aux_stream_dict: Dict[AuxStreamType, torch.cuda.Stream], - mapping_with_cp: Optional[Mapping] = None, ) -> None: super().__init__() @@ -1293,11 +1211,8 @@ def __init__( # KimiK3MLAAttention owns MLA projection/head sharding. Keep only the # final output reduction in this wrapper so the output gate remains # between attention and the row-parallel o_proj. - # Helix: mapping_with_cp (the CP original) activates the base MLA's - # helix machinery; this wrapper's allreduce over the repurposed - # mapping sums the base o_proj's tp*cp partials. mapping = model_config.mapping - reduce_output = not mapping.enable_attention_dp and mapping.tp_size > 1 + reduce_output = mapping.tp_size > 1 self._o_allreduce = ( AllReduce( mapping=mapping, @@ -1341,7 +1256,6 @@ def __init__( max_position_embeddings=max_positions, model_config=attention_config, aux_stream_dict=aux_stream_dict, - mapping_with_cp=mapping_with_cp, ) def forward( @@ -1400,8 +1314,6 @@ def __init__( layer_idx, model_config=model_config, aux_stream_dict=aux_stream_dict, - # CP original stashed by _setup_helix_mappings; None outside helix. - mapping_with_cp=getattr(model_config, "_helix_mapping_with_cp", None), ) self.is_moe = ( @@ -1414,24 +1326,17 @@ def __init__( else: situ_beta = getattr(cfg, "activation_situ_beta", None) or 1.0 situ_linear_beta = getattr(cfg, "activation_situ_linear_beta", None) - attention_dp = model_config.mapping.enable_attention_dp - if attention_dp: - self.mlp_tp_size = 1 - else: - self.mlp_tp_size = math.gcd(cfg.intermediate_size, model_config.mapping.tp_size) - # Over MNNVL (one NVLink domain across the nodes, where a cross-node all-reduce costs what a node's - # does) the MLP stays split over the whole TP group, the per-rank shapes its decode GEMVs take - # (decode_gemv.SITES); otherwise it stays within one node. - spans_nodes = self._mnnvl_allreduce() is not None - if self.mlp_tp_size > model_config.mapping.gpus_per_node and not spans_nodes: - self.mlp_tp_size = math.gcd( - self.mlp_tp_size, model_config.mapping.gpus_per_node - ) + self.mlp_tp_size = math.gcd(cfg.intermediate_size, model_config.mapping.tp_size) + # Over MNNVL (one NVLink domain across the nodes, where a cross-node all-reduce costs what a node's + # does) the MLP stays split over the whole TP group, the per-rank shapes its decode GEMVs take + # (decode_gemv.SITES); otherwise it stays within one node. + spans_nodes = self._mnnvl_allreduce() is not None + if self.mlp_tp_size > model_config.mapping.gpus_per_node and not spans_nodes: + self.mlp_tp_size = math.gcd(self.mlp_tp_size, model_config.mapping.gpus_per_node) mlp_model_config = copy.copy(model_config) mlp_model_config.quant_config = QuantConfig() # K3's dense layer is BF16, so a unit block size gives the same - # subgroup selection as DeepSeek-V3. Attention DP replicates the - # MLP because ranks own different tokens. + # subgroup selection as DeepSeek-V3. self.mlp = GatedMLP( hidden_size=cfg.hidden_size, intermediate_size=cfg.intermediate_size, @@ -2182,6 +2087,12 @@ def _check_construction(model_config: ModelConfig) -> None: assert kv_algo is None, ( f"this target's MLA kernels read a bf16 KV pool; kv_cache_config.dtype resolved to {kv_algo}" ) + quant_algo = model_config.quant_config.quant_algo + assert quant_algo == QuantAlgo.W4A16_MXFP4 and not model_config.quant_config_dict, ( + "this target loads the MXFP4 checkpoint, whose compressed-tensors config the model config reads as " + "W4A16_MXFP4 with no per-layer declarations (its routed experts run the W4A8_MXFP4_MXFP8 default); the " + f"engine read {quant_algo} with {len(model_config.quant_config_dict or {})} per-layer declarations" + ) strategy = model_config.allreduce_strategy assert strategy in (AllReduceStrategy.AUTO, AllReduceStrategy.MNNVL), ( f"this target runs its all-reduces over MNNVL; allreduce_strategy is {strategy.name}" diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_construction.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_construction.py index 0763b4667356..6c0478735a57 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_construction.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_construction.py @@ -15,6 +15,7 @@ modeling as route_b, ) from tensorrt_llm.functional import AllReduceStrategy +from tensorrt_llm.quantization.mode import QuantAlgo def _config(moe_tp, moe_ep, split_set=True, spec_config=None, attention_dp=False): @@ -32,7 +33,10 @@ def _config(moe_tp, moe_ep, split_set=True, spec_config=None, attention_dp=False mapping=mapping, spec_config=spec_config, torch_dtype=torch.bfloat16, - quant_config=types.SimpleNamespace(quant_algo=None, kv_cache_quant_algo=None), + # The MXFP4 checkpoint's quantization, as the model config reads it. + quant_config=types.SimpleNamespace( + quant_algo=QuantAlgo.W4A16_MXFP4, kv_cache_quant_algo=None + ), quant_config_dict=None, allreduce_strategy=AllReduceStrategy.AUTO, ) From de6326362ee5d9a151dc2380b790d7b6904beead Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 22:08:09 -0700 Subject: [PATCH 085/161] [None][feat] modeling_v2 Kimi K3 target: KDA attention outputs on every path K3DecodeKDA honors reduce_output=False and project_output=False off the decode branch too (unclassified steps and breakable CUDA graphs): the core is the built-in forward's (maybe_bcg_kda_core_inplace inside a breakable graph, else _forward_impl), and the partial is o_proj without the all-reduce. The fused post-attention step runs on steps of at most 16 tokens, classified or not. Signed-off-by: Vasanth Sabavat --- .../modeling.py | 28 ++++++++++++++----- 1 file changed, 21 insertions(+), 7 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index 6003fd0882c7..b5f5d0e88871 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -97,6 +97,7 @@ from tensorrt_llm._torch.modules.gated_mlp import GatedMLP from tensorrt_llm._torch.modules.kimi_k3_mla import KimiK3MLAAttention from tensorrt_llm._torch.modules.kimi_kda import KimiKDALinearAttention +from tensorrt_llm._torch.modules.kimi_kda.kimi_kda_mixer import maybe_bcg_kda_core_inplace from tensorrt_llm._torch.modules.multi_stream_utils import maybe_execute_in_parallel from tensorrt_llm._torch.modules.rms_norm import RMSNorm from tensorrt_llm._torch.modules.situ import SituAndMul @@ -179,6 +180,7 @@ "tensorrt_llm._torch.models.modeling_utils.DecoderModel", # The text model's stock modules. "tensorrt_llm._torch.modules.kimi_kda.KimiKDALinearAttention", + "tensorrt_llm._torch.modules.kimi_kda.kimi_kda_mixer.maybe_bcg_kda_core_inplace", "tensorrt_llm._torch.custom_ops.cute_dsl_kimi_k3_kda_mtp_ops", "tensorrt_llm._torch.modules.kimi_k3_mla.KimiK3MLAAttention", "tensorrt_llm._torch.moe.fused_moe.create_moe", @@ -1751,7 +1753,7 @@ def will_run_decode_branch( self, attn_metadata: AttentionMetadata, step: Optional[DecodeStep] ) -> bool: """Whether ``forward`` runs ``step`` on the decode branch: a step ``decode_step`` classifies, outside a - breakable CUDA graph. Only there does it return the unreduced ``o_proj`` output or the core.""" + breakable CUDA graph.""" return step is not None and not is_in_breakable_cuda_graph() def forward( @@ -1766,14 +1768,13 @@ def forward( on ``ssm/k3_kda_decode_attn`` on a decode step of one token per request, else the built-in dispatch; then ``o_proj`` on its decode GEMV site and the TP all-reduce. - On the decode branch only: ``reduce_output=False`` returns ``o_proj``'s TP partial (no all-reduce), and + On every path, ``reduce_output=False`` returns ``o_proj``'s TP partial (no all-reduce), and ``project_output=False`` the post-o_norm core ``[N, H * 128]`` (no ``o_proj``).""" if not self.will_run_decode_branch(attn_metadata, step): - if not (reduce_output and project_output): - raise ValueError( - "reduce_output / project_output need a step will_run_decode_branch takes" - ) - return super().forward(hidden_states, attn_metadata) + if reduce_output and project_output: + return super().forward(hidden_states, attn_metadata) + core = self._builtin_core(hidden_states, attn_metadata).reshape(-1, self.proj_size) + return self.o_proj(core) if project_output else core if ( step.decode and step.tokens_per_request == 1 @@ -1787,6 +1788,19 @@ def forward( return core.reshape(-1, self.proj_size) return self._k3_project_output(core, reduce_output) + def _builtin_core( + self, hidden_states: torch.Tensor, attn_metadata: AttentionMetadata + ) -> torch.Tensor: + """The built-in forward's core ``[N, H, 128]``: inside a breakable CUDA graph from the eager + ``maybe_bcg_kda_core_inplace``, else from ``_forward_impl``.""" + if self.register_to_config and is_in_breakable_cuda_graph(): + core = hidden_states.new_empty( + (hidden_states.shape[0], self.num_heads, self.head_dim), dtype=torch.bfloat16 + ) + maybe_bcg_kda_core_inplace(hidden_states, self.layer_idx_str, core) + return core + return self._forward_impl(hidden_states, attn_metadata) + def _k3_project_output(self, core: torch.Tensor, reduce_output: bool = True) -> torch.Tensor: """``o_proj`` on the ``o_proj`` decode GEMV site where it takes the rows (else the module), then, with ``reduce_output``, the TP all-reduce.""" From cce3f4641d3beeccb4fa753b2d1787670190664d Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 22:09:23 -0700 Subject: [PATCH 086/161] [None][test] Kimi K3 latent exchange: the one-process test checks the TP-size refusal K3LatentExchange.create now refuses a TP size other than 4, 8 or 16 before any collective step, so a group of one raises ValueError, eagerly and under CUDA-graph capture alike; the test asserts that. The capture refusal at a supported size is a collective agreement, which the 4-rank matrix checks. Signed-off-by: Vasanth Sabavat --- .../kimi_k3/test_k3_latent_reduce.py | 12 ++++++++---- 1 file changed, 8 insertions(+), 4 deletions(-) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_latent_reduce.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_latent_reduce.py index f5cf86348935..6a665ae2554c 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_latent_reduce.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_latent_reduce.py @@ -232,17 +232,21 @@ def test_k3_latent_reduce(mpi_pool_executor): assert all(row["ok"] for rows in per_rank for row in rows) -def test_latent_exchange_refuses_graph_capture(): - """The exchange is created collectively (an MNNVL multicast allocation over the TP group): creating it under - CUDA-graph capture raises instead of entering the collective. One process, a group of one.""" +def test_latent_exchange_refuses_unsupported_tp_size(): + """The exchange serves TP groups of 4, 8 or 16 ranks: another size raises ValueError before any collective step + (the size is the same on every rank of a group), eagerly and under CUDA-graph capture alike. One process, a group + of one. Its refusal of capture at a supported size is a collective agreement: the 4-rank matrix + (tests/unittest/_torch/modeling_v2/comm/_k3_latent_reduce_op_matrix.py) checks it.""" import tensorrt_llm # noqa: F401 from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import latent_op from tensorrt_llm.mapping import Mapping mapping = Mapping(world_size=1, rank=0, tp_size=1) + with pytest.raises(ValueError, match="supports 4, 8 and 16"): + latent_op.K3LatentExchange.create(mapping) graph, stream = torch.cuda.CUDAGraph(), torch.cuda.Stream() with torch.cuda.graph(graph, stream=stream): - with pytest.raises(RuntimeError, match="outside CUDA-graph capture"): + with pytest.raises(ValueError, match="supports 4, 8 and 16"): latent_op.K3LatentExchange.create(mapping) From 0dcbdafb6eeeb62fd119662dcb8a47045a04262f Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 22:13:30 -0700 Subject: [PATCH 087/161] [None][test] Kimi K3 wide MoE test: a seeded skewed routing instead of a router-dump replay The rtb2 cases replayed a recorded router dump named by K3_ROUTING_DUMP and skipped without it, so CI never ran them (256 cases). The skewed cases draw the same kind of routing in the test, seeded: per draw a skewed expert popularity (log-normal), each token's 16 distinct experts a weighted draw without replacement, and the draw's busiest EP group mapped onto this rank, at every M of 1..64. Nothing in the file reads the environment any more. Signed-off-by: Vasanth Sabavat --- .../kimi_k3/test_k3_moe_wide.py | 52 ++++++++----------- 1 file changed, 23 insertions(+), 29 deletions(-) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_wide.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_wide.py index af6c48a8b09e..bab66c81b8c2 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_wide.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_wide.py @@ -21,8 +21,8 @@ - hot: one local expert in every token's top-16 (8 groups of it at M = 64); - group_cap: 100 experts with 9 of the 64 tokens and 124 with one (324 groups at M = 64: the kernel's group capacity); - none_local: no local expert; -- rtb2 (with K3_ROUTING_DUMP naming the campaign's router dump): 4 draws of 8 decode forwards of 8 tokens each, the - step's busiest EP group mapped onto this rank (M tokens: the first M of the draw). +- skewed: 4 seeded draws of 64 tokens whose experts follow a skewed popularity, as a trained router's do (popular + experts recur across tokens), each draw's busiest EP group mapped onto this rank (M tokens: the first M of the draw). Checks: against the stock path (trtllm::kimi_k3_noaux_tc_mxfp8_quant, then the TRTLLM-Gen W4A8_MXFP4_MXFP8 MoE runner with those ids) and an fp64 reference over the dequantized MXFP4 experts (op-catalog gates: 8 ulp of the row max per element, 4 ulp relative RMS); run-to-run identical bits; the slab armed and the layer's counters zero after every @@ -32,7 +32,6 @@ import functools import math -import os from types import SimpleNamespace import pytest @@ -59,7 +58,6 @@ def _is_sm100() -> bool: E4M3_MAX = 448.0 M_MAX = 64 M_ALL = list(range(1, M_MAX + 1)) -ROUTING_DUMP = os.environ.get("K3_ROUTING_DUMP") def _ops(): @@ -253,7 +251,7 @@ def _scratch_rearmed(state, *layers): CASES = ["random", "16_local", "disjoint", "hot", "group_cap", "none_local"] -RTB2_DRAWS = 4 +SKEWED_DRAWS = 4 def _chosen_logits(chosen, gen): @@ -267,33 +265,30 @@ def _chosen_logits(chosen, gen): @functools.lru_cache(maxsize=None) -def _rtb2_steps(path, draws, seed=11): - """R decode forwards of 8 tokens each from the router dump (one random layer per draw), R = 8 (64 tokens): each - step's ids rotated so that its busiest EP group (most distinct experts) is this rank's.""" - import numpy as np - - z = np.load(path) - ntok = z["fwd_ntok"] - starts = np.concatenate([[0], np.cumsum(ntok)[:-1]]) - decode = [i for i, n in enumerate(ntok) if n == 8] - layers = [int(v) for v in z["layers"]] - rng = np.random.default_rng(seed) +def _skewed_steps(draws, seed=11): + """Per draw, 64 tokens' top-16 experts under a skewed popularity: each expert's log-popularity ~ N(0, 1.5) for the + draw, each token's 16 distinct experts the top-16 of log-popularity plus Gumbel noise (a weighted draw without + replacement), so popular experts recur across tokens. Each draw's ids are rotated so that its busiest EP group + (most distinct experts) is this rank's.""" + gen = torch.Generator().manual_seed(seed) steps = [] for _ in range(draws): - fwds = rng.choice(decode, size=M_MAX // 8, replace=False) - toks = np.concatenate([np.arange(starts[f], starts[f] + 8) for f in fwds]) - ids = z[f"ids_{rng.choice(layers)}"][toks].astype(np.int64) - counts = np.bincount(ids.reshape(-1), minlength=NUM_EXPERTS).reshape(4, E_LOCAL) - busiest = int(np.argmax((counts > 0).sum(axis=1))) + log_pop = torch.randn(NUM_EXPERTS, generator=gen) * 1.5 + uniform = torch.rand(M_MAX, NUM_EXPERTS, generator=gen).clamp_min(1e-20) + ids = (log_pop - torch.log(-torch.log(uniform))).topk(TOP_K, dim=1).indices + counts = torch.bincount(ids.reshape(-1), minlength=NUM_EXPERTS).reshape( + NUM_EXPERTS // E_LOCAL, E_LOCAL + ) + busiest = int((counts > 0).sum(dim=1).argmax()) ids = (ids + (EP_RANK - busiest) * E_LOCAL) % NUM_EXPERTS - steps.append([set(int(e) for e in row) for row in ids]) + steps.append([sorted(int(e) for e in row) for row in ids.tolist()]) return steps @functools.lru_cache(maxsize=None) def _tokens(case: str, seed: int = 7): """64 tokens: router logits and the latent rows (M tokens use the first M).""" - salt = CASES.index(case) if case in CASES else len(CASES) + int(case[5:]) + salt = CASES.index(case) if case in CASES else len(CASES) + int(case.split(".")[1]) gen = torch.Generator(device="cuda").manual_seed(seed + salt) cpu = torch.Generator().manual_seed(seed + salt) x = torch.randn(M_MAX, H, generator=gen, device="cuda").bfloat16() @@ -330,8 +325,8 @@ def pick(pool, k): chosen[t].append(e) free[t] -= 1 assert all(f == 0 for f in free) - elif case.startswith("rtb2."): - chosen = [sorted(e) for e in _rtb2_steps(ROUTING_DUMP, RTB2_DRAWS)[int(case[5:])]] + elif case.startswith("skewed."): + chosen = _skewed_steps(SKEWED_DRAWS)[int(case.split(".")[1])] else: raise ValueError(case) return _chosen_logits(chosen, gen), x @@ -396,11 +391,10 @@ def test_k3_moe_wide(case, m): _check(case, m) -@pytest.mark.skipif(not ROUTING_DUMP, reason="K3_ROUTING_DUMP names no router dump") @pytest.mark.parametrize("m", M_ALL) -@pytest.mark.parametrize("draw", range(RTB2_DRAWS)) -def test_k3_moe_wide_rtb2(draw, m): - _check(f"rtb2.{draw}", m) +@pytest.mark.parametrize("draw", range(SKEWED_DRAWS)) +def test_k3_moe_wide_skewed(draw, m): + _check(f"skewed.{draw}", m) def test_k3_moe_wide_mixed_sequence(): From ba296fc668976b193de96aa160c0e37b238a1bde Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 22:36:02 -0700 Subject: [PATCH 088/161] [None][test] Kimi K3 modeling_v2 target: run the decode-step test in the CPU stage test_modeling_v2_kimi_k3_decode_step.py is cpu_only. l0_b300 collects unittest/_torch/modeling_v2 with -m "not cpu_only", and no CPU list named it, so no CI stage ran it. l0_cpu now lists the file. Signed-off-by: Vasanth Sabavat --- tests/integration/test_lists/test-db/l0_cpu.yml | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/integration/test_lists/test-db/l0_cpu.yml b/tests/integration/test_lists/test-db/l0_cpu.yml index bfb2c39eba38..2f1a4de55524 100644 --- a/tests/integration/test_lists/test-db/l0_cpu.yml +++ b/tests/integration/test_lists/test-db/l0_cpu.yml @@ -48,6 +48,8 @@ l0_cpu: - unittest/_torch/peft - unittest/_torch/memory - unittest/_torch/modeling + # l0_b300 selects unittest/_torch/modeling_v2 with -m "not cpu_only", so this entry is what runs it. + - unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_step.py - unittest/_torch/models/test_minimax_m3.py::test_minimax_m3_fp8_indexer_rejects_different_qk_norm_epsilons - unittest/_torch/models/test_minimax_m3.py::test_minimax_m3_moe_reduces_only_local_terms - unittest/_torch/models/test_minimax_m3.py::test_minimax_m3_decoder_layer_sets_post_fusion_from_moe_scheduler From df6da9d9b0171ef0b2b126292f9124c0407bf7b4 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 22:36:03 -0700 Subject: [PATCH 089/161] [None][doc] modeling_v2 catalog: 18 entries carry sm_100 receipts The receipt status in index.yaml and the package README still said 12, the count before the Kimi K3 KDA and MLA decode entries were added. 18 contracts have an sm_100 receipt: the 16 Kimi K3 entries (k3_*, attn_res_* and ssm/kda_decode) plus cublas_mm and flashinfer_rmsnorm. Signed-off-by: Vasanth Sabavat --- .../_torch/_experimental/modeling_v2/README.md | 4 ++-- .../_experimental/modeling_v2/catalog/index.yaml | 10 +++++----- 2 files changed, 7 insertions(+), 7 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/README.md b/tensorrt_llm/_torch/_experimental/modeling_v2/README.md index 5c7448160818..87069886e5be 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/README.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/README.md @@ -181,8 +181,8 @@ Perf is measured, never gated. ## Status of every record in this tree -**The catalog is certified: 19 entries on sm_103, and 12 on sm_100 (B200 / -GB200): the 10 Kimi K3 entries, where their first caller runs, and +**The catalog is certified: 19 entries on sm_103, and 18 on sm_100 (B200 / +GB200): the 16 Kimi K3 entries, where their first caller runs, and `cublas_mm` and `flashinfer_rmsnorm`. The targets construct but have never executed.** diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml index f4e68e6e2a5d..f9103979de78 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml @@ -28,7 +28,7 @@ # limits, e.g. no float8 arithmetic, are noted in the entry docstring). # Entries wrapping trtllm ops carry all three. # -# ── RECEIPT STATUS: 19 entries certified on sm_103, 12 on sm_100 ────────── +# ── RECEIPT STATUS: 19 entries certified on sm_103, 18 on sm_100 ────────── # # A receipt says this entry's test passed on a stated GPU architecture. The # architecture is the whole key: it is a real axis -- see the @@ -67,10 +67,10 @@ # device, and rewriting them would manufacture GB300 evidence that does not # exist. Read them as provenance; the frontmatter is the certification. # -# The Kimi K3 entries (k3_*, attn_res_*) are certified on sm_100 (B200 / -# GB200), where their first caller runs, and their tests skip on any other -# architecture. cublas_mm and flashinfer_rmsnorm, which that caller also -# uses, keep their sm_103 receipts and add sm_100 ones. +# The Kimi K3 entries (k3_*, attn_res_* and ssm/kda_decode) are certified on +# sm_100 (B200 / GB200), where their first caller runs, and their tests skip on +# any other architecture. cublas_mm and flashinfer_rmsnorm, which that caller +# also uses, keep their sm_103 receipts and add sm_100 ones. # ────────────────────────────────────────────────────────────────────────── # # 30 of the source catalog's 43 entries are here — the union of what the two From db8ca58bd0dd00c98019181741555246018d0ae8 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 22:42:39 -0700 Subject: [PATCH 090/161] [None][test] modeling_v2 Kimi K3 collective and MoE entries: sm_100 receipts The nine entries get their sm_100 receipts: the seven collective ones at 4 ranks on one GB200 tray (each matrix's collected test file in its CI form, pytest starting mpirun, and the same rank bodies started by srun), moe/k3_moe and moe/k3_route_quant with their single-GPU test files' counts. comm/allgather, which the Kimi K3 LM head calls, gets an sm_100 key; its sm_103 receipt stays, taken on the same, unchanged test file. The catalog's counts follow, and the MNNVL all-reduce contract states the evidence for its one-shot kernel's change (main's test bodies, bit-identical outputs, per- call time against the previous kernel). Signed-off-by: Vasanth Sabavat --- .../_torch/_experimental/modeling_v2/README.md | 8 ++++---- .../modeling_v2/catalog/comm/allgather.md | 2 ++ .../modeling_v2/catalog/comm/k3_latent_reduce.md | 2 +- .../modeling_v2/catalog/comm/k3_sandwich_oproj.md | 2 +- .../modeling_v2/catalog/comm/k3_sandwich_plain.md | 2 +- .../modeling_v2/catalog/comm/k3_sandwich_tail.md | 2 +- .../modeling_v2/catalog/comm/mnnvl_allgather_split.md | 2 +- .../modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md | 10 +++++++--- .../_experimental/modeling_v2/catalog/index.yaml | 6 +++--- .../_experimental/modeling_v2/catalog/moe/k3_moe.md | 2 +- .../modeling_v2/catalog/moe/k3_moe_front.md | 2 +- .../modeling_v2/catalog/moe/k3_route_quant.md | 2 +- 12 files changed, 24 insertions(+), 18 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/README.md b/tensorrt_llm/_torch/_experimental/modeling_v2/README.md index dfa8d82d1352..7b0ffa818edf 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/README.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/README.md @@ -186,10 +186,10 @@ Perf is measured, never gated. ## Status of every record in this tree -**The catalog is certified: 19 entries on sm_103, and 13 on sm_100 (B200 / -GB200): the 11 Kimi K3 entries (`comm/mnnvl_allreduce_attn_res` among them), -where their first caller runs, and `cublas_mm` and `flashinfer_rmsnorm`. The -targets construct but have never executed.** +**The catalog is certified: 19 entries on sm_103, and 23 on sm_100 (B200 / +GB200): the 20 Kimi K3 entries (the MNNVL ones among them), where their +first caller runs, and `cublas_mm`, `flashinfer_rmsnorm` and `allgather`. +The targets construct but have never executed.** Two things voided every receipt in the move: each catalog test file was rewritten, and the targets moved from sm_100 (B200) to sm_103 (GB300), where diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/allgather.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/allgather.md index 82072d3843a6..45eb99368ae3 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/allgather.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/allgather.md @@ -1,6 +1,8 @@ --- receipts: + # sm_103 was certified before the sm_100 key was added, on the same test file. sm_103: {status: passed, world_size: 4} + sm_100: {status: passed, world_size: 4} --- # allgather diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md index cce0226627f0..7f171431f01c 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md @@ -1,6 +1,6 @@ --- receipts: - sm_100: {status: pending, world_size: 4} + sm_100: {status: passed, world_size: 4} --- # k3_latent_reduce diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md index b858953e3a9a..316378ce619c 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md @@ -1,6 +1,6 @@ --- receipts: - sm_100: {status: pending, world_size: 4} + sm_100: {status: passed, world_size: 4} --- # k3_sandwich_oproj diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md index 29f0b09cc826..a21cc254f931 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md @@ -1,6 +1,6 @@ --- receipts: - sm_100: {status: pending, world_size: 4} + sm_100: {status: passed, world_size: 4} --- # k3_sandwich_plain diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md index 068c1532da24..3d3a1d2e7df0 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md @@ -1,6 +1,6 @@ --- receipts: - sm_100: {status: pending, world_size: 4} + sm_100: {status: passed, world_size: 4} --- # k3_sandwich_tail diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md index b09a7c2ceb78..04b225a7d4c6 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md @@ -1,6 +1,6 @@ --- receipts: - sm_100: {status: pending, world_size: 4} + sm_100: {status: passed, world_size: 4} --- # mnnvl_allgather_split diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md index a50d3fcb2fd8..7fce1c4c2e6e 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md @@ -1,6 +1,6 @@ --- receipts: - sm_100: {status: pending, world_size: 4} + sm_100: {status: passed, world_size: 4} --- # mnnvl_fusion_allreduce @@ -233,8 +233,12 @@ follows `T`, `H` and the device's SM count, and in the fused form it sets the or so a dependent kernel (a GEMV) can launch and stream its weights while this kernel waits for its peers; a dependent still reads the output and the flags only after its own grid wait. Its Lamport reduction is the shared `reduceOneshotLamport` of the attention-residual one-shot kernel: the code is moved, not changed, and the reduction - order is the same. Results unchanged (main's MNNVL test bodies pass on this build; the kernel's SASS changes). - EVIDENCE: + order is the same. Certified on sm_100: main's MNNVL test bodies (`multi_gpu/test_mnnvl_allreduce.py`, 110 cases + at 4 ranks and 109 at 2, and its graph-capture cases) pass on this build. At `[M, 7168]` bf16, `M` 1 to 2048, + plain and fused, one-shot and two-shot, 4 ranks, this build's outputs equal the previous kernels' bit for bit, and + the per-call time of back-to-back captured calls is within -0.56 / +0.12 us of theirs (noise 0.10 us). Every + MNNVL kernel's SASS changes, since the kernel parameters gained `earlyTrigger`; in the one-shot kernel's, the + dependents' launch sits right after the grid-dependency wait. - At `W` = 16 the one-shot kernel adds the ranks in two chunks of 8, a branch a 4-rank run never reaches. The 16-rank receipt is pending. - Kimi K3's calls (its decode path, not this test): the model sets every `MNNVLAllReduce` of the target, its diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml index 6ee1d5eac663..2111035e82ee 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml @@ -28,7 +28,7 @@ # limits, e.g. no float8 arithmetic, are noted in the entry docstring). # Entries wrapping trtllm ops carry all three. # -# ── RECEIPT STATUS: 19 entries certified on sm_103, 13 on sm_100 ────────── +# ── RECEIPT STATUS: 19 entries certified on sm_103, 23 on sm_100 ────────── # # A receipt says this entry's test passed on a stated GPU architecture. The # architecture is the whole key: it is a real axis -- see the @@ -75,8 +75,8 @@ # # The Kimi K3 entries (k3_*, attn_res_*) are certified on sm_100 (B200 / # GB200), where their first caller runs, and their tests skip on any other -# architecture. cublas_mm and flashinfer_rmsnorm, which that caller also -# uses, keep their sm_103 receipts and add sm_100 ones. +# architecture. cublas_mm, flashinfer_rmsnorm and allgather, which that +# caller also uses, keep their sm_103 receipts and add sm_100 ones. # ────────────────────────────────────────────────────────────────────────── # # STATEFUL ENTRIES. When an op's result depends on state that outlives the diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md index e4f139541269..ed3ac4838778 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md @@ -1,6 +1,6 @@ --- receipts: - sm_100: {status: pending, world_size: 4} + sm_100: {status: passed, tests: 62} --- # k3_moe diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md index fdb5ad7ef751..d1c317863ff1 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md @@ -1,6 +1,6 @@ --- receipts: - sm_100: {status: pending, world_size: 4} + sm_100: {status: passed, world_size: 4} --- # k3_moe_front diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_route_quant.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_route_quant.md index 3d3f5f1851d8..cef91e814286 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_route_quant.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_route_quant.md @@ -1,6 +1,6 @@ --- receipts: - sm_100: {status: pending, tests: 18} + sm_100: {status: passed, tests: 18} --- # k3_route_quant From 37a8c925c28ece293d537aca5afb54c2a2fee7c8 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 21:31:50 -0700 Subject: [PATCH 091/161] [None][feat] Kimi K3 MoE: trtllm::k3_moe_m1 / k3_moe_m2, the routed experts of one or two decode tokens Two weight-stream CuTe DSL kernels for decode steps of one or two tokens when every expert is on the rank (moe TP16 x EP1: all 896 experts, a 192-wide slice of each intermediate zero-padded to 256): FC1 (MXFP4 x MXFP8), SiTU, the MXFP8 intermediate, FC2 and the routing-weighted combine in trtllm::k3_moe's order, so the result has k3_moe's bits except where FC1's two partial sums round an intermediate value differently. - trtllm::k3_moe_m1 (one or two tokens) and trtllm::k3_moe_m2 (two tokens, an intermediate of at most 256, per-group FC1 -> FC2 hand-offs) are torch ops over a caller-owned workspace; mutates_args names the intermediate rows, the counts, the epochs, the exchange's multicast words and out. - K3MoeM1State / K3MoeM2State own the workspace: every call re-arms the count set the next call uses and advances the CTAs' epochs. create() compiles the plain build and the push builds it is given, before capture. K3MoeM1Layer / K3MoeM2Layer hold a layer's weights. - The push build stores the partial into every rank's K3LatentExchange for trtllm::k3_latent_reduce instead of returning it. Tests: both kernels against trtllm::k3_moe on the same buffers and an fp32 reference over the dequantized experts (TP16, and TP4 x EP4 for m1 at one token), run to run, layers interleaved on one state, the epochs across the int32 wrap, and the two-token builds' refusals. Signed-off-by: Vasanth Sabavat --- .../k3_fused_moe/k3_moe_m1_kernel.py | 1193 +++++++++++++++++ .../k3_fused_moe/k3_moe_m2_kernel.py | 1129 ++++++++++++++++ .../cute_dsl_kernels/k3_fused_moe/op.py | 575 ++++++++ .../test_lists/test-db/l0_b200.yml | 2 + .../kimi_k3/test_k3_moe_m1.py | 306 +++++ .../kimi_k3/test_k3_moe_m2.py | 329 +++++ 6 files changed, 3534 insertions(+) create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_moe_m1_kernel.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_moe_m2_kernel.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_m1.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_m2.py diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_moe_m1_kernel.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_moe_m1_kernel.py new file mode 100644 index 000000000000..1803fc1fa625 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_moe_m1_kernel.py @@ -0,0 +1,1193 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Kimi K3 routed experts of one or two decode tokens (M = m_max <= 2) as a weight-stream kernel: ``k3_moe_m1``. + +At M <= 2 every expert a token routes to is a GEMV, so the kernel spreads the experts' weight rows over all CTAs +instead of k3_moe's 128-row tiles. It reads the outputs of trtllm::k3_moe_front (or trtllm::k3_route_quant): the +tokens' top-16 global ids and bf16 weights and their MXFP8 latents. It computes this rank's routed partial [M, 3584] +bf16, the tensor k3_moe returns. Weights are read in place in the TRTLLM-Gen W4A8_MXFP4_MXFP8 layout (see +k3_moe_kernel.py), including the loader's zero padding of a rank's intermediate to a multiple of 128 (i_pad, e.g. +192 -> 256 at TP16). The kernel streams the i_tp real values only. + +- Routing: one warp decodes the 16 M top-k slots into experts (the distinct local experts in lane order, each with + its tokens). The first FC1 unit's ring fill is issued from the lanes before the CTA barrier. +- FC1 (gate_up, MXFP4 x MXFP8) in 64-row units spread over the CTAs, with k3_moe's per-stage block-scaled + machinery. Two k-tiles go into each stage: the unit's 64 rows at k-tile 2j fill MMA rows 0-63 and at k-tile + 2j + 1 rows 64-127; B rows t and 8 + t hold token t at those k-tiles (N 16), and the scale atoms are spliced to + match. The epilogue adds the two partial sums per token, then applies k3_moe's SiTU and MXFP8 requantization, into + one intermediate row per (expert, token). +- Hand-off: each unit adds one release to the call's count. The count has two slots by epoch parity; CTA 0 re-arms + the other slot, and every CTA advances its epoch. The FC2 CTAs acquire the count, then load the intermediate in + one round. +- FC2 (down): each of 112 CTAs owns one 32-row block of the shuffled down projection. Four experts' 32-row slices + stack in one 128-row MMA (rows 32 j .. 32 j + 31 = expert 4 i + j) against B N 16, whose row 4 t + j is token t's + intermediate for expert 4 i + j; this is k3_moe's FC2 operand, so each expert's down projection has k3_moe's bits. + - The A tiles load into the FC1 ring as it drains, evict-first (read once per step, like FC1's weights). + - The weight scale atoms are spliced while FC1 runs. +- Combine: per token, the routing-weighted experts are summed in k3_moe's order (ascending local id in min(G, 5) + slices, an expert without the token adding zero, products and sums rounded on their own), so the output matches + k3_moe's bits except where the two FC1 partial sums round an intermediate value differently. + +Push (K3_CONFIG "push" = 1): instead of writing ``out``, the combine stores this +rank's row of token t into slot [half][t][rank x copies + c] of every rank's latent exchange +(``latent_op.K3LatentExchange``, int32 [2][8][push_world][1792], bf16 pairs, -0.0 stored as +0.0) through its multicast +mapping, half = lat_flags[0] & 1 read after the grid wait; trtllm::k3_latent_reduce sums it. ``push_copies`` > 1 +only emulates a larger group's receive side. + +Configuration is per module instance (the shapes are trace-time constants): the loader injects K3_CONFIG = +{"i_tp": ..., "i_pad": ..., "num_local": ..., "num_ctas": ..., "m_max": ..., ["push", "push_world", "push_copies"]} +before executing the module. Launched with programmatic dependent launch: the barrier setup and the TMEM +allocation run before griddepcontrol.wait, which precedes every read of the producer's outputs and every global +write. +""" + +from __future__ import annotations + +import cuda.bindings.driver as cuda_driver +import cutlass +import cutlass.cute as cute +import cutlass.experimental.cuda as cuda +from cutlass._mlir.dialects import llvm as _llvm +from cutlass._mlir.extras import types as _T +from cutlass.cutlass_dsl import dsl_user_op +from cutlass.experimental import primitives as prims +from cutlass.experimental.cuda.tensor_map import TensorMapDataType + +_CFG = globals().get("K3_CONFIG") or {} + + +def _cfg(key: str, default): + """A kernel option from the op's configuration (K3_CONFIG), else its default.""" + return type(default)(_CFG.get(key, default)) + + +# ============================================================================= +# Problem shape (trace-time constants). +# ============================================================================= +H = 3584 # latent hidden = FC1 K = FC2 rows +I_TP = _cfg("i_tp", 192) # a rank's intermediate values (the logical shard) +I_PAD = _cfg( + "i_pad", (I_TP + 127) // 128 * 128 +) # the loader's padded intermediate (the buffers' layout) +E_LOCAL = _cfg("num_local", 896) # this rank's experts +M_MAX = _cfg("m_max", 1) # tokens per call (one routing warp: 16 M lanes) +NCTA = _cfg("num_ctas", 148) # one CTA per SM +PUSH = bool(_cfg("push", 0)) # the combine pushes into the latent exchange instead of writing out +PUSH_WORLD = _cfg("push_world", 16) # slots per (half, token) of the exchange +PUSH_COPIES = _cfg("push_copies", 1) # slots this rank fills (rank x copies + c) +LAT_ROW_WORDS = 3584 // 2 # int32 words of a latent row in the exchange +TWO_I = 2 * I_TP +TOP_K = 16 +G_CAP = TOP_K * M_MAX # distinct local experts per call +THREADS = 256 +N = 16 # MMA N: FC1 B rows 0 (k-tile 2j) and 8 (k-tile 2j + 1) are the token; FC2 B rows 0-3 are a group's experts +MMA_M, MMA_TILE_K, MMA_INST_K = 128, 128, 32 +ROWS1 = 64 # FC1 rows per unit (one half of a 128-row tile) +U1 = TWO_I // ROWS1 # units per expert +K1_TILES = H // MMA_TILE_K # 28 +K1_PAIRS = K1_TILES // 2 # 14 stages per unit +NUM_KBLOCKS = MMA_TILE_K // MMA_INST_K # 4 +a_dtype = cutlass.Float4E2M1FN +b_dtype = cutlass.Float8E4M3FN +sf_dtype = cutlass.Float8E8M0FNU +a_smem_width = 8 # FP4 unpacked to 8-bit containers in shared memory +sf_vec_size = 32 +num_m0_per_sf_atom = 32 +num_m1_per_sf_atom = 4 +num_k_per_sf_atom = 4 +num_elts_atom_sf_fp16 = num_m0_per_sf_atom * num_m1_per_sf_atom * num_k_per_sf_atom // 2 +num_tmem_cols_per_sf_atom = 4 +NUM_BYTES_A = MMA_M * MMA_TILE_K * a_smem_width // 8 # 16384 +NUM_BYTES_A_HALF = ROWS1 * MMA_TILE_K * a_smem_width // 8 # 8192 +NUM_TX_A_HALF = ROWS1 * MMA_TILE_K * 4 // 8 # 4096: FP4 bytes in global memory +NUM_BYTES_B = N * MMA_TILE_K # 2048 +NUM_BYTES_SFA = 512 +NUM_BYTES_SFA_RAW = 2 * NUM_BYTES_SFA # the two k-tiles' scale atoms as loaded +NUM_BYTES_SFB = 512 +SFB_GROUP_BYTES = 16 +SFA_COLS = num_tmem_cols_per_sf_atom +SFB_COLS = num_tmem_cols_per_sf_atom +NUM_SF_IDS = num_k_per_sf_atom * sf_vec_size // MMA_INST_K # 4 +SITU_GATE_CAP = 4.0 +SITU_LINEAR_CAP = 25.0 +E4M3_MAX = 448.0 +FP8_SENTINEL_I8 = -128 +_LOG2E = 1.4426950408889634 +# FC2 +R2 = 32 # output rows per CTA: one 32-row block of the shuffled layout +FC2_CTAS = H // R2 # 112 +ROW_PAD2 = I_PAD // 2 # FP4 bytes of a down row in the buffer +KB2 = I_TP // 32 # MX blocks (K 32 MMAs) of a down row's real values +KT2 = (I_TP + MMA_TILE_K - 1) // MMA_TILE_K # 128-wide k-tiles of the real values +KA2 = ( + I_PAD // 128 +) # w2 scale atoms per 128-row block (block_scale_interleave of the padded I / 32 columns) +GROUPS2 = G_CAP // 4 # 4 experts x 32 rows per 128-row MMA +NT2 = GROUPS2 * KT2 # A tiles, B tiles and scale atoms per CTA +NUM_TX_A2 = R2 * MMA_TILE_K * 4 // 8 # 2048: FP4 bytes of one expert's 32 rows of a k-tile +NUM_BYTES_B2 = N * MMA_TILE_K # 2048 +H_ROW = ( + (I_TP + I_TP // 32 + 15) // 16 * 16 +) # an intermediate row: fp8 values + E8M0 scales, 16-byte multiple +VCH = I_TP // 16 # 16-byte value chunks of an intermediate row +SCH = (KB2 + 15) // 16 # 16-byte scale chunks +STAGE_ROUNDS = (G_CAP * M_MAX * (VCH + SCH) + 127) // 128 # load rounds of the 128 epilogue threads +# The ring: as many stages as fit next to FC2's B tiles and scale atoms (8 at TP16, 7 at TP4 x EP4). +_STAGE_BYTES = NUM_BYTES_A + NUM_BYTES_B + NUM_BYTES_SFA + NUM_BYTES_SFA_RAW + NUM_BYTES_SFB +_FIXED_BYTES = NT2 * (NUM_BYTES_B2 + NUM_BYTES_SFA + NUM_BYTES_SFB) + G_CAP * M_MAX * R2 * 4 + 4096 +STAGES = min(8, (227 * 1024 - _FIXED_BYTES) // _STAGE_BYTES) +PRE_FILL = min(STAGES, K1_PAIRS) # stages of the first unit issued before the routing barrier +# TMEM columns: FC1 accumulator, FC1 SFA / SFB per stage, FC2 accumulators (16 per group), FC2 SFA / SFB per tile. +ACC2_COL = N + 2 * STAGES * SFA_COLS +SFA2_COL = ACC2_COL + GROUPS2 * N +SFB2_COL = SFA2_COL + NT2 * SFA_COLS +TMEM_COLS = 32 +while TMEM_COLS < SFB2_COL + NT2 * SFB_COLS: + TMEM_COLS *= 2 +REF_SLICES = 5 # k3_moe's FC2 slices (fc2_slices): its combine's sum tree, kept so the bits match +_REF_SLICES_FULL = [ + (s * G_CAP // REF_SLICES, (s + 1) * G_CAP // REF_SLICES) for s in range(REF_SLICES) +] +EPI_BAR_ID = 1 +EPI_THREADS = 128 +W_WARP = 4 +X_WARP = 5 +S_WARP = 6 +MMA_WARP = 7 +EVICT_FIRST = 0x12F0000000000000 +assert TWO_I % ROWS1 == 0 and I_TP % 32 == 0 and I_PAD % 128 == 0 and I_PAD >= I_TP, (I_TP, I_PAD) +assert STAGES >= 4 and TMEM_COLS <= 512 and NCTA >= FC2_CTAS, ( + f"k3_moe_m1 does not fit i_tp {I_TP} at m_max {M_MAX}: {STAGES} ring stages, {TMEM_COLS} TMEM columns, {NCTA} CTAs" +) +assert M_MAX in (1, 2) and 4 * M_MAX <= 8, ( + M_MAX +) # FC2 B rows 4 t + j stay in one 8-row swizzle atom + + +@dsl_user_op +def _mul_rn(a, b, *, loc=None, ip=None): + """a * b rounded on its own (mul.rn is never fused into an FMA).""" + return cutlass.Float32(_llvm.inline_asm( + _T.f32(), [cutlass.Float32(a).ir_value(loc=loc, ip=ip), cutlass.Float32(b).ir_value(loc=loc, ip=ip)], + "mul.rn.f32 $0, $1, $2;", "=f,f,f", has_side_effects=False, is_align_stack=False, + asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + )) # fmt: skip + + +@dsl_user_op +def _add_rn(a, b, *, loc=None, ip=None): + """a + b rounded on its own (add.rn is never fused into an FMA).""" + return cutlass.Float32(_llvm.inline_asm( + _T.f32(), [cutlass.Float32(a).ir_value(loc=loc, ip=ip), cutlass.Float32(b).ir_value(loc=loc, ip=ip)], + "add.rn.f32 $0, $1, $2;", "=f,f,f", has_side_effects=False, is_align_stack=False, + asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + )) # fmt: skip + + +@dsl_user_op +def _red_release_add(addr, val, *, loc=None, ip=None): + _llvm.inline_asm( + None, [cutlass.Int64(addr).ir_value(loc=loc, ip=ip), cutlass.Int32(val).ir_value(loc=loc, ip=ip)], + "red.release.gpu.global.add.u32 [$0], $1;", "l,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _load_acquire(addr, *, loc=None, ip=None): + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [cutlass.Int64(addr).ir_value(loc=loc, ip=ip)], + "ld.acquire.gpu.global.u32 $0, [$1];", "=r,l", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _st_u32(addr, val, *, loc=None, ip=None): + _llvm.inline_asm( + None, [cutlass.Int64(addr).ir_value(loc=loc, ip=ip), cutlass.Int32(val).ir_value(loc=loc, ip=ip)], + "st.global.u32 [$0], $1;", "l,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _pack_bf16x2(hi, lo, *, loc=None, ip=None): + """(bf16(hi) << 16) | bf16(lo), round to nearest even.""" + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [hi.ir_value(loc=loc, ip=ip), lo.ir_value(loc=loc, ip=ip)], + "cvt.rn.bf16x2.f32 $0, $1, $2;", "=r,f,f", has_side_effects=False, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@cute.jit +def _emit(acc, lane, row2, tq, ch, out, lat_mc, lat_flags, lat_rank): + """Warp tq of an FC2 CTA (token tq's combine): lane l holds the value of output column row2 + 4 (l % 8) + l // 8, + stored to ``out`` at ``ch``; with PUSH, lanes 0, 2, 4, 6 gather columns row2 + 4 l .. row2 + 4 l + 7 and store + them as 16 bytes (bf16 pairs, -0.0 as +0.0) into this rank's slots of token tq in every rank's latent exchange + through its multicast mapping.""" + if cutlass.const_expr(PUSH): + bits = _pack_bf16x2(cutlass.Float32(0.0), acc) & cutlass.Int32(0xFFFF) + bits = cutlass.select_(bits == cutlass.Int32(0x8000), cutlass.Int32(0), bits) + b1 = cute.arch.shuffle_sync_down(bits, 8) + b2 = cute.arch.shuffle_sync_down(bits, 16) + b3 = cute.arch.shuffle_sync_down(bits, 24) + b4 = cute.arch.shuffle_sync_down(bits, 1) + b5 = cute.arch.shuffle_sync_down(bits, 9) + b6 = cute.arch.shuffle_sync_down(bits, 17) + b7 = cute.arch.shuffle_sync_down(bits, 25) + half = cutlass.Int32(lat_flags.load(idx=0, is_volatile=True)) & cutlass.Int32(1) + if (lane < 8) & (lane % 2 == 0): + vec = (bits | (b1 << cutlass.Int32(16)), b2 | (b3 << cutlass.Int32(16)), b4 | (b5 << cutlass.Int32(16)), + b6 | (b7 << cutlass.Int32(16))) # fmt: skip + for c in cutlass.range_constexpr(PUSH_COPIES): + slot = lat_rank * cutlass.Int32(PUSH_COPIES) + cutlass.Int32(c) + lat_mc.store(vec, idx=((half * cutlass.Int32(8) + tq) * cutlass.Int32(PUSH_WORLD) + slot) + * cutlass.Int32(LAT_ROW_WORDS) + (row2 + cutlass.Int32(4) * lane) // cutlass.Int32(2), + alignment=16) # fmt: skip + else: + out.store(cutlass.BFloat16(acc), idx=ch) + + +def _tanh_f32(x): + e = cute.math.exp2(cute.math.abs(x) * cutlass.Float32(-2.0 * _LOG2E), fastmath=True) + t = (cutlass.Float32(1.0) - e) * cute.arch.rcp_approx(cutlass.Float32(1.0) + e) + return cutlass.select_(x < cutlass.Float32(0.0), -t, t) + + +def _sigmoid_f32(x): + return cute.arch.rcp_approx( + cutlass.Float32(1.0) + cute.math.exp2(x * cutlass.Float32(-_LOG2E), fastmath=True) + ) + + +def _situ(gate, up): + g = ( + cutlass.Float32(SITU_GATE_CAP) + * _tanh_f32(gate * cutlass.Float32(1.0 / SITU_GATE_CAP)) + * _sigmoid_f32(gate) + ) + u = cutlass.Float32(SITU_LINEAR_CAP) * _tanh_f32(up * cutlass.Float32(1.0 / SITU_LINEAR_CAP)) + return g * u + + +def _block_e8m0(amax): + """E8M0 byte of an MX block and 2^(127 - byte) as f32 (k3_moe's ceil recipe).""" + sf = amax * cutlass.Float32(1.0 / 448.0) + sbits = cutlass.Int32(sf.bitcast(cutlass.Int32)) + sexp = (sbits >> cutlass.Int32(23)) & cutlass.Int32(0xFF) + mant = sbits & cutlass.Int32(0x7FFFFF) + byte = sexp + cutlass.select_(mant != cutlass.Int32(0), cutlass.Int32(1), cutlass.Int32(0)) + byte = cutlass.select_(byte > cutlass.Int32(0xFE), cutlass.Int32(0xFE), byte) + byte = cutlass.select_(amax > cutlass.Float32(0.0), byte, cutlass.Int32(0)) + inv = cutlass.Int32((cutlass.Int32(254) - byte) << cutlass.Int32(23)).bitcast(cutlass.Float32) + return byte, inv + + +@dsl_user_op +def _match_any(value, *, loc=None, ip=None): + """The mask of the warp's lanes holding the same 32-bit value.""" + return cutlass.Int32(_llvm.inline_asm( + _T.i32(), [cutlass.Int32(value).ir_value(loc=loc, ip=ip)], "match.any.sync.b32 $0, $1, 0xffffffff;", "=r,r", + has_side_effects=False, is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + )) # fmt: skip + + +@cute.jit +def _wait(bar: cutlass.Array, parity): + while not cute.arch.mbarrier_try_wait(bar.data_ptr(), parity): + pass + + +@cute.kernel +def k3_moe_m1_kernel( + tma_a1_desc: cutlass.GridConstant[cuda.TensorMap], + tma_b1_desc: cutlass.GridConstant[cuda.TensorMap], + tma_sfa1_desc: cutlass.GridConstant[cuda.TensorMap], + tma_a2_desc: cutlass.GridConstant[cuda.TensorMap], + sfb1_ptr: cutlass.Int64, + ids: cutlass.Array, # int32 [M, 16] + wts: cutlass.Array, # bf16 [M, 16] routing weights (as int16 bits) + w2s32: cutlass.Array, # int32 words of w2_weight_scale [E, H / 128, I_PAD / 128, 512 B] + hbuf: cutlass.Array, # int8 [G_CAP, M, H_ROW] + hbuf32: cutlass.Array, # the same memory as int32 words + counts: cutlass.Array, # int32 [2] + epochs: cutlass.Array, # int32 [NCTA] + out: cutlass.Array, # bf16 [M, 3584] (not written with PUSH) + offset: cutlass.Int32, + lat_mc: cutlass.Array, # PUSH: int32 words of the latent exchange's multicast mapping + lat_flags: cutlass.Array, # PUSH: int32 [4], [0] the reduce's call count + lat_rank: cutlass.Int32, +): + tidx, _, _ = cute.arch.thread_idx() + bx, _, _ = cute.arch.block_idx() + warp = cute.arch.make_warp_uniform(cute.arch.warp_idx()) + lane = tidx % 32 + + sA = cutlass.Array( + cutlass.Int8, NUM_BYTES_A * STAGES, space=cutlass.AddressSpace.smem, alignment=1024 + ) + sB = cutlass.Array( + cutlass.Int8, NUM_BYTES_B * STAGES, space=cutlass.AddressSpace.smem, alignment=1024 + ) + sSFA = cutlass.Array( + cutlass.Int8, NUM_BYTES_SFA * STAGES, space=cutlass.AddressSpace.smem, alignment=1024 + ) + sSFAraw = cutlass.Array(cutlass.Int32, NUM_BYTES_SFA_RAW * STAGES // 4, space=cutlass.AddressSpace.smem, + alignment=1024) # fmt: skip + sSFB = cutlass.Array( + cutlass.Int8, NUM_BYTES_SFB * STAGES, space=cutlass.AddressSpace.smem, alignment=1024 + ) + b2 = cutlass.Array( + cutlass.Int8, NT2 * NUM_BYTES_B2, space=cutlass.AddressSpace.smem, alignment=1024 + ) + sfa2 = cutlass.Array( + cutlass.Int8, NT2 * NUM_BYTES_SFA, space=cutlass.AddressSpace.smem, alignment=1024 + ) + sfb2 = cutlass.Array( + cutlass.Int8, NT2 * NUM_BYTES_SFB, space=cutlass.AddressSpace.smem, alignment=1024 + ) + s_part = cutlass.Array( + cutlass.Float32, M_MAX * ROWS1, space=cutlass.AddressSpace.smem, alignment=16 + ) + ys = cutlass.Array( + cutlass.Float32, G_CAP * M_MAX * R2, space=cutlass.AddressSpace.smem, alignment=16 + ) + ab_full = cutlass.Array(cutlass.Int64, STAGES, space=cutlass.AddressSpace.smem, alignment=8) + ab_empty = cutlass.Array(cutlass.Int64, STAGES, space=cutlass.AddressSpace.smem, alignment=8) + fc2_full = cutlass.Array(cutlass.Int64, STAGES, space=cutlass.AddressSpace.smem, alignment=8) + scales_in_tmem = cutlass.Array( + cutlass.Int64, STAGES, space=cutlass.AddressSpace.smem, alignment=8 + ) + acc_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + acc_empty = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + tmem_ready = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + sfa2_ready = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + b2_ready = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + acc2_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + localmax_smem = cutlass.Array(cutlass.Float32, 8 * M_MAX, space=cutlass.AddressSpace.smem) + tmem_ptr_i32 = cutlass.Array(cutlass.Int32, 1, space=cutlass.AddressSpace.smem) + s_el = cutlass.Array(cutlass.Int32, G_CAP + 1, space=cutlass.AddressSpace.smem, alignment=16) + s_w = cutlass.Array( + cutlass.Float32, G_CAP * M_MAX, space=cutlass.AddressSpace.smem, alignment=16 + ) + s_tmask = cutlass.Array(cutlass.Int32, G_CAP, space=cutlass.AddressSpace.smem, alignment=16) + s_perm = cutlass.Array(cutlass.Int32, G_CAP, space=cutlass.AddressSpace.smem, alignment=16) + + # ---- before the grid dependency: barriers, the TMEM allocation + if warp == 0: + if tidx < STAGES: + prims.mbarrier_init(ab_full.subview(tidx), 3) # A + SFA, B, SFB + prims.mbarrier_init(ab_empty.subview(tidx), 1) + prims.mbarrier_init(fc2_full.subview(tidx), 1) + prims.mbarrier_init(scales_in_tmem.subview(tidx), 1) + if tidx == 0: + prims.mbarrier_init(acc_full.subview(0), 1) + prims.mbarrier_init(acc_empty.subview(0), EPI_THREADS) + prims.mbarrier_init(tmem_ready.subview(0), 32) + prims.mbarrier_init(sfa2_ready.subview(0), EPI_THREADS) + prims.mbarrier_init(b2_ready.subview(0), EPI_THREADS) + prims.mbarrier_init(acc2_full.subview(0), 1) + prims.fence_mbarrier_init() + prims.barrier_cta_sync(0) + if warp == MMA_WARP: + prims.tcgen05_alloc(tmem_ptr_i32, TMEM_COLS) + prims.mbarrier_arrive(tmem_ready) + prims.tcgen05_relinquish_alloc_permit() + if warp == W_WARP: + prims.prefetch_tensormap(tma_a1_desc.get_ptr()) + prims.prefetch_tensormap(tma_sfa1_desc.get_ptr()) + prims.prefetch_tensormap(tma_a2_desc.get_ptr()) + if warp == X_WARP: + prims.prefetch_tensormap(tma_b1_desc.get_ptr()) + + cute.arch.griddepcontrol_wait() + cute.arch.griddepcontrol_launch_dependents() + + # ---- routing: lane l = top-k slot l % 16 of token l / 16. Slots: the distinct local experts in lane order (FC1 + # units, intermediate rows, FC2 groups); the combine walks them in ascending id through s_perm (below). + if warp == W_WARP: + loc_e = cutlass.Int32(-1) + wv = cutlass.Float32(0.0) + if lane < TOP_K * M_MAX: + loc_e = ids.load(idx=lane) - offset + wv = cutlass.Float32( + cutlass.Int32(cutlass.Int32(wts.load(idx=lane)) << cutlass.Int32(16)).bitcast( + cutlass.Float32 + ) + ) + is_local = ( + (lane < TOP_K * M_MAX) & (loc_e >= cutlass.Int32(0)) & (loc_e < cutlass.Int32(E_LOCAL)) + ) + first = is_local + leader = lane + if cutlass.const_expr(M_MAX > 1): + # An expert two tokens share: the lowest of its lanes stands for it; the others add their token. + same = _match_any(cutlass.select_(is_local, loc_e, cutlass.Int32(-1) - lane)) + leader = cute.arch.popc((same & (cutlass.Int32(0) - same)) - cutlass.Int32(1)) + first = is_local & (lane == leader) + bal = prims.vote_sync(0xFFFFFFFF, first, prims.VoteSync.BALLOT) + rank = cute.arch.popc(bal & cutlass.Int32(cute.arch.lanemask_lt())) + if cutlass.const_expr(M_MAX > 1): + for wz in cutlass.range_constexpr(G_CAP * M_MAX // 32): + s_w.store(cutlass.Float32(0.0), idx=lane + cutlass.Int32(32 * wz)) + cute.arch.sync_warp() + rank_l = cute.arch.shuffle_sync(rank, leader) + if is_local: + s_w.store(wv, idx=rank_l * cutlass.Int32(M_MAX) + lane // cutlass.Int32(TOP_K)) + if first: + s_el.store(loc_e, idx=rank) + s_tmask.store(cutlass.select_((same & cutlass.Int32(0xFFFF)) != cutlass.Int32(0), cutlass.Int32(1), + cutlass.Int32(0)) + | cutlass.select_((same & cutlass.Int32(-65536)) != cutlass.Int32(0), cutlass.Int32(2), + cutlass.Int32(0)), idx=rank) # fmt: skip + else: + if is_local: + s_el.store(loc_e, idx=rank) + s_w.store(wv, idx=rank) + if lane == 0: + s_el.store(cute.arch.popc(bal), idx=G_CAP) + # The first unit's ring fill before the CTA barrier: its expert (rank u / U1) comes from the lanes, not + # shared memory. + if bx < cute.arch.popc(bal) * cutlass.Int32(U1): + grp0 = bx // cutlass.Int32(U1) + m_src = prims.vote_sync(0xFFFFFFFF, first & (rank == grp0), prims.VoteSync.BALLOT) + e0 = cute.arch.shuffle_sync( + loc_e, cute.arch.popc((m_src & (cutlass.Int32(0) - m_src)) - cutlass.Int32(1)) + ) + t128_0 = (bx % cutlass.Int32(U1)) // cutlass.Int32(2) + coord_m0 = t128_0 * cutlass.Int32(MMA_M) + (bx % cutlass.Int32(2)) * cutlass.Int32( + ROWS1 + ) + if prims.elect_sync(): + for kp in cutlass.range_constexpr(PRE_FILL): + prims.mbarrier_arrive_expect_tx( + ab_full.subview(kp), 2 * (NUM_TX_A_HALF + NUM_BYTES_SFA) + ) + for q in cutlass.range_constexpr(2): + prims.cp_async_bulk_tensor_shared_cta_global( + sA.subview(cutlass.Int32(kp * NUM_BYTES_A + q * NUM_BYTES_A_HALF)), tma_a1_desc.get_ptr(), + (cutlass.Int32((2 * kp + q) * MMA_TILE_K), coord_m0, e0), ab_full.subview(kp), + l2_cache_hint=EVICT_FIRST, + ) # fmt: skip + prims.cp_async_bulk_tensor_shared_cta_global( + sSFAraw.subview(cutlass.Int32((kp * NUM_BYTES_SFA_RAW + q * NUM_BYTES_SFA) // 4)), + tma_sfa1_desc.get_ptr(), (cutlass.Int32(0), cutlass.Int32(2 * kp + q), t128_0, e0), + ab_full.subview(kp), l2_cache_hint=EVICT_FIRST, + ) # fmt: skip + # Ascending local id -> slot: s_perm[the distinct local experts below this one] = its slot. + key = cutlass.select_(first, loc_e, cutlass.Int32(1 << 30)) + rank_id = cutlass.Int32(0) + for k in range(TOP_K * M_MAX): + other = cute.arch.shuffle_sync(key, k) + if other < key: + rank_id = rank_id + cutlass.Int32(1) + if first: + s_perm.store(rank, idx=rank_id) + prims.barrier_cta_sync(0) + g_n = s_el.load(idx=G_CAP) + units = g_n * cutlass.Int32(U1) + mine = (units - bx + cutlass.Int32(NCTA - 1)) // cutlass.Int32(NCTA) + if bx >= units: + mine = cutlass.Int32(0) + ep = epochs.load(idx=bx) + slot_addr = counts.subview(0).data_ptr().toint() + cutlass.Int64((ep & 1) * 4) + row2 = bx * cutlass.Int32(R2) + is_fc2 = bx < cutlass.Int32(FC2_CTAS) + groups = (g_n + cutlass.Int32(3)) // cutlass.Int32(4) + nt2 = cutlass.select_(is_fc2, groups * cutlass.Int32(KT2), cutlass.Int32(0)) + blk2 = row2 // cutlass.Int32(128) + m1 = (row2 % cutlass.Int32(128)) // cutlass.Int32(32) + + # ---- weights producer (4): FC1 stages (the unit's 64 rows at k-tiles 2j, 2j + 1 and both scale atoms, raw), then + # this CTA's FC2 A tiles into the ring as it drains (group i, k-tile t: four experts' 32 rows at 4 KB offsets). + if warp == W_WARP: + if bx == 0: + if lane == 0: + _st_u32( + counts.subview(0).data_ptr().toint() + cutlass.Int64(((ep + 1) & 1) * 4), + cutlass.Int32(0), + ) + g = cutlass.select_( + mine > cutlass.Int32(0), cutlass.Int32(PRE_FILL), cutlass.Int32(0) + ) # issued pre-barrier + ab_empty_phase = 1 + for ui in range(mine): + u = bx + ui * cutlass.Int32(NCTA) + grp = u // cutlass.Int32(U1) + t128 = (u % cutlass.Int32(U1)) // cutlass.Int32(2) + half = u % cutlass.Int32(2) + coord_expert = s_el.load(idx=grp) + coord_m = t128 * cutlass.Int32(MMA_M) + half * cutlass.Int32(ROWS1) + for kp in cutlass.range(cutlass.select_(ui == cutlass.Int32(0), cutlass.Int32(PRE_FILL), cutlass.Int32(0)), + K1_PAIRS, unroll=1): # fmt: skip + stage = g % cutlass.Int32(STAGES) + if stage == cutlass.Int32(0) and g != cutlass.Int32(0): + ab_empty_phase = ab_empty_phase ^ 1 + while not cute.arch.mbarrier_try_wait( + ab_empty.subview(stage).data_ptr(), ab_empty_phase + ): + pass + coord_k = kp * cutlass.Int32(2 * MMA_TILE_K) + if prims.elect_sync(): + prims.mbarrier_arrive_expect_tx( + ab_full.subview(stage), 2 * (NUM_TX_A_HALF + NUM_BYTES_SFA) + ) + for q in cutlass.range_constexpr(2): + prims.cp_async_bulk_tensor_shared_cta_global( + sA.subview(stage * cutlass.Int32(NUM_BYTES_A) + cutlass.Int32(q * NUM_BYTES_A_HALF)), + tma_a1_desc.get_ptr(), (coord_k + cutlass.Int32(q * MMA_TILE_K), coord_m, coord_expert), + ab_full.subview(stage), l2_cache_hint=EVICT_FIRST, + ) # fmt: skip + prims.cp_async_bulk_tensor_shared_cta_global( + sSFAraw.subview( + (stage * cutlass.Int32(NUM_BYTES_SFA_RAW) + cutlass.Int32(q * NUM_BYTES_SFA)) // 4 + ), + tma_sfa1_desc.get_ptr(), (cutlass.Int32(0), kp * cutlass.Int32(2) + cutlass.Int32(q), t128, + coord_expert), + ab_full.subview(stage), l2_cache_hint=EVICT_FIRST, + ) # fmt: skip + g = g + cutlass.Int32(1) + for j in range(nt2): + stage = g % cutlass.Int32(STAGES) + if stage == cutlass.Int32(0) and g != cutlass.Int32(0): + ab_empty_phase = ab_empty_phase ^ 1 + while not cute.arch.mbarrier_try_wait( + ab_empty.subview(stage).data_ptr(), ab_empty_phase + ): + pass + i2 = j // cutlass.Int32(KT2) + t2 = j % cutlass.Int32(KT2) + n_here = g_n - i2 * cutlass.Int32(4) + if n_here > cutlass.Int32(4): + n_here = cutlass.Int32(4) + if prims.elect_sync(): + prims.mbarrier_arrive_expect_tx( + fc2_full.subview(stage), n_here * cutlass.Int32(NUM_TX_A2) + ) + for jj in cutlass.range_constexpr(4): + if cutlass.Int32(jj) < n_here: + prims.cp_async_bulk_tensor_shared_cta_global( + sA.subview(stage * cutlass.Int32(NUM_BYTES_A) + cutlass.Int32(jj * R2 * MMA_TILE_K)), + tma_a2_desc.get_ptr(), (t2 * cutlass.Int32(MMA_TILE_K), row2, + s_el.load(idx=i2 * cutlass.Int32(4) + cutlass.Int32(jj))), + fc2_full.subview(stage), l2_cache_hint=EVICT_FIRST, + ) # fmt: skip + g = g + cutlass.Int32(1) + + # ---- activations producer (5): the token's 128 K of k-tiles 2j, 2j + 1 into B rows 0 and 8, their scales. + if warp == X_WARP: + g = cutlass.Int32(0) + ab_empty_phase = 1 + for ui in range(mine): + for kp in cutlass.range(K1_PAIRS, unroll=1): + stage = g % cutlass.Int32(STAGES) + if stage == cutlass.Int32(0) and g != cutlass.Int32(0): + ab_empty_phase = ab_empty_phase ^ 1 + while not cute.arch.mbarrier_try_wait( + ab_empty.subview(stage).data_ptr(), ab_empty_phase + ): + pass + coord_k = kp * cutlass.Int32(2 * MMA_TILE_K) + if prims.elect_sync(): + prims.mbarrier_arrive_expect_tx(ab_full.subview(stage), 2 * M_MAX * MMA_TILE_K) + for tq in cutlass.range_constexpr(M_MAX): + for q in cutlass.range_constexpr(2): + prims.cp_async_bulk_tensor_shared_cta_global( + sB.subview( + stage * cutlass.Int32(NUM_BYTES_B) + cutlass.Int32((q * 8 + tq) * MMA_TILE_K) + ), + tma_b1_desc.get_ptr(), (coord_k + cutlass.Int32(q * MMA_TILE_K), cutlass.Int32(tq)), + ab_full.subview(stage), + ) # fmt: skip + if lane == 0: + for tq in cutlass.range_constexpr(M_MAX): + for q in cutlass.range_constexpr(2): + sfb_gmem = ( + sfb1_ptr + + cutlass.Int64(tq * (H // sf_vec_size)) + + cutlass.Int64( + (kp * cutlass.Int32(2) + cutlass.Int32(q)) * NUM_KBLOCKS + ) + ) + prims.cp_async_shared_global( + sSFB.subview(stage * cutlass.Int32(NUM_BYTES_SFB) + + cutlass.Int32((q * 8 + tq) * SFB_GROUP_BYTES)).data_ptr(), + cutlass.inttoptr(sfb_gmem, mem_space=1, dtype=sf_dtype), + size=4, modifier="ca", cp_size=4, + ) # fmt: skip + prims.cp_async_mbarrier_arrive(ab_full.subview(stage), noinc=True) + g = g + cutlass.Int32(1) + + # ---- scales to TMEM (6): the unit's half of both k-tiles' atoms spliced into one MMA atom (bytes 8 h .. 8 h + 7 + # of each 16-byte row group: k-tile 2j to MMA rows 0-63, k-tile 2j + 1 to rows 64-127), then SFA and SFB to TMEM. + if warp == S_WARP: + _wait(tmem_ready, 0) + tmem_raw_addr = tmem_ptr_i32.load() + base_col_id = tmem_raw_addr & 0xFFFF + base_row_id = tmem_raw_addr >> 16 + sfa_col_id0 = base_col_id + N + sfb_col_id0 = sfa_col_id0 + STAGES * SFA_COLS + s2t_shape, s2t_multicast = prims.S2TCopyMode.S2T_32x128b_WARPX4 + sSFA32 = cutlass.Array(sSFA.data_ptr(0), shape=(NUM_BYTES_SFA * STAGES // 4,), dtype=cutlass.Int32, + alignment=16) # fmt: skip + g = cutlass.Int32(0) + full_phase = 0 + for ui in range(mine): + u = bx + ui * cutlass.Int32(NCTA) + half = u % cutlass.Int32(2) + for kp in cutlass.range(K1_PAIRS, unroll=1): + stage = g % cutlass.Int32(STAGES) + if stage == cutlass.Int32(0) and g != cutlass.Int32(0): + full_phase = full_phase ^ 1 + while not cute.arch.mbarrier_try_wait( + ab_full.subview(stage).data_ptr(), full_phase + ): + pass + raw = ( + stage * cutlass.Int32(NUM_BYTES_SFA_RAW // 4) + + lane * cutlass.Int32(4) + + half * cutlass.Int32(2) + ) + lo = sSFAraw.load(idx=raw, vector_size=2, alignment=8) + hi = sSFAraw.load( + idx=raw + cutlass.Int32(NUM_BYTES_SFA // 4), vector_size=2, alignment=8 + ) + sSFA32.store((cutlass.Int32(lo[0]), cutlass.Int32(lo[1]), cutlass.Int32(hi[0]), cutlass.Int32(hi[1])), + idx=stage * cutlass.Int32(NUM_BYTES_SFA // 4) + lane * cutlass.Int32(4), + alignment=16) # fmt: skip + prims.fence_proxy("async_shared", space=prims.SharedSpace.shared_cta) + cute.arch.sync_warp() + prims.tcgen05_fence( + prims.Tcgen05Fence.AFTER_THREAD_SYNC + ) # the lanes' spliced scale atoms + sfa_tmem_ptr = cutlass.inttoptr( + (base_row_id << 16) | (sfa_col_id0 + stage * SFA_COLS), 6, cutlass.Int32 + ) + sfb_tmem_ptr = cutlass.inttoptr( + (base_row_id << 16) | (sfb_col_id0 + stage * SFB_COLS), 6, cutlass.Int32 + ) + desc_a = prims.Tcgen05SmemDesc.build( + sSFA.subview(stage * NUM_BYTES_SFA), leading_byte_offset=16, stride_byte_offset=128, + base_offset=0, layout=0, + ) # fmt: skip + desc_b = prims.Tcgen05SmemDesc.build( + sSFB.subview(stage * NUM_BYTES_SFB), leading_byte_offset=16, stride_byte_offset=128, + base_offset=0, layout=0, + ) # fmt: skip + if prims.elect_sync(): + prims.tcgen05_cp(s2t_shape, sfa_tmem_ptr, desc_a, multicast=s2t_multicast) + prims.tcgen05_cp(s2t_shape, sfb_tmem_ptr, desc_b, multicast=s2t_multicast) + prims.tcgen05_commit(scales_in_tmem.subview(stage)) + g = g + cutlass.Int32(1) + + # ---- MMA (7): FC1 per unit; then FC2: the spliced weight scales and (after the hand-off) the intermediate's + # scales to TMEM, then per A tile the group's K 32 MMAs into its 16 accumulator columns. + if warp == MMA_WARP: + tmem_raw_addr = tmem_ptr_i32.load() + acc_tmem_ptr = cutlass.inttoptr(tmem_raw_addr, 6, cutlass.Float32) + idesc = prims.Tcgen05MxInstrDesc.build( + a_dtype=a_dtype, b_dtype=b_dtype, scale_format=1, n_dim=N, m_dim=MMA_M + ) + base_col_id = tmem_raw_addr & 0xFFFF + base_row_id = tmem_raw_addr >> 16 + sfa_col_id0 = base_col_id + N + sfb_col_id0 = sfa_col_id0 + STAGES * SFA_COLS + g = cutlass.Int32(0) + st_phase = 0 + acc_empty_phase = 1 + for ui in range(mine): + while not cute.arch.mbarrier_try_wait(acc_empty.data_ptr(), acc_empty_phase): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc_empty_phase = acc_empty_phase ^ 1 + scale_d = False + for kp in cutlass.range(K1_PAIRS, unroll=1): + stage = g % cutlass.Int32(STAGES) + if stage == cutlass.Int32(0) and g != cutlass.Int32(0): + st_phase = st_phase ^ 1 + while not cute.arch.mbarrier_try_wait( + scales_in_tmem.subview(stage).data_ptr(), st_phase + ): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + sfa_base = (base_row_id << 16) | (sfa_col_id0 + stage * SFA_COLS) + sfb_base = (base_row_id << 16) | (sfb_col_id0 + stage * SFB_COLS) + desc_a_base = prims.Tcgen05SmemDesc.build( + sA.subview(stage * NUM_BYTES_A), leading_byte_offset=16, stride_byte_offset=1024, base_offset=0, + layout=2, + ) # fmt: skip + desc_b_base = prims.Tcgen05SmemDesc.build( + sB.subview(stage * NUM_BYTES_B), leading_byte_offset=16, stride_byte_offset=1024, base_offset=0, + layout=2, + ) # fmt: skip + for kb in cutlass.range(NUM_KBLOCKS, unroll_full=True): + sf_inside = kb % NUM_SF_IDS + sf_col = kb // NUM_SF_IDS + sfa_tmem_ptr = cutlass.inttoptr(sfa_base + sf_col * SFA_COLS, 6, cutlass.Int32) + sfb_tmem_ptr = cutlass.inttoptr(sfb_base + sf_col * SFB_COLS, 6, cutlass.Int32) + idesc_u = idesc.set_sf_ids(a_sf_id=sf_inside, b_sf_id=sf_inside) + inc = ((MMA_INST_K * a_smem_width // 8) >> 4) * kb + if prims.elect_sync(): + prims.tcgen05_mma_block_scale( + prims.MMABlockScaleKind.MXF8F6F4, prims.CTAGroup.CTA_1, acc_tmem_ptr, + desc_a_base + inc, desc_b_base + inc, idesc_u, scale_d, sfa_tmem_ptr, sfb_tmem_ptr, + ) # fmt: skip + scale_d = True + if prims.elect_sync(): + prims.tcgen05_commit(ab_empty.subview(stage)) + g = g + cutlass.Int32(1) + if prims.elect_sync(): + prims.tcgen05_commit(acc_full) + if nt2 > cutlass.Int32(0): + s2t_shape, s2t_multicast = prims.S2TCopyMode.S2T_32x128b_WARPX4 + g0 = g + if cutlass.const_expr(NT2 <= STAGES): + # Every A tile has its own ring stage: wait for all of them and copy the weight scales before B is + # ready, then one thread issues the B scale copies and every MMA back to back. + for j in cutlass.range_constexpr(NT2): + if cutlass.Int32(j) < nt2: + _wait(fc2_full.subview((g0 + cutlass.Int32(j)) % cutlass.Int32(STAGES)), 0) + _wait(sfa2_ready, 0) + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + if prims.elect_sync(): + for j in cutlass.range_constexpr(NT2): + if cutlass.Int32(j) < nt2: + prims.tcgen05_cp( + s2t_shape, + cutlass.inttoptr((base_row_id << 16) | (base_col_id + SFA2_COL + j * SFA_COLS), 6, + cutlass.Int32), + prims.Tcgen05SmemDesc.build(sfa2.subview(j * NUM_BYTES_SFA), leading_byte_offset=16, + stride_byte_offset=128, base_offset=0, layout=0), + multicast=s2t_multicast, + ) # fmt: skip + _wait(b2_ready, 0) + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + if prims.elect_sync(): + for j in cutlass.range_constexpr(NT2): + if cutlass.Int32(j) < nt2: + prims.tcgen05_cp( + s2t_shape, + cutlass.inttoptr((base_row_id << 16) | (base_col_id + SFB2_COL + j * SFB_COLS), 6, + cutlass.Int32), + prims.Tcgen05SmemDesc.build(sfb2.subview(j * NUM_BYTES_SFB), leading_byte_offset=16, + stride_byte_offset=128, base_offset=0, layout=0), + multicast=s2t_multicast, + ) # fmt: skip + for j in cutlass.range_constexpr(NT2): + if cutlass.Int32(j) < nt2: + stage = (g0 + cutlass.Int32(j)) % cutlass.Int32(STAGES) + desc_a_base = prims.Tcgen05SmemDesc.build( + sA.subview(stage * NUM_BYTES_A), leading_byte_offset=16, stride_byte_offset=1024, + base_offset=0, layout=2, + ) # fmt: skip + desc_b_base = prims.Tcgen05SmemDesc.build( + b2.subview(j * NUM_BYTES_B2), leading_byte_offset=16, stride_byte_offset=1024, + base_offset=0, layout=2, + ) # fmt: skip + acc2_ptr = cutlass.inttoptr( + (base_row_id << 16) | (base_col_id + ACC2_COL + (j // KT2) * N), 6, cutlass.Float32 + ) # fmt: skip + sfa_tmem_ptr = cutlass.inttoptr( + (base_row_id << 16) | (base_col_id + SFA2_COL + j * SFA_COLS), 6, cutlass.Int32 + ) # fmt: skip + sfb_tmem_ptr = cutlass.inttoptr( + (base_row_id << 16) | (base_col_id + SFB2_COL + j * SFB_COLS), 6, cutlass.Int32 + ) # fmt: skip + for kb2 in cutlass.range_constexpr( + min(NUM_KBLOCKS, KB2 - NUM_KBLOCKS * (j % KT2)) + ): + prims.tcgen05_mma_block_scale( + prims.MMABlockScaleKind.MXF8F6F4, prims.CTAGroup.CTA_1, acc2_ptr, + desc_a_base + ((MMA_INST_K * a_smem_width // 8) >> 4) * kb2, + desc_b_base + ((MMA_INST_K * a_smem_width // 8) >> 4) * kb2, + idesc.set_sf_ids(a_sf_id=kb2, b_sf_id=kb2), (j % KT2) != 0 or kb2 != 0, + sfa_tmem_ptr, sfb_tmem_ptr, + ) # fmt: skip + prims.tcgen05_commit(acc2_full) + else: + # More A tiles than ring stages (TP4 x EP4 shapes): tiles cycle through the ring, each stage freed by + # its MMAs' commit. + _wait(sfa2_ready, 0) + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + if prims.elect_sync(): + for j in range(nt2): + prims.tcgen05_cp( + s2t_shape, + cutlass.inttoptr( + (base_row_id << 16) | (base_col_id + SFA2_COL + j * SFA_COLS), 6, cutlass.Int32 + ), + prims.Tcgen05SmemDesc.build(sfa2.subview(j * NUM_BYTES_SFA), leading_byte_offset=16, + stride_byte_offset=128, base_offset=0, layout=0), + multicast=s2t_multicast, + ) # fmt: skip + _wait(b2_ready, 0) + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + if prims.elect_sync(): + for j in range(nt2): + prims.tcgen05_cp( + s2t_shape, + cutlass.inttoptr( + (base_row_id << 16) | (base_col_id + SFB2_COL + j * SFB_COLS), 6, cutlass.Int32 + ), + prims.Tcgen05SmemDesc.build(sfb2.subview(j * NUM_BYTES_SFB), leading_byte_offset=16, + stride_byte_offset=128, base_offset=0, layout=0), + multicast=s2t_multicast, + ) # fmt: skip + for j in range(nt2): + stage = g % cutlass.Int32(STAGES) + while not cute.arch.mbarrier_try_wait(fc2_full.subview(stage).data_ptr(), + (j // cutlass.Int32(STAGES)) % cutlass.Int32(2)): # fmt: skip + pass + i2 = j // cutlass.Int32(KT2) + t2 = j % cutlass.Int32(KT2) + acc2_ptr = cutlass.inttoptr( + (base_row_id << 16) | (base_col_id + ACC2_COL + i2 * N), 6, cutlass.Float32 + ) + sfa_tmem_ptr = cutlass.inttoptr((base_row_id << 16) | (base_col_id + SFA2_COL + j * SFA_COLS), 6, + cutlass.Int32) # fmt: skip + sfb_tmem_ptr = cutlass.inttoptr((base_row_id << 16) | (base_col_id + SFB2_COL + j * SFB_COLS), 6, + cutlass.Int32) # fmt: skip + desc_a_base = prims.Tcgen05SmemDesc.build( + sA.subview(stage * NUM_BYTES_A), leading_byte_offset=16, stride_byte_offset=1024, base_offset=0, + layout=2, + ) # fmt: skip + desc_b_base = prims.Tcgen05SmemDesc.build( + b2.subview(j * NUM_BYTES_B2), leading_byte_offset=16, stride_byte_offset=1024, base_offset=0, + layout=2, + ) # fmt: skip + nkb = cutlass.Int32(KB2) - t2 * cutlass.Int32(NUM_KBLOCKS) + for kb2 in cutlass.range_constexpr(NUM_KBLOCKS): + if cutlass.Int32(kb2) < nkb: + idesc_u = idesc.set_sf_ids(a_sf_id=kb2, b_sf_id=kb2) + inc = ((MMA_INST_K * a_smem_width // 8) >> 4) * kb2 + acc2_on = (t2 != cutlass.Int32(0)) | ( + cutlass.Int32(kb2) != cutlass.Int32(0) + ) + if prims.elect_sync(): + prims.tcgen05_mma_block_scale( + prims.MMABlockScaleKind.MXF8F6F4, prims.CTAGroup.CTA_1, acc2_ptr, + desc_a_base + inc, desc_b_base + inc, idesc_u, acc2_on, sfa_tmem_ptr, sfb_tmem_ptr, + ) # fmt: skip + if prims.elect_sync(): + prims.tcgen05_commit(ab_empty.subview(stage)) + g = g + cutlass.Int32(1) + if prims.elect_sync(): + prims.tcgen05_commit(acc2_full) + + # ---- epilogue warps (0-3): FC2 weight scales spliced during FC1; the FC1 epilogue per unit; then the FC2 hand-off. + if warp < 4: + _wait(tmem_ready, 0) + tmem_raw_addr = tmem_ptr_i32.load() + base_col_id = tmem_raw_addr & 0xFFFF + base_row_id = tmem_raw_addr >> 16 + row_id_with_warp_offset = base_row_id + warp * 32 + # FC2 weight scale atoms: MMA atom j (group i = j / KT2, k-tile t) word 4 m0 + jj = slot 4 i + jj's w2 atom + # (its 128-row block, k-atom t) word 4 m0 + m1, m1 = this CTA's 32-row block within the 128. + if is_fc2: + sfa2_32 = cutlass.Array(sfa2.data_ptr(0), shape=(NT2 * NUM_BYTES_SFA // 4,), dtype=cutlass.Int32, + alignment=16) # fmt: skip + vals = [] + for r in cutlass.range_constexpr(NT2 * (NUM_BYTES_SFA // 4) // EPI_THREADS): + w = tidx + cutlass.Int32(r * EPI_THREADS) + j = w // cutlass.Int32(NUM_BYTES_SFA // 4) + rem = w % cutlass.Int32(NUM_BYTES_SFA // 4) + slot = (j // cutlass.Int32(KT2)) * cutlass.Int32(4) + rem % cutlass.Int32(4) + ok = (j < nt2) & (slot < g_n) + e = cutlass.Int32(s_el.load(idx=cutlass.select_(ok, slot, cutlass.Int32(0)))) + src = (((e * cutlass.Int32(H // 128) + blk2) * cutlass.Int32(KA2) + j % cutlass.Int32(KT2)) + * cutlass.Int32(NUM_BYTES_SFA // 4) + + (rem // cutlass.Int32(4)) * cutlass.Int32(4) + m1) # fmt: skip + v = w2s32.load(idx=cutlass.select_(ok, src, cutlass.Int32(0))) + vals.append(cutlass.select_(ok, cutlass.Int32(v), cutlass.Int32(0))) + for r in cutlass.range_constexpr(NT2 * (NUM_BYTES_SFA // 4) // EPI_THREADS): + sfa2_32.store(vals[r], idx=tidx + cutlass.Int32(r * EPI_THREADS)) + prims.fence_proxy("async_shared", space=prims.SharedSpace.shared_cta) + prims.mbarrier_arrive(sfa2_ready) + # FC1 epilogue: lanes 0-63 column 0 (k-tiles 2j) + lanes 64-127 column 8 (k-tiles 2j + 1) = the unit's 64 rows; + # warps 0-1 apply k3_moe's SiTU + MXFP8 requant to the 32 intermediate columns (one MX block). + is_up = ((lane // 8) % 2) == 0 + up_mask = cutlass.select_(is_up, cutlass.Float32(1.0), cutlass.Float32(0.0)) + fc1_col_in_tile = warp * 16 + 2 * (lane % 8) + lane // 16 + acc_full_phase = 0 + for ui in range(mine): + u = bx + ui * cutlass.Int32(NCTA) + grp = u // cutlass.Int32(U1) + t128 = (u % cutlass.Int32(U1)) // cutlass.Int32(2) + half = u % cutlass.Int32(2) + while not cute.arch.mbarrier_try_wait(acc_full.data_ptr(), acc_full_phase): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc_full_phase = acc_full_phase ^ 1 + tmem_ld = cutlass.inttoptr( + (row_id_with_warp_offset << 16) | base_col_id, 6, cutlass.Float32 + ) + t2r_rmem = prims.tcgen05_ld("32x32b", tmem_ld, num=N) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + prims.mbarrier_arrive(acc_empty) + if warp >= 2: + for tq in cutlass.range_constexpr(M_MAX): + s_part.store( + cutlass.Float32(t2r_rmem[8 + tq]), idx=tq * ROWS1 + (warp - 2) * 32 + lane + ) + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + res = [] + for tq in cutlass.range_constexpr(M_MAX): + xn = cutlass.Float32(t2r_rmem[tq]) + s_part.load( + idx=tq * ROWS1 + (warp % 2) * 32 + lane + ) + partner = cute.arch.shuffle_sync_bfly(xn, 8) + res.append(_situ(partner, xn)) + absv = cute.math.abs(res[tq]) * up_mask + warp_amax = prims.redux_sync(absv, prims.ReductionKind.FMAX, 0xFFFFFFFF, abs=True) + if lane == 0: + localmax_smem.store(warp_amax, idx=tq * 8 + warp) + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + for tq in cutlass.range_constexpr(M_MAX): + block_amax = cute.arch.fmax( + localmax_smem.load(idx=tq * 8), localmax_smem.load(idx=tq * 8 + 1) + ) + byte, inv_scale = _block_e8m0(block_amax) + qv = cute.arch.fmax(cute.arch.fmin(res[tq] * inv_scale, cutlass.Float32(E4M3_MAX)), + cutlass.Float32(-E4M3_MAX)) # fmt: skip + fp8_i8 = cutlass.Float8E4M3FN(qv).bitcast(cutlass.Int8) + if fp8_i8 == cutlass.Int8(FP8_SENTINEL_I8): + fp8_i8 = cutlass.Int8(0) + hrow = (grp * cutlass.Int32(M_MAX) + cutlass.Int32(tq)) * cutlass.Int32(H_ROW) + if (warp < 2) & is_up: + hbuf.store(fp8_i8, idx=hrow + t128 * cutlass.Int32(MMA_M // 2) + half * cutlass.Int32(32) + + fc1_col_in_tile, alignment=1) # fmt: skip + if tidx == 0: + hbuf.store( + cutlass.Int8(byte & cutlass.Int32(0xFF)), + idx=hrow + cutlass.Int32(I_TP) + t128 * cutlass.Int32(2) + half, + alignment=1, + ) + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + if tidx == 0: + _red_release_add(slot_addr, cutlass.Int32(1)) + # FC2: the call's count reaches U1 G; the intermediate rows into the B tiles (row n of tile (i, t) = slot + # 4 i + n's K 128 t .., 128B-swizzled) and their scales into the B scale atoms (byte 16 n + k = block 4 t + k), + # one load round; then the MMAs (warp 7), the accumulators and the combine. No local expert: zero rows. + if is_fc2 & (g_n == cutlass.Int32(0)): + if warp < M_MAX: + _emit(cutlass.Float32(0.0), lane, row2, warp, warp * cutlass.Int32(H) + row2 + lane, out, lat_mc, + lat_flags, lat_rank) # fmt: skip + if is_fc2 & (g_n > cutlass.Int32(0)): + if tidx == 0: + while _load_acquire(slot_addr) < units: + pass + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + b2_32 = cutlass.Array( + b2.data_ptr(0), shape=(NT2 * NUM_BYTES_B2 // 4,), dtype=cutlass.Int32, alignment=16 + ) + sfb2_32 = cutlass.Array(sfb2.data_ptr(0), shape=(NT2 * NUM_BYTES_SFB // 4,), dtype=cutlass.Int32, + alignment=16) # fmt: skip + got = [] + for r in cutlass.range_constexpr(STAGE_ROUNDS): + qi = tidx + cutlass.Int32(r * EPI_THREADS) + rr = qi // cutlass.Int32(VCH + SCH) # intermediate row: slot rr / M, token rr % M + ok = rr < g_n * cutlass.Int32(M_MAX) + c = qi % cutlass.Int32(VCH + SCH) + v4 = prims.load_ext( + hbuf32.subview(cutlass.select_( + ok, (rr * cutlass.Int32(H_ROW) + c * cutlass.Int32(16)) // cutlass.Int32(4), cutlass.Int32(0) + )), + dtype=cutlass.Int32, count=4, order="relaxed", scope="gpu", + ) # fmt: skip + got.append(v4) + for r in cutlass.range_constexpr(STAGE_ROUNDS): + qi = tidx + cutlass.Int32(r * EPI_THREADS) + rr = qi // cutlass.Int32(VCH + SCH) + sl = rr // cutlass.Int32(M_MAX) + c = qi % cutlass.Int32(VCH + SCH) + n = (rr % cutlass.Int32(M_MAX)) * cutlass.Int32(4) + sl % cutlass.Int32( + 4 + ) # B row: 4 token + j + v4 = got[r] + if sl < g_n: + if c < cutlass.Int32(VCH): + t2 = c // cutlass.Int32(8) + cc = c % cutlass.Int32(8) + j = (sl // cutlass.Int32(4)) * cutlass.Int32(KT2) + t2 + b2_32.store((cutlass.Int32(v4[0]), cutlass.Int32(v4[1]), + cutlass.Int32(v4[2]), cutlass.Int32(v4[3])), + idx=(j * cutlass.Int32(NUM_BYTES_B2) + n * cutlass.Int32(MMA_TILE_K) + + (cc ^ n) * cutlass.Int32(16)) // cutlass.Int32(4), alignment=16) # fmt: skip + else: + for wq in cutlass.range_constexpr(4): + t2 = (c - cutlass.Int32(VCH)) * cutlass.Int32(4) + cutlass.Int32(wq) + if t2 < cutlass.Int32(KT2): + j = (sl // cutlass.Int32(4)) * cutlass.Int32(KT2) + t2 + sfb2_32.store( + cutlass.Int32(v4[wq]), + idx=(j * cutlass.Int32(NUM_BYTES_SFB) + n * cutlass.Int32(16)) + // cutlass.Int32(4), + ) + prims.fence_proxy("async_shared", space=prims.SharedSpace.shared_cta) + prims.mbarrier_arrive(b2_ready) + _wait(acc2_full, 0) + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + yv = [] + for gi in cutlass.range_constexpr(GROUPS2): + for tq in cutlass.range_constexpr(M_MAX): + yv.append(prims.tcgen05_ld("32x32b", cutlass.inttoptr( + (row_id_with_warp_offset << 16) | (base_col_id + ACC2_COL + gi * N + 4 * tq + warp), 6, + cutlass.Float32), num=1)) # fmt: skip + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + for gi in cutlass.range_constexpr(GROUPS2): + sl = cutlass.Int32(4 * gi) + warp + if sl < g_n: + for tq in cutlass.range_constexpr(M_MAX): + ys.store( + cutlass.Float32(yv[gi * M_MAX + tq][0]), + idx=(sl * cutlass.Int32(M_MAX) + cutlass.Int32(tq)) * cutlass.Int32(R2) + + lane, + ) + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + if warp < M_MAX: + tq = warp # this warp's token + ch = ( + tq * cutlass.Int32(H) + + row2 + + cutlass.Int32(4) * (lane % cutlass.Int32(8)) + + lane // cutlass.Int32(8) + ) + fast = cutlass.Boolean(False) + if cutlass.const_expr(M_MAX == 1): + fast = g_n == cutlass.Int32(G_CAP) + if fast: + # All G_CAP experts present (every TP16 call): k3_moe's slices at trace time; every product and sum + # rounded on its own, as in k3_moe (whose selects keep them apart). + acc = cutlass.Float32(0.0) + for si in cutlass.range_constexpr(REF_SLICES): + part = cutlass.Float32(0.0) + for sc in cutlass.range_constexpr( + _REF_SLICES_FULL[si][0], _REF_SLICES_FULL[si][1] + ): + slot = s_perm.load(idx=sc) + part = _add_rn( + part, + _mul_rn( + ys.load(idx=slot * cutlass.Int32(R2) + lane), s_w.load(idx=slot) + ), + ) + acc = _add_rn(acc, part) + _emit(acc, lane, row2, tq, ch, out, lat_mc, lat_flags, lat_rank) + else: + # k3_moe's sum tree: min(G, 5) slices of consecutive experts + # (slice s = experts [s G / S, (s + 1) G / S)), each summed from 0 in ascending id (an expert + # without this token adds zero), then the slices summed in order. + n_sl = cutlass.select_( + g_n < cutlass.Int32(REF_SLICES), g_n, cutlass.Int32(REF_SLICES) + ) + n_sl = cutlass.select_(n_sl < cutlass.Int32(1), cutlass.Int32(1), n_sl) + bnd = [] + for b in cutlass.range_constexpr(1, REF_SLICES): + bnd.append( + cutlass.select_( + cutlass.Int32(b) < n_sl, + cutlass.Int32(b) * g_n // n_sl, + cutlass.Int32(-1), + ) + ) + zero = cutlass.Float32(0.0) + acc = zero + part = zero + for sc in cutlass.range_constexpr(G_CAP): + if cutlass.const_expr(sc > 0): + start = bnd[0] == cutlass.Int32(sc) + for b in cutlass.range_constexpr(1, REF_SLICES - 1): + start = start | (bnd[b] == cutlass.Int32(sc)) + acc = cutlass.select_(start, acc + part, acc) + part = cutlass.select_(start, zero, part) + slot = s_perm.load( + idx=cutlass.select_( + cutlass.Int32(sc) < g_n, cutlass.Int32(sc), cutlass.Int32(0) + ) + ) + valid = cutlass.Int32(sc) < g_n + if cutlass.const_expr(M_MAX > 1): + valid = valid & ( + ((s_tmask.load(idx=slot) >> tq) & cutlass.Int32(1)) + != cutlass.Int32(0) + ) + pq = slot * cutlass.Int32(M_MAX) + tq + part = part + cutlass.select_( + valid, ys.load(idx=pq * cutlass.Int32(R2) + lane) * s_w.load(idx=pq), zero + ) # fmt: skip + acc = acc + part + _emit(acc, lane, row2, tq, ch, out, lat_mc, lat_flags, lat_rank) + + # ---- teardown + if tidx == 0: + epochs.store(ep + cutlass.Int32(1), idx=bx) + prims.barrier_cta_sync(0) + if warp == MMA_WARP: + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + prims.tcgen05_dealloc(cutlass.inttoptr(tmem_ptr_i32.load(), 6, cutlass.Float32), TMEM_COLS) + + +@cute.jit +def k3_moe_m1( + a1_tensor: cute.Tensor, # w3_w1_weight viewed (H/2, 2 I_PAD, E) FP4 bytes, K-major + b1_tensor: cute.Tensor, # MXFP8 activations viewed (H, M) FP8 + sfa1_tensor: cute.Tensor, # w3_w1_weight_scale viewed (512, H/128, 2 I_PAD/128, E) + sfb1_tensor: cute.Tensor, # activation scales (M, H/32) E8M0 + a2_tensor: cute.Tensor, # w2_weight viewed (I_PAD/2, H, E) FP4 bytes, K-major + ids: cute.Tensor, wts: cute.Tensor, w2s32: cute.Tensor, hbuf: cute.Tensor, hbuf32: cute.Tensor, + counts: cute.Tensor, epochs: cute.Tensor, out: cute.Tensor, + lat_mc: cute.Tensor, # PUSH: int32 words of the latent exchange's multicast mapping (else any int32 tensor) + lat_flags: cute.Tensor, # PUSH: int32 [4] of the exchange (else any int32 tensor) + offset: cutlass.Int32, lat_rank: cutlass.Int32, stream: cuda_driver.CUstream, +): # fmt: skip + _kpp = a1_tensor.shape[0] + _mw = a1_tensor.shape[1] + _ew = a1_tensor.shape[2] + tma_a1_desc = cuda.create_tensor_map_tiled( + global_address=a1_tensor.iterator.toint(), dtype=a_dtype, global_dims=[_kpp * 2, _mw, _ew], + global_strides=[_kpp // 16, (_mw * _kpp) // 16], box_dims=(MMA_TILE_K, ROWS1, 1), + swizzle=cuda.TensorMapSwizzle.s128b, tma_format=TensorMapDataType.f416u4_align16b, + ) # fmt: skip + tma_b1_desc = cuda.create_tensor_map_tiled_from_view( + b1_tensor, + box_dims=(MMA_TILE_K, 1), + stride_order=(0, 1), + swizzle=cuda.TensorMapSwizzle.s128b, + ) + sfa1_fp16 = cute.recast_tensor(sfa1_tensor, cutlass.Uint16) + tma_sfa1_desc = cuda.create_tensor_map_tiled_from_view( + sfa1_fp16, box_dims=(num_elts_atom_sf_fp16, 1, 1, 1), stride_order=(0, 1, 2, 3), + swizzle=cuda.TensorMapSwizzle.none, + ) # fmt: skip + _kp2 = a2_tensor.shape[0] + _h2 = a2_tensor.shape[1] + _e2 = a2_tensor.shape[2] + tma_a2_desc = cuda.create_tensor_map_tiled( + global_address=a2_tensor.iterator.toint(), dtype=a_dtype, global_dims=[_kp2 * 2, _h2, _e2], + global_strides=[_kp2 // 16, (_h2 * _kp2) // 16], box_dims=(MMA_TILE_K, R2, 1), + swizzle=cuda.TensorMapSwizzle.s128b, tma_format=TensorMapDataType.f416u4_align16b, + ) # fmt: skip + k3_moe_m1_kernel( + tma_a1_desc, tma_b1_desc, tma_sfa1_desc, tma_a2_desc, sfb1_tensor.iterator.toint(), ids, wts, w2s32, hbuf, + hbuf32, counts, epochs, out, offset, lat_mc, lat_flags, lat_rank, + ).launch(grid=(NCTA, 1, 1), block=(THREADS, 1, 1), cluster=(1, 1, 1), stream=stream, use_pdl=True) # fmt: skip diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_moe_m2_kernel.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_moe_m2_kernel.py new file mode 100644 index 000000000000..5ad79e6fec4d --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/k3_moe_m2_kernel.py @@ -0,0 +1,1129 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Kimi K3 routed experts of two decode tokens as a weight-stream kernel: ``k3_moe_m2``. + +At M = 2 every expert a token routes to is a GEMV, so the kernel spreads the experts' weight rows over all CTAs +instead of k3_moe's 128-row tiles. It reads the outputs of trtllm::k3_moe_front (or trtllm::k3_route_quant): the +tokens' top-16 global ids and bf16 weights and their MXFP8 latents. It computes this rank's routed partial [M, 3584] +bf16, the tensor k3_moe returns (M <= 2; the op builds it for 2). Weights are read in place in the TRTLLM-Gen +W4A8_MXFP4_MXFP8 layout (see k3_moe_kernel.py), including the loader's zero padding of a rank's intermediate to a +multiple of 128 (i_pad, e.g. 192 -> 256 at TP16). The kernel streams the i_tp real values only. + +- Routing: one warp holds the 2 x 16 (token, top-k) pairs; each lane ranks its expert among the warp's first + occurrences, so the distinct local experts take slots in ascending id, each slot holding both tokens' routing + weights (zero where a token does not route there). The first FC1 unit's ring fill is issued from the lanes before + the CTA barrier. +- FC1 (gate_up, MXFP4 x MXFP8) in 64-row units with k3_moe's per-stage block-scaled machinery, two k-tiles per + stage: the unit's 64 rows at k-tile 2j fill MMA rows 0-63 and at k-tile 2j + 1 rows 64-127; B rows t and 8 + t + hold token t at those k-tiles (N 16), and the scale atoms are spliced to match. Every expert's intermediate is + computed for both tokens; the MMA columns are independent, so a routed token's bits do not depend on the other. + Units past the first wave (one unit per CTA) go first to the CTAs without FC2 work. The epilogue adds the two + partial sums per token, then applies k3_moe's SiTU and MXFP8 requantization, into one intermediate row per + (expert slot, token), and adds one release to its FC2 group's counter (two counter sets by epoch parity; CTA 0 + re-arms the other set, every CTA advances its epoch). +- FC2 (down): each of 112 CTAs owns one 32-row block of the shuffled down projection. Four experts' 32-row slices + stack in one 128-row MMA (rows 32 j .. 32 j + 31 = slot 4 i + j) against B N 8, whose row 2 j + t is token t's + intermediate for slot 4 i + j; this is k3_moe's FC2 operand, so each expert's down projection has k3_moe's bits. + One warp loads each group's intermediate rows as soon as that group's counter is complete: the epilogue warps take + the groups whose FC1 units all run in the first wave, warps 5-6 (idle after FC1) the rest, so each group's MMAs + start without waiting for the others. The A tiles load into the FC1 ring as it drains, evict-first, and the weight + scale atoms are spliced while FC1 runs. +- Combine: per token, the routing-weighted experts are summed in k3_moe's order (ascending local id in min(G, 5) + slices, an expert without the token adding zero, products and sums rounded on their own), each group folded in + as soon as its MMAs complete. The output matches k3_moe's bits except where the two FC1 partial sums round an + intermediate value differently. + +Push (K3_CONFIG "push" = 1): instead of writing ``out``, the combine stores this +rank's row of token t into slot [half][t][rank x copies + c] of every rank's latent exchange +(``latent_op.K3LatentExchange``, int32 [2][8][push_world][1792], bf16 pairs, -0.0 stored as +0.0) through its multicast +mapping, half = lat_flags[0] & 1 read after the grid wait; trtllm::k3_latent_reduce sums it. ``push_copies`` > 1 +only emulates a larger group's receive side. + +Configuration is per module instance (the shapes are trace-time constants): the loader injects K3_CONFIG = +{"i_tp": ..., "i_pad": ..., "num_local": ..., "num_ctas": ..., "m_max": 2, ["push", "push_world", "push_copies"]} +before executing the module. Launched with programmatic dependent launch: the barrier setup and the TMEM +allocation run before griddepcontrol.wait, which precedes every read of the producer's outputs and every global +write. +""" + +from __future__ import annotations + +import cuda.bindings.driver as cuda_driver +import cutlass +import cutlass.cute as cute +import cutlass.experimental.cuda as cuda +from cutlass._mlir.dialects import llvm as _llvm +from cutlass._mlir.extras import types as _T +from cutlass.cutlass_dsl import dsl_user_op +from cutlass.experimental import primitives as prims +from cutlass.experimental.cuda.tensor_map import TensorMapDataType + +_CFG = globals().get("K3_CONFIG") or {} + + +def _cfg(key: str, default): + """A kernel option from the op's configuration (K3_CONFIG), else its default.""" + return type(default)(_CFG.get(key, default)) + + +H = 3584 +I_TP = _cfg("i_tp", 192) # a rank's intermediate values (the logical shard) +I_PAD = _cfg( + "i_pad", (I_TP + 127) // 128 * 128 +) # the loader's padded intermediate (the buffers' layout) +TWO_I = 2 * I_TP +TOP_K = 16 +E_LOCAL = _cfg("num_local", 896) # this rank's experts +M_MAX = _cfg("m_max", 2) # tokens per call (one routing warp: 16 M lanes) +NCTA = _cfg("num_ctas", 148) # one CTA per SM +PUSH = bool(_cfg("push", 0)) # the combine pushes into the latent exchange instead of writing out +PUSH_WORLD = _cfg("push_world", 16) # slots per (half, token) of the exchange +PUSH_COPIES = _cfg("push_copies", 1) # slots this rank fills (rank x copies + c) +LAT_ROW_WORDS = 3584 // 2 # int32 words of a latent row in the exchange +THREADS = 256 +N = 16 # FC1 MMA N: B rows 0..1 (k-tile 2j) and 8..9 (k-tile 2j + 1) are the tokens +N2 = 8 # FC2 MMA N: row 2 j + t = (slot 4 i + j, token t) +MMA_M, MMA_TILE_K, MMA_INST_K = 128, 128, 32 +ROWS1 = 64 # FC1 rows per unit (one half of a 128-row tile) +U1 = TWO_I // ROWS1 # units per expert +K1_TILES = H // MMA_TILE_K # 28 +K1_PAIRS = K1_TILES // 2 # 14 stages per unit +NUM_KBLOCKS = MMA_TILE_K // MMA_INST_K # 4 +a_dtype = cutlass.Float4E2M1FN +b_dtype = cutlass.Float8E4M3FN +sf_dtype = cutlass.Float8E8M0FNU +a_smem_width = 8 # FP4 unpacked to 8-bit containers in shared memory +sf_vec_size = 32 +num_m0_per_sf_atom = 32 +num_m1_per_sf_atom = 4 +num_k_per_sf_atom = 4 +num_elts_atom_sf_fp16 = num_m0_per_sf_atom * num_m1_per_sf_atom * num_k_per_sf_atom // 2 +num_tmem_cols_per_sf_atom = 4 +NUM_BYTES_A = MMA_M * MMA_TILE_K * a_smem_width // 8 # 16384 +NUM_BYTES_A_HALF = ROWS1 * MMA_TILE_K * a_smem_width // 8 # 8192 +NUM_TX_A_HALF = ROWS1 * MMA_TILE_K * 4 // 8 # 4096: FP4 bytes in global memory +NUM_BYTES_B = N * MMA_TILE_K # 2048 +NUM_TX_B = 2 * M_MAX * MMA_TILE_K # two token boxes of M_MAX rows per stage +NUM_BYTES_SFA = 512 +NUM_BYTES_SFA_RAW = 2 * NUM_BYTES_SFA # the two k-tiles' scale atoms as loaded +NUM_BYTES_SFB = 512 +SFB_GROUP_BYTES = 16 +SFB_ROW_BYTES = H // sf_vec_size # 112: an activation row's scales +STAGES = 8 +PRE_FILL = min(STAGES, K1_PAIRS) # stages of the first unit issued before the routing barrier +SFA_COLS = num_tmem_cols_per_sf_atom +SFB_COLS = num_tmem_cols_per_sf_atom +NUM_SF_IDS = num_k_per_sf_atom * sf_vec_size // MMA_INST_K # 4 +SITU_GATE_CAP = 4.0 +SITU_LINEAR_CAP = 25.0 +E4M3_MAX = 448.0 +FP8_SENTINEL_I8 = -128 +_LOG2E = 1.4426950408889634 +G_CAP = M_MAX * TOP_K # distinct experts at most +# FC2 +R2 = 32 # output rows per CTA: one 32-row block of the shuffled layout +FC2_CTAS = H // R2 # 112 +KB2 = I_TP // 32 # MX blocks (K 32 MMAs) of a lean down row +KT2 = (I_TP + MMA_TILE_K - 1) // MMA_TILE_K # 128-wide k-tiles of a lean down row +KA2 = ( + I_PAD // 128 +) # w2 scale atoms per 128-row block (block_scale_interleave of the padded I / 32 columns) +GROUPS2 = G_CAP // 4 # 4 experts x 32 rows per 128-row MMA +NT2 = GROUPS2 * KT2 # A tiles, B tiles and scale atoms per CTA +NUM_TX_A2 = R2 * MMA_TILE_K * 4 // 8 # 2048: FP4 bytes of one expert's 32 rows of a k-tile +NUM_BYTES_B2 = N2 * MMA_TILE_K # 1024 +H_ROW = ( + (I_TP + I_TP // 32 + 15) // 16 * 16 +) # an intermediate row: fp8 values + E8M0 scales, 16-byte multiple +VCH = I_TP // 16 # 16-byte value chunks of an intermediate row +SCH = (KB2 + 15) // 16 # 16-byte scale chunks +GROUP_CHUNKS = 4 * M_MAX * (VCH + SCH) # 16-byte chunks of a group's intermediate rows +GROUP_ROUNDS = (GROUP_CHUNKS + 31) // 32 # load rounds of the one warp that loads a group +WAVE1_GROUPS = ( + NCTA // U1 +) // 4 # FC2 groups whose FC1 units all run in the first wave (CTA = unit) +CW = 32 # words per group counter (one 128-byte line each) +# TMEM columns: FC1 accumulator, FC1 SFA / SFB per stage, FC2 accumulators (8 per group), FC2 SFA / SFB per tile. +ACC2_COL = N + 2 * STAGES * SFA_COLS +SFA2_COL = ACC2_COL + GROUPS2 * N2 +SFB2_COL = SFA2_COL + NT2 * SFA_COLS +TMEM_COLS = 32 +while TMEM_COLS < SFB2_COL + NT2 * SFB_COLS: + TMEM_COLS *= 2 +REF_SLICES = ( + 5 # k3_moe's FC2 slices (fc2_slices): its combine's sum tree, kept so the bits can match +) +EPI_BAR_ID = 1 +EPI_THREADS = 128 +W_WARP = 4 +X_WARP = 5 +S_WARP = 6 +MMA_WARP = 7 +EVICT_FIRST = 0x12F0000000000000 +assert M_MAX == 2, M_MAX +assert ( + TWO_I % ROWS1 == 0 and H % R2 == 0 and I_TP % 32 == 0 and I_PAD % 128 == 0 and I_PAD >= I_TP +), (I_TP, I_PAD) +assert TMEM_COLS <= 512 and I_TP % 16 == 0, (TMEM_COLS, I_TP) +assert NT2 > STAGES # FC2's A tiles cycle through the ring +assert NCTA >= FC2_CTAS, NCTA + + +@dsl_user_op +def _mul_rn(a, b, *, loc=None, ip=None): + """a * b rounded on its own (mul.rn is never fused into an FMA).""" + return cutlass.Float32(_llvm.inline_asm( + _T.f32(), [cutlass.Float32(a).ir_value(loc=loc, ip=ip), cutlass.Float32(b).ir_value(loc=loc, ip=ip)], + "mul.rn.f32 $0, $1, $2;", "=f,f,f", has_side_effects=False, is_align_stack=False, + asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + )) # fmt: skip + + +@dsl_user_op +def _add_rn(a, b, *, loc=None, ip=None): + """a + b rounded on its own (add.rn is never fused into an FMA).""" + return cutlass.Float32(_llvm.inline_asm( + _T.f32(), [cutlass.Float32(a).ir_value(loc=loc, ip=ip), cutlass.Float32(b).ir_value(loc=loc, ip=ip)], + "add.rn.f32 $0, $1, $2;", "=f,f,f", has_side_effects=False, is_align_stack=False, + asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + )) # fmt: skip + + +@dsl_user_op +def _red_release_add(addr, val, *, loc=None, ip=None): + _llvm.inline_asm( + None, [cutlass.Int64(addr).ir_value(loc=loc, ip=ip), cutlass.Int32(val).ir_value(loc=loc, ip=ip)], + "red.release.gpu.global.add.u32 [$0], $1;", "l,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _load_acquire(addr, *, loc=None, ip=None): + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [cutlass.Int64(addr).ir_value(loc=loc, ip=ip)], + "ld.acquire.gpu.global.u32 $0, [$1];", "=r,l", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _st_u32(addr, val, *, loc=None, ip=None): + _llvm.inline_asm( + None, [cutlass.Int64(addr).ir_value(loc=loc, ip=ip), cutlass.Int32(val).ir_value(loc=loc, ip=ip)], + "st.global.u32 [$0], $1;", "l,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +@dsl_user_op +def _pack_bf16x2(hi, lo, *, loc=None, ip=None): + """(bf16(hi) << 16) | bf16(lo), round to nearest even.""" + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [hi.ir_value(loc=loc, ip=ip), lo.ir_value(loc=loc, ip=ip)], + "cvt.rn.bf16x2.f32 $0, $1, $2;", "=r,f,f", has_side_effects=False, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@cute.jit +def _emit(acc, lane, row2, tq, ch, out, lat_mc, lat_flags, lat_rank): + """Warp tq of an FC2 CTA (token tq's combine): lane l holds the value of output column row2 + 4 (l % 8) + l // 8, + stored to ``out`` at ``ch``; with PUSH, lanes 0, 2, 4, 6 gather columns row2 + 4 l .. row2 + 4 l + 7 and store + them as 16 bytes (bf16 pairs, -0.0 as +0.0) into this rank's slots of token tq in every rank's latent exchange + through its multicast mapping.""" + if cutlass.const_expr(PUSH): + bits = _pack_bf16x2(cutlass.Float32(0.0), acc) & cutlass.Int32(0xFFFF) + bits = cutlass.select_(bits == cutlass.Int32(0x8000), cutlass.Int32(0), bits) + b1 = cute.arch.shuffle_sync_down(bits, 8) + b2 = cute.arch.shuffle_sync_down(bits, 16) + b3 = cute.arch.shuffle_sync_down(bits, 24) + b4 = cute.arch.shuffle_sync_down(bits, 1) + b5 = cute.arch.shuffle_sync_down(bits, 9) + b6 = cute.arch.shuffle_sync_down(bits, 17) + b7 = cute.arch.shuffle_sync_down(bits, 25) + half = cutlass.Int32(lat_flags.load(idx=0, is_volatile=True)) & cutlass.Int32(1) + if (lane < 8) & (lane % 2 == 0): + vec = (bits | (b1 << cutlass.Int32(16)), b2 | (b3 << cutlass.Int32(16)), b4 | (b5 << cutlass.Int32(16)), + b6 | (b7 << cutlass.Int32(16))) # fmt: skip + for c in cutlass.range_constexpr(PUSH_COPIES): + slot = lat_rank * cutlass.Int32(PUSH_COPIES) + cutlass.Int32(c) + lat_mc.store(vec, idx=((half * cutlass.Int32(8) + tq) * cutlass.Int32(PUSH_WORLD) + slot) + * cutlass.Int32(LAT_ROW_WORDS) + (row2 + cutlass.Int32(4) * lane) // cutlass.Int32(2), + alignment=16) # fmt: skip + else: + out.store(cutlass.BFloat16(acc), idx=ch) + + +def _tanh_f32(x): + e = cute.math.exp2(cute.math.abs(x) * cutlass.Float32(-2.0 * _LOG2E), fastmath=True) + t = (cutlass.Float32(1.0) - e) * cute.arch.rcp_approx(cutlass.Float32(1.0) + e) + return cutlass.select_(x < cutlass.Float32(0.0), -t, t) + + +def _sigmoid_f32(x): + return cute.arch.rcp_approx( + cutlass.Float32(1.0) + cute.math.exp2(x * cutlass.Float32(-_LOG2E), fastmath=True) + ) + + +def _situ(gate, up): + g = ( + cutlass.Float32(SITU_GATE_CAP) + * _tanh_f32(gate * cutlass.Float32(1.0 / SITU_GATE_CAP)) + * _sigmoid_f32(gate) + ) + u = cutlass.Float32(SITU_LINEAR_CAP) * _tanh_f32(up * cutlass.Float32(1.0 / SITU_LINEAR_CAP)) + return g * u + + +def _block_e8m0(amax): + """E8M0 byte of an MX block and 2^(127 - byte) as f32 (k3_moe's ceil recipe).""" + sf = amax * cutlass.Float32(1.0 / 448.0) + sbits = cutlass.Int32(sf.bitcast(cutlass.Int32)) + sexp = (sbits >> cutlass.Int32(23)) & cutlass.Int32(0xFF) + mant = sbits & cutlass.Int32(0x7FFFFF) + byte = sexp + cutlass.select_(mant != cutlass.Int32(0), cutlass.Int32(1), cutlass.Int32(0)) + byte = cutlass.select_(byte > cutlass.Int32(0xFE), cutlass.Int32(0xFE), byte) + byte = cutlass.select_(amax > cutlass.Float32(0.0), byte, cutlass.Int32(0)) + inv = cutlass.Int32((cutlass.Int32(254) - byte) << cutlass.Int32(23)).bitcast(cutlass.Float32) + return byte, inv + + +@cute.jit +def _wait(bar: cutlass.Array, parity): + while not cute.arch.mbarrier_try_wait(bar.data_ptr(), parity): + pass + + +@cute.jit +def _unit_of(ui, bx, pos2): + """This CTA's ui-th FC1 unit: its own index in the first wave, then position pos2 of every later wave (the CTAs + without FC2 work first).""" + return cutlass.select_(ui == cutlass.Int32(0), bx, ui * cutlass.Int32(NCTA) + pos2) + + +@cute.jit +def _load_group( + gi, + lane, + g_n, + cnt_base, + hbuf32: cutlass.Array, + b2_32: cutlass.Array, + sfb2_32: cutlass.Array, + b2_ready: cutlass.Array, +): + """One warp, FC2 group gi: once the group's counter reaches its experts' units (every lane polls it), the group's + intermediate rows into its B tiles (row 2 j + t of tile (gi, t2) = slot 4 gi + j, token t, K 128 t2 .., + 128B-swizzled) and their scales into the B scale atoms (byte 16 (2 j + t) + k = block 4 t2 + k), all the warp's + loads in flight at once; then the warp's 32 arrivals on b2_ready[gi].""" + n_here = g_n - gi * cutlass.Int32(4) + n_here = cutlass.select_(n_here > cutlass.Int32(4), cutlass.Int32(4), n_here) + while _load_acquire(cnt_base + cutlass.Int64(gi * (CW * 4))) < n_here * cutlass.Int32(U1): + pass + got = [] + for r in cutlass.range_constexpr(GROUP_ROUNDS): + rem = lane + cutlass.Int32(r * 32) + jr = rem // cutlass.Int32(VCH + SCH) # (slot in group, token) row 2 j + t + c = rem % cutlass.Int32(VCH + SCH) + sl = gi * cutlass.Int32(4) + jr // cutlass.Int32(M_MAX) + ok = (rem < cutlass.Int32(GROUP_CHUNKS)) & (sl < g_n) + hrow = sl * cutlass.Int32(M_MAX) + jr % cutlass.Int32(M_MAX) + v4 = prims.load_ext( + hbuf32.subview(cutlass.select_( + ok, (hrow * cutlass.Int32(H_ROW) + c * cutlass.Int32(16)) // cutlass.Int32(4), cutlass.Int32(0) + )), + dtype=cutlass.Int32, count=4, order="relaxed", scope="gpu", + ) # fmt: skip + got.append(v4) + for r in cutlass.range_constexpr(GROUP_ROUNDS): + rem = lane + cutlass.Int32(r * 32) + jr = rem // cutlass.Int32(VCH + SCH) + c = rem % cutlass.Int32(VCH + SCH) + sl = gi * cutlass.Int32(4) + jr // cutlass.Int32(M_MAX) + v4 = got[r] + if (rem < cutlass.Int32(GROUP_CHUNKS)) & (sl < g_n): + if c < cutlass.Int32(VCH): + t2 = c // cutlass.Int32(8) + cc = c % cutlass.Int32(8) + jt = gi * cutlass.Int32(KT2) + t2 + b2_32.store((cutlass.Int32(v4[0]), cutlass.Int32(v4[1]), cutlass.Int32(v4[2]), cutlass.Int32(v4[3])), + idx=(jt * cutlass.Int32(NUM_BYTES_B2) + jr * cutlass.Int32(MMA_TILE_K) + + (cc ^ jr) * cutlass.Int32(16)) // cutlass.Int32(4), alignment=16) # fmt: skip + else: + for wq in cutlass.range_constexpr(4): + t2 = (c - cutlass.Int32(VCH)) * cutlass.Int32(4) + cutlass.Int32(wq) + if t2 < cutlass.Int32(KT2): + jt = gi * cutlass.Int32(KT2) + t2 + sfb2_32.store( + cutlass.Int32(v4[wq]), + idx=(jt * cutlass.Int32(NUM_BYTES_SFB) + jr * cutlass.Int32(16)) + // cutlass.Int32(4), + ) + prims.fence_proxy("async_shared", space=prims.SharedSpace.shared_cta) + prims.mbarrier_arrive(b2_ready.subview(gi)) + + +@cute.kernel +def k3_moe_m2_kernel( + tma_a1_desc: cutlass.GridConstant[cuda.TensorMap], + tma_b1_desc: cutlass.GridConstant[cuda.TensorMap], + tma_sfa1_desc: cutlass.GridConstant[cuda.TensorMap], + tma_a2_desc: cutlass.GridConstant[cuda.TensorMap], + sfb1_ptr: cutlass.Int64, + ids: cutlass.Array, # int32 [M * 16] + wts: cutlass.Array, # bf16 [M * 16] routing weights (as int16 bits) + w2s32: cutlass.Array, # int32 words of w2_weight_scale [E, H / 128, I_PAD / 128, 512 B] + hbuf: cutlass.Array, # int8 [G_CAP, M_MAX, H_ROW] + hbuf32: cutlass.Array, # the same memory as int32 words + counts: cutlass.Array, # int32 [2, GROUPS2, CW] + epochs: cutlass.Array, # int32 [NCTA] + out: cutlass.Array, # bf16 [M, 3584] (not written with PUSH) + offset: cutlass.Int32, + m: cutlass.Int32, + lat_mc: cutlass.Array, # PUSH: int32 words of the latent exchange's multicast mapping + lat_flags: cutlass.Array, # PUSH: int32 [4], [0] the reduce's call count + lat_rank: cutlass.Int32, +): + tidx, _, _ = cute.arch.thread_idx() + bx, _, _ = cute.arch.block_idx() + warp = cute.arch.make_warp_uniform(cute.arch.warp_idx()) + lane = tidx % 32 + + sA = cutlass.Array( + cutlass.Int8, NUM_BYTES_A * STAGES, space=cutlass.AddressSpace.smem, alignment=1024 + ) + sB = cutlass.Array( + cutlass.Int8, NUM_BYTES_B * STAGES, space=cutlass.AddressSpace.smem, alignment=1024 + ) + b2 = cutlass.Array( + cutlass.Int8, NT2 * NUM_BYTES_B2, space=cutlass.AddressSpace.smem, alignment=1024 + ) + sSFA = cutlass.Array( + cutlass.Int8, NUM_BYTES_SFA * STAGES, space=cutlass.AddressSpace.smem, alignment=1024 + ) + sSFAraw = cutlass.Array(cutlass.Int32, NUM_BYTES_SFA_RAW * STAGES // 4, space=cutlass.AddressSpace.smem, + alignment=1024) # fmt: skip + sSFB = cutlass.Array( + cutlass.Int8, NUM_BYTES_SFB * STAGES, space=cutlass.AddressSpace.smem, alignment=1024 + ) + sfa2 = cutlass.Array( + cutlass.Int8, NT2 * NUM_BYTES_SFA, space=cutlass.AddressSpace.smem, alignment=1024 + ) + sfb2 = cutlass.Array( + cutlass.Int8, NT2 * NUM_BYTES_SFB, space=cutlass.AddressSpace.smem, alignment=1024 + ) + s_part = cutlass.Array( + cutlass.Float32, M_MAX * ROWS1, space=cutlass.AddressSpace.smem, alignment=16 + ) + ys = cutlass.Array( + cutlass.Float32, G_CAP * M_MAX * R2, space=cutlass.AddressSpace.smem, alignment=16 + ) + ab_full = cutlass.Array(cutlass.Int64, STAGES, space=cutlass.AddressSpace.smem, alignment=8) + ab_empty = cutlass.Array(cutlass.Int64, STAGES, space=cutlass.AddressSpace.smem, alignment=8) + fc2_full = cutlass.Array(cutlass.Int64, STAGES, space=cutlass.AddressSpace.smem, alignment=8) + scales_in_tmem = cutlass.Array( + cutlass.Int64, STAGES, space=cutlass.AddressSpace.smem, alignment=8 + ) + acc_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + acc_empty = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + tmem_ready = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + sfa2_ready = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + b2_ready = cutlass.Array(cutlass.Int64, GROUPS2, space=cutlass.AddressSpace.smem, alignment=8) + acc2_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + acc2_grp = cutlass.Array(cutlass.Int64, GROUPS2, space=cutlass.AddressSpace.smem, alignment=8) + localmax_smem = cutlass.Array( + cutlass.Float32, 2 * M_MAX, space=cutlass.AddressSpace.smem, alignment=16 + ) + tmem_ptr_i32 = cutlass.Array(cutlass.Int32, 1, space=cutlass.AddressSpace.smem) + s_el = cutlass.Array(cutlass.Int32, G_CAP + 1, space=cutlass.AddressSpace.smem, alignment=16) + s_w = cutlass.Array( + cutlass.Float32, G_CAP * M_MAX, space=cutlass.AddressSpace.smem, alignment=16 + ) + + # ---- before the grid dependency: barriers, the TMEM allocation, the routing scratch zeroed + if warp == 0: + if tidx < STAGES: + prims.mbarrier_init(ab_full.subview(tidx), 3) # A + SFA, B, SFB + prims.mbarrier_init(ab_empty.subview(tidx), 1) + prims.mbarrier_init(fc2_full.subview(tidx), 1) + prims.mbarrier_init(scales_in_tmem.subview(tidx), 1) + if tidx < GROUPS2: + # One warp loads each group: the first wave's groups the epilogue warps, the rest warps 5-6. + prims.mbarrier_init(b2_ready.subview(tidx), 32) + prims.mbarrier_init(acc2_grp.subview(tidx), 1) + if tidx == 0: + prims.mbarrier_init(acc_full.subview(0), 1) + prims.mbarrier_init(acc_empty.subview(0), EPI_THREADS) + prims.mbarrier_init(tmem_ready.subview(0), 32) + prims.mbarrier_init(sfa2_ready.subview(0), EPI_THREADS) + prims.mbarrier_init(acc2_full.subview(0), 1) + if tidx < G_CAP * M_MAX: + s_w.store(cutlass.Float32(0.0), idx=tidx) + prims.fence_mbarrier_init() + prims.barrier_cta_sync(0) + if warp == MMA_WARP: + prims.tcgen05_alloc(tmem_ptr_i32, TMEM_COLS) + prims.mbarrier_arrive(tmem_ready) + prims.tcgen05_relinquish_alloc_permit() + if warp == W_WARP: + prims.prefetch_tensormap(tma_a1_desc.get_ptr()) + prims.prefetch_tensormap(tma_sfa1_desc.get_ptr()) + prims.prefetch_tensormap(tma_a2_desc.get_ptr()) + if warp == X_WARP: + prims.prefetch_tensormap(tma_b1_desc.get_ptr()) + + cute.arch.griddepcontrol_wait() + cute.arch.griddepcontrol_launch_dependents() + pos2 = cutlass.select_(bx >= cutlass.Int32(FC2_CTAS), bx - cutlass.Int32(FC2_CTAS), + bx + cutlass.Int32(NCTA - FC2_CTAS)) # fmt: skip + + # ---- routing: lane l = (token l / 16, top-k slot l % 16); the distinct local experts take slots in ascending id + # (each lane ranks its expert among the warp's first occurrences), each slot holding both tokens' weights. + if warp == W_WARP: + tk = lane // cutlass.Int32(TOP_K) + live = tk < m + loc_e = cutlass.Int32(-1) + wv = cutlass.Float32(0.0) + if live: + loc_e = ids.load(idx=lane) - offset + wv = cutlass.Float32( + cutlass.Int32(cutlass.Int32(wts.load(idx=lane)) << cutlass.Int32(16)).bitcast( + cutlass.Float32 + ) + ) + is_local = live & (loc_e >= cutlass.Int32(0)) & (loc_e < cutlass.Int32(E_LOCAL)) + key = cutlass.select_(is_local, loc_e, cutlass.Int32(1 << 30)) + keys = [] + dup = cutlass.Boolean(False) + for k in cutlass.range_constexpr(32): + kk = cute.arch.shuffle_sync(key, k) + keys.append(kk) + dup = dup | ((kk == key) & (cutlass.Int32(k) < lane)) + is_first = cutlass.select_(dup, cutlass.Boolean(False), is_local) + fbal = prims.vote_sync(0xFFFFFFFF, is_first, prims.VoteSync.BALLOT) + slot = cutlass.Int32(0) + for k in cutlass.range_constexpr(32): + slot = slot + cutlass.select_( + (((fbal >> cutlass.Int32(k)) & cutlass.Int32(1)) != cutlass.Int32(0)) + & (keys[k] < key), + cutlass.Int32(1), + cutlass.Int32(0), + ) + g_tot = cute.arch.popc(fbal) + if is_first: + s_el.store(loc_e, idx=slot) + if is_local: + s_w.store(wv, idx=slot * cutlass.Int32(M_MAX) + tk) + if lane == 0: + s_el.store(g_tot, idx=G_CAP) + cute.arch.sync_warp() + # The first unit's ring fill before the CTA barrier. + if bx < g_tot * cutlass.Int32(U1): + grp0 = bx // cutlass.Int32(U1) + m_src = prims.vote_sync(0xFFFFFFFF, is_first & (slot == grp0), prims.VoteSync.BALLOT) + e0 = cute.arch.shuffle_sync( + loc_e, cute.arch.popc((m_src & (cutlass.Int32(0) - m_src)) - cutlass.Int32(1)) + ) + t128_0 = (bx % cutlass.Int32(U1)) // cutlass.Int32(2) + coord_m0 = t128_0 * cutlass.Int32(MMA_M) + (bx % cutlass.Int32(2)) * cutlass.Int32( + ROWS1 + ) + if prims.elect_sync(): + for kp in cutlass.range_constexpr(PRE_FILL): + prims.mbarrier_arrive_expect_tx( + ab_full.subview(kp), 2 * (NUM_TX_A_HALF + NUM_BYTES_SFA) + ) + for q in cutlass.range_constexpr(2): + prims.cp_async_bulk_tensor_shared_cta_global( + sA.subview(cutlass.Int32(kp * NUM_BYTES_A + q * NUM_BYTES_A_HALF)), tma_a1_desc.get_ptr(), + (cutlass.Int32((2 * kp + q) * MMA_TILE_K), coord_m0, e0), ab_full.subview(kp), + l2_cache_hint=EVICT_FIRST, + ) # fmt: skip + prims.cp_async_bulk_tensor_shared_cta_global( + sSFAraw.subview(cutlass.Int32((kp * NUM_BYTES_SFA_RAW + q * NUM_BYTES_SFA) // 4)), + tma_sfa1_desc.get_ptr(), (cutlass.Int32(0), cutlass.Int32(2 * kp + q), t128_0, e0), + ab_full.subview(kp), l2_cache_hint=EVICT_FIRST, + ) # fmt: skip + prims.barrier_cta_sync(0) + g_n = s_el.load(idx=G_CAP) + units = g_n * cutlass.Int32(U1) + mine = cutlass.select_(bx < units, cutlass.Int32(1), cutlass.Int32(0)) + cutlass.select_( + units > cutlass.Int32(NCTA) + pos2, + (units - cutlass.Int32(NCTA) - pos2 + cutlass.Int32(NCTA - 1)) // cutlass.Int32(NCTA), + cutlass.Int32(0), + ) + ep = epochs.load(idx=bx) + cnt_base = counts.subview(0).data_ptr().toint() + cutlass.Int64((ep & 1) * (GROUPS2 * CW * 4)) + row2 = bx * cutlass.Int32(R2) + is_fc2 = bx < cutlass.Int32(FC2_CTAS) + groups = (g_n + cutlass.Int32(3)) // cutlass.Int32(4) + nt2 = cutlass.select_(is_fc2, groups * cutlass.Int32(KT2), cutlass.Int32(0)) + blk2 = row2 // cutlass.Int32(128) + m1 = (row2 % cutlass.Int32(128)) // cutlass.Int32(32) + + # ---- weights producer (4): FC1 stages (the unit's 64 rows at k-tiles 2j, 2j + 1 and both scale atoms, raw), then + # this CTA's FC2 A tiles into the ring as it drains (group i, k-tile t: four experts' 32 rows at 4 KB offsets). + if warp == W_WARP: + if bx == 0: + if lane < cutlass.Int32(GROUPS2): + _st_u32(counts.subview(0).data_ptr().toint() + + cutlass.Int64((((ep + 1) & 1) * GROUPS2 * CW + lane * CW) * 4), cutlass.Int32(0)) # fmt: skip + g = cutlass.select_( + mine > cutlass.Int32(0), cutlass.Int32(PRE_FILL), cutlass.Int32(0) + ) # issued pre-barrier + ab_empty_phase = 1 + for ui in range(mine): + u = _unit_of(ui, bx, pos2) + grp = u // cutlass.Int32(U1) + t128 = (u % cutlass.Int32(U1)) // cutlass.Int32(2) + half = u % cutlass.Int32(2) + coord_expert = s_el.load(idx=grp) + coord_m = t128 * cutlass.Int32(MMA_M) + half * cutlass.Int32(ROWS1) + for kp in cutlass.range(cutlass.select_(ui == cutlass.Int32(0), cutlass.Int32(PRE_FILL), cutlass.Int32(0)), + K1_PAIRS, unroll=1): # fmt: skip + stage = g % cutlass.Int32(STAGES) + if stage == cutlass.Int32(0) and g != cutlass.Int32(0): + ab_empty_phase = ab_empty_phase ^ 1 + while not cute.arch.mbarrier_try_wait( + ab_empty.subview(stage).data_ptr(), ab_empty_phase + ): + pass + coord_k = kp * cutlass.Int32(2 * MMA_TILE_K) + if prims.elect_sync(): + prims.mbarrier_arrive_expect_tx( + ab_full.subview(stage), 2 * (NUM_TX_A_HALF + NUM_BYTES_SFA) + ) + for q in cutlass.range_constexpr(2): + prims.cp_async_bulk_tensor_shared_cta_global( + sA.subview(stage * cutlass.Int32(NUM_BYTES_A) + cutlass.Int32(q * NUM_BYTES_A_HALF)), + tma_a1_desc.get_ptr(), (coord_k + cutlass.Int32(q * MMA_TILE_K), coord_m, coord_expert), + ab_full.subview(stage), l2_cache_hint=EVICT_FIRST, + ) # fmt: skip + prims.cp_async_bulk_tensor_shared_cta_global( + sSFAraw.subview( + (stage * cutlass.Int32(NUM_BYTES_SFA_RAW) + cutlass.Int32(q * NUM_BYTES_SFA)) // 4 + ), + tma_sfa1_desc.get_ptr(), (cutlass.Int32(0), kp * cutlass.Int32(2) + cutlass.Int32(q), t128, + coord_expert), + ab_full.subview(stage), l2_cache_hint=EVICT_FIRST, + ) # fmt: skip + g = g + cutlass.Int32(1) + for j in range(nt2): + stage = g % cutlass.Int32(STAGES) + if stage == cutlass.Int32(0) and g != cutlass.Int32(0): + ab_empty_phase = ab_empty_phase ^ 1 + while not cute.arch.mbarrier_try_wait( + ab_empty.subview(stage).data_ptr(), ab_empty_phase + ): + pass + i2 = j // cutlass.Int32(KT2) + t2 = j % cutlass.Int32(KT2) + n_here = g_n - i2 * cutlass.Int32(4) + if n_here > cutlass.Int32(4): + n_here = cutlass.Int32(4) + if prims.elect_sync(): + prims.mbarrier_arrive_expect_tx( + fc2_full.subview(stage), n_here * cutlass.Int32(NUM_TX_A2) + ) + for jj in cutlass.range_constexpr(4): + if cutlass.Int32(jj) < n_here: + prims.cp_async_bulk_tensor_shared_cta_global( + sA.subview(stage * cutlass.Int32(NUM_BYTES_A) + cutlass.Int32(jj * R2 * MMA_TILE_K)), + tma_a2_desc.get_ptr(), (t2 * cutlass.Int32(MMA_TILE_K), row2, + s_el.load(idx=i2 * cutlass.Int32(4) + cutlass.Int32(jj))), + fc2_full.subview(stage), l2_cache_hint=EVICT_FIRST, + ) # fmt: skip + g = g + cutlass.Int32(1) + + # ---- activations producer (5): the tokens' 128 K of k-tiles 2j, 2j + 1 into B rows 0..1 and 8..9, their scales. + if warp == X_WARP: + g = cutlass.Int32(0) + ab_empty_phase = 1 + for ui in range(mine): + for kp in cutlass.range(K1_PAIRS, unroll=1): + stage = g % cutlass.Int32(STAGES) + if stage == cutlass.Int32(0) and g != cutlass.Int32(0): + ab_empty_phase = ab_empty_phase ^ 1 + while not cute.arch.mbarrier_try_wait( + ab_empty.subview(stage).data_ptr(), ab_empty_phase + ): + pass + coord_k = kp * cutlass.Int32(2 * MMA_TILE_K) + if prims.elect_sync(): + prims.mbarrier_arrive_expect_tx(ab_full.subview(stage), NUM_TX_B) + for q in cutlass.range_constexpr(2): + prims.cp_async_bulk_tensor_shared_cta_global( + sB.subview(stage * cutlass.Int32(NUM_BYTES_B) + cutlass.Int32(q * 8 * MMA_TILE_K)), + tma_b1_desc.get_ptr(), (coord_k + cutlass.Int32(q * MMA_TILE_K), cutlass.Int32(0)), + ab_full.subview(stage), + ) # fmt: skip + if lane == 0: + for q in cutlass.range_constexpr(2): + for n in cutlass.range_constexpr(M_MAX): + if cutlass.Int32(n) < m: + sfb_gmem = ( + sfb1_ptr + + cutlass.Int64(n * SFB_ROW_BYTES) + + cutlass.Int64( + (kp * cutlass.Int32(2) + cutlass.Int32(q)) * NUM_KBLOCKS + ) + ) + prims.cp_async_shared_global( + sSFB.subview(stage * cutlass.Int32(NUM_BYTES_SFB) + + cutlass.Int32((q * 8 + n) * SFB_GROUP_BYTES)).data_ptr(), + cutlass.inttoptr(sfb_gmem, mem_space=1, dtype=sf_dtype), size=4, modifier="ca", + cp_size=4, + ) # fmt: skip + prims.cp_async_mbarrier_arrive(ab_full.subview(stage), noinc=True) + g = g + cutlass.Int32(1) + + # ---- scales to TMEM (6): the unit's half of both k-tiles' atoms spliced into one MMA atom (bytes 8 h .. 8 h + 7 + # of each 16-byte row group: k-tile 2j to MMA rows 0-63, k-tile 2j + 1 to rows 64-127), then SFA and SFB to TMEM. + if warp == S_WARP: + _wait(tmem_ready, 0) + tmem_raw_addr = tmem_ptr_i32.load() + base_col_id = tmem_raw_addr & 0xFFFF + base_row_id = tmem_raw_addr >> 16 + sfa_col_id0 = base_col_id + N + sfb_col_id0 = sfa_col_id0 + STAGES * SFA_COLS + s2t_shape, s2t_multicast = prims.S2TCopyMode.S2T_32x128b_WARPX4 + sSFA32 = cutlass.Array(sSFA.data_ptr(0), shape=(NUM_BYTES_SFA * STAGES // 4,), dtype=cutlass.Int32, + alignment=16) # fmt: skip + g = cutlass.Int32(0) + full_phase = 0 + for ui in range(mine): + u = _unit_of(ui, bx, pos2) + half = u % cutlass.Int32(2) + for kp in cutlass.range(K1_PAIRS, unroll=1): + stage = g % cutlass.Int32(STAGES) + if stage == cutlass.Int32(0) and g != cutlass.Int32(0): + full_phase = full_phase ^ 1 + while not cute.arch.mbarrier_try_wait( + ab_full.subview(stage).data_ptr(), full_phase + ): + pass + raw = ( + stage * cutlass.Int32(NUM_BYTES_SFA_RAW // 4) + + lane * cutlass.Int32(4) + + half * cutlass.Int32(2) + ) + lo = sSFAraw.load(idx=raw, vector_size=2, alignment=8) + hi = sSFAraw.load( + idx=raw + cutlass.Int32(NUM_BYTES_SFA // 4), vector_size=2, alignment=8 + ) + sSFA32.store((cutlass.Int32(lo[0]), cutlass.Int32(lo[1]), cutlass.Int32(hi[0]), cutlass.Int32(hi[1])), + idx=stage * cutlass.Int32(NUM_BYTES_SFA // 4) + lane * cutlass.Int32(4), + alignment=16) # fmt: skip + prims.fence_proxy("async_shared", space=prims.SharedSpace.shared_cta) + cute.arch.sync_warp() + prims.tcgen05_fence( + prims.Tcgen05Fence.AFTER_THREAD_SYNC + ) # the lanes' spliced scale atoms + sfa_tmem_ptr = cutlass.inttoptr( + (base_row_id << 16) | (sfa_col_id0 + stage * SFA_COLS), 6, cutlass.Int32 + ) + sfb_tmem_ptr = cutlass.inttoptr( + (base_row_id << 16) | (sfb_col_id0 + stage * SFB_COLS), 6, cutlass.Int32 + ) + desc_a = prims.Tcgen05SmemDesc.build( + sSFA.subview(stage * NUM_BYTES_SFA), leading_byte_offset=16, stride_byte_offset=128, + base_offset=0, layout=0, + ) # fmt: skip + desc_b = prims.Tcgen05SmemDesc.build( + sSFB.subview(stage * NUM_BYTES_SFB), leading_byte_offset=16, stride_byte_offset=128, + base_offset=0, layout=0, + ) # fmt: skip + if prims.elect_sync(): + prims.tcgen05_cp(s2t_shape, sfa_tmem_ptr, desc_a, multicast=s2t_multicast) + prims.tcgen05_cp(s2t_shape, sfb_tmem_ptr, desc_b, multicast=s2t_multicast) + prims.tcgen05_commit(scales_in_tmem.subview(stage)) + g = g + cutlass.Int32(1) + + # ---- MMA (7): FC1 per unit; then FC2: the spliced weight scales to TMEM, then per group (once its intermediate + # is in) the B scales and per A tile the group's K 32 MMAs into its 8 accumulator columns. + if warp == MMA_WARP: + tmem_raw_addr = tmem_ptr_i32.load() + acc_tmem_ptr = cutlass.inttoptr(tmem_raw_addr, 6, cutlass.Float32) + idesc = prims.Tcgen05MxInstrDesc.build( + a_dtype=a_dtype, b_dtype=b_dtype, scale_format=1, n_dim=N, m_dim=MMA_M + ) + idesc2 = prims.Tcgen05MxInstrDesc.build( + a_dtype=a_dtype, b_dtype=b_dtype, scale_format=1, n_dim=N2, m_dim=MMA_M + ) + base_col_id = tmem_raw_addr & 0xFFFF + base_row_id = tmem_raw_addr >> 16 + sfa_col_id0 = base_col_id + N + sfb_col_id0 = sfa_col_id0 + STAGES * SFA_COLS + g = cutlass.Int32(0) + st_phase = 0 + acc_empty_phase = 1 + for ui in range(mine): + while not cute.arch.mbarrier_try_wait(acc_empty.data_ptr(), acc_empty_phase): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc_empty_phase = acc_empty_phase ^ 1 + scale_d = False + for kp in cutlass.range(K1_PAIRS, unroll=1): + stage = g % cutlass.Int32(STAGES) + if stage == cutlass.Int32(0) and g != cutlass.Int32(0): + st_phase = st_phase ^ 1 + while not cute.arch.mbarrier_try_wait( + scales_in_tmem.subview(stage).data_ptr(), st_phase + ): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + sfa_base = (base_row_id << 16) | (sfa_col_id0 + stage * SFA_COLS) + sfb_base = (base_row_id << 16) | (sfb_col_id0 + stage * SFB_COLS) + desc_a_base = prims.Tcgen05SmemDesc.build( + sA.subview(stage * NUM_BYTES_A), leading_byte_offset=16, stride_byte_offset=1024, base_offset=0, + layout=2, + ) # fmt: skip + desc_b_base = prims.Tcgen05SmemDesc.build( + sB.subview(stage * NUM_BYTES_B), leading_byte_offset=16, stride_byte_offset=1024, base_offset=0, + layout=2, + ) # fmt: skip + for kb in cutlass.range(NUM_KBLOCKS, unroll_full=True): + sf_inside = kb % NUM_SF_IDS + sf_col = kb // NUM_SF_IDS + sfa_tmem_ptr = cutlass.inttoptr(sfa_base + sf_col * SFA_COLS, 6, cutlass.Int32) + sfb_tmem_ptr = cutlass.inttoptr(sfb_base + sf_col * SFB_COLS, 6, cutlass.Int32) + idesc_u = idesc.set_sf_ids(a_sf_id=sf_inside, b_sf_id=sf_inside) + inc = ((MMA_INST_K * a_smem_width // 8) >> 4) * kb + if prims.elect_sync(): + prims.tcgen05_mma_block_scale( + prims.MMABlockScaleKind.MXF8F6F4, prims.CTAGroup.CTA_1, acc_tmem_ptr, + desc_a_base + inc, desc_b_base + inc, idesc_u, scale_d, sfa_tmem_ptr, sfb_tmem_ptr, + ) # fmt: skip + scale_d = True + if prims.elect_sync(): + prims.tcgen05_commit(ab_empty.subview(stage)) + g = g + cutlass.Int32(1) + if prims.elect_sync(): + prims.tcgen05_commit(acc_full) + if nt2 > cutlass.Int32(0): + s2t_shape, s2t_multicast = prims.S2TCopyMode.S2T_32x128b_WARPX4 + _wait(sfa2_ready, 0) + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + if prims.elect_sync(): + for j in range(nt2): + prims.tcgen05_cp( + s2t_shape, + cutlass.inttoptr( + (base_row_id << 16) | (base_col_id + SFA2_COL + j * SFA_COLS), 6, cutlass.Int32 + ), + prims.Tcgen05SmemDesc.build(sfa2.subview(j * NUM_BYTES_SFA), leading_byte_offset=16, + stride_byte_offset=128, base_offset=0, layout=0), + multicast=s2t_multicast, + ) # fmt: skip + j = cutlass.Int32(0) + for i2 in range(groups): + _wait(b2_ready.subview(i2), 0) + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + if prims.elect_sync(): + for t2c in cutlass.range_constexpr(KT2): + jt = i2 * cutlass.Int32(KT2) + cutlass.Int32(t2c) + prims.tcgen05_cp( + s2t_shape, + cutlass.inttoptr((base_row_id << 16) | (base_col_id + SFB2_COL + jt * SFB_COLS), 6, + cutlass.Int32), + prims.Tcgen05SmemDesc.build(sfb2.subview(jt * NUM_BYTES_SFB), leading_byte_offset=16, + stride_byte_offset=128, base_offset=0, layout=0), + multicast=s2t_multicast, + ) # fmt: skip + acc2_ptr = cutlass.inttoptr( + (base_row_id << 16) | (base_col_id + ACC2_COL + i2 * N2), 6, cutlass.Float32 + ) + for t2c in cutlass.range_constexpr(KT2): + stage = g % cutlass.Int32(STAGES) + while not cute.arch.mbarrier_try_wait(fc2_full.subview(stage).data_ptr(), + (j // cutlass.Int32(STAGES)) % cutlass.Int32(2)): # fmt: skip + pass + sfa_tmem_ptr = cutlass.inttoptr((base_row_id << 16) | (base_col_id + SFA2_COL + j * SFA_COLS), 6, + cutlass.Int32) # fmt: skip + sfb_tmem_ptr = cutlass.inttoptr((base_row_id << 16) | (base_col_id + SFB2_COL + j * SFB_COLS), 6, + cutlass.Int32) # fmt: skip + desc_a_base = prims.Tcgen05SmemDesc.build( + sA.subview(stage * NUM_BYTES_A), leading_byte_offset=16, stride_byte_offset=1024, base_offset=0, + layout=2, + ) # fmt: skip + desc_b_base = prims.Tcgen05SmemDesc.build( + b2.subview(j * NUM_BYTES_B2), leading_byte_offset=16, stride_byte_offset=1024, base_offset=0, + layout=2, + ) # fmt: skip + for kb2 in cutlass.range_constexpr(min(NUM_KBLOCKS, KB2 - NUM_KBLOCKS * t2c)): + if prims.elect_sync(): + prims.tcgen05_mma_block_scale( + prims.MMABlockScaleKind.MXF8F6F4, prims.CTAGroup.CTA_1, acc2_ptr, + desc_a_base + ((MMA_INST_K * a_smem_width // 8) >> 4) * kb2, + desc_b_base + ((MMA_INST_K * a_smem_width // 8) >> 4) * kb2, + idesc2.set_sf_ids(a_sf_id=kb2, b_sf_id=kb2), t2c != 0 or kb2 != 0, + sfa_tmem_ptr, sfb_tmem_ptr, + ) # fmt: skip + if prims.elect_sync(): + prims.tcgen05_commit(ab_empty.subview(stage)) + g = g + cutlass.Int32(1) + j = j + cutlass.Int32(1) + if prims.elect_sync(): + prims.tcgen05_commit(acc2_grp.subview(i2)) + if prims.elect_sync(): + prims.tcgen05_commit(acc2_full) + + # ---- FC2 intermediate loader for the groups with second-wave FC1 units (warps 5-6, after their FC1 duties, in + # parallel with the epilogue's groups and folds): each warp takes every other group, one group at a time. + if (warp == X_WARP) | (warp == S_WARP): + if is_fc2 & (g_n > cutlass.Int32(0)): + b2_32 = cutlass.Array( + b2.data_ptr(0), shape=(NT2 * NUM_BYTES_B2 // 4,), dtype=cutlass.Int32, alignment=16 + ) + sfb2_32 = cutlass.Array(sfb2.data_ptr(0), shape=(NT2 * NUM_BYTES_SFB // 4,), dtype=cutlass.Int32, + alignment=16) # fmt: skip + n_a = cutlass.select_( + groups < cutlass.Int32(WAVE1_GROUPS), groups, cutlass.Int32(WAVE1_GROUPS) + ) + for gi in cutlass.range(n_a + warp - cutlass.Int32(X_WARP), groups, 2, unroll=1): + _load_group(gi, lane, g_n, cnt_base, hbuf32, b2_32, sfb2_32, b2_ready) + + # ---- epilogue warps (0-3): FC2 weight scales spliced during FC1; the FC1 epilogue per unit; then FC2 per group. + if warp < 4: + _wait(tmem_ready, 0) + tmem_raw_addr = tmem_ptr_i32.load() + base_col_id = tmem_raw_addr & 0xFFFF + base_row_id = tmem_raw_addr >> 16 + row_id_with_warp_offset = base_row_id + warp * 32 + # FC2 weight scale atoms: MMA atom j (group i = j / KT2, k-tile t) word 4 m0 + jj = slot 4 i + jj's w2 atom + # (its 128-row block, k-atom t) word 4 m0 + m1, m1 = this CTA's 32-row block within the 128. + if is_fc2: + sfa2_32 = cutlass.Array(sfa2.data_ptr(0), shape=(NT2 * NUM_BYTES_SFA // 4,), dtype=cutlass.Int32, + alignment=16) # fmt: skip + vals = [] + for r in cutlass.range_constexpr(NT2 * (NUM_BYTES_SFA // 4) // EPI_THREADS): + w = tidx + cutlass.Int32(r * EPI_THREADS) + j = w // cutlass.Int32(NUM_BYTES_SFA // 4) + rem = w % cutlass.Int32(NUM_BYTES_SFA // 4) + slot = (j // cutlass.Int32(KT2)) * cutlass.Int32(4) + rem % cutlass.Int32(4) + ok = (j < nt2) & (slot < g_n) + e = cutlass.Int32(s_el.load(idx=cutlass.select_(ok, slot, cutlass.Int32(0)))) + src = (((e * cutlass.Int32(H // 128) + blk2) * cutlass.Int32(KA2) + j % cutlass.Int32(KT2)) + * cutlass.Int32(NUM_BYTES_SFA // 4) + + (rem // cutlass.Int32(4)) * cutlass.Int32(4) + m1) # fmt: skip + v = w2s32.load(idx=cutlass.select_(ok, src, cutlass.Int32(0))) + vals.append(cutlass.select_(ok, cutlass.Int32(v), cutlass.Int32(0))) + for r in cutlass.range_constexpr(NT2 * (NUM_BYTES_SFA // 4) // EPI_THREADS): + sfa2_32.store(vals[r], idx=tidx + cutlass.Int32(r * EPI_THREADS)) + prims.fence_proxy("async_shared", space=prims.SharedSpace.shared_cta) + prims.mbarrier_arrive(sfa2_ready) + # FC1 epilogue: lanes 0-63 columns 0..1 (k-tiles 2j) + lanes 64-127 columns 8..9 (k-tiles 2j + 1) = the unit's + # 64 rows for both tokens; warps 0-1 apply k3_moe's SiTU + MXFP8 requant to the 32 intermediate columns (one MX + # block) per token. + is_up = ((lane // 8) % 2) == 0 + up_mask = cutlass.select_(is_up, cutlass.Float32(1.0), cutlass.Float32(0.0)) + fc1_col_in_tile = warp * 16 + 2 * (lane % 8) + lane // 16 + acc_full_phase = 0 + for ui in range(mine): + u = _unit_of(ui, bx, pos2) + grp = u // cutlass.Int32(U1) + t128 = (u % cutlass.Int32(U1)) // cutlass.Int32(2) + half = u % cutlass.Int32(2) + while not cute.arch.mbarrier_try_wait(acc_full.data_ptr(), acc_full_phase): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc_full_phase = acc_full_phase ^ 1 + tmem_ld = cutlass.inttoptr( + (row_id_with_warp_offset << 16) | base_col_id, 6, cutlass.Float32 + ) + t2r_rmem = prims.tcgen05_ld("32x32b", tmem_ld, num=N) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + prims.mbarrier_arrive(acc_empty) + if warp >= 2: + for t in cutlass.range_constexpr(M_MAX): + s_part.store( + cutlass.Float32(t2r_rmem[8 + t]), idx=t * ROWS1 + (warp - 2) * 32 + lane + ) + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + res = [] + for t in cutlass.range_constexpr(M_MAX): + xn = cutlass.Float32(t2r_rmem[t]) + s_part.load( + idx=t * ROWS1 + (warp % 2) * 32 + lane + ) + partner = cute.arch.shuffle_sync_bfly(xn, 8) + rt = _situ(partner, xn) + res.append(rt) + warp_amax = prims.redux_sync(cute.math.abs(rt) * up_mask, prims.ReductionKind.FMAX, 0xFFFFFFFF, + abs=True) # fmt: skip + if lane == 0: + if warp < 2: + localmax_smem.store(warp_amax, idx=t * 2 + warp) + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + for t in cutlass.range_constexpr(M_MAX): + block_amax = cute.arch.fmax( + localmax_smem.load(idx=t * 2), localmax_smem.load(idx=t * 2 + 1) + ) + byte, inv_scale = _block_e8m0(block_amax) + qv = cute.arch.fmax( + cute.arch.fmin(res[t] * inv_scale, cutlass.Float32(E4M3_MAX)), + cutlass.Float32(-E4M3_MAX), + ) + fp8_i8 = cutlass.Float8E4M3FN(qv).bitcast(cutlass.Int8) + if fp8_i8 == cutlass.Int8(FP8_SENTINEL_I8): + fp8_i8 = cutlass.Int8(0) + row_base = (grp * cutlass.Int32(M_MAX) + cutlass.Int32(t)) * cutlass.Int32(H_ROW) + if (warp < 2) & is_up: + hbuf.store(fp8_i8, idx=row_base + t128 * cutlass.Int32(MMA_M // 2) + half * cutlass.Int32(32) + + fc1_col_in_tile, alignment=1) # fmt: skip + if tidx == 0: + hbuf.store(cutlass.Int8(byte & cutlass.Int32(0xFF)), + idx=row_base + cutlass.Int32(I_TP) + t128 * cutlass.Int32(2) + half, + alignment=1) # fmt: skip + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + if tidx == 0: + _red_release_add( + cnt_base + cutlass.Int64((grp // cutlass.Int32(4)) * (CW * 4)), cutlass.Int32(1) + ) + # FC2 per group: once its counter reaches its experts' units, the intermediate rows into the group's B tiles + # (row 2 j + t of tile (i, t2) = slot 4 i + j, token t, K 128 t2 .., 128B-swizzled) and their scales into the B + # scale atoms (byte 16 (2 j + t) + k = block 4 t2 + k), one load round; then warp 7's MMAs. No local expert: + # zero rows. + if is_fc2 & (g_n == cutlass.Int32(0)): + if warp < 2: + if warp < m: + _emit(cutlass.Float32(0.0), lane, row2, warp, warp * cutlass.Int32(H) + row2 + lane, out, lat_mc, + lat_flags, lat_rank) # fmt: skip + if is_fc2 & (g_n > cutlass.Int32(0)): + b2_32 = cutlass.Array( + b2.data_ptr(0), shape=(NT2 * NUM_BYTES_B2 // 4,), dtype=cutlass.Int32, alignment=16 + ) + sfb2_32 = cutlass.Array(sfb2.data_ptr(0), shape=(NT2 * NUM_BYTES_SFB // 4,), dtype=cutlass.Int32, + alignment=16) # fmt: skip + # The combine state of token `warp` (warps < M): k3_moe's sum tree, min(G, 5) slices of consecutive + # slots (ascending id, slice s = slots [s G / S, (s + 1) G / S)), each summed from 0 over the slots the + # token routes to, then the slices summed in order (selects keep every product and sum rounded on its + # own). Each group's 4 slots are folded in as soon as its MMAs are done. + n_sl = cutlass.select_(g_n < cutlass.Int32(REF_SLICES), g_n, cutlass.Int32(REF_SLICES)) + n_sl = cutlass.select_(n_sl < cutlass.Int32(1), cutlass.Int32(1), n_sl) + bnd = [] + for b in cutlass.range_constexpr(1, REF_SLICES): + bnd.append( + cutlass.select_( + cutlass.Int32(b) < n_sl, cutlass.Int32(b) * g_n // n_sl, cutlass.Int32(-1) + ) + ) + zero = cutlass.Float32(0.0) + acc = zero + part = zero + # The groups whose FC1 units all run in the first wave (warps 5-6 load the rest): warp w takes groups w, + # w + 4, .., each as soon as its own counter fills. + n_a = cutlass.select_( + groups < cutlass.Int32(WAVE1_GROUPS), groups, cutlass.Int32(WAVE1_GROUPS) + ) + for gi in cutlass.range(warp, n_a, 4, unroll=1): + _load_group(gi, lane, g_n, cnt_base, hbuf32, b2_32, sfb2_32, b2_ready) + # Fold each group into the combine as its MMAs complete: warp w's rows hold slot 4 gi + w, D + # columns 2 w + t. + for gi in cutlass.range(groups, unroll=1): + _wait(acc2_grp.subview(gi), 0) + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + yv = prims.tcgen05_ld("32x32b", cutlass.inttoptr( + (row_id_with_warp_offset << 16) | (base_col_id + ACC2_COL + gi * N2 + 2 * warp), 6, + cutlass.Float32), num=M_MAX) # fmt: skip + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + sl = gi * cutlass.Int32(4) + warp + if sl < g_n: + for t in cutlass.range_constexpr(M_MAX): + ys.store( + cutlass.Float32(yv[t]), + idx=(sl * cutlass.Int32(M_MAX) + cutlass.Int32(t)) * cutlass.Int32(R2) + + lane, + ) + cute.arch.barrier(barrier_id=EPI_BAR_ID, number_of_threads=EPI_THREADS) + if warp < m: + for jj in cutlass.range_constexpr(4): + sc = gi * cutlass.Int32(4) + cutlass.Int32(jj) + start = bnd[0] == sc + for b in cutlass.range_constexpr(1, REF_SLICES - 1): + start = start | (bnd[b] == sc) + acc = cutlass.select_(start, acc + part, acc) + part = cutlass.select_(start, zero, part) + wsc = s_w.load(idx=sc * cutlass.Int32(M_MAX) + warp) + yv_sc = ys.load( + idx=(sc * cutlass.Int32(M_MAX) + warp) * cutlass.Int32(R2) + lane + ) + part = part + cutlass.select_((sc < g_n) & (wsc != zero), yv_sc * wsc, zero) + if warp < m: + acc = acc + part + _emit(acc, lane, row2, warp, warp * cutlass.Int32(H) + row2 + + cutlass.Int32(4) * (lane % cutlass.Int32(8)) + lane // cutlass.Int32(8), out, lat_mc, lat_flags, + lat_rank) # fmt: skip + + # ---- teardown + if tidx == 0: + epochs.store(ep + cutlass.Int32(1), idx=bx) + prims.barrier_cta_sync(0) + if warp == MMA_WARP: + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + prims.tcgen05_dealloc(cutlass.inttoptr(tmem_ptr_i32.load(), 6, cutlass.Float32), TMEM_COLS) + + +@cute.jit +def k3_moe_m2( + a1_tensor: cute.Tensor, # w3_w1_weight viewed (H/2, 2 I_PAD, E) FP4 bytes, K-major + b1_tensor: cute.Tensor, # MXFP8 activations viewed (H, M) FP8 + sfa1_tensor: cute.Tensor, # w3_w1_weight_scale viewed (512, H/128, 2 I_PAD/128, E) + sfb1_tensor: cute.Tensor, # activation scales (M, H/32) E8M0 + a2_tensor: cute.Tensor, # w2_weight viewed (I_PAD/2, H, E) FP4 bytes, K-major + ids: cute.Tensor, wts: cute.Tensor, w2s32: cute.Tensor, hbuf: cute.Tensor, hbuf32: cute.Tensor, + counts: cute.Tensor, epochs: cute.Tensor, out: cute.Tensor, + lat_mc: cute.Tensor, # PUSH: int32 words of the latent exchange's multicast mapping (else any int32 tensor) + lat_flags: cute.Tensor, # PUSH: int32 [4] of the exchange (else any int32 tensor) + offset: cutlass.Int32, m: cutlass.Int32, lat_rank: cutlass.Int32, stream: cuda_driver.CUstream, +): # fmt: skip + _kpp = a1_tensor.shape[0] + _mw = a1_tensor.shape[1] + _ew = a1_tensor.shape[2] + tma_a1_desc = cuda.create_tensor_map_tiled( + global_address=a1_tensor.iterator.toint(), dtype=a_dtype, global_dims=[_kpp * 2, _mw, _ew], + global_strides=[_kpp // 16, (_mw * _kpp) // 16], box_dims=(MMA_TILE_K, ROWS1, 1), + swizzle=cuda.TensorMapSwizzle.s128b, tma_format=TensorMapDataType.f416u4_align16b, + ) # fmt: skip + tma_b1_desc = cuda.create_tensor_map_tiled_from_view( + b1_tensor, + box_dims=(MMA_TILE_K, M_MAX), + stride_order=(0, 1), + swizzle=cuda.TensorMapSwizzle.s128b, + ) + sfa1_fp16 = cute.recast_tensor(sfa1_tensor, cutlass.Uint16) + tma_sfa1_desc = cuda.create_tensor_map_tiled_from_view( + sfa1_fp16, box_dims=(num_elts_atom_sf_fp16, 1, 1, 1), stride_order=(0, 1, 2, 3), + swizzle=cuda.TensorMapSwizzle.none, + ) # fmt: skip + _kp2 = a2_tensor.shape[0] + _h2 = a2_tensor.shape[1] + _e2 = a2_tensor.shape[2] + tma_a2_desc = cuda.create_tensor_map_tiled( + global_address=a2_tensor.iterator.toint(), dtype=a_dtype, global_dims=[_kp2 * 2, _h2, _e2], + global_strides=[_kp2 // 16, (_h2 * _kp2) // 16], box_dims=(MMA_TILE_K, R2, 1), + swizzle=cuda.TensorMapSwizzle.s128b, tma_format=TensorMapDataType.f416u4_align16b, + ) # fmt: skip + k3_moe_m2_kernel( + tma_a1_desc, tma_b1_desc, tma_sfa1_desc, tma_a2_desc, sfb1_tensor.iterator.toint(), ids, wts, w2s32, hbuf, + hbuf32, counts, epochs, out, offset, m, lat_mc, lat_flags, lat_rank, + ).launch(grid=(NCTA, 1, 1), block=(THREADS, 1, 1), cluster=(1, 1, 1), stream=stream, use_pdl=True) # fmt: skip diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py index 69f70d200ba9..60e151a6b98e 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py @@ -226,6 +226,8 @@ def build(uc, mc, handle, comm): # (5), the head flags build's ready words and head flags (2), the latent slab (1). The builds this module compiles # read only the ready words and head flags (head_flags), so the others are given a stand-in they never touch. _OPTIONAL_ARGS = 11 +# int32 words of one slot of a K3LatentExchange: two halves of 8 tokens' bf16 [3584] rows. +_EXCHANGE_SLOT_WORDS = 2 * _TOKEN_SLOTS * (HIDDEN_SIZE // 2) # (alignment, leading dim) of every tensor argument, in the kernel's order. _ALIGNS = [16, 16, 16, 16, 16, 4, 16, 16, 16, 16, 16, 16, 16, 16, 16, 4, 4, 4] _ALIGNS += [16] * _OPTIONAL_ARGS @@ -567,3 +569,576 @@ def _(x_fp8, x_sf, topk_ids, topk_weights, w3_w1_weight, w3_w1_weight_scale, w2_ out=None): # fmt: skip rows = 0 if out is not None else topk_ids.shape[0] return x_fp8.new_empty((rows, HIDDEN_SIZE), dtype=torch.bfloat16) + +# --------------------------------------------------------------------------------------------------------------------- +# trtllm::k3_moe_m1 / trtllm::k3_moe_m2: the routed experts of one or two decode tokens as weight-stream kernels +# (k3_moe_m1_kernel.py, k3_moe_m2_kernel.py), on a caller-owned workspace. +# --------------------------------------------------------------------------------------------------------------------- +_ENGINE_KERNEL_PATHS = { + "m1": os.path.join(os.path.dirname(os.path.abspath(__file__)), "k3_moe_m1_kernel.py"), + "m2": os.path.join(os.path.dirname(os.path.abspath(__file__)), "k3_moe_m2_kernel.py"), +} +_M1_MIN_CTAS = HIDDEN_SIZE // 32 # one 32-row block of the down projection per CTA +# (alignment, leading dim) per tensor argument of k3_moe_m1 / k3_moe_m2, in _engine_args' order. +_ENGINE_ALIGNS = [16, 16, 16, 16, 16, 4, 2, 4, 16, 16, 4, 4, 2, 16, 4] +_ENGINE_LEADING = [0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0] + + +def _engine_module(kind: str, config: dict): + """One ``k3_moe_m1`` / ``k3_moe_m2`` module per configuration (shapes are trace-time constants), kept with + k3_moe's.""" + key = (kind,) + tuple(sorted(config.items())) + mod = _modules.get(key) + if mod is None: + name = f"{__name__}_{kind}_kernel_" + "_".join(f"{k}{v}" for k, v in key[1:]) + spec = importlib.util.spec_from_file_location(name, _ENGINE_KERNEL_PATHS[kind]) + mod = importlib.util.module_from_spec(spec) + mod.K3_CONFIG = dict(config) + sys.modules[name] = mod + spec.loader.exec_module(mod) + _modules[key] = mod + return mod + + +def m1_supported( + w3_w1_weight: torch.Tensor, + w3_w1_weight_scale: torch.Tensor, + w2_weight: torch.Tensor, + w2_weight_scale: torch.Tensor, + local_num_experts: int, + intermediate_size: int, +) -> Tuple[bool, str]: + """Whether these TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers fit ``k3_moe_m1``. ``intermediate_size`` is a rank's + logical intermediate; the buffers may hold it zero-padded to a multiple of 128, as TRT-LLM's loader lays out an + unaligned shard (e.g. 192 -> 256 at TP16). Only metadata is read, so the buffers may still be on the meta + device.""" + if not torch.cuda.is_available() or torch.cuda.get_device_capability() != (10, 0): + return False, "needs sm_100" + try: + import cutlass # noqa: F401 + except ImportError: + return False, "CuTe DSL (nvidia-cutlass-dsl) is not installed" + e, two_i, _ = w3_w1_weight.shape + i_pad = two_i // 2 + expected = { + "w3_w1_weight": (w3_w1_weight, (local_num_experts, two_i, HIDDEN_SIZE // 2)), + "w3_w1_weight_scale": ( + w3_w1_weight_scale, + (local_num_experts, two_i, HIDDEN_SIZE // _SF_VEC), + ), + "w2_weight": (w2_weight, (local_num_experts, HIDDEN_SIZE, i_pad // 2)), + "w2_weight_scale": (w2_weight_scale, (local_num_experts, HIDDEN_SIZE, i_pad // _SF_VEC)), + } + for name, (t, shape) in expected.items(): + if t.dtype != torch.uint8 or tuple(t.shape) != shape or not t.is_contiguous(): + return False, f"{name} is {t.dtype} {tuple(t.shape)}, expected contiguous uint8 {shape}" + if i_pad % 128 != 0 or not 0 < intermediate_size <= i_pad or intermediate_size % 32 != 0: + return False, f"intermediate {intermediate_size} in buffers of {i_pad}" + if local_num_experts > NUM_EXPERTS: + return False, f"{local_num_experts} local experts" + sms = torch.cuda.get_device_properties(torch.cuda.current_device()).multi_processor_count + if sms < _M1_MIN_CTAS: + return False, f"{sms} SMs, the kernel needs {_M1_MIN_CTAS}" + return True, "" + + +def m2_supported( + w3_w1_weight: torch.Tensor, + w3_w1_weight_scale: torch.Tensor, + w2_weight: torch.Tensor, + w2_weight_scale: torch.Tensor, + local_num_experts: int, + intermediate_size: int, +) -> Tuple[bool, str]: + """Whether these TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers fit ``k3_moe_m2``: ``k3_moe_m1``'s conditions, and an + intermediate of at most 256 (the TP16 slice; its FC2 tiles of every expert slot stay in shared memory).""" + ok, why = m1_supported( + w3_w1_weight, + w3_w1_weight_scale, + w2_weight, + w2_weight_scale, + local_num_experts, + intermediate_size, + ) + if ok and w3_w1_weight.shape[1] // 2 > 256: + return False, f"padded intermediate {w3_w1_weight.shape[1] // 2} > 256" + return ok, why + + +def _engine_config(i_tp: int, i_pad: int, num_local: int, num_ctas: int, m_max: int, push_world: int = 0, + copies: int = 1) -> dict: # fmt: skip + """The kernel options of one ``k3_moe_m1`` / ``k3_moe_m2`` build (trace-time constants). ``push_world``: the push + build for a latent exchange of that many slots, this rank filling ``copies`` of them.""" + config = { + "i_tp": i_tp, + "i_pad": i_pad, + "num_local": num_local, + "num_ctas": num_ctas, + "m_max": m_max, + } + if push_world: + config.update(push=1, push_world=push_world, push_copies=copies) + return config + + +def _engine_args( + weights, x_fp8, x_sf, topk_ids, topk_weights, hbuf, counts, epochs, y, lat_mc, lat_flags +) -> tuple: + """The tensor arguments of ``k3_moe_m1`` / ``k3_moe_m2`` in their order: the weight / activation views, then the + flat arrays (ids, weights, w2 scale words, intermediate rows and their words, counts, epochs, output, the latent + exchange's multicast words and flags).""" + w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale = weights + e, two_i, _ = w3_w1_weight.shape + m = topk_ids.shape[0] + return ( + w3_w1_weight.view(torch.int8).permute(2, 1, 0), x_fp8.view(torch.uint8).permute(1, 0), + w3_w1_weight_scale.view(e, two_i // 128, HIDDEN_SIZE // 128, 512).permute(3, 2, 1, 0), + x_sf.view(torch.uint8).view(m, HIDDEN_SIZE // _SF_VEC), w2_weight.view(torch.int8).permute(2, 1, 0), + topk_ids.view(-1), topk_weights.view(torch.int16).view(-1), w2_weight_scale.view(-1).view(torch.int32), hbuf, + hbuf.view(torch.int32), counts, epochs, y.view(-1), lat_mc.view(-1), lat_flags.view(-1), + ) # fmt: skip + + +def _engine_scalars(kind: str, local_expert_offset: int, num_tokens: int, lat_rank: int) -> tuple: + if kind == "m1": + return (local_expert_offset, lat_rank) + return (local_expert_offset, num_tokens, lat_rank) + + +def _engine_compile(kind: str, mod, args, scalars): + """TVM-FFI build of ``mod``'s kernel for these torch arguments' types and layouts.""" + import cutlass.cute as cute + + assert len(args) == len(_ENGINE_ALIGNS) + signature = [_view(t, a, d) for t, a, d in zip(args, _ENGINE_ALIGNS, _ENGINE_LEADING)] + return cute.compile( + getattr(mod, f"k3_moe_{kind}"), + *signature, + *scalars, + cute.runtime.make_fake_stream(), + options="--enable-tvm-ffi", + ) + + +def _engine_workspace_sizes(kind: str, mod, m: int) -> Tuple[int, int]: + """(intermediate row bytes, count words) of one build's workspace for ``m`` tokens.""" + if kind == "m1": + return mod.G_CAP * m * mod.H_ROW, 4 + return mod.G_CAP * mod.M_MAX * mod.H_ROW, 2 * mod.GROUPS2 * mod.CW + + +def _engine_call(kind, x_fp8, x_sf, topk_ids, topk_weights, weights, hbuf, counts, epochs, local_expert_offset, + intermediate_size, lat_mc, lat_flags, lat_rank, copies, out, compile_num_local=0): # fmt: skip + """The body of ``trtllm::k3_moe_m1`` / ``trtllm::k3_moe_m2``: checks, then the build's launch (compiled on its + first call). ``compile_num_local``: compile the build for ``compile_num_local`` local experts and launch nothing + (``weights`` may then be stand-ins of any expert count).""" + name = f"k3_moe_{kind}" + m = topk_ids.shape[0] + if (m not in (1, 2)) if kind == "m1" else m != 2: + raise ValueError( + f"{name}: {'1 or 2 tokens' if kind == 'm1' else '2 tokens'} per call, got {m}" + ) + if ( + topk_ids.dtype != torch.int32 + or tuple(topk_ids.shape) != (m, TOP_K) + or topk_weights.dtype != torch.bfloat16 + or tuple(topk_weights.shape) != (m, TOP_K) + or x_fp8.dtype != torch.float8_e4m3fn + or tuple(x_fp8.shape) != (m, HIDDEN_SIZE) + or x_sf.numel() != m * (HIDDEN_SIZE // _SF_VEC) + or not (topk_ids.is_contiguous() and topk_weights.is_contiguous()) + or not (x_fp8.is_contiguous() and x_sf.is_contiguous()) + ): + raise ValueError(f"{name}: expects {m} token(s)' routing and MXFP8 latents") + num_local = compile_num_local or weights[0].shape[0] + if not compile_num_local: + supported = m1_supported if kind == "m1" else m2_supported + ok, why = supported(*weights, num_local, intermediate_size) + if not ok: + raise ValueError(f"{name}: {why}") + push = lat_mc is not None + if push != (lat_flags is not None): + raise ValueError(f"{name}: exchange_mc and exchange_flags go together") + slots = 0 + if push: + slots = lat_mc.numel() // _EXCHANGE_SLOT_WORDS + if ( + lat_mc.dtype != torch.int32 + or slots == 0 + or lat_mc.numel() != slots * _EXCHANGE_SLOT_WORDS + or lat_flags.dtype != torch.int32 + or lat_flags.numel() < 1 + ): + raise ValueError( + f"{name}: the exchange is int32 words [2][8][slots][1792] and int32 flags" + ) + if copies < 1 or not 0 <= lat_rank * copies <= slots - copies: + raise ValueError( + f"{name}: slots {lat_rank * copies}..+{copies} outside the exchange's {slots} slots" + ) + if out is not None: + raise ValueError(f"{name}: the push build takes no out") + num_ctas = epochs.numel() + i_pad = weights[0].shape[1] // 2 + config = _engine_config( + intermediate_size, i_pad, num_local, num_ctas, m if kind == "m1" else 2, slots, copies + ) + mod = _engine_module(kind, config) + hbuf_bytes, count_words = _engine_workspace_sizes(kind, mod, m) + if ( + num_ctas < _M1_MIN_CTAS + or epochs.dtype != torch.int32 + or hbuf.dtype != torch.int8 + or hbuf.numel() != hbuf_bytes + or counts.dtype != torch.int32 + or counts.numel() != count_words + or not all(t.is_contiguous() for t in (hbuf, counts, epochs)) + ): + raise ValueError(f"{name}: hbuf / counts / epochs are not this build's workspace") + if out is None: + y = torch.empty(m, HIDDEN_SIZE, dtype=torch.bfloat16, device=x_fp8.device) + else: + if ( + out.dtype != torch.bfloat16 + or tuple(out.shape) != (m, HIDDEN_SIZE) + or not out.is_contiguous() + ): + raise ValueError(f"{name}: out must be contiguous bf16 [{m}, 3584]") + y = out + # The plain build never touches the exchange arguments: the counts stand in for them. + args = _engine_args(weights, x_fp8, x_sf, topk_ids, topk_weights, hbuf, counts, epochs, y, + lat_mc if push else counts, lat_flags if push else counts) # fmt: skip + scalars = _engine_scalars(kind, local_expert_offset, m, lat_rank) + key = (kind,) + _compile_key(x_fp8.device, config) + fn = _compiled.get(key) + if fn is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + f"trtllm::{name} compiles on its first call for each build: call it before capture" + ) + fn = _compiled[key] = _engine_compile(kind, mod, args, scalars) + if compile_num_local: + return None + fn(*args, *scalars, torch.cuda.current_stream(x_fp8.device).cuda_stream) + if out is not None or push: + return y.new_empty((0, HIDDEN_SIZE)) + return y + + +@torch.library.custom_op( + "trtllm::k3_moe_m1", + mutates_args=("hbuf", "counts", "epochs", "exchange_mc", "out"), +) +def k3_moe_m1( + x_fp8: torch.Tensor, + x_sf: torch.Tensor, + topk_ids: torch.Tensor, + topk_weights: torch.Tensor, + w3_w1_weight: torch.Tensor, + w3_w1_weight_scale: torch.Tensor, + w2_weight: torch.Tensor, + w2_weight_scale: torch.Tensor, + hbuf: torch.Tensor, + counts: torch.Tensor, + epochs: torch.Tensor, + local_expert_offset: int, + intermediate_size: int, + exchange_mc: Optional[torch.Tensor] = None, + exchange_flags: Optional[torch.Tensor] = None, + exchange_rank: int = 0, + copies: int = 1, + out: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """This rank's routed partial ``[M, 3584]`` bf16 for M = 1 or 2 decode tokens from the weight-stream kernel + ``k3_moe_m1``: FC1 + SiTU + FC2 with k3_moe's combine over this rank's experts. + + ``x_fp8`` float8_e4m3fn ``[M, 3584]``, ``x_sf`` its E8M0 scales (``M * 112`` bytes), ``topk_ids`` int32 + ``[M, 16]`` global expert ids and ``topk_weights`` bf16 ``[M, 16]``: the outputs of ``trtllm::k3_route_quant`` or + ``trtllm::k3_moe_front``. The weights are the TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers of a rank's experts (global ids + ``[local_expert_offset, local_expert_offset + E)``), read in place, holding an ``intermediate_size`` slice + zero-padded to a multiple of 128. ``hbuf``, ``counts``, ``epochs``: a :class:`K3MoeM1State`'s workspace for M + tokens; every call re-arms the count slot the next call uses and advances the epochs. ``out``: contiguous bf16 + ``[M, 3584]``; the call writes it and returns an empty ``[0, 3584]`` instead of a new tensor. + + ``exchange_mc`` / ``exchange_flags``: a ``K3LatentExchange``'s multicast words and flags. With them the call is + the push build: the partial goes into slots ``exchange_rank * copies .. + copies - 1`` of half + ``exchange_flags[0] & 1`` of every rank's exchange (bf16 pairs, -0.0 stored as +0.0) for + ``trtllm::k3_latent_reduce``, and the call returns an empty ``[0, 3584]``. Each push of M tokens is followed by one + reduce of M tokens on that exchange before the next push.""" + weights = (w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale) + return _engine_call("m1", x_fp8, x_sf, topk_ids, topk_weights, weights, hbuf, counts, epochs, + local_expert_offset, intermediate_size, exchange_mc, exchange_flags, exchange_rank, copies, + out) # fmt: skip + + +@k3_moe_m1.register_fake +def _(x_fp8, x_sf, topk_ids, topk_weights, w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale, hbuf, counts, + epochs, local_expert_offset, intermediate_size, exchange_mc=None, exchange_flags=None, exchange_rank=0, copies=1, + out=None): # fmt: skip + rows = 0 if out is not None or exchange_mc is not None else topk_ids.shape[0] + return x_fp8.new_empty((rows, HIDDEN_SIZE), dtype=torch.bfloat16) + + +@torch.library.custom_op( + "trtllm::k3_moe_m2", + mutates_args=("hbuf", "counts", "epochs", "exchange_mc", "out"), +) +def k3_moe_m2( + x_fp8: torch.Tensor, + x_sf: torch.Tensor, + topk_ids: torch.Tensor, + topk_weights: torch.Tensor, + w3_w1_weight: torch.Tensor, + w3_w1_weight_scale: torch.Tensor, + w2_weight: torch.Tensor, + w2_weight_scale: torch.Tensor, + hbuf: torch.Tensor, + counts: torch.Tensor, + epochs: torch.Tensor, + local_expert_offset: int, + intermediate_size: int, + exchange_mc: Optional[torch.Tensor] = None, + exchange_flags: Optional[torch.Tensor] = None, + exchange_rank: int = 0, + copies: int = 1, + out: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """``trtllm::k3_moe_m1``'s contract for exactly two tokens, from the weight-stream kernel ``k3_moe_m2`` (an + intermediate slice of at most 256; the FC2 tiles of every expert slot stay in shared memory) on a + :class:`K3MoeM2State`'s workspace.""" + weights = (w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale) + return _engine_call("m2", x_fp8, x_sf, topk_ids, topk_weights, weights, hbuf, counts, epochs, + local_expert_offset, intermediate_size, exchange_mc, exchange_flags, exchange_rank, copies, + out) # fmt: skip + + +@k3_moe_m2.register_fake +def _(x_fp8, x_sf, topk_ids, topk_weights, w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale, hbuf, counts, + epochs, local_expert_offset, intermediate_size, exchange_mc=None, exchange_flags=None, exchange_rank=0, copies=1, + out=None): # fmt: skip + rows = 0 if out is not None or exchange_mc is not None else topk_ids.shape[0] + return x_fp8.new_empty((rows, HIDDEN_SIZE), dtype=torch.bfloat16) + + +class _K3MoeEngineState: + """The workspace of ``trtllm::k3_moe_m1`` / ``trtllm::k3_moe_m2`` on one device, shared by the layers that run on + it: the intermediate rows (``hbuf``), the FC1 -> FC2 counts (``counts``, two sets by epoch parity) and the CTAs' + epochs (``epochs``). Every call re-arms the count set the next call uses and advances the epochs, so the + workspace never needs a reset; its layers run in one stream order. Build it with ``create`` before CUDA-graph + capture and keep it with the model: ``create`` compiles the plain build and the push builds it is given. A build + it did not compile compiles on its first call, which must come before capture too.""" + + kind = "" + + def __init__( + self, device: torch.device, i_tp: int, i_pad: int, num_local: int, num_tokens: int + ): + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + f"{type(self).__name__} allocates its workspace: build it outside CUDA-graph capture" + ) + num_ctas = torch.cuda.get_device_properties(device).multi_processor_count + self.config = _engine_config(i_tp, i_pad, num_local, num_ctas, num_tokens) + self.mod = mod = _engine_module(self.kind, self.config) + self.device = device + self.i_tp = i_tp + self.i_pad = i_pad + self.num_local = num_local + self.num_tokens = num_tokens + hbuf_bytes, count_words = _engine_workspace_sizes(self.kind, mod, num_tokens) + kw = dict(device=device) + self.hbuf = torch.zeros(hbuf_bytes, dtype=torch.int8, **kw) + self.counts = torch.zeros(count_words, dtype=torch.int32, **kw) + self.epochs = torch.zeros(num_ctas, dtype=torch.int32, **kw) + + def _create(self, push: Tuple[Tuple[int, int], ...]) -> "_K3MoeEngineState": + """Compile the plain build and a push build per (exchange slots, copies) in ``push``, launching nothing: two + experts' zero buffers and zero tokens stand in for a layer's arguments (the builds take any shapes of these + types and layouts).""" + kw = dict(device=self.device) + m, i_pad = self.num_tokens, self.i_pad + weights = ( + torch.zeros(2, 2 * i_pad, HIDDEN_SIZE // 2, dtype=torch.uint8, **kw), + torch.zeros(2, 2 * i_pad, HIDDEN_SIZE // _SF_VEC, dtype=torch.uint8, **kw), + torch.zeros(2, HIDDEN_SIZE, i_pad // 2, dtype=torch.uint8, **kw), + torch.zeros(2, HIDDEN_SIZE, i_pad // _SF_VEC, dtype=torch.uint8, **kw), + ) + tokens = ( + torch.zeros(m, HIDDEN_SIZE, dtype=torch.float8_e4m3fn, **kw), + torch.zeros(m, HIDDEN_SIZE // _SF_VEC, dtype=torch.uint8, **kw), + torch.zeros(m, TOP_K, dtype=torch.int32, **kw), + torch.zeros(m, TOP_K, dtype=torch.bfloat16, **kw), + ) + builds = [(None, None, 1)] + for slots, copies in push: + lat_mc = torch.zeros(slots * _EXCHANGE_SLOT_WORDS, dtype=torch.int32, **kw) + builds.append((lat_mc, torch.zeros(4, dtype=torch.int32, **kw), copies)) + for lat_mc, lat_flags, copies in builds: + _engine_call(self.kind, *tokens, weights, self.hbuf, self.counts, self.epochs, 0, self.i_tp, lat_mc, + lat_flags, 0, copies, None, compile_num_local=self.num_local) # fmt: skip + return self + + @property + def compiled(self) -> bool: + """Whether the plain build has been compiled on this device (by any state's ``create`` or first call).""" + return (self.kind,) + _compile_key(self.device, self.config) in _compiled + + def push_compiled(self, slots: int, copies: int = 1) -> bool: + """Whether the push build for an exchange of ``slots`` slots, this rank filling ``copies``, is compiled.""" + config = _engine_config(self.i_tp, self.i_pad, self.num_local, self.epochs.numel(), self.config["m_max"], + slots, copies) # fmt: skip + return (self.kind,) + _compile_key(self.device, config) in _compiled + + +class K3MoeM1State(_K3MoeEngineState): + """``trtllm::k3_moe_m1``'s workspace on one device for ``num_tokens`` (1 or 2) tokens per call: see + :class:`_K3MoeEngineState`. The count set is two words (one per epoch parity).""" + + kind = "m1" + + @classmethod + def create( + cls, + device: torch.device, + i_tp: int, + i_pad: int, + num_local: int, + num_tokens: int = 1, + push: Tuple[Tuple[int, int], ...] = (), + ) -> "K3MoeM1State": + """The state on ``device`` with its builds compiled: the plain build, and the push build for each (exchange + slots, copies) in ``push`` (``((tp_size, 1),)`` for a TP group's ``K3LatentExchange``). Eager: it allocates and + compiles, so it refuses to run under CUDA-graph capture. Not collective.""" + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "K3MoeM1State.create allocates and compiles: call it before CUDA-graph capture" + ) + return cls(torch.device(device), i_tp, i_pad, num_local, num_tokens)._create(push) + + def __init__( + self, device: torch.device, i_tp: int, i_pad: int, num_local: int, num_tokens: int = 1 + ): + if num_tokens not in (1, 2): + raise ValueError(f"k3_moe_m1 takes 1 or 2 tokens per call, not {num_tokens}") + super().__init__(device, i_tp, i_pad, num_local, num_tokens) + + def layer( + self, + w3_w1_weight: torch.Tensor, + w3_w1_weight_scale: torch.Tensor, + w2_weight: torch.Tensor, + w2_weight_scale: torch.Tensor, + ) -> "K3MoeM1Layer": + """A layer's handle: its experts' TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers, read in place.""" + return K3MoeM1Layer(self, w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale) + + +class K3MoeM2State(_K3MoeEngineState): + """``trtllm::k3_moe_m2``'s workspace on one device (two tokens per call): see :class:`_K3MoeEngineState`. The + count sets are per FC1 -> FC2 group.""" + + kind = "m2" + num_tokens = 2 + + @classmethod + def create( + cls, + device: torch.device, + i_tp: int, + i_pad: int, + num_local: int, + push: Tuple[Tuple[int, int], ...] = (), + ) -> "K3MoeM2State": + """The state on ``device`` with its builds compiled: as :meth:`K3MoeM1State.create`.""" + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "K3MoeM2State.create allocates and compiles: call it before CUDA-graph capture" + ) + return cls(torch.device(device), i_tp, i_pad, num_local)._create(push) + + def __init__(self, device: torch.device, i_tp: int, i_pad: int, num_local: int): + super().__init__(device, i_tp, i_pad, num_local, 2) + + def layer( + self, + w3_w1_weight: torch.Tensor, + w3_w1_weight_scale: torch.Tensor, + w2_weight: torch.Tensor, + w2_weight_scale: torch.Tensor, + ) -> "K3MoeM2Layer": + """A layer's handle: its experts' TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers, read in place.""" + return K3MoeM2Layer(self, w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale) + + +class _K3MoeEngineLayer: + """One MoE layer on a ``k3_moe_m1`` / ``k3_moe_m2`` state: its experts' weight buffers, read in place.""" + + def __init__( + self, + state: _K3MoeEngineState, + w3_w1_weight: torch.Tensor, + w3_w1_weight_scale: torch.Tensor, + w2_weight: torch.Tensor, + w2_weight_scale: torch.Tensor, + ): + supported = m1_supported if state.kind == "m1" else m2_supported + ok, why = supported( + w3_w1_weight, + w3_w1_weight_scale, + w2_weight, + w2_weight_scale, + state.num_local, + state.i_tp, + ) + if not ok or w3_w1_weight.shape[1] != 2 * state.i_pad: + raise ValueError( + f"k3_moe_{state.kind} layer: {why or 'padded intermediate differs from the state'}" + ) + self.state = state + self.weights = (w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale) + + def _op(self): + return torch.ops.trtllm.k3_moe_m1 if self.state.kind == "m1" else torch.ops.trtllm.k3_moe_m2 + + def __call__( + self, + x_fp8: torch.Tensor, + x_sf: torch.Tensor, + topk_ids: torch.Tensor, + topk_weights: torch.Tensor, + local_expert_offset: int, + out: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + """This rank's routed partial ``[M, 3584]`` bf16 for the state's M tokens (``out``, written, when given): the + op on this layer's experts and its state's workspace. See ``trtllm::k3_moe_m1``.""" + st = self.state + y = self._op()(x_fp8, x_sf, topk_ids, topk_weights, *self.weights, st.hbuf, st.counts, st.epochs, + local_expert_offset, st.i_tp, out=out) # fmt: skip + return y if out is None else out + + def push( + self, + x_fp8: torch.Tensor, + x_sf: torch.Tensor, + topk_ids: torch.Tensor, + topk_weights: torch.Tensor, + local_expert_offset: int, + lat_mc: torch.Tensor, + lat_flags: torch.Tensor, + lat_rank: int, + copies: int = 1, + ) -> None: + """The push build: the same routed partial, stored into slots ``lat_rank * copies .. + copies - 1`` of every + rank's latent exchange (``lat_mc``: the multicast int32 words of a ``K3LatentExchange``, ``lat_flags`` its + flags) for ``trtllm::k3_latent_reduce``, instead of returned. See ``trtllm::k3_moe_m1``.""" + st = self.state + self._op()(x_fp8, x_sf, topk_ids, topk_weights, *self.weights, st.hbuf, st.counts, st.epochs, + local_expert_offset, st.i_tp, lat_mc, lat_flags, lat_rank, copies) # fmt: skip + + +class K3MoeM1Layer(_K3MoeEngineLayer): + """One MoE layer on a :class:`K3MoeM1State`.""" + + +class K3MoeM2Layer(_K3MoeEngineLayer): + """One MoE layer on a :class:`K3MoeM2State`.""" diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index 0f2d55dacd77..ac6e6970ef8f 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -325,6 +325,8 @@ l0_b200: - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front_geometry.py - unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe.py - unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_route_quant.py + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_m1.py + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_m2.py # ------------- Visual Gen tests --------------- - unittest/_torch/cute_dsl_kernels/test_nvfp4_conv3d.py - unittest/_torch/visual_gen/test_media_decode.py diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_m1.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_m1.py new file mode 100644 index 000000000000..df647f7a5e47 --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_m1.py @@ -0,0 +1,306 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""k3_moe_m1 (one or two decode tokens' routed experts as a weight-stream kernel) on one GPU, at both token counts +and in the two Kimi K3 decode layouts of the routed experts: TP16 (every rank holds all 896 experts, its +192-wide intermediate slice zero-padded to 256 by TRT-LLM's loader) and TP4 x EP4 (224 local experts, +intermediate 768). Weights are random checkpoint-format +MXFP4 experts put through TRT-LLM's own loader. Per routing case, against: +- trtllm::k3_moe (K3MoeLayer) on the same buffers and routing, which it reads at the + padded size. The zero rows add exact + zeros, so the two compute the same math; k3_moe_m1 keeps k3_moe's FC2 operand and combine order, so the bits match + except where its two FC1 partial sums round an intermediate value differently. Bit identity is reported; the gate + is one bf16 ulp of the row's max; +- the fp32 reference over the dequantized experts (op-catalog gates: 8 ulp of the row max per element, 4 ulp + relative RMS). +It also checks run-to-run identical bits, and that calls of different layers sharing the state's workspace, in any +order, each give the bits of the same call alone. +""" + +import functools +import math +from types import SimpleNamespace + +import pytest +import torch + + +def _is_sm100() -> bool: + if not torch.cuda.is_available(): + return False + return torch.cuda.get_device_capability() == (10, 0) + + +pytestmark = pytest.mark.skipif(not _is_sm100(), reason="k3_moe_m1 needs sm_100") + +H, NUM_EXPERTS, TOP_K, SV = 3584, 896, 16, 32 +GATE_CAP, LINEAR_CAP = ( + 4.0, + 25.0, +) # the SiTU caps (activation_situ_beta, activation_situ_linear_beta) +RSF = 2.827 +ULP = 2.0**-8 +E4M3_MAX = 448.0 +# Layout -> (moe_tp, the rank's logical intermediate, local experts, first local id). +LAYOUTS = {"tp16": (16, 192, 896, 0), "tp4ep4": (4, 768, 224, 224)} +_E2M1 = [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0, -0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0] + + +def _ops(): + import tensorrt_llm # noqa: F401 + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op # noqa: F401 + from tensorrt_llm._torch.cute_dsl_kernels.k3_route_quant import op as _rq # noqa: F401 + + return torch.ops.trtllm + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def _rand_mxfp4(rows, k, k_full, gen): + """Random checkpoint-format MXFP4: packed [rows, k / 2] (low nibble = even k), E8M0 per 32 k, scaled so a + k_full-long dot product lands near std 3.""" + codes = torch.randint(0, 16, (rows, k), dtype=torch.uint8, device="cuda", generator=gen) + base = 127 + round(0.5 * math.log2(0.01057 / k_full)) + exps = torch.randint( + base, base + 6, (rows, k // SV), dtype=torch.uint8, device="cuda", generator=gen + ) + return (codes[:, 0::2] | (codes[:, 1::2] << 4)).contiguous(), exps + + +@functools.lru_cache(maxsize=None) +def _experts(layout: str, seed: int = 20260928): + """A rank's experts through W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod's loader (the buffers both kernels read, padded + as the loader pads an unaligned shard) and the rank's logical slices of the checkpoint tensors (what the reference + reads). Only the rank's shard is generated: an unaligned shard as rank 0 of tensors that hold exactly that shard + (the loader slices it, then pads), an aligned one as the whole tensor of a single rank.""" + from tensorrt_llm._torch.moe.fused_moe.quantization import W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod + + moe_tp, i_tp, e_local, _ = LAYOUTS[layout] + i_pad = (i_tp + 127) // 128 * 128 + tp = moe_tp if i_pad != i_tp else 1 + method = W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod() + module = SimpleNamespace(tp_size=tp, tp_rank=0, scaling_vector_size=SV, intermediate_size=i_tp * tp, + intermediate_size_per_partition=i_tp, hidden_size=H) # fmt: skip + kw = dict(dtype=torch.uint8, device="cuda") + proc = dict( + w31=torch.zeros(e_local, 2 * i_pad, H // 2, **kw), + w31s=torch.zeros(e_local, 2 * i_pad, H // SV, **kw), + w2=torch.zeros(e_local, H, i_pad // 2, **kw), + w2s=torch.zeros(e_local, H, i_pad // SV, **kw), + ) + raw = {name: [] for name in ("up", "up_s", "gate", "gate_s", "down", "down_s")} + gen = torch.Generator(device="cuda").manual_seed(seed) + for e in range(e_local): + w1, w1s = _rand_mxfp4(i_tp, H, H, gen) # gate + w3, w3s = _rand_mxfp4(i_tp, H, H, gen) # up + w2, w2s = _rand_mxfp4(H, i_tp, i_tp * moe_tp, gen) # down + method.load_expert_w3_w1_weight(module, w1, w3, proc["w31"][e]) + method.load_expert_w2_weight(module, w2, proc["w2"][e]) + method.load_expert_w3_w1_weight_scale_mxfp4(module, w1s, w3s, proc["w31s"][e]) + method.load_expert_w2_weight_scale_mxfp4(module, w2s, proc["w2s"][e]) + for name, t in zip(raw, (w3, w3s, w1, w1s, w2, w2s)): + raw[name].append(t) + torch.cuda.synchronize() + gen_b = torch.Generator(device="cuda").manual_seed(seed + 1) + bias = (torch.randn(NUM_EXPERTS, generator=gen_b, device="cuda") * 0.05).float() + return proc, raw, bias + + +def _deq_w(packed, sf): + lut = torch.tensor(_E2M1, device=packed.device) + vals = torch.empty(packed.shape[0], packed.shape[1] * 2, device=packed.device) + vals[:, 0::2] = lut[(packed & 0xF).long()] + vals[:, 1::2] = lut[(packed >> 4).long()] + return vals * torch.exp2(sf.float() - 127.0).repeat_interleave(SV, dim=1) + + +def _deq_x(x_fp8, x_sf): + rows, k = x_fp8.shape + return x_fp8.float() * torch.exp2( + x_sf.reshape(rows, k // SV).float() - 127.0 + ).repeat_interleave(SV, dim=1) + + +def _requant(act): + """The FC1 epilogue's MXFP8 requantization per 32 columns (round-up scale), dequantized.""" + rows, cols = act.shape + blocks = act.reshape(rows, cols // SV, SV) + amax = blocks.abs().amax(dim=-1, keepdim=True) + ex = torch.ceil(torch.log2(amax / E4M3_MAX)) + ex = torch.where(amax == 0, torch.full_like(amax, -127.0), ex).clamp(-127.0, 127.0) + scale = torch.exp2(ex) + q8 = (blocks / scale).clamp(-E4M3_MAX, E4M3_MAX).to(torch.float8_e4m3fn) + return (q8.float() * scale).reshape(rows, cols) + + +def _reference(layout, raw, x_deq, ids, weights): + """fp32 routed MoE over this rank's experts from the checkpoint slices (f64 GEMMs): SiTU, the MXFP8 + intermediate, the down projection, the routing-weighted sum, per token.""" + _, _, e_local, offset = LAYOUTS[layout] + out = torch.zeros(ids.shape[0], H, device="cuda") + for t in range(ids.shape[0]): + for k in range(TOP_K): + e = int(ids[t, k]) - offset + if not 0 <= e < e_local: + continue + xe = x_deq[t : t + 1].double() + up = (xe @ _deq_w(raw["up"][e], raw["up_s"][e]).double().t()).float() + gate = (xe @ _deq_w(raw["gate"][e], raw["gate_s"][e]).double().t()).float() + act = (GATE_CAP * torch.tanh(gate / GATE_CAP) * torch.sigmoid(gate) + * (LINEAR_CAP * torch.tanh(up / LINEAR_CAP))) # fmt: skip + y = ( + _requant(act).double() @ _deq_w(raw["down"][e], raw["down_s"][e]).double().t() + ).float() + out[t] += (y * weights[t, k].float())[0] + return out + + +def _row_ulp(y, ref): + """Largest |y - ref| in bf16 ulps of the row's max |ref|, and the relative RMS in ulps.""" + o, r = y.float(), ref.float() + row = r.abs().amax(dim=1, keepdim=True).clamp_min(1e-12) + elt = ((o - r).abs() / row).max().item() / ULP + rms = ((o - r).pow(2).mean().sqrt() / r.pow(2).mean().sqrt().clamp_min(1e-12)).item() / ULP + return elt, rms + + +@functools.lru_cache(maxsize=None) +def _state(layout: str, m: int): + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op + + _, i_tp, e_local, _ = LAYOUTS[layout] + return op.K3MoeM1State(torch.device("cuda", torch.cuda.current_device()), i_tp, (i_tp + 127) // 128 * 128, e_local, + num_tokens=m) # fmt: skip + + +@functools.lru_cache(maxsize=None) +def _k3_moe_state(layout: str): + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op + + _, i_tp, e_local, _ = LAYOUTS[layout] + return op.K3MoeState( + torch.device("cuda", torch.cuda.current_device()), (i_tp + 127) // 128 * 128, e_local + ) + + +def _k3_moe(layout, proc, ids, w, x_fp8, x_sf): + """trtllm::k3_moe (K3MoeLayer) on the same buffers and routing.""" + _, _, _, offset = LAYOUTS[layout] + layer = _k3_moe_state(layout).layer(proc["w31"], proc["w31s"], proc["w2"], proc["w2s"]) + return layer(x_fp8, x_sf, ids, w, offset) + + +CASES = ["random", "4_local", "none_local"] + + +def _token(layout: str, case: str, seed: int, m: int = 1): + """m tokens' router logits and latent rows. "4_local": 4 of this rank's experts in each token's top-16; + "none_local": no local expert (with all experts local, every case routes 16 local experts per token).""" + _, _, e_local, offset = LAYOUTS[layout] + gen = torch.Generator(device="cuda").manual_seed(seed) + cpu = torch.Generator().manual_seed(seed) + logits = (torch.randn(m, NUM_EXPERTS, generator=gen, device="cuda") * 3.0).float() + if e_local < NUM_EXPERTS: + if case == "4_local": + logits[:, offset : offset + e_local] -= 30.0 + for t in range(m): + logits[t, offset + torch.randperm(e_local, generator=cpu)[:4].cuda()] = 30.0 + elif case == "none_local": + logits[:, offset : offset + e_local] = -30.0 + x = torch.randn(m, H, generator=gen, device="cuda").bfloat16() + return logits, x + + +def _m1(layout, proc, bias, logits, x): + ops = _ops() + _, _, _, offset = LAYOUTS[layout] + ids, w, x_fp8, x_sf = ops.k3_route_quant(logits, bias, x, RSF, True) + layer = _state(layout, x.shape[0]).layer(proc["w31"], proc["w31s"], proc["w2"], proc["w2s"]) + return layer(x_fp8, x_sf, ids, w, offset), ids, w, x_fp8, x_sf + + +# (layout, routing case, tokens per call). With every expert local (TP16) only random routing applies; the two-token +# build is for the TP16 slice (test_k3_moe_m1_two_tokens_refuse_wide_intermediate). +CELLS = [("tp16", "random", 1), ("tp16", "random", 2)] + [("tp4ep4", case, 1) for case in CASES] + + +@pytest.mark.parametrize("layout,case,m", CELLS) +def test_k3_moe_m1(layout, case, m): + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op + + _ops() + proc, raw, bias = _experts(layout) + moe_tp, i_tp, e_local, offset = LAYOUTS[layout] + ok, why = op.m1_supported(proc["w31"], proc["w31s"], proc["w2"], proc["w2s"], e_local, i_tp) + assert ok, why + n_bit = n_local = 0 + seeds = range(8) + for seed in seeds: + logits, x = _token(layout, case, 100 * CASES.index(case) + seed, m) + y, ids, w, x_fp8, x_sf = _m1(layout, proc, bias, logits, x) + again, *_ = _m1(layout, proc, bias, logits, x) + y_k3 = _k3_moe(layout, proc, ids, w, x_fp8, x_sf) + torch.cuda.synchronize() + local = int(((ids >= offset) & (ids < offset + e_local)).sum()) + assert torch.equal(_bits(y), _bits(again)), "run-to-run bits differ" + if local == 0: + assert bool((y.float() == 0).all()), "no local expert: the partial must be zero" + continue + n_local += 1 + ref = _reference(layout, raw, _deq_x(x_fp8, x_sf), ids, w).bfloat16() + elt, rms = _row_ulp(y, ref) + elt_k3, _ = _row_ulp(y, y_k3) + same = torch.equal(_bits(y), _bits(y_k3)) + n_bit += same + print(f"OPCHECK op=k3_moe_m1 layout={layout} case={case} m={m} seed={seed} local={local} " + f"vs_ref_elt_ulp={elt:.2f} " + f"vs_ref_rms_ulp={rms:.2f} vs_k3_moe_elt_ulp={elt_k3:.2f} bit_identical_to_k3_moe={same}") # fmt: skip + assert bool(torch.isfinite(y.float()).all()) + assert elt <= 8.0 and rms <= 4.0 + assert elt_k3 <= 1.0 + print( + f"OPCHECK op=k3_moe_m1 layout={layout} case={case} m={m} bit_identical_to_k3_moe={n_bit}/{n_local} " + f"(calls with a local expert)" + ) + + +def test_k3_moe_m1_two_tokens_refuse_wide_intermediate(): + """The two-token build's FC2 tiles of both tokens do not fit at the TP4 x EP4 intermediate (768): the build refuses + it when its kernel module loads.""" + with pytest.raises(AssertionError, match="does not fit"): + _state("tp4ep4", 2) + + +def test_k3_moe_m1_layers_interleaved(): + """Two layers (different weights) sharing one state, called in turns: each call gives the bits of the same call + alone (the count slots and epochs advance with every call).""" + _ops() + proc, _, bias = _experts("tp16") + other = { + k: v.roll(1, dims=0).contiguous() for k, v in proc.items() + } # the experts of a second layer + tokens = [_token("tp16", "random", 900 + s) for s in range(3)] + alone = {} + for name, weights in (("a", proc), ("b", other)): + for t, (logits, x) in enumerate(tokens): + alone[name, t] = _m1("tp16", weights, bias, logits, x)[0] + seq = [] + for t, (logits, x) in enumerate(tokens): + for name, weights in (("b", other), ("a", proc)): + seq.append((name, t, _m1("tp16", weights, bias, logits, x)[0])) + torch.cuda.synchronize() + for name, t, y in seq: + assert torch.equal(_bits(y), _bits(alone[name, t])), (name, t) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_m2.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_m2.py new file mode 100644 index 000000000000..3cce0611a53c --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_m2.py @@ -0,0 +1,329 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""k3_moe_m2 (two decode tokens' routed experts as a weight-stream kernel) on one GPU, in the TP16 layout of the +routed experts (every rank holds all 896 experts, its 192-wide intermediate slice zero-padded to 256 by TRT-LLM's +loader). Weights are random checkpoint-format MXFP4 experts put through TRT-LLM's own loader. Per routing case, +against: +- trtllm::k3_moe (K3MoeLayer) on the same buffers and routing, which it reads at the + padded size. The zero rows add exact + zeros, so the two compute the same math; k3_moe_m2 keeps k3_moe's FC2 operand and combine order, so the bits match + except where its two FC1 partial sums round an intermediate value differently. Bit identity is reported; the gate + is one bf16 ulp of the row's max; +- the fp32 reference over the dequantized experts (op-catalog gates: 8 ulp of the row max per element, 4 ulp + relative RMS). +The routing cases cover the tokens' experts overlapping as random routing does, fully shared (16 experts) and fully +disjoint (32, the most FC1 units). It also checks run-to-run identical bits, that calls of different layers +sharing the state's workspace, in any order, each give the bits of the same call alone, and that the same holds +across the epochs' int32 wrap. +""" + +import functools +import math +from types import SimpleNamespace + +import pytest +import torch + + +def _is_sm100() -> bool: + if not torch.cuda.is_available(): + return False + return torch.cuda.get_device_capability() == (10, 0) + + +pytestmark = pytest.mark.skipif(not _is_sm100(), reason="k3_moe_m2 needs sm_100") + +H, NUM_EXPERTS, TOP_K, SV = 3584, 896, 16, 32 +GATE_CAP, LINEAR_CAP = ( + 4.0, + 25.0, +) # the SiTU caps (activation_situ_beta, activation_situ_linear_beta) +RSF = 2.827 +ULP = 2.0**-8 +E4M3_MAX = 448.0 +# Layout -> (moe_tp, the rank's logical intermediate, local experts, first local id). +LAYOUTS = {"tp16": (16, 192, 896, 0), "tp4ep4": (4, 768, 224, 224)} +_E2M1 = [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0, -0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0] + + +def _ops(): + import tensorrt_llm # noqa: F401 + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op # noqa: F401 + from tensorrt_llm._torch.cute_dsl_kernels.k3_route_quant import op as _rq # noqa: F401 + + return torch.ops.trtllm + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def _rand_mxfp4(rows, k, k_full, gen): + """Random checkpoint-format MXFP4: packed [rows, k / 2] (low nibble = even k), E8M0 per 32 k, scaled so a + k_full-long dot product lands near std 3.""" + codes = torch.randint(0, 16, (rows, k), dtype=torch.uint8, device="cuda", generator=gen) + base = 127 + round(0.5 * math.log2(0.01057 / k_full)) + exps = torch.randint( + base, base + 6, (rows, k // SV), dtype=torch.uint8, device="cuda", generator=gen + ) + return (codes[:, 0::2] | (codes[:, 1::2] << 4)).contiguous(), exps + + +@functools.lru_cache(maxsize=None) +def _experts(layout: str, seed: int = 20260928): + """A rank's experts through W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod's loader (the buffers both kernels read, padded + as the loader pads an unaligned shard) and the rank's logical slices of the checkpoint tensors (what the reference + reads). Only the rank's shard is generated: an unaligned shard as rank 0 of tensors that hold exactly that shard + (the loader slices it, then pads), an aligned one as the whole tensor of a single rank.""" + from tensorrt_llm._torch.moe.fused_moe.quantization import W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod + + moe_tp, i_tp, e_local, _ = LAYOUTS[layout] + i_pad = (i_tp + 127) // 128 * 128 + tp = moe_tp if i_pad != i_tp else 1 + method = W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod() + module = SimpleNamespace(tp_size=tp, tp_rank=0, scaling_vector_size=SV, intermediate_size=i_tp * tp, + intermediate_size_per_partition=i_tp, hidden_size=H) # fmt: skip + kw = dict(dtype=torch.uint8, device="cuda") + proc = dict( + w31=torch.zeros(e_local, 2 * i_pad, H // 2, **kw), + w31s=torch.zeros(e_local, 2 * i_pad, H // SV, **kw), + w2=torch.zeros(e_local, H, i_pad // 2, **kw), + w2s=torch.zeros(e_local, H, i_pad // SV, **kw), + ) + raw = {name: [] for name in ("up", "up_s", "gate", "gate_s", "down", "down_s")} + gen = torch.Generator(device="cuda").manual_seed(seed) + for e in range(e_local): + w1, w1s = _rand_mxfp4(i_tp, H, H, gen) # gate + w3, w3s = _rand_mxfp4(i_tp, H, H, gen) # up + w2, w2s = _rand_mxfp4(H, i_tp, i_tp * moe_tp, gen) # down + method.load_expert_w3_w1_weight(module, w1, w3, proc["w31"][e]) + method.load_expert_w2_weight(module, w2, proc["w2"][e]) + method.load_expert_w3_w1_weight_scale_mxfp4(module, w1s, w3s, proc["w31s"][e]) + method.load_expert_w2_weight_scale_mxfp4(module, w2s, proc["w2s"][e]) + for name, t in zip(raw, (w3, w3s, w1, w1s, w2, w2s)): + raw[name].append(t) + torch.cuda.synchronize() + gen_b = torch.Generator(device="cuda").manual_seed(seed + 1) + bias = (torch.randn(NUM_EXPERTS, generator=gen_b, device="cuda") * 0.05).float() + return proc, raw, bias + + +def _deq_w(packed, sf): + lut = torch.tensor(_E2M1, device=packed.device) + vals = torch.empty(packed.shape[0], packed.shape[1] * 2, device=packed.device) + vals[:, 0::2] = lut[(packed & 0xF).long()] + vals[:, 1::2] = lut[(packed >> 4).long()] + return vals * torch.exp2(sf.float() - 127.0).repeat_interleave(SV, dim=1) + + +def _deq_x(x_fp8, x_sf): + rows, k = x_fp8.shape + return x_fp8.float() * torch.exp2( + x_sf.reshape(rows, k // SV).float() - 127.0 + ).repeat_interleave(SV, dim=1) + + +def _requant(act): + """The FC1 epilogue's MXFP8 requantization per 32 columns (round-up scale), dequantized.""" + rows, cols = act.shape + blocks = act.reshape(rows, cols // SV, SV) + amax = blocks.abs().amax(dim=-1, keepdim=True) + ex = torch.ceil(torch.log2(amax / E4M3_MAX)) + ex = torch.where(amax == 0, torch.full_like(amax, -127.0), ex).clamp(-127.0, 127.0) + scale = torch.exp2(ex) + q8 = (blocks / scale).clamp(-E4M3_MAX, E4M3_MAX).to(torch.float8_e4m3fn) + return (q8.float() * scale).reshape(rows, cols) + + +def _reference(layout, raw, x_deq, ids, weights): + """fp32 routed MoE over this rank's experts from the checkpoint slices (f64 GEMMs): SiTU, the MXFP8 + intermediate, the down projection, the routing-weighted sum, per token.""" + _, _, e_local, offset = LAYOUTS[layout] + out = torch.zeros(ids.shape[0], H, device="cuda") + for t in range(ids.shape[0]): + for k in range(TOP_K): + e = int(ids[t, k]) - offset + if not 0 <= e < e_local: + continue + xe = x_deq[t : t + 1].double() + up = (xe @ _deq_w(raw["up"][e], raw["up_s"][e]).double().t()).float() + gate = (xe @ _deq_w(raw["gate"][e], raw["gate_s"][e]).double().t()).float() + act = (GATE_CAP * torch.tanh(gate / GATE_CAP) * torch.sigmoid(gate) + * (LINEAR_CAP * torch.tanh(up / LINEAR_CAP))) # fmt: skip + y = ( + _requant(act).double() @ _deq_w(raw["down"][e], raw["down_s"][e]).double().t() + ).float() + out[t] += (y * weights[t, k].float())[0] + return out + + +def _row_ulp(y, ref): + """Largest |y - ref| in bf16 ulps of the row's max |ref|, and the relative RMS in ulps.""" + o, r = y.float(), ref.float() + row = r.abs().amax(dim=1, keepdim=True).clamp_min(1e-12) + elt = ((o - r).abs() / row).max().item() / ULP + rms = ((o - r).pow(2).mean().sqrt() / r.pow(2).mean().sqrt().clamp_min(1e-12)).item() / ULP + return elt, rms + + +@functools.lru_cache(maxsize=None) +def _state(layout: str): + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op + + _, i_tp, e_local, _ = LAYOUTS[layout] + return op.K3MoeM2State( + torch.device("cuda", torch.cuda.current_device()), i_tp, (i_tp + 127) // 128 * 128, e_local + ) + + +@functools.lru_cache(maxsize=None) +def _k3_moe_state(layout: str): + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op + + _, i_tp, e_local, _ = LAYOUTS[layout] + return op.K3MoeState( + torch.device("cuda", torch.cuda.current_device()), (i_tp + 127) // 128 * 128, e_local + ) + + +def _k3_moe(layout, proc, ids, w, x_fp8, x_sf): + """trtllm::k3_moe (K3MoeLayer) on the same buffers and routing.""" + _, _, _, offset = LAYOUTS[layout] + layer = _k3_moe_state(layout).layer(proc["w31"], proc["w31s"], proc["w2"], proc["w2s"]) + return layer(x_fp8, x_sf, ids, w, offset) + + +CASES = ["random", "shared", "disjoint"] + + +def _tokens(case: str, seed: int): + """Two tokens' router logits and latent rows. "shared": both tokens route to the same 16 experts; "disjoint": to + 32 different experts.""" + gen = torch.Generator(device="cuda").manual_seed(seed) + logits = (torch.randn(2, NUM_EXPERTS, generator=gen, device="cuda") * 3.0).float() + if case == "shared": + logits[1] = logits[0] + elif case == "disjoint": + top0 = torch.topk(logits[0], TOP_K).indices + logits[1, top0] = -30.0 + x = torch.randn(2, H, generator=gen, device="cuda").bfloat16() + return logits, x + + +def _m2(layout, proc, bias, logits, x): + ops = _ops() + _, _, _, offset = LAYOUTS[layout] + ids, w, x_fp8, x_sf = ops.k3_route_quant(logits, bias, x, RSF, True) + layer = _state(layout).layer(proc["w31"], proc["w31s"], proc["w2"], proc["w2s"]) + return layer(x_fp8, x_sf, ids, w, offset), ids, w, x_fp8, x_sf + + +@pytest.mark.parametrize("case", CASES) +def test_k3_moe_m2(case): + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op + + _ops() + layout = "tp16" # k3_moe_m2 is for the TP16 slice (test_k3_moe_m2_refuses_wide_intermediate) + moe_tp, i_tp, e_local, offset = LAYOUTS[layout] + proc, raw, bias = _experts(layout) + ok, why = op.m2_supported(proc["w31"], proc["w31s"], proc["w2"], proc["w2s"], e_local, i_tp) + assert ok, why + n_bit = 0 + seeds = range(8) + for seed in seeds: + logits, x = _tokens(case, 100 * CASES.index(case) + seed) + y, ids, w, x_fp8, x_sf = _m2(layout, proc, bias, logits, x) + again, *_ = _m2(layout, proc, bias, logits, x) + y_k3 = _k3_moe(layout, proc, ids, w, x_fp8, x_sf) + torch.cuda.synchronize() + experts = len(set(ids.view(-1).tolist())) + assert torch.equal(_bits(y), _bits(again)), "run-to-run bits differ" + ref = _reference(layout, raw, _deq_x(x_fp8, x_sf), ids, w).bfloat16() + elt, rms = _row_ulp(y, ref) + elt_k3, _ = _row_ulp(y, y_k3) + same = torch.equal(_bits(y), _bits(y_k3)) + n_bit += same + print(f"OPCHECK op=k3_moe_m2 layout={layout} case={case} seed={seed} experts={experts} " + f"vs_ref_elt_ulp={elt:.2f} " + f"vs_ref_rms_ulp={rms:.2f} vs_k3_moe_elt_ulp={elt_k3:.2f} bit_identical_to_k3_moe={same}") # fmt: skip + assert bool(torch.isfinite(y.float()).all()) + assert elt <= 8.0 and rms <= 4.0 + assert elt_k3 <= 1.0 + print( + f"OPCHECK op=k3_moe_m2 layout={layout} case={case} bit_identical_to_k3_moe={n_bit}/{len(seeds)}" + ) + + +def test_k3_moe_m2_refuses_wide_intermediate(): + """k3_moe_m2 is for an intermediate of at most 256 (its FC2 tiles of every expert slot stay in shared memory): + m2_supported refuses the TP4 x EP4 slice (768). Only metadata is read.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op + + _, i_tp, e_local, _ = LAYOUTS["tp4ep4"] + kw = dict(dtype=torch.uint8, device="meta") + ok, why = op.m2_supported( + torch.empty(e_local, 2 * i_tp, H // 2, **kw), + torch.empty(e_local, 2 * i_tp, H // SV, **kw), + torch.empty(e_local, H, i_tp // 2, **kw), + torch.empty(e_local, H, i_tp // SV, **kw), + e_local, + i_tp, + ) + assert not ok and "256" in why, why + + +def test_k3_moe_m2_layers_interleaved(): + """Two layers (different weights) sharing one state, called in turns: each call gives the bits of the same call + alone (the counter sets and epochs advance with every call).""" + _ops() + proc, _, bias = _experts("tp16") + other = { + k: v.roll(1, dims=0).contiguous() for k, v in proc.items() + } # the experts of a second layer + tokens = [_tokens("random", 900 + s) for s in range(3)] + alone = {} + for name, weights in (("a", proc), ("b", other)): + for t, (logits, x) in enumerate(tokens): + alone[name, t] = _m2("tp16", weights, bias, logits, x)[0] + seq = [] + for t, (logits, x) in enumerate(tokens): + for name, weights in (("b", other), ("a", proc)): + seq.append((name, t, _m2("tp16", weights, bias, logits, x)[0])) + torch.cuda.synchronize() + for name, t, y in seq: + assert torch.equal(_bits(y), _bits(alone[name, t])), (name, t) + + +@pytest.mark.parametrize("start", [2**31 - 2, 2**31 - 1]) +def test_k3_moe_m2_epoch_wrap(start): + """The CTAs' epochs (int32, + 1 per call; only their parity picks the counter set) across the int32 wrap. The + state is preset to where ~2^31 calls leave it: every epoch at ``start`` and the counter set the next call uses + zero. Each call then gives the bits of the same call at small epochs, the epochs wrap to -2^31 and keep counting, + and every call leaves the set the next one uses zero.""" + _ops() + proc, _, bias = _experts("tp16") + st = _state("tp16") + words = st.mod.GROUPS2 * st.mod.CW + tokens = [_tokens("random", 700 + c) for c in range(5)] + alone = [_m2("tp16", proc, bias, logits, x)[0] for logits, x in tokens] + st.counts.zero_() + st.epochs.fill_(start) + for c, (logits, x) in enumerate(tokens): + y = _m2("tp16", proc, bias, logits, x)[0] + torch.cuda.synchronize() + ep = (start + c + 1 + 2**31) % 2**32 - 2**31 # int32 two's complement + assert torch.equal(_bits(y), _bits(alone[c])), c + assert bool((st.epochs == ep).all()), (c, ep) + assert bool((st.counts.view(2, words)[ep & 1] == 0).all()), (c, ep) From 45cbb249feebb00e7e85b1bf7e9322dbcb586218 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 21:32:09 -0700 Subject: [PATCH 092/161] [None][feat] Kimi K3 MoE: k3_moe's push build into the latent exchange trtllm::k3_moe takes a K3LatentExchange's words (this rank's and the multicast mapping) and flags, and a slot. With them it is the push build, the fused all-reduce's push-only mode: the routed partial goes into that slot of every rank's exchange instead of a returned tensor, for trtllm::k3_latent_reduce, and the call returns an empty tensor. mutates_args gains the exchange's words. K3MoeLayer.push makes the call after trtllm::k3_route_quant or trtllm::k3_moe_front. Signed-off-by: Vasanth Sabavat --- .../cute_dsl_kernels/k3_fused_moe/op.py | 119 +++++++++++++++--- 1 file changed, 102 insertions(+), 17 deletions(-) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py index 60e151a6b98e..2bf056c190a7 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py @@ -25,7 +25,9 @@ unchanged. Two builds: up to 8 tokens (:class:`K3MoeState`, optionally acquiring the front's outputs through the head -workspace's ready words) and up to 64 (:class:`K3MoeWideState`). The caller owns all state: the scratch a state's +workspace's ready words) and up to 64 (:class:`K3MoeWideState`). Given a ``K3LatentExchange``, the 8-token build +pushes the partial into every rank's exchange for ``trtllm::k3_latent_reduce`` instead of returning it +(:meth:`K3MoeLayer.push`). The caller owns all state: the scratch a state's layers share (the intermediate slab, left armed by every call, and the FC2 partial rows), each layer's counters (:class:`K3MoeLayer`, left zero by every call), and the head all-gather buffers of the front (:class:`K3MoeHeadWorkspace`, collective over the TP group). Build them before CUDA-graph capture; each build @@ -224,7 +226,8 @@ def build(uc, mc, handle, comm): # The kernel's tensor arguments after the 18 it always reads: the fused all-reduce's buffers (3), the fold's inputs # (5), the head flags build's ready words and head flags (2), the latent slab (1). The builds this module compiles -# read only the ready words and head flags (head_flags), so the others are given a stand-in they never touch. +# read only the ready words and head flags (head_flags) and the fused all-reduce's buffers in its push-only mode (the +# push build), so the others are given a stand-in they never touch. _OPTIONAL_ARGS = 11 # int32 words of one slot of a K3LatentExchange: two halves of 8 tokens' bf16 [3584] rows. _EXCHANGE_SLOT_WORDS = 2 * _TOKEN_SLOTS * (HIDDEN_SIZE // 2) @@ -243,10 +246,17 @@ def _part_rows(mod) -> int: def _config( - i_tp: int, num_ctas: int, num_local: int, m_max: int, use_pdl: bool, head_flags: bool + i_tp: int, + num_ctas: int, + num_local: int, + m_max: int, + use_pdl: bool, + head_flags: bool, + push_world: int = 0, ) -> dict: - """The kernel options of one build (trace-time constants).""" - return { + """The kernel options of one build (trace-time constants). ``push_world``: the push build for a latent exchange of + that many slots (the fused all-reduce's push-only mode).""" + config = { "i_tp": i_tp, "num_ctas": num_ctas, "num_local": num_local, @@ -255,6 +265,9 @@ def _config( "head_flags": int(head_flags), "lat_slab": 0, } + if push_world: + config.update(ar_world=push_world, ar_push_only=1) + return config class _K3MoeScratch: @@ -399,6 +412,33 @@ def __call__( ) # fmt: skip return y if out is None else out[: topk_ids.shape[0]] + def push( + self, + x_fp8: torch.Tensor, + x_sf: torch.Tensor, + topk_ids: torch.Tensor, + topk_weights: torch.Tensor, + local_expert_offset: int, + exchange, + slot: Optional[int] = None, + head_ready: Optional[torch.Tensor] = None, + head_flags: Optional[torch.Tensor] = None, + ) -> None: + """``trtllm::k3_moe``'s push build on this layer: the routed partial goes into slot ``slot`` (default + ``exchange.rank``) of every rank's ``exchange`` (a ``K3LatentExchange``) instead of a returned tensor, for + ``trtllm::k3_latent_reduce``. See the op. ``head_ready`` and ``head_flags`` are given for, and only for, the + layers of a ``head_flags`` state.""" + st = self.state + if (head_ready is not None) != st.head_flags: + raise ValueError( + "K3MoeLayer: head_ready / head_flags are given for, and only for, a head_flags state's layers" + ) + torch.ops.trtllm.k3_moe( + x_fp8, x_sf, topk_ids, topk_weights, *self.weights, st.c, st.cs, st.part, self.counters, + local_expert_offset, st.num_local, st.num_ctas, st.m_max, st.use_pdl, head_ready, head_flags, None, + exchange.uc, exchange.mc, exchange.flags, exchange.rank if slot is None else slot, + ) # fmt: skip + def _compile_key(device: torch.device, config: dict) -> tuple: index = device.index if device.index is not None else torch.cuda.current_device() @@ -422,8 +462,10 @@ def _compile(mod, args, scalars): @torch.library.custom_op( "trtllm::k3_moe", - mutates_args=("c", "cs", "part", "counters", "head_ready", "head_flags", "out"), -) + mutates_args=( + "c", "cs", "part", "counters", "head_ready", "head_flags", "out", "exchange_uc", "exchange_mc", + ), +) # fmt: skip def k3_moe( x_fp8: torch.Tensor, x_sf: torch.Tensor, @@ -445,6 +487,10 @@ def k3_moe( head_ready: Optional[torch.Tensor] = None, head_flags: Optional[torch.Tensor] = None, out: Optional[torch.Tensor] = None, + exchange_uc: Optional[torch.Tensor] = None, + exchange_mc: Optional[torch.Tensor] = None, + exchange_flags: Optional[torch.Tensor] = None, + exchange_slot: int = 0, ) -> torch.Tensor: """This rank's routed partial ``[M, 3584]`` bf16 from the persistent ``k3_moe`` kernel: FC1 + SiTU + FC2 with the routing-weighted, deterministic combine over this rank's experts. @@ -457,11 +503,43 @@ def k3_moe( build (``num_ctas``, ``use_pdl``); every call leaves the slab armed. ``counters``: the layer's, left zero. ``head_ready`` / ``head_flags``: a ``head_flags`` build's ready words and head flags (a ``K3MoeHeadWorkspace``'s ``ready`` and ``flags``), whose epoch the call advances. ``out``: bf16, contiguous, at least ``[M, 3584]``; the - call writes its first M rows and returns an empty ``[0, 3584]`` instead of a new tensor. 1 <= M <= ``m_max``.""" + call writes its first M rows and returns an empty ``[0, 3584]`` instead of a new tensor. 1 <= M <= ``m_max``. + + ``exchange_uc`` / ``exchange_mc`` / ``exchange_flags``: a ``K3LatentExchange``'s words (this rank's, and the same + words through the multicast mapping) and flags. With them the call is the push build (``m_max`` 8, no ``out``): + the partial goes into slot ``exchange_slot`` of half ``exchange_flags[0] & 1`` of every rank's exchange through + the multicast mapping (bf16 pairs, -0.0 stored as +0.0, zero rows when nothing is routed here) and the call returns + an empty ``[0, 3584]``. Nothing reduces the slots and the flags are not written: ``trtllm::k3_latent_reduce`` sums + the slots, empties the half it read and advances the flags. Each push of M tokens is followed by one reduce of M + tokens on that exchange before the next push.""" num_tokens = topk_ids.shape[0] head = head_ready is not None if head != (head_flags is not None): raise ValueError("k3_moe: head_ready and head_flags go together") + push = exchange_uc is not None + if push != (exchange_mc is not None) or push != (exchange_flags is not None): + raise ValueError("k3_moe: exchange_uc, exchange_mc and exchange_flags go together") + push_world = 0 + if push: + push_world = exchange_uc.numel() // _EXCHANGE_SLOT_WORDS + if ( + exchange_uc.dtype != torch.int32 + or exchange_mc.dtype != torch.int32 + or push_world == 0 + or exchange_uc.numel() != push_world * _EXCHANGE_SLOT_WORDS + or exchange_mc.numel() != exchange_uc.numel() + or exchange_flags.dtype != torch.int32 + or exchange_flags.numel() < 1 + ): + raise ValueError( + "k3_moe: the exchange is int32 words [2][8][slots][1792] (this rank's and multicast) and int32 flags" + ) + if not 0 <= exchange_slot < push_world: + raise ValueError( + f"k3_moe: slot {exchange_slot} outside the exchange's {push_world} slots" + ) + if m_max != MAX_TOKENS or out is not None: + raise ValueError("k3_moe: the push build is the m_max 8 build and takes no out") if m_max not in (MAX_TOKENS, WIDE_MAX_TOKENS) or (head and m_max != MAX_TOKENS): raise ValueError( f"k3_moe: m_max is {MAX_TOKENS} (head flags possible) or {WIDE_MAX_TOKENS}, got {m_max}" @@ -485,7 +563,7 @@ def k3_moe( raise ValueError(f"k3_moe: {why}") e, two_i, _ = w3_w1_weight.shape i_tp = two_i // 2 - config = _config(i_tp, num_ctas, num_local, m_max, use_pdl, head) + config = _config(i_tp, num_ctas, num_local, m_max, use_pdl, head, push_world) mod = _kernel_module(config) g_cap = mod.G_CAP if ( @@ -528,10 +606,15 @@ def k3_moe( .reshape(g_cap, _TOKEN_SLOTS, mod.K2_TILES, mod.SFB_GROUP_BYTES) .permute(3, 2, 1, 0) ) - # Stand-in for the options this build does not have (fused all-reduce, fold, latent slab, and the head flags - # without head_flags): the kernel never touches it. + # Stand-in for the options this build does not have (fused all-reduce without a push, fold, latent slab, and the + # head flags without head_flags): the kernel never touches it. The push build writes no y. unused = counters flag_args = (head_ready.view(-1), head_flags.view(-1)) if head else (unused, unused) + ar_args = ( + (exchange_uc.view(-1), exchange_mc.view(-1), exchange_flags.view(-1)) + if push + else (unused,) * 3 + ) args = ( w3_w1_weight.view(torch.int8).permute(2, 1, 0), x_fp8.view(torch.uint8).permute(1, 0), @@ -542,13 +625,14 @@ def k3_moe( c.view(torch.uint8).permute(2, 1, 0), w2_weight_scale.view(e, HIDDEN_SIZE // 128, i_tp // 128, 512).permute(3, 2, 1, 0), sfb2, y, y.view(torch.int32), part, topk_ids, topk_weights, counters, - unused, unused, unused, # the fused all-reduce's buffers + *ar_args, # the fused all-reduce's buffers: the exchange's words and flags (push build) unused, unused, unused, unused, unused, # the fold's inputs *flag_args, # the ready words and the head flags unused, # the latent slab ) # fmt: skip - # (tokens, local offset, local experts, all-reduce rank, routed scaling factor (fold only), slab buffer, re-arm 0) - scalars = (num_tokens, local_expert_offset, num_local, 0, 1.0, 0, 0) + # (tokens, local offset, local experts, all-reduce rank (the push's slot), routed scaling factor (fold only), slab + # buffer, re-arm 0) + scalars = (num_tokens, local_expert_offset, num_local, exchange_slot if push else 0, 1.0, 0, 0) key = _compile_key(x_fp8.device, config) fn = _compiled.get(key) if fn is None: @@ -558,7 +642,7 @@ def k3_moe( ) fn = _compiled[key] = _compile(mod, args, scalars) fn(*args, *scalars, torch.cuda.current_stream(x_fp8.device).cuda_stream) - if out is not None: + if out is not None or push: return y.new_empty((0, HIDDEN_SIZE)) return y @@ -566,10 +650,11 @@ def k3_moe( @k3_moe.register_fake def _(x_fp8, x_sf, topk_ids, topk_weights, w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale, c, cs, part, counters, local_expert_offset, num_local, num_ctas, m_max, use_pdl, head_ready=None, head_flags=None, - out=None): # fmt: skip - rows = 0 if out is not None else topk_ids.shape[0] + out=None, exchange_uc=None, exchange_mc=None, exchange_flags=None, exchange_slot=0): # fmt: skip + rows = 0 if out is not None or exchange_uc is not None else topk_ids.shape[0] return x_fp8.new_empty((rows, HIDDEN_SIZE), dtype=torch.bfloat16) + # --------------------------------------------------------------------------------------------------------------------- # trtllm::k3_moe_m1 / trtllm::k3_moe_m2: the routed experts of one or two decode tokens as weight-stream kernels # (k3_moe_m1_kernel.py, k3_moe_m2_kernel.py), on a caller-owned workspace. From 81a0235aa8060dc7be67b3ff0d334b692c3344bb Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 21:32:25 -0700 Subject: [PATCH 093/161] [None][test] Kimi K3 MoE: the push builds at 4 ranks against MNNVLAllReduce test_k3_moe_push.py, one process per GPU over the run's TP group (4 on one GB200 tray), every rank with its own TP16 experts and the same tokens: - push: trtllm::k3_moe_m1 at 1 and 2 tokens and trtllm::k3_moe_m2 at 2; - k3_moe_push: trtllm::k3_moe's push build at 1, 3 and 8 tokens, after trtllm::k3_route_quant and after trtllm::k3_moe_front. Per token count and routing set, the pushed partials' reduce equals MNNVLAllReduce's one-shot of the plain partials bit for bit, on every rank and run to run, and leaves the exchange empty with its call count advanced; the same into a 16-slot exchange filled 4 slots per rank (the order of TP16's receive side); and all of it again across the int32 wrap of the call count. Listed in l0_gb200_multi_gpus. Signed-off-by: Vasanth Sabavat --- .../test-db/l0_gb200_multi_gpus.yml | 1 + .../kimi_k3/test_k3_moe_push.py | 412 ++++++++++++++++++ 2 files changed, 413 insertions(+) create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py diff --git a/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml b/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml index 4477c83fd8c7..02a9429501c4 100644 --- a/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml +++ b/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml @@ -57,6 +57,7 @@ l0_gb200_multi_gpus: - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_latent_reduce.py + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py - unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_sandwich_oproj_op_matrix.py - unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_sandwich_tail_op_matrix.py - unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_sandwich_plain_op_matrix.py diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py new file mode 100644 index 000000000000..1f6625d1effc --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py @@ -0,0 +1,412 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""The push builds of the Kimi K3 decode routed experts: each rank's routed partial stored into a latent exchange for +trtllm::k3_latent_reduce instead of returned. One process per GPU over the TP group of this run (4 on one GB200 tray). +Every rank holds its own TP16 experts (all 896, the rank's 192-wide intermediate slice zero-padded to 256 by TRT-LLM's +loader; random checkpoint-format MXFP4, a seed per rank) and runs the same tokens. + +push trtllm::k3_moe_m1 at 1 and 2 tokens and trtllm::k3_moe_m2 at 2 (Layer.push), on tokens routed by + trtllm::k3_route_quant; +k3_moe_push trtllm::k3_moe's push build (K3MoeLayer.push) at 1, 3 and 8 tokens, after trtllm::k3_route_quant and + after trtllm::k3_moe_front. The front's head and shared experts are sharded over the run's ranks; its + shared activation must be the same bits in every call of a set. + +Per op, token count and routing set, against the plain call (Layer.__call__, K3MoeLayer.__call__ after the same +producer), whose partial must be nonzero: + exact the partial pushed into the group's exchange (one slot per rank): the reduce equals MNNVLAllReduce's + one-shot of the plain partials bit for bit, on every rank and run to run (two more push + reduce pairs); + afterwards the exchange is empty, its call count is +1 and the arrival word 0; + exact16 the same into a 16-slot exchange, every rank's partial in 4 slots (TP16's receive side on one tray): the + reduce equals the one-shot's order over 16 slots (fp32 sums of 8 slots, added in order, then bf16); + wrap all of it again with both exchanges' call counts starting at 2^31 - 2, so the int32 count wraps to -2^31 + during the sets and the halves keep alternating. + +Run under pytest (a pool of 4 MPI workers) or directly, one process per GPU: + srun -N1 -n4 --mpi=pmix python3 test_k3_moe_push.py [push k3_moe_push] +""" + +import functools +import hashlib +import math +import os +import pickle +import sys +import traceback +from types import SimpleNamespace + +import pytest +import torch + +try: + import cloudpickle + from mpi4py import MPI +except ImportError: # the test is skipped below + cloudpickle = MPI = None + +if cloudpickle is not None: + cloudpickle.register_pickle_by_value(sys.modules[__name__]) + MPI.pickle.__init__(cloudpickle.dumps, cloudpickle.loads, pickle.HIGHEST_PROTOCOL) + +WORLD = 4 +H, NUM_EXPERTS, SV = 3584, 896, 32 +# TP16: a rank's 192-wide intermediate slice, zero-padded to whole tiles (256) by the loader. +I_TP, I_PAD, MOE_TP = 192, 256, 16 +HIDDEN, SHARED_INTER = 7168, 6144 # the front's input width; two shared experts of 3072 +GATE_CAP, LINEAR_CAP = 4.0, 25.0 # the SiTU caps +RSF = 2.827 +SETS = 4 +COUNT_STARTS = (0, 2**31 - 2) +EMPTY_WORD = -(2**31) +ENGINES = (("k3_moe_m1", 1), ("k3_moe_m1", 2), ("k3_moe_m2", 2)) # (engine, tokens per call) +# What routes and quantizes K3MoeLayer.push's tokens. +K3_MOE_PRODUCERS = ("k3_route_quant", "k3_moe_front") +K3_MOE_TOKENS = (1, 3, 8) +K3_MOE_SETS = 2 + + +def _supported() -> bool: + if MPI is None or not torch.cuda.is_available() or torch.cuda.device_count() < WORLD: + return False + return torch.cuda.get_device_capability() == (10, 0) + + +pytestmark = [ + pytest.mark.threadleak(enabled=False), + pytest.mark.skipif(not _supported(), reason=f"needs {WORLD} SM100 GPUs with MNNVL and mpi4py"), +] + + +def _bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def _same(a: torch.Tensor, b: torch.Tensor) -> bool: + return a.shape == b.shape and a.dtype == b.dtype and torch.equal(_bits(a), _bits(b)) + + +def _digest(t: torch.Tensor) -> str: + return hashlib.sha256(_bits(t).cpu().numpy().tobytes()).hexdigest() + + +def _rand_mxfp4(rows, k, k_full, gen): + """Random checkpoint-format MXFP4: packed [rows, k / 2] (low nibble = even k), E8M0 per 32 k, scaled so a + k_full-long dot product lands near std 3.""" + codes = torch.randint(0, 16, (rows, k), dtype=torch.uint8, device="cuda", generator=gen) + base = 127 + round(0.5 * math.log2(0.01057 / k_full)) + exps = torch.randint( + base, base + 6, (rows, k // SV), dtype=torch.uint8, device="cuda", generator=gen + ) + return (codes[:, 0::2] | (codes[:, 1::2] << 4)).contiguous(), exps + + +@functools.lru_cache(maxsize=None) +def _experts(seed: int): + """This rank's TP16 experts through W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod's loader: the buffers the engines read, + the 192-wide shard generated as rank 0 of tensors that hold exactly it (the loader slices it, then pads it to 256). + """ + from tensorrt_llm._torch.moe.fused_moe.quantization import W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod + + method = W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod() + module = SimpleNamespace(tp_size=MOE_TP, tp_rank=0, scaling_vector_size=SV, intermediate_size=I_TP * MOE_TP, + intermediate_size_per_partition=I_TP, hidden_size=H) # fmt: skip + kw = dict(dtype=torch.uint8, device="cuda") + w31 = torch.zeros(NUM_EXPERTS, 2 * I_PAD, H // 2, **kw) + w31s = torch.zeros(NUM_EXPERTS, 2 * I_PAD, H // SV, **kw) + w2 = torch.zeros(NUM_EXPERTS, H, I_PAD // 2, **kw) + w2s = torch.zeros(NUM_EXPERTS, H, I_PAD // SV, **kw) + gen = torch.Generator(device="cuda").manual_seed(seed) + for e in range(NUM_EXPERTS): + gate, gate_s = _rand_mxfp4(I_TP, H, H, gen) + up, up_s = _rand_mxfp4(I_TP, H, H, gen) + down, down_s = _rand_mxfp4(H, I_TP, I_TP * MOE_TP, gen) + method.load_expert_w3_w1_weight(module, gate, up, w31[e]) + method.load_expert_w2_weight(module, down, w2[e]) + method.load_expert_w3_w1_weight_scale_mxfp4(module, gate_s, up_s, w31s[e]) + method.load_expert_w2_weight_scale_mxfp4(module, down_s, w2s[e]) + torch.cuda.synchronize() + return w31, w31s, w2, w2s + + +def _tokens(m: int, seed: int): + """The same tokens on every rank: hidden rows, fp32 router logits and the routing bias, quantized and routed by + trtllm::k3_route_quant as the engines' callers do.""" + gen = torch.Generator(device="cuda").manual_seed(seed) + x = torch.randn(m, H, generator=gen, device="cuda").bfloat16() + logits = torch.randn(m, NUM_EXPERTS, generator=gen, device="cuda") + bias = (torch.randn(NUM_EXPERTS, generator=gen, device="cuda") * 0.05).float() + ids, weights, x_fp8, x_sf = torch.ops.trtllm.k3_route_quant(logits, bias, x, RSF, True) + return x_fp8, x_sf, ids, weights + + +class _Exchange: + """A latent exchange over this run's ranks with ``slots`` slots: K3LatentExchange's buffers (int32 words + [2][8][slots][1792] behind one multicast mapping, every word 0x80000000; int32 flags[4]) at any slot count, its + call count starting at ``count``.""" + + def __init__(self, ctx, slots: int, count: int): + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import k3_latent_reduce as kernel + from tensorrt_llm._torch.distributed.ops import ( + _get_mnnvl_workspace_comm, + _make_mnnvl_mcast_buffer, + ) + + words = kernel.buffer_words(slots) + fabric = os.environ.get("TRTLLM_FORCE_MNNVL_AR", "0") == "1" or ctx.mapping.is_multi_node() + self.handle = _make_mnnvl_mcast_buffer( + _get_mnnvl_workspace_comm(ctx.mapping), words * 4, ctx.mapping, fabric + ) + self.uc = self.handle.get_uc_buffer(ctx.rank, (words,), torch.int32, 0) + self.mc = self.handle.get_mc_buffer((words,), torch.int32, 0) + self.uc.fill_(EMPTY_WORD) + self.flags = torch.zeros(4, dtype=torch.int32, device="cuda") + self.flags[0] = count + self.slots, self.count, self.rank = slots, count, ctx.rank + torch.cuda.synchronize() + ctx.comm.Barrier() + + def reduce(self, m: int) -> torch.Tensor: + out = torch.ops.trtllm.k3_latent_reduce(self.uc, self.flags, m, 0) + self.count = (self.count + 1 + 2**31) % 2**32 - 2**31 # int32 two's complement + return out + + def state_ok(self, ctx) -> bool: + """Every rank's last reduce done and nothing of the next call pushed yet: the whole buffer is empty, the call + count advanced, the arrival word cleared.""" + torch.cuda.synchronize() + ctx.comm.Barrier() + flags = self.flags.tolist() + ok = bool((self.uc == EMPTY_WORD).all().item()) and flags[0] == self.count and flags[2] == 0 + ctx.comm.Barrier() + return ok + + +def _order16(rows, copies): + """The one-shot's order over len(rows) x copies slots (slot s holds rank s // copies's row): fp32 sums of 8 slots + from slot 0, the sums added in order, then bf16 (round to nearest even).""" + slots = [rows[s // copies].float() for s in range(len(rows) * copies)] + total = torch.zeros_like(slots[0]) + for first in range(0, len(slots), 8): + chunk = torch.zeros_like(slots[0]) + for s in slots[first : first + 8]: + chunk = chunk + s + total = total + chunk + return total.bfloat16() + + +def _context(): + os.environ.setdefault("TRTLLM_FORCE_MNNVL_AR", "1") + comm = MPI.COMM_WORLD + rank, world = comm.Get_rank(), comm.Get_size() + gpus = torch.cuda.device_count() + torch.cuda.set_device(rank % gpus) + import tensorrt_llm # noqa: F401 + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import ( + latent_op, # noqa: F401 (registers the reduce) + ) + from tensorrt_llm._torch.cute_dsl_kernels.k3_route_quant import op as _rq # noqa: F401 + from tensorrt_llm._torch.distributed.ops import MNNVLAllReduce + from tensorrt_llm.mapping import Mapping + + mapping = Mapping(world_size=world, rank=rank, gpus_per_node=gpus, tp_size=world) + return SimpleNamespace(comm=comm, rank=rank, world=world, mapping=mapping, + mnnvl=MNNVLAllReduce(mapping, torch.bfloat16)) # fmt: skip + + +def _allreduce(ctx, y: torch.Tensor) -> torch.Tensor: + """MNNVLAllReduce sent one-shot (the order the reduce reproduces).""" + from tensorrt_llm._torch.distributed import AllReduceParams + + return ctx.mnnvl( + y, AllReduceParams(), one_shot_max_bytes=y.numel() * ctx.world * y.element_size() + ) + + +def _row(ctx, op, case, m, y, ref, got, state, rows, copies16, got16, state16, **extra): + """One result row; ``ok`` when every check holds on every rank.""" + row = dict(op=op, case=case, M=m, nonzero=bool((y != 0).any().item()), exact=_same(got[0], ref), + det=all(_same(g, got[0]) for g in got[1:]), + ranks_agree=len(set(ctx.comm.allgather(_digest(got[0])))) == 1, state=state, + exact16=_same(got16, _order16(rows, copies16)), state16=state16, **extra) # fmt: skip + row["ok"] = all( + ctx.comm.allgather(all(v for k, v in row.items() if k not in ("op", "case", "M"))) + ) + return row + + +def _layer(name: str, m: int, weights): + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op + + device = torch.device("cuda") + if name == "k3_moe_m1": + state = op.K3MoeM1State(device, I_TP, I_PAD, NUM_EXPERTS, num_tokens=m) + else: + state = op.K3MoeM2State(device, I_TP, I_PAD, NUM_EXPERTS) + return state.layer(*weights) + + +def check_push(ctx): + results = [] + weights = _experts(20260928 + ctx.rank) + # Each state compiles its builds once. + layers = {(name, m): _layer(name, m, weights) for name, m in ENGINES} + copies16 = 16 // ctx.world + for start in COUNT_STARTS: + ex, ex16 = _Exchange(ctx, ctx.world, start), _Exchange(ctx, 16, start) + for (name, m), layer in layers.items(): + for si in range(SETS): + x_fp8, x_sf, ids, w = _tokens(m, 1000 + 10 * si + m) + y = layer(x_fp8, x_sf, ids, w, 0) + ref = _allreduce(ctx, y) + got = [] + for _ in range(3): + layer.push(x_fp8, x_sf, ids, w, 0, ex.mc, ex.flags, ctx.rank) + got.append(ex.reduce(m)) + state = ex.state_ok(ctx) + rows = [t.cuda() for t in ctx.comm.allgather(y.cpu())] + layer.push(x_fp8, x_sf, ids, w, 0, ex16.mc, ex16.flags, ctx.rank, copies16) + got16 = ex16.reduce(m) + state16 = ex16.state_ok(ctx) + results.append(_row(ctx, name, f"set{si}_count{start}", m, y, ref, got, state, rows, copies16, + got16, state16)) # fmt: skip + return results + + +def _front(ctx): + """The MoE front's weights on this rank: its head rows (latent-down, then router) and the shared experts' gate_up + rows, sharded over the run's ranks, and the head all-gather's workspace.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import front_op + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.op import K3MoeHeadWorkspace + + wl, we, inter = H // ctx.world, NUM_EXPERTS // ctx.world, SHARED_INTER // ctx.world + gen = torch.Generator(device="cuda").manual_seed(111 + ctx.rank) + head = (torch.randn(wl + we, HIDDEN, device="cuda", generator=gen) * 0.02).bfloat16() + head[wl:] *= 8.0 # router rows: logits of a few units + gate_up = (torch.randn(2 * inter, HIDDEN, device="cuda", generator=gen) * 0.02).bfloat16() + assert front_op.weight_supported( + ctx.world, inter, HIDDEN, torch.device("cuda", torch.cuda.current_device()) + ) + fabric = os.environ.get("TRTLLM_FORCE_MNNVL_AR", "0") == "1" or ctx.mapping.is_multi_node() + return SimpleNamespace(weight=front_op.front_weight(head, gate_up), inter=inter, + head=K3MoeHeadWorkspace.create(ctx.mapping, fabric_handle=fabric)) # fmt: skip + + +def _k3_moe_layer(weights): + """This rank's experts as a K3MoeLayer of the 8-token build; a call with an exchange takes its push build.""" + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.op import K3MoeState + + device = torch.device("cuda", torch.cuda.current_device()) + return K3MoeState(device, I_PAD, NUM_EXPERTS).layer(*weights) + + +def _k3_moe_inputs(producer, m, seed): + """The producer's inputs: the latent and router logits (k3_route_quant), or the MoE input (k3_moe_front), with the + routing bias.""" + gen = torch.Generator(device="cuda").manual_seed(seed) + bias = (torch.randn(NUM_EXPERTS, generator=gen, device="cuda") * 0.05).float() + if producer == "k3_route_quant": + x = torch.randn(m, H, generator=gen, device="cuda").bfloat16() + return x, torch.randn(m, NUM_EXPERTS, generator=gen, device="cuda"), bias + return torch.randn(m, HIDDEN, generator=gen, device="cuda").bfloat16(), bias + + +def _k3_moe_routed(producer, inputs, front): + """(MXFP8 rows, their scales, ids, weights, shared activation or None) from the producer, in K3MoeLayer's + argument order. k3_moe_front is collective over the run's ranks (the head all-gather on ``front.head``).""" + if producer == "k3_route_quant": + x, logits, bias = inputs + ids, weights, x_fp8, x_sf = torch.ops.trtllm.k3_route_quant(logits, bias, x, RSF, True) + return x_fp8, x_sf, ids, weights, None + x, bias = inputs + head = front.head + ids, weights, x_fp8, x_sf, shared = torch.ops.trtllm.k3_moe_front( + x, front.weight, bias, RSF, front.inter, GATE_CAP, LINEAR_CAP, head.uc, head.mc, head.flags, head.rank, + head.world_size, + ) # fmt: skip + return x_fp8, x_sf, ids, weights, shared + + +def check_k3_moe_push(ctx): + results = [] + front = _front(ctx) + layer = _k3_moe_layer(_experts(20260928 + ctx.rank)) + copies16 = 16 // ctx.world + for start in COUNT_STARTS: + ex, ex16 = _Exchange(ctx, ctx.world, start), _Exchange(ctx, 16, start) + for producer in K3_MOE_PRODUCERS: + for m in K3_MOE_TOKENS: + for si in range(K3_MOE_SETS): + inputs = _k3_moe_inputs(producer, m, 2000 + 10 * si + m) + *routed, shared = _k3_moe_routed(producer, inputs, front) + y = layer(*routed, 0) + ref = _allreduce(ctx, y) + got, pushed_shared = [], [] + for _ in range(3): + *routed, pushed = _k3_moe_routed(producer, inputs, front) + layer.push(*routed, 0, ex, ctx.rank) + pushed_shared.append(pushed) + got.append(ex.reduce(m)) + state = ex.state_ok(ctx) + rows = [t.cuda() for t in ctx.comm.allgather(y.cpu())] + # TP16's receive side: this rank's partial in slots 4 r .. 4 r + 3, one push each. + for c in range(copies16): + *routed, pushed = _k3_moe_routed(producer, inputs, front) + layer.push(*routed, 0, ex16, ctx.rank * copies16 + c) + pushed_shared.append(pushed) + got16 = ex16.reduce(m) + state16 = ex16.state_ok(ctx) + shared_eq = shared is None or all(_same(s, shared) for s in pushed_shared) + results.append(_row(ctx, f"K3MoeLayer.push after {producer}", f"set{si}_count{start}", m, y, ref, + got, state, rows, copies16, got16, state16, shared_eq=shared_eq)) # fmt: skip + return results + + +CHECKS = {"push": check_push, "k3_moe_push": check_k3_moe_push} + + +def _run_checks(names): + try: + ctx = _context() + with torch.inference_mode(): + return [row for name in names for row in CHECKS[name](ctx)] + except Exception: + traceback.print_exc() + raise + + +def _report(rows): + for row in rows: + fields = " ".join(f"{k}={v}" for k, v in row.items() if k not in ("op", "case", "M")) + print(f"OPCHECK op={row['op']} case={row['case']} M={row['M']} {fields}", flush=True) + + +@pytest.mark.parametrize("mpi_pool_executor", [WORLD], indirect=True) +@pytest.mark.parametrize("check", list(CHECKS)) +def test_k3_moe_push(mpi_pool_executor, check): + per_rank = list(mpi_pool_executor.map(_run_checks, [[check]] * WORLD)) + _report(per_rank[0]) + assert all(row["ok"] for rows in per_rank for row in rows) + + +def main() -> int: + names = sys.argv[1:] or list(CHECKS) + rows = _run_checks(names) + if MPI.COMM_WORLD.Get_rank() == 0: + _report(rows) + print("PASS" if all(r["ok"] for r in rows) else "FAIL", flush=True) + return 0 if all(r["ok"] for r in rows) else 1 + + +if __name__ == "__main__": + sys.exit(main()) From 82fc8ad43dc6ca9b59ea9e6a6a97aa16bc697dce Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 22:44:54 -0700 Subject: [PATCH 094/161] [None][feat] modeling_v2 catalog: moe/k3_moe_m1 and moe/k3_moe_m2 over caller-owned states Two stateful entries for the routed experts of one or two decode tokens at moe TP16 x EP1, each with its contract (a State section for the caller-owned K3MoeM1State / K3MoeM2State: contents, creator, sharing, call order, what a later call reads, re-arm) and its wrappers (the plain call and the push form into a K3LatentExchange), in index.yaml. Tests (one GPU, l0_b200): single calls against the fp32 reference, and the call sequences the contracts certify: 12 steps of 3 layers on one state, capture and replay with rewritten inputs, two states interleaved, the epochs across the int32 wrap, and create() refusing capture. test_k3_moe_push.py gains the push form's call sequences at 4 ranks, with a negative control (one rank pushing two token sets in the other order). Both entries carry sm_100 receipts with their test counts, and the README's sm_100 count follows (25 entries, 22 of them Kimi K3). Signed-off-by: Vasanth Sabavat --- .../_experimental/modeling_v2/README.md | 4 +- .../modeling_v2/catalog/index.yaml | 6 + .../modeling_v2/catalog/moe/k3_moe_m1.md | 163 +++++++++ .../modeling_v2/catalog/moe/k3_moe_m1.py | 51 +++ .../modeling_v2/catalog/moe/k3_moe_m2.md | 164 +++++++++ .../modeling_v2/catalog/moe/k3_moe_m2.py | 51 +++ .../test_lists/test-db/l0_b200.yml | 2 + .../kimi_k3/test_k3_moe_push.py | 84 ++++- .../_torch/modeling_v2/moe/_k3_moe_engines.py | 333 ++++++++++++++++++ .../moe/test_modeling_v2_k3_moe_m1.py | 44 +++ .../moe/test_modeling_v2_k3_moe_m2.py | 37 ++ 11 files changed, 936 insertions(+), 3 deletions(-) create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m1.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m1.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m2.md create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m2.py create mode 100644 tests/unittest/_torch/modeling_v2/moe/_k3_moe_engines.py create mode 100644 tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_m1.py create mode 100644 tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_m2.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/README.md b/tensorrt_llm/_torch/_experimental/modeling_v2/README.md index 7b0ffa818edf..93cebcb77a9b 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/README.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/README.md @@ -186,8 +186,8 @@ Perf is measured, never gated. ## Status of every record in this tree -**The catalog is certified: 19 entries on sm_103, and 23 on sm_100 (B200 / -GB200): the 20 Kimi K3 entries (the MNNVL ones among them), where their +**The catalog is certified: 19 entries on sm_103, and 25 on sm_100 (B200 / +GB200): the 22 Kimi K3 entries (the MNNVL ones among them), where their first caller runs, and `cublas_mm`, `flashinfer_rmsnorm` and `allgather`. The targets construct but have never executed.** diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml index 2111035e82ee..5b6dbbf8ab45 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml @@ -316,3 +316,9 @@ entries: - path: moe/k3_moe.py impl: torch.ops.trtllm.k3_moe summary: "Kimi K3's routed experts at decode size: this rank's routed partial from the persistent CuTe DSL kernel k3_moe (FC1 + SiTU + FC2 with the routing-weighted combine over the TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers, read in place) on the routing and MXFP8 latent of k3_route_quant or the MoE front, over caller-owned per-rank state: a K3MoeState (up to 8 tokens; its head_flags build acquires a publishing front's ready words on the group's K3MoeHeadWorkspace) or a K3MoeWideState (up to 64), and one K3MoeLayer per MoE layer; a state's layers share its scratch (left armed by every call) and run in one stream order, each layer's counters are left zero" + - path: moe/k3_moe_m1.py + impl: torch.ops.trtllm.k3_moe_m1 + summary: "Kimi K3's routed experts of one or two decode tokens with every expert on the rank (moe TP16 x EP1) as one weight-stream CuTe DSL kernel (trtllm::k3_moe_m1): FC1 (MXFP4 x MXFP8), SiTU, the MXFP8 intermediate, FC2 and the routing-weighted combine in trtllm::k3_moe's order, returned as this rank's routed partial or pushed into every rank's latent exchange for trtllm::k3_latent_reduce (bit-exact against MNNVLAllReduce's one-shot); over a caller-owned K3MoeM1State, created eagerly with its builds compiled before capture, whose workspace every layer shares in one stream order and every call re-arms for the next" + - path: moe/k3_moe_m2.py + impl: torch.ops.trtllm.k3_moe_m2 + summary: "Kimi K3's routed experts of two decode tokens with every expert on the rank and an intermediate of at most 256 (moe TP16 x EP1) as one weight-stream CuTe DSL kernel (trtllm::k3_moe_m2), k3_moe_m1's function with per-group FC1 -> FC2 hand-offs: returned as this rank's routed partial or pushed into every rank's latent exchange for trtllm::k3_latent_reduce; over a caller-owned K3MoeM2State, created eagerly with its builds compiled before capture, whose workspace every layer shares in one stream order and every call re-arms for the next" diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m1.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m1.md new file mode 100644 index 000000000000..a506c0704cd7 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m1.md @@ -0,0 +1,163 @@ +--- +receipts: + sm_100: {status: passed, tests: 14} +--- + +# k3_moe_m1 + +**Wraps** `torch.ops.trtllm.k3_moe_m1` through `K3MoeM1Layer.__call__` and `K3MoeM1Layer.push` +(`tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py`; one CuTe DSL kernel, `k3_moe_m1_kernel.py`, launched +with programmatic dependent launch). + +A stateful entry: its correctness depends on a caller-owned state object, `K3MoeM1State`, which the `layer` argument +carries; see *State*. + +## Semantics + +Kimi K3's routed experts for M = 1 or 2 decode tokens when every expert is on this rank (the routed experts at moe +TP16 x EP1: each rank holds all 896 experts, a 192-wide slice of each intermediate). The inputs are what +`trtllm::k3_moe_front` or `trtllm::k3_route_quant` return. The output is this rank's routed partial, the tensor +`trtllm::k3_moe` returns for the same tokens. Per token `t`, over its 16 routed experts `e` whose global id is +in `[local_expert_offset, local_expert_offset + num_local)`: + +``` +x = x_fp8[t] * 2^(x_sf[t, k // 32] - 127) # the MXFP8 latent, [3584] +gate = x @ W1[e].T ; up = x @ W3[e].T # MXFP4 weights, fp32 accumulation, [i_tp] +act = 4 tanh(gate / 4) sigmoid(gate) * 25 tanh(up / 25) # SiTU +h = MXFP8(act) # per 32 columns, round-up E8M0 scale +out[t] = bf16(sum over e of topk_weights[t, e] * (h @ W2[e].T)) # [3584] +``` + +The experts' weights are the TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers, read in place, with the loader's zero padding of +the intermediate to `i_pad` (192 -> 256): the kernel streams the `i_tp` real values only. The combine sums a token's +experts in `trtllm::k3_moe`'s order (ascending local id, in up to 5 slices; products and sums rounded on their +own), so the result has k3_moe's bits except where FC1's two partial sums round an intermediate value differently +(one bf16 ulp of the row's max at most in the op tests). The order is fixed: the result is deterministic (certified, +run to run). A token with no local expert gets a zero row. + +**Push form.** `k3_moe_m1_push` computes the same partial and, instead of returning it, stores token `t`'s row into +slot `rank` of half `exchange.flags[0] & 1` of every rank's latent exchange (`K3LatentExchange`, int32 +`[2][8][world][1792]`: bf16 pairs, `0x80000000` empty, -0.0 stored as +0.0) through its multicast mapping. It reads +the half after its grid-dependency wait and writes no flags word. `trtllm::k3_latent_reduce` sums the slots in the +MNNVL one-shot's order, empties the half it read and advances the count. A push and its reduce equal +`MNNVLAllReduce`'s one-shot of the plain partials bit for bit (certified at 4 ranks). + +Fusion boundary. Inside: FC1, SiTU, the MXFP8 intermediate, FC2 and the routing-weighted combine of this rank's +experts. Outside: the routing and input quantization (their producer), the sum over ranks (MNNVL all-reduce, or the +push form plus `trtllm::k3_latent_reduce`), the shared expert. + +## Signature + +```python +def k3_moe_m1( + x_fp8: torch.Tensor, + x_sf: torch.Tensor, + topk_ids: torch.Tensor, + topk_weights: torch.Tensor, + local_expert_offset: int, + layer: K3MoeM1Layer, + out: Optional[torch.Tensor] = None, +) -> torch.Tensor + +def k3_moe_m1_push( + x_fp8: torch.Tensor, + x_sf: torch.Tensor, + topk_ids: torch.Tensor, + topk_weights: torch.Tensor, + local_expert_offset: int, + layer: K3MoeM1Layer, + exchange: K3LatentExchange, +) -> None +``` + +`layer = state.layer(w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale)`: one MoE layer's buffers on a +state (it checks them and keeps them; it allocates nothing). + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `x_fp8` | `[M, 3584]`, `M` = `state.num_tokens` (1 and 2 certified) | float8_e4m3fn | contiguous | CUDA, the state's device | +| `x_sf` | `M x 112` E8M0 bytes, any shape | uint8 | contiguous | CUDA | +| `topk_ids` | `[M, 16]`, global expert ids | int32 | contiguous | CUDA | +| `topk_weights` | `[M, 16]` | bf16 | contiguous | CUDA | +| `local_expert_offset` | the first global id of the layer's experts (0 at TP16) | Python int | — | — | +| `layer` weights | `w3_w1_weight [E, 2 i_pad, 1792]`, `w3_w1_weight_scale [E, 2 i_pad, 112]`, `w2_weight [E, 3584, i_pad / 2]`, `w2_weight_scale [E, 3584, i_pad / 32]`; `E` = 896, `i_tp` = 192, `i_pad` = 256 certified | uint8 | contiguous (TRTLLM-Gen layout) | CUDA | +| `out` | `None`, or `[M, 3584]` | bf16 | contiguous | CUDA | +| `exchange` (push) | a `K3LatentExchange` of this rank's TP group, 4 ranks certified | — | — | — | +| returns | `[M, 3584]` (`out` when given); `None` for the push form | bf16 | contiguous | the state's device | + +The inputs are read only. A call writes the state's workspace; the push form also writes every rank's exchange. +The op's `mutates_args` names them: `hbuf`, `counts`, `epochs`, `exchange_mc` (the push form) and `out`. + +## State + +**Object.** `K3MoeM1State` (`k3_fused_moe/op.py`), one per device and token count, owned by the caller. Every +layer's handle (`K3MoeM1Layer`) points at it. + +**Contents and size.** The workspace its layers share: `hbuf`, the intermediate rows of one call (int8 +`[16 M x M x 208]`, 3.3 KB at M 1, 13 KB at M 2); `counts`, the call's FC1 -> FC2 hand-off count in two slots by +epoch parity (int32 `[4]`, 2 used); `epochs`, one word per CTA (int32 `[SMs]`); The compiled builds (the plain build and one +push build per (exchange slots, copies)) live in the module's cache, keyed by device and configuration. + +**Who creates it, and when.** The target, in `post_load_weights`, with +`K3MoeM1State.create(device, i_tp, i_pad, num_local, num_tokens, push=((tp_size, 1),))`: + +- eager: it allocates the workspace (all zeros) and compiles the plain build and each listed push build without + launching anything, so it refuses to run under CUDA-graph capture; +- not collective: the state is this device's alone (the exchange is a separate, collective object); +- a build it did not compile compiles on its first call, which must also come before capture. + +**Which ops may share one object.** Every layer of the model and both forms (plain and push) share one state per +device and token count. A `K3MoeM2State` is a separate workspace. Two states are independent: calls alternating +between two states in an irregular order are all correct (certified). + +**Call-order invariant.** The calls on one state run one after another in one stream order: eager calls and graph +replays alike, whichever layer they belong to. Each call is a complete kernel: its CTAs read the count slot of their +epoch's parity, CTA 0 zeroes the other slot (the next call's), and every CTA advances its epoch. With programmatic +dependent launch the kernel's `griddepcontrol.wait` precedes every read of the producer's outputs and every global +write, so a call may start early beside its predecessor but touches the workspace only after it. + +**What a later launch reads.** The CTAs' epochs (their parity picks the slot) and the slot the previous call zeroed. +The intermediate rows are written before they are read within one call; no call reads another call's rows. + +**How it is re-armed.** Every call zeroes the slot the next call uses and advances the epochs, so the workspace +never needs a reset. Only the parity of an epoch matters, so the int32 wrap after 2^31 calls is harmless +(certified: epochs preset at 2^31 - 2 and 2^31 - 1). + +**Why the test drives call sequences.** The re-arm and the parity act only across calls: a call that left the next +slot armed wrongly, or a parity that broke at the wrap, passes every single-call test. The test runs 12 steps of 3 +layers on one state with the tokens changing every step, a captured step replayed with rewritten inputs and eager +calls between replays, two states interleaved, and the wrap; every call against the same call on a reference state, +bit for bit, and the epochs and the next slot checked after each call. + +**What a wrong order does.** Not exercised on the state: two calls that overlap on one state (two streams without +ordering) are outside the contract. The push form's order across ranks is the exchange's (`trtllm::k3_latent_reduce`: +each push followed by one reduce of the same token count, in the same order on every rank). + +## Metadata consumed + +None besides `layer` (its state and weights) and `exchange`. The kernel modules and their compiled builds are +cached per configuration in the process (code only, no state). + +## Preconditions + +- sm_100 and the CuTe DSL (`nvidia-cutlass-dsl`); at least 112 SMs. +- `op.m1_supported(...)` holds for the layer's buffers: contiguous uint8 of the shapes above, `i_pad` a multiple of + 128, `0 < i_tp <= i_pad`, `i_tp` a multiple of 32, `E` <= 896. The layer raises `ValueError` otherwise, and when + its padded intermediate differs from the state's. +- `M` = the state's token count (1 or 2), inputs contiguous; a call raises `ValueError` otherwise. +- The state was created before any capture; calls may be captured (certified: 3 layers captured, replayed 4 times). +- Push form: the exchange belongs to this rank's TP group, and one reduce of `M` tokens on it follows each push + before the next, on every rank in the same order. + +## Notes + +- Certified path: one GB200 GPU (sm_100), TP16 shapes (896 experts, `i_tp` 192 in buffers of 256), M 1 and 2. + Tests: `tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_m1.py` (the reference is fp32 over the + dequantized checkpoint slices: 8 ulp of the row's max per element, 4 ulp relative RMS); + `tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_m1.py` (also against `trtllm::k3_moe` and at + TP4 x EP4 shapes, M 1); the push form at 4 ranks in `test_k3_moe_push.py` (bit-exact against `MNNVLAllReduce`, + a 16-slot exchange filled 4 slots per rank, the exchange's count across the int32 wrap). +- Kimi K3 runs the push form over 16 ranks; a 4-rank run fills a 16-slot exchange only by repeating each rank's + partial. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m1.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m1.py new file mode 100644 index 000000000000..c2b2eb8b6ef9 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m1.py @@ -0,0 +1,51 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3's routed experts of one or two decode tokens with every expert on the rank (moe TP16 x EP1): one +weight-stream CuTe DSL kernel over a caller-owned :class:`K3MoeM1State`, returning this rank's routed partial or +pushing it into a latent exchange for ``trtllm::k3_latent_reduce``.""" + +from typing import Optional + +import torch + +from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.op import K3MoeM1Layer, K3MoeM1State + +__all__ = ["K3MoeM1Layer", "K3MoeM1State", "k3_moe_m1", "k3_moe_m1_push"] + + +def k3_moe_m1( + x_fp8: torch.Tensor, + x_sf: torch.Tensor, + topk_ids: torch.Tensor, + topk_weights: torch.Tensor, + local_expert_offset: int, + layer: K3MoeM1Layer, + out: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """This rank's routed partial ``[M, 3584]`` bf16 of ``layer``'s experts for the M tokens of its state. Advances + the state's workspace by one call: the calls on one state run in one stream order.""" + return layer(x_fp8, x_sf, topk_ids, topk_weights, local_expert_offset, out) + + +def k3_moe_m1_push( + x_fp8: torch.Tensor, + x_sf: torch.Tensor, + topk_ids: torch.Tensor, + topk_weights: torch.Tensor, + local_expert_offset: int, + layer: K3MoeM1Layer, + exchange, +) -> None: + """The same partial stored into this rank's slot of every rank's ``exchange`` (a TP group's + ``K3LatentExchange``) instead of returned; one ``trtllm::k3_latent_reduce`` of the M tokens on that exchange must + follow before the next push. Advances the state's workspace by one call.""" + layer.push( + x_fp8, + x_sf, + topk_ids, + topk_weights, + local_expert_offset, + exchange.mc, + exchange.flags, + exchange.rank, + ) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m2.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m2.md new file mode 100644 index 000000000000..07aaa35f6334 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m2.md @@ -0,0 +1,164 @@ +--- +receipts: + sm_100: {status: passed, tests: 7} +--- + +# k3_moe_m2 + +**Wraps** `torch.ops.trtllm.k3_moe_m2` through `K3MoeM2Layer.__call__` and `K3MoeM2Layer.push` +(`tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py`; one CuTe DSL kernel, `k3_moe_m2_kernel.py`, launched +with programmatic dependent launch). + +A stateful entry: its correctness depends on a caller-owned state object, `K3MoeM2State`, which the `layer` argument +carries; see *State*. + +## Semantics + +Kimi K3's routed experts for two decode tokens when every expert is on this rank and a rank's intermediate is at +most 256 (the routed experts at moe TP16 x EP1: all 896 experts, a 192-wide slice of each intermediate). It computes +the same function as `moe/k3_moe_m1` at M = 2, with a different schedule: every distinct local expert of the two +tokens takes one slot in ascending id, each slot's intermediate is computed for both tokens (the MMA columns are +independent, so a token's bits do not depend on the other token), and FC2 starts per group of 4 slots as soon as +that group's FC1 rows are complete. Per token `t`, over its 16 routed experts `e` whose +global id is in `[local_expert_offset, local_expert_offset + num_local)`: + +``` +x = x_fp8[t] * 2^(x_sf[t, k // 32] - 127) # the MXFP8 latent, [3584] +gate = x @ W1[e].T ; up = x @ W3[e].T # MXFP4 weights, fp32 accumulation, [i_tp] +act = 4 tanh(gate / 4) sigmoid(gate) * 25 tanh(up / 25) # SiTU +h = MXFP8(act) # per 32 columns, round-up E8M0 scale +out[t] = bf16(sum over e of topk_weights[t, e] * (h @ W2[e].T)) # [3584] +``` + +The experts' weights are the TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers, read in place, with the loader's zero padding of +the intermediate to `i_pad` (192 -> 256): the kernel streams the `i_tp` real values only. The combine sums a token's +experts in `trtllm::k3_moe`'s order (ascending local id, in up to 5 slices; products and sums rounded on their +own), so the result has k3_moe's bits except where FC1's two partial sums round an intermediate value differently +(one bf16 ulp of the row's max at most in the op tests). The order is fixed: the result is deterministic (certified, +run to run). + +**Push form.** `k3_moe_m2_push` computes the same partials and, instead of returning them, stores each token's row +into slot `rank` of half `exchange.flags[0] & 1` of every rank's latent exchange (`K3LatentExchange`, int32 +`[2][8][world][1792]`: bf16 pairs, `0x80000000` empty, -0.0 stored as +0.0) through its multicast mapping. It reads +the half after its grid-dependency wait and writes no flags word. `trtllm::k3_latent_reduce` sums the slots in the +MNNVL one-shot's order, empties the half it read and advances the count. A push and its reduce equal +`MNNVLAllReduce`'s one-shot of the plain partials bit for bit (certified at 4 ranks). + +Fusion boundary. Inside: FC1, SiTU, the MXFP8 intermediate, FC2 and the routing-weighted combine of this rank's +experts. Outside: the routing and input quantization (their producer), the sum over ranks (MNNVL all-reduce, or the +push form plus `trtllm::k3_latent_reduce`), the shared expert. + +## Signature + +```python +def k3_moe_m2( + x_fp8: torch.Tensor, + x_sf: torch.Tensor, + topk_ids: torch.Tensor, + topk_weights: torch.Tensor, + local_expert_offset: int, + layer: K3MoeM2Layer, + out: Optional[torch.Tensor] = None, +) -> torch.Tensor + +def k3_moe_m2_push( + x_fp8: torch.Tensor, + x_sf: torch.Tensor, + topk_ids: torch.Tensor, + topk_weights: torch.Tensor, + local_expert_offset: int, + layer: K3MoeM2Layer, + exchange: K3LatentExchange, +) -> None +``` + +`layer = state.layer(w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale)`: one MoE layer's buffers on a +state (it checks them and keeps them; it allocates nothing). + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `x_fp8` | `[2, 3584]` | float8_e4m3fn | contiguous | CUDA, the state's device | +| `x_sf` | `2 x 112` E8M0 bytes, any shape | uint8 | contiguous | CUDA | +| `topk_ids` | `[2, 16]`, global expert ids | int32 | contiguous | CUDA | +| `topk_weights` | `[2, 16]` | bf16 | contiguous | CUDA | +| `local_expert_offset` | the first global id of the layer's experts (0 at TP16) | Python int | — | — | +| `layer` weights | `w3_w1_weight [E, 2 i_pad, 1792]`, `w3_w1_weight_scale [E, 2 i_pad, 112]`, `w2_weight [E, 3584, i_pad / 2]`, `w2_weight_scale [E, 3584, i_pad / 32]`; `E` = 896, `i_tp` = 192, `i_pad` = 256 certified | uint8 | contiguous (TRTLLM-Gen layout) | CUDA | +| `out` | `None`, or `[2, 3584]` | bf16 | contiguous | CUDA | +| `exchange` (push) | a `K3LatentExchange` of this rank's TP group, 4 ranks certified | — | — | — | +| returns | `[2, 3584]` (`out` when given); `None` for the push form | bf16 | contiguous | the state's device | + +The inputs are read only. A call writes the state's workspace; the push form also writes every rank's exchange. +The op's `mutates_args` names them: `hbuf`, `counts`, `epochs`, `exchange_mc` (the push form) and `out`. + +## State + +**Object.** `K3MoeM2State` (`k3_fused_moe/op.py`), one per device, owned by the caller. Every layer's handle +(`K3MoeM2Layer`) points at it. + +**Contents and size.** The workspace its layers share: `hbuf`, the intermediate rows of one call (int8 +`[32 x 2 x 208]`, 13 KB); `counts`, the FC1 -> FC2 counters of the 8 groups, one 128-byte line each, in two sets by +epoch parity (int32 `[2, 8, 32]`); `epochs`, one word per CTA (int32 `[SMs]`); The compiled builds (the plain build and one push build per (exchange slots, copies)) live in the module's cache, +keyed by device and configuration. + +**Who creates it, and when.** The target, in `post_load_weights`, with +`K3MoeM2State.create(device, i_tp, i_pad, num_local, push=((tp_size, 1),))`: + +- eager: it allocates the workspace (all zeros) and compiles the plain build and each listed push build without + launching anything, so it refuses to run under CUDA-graph capture; +- not collective: the state is this device's alone (the exchange is a separate, collective object); +- a build it did not compile compiles on its first call, which must also come before capture. + +**Which ops may share one object.** Every layer of the model and both forms (plain and push) share one state per +device. A `K3MoeM1State` is a separate workspace. Two states are independent: calls alternating between two states in +an irregular order are all correct (certified). + +**Call-order invariant.** The calls on one state run one after another in one stream order: eager calls and graph +replays alike, whichever layer they belong to. Each call is a complete kernel: its CTAs use the counter set of their +epoch's parity, CTA 0 zeroes the other set (the next call's), and every CTA advances its epoch. With programmatic +dependent launch the kernel's `griddepcontrol.wait` precedes every read of the producer's outputs and every global +write, so a call may start early beside its predecessor but touches the workspace only after it. + +**What a later launch reads.** The CTAs' epochs (their parity picks the set) and the set the previous call zeroed. +The intermediate rows are written before they are read within one call; no call reads another call's rows. + +**How it is re-armed.** Every call zeroes the set the next call uses and advances the epochs, so the workspace never +needs a reset. Only the parity of an epoch matters, so the int32 wrap after 2^31 calls is harmless (certified: +epochs preset at 2^31 - 2 and 2^31 - 1). + +**Why the test drives call sequences.** The re-arm and the parity act only across calls: a call that left the next +set armed wrongly, or a parity that broke at the wrap, passes every single-call test. The test runs 12 steps of 3 +layers on one state with the tokens changing every step, a captured step replayed with rewritten inputs and eager +calls between replays, two states interleaved, and the wrap; every call against the same call on a reference state, +bit for bit, and the epochs and the next set checked after each call. + +**What a wrong order does.** Not exercised on the state: two calls that overlap on one state (two streams without +ordering) are outside the contract. The push form's order across ranks is the exchange's (`trtllm::k3_latent_reduce`: +each push followed by one reduce of the same token count, in the same order on every rank). + +## Metadata consumed + +None besides `layer` (its state and weights) and `exchange`. The kernel modules and their compiled builds are +cached per configuration in the process (code only, no state). + +## Preconditions + +- sm_100 and the CuTe DSL (`nvidia-cutlass-dsl`); at least 112 SMs. +- `op.m2_supported(...)` holds for the layer's buffers: `op.m1_supported`'s conditions and a padded intermediate of + at most 256. The layer raises `ValueError` otherwise, and when its padded intermediate differs from the state's. +- Two tokens, inputs contiguous; a call raises `ValueError` otherwise. +- The state was created before any capture; calls may be captured (certified: 3 layers captured, replayed 4 times). +- Push form: the exchange belongs to this rank's TP group, and one reduce of 2 tokens on it follows each push before + the next, on every rank in the same order. + +## Notes + +- Certified path: one GB200 GPU (sm_100), TP16 shapes (896 experts, `i_tp` 192 in buffers of 256). Tests: + `tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_m2.py` (the reference is fp32 over the dequantized + checkpoint slices: 8 ulp of the row's max per element, 4 ulp relative RMS); + `tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_m2.py` (also against `trtllm::k3_moe`, two + tokens routed to the same and to disjoint experts); the push form at 4 ranks in `test_k3_moe_push.py` (bit-exact + against `MNNVLAllReduce`, a 16-slot exchange filled 4 slots per rank, the exchange's count across the int32 wrap). +- Kimi K3 runs the push form over 16 ranks; a 4-rank run fills a 16-slot exchange only by repeating each rank's + partial. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m2.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m2.py new file mode 100644 index 000000000000..68d0854310ed --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m2.py @@ -0,0 +1,51 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Kimi K3's routed experts of two decode tokens with every expert on the rank (moe TP16 x EP1, intermediate <= 256): +one weight-stream CuTe DSL kernel over a caller-owned :class:`K3MoeM2State`, returning this rank's routed partial or +pushing it into a latent exchange for ``trtllm::k3_latent_reduce``.""" + +from typing import Optional + +import torch + +from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.op import K3MoeM2Layer, K3MoeM2State + +__all__ = ["K3MoeM2Layer", "K3MoeM2State", "k3_moe_m2", "k3_moe_m2_push"] + + +def k3_moe_m2( + x_fp8: torch.Tensor, + x_sf: torch.Tensor, + topk_ids: torch.Tensor, + topk_weights: torch.Tensor, + local_expert_offset: int, + layer: K3MoeM2Layer, + out: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """This rank's routed partial ``[2, 3584]`` bf16 of ``layer``'s experts for two tokens. Advances the state's + workspace by one call: the calls on one state run in one stream order.""" + return layer(x_fp8, x_sf, topk_ids, topk_weights, local_expert_offset, out) + + +def k3_moe_m2_push( + x_fp8: torch.Tensor, + x_sf: torch.Tensor, + topk_ids: torch.Tensor, + topk_weights: torch.Tensor, + local_expert_offset: int, + layer: K3MoeM2Layer, + exchange, +) -> None: + """The same partials stored into this rank's slot of both tokens' rows of every rank's ``exchange`` (a TP + group's ``K3LatentExchange``) instead of returned; one ``trtllm::k3_latent_reduce`` of the two tokens on that + exchange must follow before the next push. Advances the state's workspace by one call.""" + layer.push( + x_fp8, + x_sf, + topk_ids, + topk_weights, + local_expert_offset, + exchange.mc, + exchange.flags, + exchange.rank, + ) diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index ac6e6970ef8f..1b126315b963 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -327,6 +327,8 @@ l0_b200: - unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_route_quant.py - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_m1.py - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_m2.py + - unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_m1.py + - unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_m2.py # ------------- Visual Gen tests --------------- - unittest/_torch/cute_dsl_kernels/test_nvfp4_conv3d.py - unittest/_torch/visual_gen/test_media_decode.py diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py index 1f6625d1effc..66712186a648 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py @@ -372,7 +372,89 @@ def check_k3_moe_push(ctx): return results -CHECKS = {"push": check_push, "k3_moe_push": check_k3_moe_push} +def check_sequences(ctx): + """The push form of the moe/k3_moe_m1 entry (catalog wrappers, one token) on one created state and one exchange, + each push followed by one reduce, against MNNVLAllReduce's one-shot of the plain partials: + steps 12 steps of 3 layers (the experts in 3 orders), the tokens changing every step, a random rank 5 ms + late at every call; + replay one step captured and replayed 4 times with rewritten inputs, an eager push + reduce between replays; + swapped the negative control: rank 0 pushes two token sets in the other order. Nothing raises or hangs, but + every rank's two sums are wrong (each reduce sums rank 0's partial of the other set); the next pair, + in the same order on every rank, is right again.""" + import random + import time + + from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.k3_moe_m1 import ( + k3_moe_m1, + k3_moe_m1_push, + ) + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op + + device = torch.device("cuda", torch.cuda.current_device()) + base = _experts(20260928 + ctx.rank) + weights = [ + tuple(t.roll(li, dims=0).contiguous() for t in base) if li else base for li in range(3) + ] + state = op.K3MoeM1State.create( + device, I_TP, I_PAD, NUM_EXPERTS, num_tokens=1, push=((ctx.world, 1),) + ) + plain_state = op.K3MoeM1State.create(device, I_TP, I_PAD, NUM_EXPERTS, num_tokens=1) + layers = [state.layer(*w) for w in weights] + plain_layers = [plain_state.layer(*w) for w in weights] + sets = [_tokens(1, 3000 + s) for s in range(5)] + refs = {(li, si): _allreduce(ctx, k3_moe_m1(*sets[si], 0, plain_layers[li])) + for li in range(3) for si in range(5)} # fmt: skip + ex = _Exchange(ctx, ctx.world, 0) + late = random.Random(7) # the same draws on every rank + results = [] + + def push_reduce(li, tokens): + k3_moe_m1_push(*tokens, 0, layers[li], ex) + return ex.reduce(1) + + good = True + for step in range(12): + for li in range(3): + if late.randrange(ctx.world) == ctx.rank: + time.sleep(0.005) + good &= _same(push_reduce(li, sets[step % 4]), refs[li, step % 4]) + results.append(dict(op="k3_moe_m1_push", case="steps", M=1, exact=good, state=ex.state_ok(ctx))) + + static = [tuple(t.clone() for t in sets[0]) for _ in range(3)] + graph = torch.cuda.CUDAGraph() + ctx.comm.Barrier() + with torch.cuda.graph(graph): + outs = [push_reduce(li, static[li]) for li in range(3)] + ex.count = (ex.count - 3 + 2**31) % 2**32 - 2**31 # capture launched nothing + good = True + for rep in range(4): + for li in range(3): + for dst, src in zip(static[li], sets[(rep + li) % 4]): + dst.copy_(src) + ctx.comm.Barrier() + graph.replay() + ex.count = (ex.count + 3 + 2**31) % 2**32 - 2**31 + good &= all(_same(outs[li], refs[li, (rep + li) % 4]) for li in range(3)) + good &= _same(push_reduce(rep % 3, sets[4]), refs[rep % 3, 4]) + results.append( + dict(op="k3_moe_m1_push", case="replay", M=1, exact=good, state=ex.state_ok(ctx)) + ) + del graph + + first, second = (sets[1], sets[0]) if ctx.rank == 0 else (sets[0], sets[1]) + got0, got1 = push_reduce(0, first), push_reduce(0, second) + wrong = [(got0 != refs[0, 0]).float().mean().item(), (got1 != refs[0, 1]).float().mean().item()] + after = _same(push_reduce(0, sets[2]), refs[0, 2]) + results.append(dict(op="k3_moe_m1_push", case="swapped", M=1, detected=min(wrong) > 0.5, after=after, + state=ex.state_ok(ctx))) # fmt: skip + for row in results: + row["ok"] = all( + ctx.comm.allgather(all(v for k, v in row.items() if k not in ("op", "case", "M"))) + ) + return results + + +CHECKS = {"push": check_push, "k3_moe_push": check_k3_moe_push, "sequences": check_sequences} def _run_checks(names): diff --git a/tests/unittest/_torch/modeling_v2/moe/_k3_moe_engines.py b/tests/unittest/_torch/modeling_v2/moe/_k3_moe_engines.py new file mode 100644 index 000000000000..192e4512300d --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/moe/_k3_moe_engines.py @@ -0,0 +1,333 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Shared pieces of the moe/k3_moe_m1 and moe/k3_moe_m2 tests: a rank's TP16 experts through TRT-LLM's loader, +routed tokens, the fp32 reference, and the checks both stateful entries run over their caller-owned state (single +calls, call sequences across layers and steps, capture and replay, two states interleaved, the epochs across the +int32 wrap, the explicit constructor). + +Imported by the two test files beside it (this tree is not a package). +""" + +import functools +import math +from types import SimpleNamespace + +import torch + +H, NUM_EXPERTS, TOP_K, SV = 3584, 896, 16, 32 +# TP16: every expert on the rank, its 192-wide intermediate slice zero-padded to whole tiles (256) by the loader. +I_TP, I_PAD, MOE_TP = 192, 256, 16 +GATE_CAP, LINEAR_CAP = ( + 4.0, + 25.0, +) # the SiTU caps (activation_situ_beta, activation_situ_linear_beta) +RSF = 2.827 +ULP = 2.0**-8 +E4M3_MAX = 448.0 +_E2M1 = [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0, -0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0] +LAYERS = 3 # distinct weight sets; the sequences cycle through them +TOKEN_SETS = 4 +STEPS = 12 + + +def is_sm100() -> bool: + return torch.cuda.is_available() and torch.cuda.get_device_capability() == (10, 0) + + +def bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int16) + + +def same(a: torch.Tensor, b: torch.Tensor) -> bool: + return a.shape == b.shape and torch.equal(bits(a), bits(b)) + + +def _rand_mxfp4(rows, k, k_full, gen): + """Random checkpoint-format MXFP4: packed [rows, k / 2] (low nibble = even k), E8M0 per 32 k, scaled so a + k_full-long dot product lands near std 3.""" + codes = torch.randint(0, 16, (rows, k), dtype=torch.uint8, device="cuda", generator=gen) + base = 127 + round(0.5 * math.log2(0.01057 / k_full)) + exps = torch.randint( + base, base + 6, (rows, k // SV), dtype=torch.uint8, device="cuda", generator=gen + ) + return (codes[:, 0::2] | (codes[:, 1::2] << 4)).contiguous(), exps + + +@functools.lru_cache(maxsize=None) +def experts(seed: int = 20260928): + """Rank 0's TP16 experts through W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod's loader (the buffers the engines read, + the 192-wide shard sliced from tensors that hold exactly it, then padded to 256) and the checkpoint slices the + reference reads.""" + from tensorrt_llm._torch.moe.fused_moe.quantization import W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod + + method = W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod() + module = SimpleNamespace(tp_size=MOE_TP, tp_rank=0, scaling_vector_size=SV, intermediate_size=I_TP * MOE_TP, + intermediate_size_per_partition=I_TP, hidden_size=H) # fmt: skip + kw = dict(dtype=torch.uint8, device="cuda") + proc = ( + torch.zeros(NUM_EXPERTS, 2 * I_PAD, H // 2, **kw), + torch.zeros(NUM_EXPERTS, 2 * I_PAD, H // SV, **kw), + torch.zeros(NUM_EXPERTS, H, I_PAD // 2, **kw), + torch.zeros(NUM_EXPERTS, H, I_PAD // SV, **kw), + ) + raw = {name: [] for name in ("up", "up_s", "gate", "gate_s", "down", "down_s")} + gen = torch.Generator(device="cuda").manual_seed(seed) + for e in range(NUM_EXPERTS): + gate, gate_s = _rand_mxfp4(I_TP, H, H, gen) + up, up_s = _rand_mxfp4(I_TP, H, H, gen) + down, down_s = _rand_mxfp4(H, I_TP, I_TP * MOE_TP, gen) + method.load_expert_w3_w1_weight(module, gate, up, proc[0][e]) + method.load_expert_w2_weight(module, down, proc[2][e]) + method.load_expert_w3_w1_weight_scale_mxfp4(module, gate_s, up_s, proc[1][e]) + method.load_expert_w2_weight_scale_mxfp4(module, down_s, proc[3][e]) + for name, t in zip(raw, (up, up_s, gate, gate_s, down, down_s)): + raw[name].append(t) + torch.cuda.synchronize() + return proc, raw + + +@functools.lru_cache(maxsize=None) +def layer_weights(index: int): + """The buffers of layer ``index``: layer 0's experts in another order (a different layer to the engine).""" + proc, _ = experts() + return tuple(t.roll(index, dims=0).contiguous() for t in proc) if index else proc + + +def routed(m: int, seed: int): + """m tokens' MXFP8 latents and routing, as trtllm::k3_route_quant produces them (the engines' caller).""" + import tensorrt_llm # noqa: F401 + from tensorrt_llm._torch.cute_dsl_kernels.k3_route_quant import op as _rq # noqa: F401 + + gen = torch.Generator(device="cuda").manual_seed(seed) + logits = (torch.randn(m, NUM_EXPERTS, generator=gen, device="cuda") * 3.0).float() + x = torch.randn(m, H, generator=gen, device="cuda").bfloat16() + bias = (torch.randn(NUM_EXPERTS, generator=gen, device="cuda") * 0.05).float() + ids, weights, x_fp8, x_sf = torch.ops.trtllm.k3_route_quant(logits, bias, x, RSF, True) + return x_fp8, x_sf, ids, weights + + +def _deq_w(packed, sf): + lut = torch.tensor(_E2M1, device=packed.device) + vals = torch.empty(packed.shape[0], packed.shape[1] * 2, device=packed.device) + vals[:, 0::2] = lut[(packed & 0xF).long()] + vals[:, 1::2] = lut[(packed >> 4).long()] + return vals * torch.exp2(sf.float() - 127.0).repeat_interleave(SV, dim=1) + + +def _requant(act): + """The FC1 epilogue's MXFP8 requantization per 32 columns (round-up scale), dequantized.""" + rows, cols = act.shape + blocks = act.reshape(rows, cols // SV, SV) + amax = blocks.abs().amax(dim=-1, keepdim=True) + ex = torch.ceil(torch.log2(amax / E4M3_MAX)) + ex = torch.where(amax == 0, torch.full_like(amax, -127.0), ex).clamp(-127.0, 127.0) + scale = torch.exp2(ex) + q8 = (blocks / scale).clamp(-E4M3_MAX, E4M3_MAX).to(torch.float8_e4m3fn) + return (q8.float() * scale).reshape(rows, cols) + + +def reference(x_fp8, x_sf, ids, weights): + """fp32 routed MoE over layer 0's experts from the checkpoint slices (f64 GEMMs): SiTU, the MXFP8 intermediate, + the down projection, the routing-weighted sum, per token; bf16 out.""" + _, raw = experts() + rows = x_fp8.shape[0] + x = x_fp8.float() * torch.exp2(x_sf.reshape(rows, H // SV).float() - 127.0).repeat_interleave( + SV, dim=1 + ) + out = torch.zeros(rows, H, device="cuda") + for t in range(rows): + for k in range(TOP_K): + e = int(ids[t, k]) + xe = x[t : t + 1].double() + up = (xe @ _deq_w(raw["up"][e], raw["up_s"][e]).double().t()).float() + gate = (xe @ _deq_w(raw["gate"][e], raw["gate_s"][e]).double().t()).float() + act = (GATE_CAP * torch.tanh(gate / GATE_CAP) * torch.sigmoid(gate) + * (LINEAR_CAP * torch.tanh(up / LINEAR_CAP))) # fmt: skip + y = ( + _requant(act).double() @ _deq_w(raw["down"][e], raw["down_s"][e]).double().t() + ).float() + out[t] += (y * weights[t, k].float())[0] + return out.bfloat16() + + +def row_ulp(y, ref): + """Largest |y - ref| in bf16 ulps of the row's max |ref|, and the relative RMS in ulps.""" + o, r = y.float(), ref.float() + row = r.abs().amax(dim=1, keepdim=True).clamp_min(1e-12) + elt = ((o - r).abs() / row).max().item() / ULP + rms = ((o - r).pow(2).mean().sqrt() / r.pow(2).mean().sqrt().clamp_min(1e-12)).item() / ULP + return elt, rms + + +class Engine: + """One entry at one token count: its state constructor and wrapper, and how to read the state's workspace (the + epochs every CTA advances, and the count words the next call uses, which every call leaves zero).""" + + def __init__(self, name: str, m: int): + self.name, self.m = name, m + + def create(self, push=()): + from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe import op + + device = torch.device("cuda", torch.cuda.current_device()) + if self.name == "k3_moe_m1": + return op.K3MoeM1State.create( + device, I_TP, I_PAD, NUM_EXPERTS, num_tokens=self.m, push=push + ) + return op.K3MoeM2State.create(device, I_TP, I_PAD, NUM_EXPERTS, push=push) + + def call(self, layer, tokens, out=None): + from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe import k3_moe_m1, k3_moe_m2 + + entry = k3_moe_m1.k3_moe_m1 if self.name == "k3_moe_m1" else k3_moe_m2.k3_moe_m2 + return entry(*tokens, 0, layer, out) + + def next_counts(self, state) -> torch.Tensor: + """The count words the next call uses (by the parity of CTA 0's epoch).""" + ep = int(state.epochs[0].item()) + if self.name == "k3_moe_m1": + return state.counts[:2][ep & 1 : (ep & 1) + 1] + words = state.mod.GROUPS2 * state.mod.CW + return state.counts.view(2, words)[ep & 1] + + +@functools.lru_cache(maxsize=None) +def _reference_calls(name: str, m: int): + """Every (layer, token set) call once, in order, on a reference state of its own: the bits each call must give + in any sequence on any state.""" + engine = Engine(name, m) + state = engine.create() + out = {} + for li in range(LAYERS): + layer = state.layer(*layer_weights(li)) + for si in range(TOKEN_SETS): + out[li, si] = engine.call(layer, routed(m, 100 + si)) + torch.cuda.synchronize() + return out + + +def _alone(engine): + return _reference_calls(engine.name, engine.m) + + +def check_single_calls(engine, seeds=range(8)): + """Single calls over the grid against the fp32 reference (8 ulp of the row max per element, 4 ulp relative + RMS), and run to run bit for bit.""" + state = engine.create() + layer = state.layer(*layer_weights(0)) + for seed in seeds: + tokens = routed(engine.m, seed) + y = engine.call(layer, tokens) + again = engine.call(layer, tokens) + elt, rms = row_ulp(y, reference(*tokens)) + print( + f"OPCHECK op={engine.name} M={engine.m} seed={seed} vs_ref_elt_ulp={elt:.2f} vs_ref_rms_ulp={rms:.2f}" + ) + assert bool(torch.isfinite(y.float()).all()) + assert elt <= 8.0 and rms <= 4.0, (elt, rms) + assert same(y, again), "run-to-run bits differ" + + +def check_call_sequences(engine): + """12 steps of the 3 layers on one state, the token set changing every step: each call gives the bits of the + same call on the reference state, the epochs count every call, and every call leaves the next call's count + words zero.""" + alone = _alone(engine) + state = engine.create() + layers = [state.layer(*layer_weights(li)) for li in range(LAYERS)] + calls = 0 + for step in range(STEPS): + si = step % TOKEN_SETS + for li in range(LAYERS): + y = engine.call(layers[li], routed(engine.m, 100 + si)) + calls += 1 + assert same(y, alone[li, si]), (step, li) + assert bool((engine.next_counts(state) == 0).all()), (step, li) + assert bool((state.epochs == calls).all()), "every CTA's epoch counts every call" + + +def check_capture_replay(engine): + """One step of the 3 layers captured on a created state (create compiled the build, so capture compiles nothing), + replayed 4 times with rewritten inputs and an eager call on the same state between replays: every replayed and + eager call gives the bits of the same call on the reference state.""" + alone = _alone(engine) + state = engine.create() + layers = [state.layer(*layer_weights(li)) for li in range(LAYERS)] + static = [tuple(t.clone() for t in routed(engine.m, 100)) for _ in range(LAYERS)] + outs = [torch.empty(engine.m, H, dtype=torch.bfloat16, device="cuda") for _ in range(LAYERS)] + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + for li in range(LAYERS): + engine.call(layers[li], static[li], outs[li]) + for r in range(TOKEN_SETS): + for li in range(LAYERS): + for dst, src in zip(static[li], routed(engine.m, 100 + r)): + dst.copy_(src) + graph.replay() + for li in range(LAYERS): + assert same(outs[li], alone[li, r]), ("replay", r, li) + eager_set = (r + 1) % TOKEN_SETS + assert same( + engine.call(layers[1], routed(engine.m, 100 + eager_set)), alone[1, eager_set] + ), ("eager", r) + + +def check_two_states(engine): + """Two states, each with its own workspace, called in an irregular order (A A B A B B A B A): every call gives the + bits of the same call on the reference state, and each state's epochs count only its own calls.""" + alone = _alone(engine) + states = {"A": engine.create(), "B": engine.create()} + layers = {name: st.layer(*layer_weights(0)) for name, st in states.items()} + counts = {"A": 0, "B": 0} + for i, name in enumerate("AABABBABA"): + si = i % TOKEN_SETS + assert same(engine.call(layers[name], routed(engine.m, 100 + si)), alone[0, si]), (i, name) + counts[name] += 1 + for name, st in states.items(): + assert bool((st.epochs == counts[name]).all()), name + + +def check_epoch_wrap(engine, start): + """The CTAs' epochs (int32, + 1 per call; their parity picks the count words) across the int32 wrap: the state + preset to where ~2^31 calls leave it (every epoch at ``start``, the next call's count words zero). Each call gives + the bits of the same call on the reference state, the epochs wrap to -2^31 and keep counting, and every call + leaves the next call's count words zero.""" + alone = _alone(engine) + state = engine.create() + layer = state.layer(*layer_weights(0)) + state.counts.zero_() + state.epochs.fill_(start) + for c in range(TOKEN_SETS + 1): + si = c % TOKEN_SETS + y = engine.call(layer, routed(engine.m, 100 + si)) + ep = (start + c + 1 + 2**31) % 2**32 - 2**31 # int32 two's complement + assert same(y, alone[0, si]), c + assert bool((state.epochs == ep).all()), (c, ep) + assert bool((engine.next_counts(state) == 0).all()), (c, ep) + + +def check_create(engine): + """``create`` refuses to run under capture, compiles the plain build and every push build it is given, and a + created state's first call can be captured.""" + stream = torch.cuda.Stream() + graph = torch.cuda.CUDAGraph() + raised = False + with torch.cuda.stream(stream): + try: + with torch.cuda.graph(graph, stream=stream): + engine.create() + except RuntimeError as exc: + raised = "before CUDA-graph capture" in str(exc) + assert raised, "create must refuse to run under capture" + state = engine.create(push=((4, 1), (16, 4))) + assert state.compiled + assert all(state.push_compiled(slots, copies) for slots, copies in ((4, 1), (16, 4))) + layer = state.layer(*layer_weights(0)) + tokens = tuple(t.clone() for t in routed(engine.m, 100)) + out = torch.empty(engine.m, H, dtype=torch.bfloat16, device="cuda") + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + engine.call(layer, tokens, out) + graph.replay() + alone = _alone(engine) + assert same(out, alone[0, 0]) diff --git a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_m1.py b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_m1.py new file mode 100644 index 000000000000..ba3fc2eff9d8 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_m1.py @@ -0,0 +1,44 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the k3_moe_m1 catalog entry (one GPU): single calls against the fp32 reference, and the call +sequences over its caller-owned K3MoeM1State that its contract certifies. The checks are in ``_k3_moe_engines.py`` +beside this file.""" + +import _k3_moe_engines as engines +import pytest + +pytestmark = pytest.mark.skipif(not engines.is_sm100(), reason="k3_moe_m1 needs sm_100") + +ENGINES = [engines.Engine("k3_moe_m1", 1), engines.Engine("k3_moe_m1", 2)] +IDS = ["one_token", "two_tokens"] + + +@pytest.mark.parametrize("engine", ENGINES, ids=IDS) +def test_single_calls(engine): + engines.check_single_calls(engine) + + +@pytest.mark.parametrize("engine", ENGINES, ids=IDS) +def test_call_sequences(engine): + engines.check_call_sequences(engine) + + +@pytest.mark.parametrize("engine", ENGINES, ids=IDS) +def test_capture_replay(engine): + engines.check_capture_replay(engine) + + +@pytest.mark.parametrize("engine", ENGINES, ids=IDS) +def test_two_states(engine): + engines.check_two_states(engine) + + +@pytest.mark.parametrize("start", [2**31 - 2, 2**31 - 1]) +@pytest.mark.parametrize("engine", ENGINES, ids=IDS) +def test_epoch_wrap(engine, start): + engines.check_epoch_wrap(engine, start) + + +@pytest.mark.parametrize("engine", ENGINES, ids=IDS) +def test_create(engine): + engines.check_create(engine) diff --git a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_m2.py b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_m2.py new file mode 100644 index 000000000000..3472349f997a --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_m2.py @@ -0,0 +1,37 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the k3_moe_m2 catalog entry (one GPU): single calls against the fp32 reference, and the call +sequences over its caller-owned K3MoeM2State that its contract certifies. The checks are in ``_k3_moe_engines.py`` +beside this file.""" + +import _k3_moe_engines as engines +import pytest + +pytestmark = pytest.mark.skipif(not engines.is_sm100(), reason="k3_moe_m2 needs sm_100") + +ENGINE = engines.Engine("k3_moe_m2", 2) + + +def test_single_calls(): + engines.check_single_calls(ENGINE) + + +def test_call_sequences(): + engines.check_call_sequences(ENGINE) + + +def test_capture_replay(): + engines.check_capture_replay(ENGINE) + + +def test_two_states(): + engines.check_two_states(ENGINE) + + +@pytest.mark.parametrize("start", [2**31 - 2, 2**31 - 1]) +def test_epoch_wrap(start): + engines.check_epoch_wrap(ENGINE, start) + + +def test_create(): + engines.check_create(ENGINE) From a04ec2e87b02a7330fb7fea0e75ca6aff50c129a Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 21:16:16 -0700 Subject: [PATCH 095/161] [None][test] Kimi K3 kernel lints: the k3_moe_m1 and k3_moe_m2 kernels test_k3_tcgen05_fences.py reads the two weight-stream kernels' sources too: every tcgen05 operation after a wait follows fence::after_thread_sync, and every arrive or barrier after a TMEM load wait follows fence::before_thread_sync. Signed-off-by: Vasanth Sabavat --- .../_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py index 6abcb564c141..14edccd65cc6 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py @@ -33,6 +33,8 @@ "tensorrt_llm._torch.cute_dsl_kernels.k3_decode_gemv.k3_decode_gemv_kernel", "tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.k3_moe_front", "tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.k3_moe_kernel", + "tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.k3_moe_m1_kernel", + "tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.k3_moe_m2_kernel", "tensorrt_llm._torch.cute_dsl_kernels.k3_sandwich.k3_sandwich_kernel", ] From f0fd94cfbd0a08b57bb7faec8208257bc2e18b1f Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 21:38:26 -0700 Subject: [PATCH 096/161] [None][feat] modeling_v2 catalog: moe/k3_moe's push form k3_moe_push: the M <= 8 partial stored into a slot of every rank's K3LatentExchange for comm/k3_latent_reduce instead of returned, through trtllm::k3_moe's push build. The contract gains the push form's semantics, signature and arguments, the exchange's call order and its preconditions. test_k3_moe_push.py's k3_moe_push check pushes into the 16-slot exchange through the entry (K3MoeLayer.push still fills the run's exchange). Signed-off-by: Vasanth Sabavat --- .../modeling_v2/catalog/moe/k3_moe.md | 49 ++++++++++++++++--- .../modeling_v2/catalog/moe/k3_moe.py | 35 ++++++++++++- .../kimi_k3/test_k3_moe_push.py | 11 +++-- 3 files changed, 84 insertions(+), 11 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md index ed3ac4838778..450b5c36e027 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md @@ -7,7 +7,7 @@ receipts: **Wraps** `torch.ops.trtllm.k3_moe` (one call), on caller-owned state: a `K3MoeState` (up to 8 tokens) or `K3MoeWideState` (up to 64) and one `K3MoeLayer` per MoE layer; for a head_flags state, also the TP group's -`K3MoeHeadWorkspace`. +`K3MoeHeadWorkspace`; for the push form (`k3_moe_push`), also the TP group's `K3LatentExchange`. ## Semantics @@ -50,10 +50,22 @@ the wide build: With `moe/k3_moe_front` as the producer (certified in that entry's 4-rank matrix): `y` within the op-catalog gates of the stock runner on the front's own routing and latent, and on a head_flags state bit for bit the plain state's. +**Push form.** `k3_moe_push` (`M` <= 8, a `K3MoeState`) computes the same partial and, instead of returning it, +stores token `t`'s row into slot `slot` (default `exchange.rank`) of half `exchange.flags[0] & 1` of every rank's +`K3LatentExchange` (int32 `[2][8][slots][1792]`: bf16 pairs, `0x80000000` empty, -0.0 stored as +0.0, zero rows when +nothing is routed here) through its multicast mapping: the kernel's fused all-reduce in its push-only mode. It reads +the half after its grid-dependency wait and writes no flags word; `comm/k3_latent_reduce` sums the slots in the MNNVL +one-shot's order, empties the half it read and advances the count. Certified at 4 ranks, `M` 1, 3 and 8, after +`moe/k3_route_quant` and after `moe/k3_moe_front`: a push and its reduce equal `MNNVLAllReduce`'s one-shot of the +plain partials bit for bit, on every rank and run to run, and leave the exchange empty with its count advanced; the +same into a 16-slot exchange filled 4 slots per rank (the one-shot's 16-slot order), and with the exchange's call +count across the int32 wrap. + Fusion boundary. Inside: the grouping of (expert, token) pairs, FC1, SiTU, the MXFP8 intermediate, FC2, the routing-weighted combine. Outside: the routing and the latent's MXFP8 quantization (`moe/k3_route_quant` or `moe/k3_moe_front`); the sum of the routed partials over the ranks that hold the other experts and intermediate -slices (the routed-latent all-reduce); the latent-up projection; the shared experts; the residual. +slices (the routed-latent all-reduce, or the push form plus `comm/k3_latent_reduce`); the latent-up projection; the +shared experts; the residual. ## Signature @@ -68,13 +80,26 @@ def k3_moe( head: Optional[K3MoeHeadWorkspace] = None, out: Optional[torch.Tensor] = None, ) -> torch.Tensor + +def k3_moe_push( + x_fp8: torch.Tensor, + x_sf: torch.Tensor, + topk_ids: torch.Tensor, + topk_weights: torch.Tensor, + local_expert_offset: int, + layer: K3MoeLayer, + exchange: K3LatentExchange, + slot: Optional[int] = None, + head: Optional[K3MoeHeadWorkspace] = None, +) -> None ``` The entry passes the layer's weight buffers and counters and its state's scratch and build options to the op: `trtllm::k3_moe(x_fp8, x_sf, topk_ids, topk_weights, w3_w1_weight, w3_w1_weight_scale, w2_weight, w2_weight_scale, c, cs, part, counters, local_expert_offset, num_local, num_ctas, m_max, use_pdl, head_ready=None, head_flags=None, -out=None)`, with `mutates_args = (c, cs, part, counters, head_ready, head_flags, out)`: every buffer the kernel -writes. The wrapper module re-exports `K3MoeState`, `K3MoeWideState`, `K3MoeLayer`, `K3MoeHeadWorkspace` and +out=None, exchange_uc=None, exchange_mc=None, exchange_flags=None, exchange_slot=0)`, with `mutates_args = (c, cs, +part, counters, head_ready, head_flags, out, exchange_uc, exchange_mc)`: every buffer the kernel writes. The push form +passes the exchange's `uc`, `mc` and `flags` and the slot (`K3MoeLayer.push` makes the same call). The wrapper module re-exports `K3MoeState`, `K3MoeWideState`, `K3MoeLayer`, `K3MoeHeadWorkspace` and `is_supported`. ### Certified arguments @@ -89,7 +114,9 @@ writes. The wrapper module re-exports `K3MoeState`, `K3MoeWideState`, `K3MoeLaye | `layer` | a `K3MoeLayer` of a plain or head_flags `K3MoeState`, or of a `K3MoeWideState` | — | — | — | | `head` | `None`; for and only for a head_flags state's layers, the TP group's `K3MoeHeadWorkspace` (certified in `moe/k3_moe_front`'s matrix) | — | — | — | | `out` | `None`, or `[>= M, 3584]`: the call writes `out[:M]` and returns an empty `[0, 3584]`; rows past `M` untouched (certified at `M` 3, 8 and 9, 64 on the two builds) | bf16 | contiguous | CUDA | -| returns | `y [M, 3584]`, or `[0, 3584]` with `out` | bf16 | contiguous, newly allocated | the inputs' device | +| `exchange` (push form) | the TP group's `K3LatentExchange` (4 ranks certified, and a 16-slot exchange) | — | — | — | +| `slot` (push form) | `None` (this rank's) or a slot of the exchange (certified: each of 16 slots, 4 per rank) | Python int | — | — | +| returns | `y [M, 3584]`, or `[0, 3584]` with `out`; `None` for the push form | bf16 | contiguous, newly allocated | the inputs' device | The four inputs are `moe/k3_route_quant`'s outputs for `M` tokens (certified) or `moe/k3_moe_front`'s (certified in its matrix). State construction certified: `K3MoeState(device, 768, 224)` (`head_flags` False or True; `use_pdl` and @@ -140,7 +167,9 @@ launch (certified). The op itself picks the head_flags build by whether `head_re check, and the same one in `K3MoeLayer`, is what ties the build to the state. A head_flags state's calls pair with the front's on one `K3MoeHeadWorkspace`: the front call before each must publish the workspace's ready words (the `moe/k3_moe_front` entry with `publish=True`), and each publishing front call must be followed by exactly one such -`k3_moe` call on that workspace (*Preconditions*). +`k3_moe` call on that workspace (*Preconditions*). The push form shares the state with the plain calls. Its exchange +is a separate, collective object (`comm/k3_latent_reduce`'s *State*): each push of `M` tokens is followed by one +reduce of `M` tokens on that exchange before the next push, in the same order on every rank. **Call-order invariant.** The calls on all layers of one state run one after the other in one stream order. Every call needs the slab armed and its layer's counters at zero, which only the end of the previous call on the state @@ -236,6 +265,11 @@ Besides the state objects (explicit arguments): - Either build lets its own dependents launch early (plain builds right after the wait, the head_flags build at launch): a consumer of `y` must wait for `k3_moe`'s grid (`griddepcontrol.wait`, or plain stream order) before reading it. +- Push form: `M` <= 8 on a `K3MoeState`, no `out`; the exchange's words are int32 `[2][8][slots][1792]` (this + rank's and multicast) with int32 flags, and `slot` is one of its slots (`ValueError` otherwise, before any launch). + The exchange belongs to this rank's TP group, and one reduce of `M` tokens on it follows each push before the next, + on every rank in the same order. The push build compiles on its first push, which must be eager like the plain + build's. - Calls may be captured once the build's first call has run eagerly: certified with the captured step above, and in `moe/k3_moe_front`'s matrix on both K3MoeStates. @@ -247,6 +281,9 @@ Besides the state objects (explicit arguments): `moe/k3_route_quant`. The head_flags build and the front as producer: 4 ranks of one GB200 tray in `tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py` (entry point `moe/test_modeling_v2_k3_moe_front_op_matrix.py`); its 16-rank receipt is pending with `moe/k3_moe_front`'s. +- The push form: 4 ranks of one GB200 tray in `tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py` + (`k3_moe_push`: `K3MoeLayer.push` into a 4-slot exchange, this entry's `k3_moe_push` into a 16-slot one); its + 16-rank receipt is pending. - References: an fp64 reference over the dequantized experts (from the checkpoint-format tensors) and the stock path. The kernel tests (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py`, `test_k3_moe_wide.py`) remain the exhaustive numerics; this entry's test copies their references. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.py index 96ace8337365..ddd44bed1a4d 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.py @@ -3,12 +3,15 @@ """Kimi K3's routed experts at decode size: ``trtllm::k3_moe``, this rank's routed partial from the persistent CuTe DSL kernel (FC1 + SiTU + FC2 with the routing-weighted combine over the TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers, read in place), on caller-owned state: a :class:`K3MoeState` (up to 8 tokens) or :class:`K3MoeWideState` (up to 64) and one -:class:`K3MoeLayer` per MoE layer. Its inputs are the outputs of ``moe/k3_route_quant`` or ``moe/k3_moe_front``.""" +:class:`K3MoeLayer` per MoE layer. Its inputs are the outputs of ``moe/k3_route_quant`` or ``moe/k3_moe_front``. The +push form stores the partial into every rank's ``K3LatentExchange`` for ``comm/k3_latent_reduce`` instead.""" from typing import Optional import torch +from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.latent_op import K3LatentExchange + # The state types (they launch nothing per call); importing the op module registers trtllm::k3_moe. is_supported reads # metadata only. from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.op import ( @@ -26,6 +29,7 @@ "K3MoeWideState", "is_supported", "k3_moe", + "k3_moe_push", ] @@ -59,3 +63,32 @@ def k3_moe( local_expert_offset, state.num_local, state.num_ctas, state.m_max, state.use_pdl, None if head is None else head.ready, None if head is None else head.flags, out, ) # fmt: skip + + +def k3_moe_push( + x_fp8: torch.Tensor, + x_sf: torch.Tensor, + topk_ids: torch.Tensor, + topk_weights: torch.Tensor, + local_expert_offset: int, + layer: K3MoeLayer, + exchange: K3LatentExchange, + slot: Optional[int] = None, + head: Optional[K3MoeHeadWorkspace] = None, +) -> None: + """The push form of :func:`k3_moe` for ``M <= 8`` on a K3MoeState: the same partial, stored into slot ``slot`` + (default ``exchange.rank``) of every rank's ``exchange`` (a TP group's ``K3LatentExchange``) instead of returned. + One ``comm/k3_latent_reduce`` of the M tokens on that exchange must follow before the next push, on every rank in + the same order. ``head`` as in :func:`k3_moe`. Writes the state's slab (left armed) and partial rows, the layer's + counters (left zero), and every rank's exchange.""" + state = layer.state + if (head is not None) != state.head_flags: + raise ValueError( + "k3_moe_push: head is given for, and only for, the layers of a head_flags K3MoeState" + ) + torch.ops.trtllm.k3_moe( + x_fp8, x_sf, topk_ids, topk_weights, *layer.weights, state.c, state.cs, state.part, layer.counters, + local_expert_offset, state.num_local, state.num_ctas, state.m_max, state.use_pdl, + None if head is None else head.ready, None if head is None else head.flags, None, + exchange.uc, exchange.mc, exchange.flags, exchange.rank if slot is None else slot, + ) # fmt: skip diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py index 66712186a648..c0a8176c33e1 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py @@ -19,9 +19,10 @@ push trtllm::k3_moe_m1 at 1 and 2 tokens and trtllm::k3_moe_m2 at 2 (Layer.push), on tokens routed by trtllm::k3_route_quant; -k3_moe_push trtllm::k3_moe's push build (K3MoeLayer.push) at 1, 3 and 8 tokens, after trtllm::k3_route_quant and - after trtllm::k3_moe_front. The front's head and shared experts are sharded over the run's ranks; its - shared activation must be the same bits in every call of a set. +k3_moe_push trtllm::k3_moe's push build at 1, 3 and 8 tokens, after trtllm::k3_route_quant and after + trtllm::k3_moe_front: K3MoeLayer.push into the run's exchange, the moe/k3_moe entry's k3_moe_push into + the 16-slot one. The front's head and shared experts are sharded over the run's ranks; its shared + activation must be the same bits in every call of a set. Per op, token count and routing set, against the plain call (Layer.__call__, K3MoeLayer.__call__ after the same producer), whose partial must be nonzero: @@ -338,6 +339,8 @@ def _k3_moe_routed(producer, inputs, front): def check_k3_moe_push(ctx): + from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.k3_moe import k3_moe_push + results = [] front = _front(ctx) layer = _k3_moe_layer(_experts(20260928 + ctx.rank)) @@ -362,7 +365,7 @@ def check_k3_moe_push(ctx): # TP16's receive side: this rank's partial in slots 4 r .. 4 r + 3, one push each. for c in range(copies16): *routed, pushed = _k3_moe_routed(producer, inputs, front) - layer.push(*routed, 0, ex16, ctx.rank * copies16 + c) + k3_moe_push(*routed, 0, layer, ex16, ctx.rank * copies16 + c) pushed_shared.append(pushed) got16 = ex16.reduce(m) state16 = ex16.state_ok(ctx) From ebfd38001c6327c8f51b827736ccd5c47214a15b Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 21:25:24 -0700 Subject: [PATCH 097/161] [None][fix] modeling_v2 Kimi K3 target: declare the long GEMV and SiTU ops it calls REQUIRED_TRTLLM_OPS names every trtllm op the target calls, and test_modeling_v2_target_contract.py asserts each one is registered in the build. The decode GEMV sites call trtllm::k3_ctm_gemv_long (MLA's [W_a; W_g], KDA's fused projection and the dense MLP's two projections), and the dense MLP calls trtllm::k3_situ_mul, but neither was listed: a build without them passed the check and failed at the first decode step. Signed-off-by: Vasanth Sabavat --- .../kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index ca8876712688..7f50539b2ea1 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -146,9 +146,11 @@ "k3_kda_verify", "k3_mla_qkv", "k3_mla_attn_vb_out", - # The decode path's GEMVs, LM head and embedding (decode_gemv.py). + # The decode path's GEMVs, the dense MLP's activation, the LM head and the embedding (decode_gemv.py). "k3_decode_gemv", "k3_ctm_gemv_wide", + "k3_ctm_gemv_long", + "k3_situ_mul", "k3_head_gemv", "k3_embed_norm", "allgather", From f32084b8806ba25d43178adfab19effc2ed7907c Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 23:02:00 -0700 Subject: [PATCH 098/161] [None][fix] modeling_v2 Kimi K3 tp16_moetp16ep1: declare the long GEMV and SiTU ops it calls Route A's ebfd38001c, carried into route B's copy: its decode path is the same code (layer 0's dense MLP runs k3_ctm_gemv_long and k3_situ_mul), and test_modeling_v2_kimi_k3_drift.py requires the copy to equal route A outside its route B blocks. Signed-off-by: Vasanth Sabavat --- .../kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py index 638ccc2bb741..f357e1b4854c 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py @@ -151,9 +151,11 @@ # <<< route B "k3_mla_qkv", "k3_mla_attn_vb_out", - # The decode path's GEMVs, LM head and embedding (decode_gemv.py). + # The decode path's GEMVs, the dense MLP's activation, the LM head and the embedding (decode_gemv.py). "k3_decode_gemv", "k3_ctm_gemv_wide", + "k3_ctm_gemv_long", + "k3_situ_mul", "k3_head_gemv", "k3_embed_norm", "allgather", From 4bf23781fb4dbcff8587206a904a38c60849452f Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 23:55:45 -0700 Subject: [PATCH 099/161] [None][fix] DSpark GQA drafter: keep the checkpoint's trained mask-token embedding GQADSparkForCausalLM takes the target's embedding in load_weights_from_target_model, so the DFlash worker built every masked block slot from the target's row for the mask token, which the target never trained. A drafter checkpoint that ships its own embedding has trained that row. load_weights now keeps it as mask_token_embedding, and the worker's noise block (dflash_noise_block_embedding, the same computation moved into a function) uses it for every masked slot when the drafter has one for the mask id in use; slot 0 stays the bonus token's shared lookup. Every rank still embeds the mask id through embed_tokens.forward, as before. Drafts change; under greedy verification the target's tokens do not. Plain DFlash and MLA DSpark drafters set no mask_token_embedding and are unchanged. Signed-off-by: Vasanth Sabavat --- tensorrt_llm/_torch/models/modeling_dspark.py | 15 +++++++ tensorrt_llm/_torch/speculative/dflash.py | 43 ++++++++++++++---- .../test_kimi_k3_dspark_semantics.py | 44 +++++++++++++++++++ 3 files changed, 94 insertions(+), 8 deletions(-) diff --git a/tensorrt_llm/_torch/models/modeling_dspark.py b/tensorrt_llm/_torch/models/modeling_dspark.py index 10aa541a6ff4..1855c56b4312 100644 --- a/tensorrt_llm/_torch/models/modeling_dspark.py +++ b/tensorrt_llm/_torch/models/modeling_dspark.py @@ -2673,8 +2673,23 @@ def __init__(self, draft_config, *, dflash_attention_backend: str = "AUTO"): def load_weights(self, weights: Dict, weight_mapper=None, **kwargs): """Take the DSpark head weights, then hand the rest to DFlash.""" weights, _ = self._take_dspark_head_weights(weights) + self._keep_trained_mask_embedding(weights) return super().load_weights(weights, weight_mapper=weight_mapper, **kwargs) + def _keep_trained_mask_embedding(self, weights: Dict) -> None: + """Keep the checkpoint's own embedding row for the mask token. + + This drafter takes the target's embedding (``load_weights_from_target_model``). A drafter checkpoint that ships + an embedding has trained the mask token's row, which the target's embedding does not have: the block decode + reads ``mask_token_embedding`` for every masked slot instead of the shared lookup. + """ + names = [k for k in ("embed_tokens.weight", "model.embed_tokens.weight") if k in weights] + if not names: + return + weight = weights[names[0]] + row = weight[self.mask_token_id : self.mask_token_id + 1] + self.mask_token_embedding = torch.as_tensor(row).reshape(-1).to("cuda") + class MLADSparkForCausalLM(_DSparkHeadMixin, DFlashForCausalLM): """DSpark drafter on an MLA-shaped backbone, from a standalone checkpoint. diff --git a/tensorrt_llm/_torch/speculative/dflash.py b/tensorrt_llm/_torch/speculative/dflash.py index d760e15fbc78..eb1204d31b26 100644 --- a/tensorrt_llm/_torch/speculative/dflash.py +++ b/tensorrt_llm/_torch/speculative/dflash.py @@ -432,6 +432,32 @@ def dflash_draft_slot_ids( return (request_bases.unsqueeze(1) + first_slot + offsets.unsqueeze(0)).flatten() +def dflash_noise_block_embedding( + embed_tokens: nn.Module, + bonus: torch.Tensor, + mask_token_id: int, + block_size: int, + trained_mask_embedding: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """The drafter's input block per request, [len(bonus), block_size, hidden]: slot 0 the bonus token's embedding, + every other slot the mask token's. + + The mask row is the drafter's own trained one when it has one (``trained_mask_embedding``); otherwise the shared + lookup's. Both rows go through ``embed_tokens.forward`` (NOT ``.weight[...]``) so TP-sharded vocabs mask out + ranks that don't own the token id and all-reduce, the mask row included, so every rank makes the same call. + """ + num_gens = bonus.shape[0] + mask_tok = torch.full((1,), int(mask_token_id), dtype=torch.long, device=bonus.device) + combined_embed = embed_tokens(torch.cat([bonus, mask_tok], dim=0)) + embed_bonus = combined_embed[:num_gens] + embed_mask = combined_embed[num_gens] + if trained_mask_embedding is not None: + embed_mask = trained_mask_embedding.to(embed_mask.dtype) + noise_embed_2d = embed_mask.expand(num_gens, block_size, -1).clone() + noise_embed_2d[:, 0, :] = embed_bonus + return noise_embed_2d + + def dflash_position_ceiling(max_ctx: int, block_size: int, max_draft_len: int) -> int: """Positions a drafter that indexes absolute positions must be able to encode. @@ -1999,14 +2025,15 @@ def prepare_1st_drafter_inputs( query_position_ids = ctx_len_now.unsqueeze(1) + j_block.unsqueeze(0) ctx_position_ids = ctx_len_gen.unsqueeze(1) + offsets_kp1.unsqueeze(0) - # Go through embed_tokens.forward (NOT .weight[...]) so TP-sharded - # vocabs mask out ranks that don't own the token id and all-reduce. - mask_tok = torch.full((1,), int(mask_token_id), dtype=torch.long, device="cuda") - combined_embed = embed_tokens(torch.cat([bonus, mask_tok], dim=0)) - embed_bonus = combined_embed[:num_gens] - embed_mask = combined_embed[num_gens] - noise_embed_2d = embed_mask.expand(num_gens, query_tokens_per_req, -1).clone() - noise_embed_2d[:, 0, :] = embed_bonus + # The drafter's own trained mask row, if it kept one for the mask id in use. + trained_mask_embedding = ( + getattr(draft_model, "mask_token_embedding", None) + if mask_token_id == getattr(draft_model, "mask_token_id", None) + else None + ) + noise_embed_2d = dflash_noise_block_embedding( + embed_tokens, bonus, mask_token_id, query_tokens_per_req, trained_mask_embedding + ) # Accumulate new accepted features into context buffers if has_target_features: diff --git a/tests/unittest/_torch/speculative/hw_agnostic/test_kimi_k3_dspark_semantics.py b/tests/unittest/_torch/speculative/hw_agnostic/test_kimi_k3_dspark_semantics.py index 70a6688c4fa2..0939bfdca109 100644 --- a/tests/unittest/_torch/speculative/hw_agnostic/test_kimi_k3_dspark_semantics.py +++ b/tests/unittest/_torch/speculative/hw_agnostic/test_kimi_k3_dspark_semantics.py @@ -454,6 +454,50 @@ def test_dspark_drafter_loads_head_weights_and_parses_config(): assert drafter.confidence_proj_bias is not None +@needs_gpu +def test_gqa_dspark_keeps_its_trained_mask_row_over_the_target_embedding(): + """The GQA drafter takes the target's embedding (load_weights_from_target_model), but a checkpoint that ships its + own embedding trained the mask token's row, which the target's embedding does not have. The drafter keeps that row + as ``mask_token_embedding``, and the worker's noise block puts it in every masked slot; slot 0 stays the bonus + token's embedding from the shared lookup. + """ + h = TINY["hidden_size"] + weights = _tiny_weights() + g = torch.Generator().manual_seed(23) + weights["embed_tokens.weight"] = (torch.randn(VOCAB, h, generator=g) * 0.05).to(torch.bfloat16) + drafter = _build_drafter(True, weights) + target_embed = torch.nn.Embedding(VOCAB, h).to("cuda", torch.bfloat16) + target = SimpleNamespace( + model=SimpleNamespace(embed_tokens=target_embed), + lm_head=torch.nn.Linear(h, VOCAB, bias=False), + ) + drafter.load_weights_from_target_model(target) + + mask_id = drafter.mask_token_id + trained = weights["embed_tokens.weight"][mask_id] + assert not torch.equal(target_embed.weight[mask_id].detach().cpu(), trained) + torch.testing.assert_close(drafter.mask_token_embedding.cpu(), trained, rtol=0, atol=0) + + from tensorrt_llm._torch.speculative.dflash import dflash_noise_block_embedding + + bonus = torch.tensor([3, 5], dtype=torch.long, device="cuda") + block = 4 + with torch.no_grad(): + noise = dflash_noise_block_embedding( + drafter.draft_model_full.model.embed_tokens, + bonus, + mask_id, + block, + drafter.mask_token_embedding, + ) + assert tuple(noise.shape) == (2, block, h) + torch.testing.assert_close( + noise[:, 0].cpu(), target_embed.weight[bonus].detach().cpu(), rtol=0, atol=0 + ) + for j in range(1, block): + torch.testing.assert_close(noise[:, j].cpu(), trained.expand(2, -1), rtol=0, atol=0) + + @needs_gpu def test_published_drafter_spelling_activates_the_heads(): """Both public K3 DSpark checkpoints load with their heads live. From 05ab11d9e8965f05ae63195b4cb2f2d8d56c7c1a Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 00:15:44 -0700 Subject: [PATCH 100/161] [None][doc] Kimi K3 collective contracts: the 16-rank runs passed The seven stateful matrices passed at 16 ranks on four GB200 trays of one rack (fabric handles), every check on every rank, in a recorded run of this branch: the contracts say so in place of "pending". The MoE front's 16-rank run is the one that takes the half-tile head path. Signed-off-by: Vasanth Sabavat --- .../modeling_v2/catalog/comm/k3_latent_reduce.md | 3 ++- .../modeling_v2/catalog/comm/k3_sandwich_oproj.md | 9 +++++---- .../modeling_v2/catalog/comm/k3_sandwich_plain.md | 4 ++-- .../modeling_v2/catalog/comm/k3_sandwich_tail.md | 4 ++-- .../modeling_v2/catalog/comm/mnnvl_allgather_split.md | 3 ++- .../modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md | 4 ++-- .../_experimental/modeling_v2/catalog/moe/k3_moe.md | 2 +- .../modeling_v2/catalog/moe/k3_moe_front.md | 8 ++++---- 8 files changed, 20 insertions(+), 17 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md index 7f171431f01c..cd0bb97d1d82 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_latent_reduce.md @@ -224,6 +224,7 @@ scheduling, not results (the op's statement; the test runs the default). time. - World size: the matrix takes `--world-size` and `--launcher` (`mpirun` on one node, `srun` across nodes). Kimi K3 runs the op over 16 ranks on four trays, where the kernel sums the ranks in two chunks of 8 and uses 14 CTAs per - token row; a 4-rank run reaches neither. The 16-rank receipt is pending; `W` = 8 is not run. + token row; a 4-rank run reaches neither. At 16 ranks (four GB200 trays of one rack, fabric handles) the matrix + passes, every check on every rank, in a recorded run; `W` = 8 is not run. - The op writes `lat_uc` (it empties words) and `lat_flags`, and its schema declares both mutable; the producers' writes through `mc` are theirs to declare. The compile cache is a module-level dict (result-neutral, above). diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md index 316378ce619c..8c186a7916b0 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md @@ -186,10 +186,11 @@ it changes scheduling, not results. `normed` against an fp32 reference within 2e-2 of its largest magnitude. - World sizes: the matrix takes `--world-size` and `--launcher` (`mpirun` on one tray, `srun` across trays). The kernel sums ranks in chunks of 8, so a run at `W` <= 8 exercises one chunk; Kimi K3 runs `W` = 16 over four trays. - This entry's 16-rank receipt is pending. The op's kernel test - (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py`; this op bit for bit against `o_proj` and - the MNNVL one-shot) passed every case at 16 ranks on four trays in a recorded run: the kernel's record, not this - entry's receipt. + At 16 ranks (four GB200 trays of one rack, fabric handles) this entry's matrix passes, every check on every rank, in + a recorded run; CI runs it at 4. The op's kernel test + (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py`; this op bit for bit against `o_proj` and the + MNNVL one-shot) passed every case at 16 ranks on four trays in a recorded run: the kernel's record, not this entry's + receipt. - State: `mutates_args` names `ws_uc`, `ws_mc` and `ws_flags` — every call pushes through `ws_mc` into every rank's `ws_uc`, empties the words it read in `ws_uc` and advances `ws_flags` — and `x_slab`. The op module keeps no workspace registry: the caller passes the object it created. The compile cache is the documented process-wide cache diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md index a21cc254f931..48527bd2b6df 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_plain.md @@ -171,8 +171,8 @@ forms are two kernels, each compiled (seconds) on its first call, which must be is compared bit for bit in both forms. `normed` is compared with torch's fp32 RMSNorm within 2e-2 of its largest magnitude (the one-shot rounds the squares to bf16 and sums them in its own tree). SiLU-and-mul's rounding on general inputs is the kernel test's to certify (bit for bit against `k3_ctm_gemv_swiglu`), not this matrix's. -- World sizes: as for `k3_sandwich_oproj`; this entry's 16-rank receipt is pending. The op's kernel test - (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py`) passed every case at 16 ranks on four +- World sizes: as for `k3_sandwich_oproj`; at 16 ranks this entry's matrix passes in a recorded run. The op's kernel + test (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py`) passed every case at 16 ranks on four trays in a recorded run: the kernel's record, not this entry's receipt. - State: `mutates_args` names `ws_uc`, `ws_mc` and `ws_flags`, every buffer the op writes. The compile cache is the documented process-wide cache above. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md index 3d3a1d2e7df0..e1a703e7d932 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md @@ -210,8 +210,8 @@ outside CUDA-graph capture first". `updated_out` is not part of the key. The cac within 2e-2 of an fp32 reference; the tapped `updated` bit for bit with the returned one; every output bitwise across the ranks. The negative control's wrong pairing puts more than half the elements outside the 8e-3 bound, the largest error over 10 times it. -- World sizes: as for `k3_sandwich_oproj`; this entry's 16-rank receipt is pending. The op's kernel test - (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py`) passed every case at 16 ranks on four +- World sizes: as for `k3_sandwich_oproj`; at 16 ranks this entry's matrix passes in a recorded run. The op's kernel + test (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py`) passed every case at 16 ranks on four trays in a recorded run: the kernel's record, not this entry's receipt. - State: `mutates_args` names every buffer the op can write: `ws_uc`, `ws_mc`, `ws_flags`, `x_slab`, `lat_uc`, `lat_flags`, `tap` and `updated_out`. The compile cache is the documented process-wide cache above. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md index 04b225a7d4c6..c1068e80cf91 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_allgather_split.md @@ -198,7 +198,8 @@ depend on it. call sequences on real state (layers x steps, calls queued without host synchronization, capture + replay, two objects interleaved) plus a negative control; every written buffer named in the schema; the matrix takes `--world-size` and `--launcher` (`mpirun` on one node, `srun` across trays) and CI runs it at 4 ranks on one GB200 - tray; one `MnnvlWorkspace` shared by every MNNVL entry of the TP group. The 16-rank receipt is pending. + tray; one `MnnvlWorkspace` shared by every MNNVL entry of the TP group. At 16 ranks (four GB200 trays of one + rack, fabric handles) the matrix passes, every check on every rank, in a recorded run. - Not exercised: `B` = 0 or `F` = 0 (the op accepts both), denormal and non-finite values in the bf16 columns, an accepted call of more than 64 tokens, the fake implementation. - In the model today the call is `MNNVLAllReduce.allgather_split(input, bf16_columns)` on `MNNVLAllReduce`'s workspace diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md index 7fce1c4c2e6e..002ae345b220 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/mnnvl_fusion_allreduce.md @@ -239,8 +239,8 @@ follows `T`, `H` and the device's SM count, and in the fused form it sets the or the per-call time of back-to-back captured calls is within -0.56 / +0.12 us of theirs (noise 0.10 us). Every MNNVL kernel's SASS changes, since the kernel parameters gained `earlyTrigger`; in the one-shot kernel's, the dependents' launch sits right after the grid-dependency wait. -- At `W` = 16 the one-shot kernel adds the ranks in two chunks of 8, a branch a 4-rank run never reaches. The 16-rank - receipt is pending. +- At `W` = 16 the one-shot kernel adds the ranks in two chunks of 8, a branch a 4-rank run never reaches. At 16 ranks + (four GB200 trays of one rack, fabric handles) the matrix passes, every check on every rank, in a recorded run. - Kimi K3's calls (its decode path, not this test): the model sets every `MNNVLAllReduce` of the target, its LM head and a drafter to `one_shot_max_bytes` = 4 MiB (`DECODE_AR_ONE_SHOT_MAX_BYTES`, against main's 1 MiB), and a wide decode step (9 to 64 tokens) passes 1 MiB per call (`WIDE_AR_ONE_SHOT_MAX_BYTES`). Plain: the routed-latent diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md index ed3ac4838778..807c129f9fa5 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md @@ -246,7 +246,7 @@ Besides the state objects (explicit arguments): intermediate 768), with random checkpoint-format MXFP4 experts put through TRT-LLM's own loader, routed by `moe/k3_route_quant`. The head_flags build and the front as producer: 4 ranks of one GB200 tray in `tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py` (entry point - `moe/test_modeling_v2_k3_moe_front_op_matrix.py`); its 16-rank receipt is pending with `moe/k3_moe_front`'s. + `moe/test_modeling_v2_k3_moe_front_op_matrix.py`), which also passes at 16 ranks in a recorded run. - References: an fp64 reference over the dequantized experts (from the checkpoint-format tensors) and the stock path. The kernel tests (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py`, `test_k3_moe_wide.py`) remain the exhaustive numerics; this entry's test copies their references. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md index d1c317863ff1..25122aab73da 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_front.md @@ -211,9 +211,9 @@ Besides `workspace` (an explicit argument): GEMV clusters of the 13 left. At `W` 8 (560 rows, 9 half-tiles) and `W` 4 (1120 rows, 18) they do not fit, and the head runs as 5 and 9 tiles of 128 rows. The CPU test `tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front_geometry.py` checks this choice and that - `front_weight`'s rows are the rows each plan reads. So this entry's 4-rank matrix runs 128-row tiles only; the - half-tile path's 16-rank record is the kernel test `test_k3_moe_front.py` run as 16 processes (the model's TP16 - shapes), not this matrix. + `front_weight`'s rows are the rows each plan reads. So this entry's 4-rank matrix runs 128-row tiles only; at 16 + ranks (a recorded run on four GB200 trays of one rack) the matrix runs the half-tile path and passes, every check + on every rank, as does the kernel test `test_k3_moe_front.py` run as 16 processes. ## Preconditions @@ -237,7 +237,7 @@ Besides `workspace` (an explicit argument): - Certified path: 4 ranks of one GB200 tray (sm_100), one rank per GPU, the head sharded over those 4 ranks (1120 rows per rank, 128-row tiles), the shared activation at TP16's per-rank width (384). Kimi K3 TP16 shards the head over 16 ranks on four trays (280 rows per rank, the half-tile geometry: *Metadata consumed*); the matrix takes - `--world-size` and `--launcher`, and its 16-rank receipt is pending. + `--world-size` and `--launcher`; at 16 ranks it passes in a recorded run (above). - Test: `tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py` (rank body), collected by `tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_front_op_matrix.py`. It also certifies `moe/k3_moe` behind the front, on a plain and on a head_flags `K3MoeState`. From a55d13f0fc2ac2456466fceed98d8c6ab4bbc0de Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 00:27:35 -0700 Subject: [PATCH 101/161] [None][doc] Kimi K3 MoE push form: the 16-rank runs passed test_k3_moe_push.py's three checks (the m1 / m2 engines' push, k3_moe's push through the MoE front, and the multi-layer sequences) passed at 16 ranks on four GB200 trays of one rack, one exchange slot per rank, every check on every rank, in a recorded run of this branch. moe/k3_moe says so in place of "pending"; moe/k3_moe_m1 and moe/k3_moe_m2 add it to their note on 4-rank runs. Signed-off-by: Vasanth Sabavat --- .../_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md | 4 ++-- .../_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m1.md | 3 ++- .../_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m2.md | 3 ++- 3 files changed, 6 insertions(+), 4 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md index 450b5c36e027..f9e9d66c4043 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.md @@ -282,8 +282,8 @@ Besides the state objects (explicit arguments): `tests/unittest/_torch/modeling_v2/comm/_k3_moe_front_op_matrix.py` (entry point `moe/test_modeling_v2_k3_moe_front_op_matrix.py`); its 16-rank receipt is pending with `moe/k3_moe_front`'s. - The push form: 4 ranks of one GB200 tray in `tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py` - (`k3_moe_push`: `K3MoeLayer.push` into a 4-slot exchange, this entry's `k3_moe_push` into a 16-slot one); its - 16-rank receipt is pending. + (`k3_moe_push`: `K3MoeLayer.push` into a 4-slot exchange, this entry's `k3_moe_push` into a 16-slot one). At 16 + ranks (four GB200 trays of one rack, one exchange slot per rank) it passes in a recorded run. - References: an fp64 reference over the dequantized experts (from the checkpoint-format tensors) and the stock path. The kernel tests (`tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_fused_moe.py`, `test_k3_moe_wide.py`) remain the exhaustive numerics; this entry's test copies their references. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m1.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m1.md index a506c0704cd7..d6ae99cae588 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m1.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m1.md @@ -160,4 +160,5 @@ cached per configuration in the process (code only, no state). TP4 x EP4 shapes, M 1); the push form at 4 ranks in `test_k3_moe_push.py` (bit-exact against `MNNVLAllReduce`, a 16-slot exchange filled 4 slots per rank, the exchange's count across the int32 wrap). - Kimi K3 runs the push form over 16 ranks; a 4-rank run fills a 16-slot exchange only by repeating each rank's - partial. + partial. At 16 ranks (four GB200 trays of one rack, one distinct partial per slot) `test_k3_moe_push.py` passes in + a recorded run. diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m2.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m2.md index 07aaa35f6334..de4b95292553 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m2.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe_m2.md @@ -161,4 +161,5 @@ cached per configuration in the process (code only, no state). tokens routed to the same and to disjoint experts); the push form at 4 ranks in `test_k3_moe_push.py` (bit-exact against `MNNVLAllReduce`, a 16-slot exchange filled 4 slots per rank, the exchange's count across the int32 wrap). - Kimi K3 runs the push form over 16 ranks; a 4-rank run fills a 16-slot exchange only by repeating each rank's - partial. + partial. At 16 ranks (four GB200 trays of one rack, one distinct partial per slot) `test_k3_moe_push.py` passes in + a recorded run. From 83924f82818c284e9ad4d26f8710e6fe7b22cb01 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 01:02:12 -0700 Subject: [PATCH 102/161] [None][fix] Kimi K3 k3_moe: compile on routing that requires grad trtllm::k3_moe compiles its kernel on its first call from DLPack views of its arguments, and DLPack refuses tensors that require grad. A model that warms the op up outside inference mode, with a routing bias that is an nn.Parameter, passes routing weights that require grad, and that first call raised BufferError. The views are now taken of detached tensors, as in the other Kimi K3 CuTe ops. test_routing_from_a_parameter_that_requires_grad makes that warm-up's calls (the compiling one, then a compiled one) and checks their bits against the same call in inference mode. Signed-off-by: Vasanth Sabavat --- .../cute_dsl_kernels/k3_fused_moe/op.py | 3 +- .../moe/test_modeling_v2_k3_moe.py | 28 +++++++++++++++++++ 2 files changed, 30 insertions(+), 1 deletion(-) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py index 69f70d200ba9..5ecc8b6b22c4 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_fused_moe/op.py @@ -73,7 +73,8 @@ def _kernel_module(config: dict): def _view(t: torch.Tensor, align: int, leading_dim: int, element_type=None): from cutlass.cute.runtime import from_dlpack - v = from_dlpack(t, assumed_align=align).mark_layout_dynamic(leading_dim=leading_dim) + # DLPack refuses tensors that require grad. + v = from_dlpack(t.detach(), assumed_align=align).mark_layout_dynamic(leading_dim=leading_dim) if element_type is not None: v.element_type = element_type return v diff --git a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe.py b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe.py index 4e85a0f2696f..2aa78156038c 100644 --- a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe.py +++ b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe.py @@ -607,6 +607,34 @@ def test_k3_moe_wide_single_call(case, m): assert det and rearmed +def test_routing_from_a_parameter_that_requires_grad(): + """A model's first k3_moe calls (a load-time warm-up) can run outside inference mode, routed with a bias that is an + ``nn.Parameter``: the routing weights then require grad. The op hands the kernel tensors DLPack exports (it refuses + tensors that require grad), so the first call compiles, and it and a compiled call return the bits of the same call + in inference mode.""" + proc, _, _ = _experts() + logits, x = _draw(8100, DECODE_MAX) + want = _call(_decode()[1][0], logits, x) + got = [] + with torch.inference_mode(False): + state = K3MoeState(_device(), I_TP, E_LOCAL) + layer = state.layer(*_weights(proc)) + bias = torch.nn.Parameter(_bias().clone()) + with _cold_k3_moe_cache(): + for _ in range(2): # the compiling call, then a compiled one + ids, weights, x_fp8, x_sf = k3_route_quant( + logits.clone(), bias, x.clone(), RSF, early_trigger=state.use_pdl + ) + assert weights.requires_grad, ( + "the routing weights do not require grad: the check misses its case" + ) + got.append(k3_moe(x_fp8, x_sf, ids, weights, OFFSET, layer).detach()) + torch.cuda.synchronize() + assert all(_same(y, want) for y in got), ( + "a call on routing that requires grad differs from the same call in inference mode" + ) + + @pytest.mark.parametrize("build,m", [("m8", 3), ("m8", 8), ("m64", 9), ("m64", 64)]) def test_k3_moe_out_buffer(build, m): """``out``: the call writes out[:M], the bits of the call without ``out``, and returns an empty [0, 3584] tensor. From eefd22557eee5353f9fbc8da6327b299aea84f38 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 21:26:49 -0700 Subject: [PATCH 103/161] [None][feat] modeling_v2 Kimi K3 target: a decode step's attention-residual epilogues on the fused kernels The attention-residual selection and the RMSNorm after it (with the add of the attention output after the attention) run as one fused kernel, attn_res_rmsnorm_fwd or attn_res_add_rmsnorm_fwd, up to a token ceiling and as the unfused add -> attn_res -> RMSNorm chain above it. The generic path keeps the built-in model's ceiling (KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS, 1 by default). A step decode_step classifies now takes the fused kernels up to one token tile (8 tokens), and a wide decode step up to 32 tokens: the pre- and post-attention epilogues of every layer and the final norm's. These are the ceilings of the decode layout the target's kernels come from. The fused kernel does not round bit for bit like the chain, so a classified step now computes what that layout computes. Signed-off-by: Vasanth Sabavat --- .../modeling.py | 84 +++++++++++++------ .../test_modeling_v2_kimi_k3_decode_step.py | 20 ++++- 2 files changed, 78 insertions(+), 26 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index bc8cc3da3c4d..1bc6af7adf35 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -40,11 +40,12 @@ then `attention/k3_mla_attn_vb_out` (the attention, v_b and the output gate in one launch). The projections around them run on the decode GEMV sites of `decode_gemv.py` (the [W_a; W_g] and KDA verify-row -projections, `o_proj` on every classified step), as do the LM head, the embedding and layer 0's dense MLP. The state -those kernels share (the KDA projection's Lamport buffers, the MLA attention workspace, the decode GEMVs' state) -lives in typed objects this target creates in `post_load_weights`, before any graph capture. The MoE front and routed -experts, the sandwiches and the residual epilogues come with their own entries; until then they run the generic path -on every step. +projections, `o_proj` on every classified step), as do the LM head, the embedding and layer 0's dense MLP. A +classified step's attention-residual epilogues (the selection and the RMSNorm after it) take the fused kernels up to +one token tile, 32 tokens on a wide decode step. The state those kernels share (the KDA projection's Lamport buffers, +the MLA attention workspace, the decode GEMVs' state) lives in typed objects this target creates in +`post_load_weights`, before any graph capture. The MoE front and routed experts and the sandwiches come with their +own entries; until then they run the generic path on every step. **What this target asserts rather than adapts**: SM 10.0; the topology above; the MXFP4 checkpoint's quantization (W4A16_MXFP4 with no per-layer declarations, so the routed experts run the W4A8_MXFP4_MXFP8 default and the excluded @@ -448,20 +449,20 @@ def _persistent_attn_res_applicable(M: int, H: int, N: int) -> bool: return H == 7168 and 2 <= N <= 9 -def _use_persistent_attn_res(M: int, H: int, N: int) -> bool: +def _use_persistent_attn_res(M: int, H: int, N: int, max_fused_tokens: int) -> bool: """Pick between the two fused kernels for this call site. ``persistent`` takes the persistent kernel at every shape it implements; - ``split`` takes it only above ``KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS`` tokens, - which stands in for the prefill/decode boundary. Shapes the persistent - kernel does not implement fall through to the caller's existing gate and - land on the unfused path. + ``split`` takes it only above ``max_fused_tokens`` tokens, which stands in + for the prefill/decode boundary. Shapes the persistent kernel does not + implement fall through to the caller's existing gate and land on the + unfused path. """ if not _persistent_attn_res_applicable(M, H, N): return False if _ATTN_RES_TOPOLOGY == "persistent": return True - return _ATTN_RES_TOPOLOGY == "split" and M > _FUSED_ATTN_RES_MAX_TOKENS + return _ATTN_RES_TOPOLOGY == "split" and M > max_fused_tokens def _apply_attn_res_fused( @@ -528,8 +529,12 @@ def _apply_attn_res_rmsnorm_fused( proj: nn.Linear, norm: KimiK3RMSNorm, output_norm: nn.Module, + max_fused_tokens: Optional[int] = None, ) -> Optional[torch.Tensor]: - """Fuse attention-residual mixing with its immediately following norm.""" + """Fuse attention-residual mixing with its immediately following norm (at most ``max_fused_tokens`` tokens, + default ``KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS``).""" + if max_fused_tokens is None: + max_fused_tokens = _FUSED_ATTN_RES_MAX_TOKENS if ( prefix_sum.dtype is not torch.bfloat16 or not prefix_sum.is_cuda @@ -539,10 +544,10 @@ def _apply_attn_res_rmsnorm_fused( M, H = prefix_sum.shape K = int(block_residual.shape[0]) N = K + 1 - # The fused path is taken for M <= _FUSED_ATTN_RES_MAX_TOKENS, H == 7168 and - # N <= 12, which is the measured window; larger token counts have not been - # measured and fall back to the unfused add + attn_res_fwd + RMSNorm path. - if _use_persistent_attn_res(M, H, N): + # The fused path is taken for M <= max_fused_tokens, H == 7168 and N <= 12, + # which is the measured window; larger token counts have not been measured + # and fall back to the unfused add + attn_res_fwd + RMSNorm path. + if _use_persistent_attn_res(M, H, N, max_fused_tokens): try: persistent_op = torch.ops.trtllm.attn_res_add_rmsnorm_persistent_fwd except (AttributeError, RuntimeError): @@ -560,7 +565,7 @@ def _apply_attn_res_rmsnorm_fused( _note_attn_res_fusion("attn_res+norm/persistent", True, M, H, N) return output.reshape(M, H) - if M > _FUSED_ATTN_RES_MAX_TOKENS or H != 7168 or N > 12: + if M > max_fused_tokens or H != 7168 or N > 12: _note_attn_res_fusion("attn_res+norm", False, M, H, N) return None try: @@ -589,8 +594,10 @@ def _apply_attn_res_add_rmsnorm_fused( proj: nn.Linear, norm: KimiK3RMSNorm, output_norm: nn.Module, + max_fused_tokens: Optional[int] = None, ) -> Optional[Tuple[torch.Tensor, torch.Tensor]]: - """Fuse ``prefix_sum + addend``, attention-residual, and trailing norm. + """Fuse ``prefix_sum + addend``, attention-residual, and trailing norm (``max_fused_tokens`` as in + ``_apply_attn_res_rmsnorm_fused``). The production residual add produces a BF16 tensor that remains live across the following MLP. The kernel therefore returns that materialized, @@ -598,6 +605,8 @@ def _apply_attn_res_add_rmsnorm_fused( while avoiding a separate add launch and a re-read of the intermediate by attention-residual selection. """ + if max_fused_tokens is None: + max_fused_tokens = _FUSED_ATTN_RES_MAX_TOKENS if ( prefix_sum.dtype is not torch.bfloat16 or addend.dtype is not torch.bfloat16 @@ -611,7 +620,7 @@ def _apply_attn_res_add_rmsnorm_fused( K = int(block_residual.shape[0]) N = K + 1 # Same measured window as _apply_attn_res_rmsnorm_fused above. - if _use_persistent_attn_res(M, H, N): + if _use_persistent_attn_res(M, H, N, max_fused_tokens): try: persistent_op = torch.ops.trtllm.attn_res_add_rmsnorm_persistent_fwd except (AttributeError, RuntimeError): @@ -629,7 +638,7 @@ def _apply_attn_res_add_rmsnorm_fused( _note_attn_res_fusion("add+attn_res+norm/persistent", True, M, H, N) return updated_prefix_sum.reshape(M, H), output.reshape(M, H) - if M > _FUSED_ATTN_RES_MAX_TOKENS or H != 7168 or N > 12: + if M > max_fused_tokens or H != 7168 or N > 12: _note_attn_res_fusion("add+attn_res+norm", False, M, H, N) return None try: @@ -687,10 +696,14 @@ def _apply_attn_res_and_rmsnorm( proj: nn.Linear, norm: KimiK3RMSNorm, output_norm: nn.Module, + max_fused_tokens: Optional[int] = None, ) -> torch.Tensor: - """Apply attention-residual selection and the next RMSNorm.""" + """Apply attention-residual selection and the next RMSNorm. ``max_fused_tokens``: the largest token count the + fused kernel takes (default ``KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS``).""" if _FUSED_ATTN_RES_ENABLED: - fused = _apply_attn_res_rmsnorm_fused(prefix_sum, block_residual, proj, norm, output_norm) + fused = _apply_attn_res_rmsnorm_fused( + prefix_sum, block_residual, proj, norm, output_norm, max_fused_tokens + ) if fused is not None: return fused return output_norm(_apply_attn_res(prefix_sum, block_residual, proj, norm)) @@ -703,20 +716,36 @@ def _apply_attn_res_add_and_rmsnorm( proj: nn.Linear, norm: KimiK3RMSNorm, output_norm: nn.Module, + max_fused_tokens: Optional[int] = None, ) -> Tuple[torch.Tensor, torch.Tensor]: - """Add an attention output to the running residual, then select and norm.""" + """Add an attention output to the running residual, then select and norm (``max_fused_tokens`` as in + ``_apply_attn_res_and_rmsnorm``).""" if _FUSED_ATTN_RES_ENABLED: fused = _apply_attn_res_add_rmsnorm_fused( - prefix_sum, addend, block_residual, proj, norm, output_norm + prefix_sum, addend, block_residual, proj, norm, output_norm, max_fused_tokens ) if fused is not None: return fused updated_prefix_sum = prefix_sum + addend return updated_prefix_sum, _apply_attn_res_and_rmsnorm( - updated_prefix_sum, block_residual, proj, norm, output_norm + updated_prefix_sum, block_residual, proj, norm, output_norm, max_fused_tokens ) +# A wide decode step's residual epilogues take the fused add + attn_res + RMSNorm kernels up to this many tokens, where +# they are faster than the add -> attn_res -> RMSNorm chain. +_WIDE_ATTN_RES_MAX_TOKENS = 32 + + +def _attn_res_max_tokens(step: Optional[DecodeStep]) -> Optional[int]: + """The most tokens of ``step`` the fused attn_res kernels take: one token tile (``DECODE_MAX_TOKENS``) on a step + ``decode_step`` classifies, ``_WIDE_ATTN_RES_MAX_TOKENS`` on a wide decode step; None on any other step (the + generic path's ``KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS``).""" + if step is None: + return None + return _WIDE_ATTN_RES_MAX_TOKENS if step.wide else DECODE_MAX_TOKENS + + _K3_ROUTED_EXPERT_KEY_SUFFIXES = ("block_sparse_moe.experts", "mlp.experts") @@ -1421,6 +1450,7 @@ def forward( """ prefix_sum = hidden_states valid_block_residual = block_residual[:num_snapshots] + attn_res_max_tokens = _attn_res_max_tokens(step) if prenormed: assert num_snapshots == 0 and self.layer_idx % self.attn_res_block_size == 0 @@ -1448,6 +1478,7 @@ def forward( self.self_attention_res_proj, self.self_attention_res_norm, self.input_layernorm, + attn_res_max_tokens, ) else: hidden_states = self.input_layernorm(hidden_states) @@ -1471,6 +1502,7 @@ def forward( self.mlp_res_proj, self.mlp_res_norm, self.post_attention_layernorm, + attn_res_max_tokens, ) else: prefix_sum, hidden_states = _apply_attn_res_add_and_rmsnorm( @@ -1480,6 +1512,7 @@ def forward( self.mlp_res_proj, self.mlp_res_norm, self.post_attention_layernorm, + attn_res_max_tokens, ) if self.is_moe: hidden_states = self.block_sparse_moe( @@ -1689,6 +1722,7 @@ def forward( self.output_attn_res_proj, self.output_attn_res_norm, self.norm, + _attn_res_max_tokens(step), ) diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_step.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_step.py index 345f803fac8f..4f221e60a1dd 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_step.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_step.py @@ -1,7 +1,7 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 """How the Kimi K3 target classifies a step for its decode kernels (host-side, no GPU): ``decode_step`` of -``kimi_k3_mxfp4__sm_100__tp16_moetp4ep4``.""" +``kimi_k3_mxfp4__sm_100__tp16_moetp4ep4``, and the attention-residual epilogue ceiling it sets.""" import types @@ -10,6 +10,7 @@ from tensorrt_llm._torch._experimental.modeling_v2.models.kimi_k3_vl.kimi_k3_mxfp4__sm_100__tp16_moetp4ep4.modeling import ( # noqa: E501 DecodeStep, + _attn_res_max_tokens, decode_step, ) @@ -67,3 +68,20 @@ def test_short_steps_are_small_only(seq_lens, num_contexts, rows): ) def test_other_steps_take_the_generic_path(seq_lens, num_contexts, rows): assert decode_step(_metadata(seq_lens, num_contexts), rows) is None + + +@pytest.mark.parametrize( + "step,ceiling", + [ + (None, None), # the generic path: KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS + (DecodeStep(5), 8), # a short prefill or mixed step + (DecodeStep(8, 8, 1), 8), # a decode step of one token tile + (DecodeStep(6, 1, 6), 8), + (DecodeStep(16, 2, 8), 32), # a wide decode step + (DecodeStep(64, 8, 8), 32), + ], +) +def test_attn_res_epilogue_ceiling(step, ceiling): + """The fused attn_res kernels take a classified step up to one token tile and a wide decode step up to 32 tokens; + any other step keeps the generic path's ceiling.""" + assert _attn_res_max_tokens(step) == ceiling From 93f0ad3f8fdfe9ab6b407a919f8d06496f831444 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 21:32:15 -0700 Subject: [PATCH 104/161] [None][feat] modeling_v2 Kimi K3 target: the post-attention all-reduce with the residual update in its epilogue A layer's post-attention step is the TP all-reduce of the attention output, then the residual update: the add to the running prefix sum, the attention-residual selection and the post-attention RMSNorm. On a step of at most 16 tokens, wide decode steps aside, the attention now hands over its unreduced o_proj partial (reduce_output=False) and one comm/mnnvl_allreduce_attn_res call runs both: the one-shot MNNVL all-reduce with the residual update as its epilogue. It returns the post-attention norm's output and the updated prefix sum. decode_comm.py holds the state, K3DecodeComm: one MnnvlWorkspace of the TP group with 4 MiB buffers (16 tokens of 7168 from 16 ranks fit one). The target creates it in post_load_weights, collectively and before any CUDA-graph capture, where every attention all-reduce runs over MNNVL, and hands it to every layer. Whether a step takes the fused call depends on its token count and kind only, which every rank shares, so every rank makes the same calls on the workspace in the same order. Signed-off-by: Vasanth Sabavat --- .../decode_comm.py | 96 +++++++++++++++++++ .../modeling.py | 62 +++++++++--- 2 files changed, 143 insertions(+), 15 deletions(-) create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py new file mode 100644 index 000000000000..d5723b047465 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py @@ -0,0 +1,96 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""The decode path's residual-update collectives on the catalog's Kimi K3 MNNVL entry. + +A layer's post-attention step adds the attention output to the running prefix sum, selects the attention residual +and applies the post-attention RMSNorm. Under TP the attention output is a sum over the TP group, so the step is that +all-reduce followed by the residual update. On a step of at most `AR_ATTN_RES_MAX_TOKENS` tokens (wide decode steps +aside) the attention hands over its unreduced o_proj partial, and `K3DecodeComm.allreduce_attn_res` runs both as one +`comm/mnnvl_allreduce_attn_res` call: the one-shot MNNVL all-reduce with the residual update as its epilogue. + +The state is one `MnnvlWorkspace` of the TP group (`K3DecodeComm.create`): collective over the group and eager, built +by the target in `post_load_weights` before any CUDA-graph capture. Every rank must make the same calls on it in the +same order. Whether a step takes the call is decided from its token count and kind alone, which every rank of the +group shares. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Optional, Tuple + +import torch +from torch import nn + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.mnnvl_allreduce_attn_res import ( + mnnvl_allreduce_attn_res, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.mnnvl_workspace import ( + MnnvlWorkspace, +) + +# The most tokens of a step whose post-attention all-reduce runs one-shot with the residual update in its epilogue. +AR_ATTN_RES_MAX_TOKENS = 16 + +# The workspace's buffer size: a one-shot call pushes T x 7168 bf16 from each of 16 ranks, 3.5 MiB at T = 16. +MNNVL_BUFFER_BYTES = 4 << 20 + + +def _eps(norm: nn.Module) -> float: + """The epsilon of a KimiK3RMSNorm (``eps``) or a stock RMSNorm (``variance_epsilon``).""" + return float(norm.eps if hasattr(norm, "eps") else norm.variance_epsilon) + + +@dataclass(eq=False) +class K3DecodeComm: + """The decode path's collective state for the TP group: the `MnnvlWorkspace` the post-attention all-reduces run + on. Built by `create`; owned by the target and shared by every layer.""" + + mnnvl: MnnvlWorkspace + + @classmethod + def create(cls, mapping) -> "K3DecodeComm": + """The state for ``mapping``'s TP group. Collective: every rank of the group calls it at the same point, + eagerly, before any CUDA-graph capture; it fails on every rank or on none (``MnnvlWorkspace.create``).""" + return cls(MnnvlWorkspace.create(mapping, MNNVL_BUFFER_BYTES)) + + def takes_post_attention(self, hidden_states: torch.Tensor, step) -> bool: + """Whether ``allreduce_attn_res`` runs the post-attention step of a layer whose attention input is + ``hidden_states``: at most `AR_ATTN_RES_MAX_TOKENS` bf16 rows of a hidden size the op takes, on any step but + a wide decode step (``step.wide``), whose post-attention update stays the all-reduce and the fused add + + attn_res + RMSNorm.""" + if step is not None and step.wide: + return False + if hidden_states.dim() != 2 or hidden_states.dtype != torch.bfloat16: + return False + rows, hidden = hidden_states.shape + return ( + hidden % 1024 == 0 + and hidden <= 8192 + and 0 < rows <= min(AR_ATTN_RES_MAX_TOKENS, self.mnnvl.max_one_shot_tokens(hidden)) + ) + + def allreduce_attn_res( + self, + partial: torch.Tensor, + prefix_sum: Optional[torch.Tensor], + block_residual: torch.Tensor, + res_proj: nn.Module, + res_norm: nn.Module, + out_norm: nn.Module, + ) -> Tuple[torch.Tensor, torch.Tensor]: + """``(normed, updated)`` in one ``comm/mnnvl_allreduce_attn_res`` call: ``updated = prefix_sum + + allreduce(partial)`` (the sum alone without ``prefix_sum``), ``normed = out_norm(attn_res(block_residual..., + updated))``, the attention residual selected with ``res_proj`` and ``res_norm``. ``block_residual`` holds the + valid snapshots, ``[S, T, H]``.""" + return mnnvl_allreduce_attn_res( + partial.contiguous(), + prefix_sum, + block_residual, + res_proj.weight.reshape(-1), + res_norm.weight, + out_norm.weight, + _eps(res_norm), + _eps(out_norm), + self.mnnvl, + ) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index 1bc6af7adf35..ab050d0e53d4 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -42,10 +42,12 @@ The projections around them run on the decode GEMV sites of `decode_gemv.py` (the [W_a; W_g] and KDA verify-row projections, `o_proj` on every classified step), as do the LM head, the embedding and layer 0's dense MLP. A classified step's attention-residual epilogues (the selection and the RMSNorm after it) take the fused kernels up to -one token tile, 32 tokens on a wide decode step. The state those kernels share (the KDA projection's Lamport buffers, -the MLA attention workspace, the decode GEMVs' state) lives in typed objects this target creates in -`post_load_weights`, before any graph capture. The MoE front and routed experts and the sandwiches come with their -own entries; until then they run the generic path on every step. +one token tile, 32 tokens on a wide decode step. On any step of at most 16 tokens but a wide one, the post-attention +all-reduce carries the residual update in its epilogue (`decode_comm.py`, on `comm/mnnvl_allreduce_attn_res`). The +state those kernels share (the KDA projection's Lamport buffers, the MLA attention workspace, the decode GEMVs' state, +the TP group's MNNVL workspace) lives in typed objects this target creates in `post_load_weights`, before any graph +capture. The MoE front and routed experts and the sandwiches come with their own entries; until then they run the +generic path on every step. **What this target asserts rather than adapts**: SM 10.0; the topology above; the MXFP4 checkpoint's quantization (W4A16_MXFP4 with no per-layer declarations, so the routed experts run the W4A8_MXFP4_MXFP8 default and the excluded @@ -119,6 +121,7 @@ from tensorrt_llm.mapping import Mapping from tensorrt_llm.models.modeling_utils import QuantAlgo, QuantConfig +from . import decode_comm as _decode_comm from . import decode_gemv as _decode_gemv from . import weights as _weights @@ -156,6 +159,8 @@ "k3_head_gemv", "k3_embed_norm", "allgather", + # The decode path's collectives (decode_comm.py). + "mnnvl_allreduce_attn_res", ) #: The engine surface the first forward checks before this target relies on it: per object, the attributes read. @@ -1417,6 +1422,8 @@ def __init__( self.mlp_res_norm = KimiK3RMSNorm(cfg.hidden_size, eps=cfg.rms_norm_eps, dtype=dtype) self.self_attention_res_proj = nn.Linear(cfg.hidden_size, 1, bias=False, dtype=dtype) self.mlp_res_proj = nn.Linear(cfg.hidden_size, 1, bias=False, dtype=dtype) + # The decode path's collectives (decode_comm.py), set by the target's post_load_weights. + self.decode_comm: Optional[_decode_comm.K3DecodeComm] = None def forward( self, @@ -1447,6 +1454,11 @@ def forward( ``prenormed`` (layer 0 on a decode step): ``hidden_states`` already is this layer's input norm, and the layer's input, the step's embedding, already is in ``block_residual[0]`` (``K3DecodeGemvs.embed_norm``). + + The post-attention step of at most ``AR_ATTN_RES_MAX_TOKENS`` tokens + runs the attention's all-reduce and the residual update in one + collective (``K3DecodeComm.allreduce_attn_res``) once the target has + built its ``decode_comm``. """ prefix_sum = hidden_states valid_block_residual = block_residual[:num_snapshots] @@ -1489,13 +1501,20 @@ def forward( num_snapshots += 1 valid_block_residual = block_residual[:num_snapshots] prefix_sum = None - if self.is_kda: - hidden_states = self.linear_attn(hidden_states, attn_metadata, step=step) - else: - hidden_states = self.self_attn(hidden_states, attn_metadata, step=step) - - if prefix_sum is None: - prefix_sum = hidden_states + attention = self.linear_attn if self.is_kda else self.self_attn + comm = self.decode_comm + if comm is not None and comm.takes_post_attention(hidden_states, step): + partial = attention(hidden_states, attn_metadata, step=step, reduce_output=False) + hidden_states, prefix_sum = comm.allreduce_attn_res( + partial, + prefix_sum, + valid_block_residual, + self.mlp_res_proj, + self.mlp_res_norm, + self.post_attention_layernorm, + ) + elif prefix_sum is None: + prefix_sum = attention(hidden_states, attn_metadata, step=step) hidden_states = _apply_attn_res_and_rmsnorm( prefix_sum, valid_block_residual, @@ -1507,7 +1526,7 @@ def forward( else: prefix_sum, hidden_states = _apply_attn_res_add_and_rmsnorm( prefix_sum, - hidden_states, + attention(hidden_states, attn_metadata, step=step), valid_block_residual, self.mlp_res_proj, self.mlp_res_norm, @@ -2349,8 +2368,10 @@ def cache_derived_state(self) -> None: def post_load_weights(self) -> None: """The state the decode kernels share, built once per device before any CUDA-graph capture and handed to - every layer of its kind: the KDA projection's Lamport buffers and the MLA decode attention's workspace; and - the decode GEMVs' state (built by ``cache_derived_state``) handed to every attention module.""" + every layer of its kind: the KDA projection's Lamport buffers and the MLA decode attention's workspace; the + decode GEMVs' state (built by ``cache_derived_state``) handed to every attention module; and, where every + attention all-reduce runs over MNNVL, the TP group's collective state (``K3DecodeComm``, collective: every + rank builds it here) handed to every layer.""" super().post_load_weights() kda = [layer.linear_attn for layer in self.model.layers if layer.is_kda] mla = [layer.self_attn.mixer for layer in self.model.layers if not layer.is_kda] @@ -2365,10 +2386,21 @@ def post_load_weights(self) -> None: module.k3_workspace = workspace for module in kda + mla: module.decode_gemvs = self.model.decode_gemvs + comm = None + if all(layer._mnnvl_allreduce() is not None for layer in self.model.layers): + comm = _decode_comm.K3DecodeComm.create(self.model_config.mapping) + for layer in self.model.layers: + layer.decode_comm = comm logger.info( "Kimi K3 decode kernels: KDA on k3_kda_decode_attn, k3_kda_attn and k3_kda_verify " f"({sum(m.takes_k3_kernels for m in kda)} / {len(kda)} layers take them), MLA on k3_mla_qkv and " - f"k3_mla_attn_vb_out ({len(mla)} layers)" + f"k3_mla_attn_vb_out ({len(mla)} layers); the post-attention all-reduce of at most " + f"{_decode_comm.AR_ATTN_RES_MAX_TOKENS} tokens " + + ( + "with the residual update (mnnvl_allreduce_attn_res)" + if comm is not None + else "unfused (an attention all-reduce does not run over MNNVL)" + ) ) def _check_step_contract(self, attn_metadata: AttentionMetadata) -> None: From d1f297a5defb747c4fc0ef4928d6e202458b918b Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 21:34:33 -0700 Subject: [PATCH 105/161] [None][feat] modeling_v2 Kimi K3 target: the post-attention step of a decode step as one sandwich kernel At most 8 tokens, where the attention runs its decode branch, o_proj, its TP all-reduce and the residual update now run as one comm/k3_sandwich_oproj kernel. The attention hands over its gated o_proj input (project_output=False). The kernel is bit for bit o_proj followed by comm/mnnvl_allreduce_attn_res, which keeps the other steps of at most 16 tokens. K3DecodeComm gains the TP group's K3SandwichWorkspace, created with the MNNVL workspace in post_load_weights. The kernel compiles there, with one call on a zero row, so no CUDA-graph capture compiles it; every rank makes that call. Whether the sandwich takes a layer is decided once the weights are final: a bias-free bf16 o_proj of the TP16 per-rank shape [7168, 768] and a snapshot bank within the kernel's candidate count. The attention's decode branch is asked for first (will_run_decode_branch): MLA's KV writes allow one attention run per step, and only that branch hands over the o_proj input. Signed-off-by: Vasanth Sabavat --- .../decode_comm.py | 131 +++++++++++++++--- .../modeling.py | 80 ++++++++--- 2 files changed, 166 insertions(+), 45 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py index d5723b047465..e536fc78738e 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py @@ -1,17 +1,22 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""The decode path's residual-update collectives on the catalog's Kimi K3 MNNVL entry. +"""The decode path's residual-update collectives on the catalog's Kimi K3 MNNVL and sandwich entries. A layer's post-attention step adds the attention output to the running prefix sum, selects the attention residual and applies the post-attention RMSNorm. Under TP the attention output is a sum over the TP group, so the step is that all-reduce followed by the residual update. On a step of at most `AR_ATTN_RES_MAX_TOKENS` tokens (wide decode steps -aside) the attention hands over its unreduced o_proj partial, and `K3DecodeComm.allreduce_attn_res` runs both as one -`comm/mnnvl_allreduce_attn_res` call: the one-shot MNNVL all-reduce with the residual update as its epilogue. - -The state is one `MnnvlWorkspace` of the TP group (`K3DecodeComm.create`): collective over the group and eager, built -by the target in `post_load_weights` before any CUDA-graph capture. Every rank must make the same calls on it in the -same order. Whether a step takes the call is decided from its token count and kind alone, which every rank of the -group shares. +aside) one collective runs both: + +* `K3DecodeComm.sandwich_oproj`, `comm/k3_sandwich_oproj`: o_proj, its all-reduce and the residual update in one + kernel, at most `SANDWICH_MAX_TOKENS` tokens of an o_proj of the TP16 per-rank shape [7168, 768]. The attention + hands over its gated o_proj input. +* `K3DecodeComm.allreduce_attn_res`, `comm/mnnvl_allreduce_attn_res`, everywhere else: the one-shot MNNVL all-reduce + of the attention's unreduced o_proj partial, with the residual update as its epilogue. + +The state is one `MnnvlWorkspace` and one `K3SandwichWorkspace` of the TP group (`K3DecodeComm.create`): collective +over the group and eager, built by the target in `post_load_weights` before any CUDA-graph capture. Every rank must +make the same calls on each in the same order. Which call a step takes is decided from its token count and kind and +from the load-time layout alone, which every rank of the group shares. """ from __future__ import annotations @@ -22,6 +27,10 @@ import torch from torch import nn +from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.k3_sandwich_oproj import ( + K3SandwichWorkspace, + k3_sandwich_oproj, +) from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.mnnvl_allreduce_attn_res import ( mnnvl_allreduce_attn_res, ) @@ -29,10 +38,16 @@ MnnvlWorkspace, ) +# The kernel's support predicate: metadata reads only. +from tensorrt_llm._torch.cute_dsl_kernels.k3_sandwich import op as _sandwich_op + # The most tokens of a step whose post-attention all-reduce runs one-shot with the residual update in its epilogue. AR_ATTN_RES_MAX_TOKENS = 16 -# The workspace's buffer size: a one-shot call pushes T x 7168 bf16 from each of 16 ranks, 3.5 MiB at T = 16. +# The most tokens the sandwich kernel takes (one token tile). +SANDWICH_MAX_TOKENS = _sandwich_op.MAX_TOKENS + +# The MNNVL workspace's buffer size: a one-shot call pushes T x 7168 bf16 from each of 16 ranks, 3.5 MiB at T = 16. MNNVL_BUFFER_BYTES = 4 << 20 @@ -41,24 +56,61 @@ def _eps(norm: nn.Module) -> float: return float(norm.eps if hasattr(norm, "eps") else norm.variance_epsilon) +def _res_args(res_proj: nn.Module, res_norm: nn.Module, out_norm: nn.Module) -> tuple: + """The residual update's weights and epsilons in the entries' order.""" + return ( + res_proj.weight.reshape(-1), + res_norm.weight, + out_norm.weight, + _eps(res_norm), + _eps(out_norm), + ) + + @dataclass(eq=False) class K3DecodeComm: - """The decode path's collective state for the TP group: the `MnnvlWorkspace` the post-attention all-reduces run - on. Built by `create`; owned by the target and shared by every layer.""" + """The decode path's collective state for the TP group: the `MnnvlWorkspace` and the `K3SandwichWorkspace` the + post-attention steps run on. Built by `create`; owned by the target and shared by every layer.""" mnnvl: MnnvlWorkspace + sandwich: K3SandwichWorkspace @classmethod - def create(cls, mapping) -> "K3DecodeComm": + def create(cls, mapping, oproj_weight: Optional[torch.Tensor] = None) -> "K3DecodeComm": """The state for ``mapping``'s TP group. Collective: every rank of the group calls it at the same point, - eagerly, before any CUDA-graph capture; it fails on every rank or on none (``MnnvlWorkspace.create``).""" - return cls(MnnvlWorkspace.create(mapping, MNNVL_BUFFER_BYTES)) + eagerly, before any CUDA-graph capture; each workspace fails on every rank or on none. + + ``oproj_weight``: an o_proj weight ``takes_oproj`` holds for. The sandwich kernel compiles here for its + shape, with one call on a zero row of a zero weight, so no capture compiles it; the call advances the + sandwich workspace on every rank alike.""" + state = cls( + MnnvlWorkspace.create(mapping, MNNVL_BUFFER_BYTES), + K3SandwichWorkspace.create(mapping), + ) + if oproj_weight is not None: + weight = torch.zeros_like(oproj_weight) + hidden = weight.shape[0] + ones = weight.new_ones(hidden) + k3_sandwich_oproj( + weight.new_zeros(1, weight.shape[1]), + weight, + None, + weight.new_zeros(0, 1, hidden), + weight.new_zeros(hidden), + ones, + ones, + 1e-6, + 1e-6, + state.sandwich, + ) + torch.cuda.synchronize(weight.device) + return state def takes_post_attention(self, hidden_states: torch.Tensor, step) -> bool: - """Whether ``allreduce_attn_res`` runs the post-attention step of a layer whose attention input is - ``hidden_states``: at most `AR_ATTN_RES_MAX_TOKENS` bf16 rows of a hidden size the op takes, on any step but - a wide decode step (``step.wide``), whose post-attention update stays the all-reduce and the fused add + - attn_res + RMSNorm.""" + """Whether one of this state's collectives runs the post-attention step of a layer whose attention input is + ``hidden_states``: at most `AR_ATTN_RES_MAX_TOKENS` bf16 rows of a hidden size the MNNVL entry takes, on any + step but a wide decode step (``step.wide``), whose post-attention update stays the all-reduce and the fused + add + attn_res + RMSNorm.""" if step is not None and step.wide: return False if hidden_states.dim() != 2 or hidden_states.dtype != torch.bfloat16: @@ -70,6 +122,21 @@ def takes_post_attention(self, hidden_states: torch.Tensor, step) -> bool: and 0 < rows <= min(AR_ATTN_RES_MAX_TOKENS, self.mnnvl.max_one_shot_tokens(hidden)) ) + @staticmethod + def takes_oproj(o_proj: nn.Module, max_snapshots: int) -> bool: + """Whether ``sandwich_oproj`` takes the layer of output projection ``o_proj`` on a step of at most + `SANDWICH_MAX_TOKENS` tokens, decided once the weights are final: a bias-free, contiguous bf16 weight of the + kernel's shape (the TP16 per-rank [7168, 768]) and a snapshot bank of at most ``max_snapshots`` rows the + kernel's candidate count holds.""" + weight = getattr(o_proj, "weight", None) + return ( + getattr(o_proj, "bias", None) is None + and isinstance(weight, torch.Tensor) + and weight.dim() == 2 + and weight.is_cuda + and _sandwich_op.supports(weight.new_empty((1, weight.shape[1])), weight, max_snapshots) + ) + def allreduce_attn_res( self, partial: torch.Tensor, @@ -87,10 +154,28 @@ def allreduce_attn_res( partial.contiguous(), prefix_sum, block_residual, - res_proj.weight.reshape(-1), - res_norm.weight, - out_norm.weight, - _eps(res_norm), - _eps(out_norm), + *_res_args(res_proj, res_norm, out_norm), self.mnnvl, ) + + def sandwich_oproj( + self, + core: torch.Tensor, + o_weight: torch.Tensor, + prefix_sum: Optional[torch.Tensor], + block_residual: torch.Tensor, + res_proj: nn.Module, + res_norm: nn.Module, + out_norm: nn.Module, + ) -> Tuple[torch.Tensor, torch.Tensor]: + """``(normed, updated)`` as ``allreduce_attn_res`` of ``core @ o_weight.T``, in one ``comm/k3_sandwich_oproj`` + call: ``core`` is this rank's gated o_proj input ``[T, 768]``, ``o_weight`` its o_proj slice. Bit for bit + o_proj followed by ``allreduce_attn_res``.""" + return k3_sandwich_oproj( + core.contiguous(), + o_weight, + prefix_sum, + block_residual, + *_res_args(res_proj, res_norm, out_norm), + self.sandwich, + ) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index ab050d0e53d4..90c10c94f4be 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -43,11 +43,12 @@ projections, `o_proj` on every classified step), as do the LM head, the embedding and layer 0's dense MLP. A classified step's attention-residual epilogues (the selection and the RMSNorm after it) take the fused kernels up to one token tile, 32 tokens on a wide decode step. On any step of at most 16 tokens but a wide one, the post-attention -all-reduce carries the residual update in its epilogue (`decode_comm.py`, on `comm/mnnvl_allreduce_attn_res`). The -state those kernels share (the KDA projection's Lamport buffers, the MLA attention workspace, the decode GEMVs' state, -the TP group's MNNVL workspace) lives in typed objects this target creates in `post_load_weights`, before any graph -capture. The MoE front and routed experts and the sandwiches come with their own entries; until then they run the -generic path on every step. +all-reduce carries the residual update (`decode_comm.py`): at most 8 tokens on an attention decode branch, o_proj, the +all-reduce and the update are one `comm/k3_sandwich_oproj` kernel; otherwise the attention's unreduced o_proj output +goes through `comm/mnnvl_allreduce_attn_res`. The state those kernels share (the KDA projection's Lamport buffers, the +MLA attention workspace, the decode GEMVs' state, the TP group's MNNVL and sandwich workspaces) lives in typed objects +this target creates in `post_load_weights`, before any graph capture. The MoE front and routed experts come with +their own entries; until then they run the generic path on every step. **What this target asserts rather than adapts**: SM 10.0; the topology above; the MXFP4 checkpoint's quantization (W4A16_MXFP4 with no per-layer declarations, so the routed experts run the W4A8_MXFP4_MXFP8 default and the excluded @@ -161,6 +162,7 @@ "allgather", # The decode path's collectives (decode_comm.py). "mnnvl_allreduce_attn_res", + "k3_sandwich_oproj", ) #: The engine surface the first forward checks before this target relies on it: per object, the attributes read. @@ -1422,8 +1424,10 @@ def __init__( self.mlp_res_norm = KimiK3RMSNorm(cfg.hidden_size, eps=cfg.rms_norm_eps, dtype=dtype) self.self_attention_res_proj = nn.Linear(cfg.hidden_size, 1, bias=False, dtype=dtype) self.mlp_res_proj = nn.Linear(cfg.hidden_size, 1, bias=False, dtype=dtype) - # The decode path's collectives (decode_comm.py), set by the target's post_load_weights. + # The decode path's collectives (decode_comm.py), and whether their sandwich takes this layer's o_proj; set + # by the target's post_load_weights. self.decode_comm: Optional[_decode_comm.K3DecodeComm] = None + self.sandwich_oproj = False def forward( self, @@ -1457,8 +1461,10 @@ def forward( The post-attention step of at most ``AR_ATTN_RES_MAX_TOKENS`` tokens runs the attention's all-reduce and the residual update in one - collective (``K3DecodeComm.allreduce_attn_res``) once the target has - built its ``decode_comm``. + collective once the target has built its ``decode_comm``: with o_proj + (``K3DecodeComm.sandwich_oproj``) where the sandwich takes the layer, + the step and the attention's decode branch, else on the attention's + unreduced o_proj output (``K3DecodeComm.allreduce_attn_res``). """ prefix_sum = hidden_states valid_block_residual = block_residual[:num_snapshots] @@ -1504,15 +1510,32 @@ def forward( attention = self.linear_attn if self.is_kda else self.self_attn comm = self.decode_comm if comm is not None and comm.takes_post_attention(hidden_states, step): - partial = attention(hidden_states, attn_metadata, step=step, reduce_output=False) - hidden_states, prefix_sum = comm.allreduce_attn_res( - partial, - prefix_sum, - valid_block_residual, - self.mlp_res_proj, - self.mlp_res_norm, - self.post_attention_layernorm, - ) + if ( + self.sandwich_oproj + and step is not None + and step.num_tokens <= _decode_comm.SANDWICH_MAX_TOKENS + and attention.will_run_decode_branch(attn_metadata, step) + ): + core = attention(hidden_states, attn_metadata, step=step, project_output=False) + hidden_states, prefix_sum = comm.sandwich_oproj( + core, + self._o_proj().weight, + prefix_sum, + valid_block_residual, + self.mlp_res_proj, + self.mlp_res_norm, + self.post_attention_layernorm, + ) + else: + partial = attention(hidden_states, attn_metadata, step=step, reduce_output=False) + hidden_states, prefix_sum = comm.allreduce_attn_res( + partial, + prefix_sum, + valid_block_residual, + self.mlp_res_proj, + self.mlp_res_norm, + self.post_attention_layernorm, + ) elif prefix_sum is None: prefix_sum = attention(hidden_states, attn_metadata, step=step) hidden_states = _apply_attn_res_and_rmsnorm( @@ -1563,6 +1586,10 @@ def _mnnvl_allreduce(self): attention = self.linear_attn if self.is_kda else self.self_attn return getattr(getattr(attention, "_o_allreduce", None), "mnnvl_allreduce", None) + def _o_proj(self) -> nn.Module: + """This layer's attention output projection (row parallel).""" + return self.linear_attn.o_proj if self.is_kda else self.self_attn.mixer.o_proj + def skip_forward( self, hidden_states: torch.Tensor, @@ -2371,7 +2398,7 @@ def post_load_weights(self) -> None: every layer of its kind: the KDA projection's Lamport buffers and the MLA decode attention's workspace; the decode GEMVs' state (built by ``cache_derived_state``) handed to every attention module; and, where every attention all-reduce runs over MNNVL, the TP group's collective state (``K3DecodeComm``, collective: every - rank builds it here) handed to every layer.""" + rank builds it here) handed to every layer, with whether its sandwich takes the layer's o_proj.""" super().post_load_weights() kda = [layer.linear_attn for layer in self.model.layers if layer.is_kda] mla = [layer.self_attn.mixer for layer in self.model.layers if not layer.is_kda] @@ -2386,10 +2413,16 @@ def post_load_weights(self) -> None: module.k3_workspace = workspace for module in kda + mla: module.decode_gemvs = self.model.decode_gemvs + layers = self.model.layers comm = None - if all(layer._mnnvl_allreduce() is not None for layer in self.model.layers): - comm = _decode_comm.K3DecodeComm.create(self.model_config.mapping) - for layer in self.model.layers: + if all(layer._mnnvl_allreduce() is not None for layer in layers): + for layer in layers: + layer.sandwich_oproj = _decode_comm.K3DecodeComm.takes_oproj( + layer._o_proj(), self.model.num_attn_res_snapshots + ) + oproj = next((layer._o_proj().weight for layer in layers if layer.sandwich_oproj), None) + comm = _decode_comm.K3DecodeComm.create(self.model_config.mapping, oproj) + for layer in layers: layer.decode_comm = comm logger.info( "Kimi K3 decode kernels: KDA on k3_kda_decode_attn, k3_kda_attn and k3_kda_verify " @@ -2397,7 +2430,10 @@ def post_load_weights(self) -> None: f"k3_mla_attn_vb_out ({len(mla)} layers); the post-attention all-reduce of at most " f"{_decode_comm.AR_ATTN_RES_MAX_TOKENS} tokens " + ( - "with the residual update (mnnvl_allreduce_attn_res)" + "with the residual update: k3_sandwich_oproj with o_proj at most " + f"{_decode_comm.SANDWICH_MAX_TOKENS} decode tokens " + f"({sum(layer.sandwich_oproj for layer in layers)} / {len(layers)} layers), else " + "mnnvl_allreduce_attn_res" if comm is not None else "unfused (an attention all-reduce does not run over MNNVL)" ) From 3ad8c5a9e70ac8eeafbe9f50a1f4f5ff06e3c44d Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 21:43:03 -0700 Subject: [PATCH 106/161] [None][feat] modeling_v2 Kimi K3 target: the decode path's one-shot ceilings for the stock MNNVL all-reduces The all-reduces the decode path still runs on the stock modules (layer 0's dense MLP, the MoE's, a wide step's attention output) choose one-shot or two-shot by message size; the stock ceiling is 1 MiB. At 16 ranks a decode step's [8, 7168] bf16 rows are 1.75 MiB, so they went two-shot, which sums in another order than the one-shot kernels of the fused steps. Where the target builds its collective state, every stock MNNVL all-reduce now sends one-shot up to 4 MiB (use_decode_one_shot), and a wide decode step's attention all-reduce keeps the stock 1 MiB per call (wide_all_reduce): at 16 ranks two-shot is faster for those rows. The attention hands the wide step's o_proj partial over unreduced for that. These are the ceilings the decode kernels were developed and checked with. Signed-off-by: Vasanth Sabavat --- .../decode_comm.py | 33 +++++++++++ .../modeling.py | 57 +++++++++++-------- 2 files changed, 66 insertions(+), 24 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py index e536fc78738e..36b8ba2a8cf4 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py @@ -17,6 +17,9 @@ over the group and eager, built by the target in `post_load_weights` before any CUDA-graph capture. Every rank must make the same calls on each in the same order. Which call a step takes is decided from its token count and kind and from the load-time layout alone, which every rank of the group shares. + +The plain MNNVL all-reduces the decode path keeps (the stock modules') send one-shot up to +`DECODE_AR_ONE_SHOT_MAX_BYTES` (`use_decode_one_shot`), except a wide decode step's (`wide_all_reduce`). """ from __future__ import annotations @@ -40,6 +43,7 @@ # The kernel's support predicate: metadata reads only. from tensorrt_llm._torch.cute_dsl_kernels.k3_sandwich import op as _sandwich_op +from tensorrt_llm._torch.distributed import AllReduceParams # The most tokens of a step whose post-attention all-reduce runs one-shot with the residual update in its epilogue. AR_ATTN_RES_MAX_TOKENS = 16 @@ -50,6 +54,35 @@ # The MNNVL workspace's buffer size: a one-shot call pushes T x 7168 bf16 from each of 16 ranks, 3.5 MiB at T = 16. MNNVL_BUFFER_BYTES = 4 << 20 +# The one-shot ceiling of the stock MNNVL all-reduces on the decode path: 8 tokens x 7168 x 16 ranks x 2 B is +# 1.75 MiB, which the stock 1 MiB ceiling would send two-shot. +DECODE_AR_ONE_SHOT_MAX_BYTES = 4 << 20 + +# A wide decode step's ceiling: the stock 1 MiB. At 16 ranks two-shot is faster for every wide step's rows. +WIDE_AR_ONE_SHOT_MAX_BYTES = 1 << 20 + + +def use_decode_one_shot(model: nn.Module) -> None: + """Every stock MNNVL all-reduce of ``model`` sends one-shot up to `DECODE_AR_ONE_SHOT_MAX_BYTES`. Each grows its + workspace on its first eager call of a larger size, before the capture of that size.""" + for module in model.modules(): + mnnvl = getattr(module, "mnnvl_allreduce", None) + if mnnvl is not None: + mnnvl.one_shot_max_bytes = DECODE_AR_ONE_SHOT_MAX_BYTES + + +def wide_all_reduce(all_reduce: nn.Module, x: torch.Tensor) -> torch.Tensor: + """``all_reduce(x)`` of a wide decode step (a stock ``AllReduce`` module, no fusion): its MNNVL all-reduce with + the `WIDE_AR_ONE_SHOT_MAX_BYTES` ceiling, else the module itself.""" + mnnvl = getattr(all_reduce, "mnnvl_allreduce", None) + if mnnvl is not None: + out = mnnvl( + x.contiguous(), AllReduceParams(), one_shot_max_bytes=WIDE_AR_ONE_SHOT_MAX_BYTES + ) + if out is not None: + return out + return all_reduce(x) + def _eps(norm: nn.Module) -> float: """The epsilon of a KimiK3RMSNorm (``eps``) or a stock RMSNorm (``variance_epsilon``).""" diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index 90c10c94f4be..7ad006ecb78c 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -45,10 +45,11 @@ one token tile, 32 tokens on a wide decode step. On any step of at most 16 tokens but a wide one, the post-attention all-reduce carries the residual update (`decode_comm.py`): at most 8 tokens on an attention decode branch, o_proj, the all-reduce and the update are one `comm/k3_sandwich_oproj` kernel; otherwise the attention's unreduced o_proj output -goes through `comm/mnnvl_allreduce_attn_res`. The state those kernels share (the KDA projection's Lamport buffers, the -MLA attention workspace, the decode GEMVs' state, the TP group's MNNVL and sandwich workspaces) lives in typed objects -this target creates in `post_load_weights`, before any graph capture. The MoE front and routed experts come with -their own entries; until then they run the generic path on every step. +goes through `comm/mnnvl_allreduce_attn_res`. The stock MNNVL all-reduces send one-shot up to 4 MiB, a wide decode +step's up to the stock 1 MiB. The state those kernels share (the KDA projection's Lamport buffers, the MLA attention +workspace, the decode GEMVs' state, the TP group's MNNVL and sandwich workspaces) lives in typed objects this target +creates in `post_load_weights`, before any graph capture. The MoE front and routed experts come with their own +entries; until then they run the generic path on every step. **What this target asserts rather than adapts**: SM 10.0; the topology above; the MXFP4 checkpoint's quantization (W4A16_MXFP4 with no per-layer declarations, so the routed experts run the W4A8_MXFP4_MXFP8 default and the excluded @@ -1536,26 +1537,32 @@ def forward( self.mlp_res_norm, self.post_attention_layernorm, ) - elif prefix_sum is None: - prefix_sum = attention(hidden_states, attn_metadata, step=step) - hidden_states = _apply_attn_res_and_rmsnorm( - prefix_sum, - valid_block_residual, - self.mlp_res_proj, - self.mlp_res_norm, - self.post_attention_layernorm, - attn_res_max_tokens, - ) else: - prefix_sum, hidden_states = _apply_attn_res_add_and_rmsnorm( - prefix_sum, - attention(hidden_states, attn_metadata, step=step), - valid_block_residual, - self.mlp_res_proj, - self.mlp_res_norm, - self.post_attention_layernorm, - attn_res_max_tokens, - ) + if comm is not None and step is not None and step.wide: + partial = attention(hidden_states, attn_metadata, step=step, reduce_output=False) + attention_output = _decode_comm.wide_all_reduce(attention._o_allreduce, partial) + else: + attention_output = attention(hidden_states, attn_metadata, step=step) + if prefix_sum is None: + prefix_sum = attention_output + hidden_states = _apply_attn_res_and_rmsnorm( + prefix_sum, + valid_block_residual, + self.mlp_res_proj, + self.mlp_res_norm, + self.post_attention_layernorm, + attn_res_max_tokens, + ) + else: + prefix_sum, hidden_states = _apply_attn_res_add_and_rmsnorm( + prefix_sum, + attention_output, + valid_block_residual, + self.mlp_res_proj, + self.mlp_res_norm, + self.post_attention_layernorm, + attn_res_max_tokens, + ) if self.is_moe: hidden_states = self.block_sparse_moe( hidden_states, getattr(attn_metadata, "all_rank_num_tokens", None) @@ -2398,7 +2405,8 @@ def post_load_weights(self) -> None: every layer of its kind: the KDA projection's Lamport buffers and the MLA decode attention's workspace; the decode GEMVs' state (built by ``cache_derived_state``) handed to every attention module; and, where every attention all-reduce runs over MNNVL, the TP group's collective state (``K3DecodeComm``, collective: every - rank builds it here) handed to every layer, with whether its sandwich takes the layer's o_proj.""" + rank builds it here) handed to every layer, with whether its sandwich takes the layer's o_proj, and the + decode path's one-shot ceiling on every stock MNNVL all-reduce (``use_decode_one_shot``).""" super().post_load_weights() kda = [layer.linear_attn for layer in self.model.layers if layer.is_kda] mla = [layer.self_attn.mixer for layer in self.model.layers if not layer.is_kda] @@ -2424,6 +2432,7 @@ def post_load_weights(self) -> None: comm = _decode_comm.K3DecodeComm.create(self.model_config.mapping, oproj) for layer in layers: layer.decode_comm = comm + _decode_comm.use_decode_one_shot(self) logger.info( "Kimi K3 decode kernels: KDA on k3_kda_decode_attn, k3_kda_attn and k3_kda_verify " f"({sum(m.takes_k3_kernels for m in kda)} / {len(kda)} layers take them), MLA on k3_mla_qkv and " From d99cbd60d7bcf6f6f7f439f9c2d9fb3cac2b9957 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 22:06:06 -0700 Subject: [PATCH 107/161] [None][feat] modeling_v2 Kimi K3 target: the MoE layers of a decode step on the K3 MoE kernels decode_moe.py runs a MoE layer of a step of at most 8 tokens on the catalog's Kimi K3 MoE entries: - moe/k3_moe_front, one kernel: this rank's slice of the MoE head (its latent-down and router rows) as one GEMV, the slices' all-gather over the TP group's K3MoeHeadWorkspace, the top-16 routing, the MXFP8 latent, and the shared experts' gate_up + SiTU; - moe/k3_moe on a K3MoeState: this rank's routed partial; - the latent all-reduce (the routed experts' all-reduce, one-shot); - the tail. The latent norm's weight is folded into the latent up projection at load, so [RMSNorm(latent) slice | shared activation] @ [latent up columns | shared down] is this rank's share of the output. A layer whose output the next layer's pre-attention step (or the final norm) can reduce hands it on unreduced, as a loop-local PendingTail. The consumer runs the tail, its all-reduce and the residual update as one comm/k3_sandwich_tail kernel (K3DecodeComm.sandwich_tail); a snapshot layer has the kernel store the prefix sum into the bank row it takes. A layer DSpark taps defers too, and its consumer taps the reduced value's pre-norm mixture. An unknown capture set and a tapped last layer keep the replicated tail: the latent RMS on the fp32 output of one GEMV with the folded latent up weight, plus the shared experts' down projection. A wide decode step's MoE (9 to 64 tokens) keeps the sharded head and the row-parallel tail on M-general ops: - the head GEMV and comm/mnnvl_allgather_split on the TP group's MnnvlWorkspace; - moe/k3_route_quant and moe/k3_moe on a K3MoeWideState, beside the shared gate_up + SiTU; - the latent all-reduce, then one GEMV with the tail weight. Its unreduced output is reduced by the consumer's all-reduce, followed by the fused add + attn_res + RMSNorm. The GEMVs run on four new decode GEMV sites: the head slice and the replicated tail's latent up projection with fp32 outputs, the shared gate_up, and the tail. The wide kernel takes the head slice and the tail up to 32 rows. Site gains out_fp32 and wide_rows. The target builds it in post_load_weights on every MoE layer it takes (layout_gaps: the TP16 shapes the front and the sandwich tail take, the MXFP4 buffers k3_moe reads), with the state shared by those layers. The head workspace is collective. It then runs every kernel once: the front and the sandwich tail on every rank, both k3_moe builds and k3_route_quant. So no capture compiles one. Signed-off-by: Vasanth Sabavat --- .../decode_comm.py | 81 +++- .../decode_gemv.py | 45 ++- .../decode_moe.py | 357 ++++++++++++++++++ .../modeling.py | 245 +++++++++++- .../test_modeling_v2_kimi_k3_decode_gemv.py | 18 +- 5 files changed, 705 insertions(+), 41 deletions(-) create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_moe.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py index 36b8ba2a8cf4..2ef7dbe27458 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py @@ -13,6 +13,10 @@ * `K3DecodeComm.allreduce_attn_res`, `comm/mnnvl_allreduce_attn_res`, everywhere else: the one-shot MNNVL all-reduce of the attention's unreduced o_proj partial, with the residual update as its epilogue. +The pre-attention step after a MoE layer whose row-parallel tail was handed on (a `PendingTail`, `decode_moe.py`) is +the same kind of collective: `K3DecodeComm.sandwich_tail`, `comm/k3_sandwich_tail`, runs the tail GEMV, its +all-reduce and the next layer's residual update (the final norm's, after the last layer) in one kernel. + The state is one `MnnvlWorkspace` and one `K3SandwichWorkspace` of the TP group (`K3DecodeComm.create`): collective over the group and eager, built by the target in `post_load_weights` before any CUDA-graph capture. Every rank must make the same calls on each in the same order. Which call a step takes is decided from its token count and kind and @@ -25,7 +29,7 @@ from __future__ import annotations from dataclasses import dataclass -from typing import Optional, Tuple +from typing import NamedTuple, Optional, Tuple import torch from torch import nn @@ -34,6 +38,9 @@ K3SandwichWorkspace, k3_sandwich_oproj, ) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.k3_sandwich_tail import ( + k3_sandwich_tail, +) from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.mnnvl_allreduce_attn_res import ( mnnvl_allreduce_attn_res, ) @@ -84,6 +91,18 @@ def wide_all_reduce(all_reduce: nn.Module, x: torch.Tensor) -> torch.Tensor: return all_reduce(x) +class PendingTail(NamedTuple): + """A MoE layer's row-parallel tail left to its consumer's fused pre-attention step (``K3DecodeComm.sandwich_tail``): + the reduced latent ``[T, 3584]``, the shared experts' activation, the tail weight ``[latent up columns | padding | + shared down]``, this rank's first latent column and the latent norm's epsilon.""" + + latent: torch.Tensor + act: torch.Tensor + weight: torch.Tensor + lo: int + lat_eps: float + + def _eps(norm: nn.Module) -> float: """The epsilon of a KimiK3RMSNorm (``eps``) or a stock RMSNorm (``variance_epsilon``).""" return float(norm.eps if hasattr(norm, "eps") else norm.variance_epsilon) @@ -139,6 +158,30 @@ def create(cls, mapping, oproj_weight: Optional[torch.Tensor] = None) -> "K3Deco torch.cuda.synchronize(weight.device) return state + def compile_tail(self, latent_size: int, act_size: int, tail_weight: torch.Tensor) -> None: + """Compile the sandwich tail kernel for a MoE tail of ``latent_size`` latent and ``act_size`` shared columns + and ``tail_weight``'s shape, with one call on a zero row of a zero weight, before any capture. Collective: every + rank of the group makes the call; it advances the sandwich workspace on every rank alike.""" + weight = torch.zeros_like(tail_weight) + hidden = weight.shape[0] + ones = weight.new_ones(hidden) + k3_sandwich_tail( + weight.new_zeros(1, latent_size), + weight.new_zeros(1, act_size), + weight, + 0, + 1e-6, + None, + weight.new_zeros(0, 1, hidden), + weight.new_zeros(hidden), + ones, + ones, + 1e-6, + 1e-6, + self.sandwich, + ) + torch.cuda.synchronize(weight.device) + def takes_post_attention(self, hidden_states: torch.Tensor, step) -> bool: """Whether one of this state's collectives runs the post-attention step of a layer whose attention input is ``hidden_states``: at most `AR_ATTN_RES_MAX_TOKENS` bf16 rows of a hidden size the MNNVL entry takes, on any @@ -170,6 +213,15 @@ def takes_oproj(o_proj: nn.Module, max_snapshots: int) -> bool: and _sandwich_op.supports(weight.new_empty((1, weight.shape[1])), weight, max_snapshots) ) + @staticmethod + def takes_tail( + latent: torch.Tensor, act: torch.Tensor, tail_weight: torch.Tensor, max_snapshots: int + ) -> bool: + """Whether ``sandwich_tail`` takes a MoE tail of these tensors' shapes (rows of ``latent`` and ``act``, at most + `SANDWICH_MAX_TOKENS`; ``tail_weight``) with a snapshot bank of at most ``max_snapshots`` rows: the TP16 + per-rank shapes.""" + return _sandwich_op.supports_tail(latent, act, tail_weight, max_snapshots) + def allreduce_attn_res( self, partial: torch.Tensor, @@ -191,6 +243,33 @@ def allreduce_attn_res( self.mnnvl, ) + def sandwich_tail( + self, + pending: PendingTail, + prefix_sum: Optional[torch.Tensor], + block_residual: torch.Tensor, + res_proj: nn.Module, + res_norm: nn.Module, + out_norm: nn.Module, + updated_out: Optional[torch.Tensor] = None, + ) -> Tuple[torch.Tensor, torch.Tensor]: + """``(normed, updated)`` as ``allreduce_attn_res`` of a MoE layer's row-parallel tail ``pending`` + (``[RMSNorm(latent)[:, lo:lo + 224] | act] @ weight.T``), in one ``comm/k3_sandwich_tail`` call. + ``updated_out``: a bf16 ``[T, H]`` tensor the call stores ``updated`` into (the consumer's snapshot bank row), + returned as ``updated``.""" + return k3_sandwich_tail( + pending.latent, + pending.act, + pending.weight, + pending.lo, + pending.lat_eps, + prefix_sum, + block_residual, + *_res_args(res_proj, res_norm, out_norm), + self.sandwich, + updated_out=updated_out, + ) + def sandwich_oproj( self, core: torch.Tensor, diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py index b384010b65dd..0e26e3971d0e 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py @@ -4,10 +4,11 @@ * **Per-site GEMVs** (`K3DecodeGemvs.project`): a projection of a decode step runs on the kernel measured fastest at its call site's weight shape (`SITES`): at most `MAX_ROWS` rows on `gemm/k3_decode_gemv`, - `gemm/k3_ctm_gemv_wide` or `gemm/k3_ctm_gemv_long`, and, where the site lists it, up to `WIDE_ROWS` rows on - `gemm/k3_ctm_gemv_wide`. The sites are the decode kernels' fused projections (MLA's [W_a; W_g] with the gate rows - through a sigmoid, KDA's [q | k | v | g | f_a | b]), the attention output projection, and the built-in MLA path's - q_a / kv_a, q_b and gate projections. + `gemm/k3_ctm_gemv_wide` or `gemm/k3_ctm_gemv_long`, and, where the site lists it, more rows (up to `WIDE_ROWS`) + on `gemm/k3_ctm_gemv_wide`. The sites are the decode kernels' fused projections (MLA's [W_a; W_g] with the gate + rows through a sigmoid, KDA's [q | k | v | g | f_a | b]), the attention output projection, the built-in MLA path's + q_a / kv_a, q_b and gate projections, and the MoE decode path's projections (`decode_moe.py`), two of them with + fp32 outputs. * **LM head** (`K3LogitsProcessor`): at most `MAX_ROWS` rows of this rank's vocabulary shard on `gemm/k3_head_gemv` over the target's `K3HeadGemvWorkspace`, then the shards gathered (`comm/allgather`) as the stock head gathers them. It is the shell's logits processor, so the speculative worker's target logits and the @@ -72,9 +73,10 @@ @dataclass(frozen=True) class Site: """A call site's weight shape (this target's per-rank shapes) and its kernels: ``small`` at 1..MAX_ROWS rows - ("decode", "wide" or "long"), and k3_ctm_gemv_wide at MAX_ROWS+1..WIDE_ROWS rows where ``wide``. Output columns - from ``sig_col0`` on are stored through a sigmoid. ``split`` / ``ring`` / ``push``: k3_ctm_gemv_long's CTAs per - 128-row weight tile, weight-ring stages, and whether the partial sums are pushed to each row's owner.""" + ("decode", "wide" or "long"), and k3_ctm_gemv_wide at MAX_ROWS+1..``wide_rows`` rows where ``wide``. Output columns + from ``sig_col0`` on are stored through a sigmoid; ``out_fp32``: an fp32 output (k3_ctm_gemv_wide only). + ``split`` / ``ring`` / ``push``: k3_ctm_gemv_long's CTAs per 128-row weight tile, weight-ring stages, and whether + the partial sums are pushed to each row's owner.""" n: int k: int @@ -84,6 +86,8 @@ class Site: split: int = 0 ring: int = 0 push: bool = False + out_fp32: bool = False + wide_rows: int = WIDE_ROWS SITES: Dict[str, Site] = { @@ -101,6 +105,14 @@ class Site: # Layer 0's dense MLP split over the 16-way TP group: gate_up [gate 2112 | up 2112] and down. "dense_gate_up": Site(4224, 7168, "long", split=4, ring=5), "dense_down": Site(7168, 2112, "long", split=2, ring=6), + # The MoE decode path (decode_moe.py): this rank's head slice [latent down 224 | router 56] with an fp32 output, + # the shared experts' gate_up, the row-parallel tail [latent up 224 | padding 32 | shared down 384] and the + # replicated tail's latent up projection with an fp32 output. The wide kernel takes the head slice and the tail up + # to 32 rows, where it is faster than the stock GEMM. + "moe_head": Site(280, 7168, "wide", wide=True, out_fp32=True, wide_rows=32), + "moe_shared_gate_up": Site(768, 7168, "wide", wide=True), + "moe_tail": Site(7168, 640, "wide", wide=True, wide_rows=32), + "moe_up": Site(7168, 3584, "wide", out_fp32=True), } @@ -115,13 +127,15 @@ def _run( ) -> Optional[torch.Tensor]: """``spec``'s ``kernel`` on dense rows ``x2d``, or None where it does not take them.""" if kernel == "decode": - if spec.sig_col0 >= 0 or not _decode_op.supports(x2d, weight): + if spec.sig_col0 >= 0 or spec.out_fp32 or not _decode_op.supports(x2d, weight): return None return k3_decode_gemv(x2d, weight) if kernel == "wide": - if not _ctm_op.supports_wide(x2d, weight, spec.sig_col0, False): + if not _ctm_op.supports_wide(x2d, weight, spec.sig_col0, spec.out_fp32): return None - return k3_ctm_gemv_wide(x2d, weight, sig_col0=spec.sig_col0) + return k3_ctm_gemv_wide(x2d, weight, sig_col0=spec.sig_col0, out_fp32=spec.out_fp32) + if spec.out_fp32: + return None # One wave of the GPU's SMs: beyond it the long GEMV loses to the others. sms = torch.cuda.get_device_properties(x2d.device).multi_processor_count if math.ceil(spec.n / 128) * spec.split > sms or not _ctm_op.supports_long( @@ -224,6 +238,8 @@ def create( spec = SITES[site] weight = torch.zeros(spec.n, spec.k, dtype=torch.bfloat16, device=device) for rows in (1, 16, 32, 64) if spec.wide else (1,): + if rows > spec.wide_rows: + continue state._project(site, weight.new_zeros(rows, spec.k), weight, warm=True) del weight if "dense_gate_up" in sites: @@ -245,9 +261,10 @@ def create( return state def project(self, site: str, x: torch.Tensor, weight: torch.Tensor) -> Optional[torch.Tensor]: - """``x @ weight.T`` (bf16 ``[..., N]``, the site's sigmoid columns through the sigmoid) for ``site``'s weight - on its decode kernel, or None where none takes the call: more rows than the site's kernels take, another - shape or dtype, or, under capture, a kernel that has not run eagerly. The caller then runs its GEMM.""" + """``x @ weight.T`` (``[..., N]``: bf16, the site's sigmoid columns through the sigmoid; fp32 at an + ``out_fp32`` site) for ``site``'s weight on its decode kernel, or None where none takes the call: more rows + than the site's kernels take, another shape or dtype, or, under capture, a kernel that has not run eagerly. + The caller then runs its GEMM.""" return self._project(site, x, weight, warm=False) def _project( @@ -266,7 +283,7 @@ def _project( rows = x.numel() // spec.k if 0 < rows <= MAX_ROWS: kernel = spec.small - elif spec.wide and MAX_ROWS < rows <= WIDE_ROWS: + elif spec.wide and MAX_ROWS < rows <= spec.wide_rows: kernel = "wide" else: return None diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_moe.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_moe.py new file mode 100644 index 000000000000..75f5b06ed66e --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_moe.py @@ -0,0 +1,357 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""The decode path's latent MoE on the catalog's Kimi K3 MoE entries. + +**At most 8 tokens** (`MAX_TOKENS`), a MoE layer runs: + +* `moe/k3_moe_front`, one kernel: this rank's slice of the MoE head (its latent-down rows and its router rows) as one + GEMV, the slices' all-gather over the TP group's `K3MoeHeadWorkspace`, the top-16 routing, the MXFP8 latent, and + the shared experts' gate_up + SiTU; +* `moe/k3_moe`: this rank's routed partial over its experts (on a `K3MoeState`); +* the latent all-reduce: the routed experts' all-reduce (one-shot, as `decode_comm.use_decode_one_shot` sets); +* the tail. The latent norm's weight is folded into the latent up projection at load, so + `[RMSNorm(latent) slice | shared activation] @ [latent up columns | shared down]` is this rank's share of the MoE + output. Where the next layer's pre-attention step (or the final norm) reduces it, the layer hands it on unreduced + as a `PendingTail`, which the consumer runs with its all-reduce and residual update as one `comm/k3_sandwich_tail` + kernel (`decode_comm.py`). Elsewhere the replicated tail runs: the latent RMS applied to the fp32 output of one + GEMV with the folded latent up weight, plus the shared experts' down projection and its all-reduce. + +**A wide decode step** (9 to 64 tokens, `WIDE_MAX_TOKENS`) keeps the sharded head and the row-parallel tail on +M-general ops: the head GEMV, `comm/mnnvl_allgather_split`, then `moe/k3_route_quant` and `moe/k3_moe` (on a +`K3MoeWideState`) beside the shared gate_up + SiTU, the latent all-reduce, and one GEMV of +`[RMSNorm(latent) slice | padding | shared activation]` with the tail weight. That is this rank's unreduced share, +which the consumer reduces with a plain all-reduce. + +The GEMVs run on the decode GEMV sites of `decode_gemv.py` where they take the call, else on the stock GEMM ops. + +`K3DecodeMoe` holds what every MoE layer shares: the head workspace (collective over the TP group), the two +`k3_moe` builds' scratch, and the TP group's MNNVL workspace (`decode_comm.K3DecodeComm`'s). `K3DecodeMoeLayer` +holds one layer's decode weights and its `k3_moe` counters. The target builds both in `post_load_weights`, before +any CUDA-graph capture, and runs every kernel once there so none compiles under a capture. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Optional, Union + +import torch +from torch import nn + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.mnnvl_allgather_split import ( + mnnvl_allgather_split, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.mnnvl_workspace import ( + MnnvlWorkspace, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.k3_moe import ( + K3MoeHeadWorkspace, + K3MoeLayer, + K3MoeState, + K3MoeWideState, + is_supported, + k3_moe, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.k3_moe_front import ( + front_weight, + k3_moe_front, + weight_supported, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.k3_route_quant import k3_route_quant +from tensorrt_llm._torch.modules.multi_stream_utils import maybe_execute_in_parallel + +from .decode_comm import K3DecodeComm, PendingTail, wide_all_reduce + +# The most tokens of the front and the small k3_moe build (one token tile), and of the wide build. +MAX_TOKENS = 8 +WIDE_MAX_TOKENS = 64 + +# The tail weight's latent columns are zero-padded to whole 128-column k-tiles of the tail kernels +# (comm/k3_sandwich_tail takes the TP16 tail weight [7168, 256 + 384]). +TAIL_K_TILE = 256 + + +@dataclass(eq=False) +class K3DecodeMoe: + """What every MoE layer's decode path shares on one device: the TP group's ``K3MoeHeadWorkspace`` (the front's + all-gather), the ``k3_moe`` builds for up to 8 and up to 64 tokens (their scratch; the layers run one at a time on + one stream), and the TP group's ``MnnvlWorkspace`` for a wide step's head all-gather. Built by `create`.""" + + head: K3MoeHeadWorkspace + small: K3MoeState + wide: K3MoeWideState + mnnvl: MnnvlWorkspace + + @classmethod + def create( + cls, mapping, device, i_tp: int, num_local: int, mnnvl: MnnvlWorkspace + ) -> "K3DecodeMoe": + """The state for ``mapping``'s TP group on ``device``, for experts of ``i_tp`` intermediate columns per rank, + ``num_local`` of them on this rank. Collective (the head workspace): every rank of the group calls it at the + same point, eagerly, before any CUDA-graph capture.""" + return cls( + K3MoeHeadWorkspace.create(mapping), + K3MoeState(device, i_tp, num_local), + K3MoeWideState(device, i_tp, num_local), + mnnvl, + ) + + +def _experts(moe: nn.Module) -> tuple: + """The routed experts' TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers ``k3_moe`` reads in place.""" + backend = moe.routed_experts.backend + return ( + backend.w3_w1_weight, + backend.w3_w1_weight_scale, + backend.w2_weight, + backend.w2_weight_scale, + ) + + +def layout_gaps(moe: nn.Module, tp_size: int, max_snapshots: int) -> list: + """Why the decode path does not take MoE layer ``moe`` (empty when it takes it), read once the routed experts' + buffers exist: the checks in order, up to the first that fails. ``max_snapshots``: the model's snapshot bank + rows.""" + gate, backend = moe.gate, moe.routed_experts.backend + shared = moe.shared_experts + gate_up, down = shared.gate_up_proj.weight, shared.down_proj.weight + if not moe._reduce_routed_output: + return ["a routed output the model does not reduce"] + if (moe.num_experts, moe.top_k, moe.moe_hidden_size, moe.hidden_size) != (896, 16, 3584, 7168): + return [ + f"experts / top-k / latent / hidden {(moe.num_experts, moe.top_k, moe.moe_hidden_size)}" + ] + if (gate.num_expert_group, gate.topk_group) != (1, 1): + return ["grouped routing"] + if moe.moe_hidden_size % (8 * tp_size) or moe.num_experts % (4 * tp_size): + return [f"a head slice of TP {tp_size}"] + latent = (gate.weight, moe.routed_expert_down_proj.weight, moe.routed_expert_up_proj.weight) + if not isinstance(moe.routed_expert_up_proj, nn.Linear) or any( + w.dtype != torch.bfloat16 for w in latent + ): + return ["latent projections or a router other than bf16"] + if ( + gate_up.dtype != torch.bfloat16 + or down.dtype != torch.bfloat16 + or shared.gate_up_proj.bias is not None + or shared.down_proj.bias is not None + or down.shape[0] != moe.hidden_size + ): + return ["a shared expert other than bf16 and unbiased"] + names = ( + "w3_w1_weight", + "w3_w1_weight_scale", + "w2_weight", + "w2_weight_scale", + "expert_size_per_partition", + ) + if ( + not all(hasattr(backend, name) for name in names) + or not is_supported(*_experts(moe), backend.expert_size_per_partition)[0] + ): + return ["routed-expert buffers k3_moe does not read"] + shared_cols, width = gate_up.shape[0] // 2, moe.moe_hidden_size // tp_size + if not weight_supported(tp_size, shared_cols, gate_up.shape[1], gate_up.device): + return ["a MoE front of this TP size and shared width"] + tail_cols = width + (-width % TAIL_K_TILE) + shared_cols + probe = gate_up.new_empty + if not K3DecodeComm.takes_tail( + probe(1, moe.moe_hidden_size), + probe(1, shared_cols), + probe(moe.hidden_size, tail_cols), + max_snapshots, + ): + return ["a row-parallel tail the sandwich tail kernel does not take"] + return [] + + +def fold_latent_norm(moe: nn.Module) -> None: + """Fold the latent RMSNorm's weight into the latent up projection's columns; the norm keeps a weight of ones, so + every step computes the same function and the tails may normalize before slicing.""" + up, norm = moe.routed_expert_up_proj, moe.routed_expert_norm + with torch.no_grad(): + up.weight.mul_(norm.weight.to(up.weight.dtype)[None, :]) + norm.weight.fill_(1) + + +@dataclass(eq=False) +class K3DecodeMoeLayer: + """One MoE layer's decode path: the front weight (this rank's head slice padded to whole tiles, then the shared + gate_up re-ordered), the head slice (a view of it), the row-parallel tail weight + ``[latent up columns lo:lo+width | padding | shared down]``, and its ``k3_moe`` handles on the shared builds. + Built by `create` once the weights are final (and the latent norm folded).""" + + state: K3DecodeMoe + front_weight: torch.Tensor + head_weight: torch.Tensor + tail_weight: torch.Tensor + tail_pad: Optional[torch.Tensor] + lo: int + width: int + shared_cols: int + small: K3MoeLayer + wide: K3MoeLayer + + @classmethod + def create( + cls, moe: nn.Module, state: K3DecodeMoe, tp_rank: int, tp_size: int + ) -> "K3DecodeMoeLayer": + """Build MoE layer ``moe``'s decode weights and its handles on ``state``'s builds (``layout_gaps`` empty).""" + width = moe.moe_hidden_size // tp_size + experts = moe.num_experts // tp_size + gate_up = moe.shared_experts.gate_up_proj.weight + shared_down = moe.shared_experts.down_proj.weight + up = moe.routed_expert_up_proj.weight + lo = tp_rank * width + pad = -width % TAIL_K_TILE + with torch.no_grad(): + head = torch.cat( + [ + moe.routed_expert_down_proj.weight[lo : lo + width], + moe.gate.weight[tp_rank * experts : (tp_rank + 1) * experts], + ] + ) + front = front_weight(head, gate_up) + parts = [up[:, lo : lo + width]] + if pad: + parts.append(up.new_zeros(up.shape[0], pad)) + parts.append(shared_down) + tail = torch.cat(parts, dim=1).contiguous() + weights = _experts(moe) + return cls( + state=state, + front_weight=front, + head_weight=front[: head.shape[0]], + tail_weight=tail, + tail_pad=up.new_zeros(WIDE_MAX_TOKENS, pad) if pad else None, + lo=lo, + width=width, + shared_cols=gate_up.shape[0] // 2, + small=state.small.layer(*weights), + wide=state.wide.layer(*weights), + ) + + def warm_up(self, moe: nn.Module) -> None: + """One call of every kernel of the decode path on zero inputs (M = 1), so none compiles under a capture: the + front (collective: every rank makes the same call), both ``k3_moe`` builds and ``k3_route_quant``.""" + device = self.front_weight.device + x = torch.zeros(1, moe.hidden_size, dtype=torch.bfloat16, device=device) + ids, weights, x_fp8, x_sf, _ = self._front(moe, x) + offset = moe.routed_experts.backend.slot_start + k3_moe(x_fp8, x_sf, ids, weights, offset, self.small) + logits = torch.zeros(1, moe.num_experts, dtype=torch.float32, device=device) + latent = torch.zeros(1, moe.moe_hidden_size, dtype=torch.bfloat16, device=device) + ids, weights, x_fp8, x_sf = k3_route_quant( + logits, moe.gate.e_score_correction_bias, latent, float(moe.gate.routed_scaling_factor), + early_trigger=True, + ) # fmt: skip + k3_moe(x_fp8, x_sf, ids, weights, offset, self.wide) + torch.cuda.synchronize(device) + + def _front(self, moe: nn.Module, x: torch.Tensor): + return k3_moe_front( + x.contiguous(), + self.front_weight, + moe.gate.e_score_correction_bias, + float(moe.gate.routed_scaling_factor), + self.shared_cols, + *moe._situ_betas, + self.state.head, + ) + + def takes(self, hidden_states: torch.Tensor, step, partial_tail: bool) -> bool: + """Whether this decode path runs the layer on ``step`` (bf16 rows ``hidden_states``): any step of at most + `MAX_TOKENS` tokens; with ``partial_tail``, also a wide decode step.""" + rows = hidden_states.shape[0] + if hidden_states.dtype != torch.bfloat16 or step is None: + return False + return 0 < rows <= MAX_TOKENS or (partial_tail and step.wide and rows <= WIDE_MAX_TOKENS) + + def forward( + self, moe: nn.Module, hidden_states: torch.Tensor, gemvs, partial_tail: bool + ) -> Union[torch.Tensor, PendingTail]: + """The MoE output of ``hidden_states`` (``takes`` holds): with ``partial_tail``, this rank's unreduced share + (a ``PendingTail`` at most `MAX_TOKENS` tokens); else the reduced output. ``gemvs``: the decode GEMVs' state, + or None.""" + if hidden_states.shape[0] > MAX_TOKENS: + return self._wide(moe, hidden_states, gemvs) + ids, weights, x_fp8, x_sf, shared_act = self._front(moe, hidden_states) + routed = k3_moe( + x_fp8, x_sf, ids, weights, moe.routed_experts.backend.slot_start, self.small + ) + latent = moe.routed_experts.all_reduce(routed) + if partial_tail: + return PendingTail( + latent.contiguous(), + shared_act.contiguous(), + self.tail_weight, + self.lo, + float(moe.routed_expert_norm.variance_epsilon), + ) + # The replicated tail: the latent RMS on the fp32 accumulator of the folded latent up projection, plus the + # shared experts' reduced output, rounded to bf16 once. + shared = moe.shared_experts + shared_out = shared.down_proj(shared_act, layer_idx=shared.layer_idx) + up = _gemv(gemvs, "moe_up", latent, moe.routed_expert_up_proj.weight, out_fp32=True) + scale = torch.rsqrt( + latent.float().pow(2).mean(-1, keepdim=True) + moe.routed_expert_norm.variance_epsilon + ) + return (up * scale + shared_out.float()).bfloat16() + + def _wide(self, moe: nn.Module, hidden_states: torch.Tensor, gemvs) -> torch.Tensor: + """A wide decode step's MoE: this rank's unreduced share of the output, ``[M, hidden]`` bf16.""" + x = hidden_states.contiguous() + head = _gemv(gemvs, "moe_head", x, self.head_weight, out_fp32=True) + routed_in, router_logits = mnnvl_allgather_split(head, self.width, self.state.mnnvl) + shared = moe.shared_experts + + def _routed_partial(): + ids, weights, x_fp8, x_sf = k3_route_quant( + router_logits, moe.gate.e_score_correction_bias, routed_in.contiguous(), + float(moe.gate.routed_scaling_factor), early_trigger=True, + ) # fmt: skip + return k3_moe( + x_fp8, x_sf, ids, weights, moe.routed_experts.backend.slot_start, self.wide + ) + + def _shared_activation(): + return shared._apply_activation( + _gemv(gemvs, "moe_shared_gate_up", x, shared.gate_up_proj.weight) + ) + + routed, shared_act = maybe_execute_in_parallel( + _routed_partial, + _shared_activation, + moe.moe_main_event, + moe.moe_shared_event, + moe.shared_expert_stream, + disable_on_compile=True, + ) + # The latent norm's weight is folded into the tail weight: normalize the whole latent row, keep this rank's + # columns. + normed = moe.routed_expert_norm(wide_all_reduce(moe.routed_experts.all_reduce, routed)) + parts = [normed[:, self.lo : self.lo + self.width]] + if self.tail_pad is not None: + parts.append(self.tail_pad[: x.shape[0]]) + parts.append(shared_act) + return _gemv(gemvs, "moe_tail", torch.cat(parts, dim=1), self.tail_weight) + + +def _gemv( + gemvs, site: str, x: torch.Tensor, weight: torch.Tensor, out_fp32: bool = False +) -> torch.Tensor: + """``x @ weight.T`` on ``site``'s decode GEMV where it takes the call (``decode_gemv.K3DecodeGemvs.project``), + else on the stock GEMM: cuBLAS through ``trtllm::dsv3_router_gemm_op`` for an fp32 output, ``F.linear`` else.""" + y = None if gemvs is None else gemvs.project(site, x, weight) + if y is not None: + return y + if out_fp32: + # The op reads the weight with a leading dimension of K: a strided weight would give wrong values. + if not weight.is_contiguous(): + raise ValueError( + f"the fp32 GEMM needs a contiguous weight, got strides {weight.stride()}" + ) + return torch.ops.trtllm.dsv3_router_gemm_op( + x.contiguous(), weight.t(), bias=None, out_dtype=torch.float32 + ) + return torch.nn.functional.linear(x, weight) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index 7ad006ecb78c..c259ab07f728 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -25,9 +25,10 @@ layout's MoE head and tail, on M-general ops. Every other step (prefill, mixed steps, decode steps above those bounds) runs the **generic path**: this target's -text model (`KimiLinearModel` below: decoder layers, attention residuals, the MLA / KDA / MoE runtimes), computed -exactly as the built-in Kimi K3 text model computes it, on stock modules and ops that have no catalog entries yet. -`UNCERTIFIED_GENERIC_CALLS` names them. +text model (`KimiLinearModel` below: decoder layers, attention residuals, the MLA / KDA / MoE runtimes), computed as +the built-in Kimi K3 text model computes it, on stock modules and ops that have no catalog entries yet. One weight +differs: where the MoE decode path takes a layer, its latent norm's weight is folded into the latent up projection +(the same function, rounded differently). `UNCERTIFIED_GENERIC_CALLS` names the stock code. The text model hands each step's classification to its attention modules, which run a **decode step** on the K3 decode kernels' catalog entries: @@ -47,9 +48,14 @@ all-reduce and the update are one `comm/k3_sandwich_oproj` kernel; otherwise the attention's unreduced o_proj output goes through `comm/mnnvl_allreduce_attn_res`. The stock MNNVL all-reduces send one-shot up to 4 MiB, a wide decode step's up to the stock 1 MiB. The state those kernels share (the KDA projection's Lamport buffers, the MLA attention -workspace, the decode GEMVs' state, the TP group's MNNVL and sandwich workspaces) lives in typed objects this target -creates in `post_load_weights`, before any graph capture. The MoE front and routed experts come with their own -entries; until then they run the generic path on every step. +workspace, the decode GEMVs' state, the TP group's MNNVL and sandwich workspaces, the MoE path's) lives in typed +objects this target creates in `post_load_weights`, before any graph capture. + +A MoE layer on a step of at most 8 tokens runs `decode_moe.py`: `moe/k3_moe_front` (this rank's head slice, its +all-gather, the routing, the MXFP8 latent and the shared experts' gate_up + SiTU in one kernel), `moe/k3_moe`, the +latent all-reduce, then the row-parallel tail, which the next layer's pre-attention step (the final norm's, after the +last layer) runs with its all-reduce and residual update as one `comm/k3_sandwich_tail` kernel. A wide decode step's +MoE keeps the sharded head and the row-parallel tail, on `moe/k3_route_quant`, `moe/k3_moe` and M-general ops. **What this target asserts rather than adapts**: SM 10.0; the topology above; the MXFP4 checkpoint's quantization (W4A16_MXFP4 with no per-layer declarations, so the routed experts run the W4A8_MXFP4_MXFP8 default and the excluded @@ -72,7 +78,7 @@ import math import os from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, Callable, Dict, Literal, NamedTuple, Optional, Tuple +from typing import TYPE_CHECKING, Any, Callable, Dict, Literal, NamedTuple, Optional, Tuple, Union import torch from torch import nn @@ -125,6 +131,7 @@ from . import decode_comm as _decode_comm from . import decode_gemv as _decode_gemv +from . import decode_moe as _decode_moe from . import weights as _weights if TYPE_CHECKING: @@ -161,9 +168,14 @@ "k3_head_gemv", "k3_embed_norm", "allgather", - # The decode path's collectives (decode_comm.py). + # The decode path's collectives (decode_comm.py) and MoE (decode_moe.py). "mnnvl_allreduce_attn_res", "k3_sandwich_oproj", + "k3_sandwich_tail", + "k3_moe_front", + "k3_moe", + "k3_route_quant", + "mnnvl_allgather_split", ) #: The engine surface the first forward checks before this target relies on it: per object, the attributes read. @@ -851,6 +863,7 @@ def __init__( raise ValueError("Kimi K3 runtime expects latent_moe_use_norm=True") situ_beta, situ_linear_beta = _resolve_kimi_situ_betas(cfg) + self._situ_betas = (situ_beta, situ_linear_beta) dtype = torch.bfloat16 # Routing scores stay fp32; the gate GEMM runs bf16xbf16 with fp32 @@ -975,6 +988,9 @@ def __init__( self.routed_expert_norm = RMSNorm( hidden_size=self.moe_hidden_size, eps=cfg.rms_norm_eps, dtype=dtype ) + # The decode path (decode_moe.py) and the decode GEMVs' state, set by the target's post_load_weights. + self.decode_moe: Optional[_decode_moe.K3DecodeMoeLayer] = None + self.decode_gemvs: Optional[_decode_gemv.K3DecodeGemvs] = None @staticmethod def _routed_projection(hidden_states: torch.Tensor, projection: nn.Module) -> torch.Tensor: @@ -1136,8 +1152,27 @@ def _routed_moe_model_config(model_config: ModelConfig) -> ModelConfig: routed_model_config._frozen = True return routed_model_config - def forward(self, hidden_states: torch.Tensor, all_rank_num_tokens=None) -> torch.Tensor: - """``hidden_states``: ``[num_tokens, hidden_size]`` bf16.""" + def tail_rp_eligible(self, hidden_states: torch.Tensor, step: Optional[DecodeStep]) -> bool: + """Whether this layer's forward on ``step`` can hand its output on as this rank's unreduced row-parallel + share (``partial_tail``): the decode path takes the step with that tail.""" + return self.decode_moe is not None and self.decode_moe.takes(hidden_states, step, True) + + def forward( + self, + hidden_states: torch.Tensor, + all_rank_num_tokens=None, + partial_tail: bool = False, + step: Optional[DecodeStep] = None, + ) -> Union[torch.Tensor, _decode_comm.PendingTail]: + """``hidden_states``: ``[num_tokens, hidden_size]`` bf16. ``step``: the step's classification + (``decode_step``); the decode path (``decode_moe.py``) runs the steps it takes. ``partial_tail`` (only where + ``tail_rp_eligible`` holds): return this rank's unreduced share of the output instead, a ``PendingTail`` at + most 8 tokens, a tensor on a wide decode step.""" + decode = self.decode_moe + if decode is not None and decode.takes(hidden_states, step, partial_tail): + return decode.forward(self, hidden_states, self.decode_gemvs, partial_tail) + if partial_tail: + raise RuntimeError("the row-parallel MoE tail needs the decode path to take the step") identity = hidden_states router_logits = self.gate.compute_logits(hidden_states) moe_all_reduce = self.routed_experts.all_reduce if self._reduce_routed_output else None @@ -1439,7 +1474,9 @@ def forward( capture: Optional[Tuple[Any, int]] = None, step: Optional[DecodeStep] = None, prenormed: bool = False, - ) -> Tuple[torch.Tensor, int]: + pending_moe_partial: Optional[Union[torch.Tensor, _decode_comm.PendingTail]] = None, + defer_moe_tail: bool = False, + ) -> Union[Tuple[torch.Tensor, int], Tuple[torch.Tensor, int, Any]]: """Port of HF ``KimiDecoderLayer._forward_attn_residual`` (per token). ``block_residual`` is a preallocated snapshot bank in kernel-native @@ -1466,13 +1503,55 @@ def forward( (``K3DecodeComm.sandwich_oproj``) where the sandwich takes the layer, the step and the attention's decode branch, else on the attention's unreduced o_proj output (``K3DecodeComm.allreduce_attn_res``). + + ``pending_moe_partial`` (only where ``accepts_moe_partial`` held): the + previous layer's MoE output, unreduced (``defer_moe_tail``); then + ``hidden_states`` is the prefix sum without it, and this layer's + pre-attention step reduces and adds it: a ``PendingTail`` in one + ``K3DecodeComm.sandwich_tail`` call, a wide decode step's tensor by its + all-reduce and the fused add + attn_res + RMSNorm. + + ``defer_moe_tail``: return ``(prefix_sum, num_snapshots, partial)`` + instead, ``partial`` this layer's MoE output unreduced, for the next + consumer to reduce. """ prefix_sum = hidden_states valid_block_residual = block_residual[:num_snapshots] attn_res_max_tokens = _attn_res_max_tokens(step) + tail = ( + pending_moe_partial + if isinstance(pending_moe_partial, _decode_comm.PendingTail) + else None + ) + # A snapshot layer whose pre-attention step is the sandwich tail has the kernel store the running prefix sum + # straight into the bank row it snapshots. + snapshot_row = None + if tail is not None and self.layer_idx % self.attn_res_block_size == 0: + snapshot_row = block_residual[num_snapshots] if prenormed: assert num_snapshots == 0 and self.layer_idx % self.attn_res_block_size == 0 + assert pending_moe_partial is None and capture is None + elif tail is not None: + hidden_states, prefix_sum = self.decode_comm.sandwich_tail( + tail, + prefix_sum, + valid_block_residual, + self.self_attention_res_proj, + self.self_attention_res_norm, + self.input_layernorm, + updated_out=snapshot_row, + ) + elif pending_moe_partial is not None: + prefix_sum, hidden_states = _apply_attn_res_add_and_rmsnorm( + prefix_sum, + _decode_comm.wide_all_reduce(self._o_allreduce(), pending_moe_partial), + valid_block_residual, + self.self_attention_res_proj, + self.self_attention_res_norm, + self.input_layernorm, + attn_res_max_tokens, + ) elif capture is not None: # The mixture tap needs the PRE-norm value, which the fused # attn-res + RMSNorm kernel does not expose. Keep the two steps @@ -1502,8 +1581,23 @@ def forward( else: hidden_states = self.input_layernorm(hidden_states) + if capture is not None and pending_moe_partial is not None: + # The tapped layer handed its MoE output on: the step above reduced it into prefix_sum. Tap that value's + # pre-norm attn_res mixture, what the split path captures. + tapped = ( + _apply_attn_res( + prefix_sum, + valid_block_residual, + self.self_attention_res_proj, + self.self_attention_res_norm, + ) + if _AUX_ATTN_RES_STREAM_ENABLED + else prefix_sum + ) + capture[0].maybe_capture_hidden_states(capture[1], tapped, None) + if self.layer_idx % self.attn_res_block_size == 0: - if not prenormed: + if not prenormed and snapshot_row is None: block_residual[num_snapshots].copy_(prefix_sum) num_snapshots += 1 valid_block_residual = block_residual[:num_snapshots] @@ -1540,7 +1634,7 @@ def forward( else: if comm is not None and step is not None and step.wide: partial = attention(hidden_states, attn_metadata, step=step, reduce_output=False) - attention_output = _decode_comm.wide_all_reduce(attention._o_allreduce, partial) + attention_output = _decode_comm.wide_all_reduce(self._o_allreduce(), partial) else: attention_output = attention(hidden_states, attn_metadata, step=step) if prefix_sum is None: @@ -1563,10 +1657,14 @@ def forward( self.post_attention_layernorm, attn_res_max_tokens, ) - if self.is_moe: - hidden_states = self.block_sparse_moe( - hidden_states, getattr(attn_metadata, "all_rank_num_tokens", None) + all_rank_num_tokens = getattr(attn_metadata, "all_rank_num_tokens", None) + if self.is_moe and defer_moe_tail: + partial = self.block_sparse_moe( + hidden_states, all_rank_num_tokens, partial_tail=True, step=step ) + return prefix_sum, num_snapshots, partial + if self.is_moe: + hidden_states = self.block_sparse_moe(hidden_states, all_rank_num_tokens, step=step) else: hidden_states = self._dense_mlp(hidden_states, step) @@ -1597,6 +1695,15 @@ def _o_proj(self) -> nn.Module: """This layer's attention output projection (row parallel).""" return self.linear_attn.o_proj if self.is_kda else self.self_attn.mixer.o_proj + def _o_allreduce(self) -> AllReduce: + """The all-reduce module of this layer's attention output; it also reduces a wide step's MoE partial.""" + return (self.linear_attn if self.is_kda else self.self_attn)._o_allreduce + + def accepts_moe_partial(self, num_snapshots: int) -> bool: + """Whether this layer's pre-attention step can reduce the previous layer's unreduced MoE output: the target + built its collective state, and the snapshot bank is not empty.""" + return num_snapshots > 0 and self.decode_comm is not None + def skip_forward( self, hidden_states: torch.Tensor, @@ -1676,6 +1783,29 @@ def kda_token_states(self) -> bool: and all(layer.linear_attn.takes_k3_kernels for layer in self.layers if layer.is_kda) ) + def _defer_moe_tail( + self, + i: int, + hidden_states: torch.Tensor, + num_snapshots: int, + spec_metadata, + capture_set, + step: Optional[DecodeStep], + ) -> bool: + """Whether layer ``i`` hands its MoE output on unreduced (the row-parallel tail, ``decode_moe.py``): its MoE + takes the step with that tail, and its consumer, the next layer's pre-attention step or the final norm after + the last layer, accepts it. A layer DSpark taps defers too: the next layer's step reduces the output, then taps + its pre-norm mixture. An unknown capture set (every layer tapped) and a tapped last layer keep the replicated + tail.""" + layer = self.layers[i] + if not (layer.is_moe and layer.block_sparse_moe.tail_rp_eligible(hidden_states, step)): + return False + if spec_metadata is not None and ( + capture_set is None or (layer.layer_idx in capture_set and i == len(self.layers) - 1) + ): + return False + return self.layers[min(i + 1, len(self.layers) - 1)].accepts_moe_partial(num_snapshots) + def forward( self, attn_metadata: AttentionMetadata, @@ -1721,6 +1851,9 @@ def forward( if spec_metadata is not None else None ) + # A MoE layer's output handed on unreduced, which the next layer's pre-attention step (or the final norm) + # reduces: a decode_comm.PendingTail, or a wide decode step's tensor. + pending_moe_partial = None for i, layer in enumerate(self.layers): # DFlash/DSpark hidden-state capture. The drafter is distilled on # the aggregated stream value -- the pre-norm softmax mixture its @@ -1739,7 +1872,10 @@ def forward( and (capture_set is None or self.layers[i - 1].layer_idx in capture_set) ): capture = (spec_metadata, self.layers[i - 1].layer_idx) - hidden_states, num_snapshots = layer( + defer_moe_tail = self._defer_moe_tail( + i, hidden_states, num_snapshots, spec_metadata, capture_set, step + ) + outputs = layer( hidden_states, block_residual, num_snapshots, @@ -1747,7 +1883,35 @@ def forward( capture=capture, step=step, prenormed=i == 0 and prenormed is not None, + pending_moe_partial=pending_moe_partial, + defer_moe_tail=defer_moe_tail, + ) + if defer_moe_tail: + hidden_states, num_snapshots, pending_moe_partial = outputs + else: + (hidden_states, num_snapshots), pending_moe_partial = outputs, None + + if isinstance(pending_moe_partial, _decode_comm.PendingTail): + normed, _ = self.layers[-1].decode_comm.sandwich_tail( + pending_moe_partial, + hidden_states, + block_residual[:num_snapshots], + self.output_attn_res_proj, + self.output_attn_res_norm, + self.norm, + ) + return normed + if pending_moe_partial is not None: + _, normed = _apply_attn_res_add_and_rmsnorm( + hidden_states, + _decode_comm.wide_all_reduce(self.layers[-1]._o_allreduce(), pending_moe_partial), + block_residual[:num_snapshots], + self.output_attn_res_proj, + self.output_attn_res_norm, + self.norm, + _attn_res_max_tokens(step), ) + return normed # The last layer has no successor, so this one recompute is # unavoidable -- output-side score weights, matching SGLang's @@ -2405,8 +2569,9 @@ def post_load_weights(self) -> None: every layer of its kind: the KDA projection's Lamport buffers and the MLA decode attention's workspace; the decode GEMVs' state (built by ``cache_derived_state``) handed to every attention module; and, where every attention all-reduce runs over MNNVL, the TP group's collective state (``K3DecodeComm``, collective: every - rank builds it here) handed to every layer, with whether its sandwich takes the layer's o_proj, and the - decode path's one-shot ceiling on every stock MNNVL all-reduce (``use_decode_one_shot``).""" + rank builds it here) handed to every layer, with whether its sandwich takes the layer's o_proj, the decode + path's one-shot ceiling on every stock MNNVL all-reduce (``use_decode_one_shot``), and the MoE decode path + (``_build_decode_moe``).""" super().post_load_weights() kda = [layer.linear_attn for layer in self.model.layers if layer.is_kda] mla = [layer.self_attn.mixer for layer in self.model.layers if not layer.is_kda] @@ -2433,6 +2598,7 @@ def post_load_weights(self) -> None: for layer in layers: layer.decode_comm = comm _decode_comm.use_decode_one_shot(self) + moe_layers = 0 if comm is None else self._build_decode_moe(comm) logger.info( "Kimi K3 decode kernels: KDA on k3_kda_decode_attn, k3_kda_attn and k3_kda_verify " f"({sum(m.takes_k3_kernels for m in kda)} / {len(kda)} layers take them), MLA on k3_mla_qkv and " @@ -2446,8 +2612,49 @@ def post_load_weights(self) -> None: if comm is not None else "unfused (an attention all-reduce does not run over MNNVL)" ) + + f"; MoE on k3_moe_front, k3_moe and the row-parallel tail ({moe_layers} layers)" ) + def _build_decode_moe(self, comm: _decode_comm.K3DecodeComm) -> int: + """The MoE decode path (``decode_moe.py``) on every MoE layer it takes: the shared state (collective: the + front's head workspace), each layer's decode weights with its latent norm folded into the latent up + projection, the decode GEMVs' state; then one call of every kernel of the path (collective: the front and + the sandwich tail). Every rank builds it here. Returns the number of MoE layers it takes.""" + mapping = self.model_config.mapping + moes = [layer.block_sparse_moe for layer in self.model.layers if layer.is_moe] + gaps = { + id(moe): _decode_moe.layout_gaps( + moe, mapping.tp_size, self.model.num_attn_res_snapshots + ) + for moe in moes + } + takes = [moe for moe in moes if not gaps[id(moe)]] + for reason in sorted({gap for moe in moes for gap in gaps[id(moe)]}): + logger.info_once( + f"Kimi K3 MoE decode path off on some layers: {reason}", + key=f"k3_decode_moe_off_{reason}", + ) + if not takes: + return 0 + backend = takes[0].routed_experts.backend + state = _decode_moe.K3DecodeMoe.create( + mapping, + backend.w3_w1_weight.device, + backend.w3_w1_weight.shape[1] // 2, + backend.expert_size_per_partition, + comm.mnnvl, + ) + for moe in takes: + _decode_moe.fold_latent_norm(moe) + moe.decode_moe = _decode_moe.K3DecodeMoeLayer.create( + moe, state, mapping.tp_rank, mapping.tp_size + ) + moe.decode_gemvs = self.model.decode_gemvs + first = takes[0].decode_moe + first.warm_up(takes[0]) + comm.compile_tail(takes[0].moe_hidden_size, first.shared_cols, first.tail_weight) + return len(takes) + def _check_step_contract(self, attn_metadata: AttentionMetadata) -> None: """First-forward checks of the engine surface and the per-engine settings.""" objects = { diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_gemv.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_gemv.py index f87c8c6a6832..135166cde8b2 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_gemv.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_gemv.py @@ -3,9 +3,10 @@ """The Kimi K3 target's decode GEMVs, LM head and embedding on the catalog's single-GPU entries (``decode_gemv.py`` of ``kimi_k3_mxfp4__sm_100__tp16_moetp4ep4``), on one GPU. -* Each GEMV site at every row count it takes (1..8, and 9..64 where k3_ctm_gemv_wide takes it): the bits of its - catalog entry's call, within 8e-3 of ``max |ref|`` of a float64 product (the sigmoid columns within 1e-2 of the - sigmoid of it); declined above its rows, at another shape or dtype, and under capture before an eager call. +* Each GEMV site at every row count it takes (1..8, and 9..its ``wide_rows`` where k3_ctm_gemv_wide takes it): the + bits of its catalog entry's call (fp32 at an ``out_fp32`` site), within 8e-3 of ``max |ref|`` of a float64 product + (the sigmoid columns within 1e-2 of the sigmoid of it); declined above its rows, at another shape or dtype, and + under capture before an eager call. * The LM head through ``K3LogitsProcessor`` on a real ``LMHead`` (one rank): fp32 logits from ``k3_head_gemv`` within the same bound, the stock processor's rows selected, the stock path above 8 rows and before the state is built. @@ -81,7 +82,7 @@ def _bits(t): def _entry(spec, rows, x, w): """The site's catalog entry called directly: what ``project`` must reproduce bit for bit.""" if rows > decode_gemv.MAX_ROWS or spec.small == "wide": - return k3_ctm_gemv_wide(x, w, sig_col0=spec.sig_col0) + return k3_ctm_gemv_wide(x, w, sig_col0=spec.sig_col0, out_fp32=spec.out_fp32) if spec.small == "decode": return k3_decode_gemv(x, w) return k3_ctm_gemv_long( @@ -90,7 +91,9 @@ def _entry(spec, rows, x, w): def _site_rows(spec): - return list(ROWS) + ([9, 16, 24, 40, 64] if spec.wide else []) + return list(ROWS) + ( + [m for m in (9, 16, 24, 32, 40, 64) if m <= spec.wide_rows] if spec.wide else [] + ) @pytest.mark.parametrize("site", list(decode_gemv.SITES)) @@ -103,18 +106,19 @@ def test_site(site): x = _rows(m, spec.k, seed=m) y = gemvs.project(site, x, w) assert y is not None and y.shape == (m, spec.n), (site, m) + assert y.dtype == (torch.float32 if spec.out_fp32 else torch.bfloat16), (site, y.dtype) assert torch.equal(_bits(y), _bits(_entry(spec, m, x, w))), (site, m) _check_product(y, x, w, spec.sig_col0) -@pytest.mark.parametrize("site", ["kv_a", "mla_ag"]) +@pytest.mark.parametrize("site", ["kv_a", "mla_ag", "moe_head"]) def test_site_declines(site): """More rows than the site takes, another weight shape, a non-bf16 input, and under capture before any eager call: None, nothing launched.""" spec = decode_gemv.SITES[site] w = _weight(spec.n, spec.k, seed=1) gemvs = decode_gemv.K3DecodeGemvs.create(None, sites=[site]) - too_many = decode_gemv.WIDE_ROWS + 1 if spec.wide else decode_gemv.MAX_ROWS + 1 + too_many = spec.wide_rows + 1 if spec.wide else decode_gemv.MAX_ROWS + 1 assert gemvs.project(site, _rows(too_many, spec.k, seed=9), w) is None assert ( gemvs.project(site, _rows(4, spec.k, seed=4), _weight(spec.n + 128, spec.k, seed=2)) is None From 1c9c28770423eb12a9348f339d104db600669aee Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 22:30:58 -0700 Subject: [PATCH 108/161] [None][test] modeling_v2 Kimi K3 target: the decode path's collectives at 4 ranks The target's own decoder layers at the TP16 per-rank shapes, on its real collective state (K3DecodeComm: the MNNVL and sandwich workspaces, built collectively) and the real MNNVL all-reduce. The attention and MoE modules are stand-ins that keep the interface the layer calls. The checks: - the fused post-attention step (comm/mnnvl_allreduce_attn_res, comm/k3_sandwich_oproj) against the built-in all-reduce + residual update, on decode, DSpark, prefill, unclassified and wide steps, including the path each step takes; - the sandwich against the MNNVL entry, bit for bit; - a deferred MoE tail run by the next layer's comm/k3_sandwich_tail, against the tail reduced in torch; - CUDA-graph capture and replay against eager, bit for bit; - every output equal across the ranks. One W-rank job, run through _rank_job like the catalog's collective matrices, and listed in l0_gb200_multi_gpus.yml. Signed-off-by: Vasanth Sabavat --- .../test-db/l0_gb200_multi_gpus.yml | 2 + .../comm/_kimi_k3_decode_comm_op_matrix.py | 612 ++++++++++++++++++ ...deling_v2_kimi_k3_decode_comm_op_matrix.py | 32 + 3 files changed, 646 insertions(+) create mode 100644 tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py create mode 100644 tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_kimi_k3_decode_comm_op_matrix.py diff --git a/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml b/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml index 02a9429501c4..56b07476585b 100644 --- a/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml +++ b/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml @@ -66,6 +66,8 @@ l0_gb200_multi_gpus: - unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_latent_reduce_op_matrix.py - unittest/_torch/modeling_v2/comm/test_modeling_v2_allgather_op_matrix.py - unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_front_op_matrix.py + # The Kimi K3 target's decode-path collectives over those entries, 4 ranks + - unittest/_torch/modeling_v2/comm/test_modeling_v2_kimi_k3_decode_comm_op_matrix.py - unittest/_torch/moe/test_moe_comm.py::TestMoEComm::test_mnnvl_checkpoint_preserves_moe_graph_addresses - unittest/_torch/moe/test_moe_comm.py::TestMoEComm::test_mnnvl_engine_checkpoint_coordination - unittest/_torch/moe/test_moe_comm.py::TestMoEComm::test_mnnvl_checkpoint_failure_is_collective_and_bounded diff --git a/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py new file mode 100644 index 000000000000..be6448fa7f28 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py @@ -0,0 +1,612 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""The Kimi K3 target's decode-path collectives (``decode_comm.py`` of ``kimi_k3_mxfp4__sm_100__tp16_moetp4ep4``) +against its built-in path, at W ranks of one GB200 tray (default 4), at the TP16 per-rank shapes. + + CUDA_VISIBLE_DEVICES=0,1,2,3 python _kimi_k3_decode_comm_op_matrix.py [--world-size 4] + srun -n 4 --mpi=pmix python _kimi_k3_decode_comm_op_matrix.py --launcher srun --world-size 4 + +Not a pytest module: one fixed sequence of checks inside one W-rank job, sharing the target's collective state. The +collected entry point is ``test_modeling_v2_kimi_k3_decode_comm_op_matrix.py``. + +Every rank builds the target's own decoder layers (KimiLinearDecoderLayer): KDA layer 12 (a snapshot layer, so the +post-attention step has no prefix sum), KDA layer 13 and MLA layer 15 with dense MLPs, and the MoE layers 22 (KDA), +23 (MLA) and 24 (a KDA snapshot layer). Their attention modules are stand-ins that keep the interface the layer calls: +``will_run_decode_branch`` (KDA: any step decode_step classifies; MLA: a decode step only), +``forward(..., reduce_output, project_output)``, a real bf16 o_proj [7168, 768] and the real stock AllReduce over +MNNVL; the core they hand over is set by the check. A dense layer's MLP is a recorder that returns zeros, so the layer +returns the post-attention ``updated`` and the recorder holds ``normed``. A MoE layer's MoE is a stand-in that hands +on a PendingTail of the TP16 per-rank shapes (latent [T, 3584], shared activation [T, 384], tail weight +[7168, 256 + 384]), or returns that tail reduced in torch and the stock all-reduce. + +Checks: + * takes_oproj: the TP16 o_proj shape only, bias-free, within the kernel's candidate count. + * use_decode_one_shot: every stock MNNVL all-reduce takes the decode ceiling. + * fused vs built-in: for each dense layer and step (decode of 1 / 3 / 8 tokens, DSpark 1 x 6, a 5-token prefill, + 12 / 16 unclassified tokens, a wide 2 x 8 step), the layer with the target's K3DecodeComm against the same layer + without it: ``updated`` and ``normed`` within TOL, and the path the call took (the attention's reduce_output / + project_output). + * sandwich vs the MNNVL entry on exact payloads (products and sums exact in fp32): bit for bit. + * graph: three layers chained, captured at an 8-token decode step (sandwich) and a 12-token unclassified step (MNNVL + entry), replayed with rewritten inputs: each replay equal to an eager run, bit for bit. + * MoE tail deferral: layer 22 hands its tail to layer 23 or 24, whose pre-attention step runs it as + comm/k3_sandwich_tail (layer 24's kernel stores the prefix sum into the bank row it takes), against the same layers + with the tail reduced in torch and added before the built-in pre-attention step: the consumer's attention input, + its outputs and the bank within TOL; the deferred chain captured and replayed equals eager bit for bit. + * Every output bitwise equal across the ranks. + +Every rank draws the replicated tensors (prefix sum, snapshot bank, norms) from one seed and its own o_proj, core and +tail from a rank seed. +""" + +import sys +from pathlib import Path +from types import SimpleNamespace + +import torch +from torch import nn + +sys.path.insert(0, str(Path(__file__).resolve().parent)) +import _lockstep as ls # noqa: E402 + +assert torch.cuda.is_available(), "the Kimi K3 target's decode collectives require CUDA devices" + +DEADLINE_S = 1500 +H = 7168 +K_IN = 768 +LATENT = 3584 +LAT_SLICE = 224 # a TP16 rank's latent columns +TAIL_LAT = 256 # the tail weight's latent part, zero-padded +TAIL_ACT = 384 # a TP16 rank's shared activation columns +MOE_LAYERS = (22, 23, 24) # KDA, MLA, KDA (a snapshot layer) +SNAPSHOTS = 2 # valid bank rows before a non-snapshot layer +BANK = 8 +TOL = 2e-2 # normed and updated (fused vs built-in): max |err| / max |ref| +LAYERS = (12, 13, 15) + +R = None +T = None # the target module +STATS = { + "bitwise_updated": 0, + "bitwise_normed": 0, + "cases": 0, + "max_err_normed": 0.0, + "max_err_updated": 0.0, +} + + +class _Attention(nn.Module): + """Stand-in for the target's attention module (see the module docstring).""" + + def __init__(self, mapping, strategy, decode_only: bool): + super().__init__() + from tensorrt_llm._torch.distributed import AllReduce + + self.o_proj = nn.Linear(K_IN, H, bias=False, dtype=torch.bfloat16, device="cuda") + self._o_allreduce = AllReduce(mapping=mapping, strategy=strategy, dtype=torch.bfloat16) + self.decode_only = decode_only + self.core = None + self.calls = [] + + def will_run_decode_branch(self, attn_metadata, step) -> bool: + return step is not None and (step.decode or not self.decode_only) + + def forward( + self, hidden_states, attn_metadata, step=None, reduce_output=True, project_output=True + ): + self.calls.append((reduce_output, project_output)) + self.last_input = hidden_states.clone() + core = self.core + assert core.shape[0] == hidden_states.shape[0], (core.shape, hidden_states.shape) + if not project_output: + if not self.will_run_decode_branch(attn_metadata, step): + raise ValueError("project_output=False off the decode branch") + return core + out = self.o_proj(core) + return self._o_allreduce(out) if reduce_output else out + + +class _KDA(_Attention): + def __init__( + self, + cfg, + layer_idx, + mapping=None, + allreduce_strategy=None, + aux_stream=None, + model_config=None, + ): + super().__init__(mapping, allreduce_strategy, decode_only=False) + + +class _MLARuntime(nn.Module): + def __init__(self, cfg, layer_idx, model_config, aux_stream_dict): + super().__init__() + self.mixer = _Attention( + model_config.mapping, model_config.allreduce_strategy, decode_only=True + ) + self._o_allreduce = self.mixer._o_allreduce + + def will_run_decode_branch(self, attn_metadata, step) -> bool: + return self.mixer.will_run_decode_branch(attn_metadata, step) + + def forward( + self, hidden_states, attn_metadata, step=None, reduce_output=True, project_output=True + ): + return self.mixer(hidden_states, attn_metadata, step, reduce_output, project_output) + + +class _Moe(nn.Module): + """Stand-in for KimiK3MoERuntime. With ``partial_tail`` it hands on ``pending`` (a PendingTail set by the check); + otherwise it returns that tail reduced in torch and the stock all-reduce (zeros without one) and records its input. + The reference tail: ``[RMSNorm(latent)[:, lo:lo + 224] | act] @ weight.T``, the RMS applied to the fp32 product + of the latent slice, rounded to bf16 per rank, then summed over the ranks.""" + + def __init__(self, model_config, cfg, layer_idx, aux_stream_dict): + super().__init__() + from tensorrt_llm._torch.distributed import AllReduce + + self.all_reduce = AllReduce(mapping=model_config.mapping, strategy=model_config.allreduce_strategy, + dtype=torch.bfloat16) # fmt: skip + self.pending = None + self.last_input = None + + def tail_rp_eligible(self, hidden_states, step) -> bool: + return self.pending is not None and step is not None and hidden_states.shape[0] <= 8 + + def forward(self, hidden_states, all_rank_num_tokens=None, partial_tail=False, step=None): + self.last_input = hidden_states.clone() + p = self.pending + if partial_tail: + return p + if p is None: + return torch.zeros_like(hidden_states) + lat, w = p.latent.float(), p.weight.float() + rs = torch.rsqrt(lat.pow(2).mean(-1, keepdim=True) + p.lat_eps) + y = (lat[:, p.lo : p.lo + LAT_SLICE] @ w[:, :LAT_SLICE].t()) * rs + p.act.float() @ w[ + :, TAIL_LAT: + ].t() + return self.all_reduce(y.bfloat16()) + + +def _attention(layer) -> _Attention: + return layer.linear_attn if layer.is_kda else layer.self_attn.mixer + + +def _gen(seed): + return torch.Generator(device="cuda").manual_seed(seed) + + +def _normal(g, shape, scale=1.0, offset=0.0): + return (offset + scale * torch.randn(shape, generator=g, device="cuda")).bfloat16() + + +class _Recorder: + """The layer's MLP: records its input (``normed``) and returns zeros, so the layer returns ``updated``.""" + + def __init__(self): + self.normed = None + + def __call__(self, hidden_states, step): + self.normed = hidden_states.clone() + return torch.zeros_like(hidden_states) + + +def build_layers(): + from tensorrt_llm._torch.model_config import ModelConfig + from tensorrt_llm._torch.utils import AuxStreamType + from tensorrt_llm.functional import AllReduceStrategy + + T.K3DecodeKDA = _KDA + T.KimiMLARuntime = _MLARuntime + T.KimiK3MoERuntime = _Moe + cfg = SimpleNamespace( + hidden_size=H, + num_experts=None, + first_k_dense_replace=1, + moe_layer_freq=1, + intermediate_size=2112 * R.world, + activation_situ_beta=1.0, + activation_situ_linear_beta=8.0, + rms_norm_eps=1e-5, + attn_res_block_size=12, + num_hidden_layers=93, + linear_attn_config=dict( + kda_layers=[i for i in range(1, 94) if i % 4], + full_attn_layers=[i for i in range(1, 94) if i % 4 == 0], + ), + ) + model_config = ModelConfig(mapping=R.mapping, allreduce_strategy=AllReduceStrategy.MNNVL) + stream = torch.cuda.Stream() + streams = {kind: stream for kind in AuxStreamType} + moe_cfg = SimpleNamespace(**{**vars(cfg), "num_experts": 896}) + layers = {} + with torch.device("cuda"): + for idx in LAYERS + MOE_LAYERS: + layer = T.KimiLinearDecoderLayer( + model_config, moe_cfg if idx in MOE_LAYERS else cfg, idx, streams + ) + assert layer._mnnvl_allreduce() is not None, ( + f"layer {idx}: the attention all-reduce is not MNNVL" + ) + assert layer.is_moe == (idx in MOE_LAYERS), idx + if not layer.is_moe: + layer.recorder = _Recorder() + layer._dense_mlp = layer.recorder + g = _gen(1000 + idx) # replicated + with torch.no_grad(): + for norm in ( + layer.input_layernorm, + layer.post_attention_layernorm, + layer.self_attention_res_norm, + layer.mlp_res_norm, + ): + norm.weight.copy_(_normal(g, norm.weight.shape, 0.1, 1.0)) + for proj in (layer.self_attention_res_proj, layer.mlp_res_proj): + proj.weight.copy_(_normal(g, proj.weight.shape, 0.05)) + # This rank's o_proj slice: exact payloads, so products and their sums are exact in fp32. + gr = _gen(2000 + 100 * idx + R.rank) + _attention(layer).o_proj.weight.copy_(ls.exact_bf16(gr, (H, K_IN), -4, 5, 1 / 16)) + layers[idx] = layer + return layers + + +def make_comm(layers): + takes = { + idx: T._decode_comm.K3DecodeComm.takes_oproj(layer._o_proj(), BANK) + for idx, layer in layers.items() + } + assert all(takes.values()), takes + oproj = layers[LAYERS[0]]._o_proj().weight + comm = T._decode_comm.K3DecodeComm.create(R.mapping, oproj) + comm.compile_tail( + LATENT, TAIL_ACT, torch.zeros(H, TAIL_LAT + TAIL_ACT, dtype=torch.bfloat16, device="cuda") + ) + return comm + + +class Case: + """One layer call's inputs: the prefix sum and bank (replicated) and this rank's core.""" + + def __init__(self, seed, idx, tokens): + g = _gen(seed) + self.x = _normal(g, (tokens, H), 1.0) + self.bank = _normal(g, (BANK, tokens, H), 1.0) + self.snapshots = 0 if idx % 12 == 0 else SNAPSHOTS + gr = _gen(seed * 31 + 7 + R.rank) + self.core = ls.exact_bf16(gr, (tokens, K_IN), -4, 5, 1 / 8) + + +def run(layer, case, step, comm, sandwich=True): + """``(updated, normed, calls)`` of one layer call on clones of the case's tensors.""" + attention = _attention(layer) + attention.core = case.core + attention.calls.clear() + layer.decode_comm = comm + layer.sandwich_oproj = sandwich + bank = case.bank.clone() + updated, _ = layer(case.x.clone(), bank, case.snapshots, SimpleNamespace(), step=step) + torch.cuda.synchronize() + return updated.clone(), layer.recorder.normed, list(attention.calls) + + +def _err(a, b): + return ((a.float() - b.float()).abs().max() / b.float().abs().max().clamp_min(1e-6)).item() + + +def STEPS(): + DS = T.DecodeStep + return [ + ("decode1", 1, DS(1, 1, 1)), + ("decode3", 3, DS(3, 3, 1)), + ("dspark6", 6, DS(6, 1, 6)), + ("decode8", 8, DS(8, 8, 1)), + ("prefill5", 5, DS(5)), # small, not a decode step + ("unclassified12", 12, None), + ("unclassified16", 16, None), + ("wide16", 16, DS(16, 2, 8)), + ] + + +def expected_call(layer, step): + """(reduce_output, project_output) the layer should pass on ``step`` with the comm.""" + if step is not None and step.wide: + return ( + False, + True, + ) # the partial, for the wide step's all-reduce ceiling (wide_all_reduce) + if step is not None and step.num_tokens <= 8 and (step.decode or layer.is_kda): + return (True, False) # the sandwich: the core + return (False, True) # the MNNVL entry: the unreduced partial + + +def check_decode_one_shot(): + """use_decode_one_shot raises every stock MNNVL all-reduce's ceiling; the check then restores the stock one, so + the built-in reference keeps main's choice.""" + mods = [m.mnnvl_allreduce for layer in LAYERS_BUILT.values() for m in layer.modules() + if getattr(m, "mnnvl_allreduce", None) is not None] # fmt: skip + assert len(mods) >= 2 * len(LAYERS), len(mods) # attention + dense MLP all-reduce per layer + stock = {id(m): m.one_shot_max_bytes for m in mods} + for layer in LAYERS_BUILT.values(): + T._decode_comm.use_decode_one_shot(layer) + assert all(m.one_shot_max_bytes == T._decode_comm.DECODE_AR_ONE_SHOT_MAX_BYTES for m in mods) + for m in mods: + m.one_shot_max_bytes = stock[id(m)] + + +def check_takes_oproj(): + K3DecodeComm = T._decode_comm.K3DecodeComm + good = nn.Linear(K_IN, H, bias=False, dtype=torch.bfloat16, device="cuda") + wide = nn.Linear(1024, H, bias=False, dtype=torch.bfloat16, device="cuda") + biased = nn.Linear(K_IN, H, bias=True, dtype=torch.bfloat16, device="cuda") + fp32 = nn.Linear(K_IN, H, bias=False, dtype=torch.float32, device="cuda") + assert K3DecodeComm.takes_oproj(good, BANK) + assert not K3DecodeComm.takes_oproj(wide, BANK) + assert not K3DecodeComm.takes_oproj(biased, BANK) + assert not K3DecodeComm.takes_oproj(fp32, BANK) + assert not K3DecodeComm.takes_oproj(good, 9), ( + "9 snapshots + the update exceed the kernel's 9 candidates" + ) + + +def check_fused_vs_builtin(): + seed = 1 + for idx in LAYERS: + layer = LAYERS_BUILT[idx] + for name, tokens, step in STEPS(): + seed += 1 + case = Case(seed, idx, tokens) + ref_u, ref_n, ref_calls = run(layer, case, step, None) + got_u, got_n, got_calls = run(layer, case, step, COMM) + assert ref_calls == [(True, True)], (idx, name, ref_calls) + assert got_calls == [expected_call(layer, step)], (idx, name, got_calls) + eu, en = _err(got_u, ref_u), _err(got_n, ref_n) + same_u, same_n = torch.equal(got_u, ref_u), torch.equal(got_n, ref_n) + STATS["cases"] += 1 + STATS["bitwise_updated"] += same_u + STATS["bitwise_normed"] += same_n + STATS["max_err_updated"] = max(STATS["max_err_updated"], eu) + STATS["max_err_normed"] = max(STATS["max_err_normed"], en) + if R.rank == 0: + print( + f"[rank 0] layer {idx} {name}: path {got_calls[0]} updated err {eu:.2e} bitwise {same_u}, " + f"normed err {en:.2e} bitwise {same_n}", + flush=True, + ) + assert torch.isfinite(got_n.float()).all() and torch.isfinite(got_u.float()).all() + assert eu < TOL and en < TOL, (idx, name, eu, en) + assert R.same_on_ranks(got_u, got_n), (idx, name, "ranks differ") + if step is not None and step.wide: + # The same all-reduce (one-shot at W <= 4 under either ceiling) and the same epilogue. + assert same_u and same_n, (idx, name, "a wide step keeps the built-in arithmetic") + + +def check_sandwich_vs_mnnvl_bitwise(): + seed = 500 + for idx in LAYERS: + layer = LAYERS_BUILT[idx] + for name, tokens, step in STEPS(): + if expected_call(layer, step) != (True, False): + continue + seed += 1 + case = Case(seed, idx, tokens) + sw_u, sw_n, sw_calls = run(layer, case, step, COMM, sandwich=True) + mn_u, mn_n, mn_calls = run(layer, case, step, COMM, sandwich=False) + assert sw_calls == [(True, False)] and mn_calls == [(False, True)], (sw_calls, mn_calls) + assert torch.equal(sw_u, mn_u), (idx, name, "updated", _err(sw_u, mn_u)) + assert torch.equal(sw_n, mn_n), (idx, name, "normed", _err(sw_n, mn_n)) + if R.rank == 0: + print(f"[rank 0] layer {idx} {name}: sandwich == mnnvl bit for bit", flush=True) + + +def _chain(cases, step): + """The three layers in order, each layer's input the previous one's output (the first takes cases[0].x).""" + x = cases[0].x_buf + outs = [] + for idx, case in zip(LAYERS, cases): + layer = LAYERS_BUILT[idx] + _attention(layer).core = case.core_buf + layer.decode_comm = COMM + layer.sandwich_oproj = True + x, _ = layer(x, case.bank_buf, case.snapshots, SimpleNamespace(), step=step) + outs.append((x, layer.recorder.normed)) + return outs + + +def check_graph_capture_and_replay(): + for name, tokens, step in (("decode8", 8, T.DecodeStep(8, 8, 1)), ("unclassified12", 12, None)): + cases = [Case(900 + i, idx, tokens) for i, idx in enumerate(LAYERS)] + for case in cases: + case.x_buf, case.bank_buf, case.core_buf = ( + case.x.clone(), + case.bank.clone(), + case.core.clone(), + ) + # Warm-up eagerly (the stock all-reduce sizes its workspace on its first call of a size), then capture. + _chain(cases, step) + torch.cuda.synchronize() + R.barrier() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + outs = _chain(cases, step) + for replay in range(3): + fresh = [Case(950 + 10 * replay + i, idx, tokens) for i, idx in enumerate(LAYERS)] + for case, new in zip(cases, fresh): + case.x_buf.copy_(new.x) + case.bank_buf.copy_(new.bank) + case.core_buf.copy_(new.core) + graph.replay() + torch.cuda.synchronize() + got = [(u.clone(), n.clone()) for u, n in outs] + # The eager reference on the same inputs (fresh buffers: the snapshot layer writes its bank row). + for case, new in zip(cases, fresh): + case.x_buf.copy_(new.x) + case.bank_buf.copy_(new.bank) + case.core_buf.copy_(new.core) + want = [(u.clone(), n.clone()) for u, n in _chain(cases, step)] + torch.cuda.synchronize() + for (gu, gn), (wu, wn) in zip(got, want): + assert torch.equal(gu, wu) and torch.equal(gn, wn), (name, replay) + assert R.same_on_ranks(*[t for pair in got for t in pair]), ( + name, + replay, + "ranks differ", + ) + if R.rank == 0: + print(f"[rank 0] graph {name}: 3 replays == eager bit for bit", flush=True) + del graph + + +def _pending(seed, tokens): + """A PendingTail of the TP16 per-rank shapes: the reduced latent replicated, this rank's shared activation and + tail weight (its latent padding columns zero), this rank's first latent column.""" + g = _gen(seed) + latent = _normal(g, (tokens, LATENT), 1.0) + gr = _gen(seed * 17 + 3 + R.rank) + act = _normal(gr, (tokens, TAIL_ACT), 0.5) + weight = _normal(gr, (H, TAIL_LAT + TAIL_ACT), 0.03) + weight[:, LAT_SLICE:TAIL_LAT] = 0 + return T._decode_comm.PendingTail(latent, act, weight.contiguous(), R.rank * LAT_SLICE, 1e-5) + + +def _tail_chain(producer, consumer, x, bank, step, pending, defer, consumer_core, producer_core): + """The producer layer (its MoE handing ``pending`` on with ``defer``, else adding its reduced tail), then the + consumer layer. Returns the consumer's attention input, its returned prefix sum, its MoE input and the bank.""" + _attention(producer).core = producer_core + _attention(consumer).core = consumer_core + producer.block_sparse_moe.pending = pending + consumer.block_sparse_moe.pending = None + for layer in (producer, consumer): + layer.decode_comm = COMM + layer.sandwich_oproj = True + snapshots = SNAPSHOTS + out = producer(x, bank, snapshots, SimpleNamespace(), step=step, defer_moe_tail=defer) + if defer: + prefix, snapshots, partial = out + assert partial is pending + else: + (prefix, snapshots), partial = out, None + prefix_c, snapshots_c = consumer( + prefix, bank, snapshots, SimpleNamespace(), step=step, pending_moe_partial=partial + ) + return ( + _attention(consumer).last_input, + prefix_c, + consumer.block_sparse_moe.last_input, + bank[:snapshots_c], + ) + + +def check_moe_tail_deferral(): + seed = 3000 + for consumer_idx in (23, 24): + producer, consumer = LAYERS_BUILT[22], LAYERS_BUILT[consumer_idx] + for tokens in (1, 3, 8): + seed += 1 + step = T.DecodeStep(tokens, tokens, 1) + case = Case(seed, 22, tokens) + gr = _gen(seed * 13 + 5 + R.rank) + core_c = ls.exact_bf16(gr, (tokens, K_IN), -4, 5, 1 / 8) + pending = _pending(seed, tokens) + got = _tail_chain(producer, consumer, case.x.clone(), case.bank.clone(), step, pending, True, core_c, + case.core) # fmt: skip + got = [t.clone() for t in got] + want = _tail_chain(producer, consumer, case.x.clone(), case.bank.clone(), step, pending, False, core_c, + case.core) # fmt: skip + torch.cuda.synchronize() + names = ("attention input", "prefix sum", "MoE input", "bank") + errs = [_err(a, b) for a, b in zip(got, want)] + if R.rank == 0: + print( + f"[rank 0] tail {22}->{consumer_idx} T={tokens}: " + + ", ".join(f"{n} {e:.2e}" for n, e in zip(names, errs)), + flush=True, + ) + assert all(torch.isfinite(t.float()).all() for t in got) + assert all(e < TOL for e in errs), (consumer_idx, tokens, dict(zip(names, errs))) + assert R.same_on_ranks(*got), (consumer_idx, tokens, "ranks differ") + + +def check_moe_tail_graph(): + producer, consumer = LAYERS_BUILT[22], LAYERS_BUILT[24] + tokens = 8 + step = T.DecodeStep(tokens, tokens, 1) + case = Case(4000, 22, tokens) + x, bank = case.x.clone(), case.bank.clone() + core_p, core_c = case.core.clone(), case.core.clone() + pending = _pending(4000, tokens) + + def chain(): + return _tail_chain(producer, consumer, x, bank, step, pending, True, core_c, core_p) + + chain() # eager warm-up + torch.cuda.synchronize() + R.barrier() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + outs = chain() + for replay in range(3): + fresh_case = Case(4100 + replay, 22, tokens) + fresh = _pending(4100 + replay, tokens) + + def load(): + x.copy_(fresh_case.x) + bank.copy_(fresh_case.bank) + core_p.copy_(fresh_case.core) + core_c.copy_(fresh_case.core.flip(0)) + for a, b in zip(pending[:3], fresh[:3]): + a.copy_(b) + + load() + graph.replay() + torch.cuda.synchronize() + got = [t.clone() for t in outs] + load() + want = [t.clone() for t in chain()] + torch.cuda.synchronize() + for a, b in zip(got, want): + assert torch.equal(a, b), replay + assert R.same_on_ranks(*got), (replay, "ranks differ") + if R.rank == 0: + print("[rank 0] graph tail 22->24 T=8: 3 replays == eager bit for bit", flush=True) + del graph + + +CHECKS = [ + check_takes_oproj, + check_decode_one_shot, + check_fused_vs_builtin, + check_sandwich_vs_mnnvl_bitwise, + check_graph_capture_and_replay, + check_moe_tail_deferral, + check_moe_tail_graph, +] + +LAYERS_BUILT = None +COMM = None + + +def _run_one_rank(args) -> int: + global R, T, LAYERS_BUILT, COMM + R = ls.Rank(args) + import importlib + + T = importlib.import_module( + "tensorrt_llm._torch._experimental.modeling_v2.models.kimi_k3_vl." + "kimi_k3_mxfp4__sm_100__tp16_moetp4ep4.modeling" + ) + with torch.inference_mode(): + LAYERS_BUILT = build_layers() + COMM = make_comm(LAYERS_BUILT) + code = ls.run_checks(R, CHECKS) + if R.rank == 0: + print(f"[rank 0] world {R.world}; {STATS}", flush=True) + return code + + +if __name__ == "__main__": + ARGS = ls.parse_args(sys.argv[1:]) + if ARGS.rank_worker or ARGS.launcher == "srun": + sys.exit(_run_one_rank(ARGS)) + ls.spawn(__file__, ARGS, DEADLINE_S) + print("OK") diff --git a/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_kimi_k3_decode_comm_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_kimi_k3_decode_comm_op_matrix.py new file mode 100644 index 000000000000..74f1c9733040 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_kimi_k3_decode_comm_op_matrix.py @@ -0,0 +1,32 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Collected entry point for the Kimi K3 target's decode-path collectives (``decode_comm.py`` of +``kimi_k3_mxfp4__sm_100__tp16_moetp4ep4``). + +The checks are ``_kimi_k3_decode_comm_op_matrix.py`` beside this file, its own W-rank launcher (the layers share the +target's collective state, so one job, not independent cases); see ``_rank_job`` for why that is left intact. +""" + +import _rank_job +import pytest +import torch + +assert torch.cuda.is_available(), "the Kimi K3 target's decode collectives require CUDA devices" + +if torch.cuda.get_device_capability() != (10, 0): + # The target is certified on sm_100 (GB200) only, as are the catalog entries it calls. + pytest.skip("the Kimi K3 target runs on sm_100 only", allow_module_level=True) + +if torch.cuda.device_count() < _rank_job.WORLD_SIZE: + # One rank per device: fewer visible devices cannot host the check's world size. + pytest.skip( + f"the check runs {_rank_job.WORLD_SIZE} ranks, one per device; " + f"{torch.cuda.device_count()} visible", + allow_module_level=True, + ) + + +# The case starts its own W-rank mpirun over the visible devices; under xdist several workers would fight for them. +@pytest.mark.no_xdist +def test_kimi_k3_decode_comm_op_matrix() -> None: + _rank_job.run("kimi_k3_decode_comm") From ba1db7e4ea7406412a27005cac0bf205d236889d Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 23:07:15 -0700 Subject: [PATCH 109/161] [None][chore] modeling_v2 Kimi K3 tp16_moetp16ep1: carry C3 into the copy The tp16_moetp4ep4 target gains the decode path's collectives (decode_comm.py: the post-attention all-reduce with the residual update, the sandwich kernels, the one-shot ceilings) and its MoE decode path (decode_moe.py). The copy here takes them: both modules byte for byte, and the changes to decode_gemv.py and modeling.py outside route B's blocks. One new route B block: this target builds no MoE decode path (moe_layers = 0 where tp16_moetp4ep4 calls _build_decode_moe). The path's state includes the wide k3_moe build, whose per-group tables do not fit 896 local experts, so the MoE layers keep the generic path on every step and no MoE tail is deferred. _build_decode_moe stays as tp16_moetp4ep4 has it, uncalled, so it logs no refusal reasons here; the startup line reports the MoE path on 0 layers. The module docstring's block describes the collectives without wide steps. Signed-off-by: Vasanth Sabavat --- .../decode_comm.py | 293 +++++++++ .../decode_gemv.py | 45 +- .../decode_moe.py | 357 +++++++++++ .../modeling.py | 559 +++++++++++++++--- 4 files changed, 1153 insertions(+), 101 deletions(-) create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_comm.py create mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_comm.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_comm.py new file mode 100644 index 000000000000..2ef7dbe27458 --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_comm.py @@ -0,0 +1,293 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""The decode path's residual-update collectives on the catalog's Kimi K3 MNNVL and sandwich entries. + +A layer's post-attention step adds the attention output to the running prefix sum, selects the attention residual +and applies the post-attention RMSNorm. Under TP the attention output is a sum over the TP group, so the step is that +all-reduce followed by the residual update. On a step of at most `AR_ATTN_RES_MAX_TOKENS` tokens (wide decode steps +aside) one collective runs both: + +* `K3DecodeComm.sandwich_oproj`, `comm/k3_sandwich_oproj`: o_proj, its all-reduce and the residual update in one + kernel, at most `SANDWICH_MAX_TOKENS` tokens of an o_proj of the TP16 per-rank shape [7168, 768]. The attention + hands over its gated o_proj input. +* `K3DecodeComm.allreduce_attn_res`, `comm/mnnvl_allreduce_attn_res`, everywhere else: the one-shot MNNVL all-reduce + of the attention's unreduced o_proj partial, with the residual update as its epilogue. + +The pre-attention step after a MoE layer whose row-parallel tail was handed on (a `PendingTail`, `decode_moe.py`) is +the same kind of collective: `K3DecodeComm.sandwich_tail`, `comm/k3_sandwich_tail`, runs the tail GEMV, its +all-reduce and the next layer's residual update (the final norm's, after the last layer) in one kernel. + +The state is one `MnnvlWorkspace` and one `K3SandwichWorkspace` of the TP group (`K3DecodeComm.create`): collective +over the group and eager, built by the target in `post_load_weights` before any CUDA-graph capture. Every rank must +make the same calls on each in the same order. Which call a step takes is decided from its token count and kind and +from the load-time layout alone, which every rank of the group shares. + +The plain MNNVL all-reduces the decode path keeps (the stock modules') send one-shot up to +`DECODE_AR_ONE_SHOT_MAX_BYTES` (`use_decode_one_shot`), except a wide decode step's (`wide_all_reduce`). +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import NamedTuple, Optional, Tuple + +import torch +from torch import nn + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.k3_sandwich_oproj import ( + K3SandwichWorkspace, + k3_sandwich_oproj, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.k3_sandwich_tail import ( + k3_sandwich_tail, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.mnnvl_allreduce_attn_res import ( + mnnvl_allreduce_attn_res, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.mnnvl_workspace import ( + MnnvlWorkspace, +) + +# The kernel's support predicate: metadata reads only. +from tensorrt_llm._torch.cute_dsl_kernels.k3_sandwich import op as _sandwich_op +from tensorrt_llm._torch.distributed import AllReduceParams + +# The most tokens of a step whose post-attention all-reduce runs one-shot with the residual update in its epilogue. +AR_ATTN_RES_MAX_TOKENS = 16 + +# The most tokens the sandwich kernel takes (one token tile). +SANDWICH_MAX_TOKENS = _sandwich_op.MAX_TOKENS + +# The MNNVL workspace's buffer size: a one-shot call pushes T x 7168 bf16 from each of 16 ranks, 3.5 MiB at T = 16. +MNNVL_BUFFER_BYTES = 4 << 20 + +# The one-shot ceiling of the stock MNNVL all-reduces on the decode path: 8 tokens x 7168 x 16 ranks x 2 B is +# 1.75 MiB, which the stock 1 MiB ceiling would send two-shot. +DECODE_AR_ONE_SHOT_MAX_BYTES = 4 << 20 + +# A wide decode step's ceiling: the stock 1 MiB. At 16 ranks two-shot is faster for every wide step's rows. +WIDE_AR_ONE_SHOT_MAX_BYTES = 1 << 20 + + +def use_decode_one_shot(model: nn.Module) -> None: + """Every stock MNNVL all-reduce of ``model`` sends one-shot up to `DECODE_AR_ONE_SHOT_MAX_BYTES`. Each grows its + workspace on its first eager call of a larger size, before the capture of that size.""" + for module in model.modules(): + mnnvl = getattr(module, "mnnvl_allreduce", None) + if mnnvl is not None: + mnnvl.one_shot_max_bytes = DECODE_AR_ONE_SHOT_MAX_BYTES + + +def wide_all_reduce(all_reduce: nn.Module, x: torch.Tensor) -> torch.Tensor: + """``all_reduce(x)`` of a wide decode step (a stock ``AllReduce`` module, no fusion): its MNNVL all-reduce with + the `WIDE_AR_ONE_SHOT_MAX_BYTES` ceiling, else the module itself.""" + mnnvl = getattr(all_reduce, "mnnvl_allreduce", None) + if mnnvl is not None: + out = mnnvl( + x.contiguous(), AllReduceParams(), one_shot_max_bytes=WIDE_AR_ONE_SHOT_MAX_BYTES + ) + if out is not None: + return out + return all_reduce(x) + + +class PendingTail(NamedTuple): + """A MoE layer's row-parallel tail left to its consumer's fused pre-attention step (``K3DecodeComm.sandwich_tail``): + the reduced latent ``[T, 3584]``, the shared experts' activation, the tail weight ``[latent up columns | padding | + shared down]``, this rank's first latent column and the latent norm's epsilon.""" + + latent: torch.Tensor + act: torch.Tensor + weight: torch.Tensor + lo: int + lat_eps: float + + +def _eps(norm: nn.Module) -> float: + """The epsilon of a KimiK3RMSNorm (``eps``) or a stock RMSNorm (``variance_epsilon``).""" + return float(norm.eps if hasattr(norm, "eps") else norm.variance_epsilon) + + +def _res_args(res_proj: nn.Module, res_norm: nn.Module, out_norm: nn.Module) -> tuple: + """The residual update's weights and epsilons in the entries' order.""" + return ( + res_proj.weight.reshape(-1), + res_norm.weight, + out_norm.weight, + _eps(res_norm), + _eps(out_norm), + ) + + +@dataclass(eq=False) +class K3DecodeComm: + """The decode path's collective state for the TP group: the `MnnvlWorkspace` and the `K3SandwichWorkspace` the + post-attention steps run on. Built by `create`; owned by the target and shared by every layer.""" + + mnnvl: MnnvlWorkspace + sandwich: K3SandwichWorkspace + + @classmethod + def create(cls, mapping, oproj_weight: Optional[torch.Tensor] = None) -> "K3DecodeComm": + """The state for ``mapping``'s TP group. Collective: every rank of the group calls it at the same point, + eagerly, before any CUDA-graph capture; each workspace fails on every rank or on none. + + ``oproj_weight``: an o_proj weight ``takes_oproj`` holds for. The sandwich kernel compiles here for its + shape, with one call on a zero row of a zero weight, so no capture compiles it; the call advances the + sandwich workspace on every rank alike.""" + state = cls( + MnnvlWorkspace.create(mapping, MNNVL_BUFFER_BYTES), + K3SandwichWorkspace.create(mapping), + ) + if oproj_weight is not None: + weight = torch.zeros_like(oproj_weight) + hidden = weight.shape[0] + ones = weight.new_ones(hidden) + k3_sandwich_oproj( + weight.new_zeros(1, weight.shape[1]), + weight, + None, + weight.new_zeros(0, 1, hidden), + weight.new_zeros(hidden), + ones, + ones, + 1e-6, + 1e-6, + state.sandwich, + ) + torch.cuda.synchronize(weight.device) + return state + + def compile_tail(self, latent_size: int, act_size: int, tail_weight: torch.Tensor) -> None: + """Compile the sandwich tail kernel for a MoE tail of ``latent_size`` latent and ``act_size`` shared columns + and ``tail_weight``'s shape, with one call on a zero row of a zero weight, before any capture. Collective: every + rank of the group makes the call; it advances the sandwich workspace on every rank alike.""" + weight = torch.zeros_like(tail_weight) + hidden = weight.shape[0] + ones = weight.new_ones(hidden) + k3_sandwich_tail( + weight.new_zeros(1, latent_size), + weight.new_zeros(1, act_size), + weight, + 0, + 1e-6, + None, + weight.new_zeros(0, 1, hidden), + weight.new_zeros(hidden), + ones, + ones, + 1e-6, + 1e-6, + self.sandwich, + ) + torch.cuda.synchronize(weight.device) + + def takes_post_attention(self, hidden_states: torch.Tensor, step) -> bool: + """Whether one of this state's collectives runs the post-attention step of a layer whose attention input is + ``hidden_states``: at most `AR_ATTN_RES_MAX_TOKENS` bf16 rows of a hidden size the MNNVL entry takes, on any + step but a wide decode step (``step.wide``), whose post-attention update stays the all-reduce and the fused + add + attn_res + RMSNorm.""" + if step is not None and step.wide: + return False + if hidden_states.dim() != 2 or hidden_states.dtype != torch.bfloat16: + return False + rows, hidden = hidden_states.shape + return ( + hidden % 1024 == 0 + and hidden <= 8192 + and 0 < rows <= min(AR_ATTN_RES_MAX_TOKENS, self.mnnvl.max_one_shot_tokens(hidden)) + ) + + @staticmethod + def takes_oproj(o_proj: nn.Module, max_snapshots: int) -> bool: + """Whether ``sandwich_oproj`` takes the layer of output projection ``o_proj`` on a step of at most + `SANDWICH_MAX_TOKENS` tokens, decided once the weights are final: a bias-free, contiguous bf16 weight of the + kernel's shape (the TP16 per-rank [7168, 768]) and a snapshot bank of at most ``max_snapshots`` rows the + kernel's candidate count holds.""" + weight = getattr(o_proj, "weight", None) + return ( + getattr(o_proj, "bias", None) is None + and isinstance(weight, torch.Tensor) + and weight.dim() == 2 + and weight.is_cuda + and _sandwich_op.supports(weight.new_empty((1, weight.shape[1])), weight, max_snapshots) + ) + + @staticmethod + def takes_tail( + latent: torch.Tensor, act: torch.Tensor, tail_weight: torch.Tensor, max_snapshots: int + ) -> bool: + """Whether ``sandwich_tail`` takes a MoE tail of these tensors' shapes (rows of ``latent`` and ``act``, at most + `SANDWICH_MAX_TOKENS`; ``tail_weight``) with a snapshot bank of at most ``max_snapshots`` rows: the TP16 + per-rank shapes.""" + return _sandwich_op.supports_tail(latent, act, tail_weight, max_snapshots) + + def allreduce_attn_res( + self, + partial: torch.Tensor, + prefix_sum: Optional[torch.Tensor], + block_residual: torch.Tensor, + res_proj: nn.Module, + res_norm: nn.Module, + out_norm: nn.Module, + ) -> Tuple[torch.Tensor, torch.Tensor]: + """``(normed, updated)`` in one ``comm/mnnvl_allreduce_attn_res`` call: ``updated = prefix_sum + + allreduce(partial)`` (the sum alone without ``prefix_sum``), ``normed = out_norm(attn_res(block_residual..., + updated))``, the attention residual selected with ``res_proj`` and ``res_norm``. ``block_residual`` holds the + valid snapshots, ``[S, T, H]``.""" + return mnnvl_allreduce_attn_res( + partial.contiguous(), + prefix_sum, + block_residual, + *_res_args(res_proj, res_norm, out_norm), + self.mnnvl, + ) + + def sandwich_tail( + self, + pending: PendingTail, + prefix_sum: Optional[torch.Tensor], + block_residual: torch.Tensor, + res_proj: nn.Module, + res_norm: nn.Module, + out_norm: nn.Module, + updated_out: Optional[torch.Tensor] = None, + ) -> Tuple[torch.Tensor, torch.Tensor]: + """``(normed, updated)`` as ``allreduce_attn_res`` of a MoE layer's row-parallel tail ``pending`` + (``[RMSNorm(latent)[:, lo:lo + 224] | act] @ weight.T``), in one ``comm/k3_sandwich_tail`` call. + ``updated_out``: a bf16 ``[T, H]`` tensor the call stores ``updated`` into (the consumer's snapshot bank row), + returned as ``updated``.""" + return k3_sandwich_tail( + pending.latent, + pending.act, + pending.weight, + pending.lo, + pending.lat_eps, + prefix_sum, + block_residual, + *_res_args(res_proj, res_norm, out_norm), + self.sandwich, + updated_out=updated_out, + ) + + def sandwich_oproj( + self, + core: torch.Tensor, + o_weight: torch.Tensor, + prefix_sum: Optional[torch.Tensor], + block_residual: torch.Tensor, + res_proj: nn.Module, + res_norm: nn.Module, + out_norm: nn.Module, + ) -> Tuple[torch.Tensor, torch.Tensor]: + """``(normed, updated)`` as ``allreduce_attn_res`` of ``core @ o_weight.T``, in one ``comm/k3_sandwich_oproj`` + call: ``core`` is this rank's gated o_proj input ``[T, 768]``, ``o_weight`` its o_proj slice. Bit for bit + o_proj followed by ``allreduce_attn_res``.""" + return k3_sandwich_oproj( + core.contiguous(), + o_weight, + prefix_sum, + block_residual, + *_res_args(res_proj, res_norm, out_norm), + self.sandwich, + ) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_gemv.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_gemv.py index b384010b65dd..0e26e3971d0e 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_gemv.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_gemv.py @@ -4,10 +4,11 @@ * **Per-site GEMVs** (`K3DecodeGemvs.project`): a projection of a decode step runs on the kernel measured fastest at its call site's weight shape (`SITES`): at most `MAX_ROWS` rows on `gemm/k3_decode_gemv`, - `gemm/k3_ctm_gemv_wide` or `gemm/k3_ctm_gemv_long`, and, where the site lists it, up to `WIDE_ROWS` rows on - `gemm/k3_ctm_gemv_wide`. The sites are the decode kernels' fused projections (MLA's [W_a; W_g] with the gate rows - through a sigmoid, KDA's [q | k | v | g | f_a | b]), the attention output projection, and the built-in MLA path's - q_a / kv_a, q_b and gate projections. + `gemm/k3_ctm_gemv_wide` or `gemm/k3_ctm_gemv_long`, and, where the site lists it, more rows (up to `WIDE_ROWS`) + on `gemm/k3_ctm_gemv_wide`. The sites are the decode kernels' fused projections (MLA's [W_a; W_g] with the gate + rows through a sigmoid, KDA's [q | k | v | g | f_a | b]), the attention output projection, the built-in MLA path's + q_a / kv_a, q_b and gate projections, and the MoE decode path's projections (`decode_moe.py`), two of them with + fp32 outputs. * **LM head** (`K3LogitsProcessor`): at most `MAX_ROWS` rows of this rank's vocabulary shard on `gemm/k3_head_gemv` over the target's `K3HeadGemvWorkspace`, then the shards gathered (`comm/allgather`) as the stock head gathers them. It is the shell's logits processor, so the speculative worker's target logits and the @@ -72,9 +73,10 @@ @dataclass(frozen=True) class Site: """A call site's weight shape (this target's per-rank shapes) and its kernels: ``small`` at 1..MAX_ROWS rows - ("decode", "wide" or "long"), and k3_ctm_gemv_wide at MAX_ROWS+1..WIDE_ROWS rows where ``wide``. Output columns - from ``sig_col0`` on are stored through a sigmoid. ``split`` / ``ring`` / ``push``: k3_ctm_gemv_long's CTAs per - 128-row weight tile, weight-ring stages, and whether the partial sums are pushed to each row's owner.""" + ("decode", "wide" or "long"), and k3_ctm_gemv_wide at MAX_ROWS+1..``wide_rows`` rows where ``wide``. Output columns + from ``sig_col0`` on are stored through a sigmoid; ``out_fp32``: an fp32 output (k3_ctm_gemv_wide only). + ``split`` / ``ring`` / ``push``: k3_ctm_gemv_long's CTAs per 128-row weight tile, weight-ring stages, and whether + the partial sums are pushed to each row's owner.""" n: int k: int @@ -84,6 +86,8 @@ class Site: split: int = 0 ring: int = 0 push: bool = False + out_fp32: bool = False + wide_rows: int = WIDE_ROWS SITES: Dict[str, Site] = { @@ -101,6 +105,14 @@ class Site: # Layer 0's dense MLP split over the 16-way TP group: gate_up [gate 2112 | up 2112] and down. "dense_gate_up": Site(4224, 7168, "long", split=4, ring=5), "dense_down": Site(7168, 2112, "long", split=2, ring=6), + # The MoE decode path (decode_moe.py): this rank's head slice [latent down 224 | router 56] with an fp32 output, + # the shared experts' gate_up, the row-parallel tail [latent up 224 | padding 32 | shared down 384] and the + # replicated tail's latent up projection with an fp32 output. The wide kernel takes the head slice and the tail up + # to 32 rows, where it is faster than the stock GEMM. + "moe_head": Site(280, 7168, "wide", wide=True, out_fp32=True, wide_rows=32), + "moe_shared_gate_up": Site(768, 7168, "wide", wide=True), + "moe_tail": Site(7168, 640, "wide", wide=True, wide_rows=32), + "moe_up": Site(7168, 3584, "wide", out_fp32=True), } @@ -115,13 +127,15 @@ def _run( ) -> Optional[torch.Tensor]: """``spec``'s ``kernel`` on dense rows ``x2d``, or None where it does not take them.""" if kernel == "decode": - if spec.sig_col0 >= 0 or not _decode_op.supports(x2d, weight): + if spec.sig_col0 >= 0 or spec.out_fp32 or not _decode_op.supports(x2d, weight): return None return k3_decode_gemv(x2d, weight) if kernel == "wide": - if not _ctm_op.supports_wide(x2d, weight, spec.sig_col0, False): + if not _ctm_op.supports_wide(x2d, weight, spec.sig_col0, spec.out_fp32): return None - return k3_ctm_gemv_wide(x2d, weight, sig_col0=spec.sig_col0) + return k3_ctm_gemv_wide(x2d, weight, sig_col0=spec.sig_col0, out_fp32=spec.out_fp32) + if spec.out_fp32: + return None # One wave of the GPU's SMs: beyond it the long GEMV loses to the others. sms = torch.cuda.get_device_properties(x2d.device).multi_processor_count if math.ceil(spec.n / 128) * spec.split > sms or not _ctm_op.supports_long( @@ -224,6 +238,8 @@ def create( spec = SITES[site] weight = torch.zeros(spec.n, spec.k, dtype=torch.bfloat16, device=device) for rows in (1, 16, 32, 64) if spec.wide else (1,): + if rows > spec.wide_rows: + continue state._project(site, weight.new_zeros(rows, spec.k), weight, warm=True) del weight if "dense_gate_up" in sites: @@ -245,9 +261,10 @@ def create( return state def project(self, site: str, x: torch.Tensor, weight: torch.Tensor) -> Optional[torch.Tensor]: - """``x @ weight.T`` (bf16 ``[..., N]``, the site's sigmoid columns through the sigmoid) for ``site``'s weight - on its decode kernel, or None where none takes the call: more rows than the site's kernels take, another - shape or dtype, or, under capture, a kernel that has not run eagerly. The caller then runs its GEMM.""" + """``x @ weight.T`` (``[..., N]``: bf16, the site's sigmoid columns through the sigmoid; fp32 at an + ``out_fp32`` site) for ``site``'s weight on its decode kernel, or None where none takes the call: more rows + than the site's kernels take, another shape or dtype, or, under capture, a kernel that has not run eagerly. + The caller then runs its GEMM.""" return self._project(site, x, weight, warm=False) def _project( @@ -266,7 +283,7 @@ def _project( rows = x.numel() // spec.k if 0 < rows <= MAX_ROWS: kernel = spec.small - elif spec.wide and MAX_ROWS < rows <= WIDE_ROWS: + elif spec.wide and MAX_ROWS < rows <= spec.wide_rows: kernel = "wide" else: return None diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py new file mode 100644 index 000000000000..75f5b06ed66e --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py @@ -0,0 +1,357 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""The decode path's latent MoE on the catalog's Kimi K3 MoE entries. + +**At most 8 tokens** (`MAX_TOKENS`), a MoE layer runs: + +* `moe/k3_moe_front`, one kernel: this rank's slice of the MoE head (its latent-down rows and its router rows) as one + GEMV, the slices' all-gather over the TP group's `K3MoeHeadWorkspace`, the top-16 routing, the MXFP8 latent, and + the shared experts' gate_up + SiTU; +* `moe/k3_moe`: this rank's routed partial over its experts (on a `K3MoeState`); +* the latent all-reduce: the routed experts' all-reduce (one-shot, as `decode_comm.use_decode_one_shot` sets); +* the tail. The latent norm's weight is folded into the latent up projection at load, so + `[RMSNorm(latent) slice | shared activation] @ [latent up columns | shared down]` is this rank's share of the MoE + output. Where the next layer's pre-attention step (or the final norm) reduces it, the layer hands it on unreduced + as a `PendingTail`, which the consumer runs with its all-reduce and residual update as one `comm/k3_sandwich_tail` + kernel (`decode_comm.py`). Elsewhere the replicated tail runs: the latent RMS applied to the fp32 output of one + GEMV with the folded latent up weight, plus the shared experts' down projection and its all-reduce. + +**A wide decode step** (9 to 64 tokens, `WIDE_MAX_TOKENS`) keeps the sharded head and the row-parallel tail on +M-general ops: the head GEMV, `comm/mnnvl_allgather_split`, then `moe/k3_route_quant` and `moe/k3_moe` (on a +`K3MoeWideState`) beside the shared gate_up + SiTU, the latent all-reduce, and one GEMV of +`[RMSNorm(latent) slice | padding | shared activation]` with the tail weight. That is this rank's unreduced share, +which the consumer reduces with a plain all-reduce. + +The GEMVs run on the decode GEMV sites of `decode_gemv.py` where they take the call, else on the stock GEMM ops. + +`K3DecodeMoe` holds what every MoE layer shares: the head workspace (collective over the TP group), the two +`k3_moe` builds' scratch, and the TP group's MNNVL workspace (`decode_comm.K3DecodeComm`'s). `K3DecodeMoeLayer` +holds one layer's decode weights and its `k3_moe` counters. The target builds both in `post_load_weights`, before +any CUDA-graph capture, and runs every kernel once there so none compiles under a capture. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Optional, Union + +import torch +from torch import nn + +from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.mnnvl_allgather_split import ( + mnnvl_allgather_split, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.mnnvl_workspace import ( + MnnvlWorkspace, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.k3_moe import ( + K3MoeHeadWorkspace, + K3MoeLayer, + K3MoeState, + K3MoeWideState, + is_supported, + k3_moe, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.k3_moe_front import ( + front_weight, + k3_moe_front, + weight_supported, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.k3_route_quant import k3_route_quant +from tensorrt_llm._torch.modules.multi_stream_utils import maybe_execute_in_parallel + +from .decode_comm import K3DecodeComm, PendingTail, wide_all_reduce + +# The most tokens of the front and the small k3_moe build (one token tile), and of the wide build. +MAX_TOKENS = 8 +WIDE_MAX_TOKENS = 64 + +# The tail weight's latent columns are zero-padded to whole 128-column k-tiles of the tail kernels +# (comm/k3_sandwich_tail takes the TP16 tail weight [7168, 256 + 384]). +TAIL_K_TILE = 256 + + +@dataclass(eq=False) +class K3DecodeMoe: + """What every MoE layer's decode path shares on one device: the TP group's ``K3MoeHeadWorkspace`` (the front's + all-gather), the ``k3_moe`` builds for up to 8 and up to 64 tokens (their scratch; the layers run one at a time on + one stream), and the TP group's ``MnnvlWorkspace`` for a wide step's head all-gather. Built by `create`.""" + + head: K3MoeHeadWorkspace + small: K3MoeState + wide: K3MoeWideState + mnnvl: MnnvlWorkspace + + @classmethod + def create( + cls, mapping, device, i_tp: int, num_local: int, mnnvl: MnnvlWorkspace + ) -> "K3DecodeMoe": + """The state for ``mapping``'s TP group on ``device``, for experts of ``i_tp`` intermediate columns per rank, + ``num_local`` of them on this rank. Collective (the head workspace): every rank of the group calls it at the + same point, eagerly, before any CUDA-graph capture.""" + return cls( + K3MoeHeadWorkspace.create(mapping), + K3MoeState(device, i_tp, num_local), + K3MoeWideState(device, i_tp, num_local), + mnnvl, + ) + + +def _experts(moe: nn.Module) -> tuple: + """The routed experts' TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers ``k3_moe`` reads in place.""" + backend = moe.routed_experts.backend + return ( + backend.w3_w1_weight, + backend.w3_w1_weight_scale, + backend.w2_weight, + backend.w2_weight_scale, + ) + + +def layout_gaps(moe: nn.Module, tp_size: int, max_snapshots: int) -> list: + """Why the decode path does not take MoE layer ``moe`` (empty when it takes it), read once the routed experts' + buffers exist: the checks in order, up to the first that fails. ``max_snapshots``: the model's snapshot bank + rows.""" + gate, backend = moe.gate, moe.routed_experts.backend + shared = moe.shared_experts + gate_up, down = shared.gate_up_proj.weight, shared.down_proj.weight + if not moe._reduce_routed_output: + return ["a routed output the model does not reduce"] + if (moe.num_experts, moe.top_k, moe.moe_hidden_size, moe.hidden_size) != (896, 16, 3584, 7168): + return [ + f"experts / top-k / latent / hidden {(moe.num_experts, moe.top_k, moe.moe_hidden_size)}" + ] + if (gate.num_expert_group, gate.topk_group) != (1, 1): + return ["grouped routing"] + if moe.moe_hidden_size % (8 * tp_size) or moe.num_experts % (4 * tp_size): + return [f"a head slice of TP {tp_size}"] + latent = (gate.weight, moe.routed_expert_down_proj.weight, moe.routed_expert_up_proj.weight) + if not isinstance(moe.routed_expert_up_proj, nn.Linear) or any( + w.dtype != torch.bfloat16 for w in latent + ): + return ["latent projections or a router other than bf16"] + if ( + gate_up.dtype != torch.bfloat16 + or down.dtype != torch.bfloat16 + or shared.gate_up_proj.bias is not None + or shared.down_proj.bias is not None + or down.shape[0] != moe.hidden_size + ): + return ["a shared expert other than bf16 and unbiased"] + names = ( + "w3_w1_weight", + "w3_w1_weight_scale", + "w2_weight", + "w2_weight_scale", + "expert_size_per_partition", + ) + if ( + not all(hasattr(backend, name) for name in names) + or not is_supported(*_experts(moe), backend.expert_size_per_partition)[0] + ): + return ["routed-expert buffers k3_moe does not read"] + shared_cols, width = gate_up.shape[0] // 2, moe.moe_hidden_size // tp_size + if not weight_supported(tp_size, shared_cols, gate_up.shape[1], gate_up.device): + return ["a MoE front of this TP size and shared width"] + tail_cols = width + (-width % TAIL_K_TILE) + shared_cols + probe = gate_up.new_empty + if not K3DecodeComm.takes_tail( + probe(1, moe.moe_hidden_size), + probe(1, shared_cols), + probe(moe.hidden_size, tail_cols), + max_snapshots, + ): + return ["a row-parallel tail the sandwich tail kernel does not take"] + return [] + + +def fold_latent_norm(moe: nn.Module) -> None: + """Fold the latent RMSNorm's weight into the latent up projection's columns; the norm keeps a weight of ones, so + every step computes the same function and the tails may normalize before slicing.""" + up, norm = moe.routed_expert_up_proj, moe.routed_expert_norm + with torch.no_grad(): + up.weight.mul_(norm.weight.to(up.weight.dtype)[None, :]) + norm.weight.fill_(1) + + +@dataclass(eq=False) +class K3DecodeMoeLayer: + """One MoE layer's decode path: the front weight (this rank's head slice padded to whole tiles, then the shared + gate_up re-ordered), the head slice (a view of it), the row-parallel tail weight + ``[latent up columns lo:lo+width | padding | shared down]``, and its ``k3_moe`` handles on the shared builds. + Built by `create` once the weights are final (and the latent norm folded).""" + + state: K3DecodeMoe + front_weight: torch.Tensor + head_weight: torch.Tensor + tail_weight: torch.Tensor + tail_pad: Optional[torch.Tensor] + lo: int + width: int + shared_cols: int + small: K3MoeLayer + wide: K3MoeLayer + + @classmethod + def create( + cls, moe: nn.Module, state: K3DecodeMoe, tp_rank: int, tp_size: int + ) -> "K3DecodeMoeLayer": + """Build MoE layer ``moe``'s decode weights and its handles on ``state``'s builds (``layout_gaps`` empty).""" + width = moe.moe_hidden_size // tp_size + experts = moe.num_experts // tp_size + gate_up = moe.shared_experts.gate_up_proj.weight + shared_down = moe.shared_experts.down_proj.weight + up = moe.routed_expert_up_proj.weight + lo = tp_rank * width + pad = -width % TAIL_K_TILE + with torch.no_grad(): + head = torch.cat( + [ + moe.routed_expert_down_proj.weight[lo : lo + width], + moe.gate.weight[tp_rank * experts : (tp_rank + 1) * experts], + ] + ) + front = front_weight(head, gate_up) + parts = [up[:, lo : lo + width]] + if pad: + parts.append(up.new_zeros(up.shape[0], pad)) + parts.append(shared_down) + tail = torch.cat(parts, dim=1).contiguous() + weights = _experts(moe) + return cls( + state=state, + front_weight=front, + head_weight=front[: head.shape[0]], + tail_weight=tail, + tail_pad=up.new_zeros(WIDE_MAX_TOKENS, pad) if pad else None, + lo=lo, + width=width, + shared_cols=gate_up.shape[0] // 2, + small=state.small.layer(*weights), + wide=state.wide.layer(*weights), + ) + + def warm_up(self, moe: nn.Module) -> None: + """One call of every kernel of the decode path on zero inputs (M = 1), so none compiles under a capture: the + front (collective: every rank makes the same call), both ``k3_moe`` builds and ``k3_route_quant``.""" + device = self.front_weight.device + x = torch.zeros(1, moe.hidden_size, dtype=torch.bfloat16, device=device) + ids, weights, x_fp8, x_sf, _ = self._front(moe, x) + offset = moe.routed_experts.backend.slot_start + k3_moe(x_fp8, x_sf, ids, weights, offset, self.small) + logits = torch.zeros(1, moe.num_experts, dtype=torch.float32, device=device) + latent = torch.zeros(1, moe.moe_hidden_size, dtype=torch.bfloat16, device=device) + ids, weights, x_fp8, x_sf = k3_route_quant( + logits, moe.gate.e_score_correction_bias, latent, float(moe.gate.routed_scaling_factor), + early_trigger=True, + ) # fmt: skip + k3_moe(x_fp8, x_sf, ids, weights, offset, self.wide) + torch.cuda.synchronize(device) + + def _front(self, moe: nn.Module, x: torch.Tensor): + return k3_moe_front( + x.contiguous(), + self.front_weight, + moe.gate.e_score_correction_bias, + float(moe.gate.routed_scaling_factor), + self.shared_cols, + *moe._situ_betas, + self.state.head, + ) + + def takes(self, hidden_states: torch.Tensor, step, partial_tail: bool) -> bool: + """Whether this decode path runs the layer on ``step`` (bf16 rows ``hidden_states``): any step of at most + `MAX_TOKENS` tokens; with ``partial_tail``, also a wide decode step.""" + rows = hidden_states.shape[0] + if hidden_states.dtype != torch.bfloat16 or step is None: + return False + return 0 < rows <= MAX_TOKENS or (partial_tail and step.wide and rows <= WIDE_MAX_TOKENS) + + def forward( + self, moe: nn.Module, hidden_states: torch.Tensor, gemvs, partial_tail: bool + ) -> Union[torch.Tensor, PendingTail]: + """The MoE output of ``hidden_states`` (``takes`` holds): with ``partial_tail``, this rank's unreduced share + (a ``PendingTail`` at most `MAX_TOKENS` tokens); else the reduced output. ``gemvs``: the decode GEMVs' state, + or None.""" + if hidden_states.shape[0] > MAX_TOKENS: + return self._wide(moe, hidden_states, gemvs) + ids, weights, x_fp8, x_sf, shared_act = self._front(moe, hidden_states) + routed = k3_moe( + x_fp8, x_sf, ids, weights, moe.routed_experts.backend.slot_start, self.small + ) + latent = moe.routed_experts.all_reduce(routed) + if partial_tail: + return PendingTail( + latent.contiguous(), + shared_act.contiguous(), + self.tail_weight, + self.lo, + float(moe.routed_expert_norm.variance_epsilon), + ) + # The replicated tail: the latent RMS on the fp32 accumulator of the folded latent up projection, plus the + # shared experts' reduced output, rounded to bf16 once. + shared = moe.shared_experts + shared_out = shared.down_proj(shared_act, layer_idx=shared.layer_idx) + up = _gemv(gemvs, "moe_up", latent, moe.routed_expert_up_proj.weight, out_fp32=True) + scale = torch.rsqrt( + latent.float().pow(2).mean(-1, keepdim=True) + moe.routed_expert_norm.variance_epsilon + ) + return (up * scale + shared_out.float()).bfloat16() + + def _wide(self, moe: nn.Module, hidden_states: torch.Tensor, gemvs) -> torch.Tensor: + """A wide decode step's MoE: this rank's unreduced share of the output, ``[M, hidden]`` bf16.""" + x = hidden_states.contiguous() + head = _gemv(gemvs, "moe_head", x, self.head_weight, out_fp32=True) + routed_in, router_logits = mnnvl_allgather_split(head, self.width, self.state.mnnvl) + shared = moe.shared_experts + + def _routed_partial(): + ids, weights, x_fp8, x_sf = k3_route_quant( + router_logits, moe.gate.e_score_correction_bias, routed_in.contiguous(), + float(moe.gate.routed_scaling_factor), early_trigger=True, + ) # fmt: skip + return k3_moe( + x_fp8, x_sf, ids, weights, moe.routed_experts.backend.slot_start, self.wide + ) + + def _shared_activation(): + return shared._apply_activation( + _gemv(gemvs, "moe_shared_gate_up", x, shared.gate_up_proj.weight) + ) + + routed, shared_act = maybe_execute_in_parallel( + _routed_partial, + _shared_activation, + moe.moe_main_event, + moe.moe_shared_event, + moe.shared_expert_stream, + disable_on_compile=True, + ) + # The latent norm's weight is folded into the tail weight: normalize the whole latent row, keep this rank's + # columns. + normed = moe.routed_expert_norm(wide_all_reduce(moe.routed_experts.all_reduce, routed)) + parts = [normed[:, self.lo : self.lo + self.width]] + if self.tail_pad is not None: + parts.append(self.tail_pad[: x.shape[0]]) + parts.append(shared_act) + return _gemv(gemvs, "moe_tail", torch.cat(parts, dim=1), self.tail_weight) + + +def _gemv( + gemvs, site: str, x: torch.Tensor, weight: torch.Tensor, out_fp32: bool = False +) -> torch.Tensor: + """``x @ weight.T`` on ``site``'s decode GEMV where it takes the call (``decode_gemv.K3DecodeGemvs.project``), + else on the stock GEMM: cuBLAS through ``trtllm::dsv3_router_gemm_op`` for an fp32 output, ``F.linear`` else.""" + y = None if gemvs is None else gemvs.project(site, x, weight) + if y is not None: + return y + if out_fp32: + # The op reads the weight with a leading dimension of K: a strided weight would give wrong values. + if not weight.is_contiguous(): + raise ValueError( + f"the fp32 GEMM needs a contiguous weight, got strides {weight.stride()}" + ) + return torch.ops.trtllm.dsv3_router_gemm_op( + x.contiguous(), weight.t(), bias=None, out_dtype=torch.float32 + ) + return torch.nn.functional.linear(x, weight) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py index f357e1b4854c..9bf03a16eb78 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py @@ -24,8 +24,7 @@ classification decides which kernels each module runs: * **small**: at most 8 tokens (`DECODE_MAX_TOKENS`, one token tile), context requests included. The token-count - kernels take it: decode GEMVs, the MoE front and routed experts, the sandwiches, the embedding and residual - epilogues. + kernels take it: decode GEMVs, the sandwiches, the embedding and residual epilogues. * **decode**: a pure decode step of R <= 8 generation requests and no context request; without speculation each request has one token. The request-aware kernels take it: MLA attention and its KV store. @@ -43,11 +42,19 @@ then `attention/k3_mla_attn_vb_out` (the attention, v_b and the output gate in one launch). The projections around them run on the decode GEMV sites of `decode_gemv.py` (the [W_a; W_g] projection, `o_proj` -on every classified step), as do the LM head, the embedding and layer 0's dense MLP. The state those kernels share -(the KDA projection's Lamport buffers, the MLA attention workspace, the decode GEMVs' state) lives in typed objects -this target creates in `post_load_weights`, before any graph capture. The MoE front and routed experts, the -sandwiches and the residual epilogues come with their own entries; until then they run the generic path on every -step. +on every classified step), as do the LM head, the embedding and layer 0's dense MLP. A classified step's +attention-residual epilogues (the selection and the RMSNorm after it) take the fused kernels up to one token tile. On +any step of at most 16 tokens, the post-attention all-reduce carries the residual update (`decode_comm.py`): at most 8 +tokens on an attention decode branch, o_proj, the all-reduce and the update are one `comm/k3_sandwich_oproj` kernel; +otherwise the attention's unreduced o_proj output goes through `comm/mnnvl_allreduce_attn_res`. The stock MNNVL +all-reduces send one-shot up to 4 MiB. The state those kernels share (the KDA projection's Lamport buffers, the MLA +attention workspace, the decode GEMVs' state, the TP group's MNNVL and sandwich workspaces) lives in typed objects this +target creates in `post_load_weights`, before any graph capture. + +The MoE layers run the generic path on every step. `decode_moe.py` is `tp16_moetp4ep4`'s MoE decode path, which this +target does not build (its wide `k3_moe` build does not fit 896 local experts); this layout's own MoE engines +(`moe/k3_moe_m1` and `moe/k3_moe_m2` at one and two tokens, `moe/k3_moe` over all 896 experts up to 8) come with their +wiring as route B blocks. **What this target asserts rather than adapts**: SM 10.0; the topology above, with the expert split set explicitly; no speculative decoding; the MXFP4 checkpoint's quantization (W4A16_MXFP4 with no per-layer declarations, so the @@ -74,7 +81,7 @@ import math import os from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, Callable, Dict, Literal, NamedTuple, Optional, Tuple +from typing import TYPE_CHECKING, Any, Callable, Dict, Literal, NamedTuple, Optional, Tuple, Union import torch from torch import nn @@ -104,6 +111,7 @@ from tensorrt_llm._torch.modules.gated_mlp import GatedMLP from tensorrt_llm._torch.modules.kimi_k3_mla import KimiK3MLAAttention from tensorrt_llm._torch.modules.kimi_kda import KimiKDALinearAttention +from tensorrt_llm._torch.modules.kimi_kda.kimi_kda_mixer import maybe_bcg_kda_core_inplace from tensorrt_llm._torch.modules.multi_stream_utils import maybe_execute_in_parallel from tensorrt_llm._torch.modules.rms_norm import RMSNorm from tensorrt_llm._torch.modules.situ import SituAndMul @@ -123,7 +131,9 @@ from tensorrt_llm.mapping import Mapping from tensorrt_llm.models.modeling_utils import QuantAlgo, QuantConfig +from . import decode_comm as _decode_comm from . import decode_gemv as _decode_gemv +from . import decode_moe as _decode_moe from . import weights as _weights if TYPE_CHECKING: @@ -159,6 +169,14 @@ "k3_head_gemv", "k3_embed_norm", "allgather", + # The decode path's collectives (decode_comm.py) and MoE (decode_moe.py). + "mnnvl_allreduce_attn_res", + "k3_sandwich_oproj", + "k3_sandwich_tail", + "k3_moe_front", + "k3_moe", + "k3_route_quant", + "mnnvl_allgather_split", ) #: The engine surface the first forward checks before this target relies on it: per object, the attributes read. @@ -187,6 +205,7 @@ "tensorrt_llm._torch.models.modeling_utils.DecoderModel", # The text model's stock modules. "tensorrt_llm._torch.modules.kimi_kda.KimiKDALinearAttention", + "tensorrt_llm._torch.modules.kimi_kda.kimi_kda_mixer.maybe_bcg_kda_core_inplace", # >>> route B: no trtllm::kda_mtp_decode registration (the built-in KDA verify's) # <<< route B "tensorrt_llm._torch.modules.kimi_k3_mla.KimiK3MLAAttention", @@ -452,20 +471,20 @@ def _persistent_attn_res_applicable(M: int, H: int, N: int) -> bool: return H == 7168 and 2 <= N <= 9 -def _use_persistent_attn_res(M: int, H: int, N: int) -> bool: +def _use_persistent_attn_res(M: int, H: int, N: int, max_fused_tokens: int) -> bool: """Pick between the two fused kernels for this call site. ``persistent`` takes the persistent kernel at every shape it implements; - ``split`` takes it only above ``KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS`` tokens, - which stands in for the prefill/decode boundary. Shapes the persistent - kernel does not implement fall through to the caller's existing gate and - land on the unfused path. + ``split`` takes it only above ``max_fused_tokens`` tokens, which stands in + for the prefill/decode boundary. Shapes the persistent kernel does not + implement fall through to the caller's existing gate and land on the + unfused path. """ if not _persistent_attn_res_applicable(M, H, N): return False if _ATTN_RES_TOPOLOGY == "persistent": return True - return _ATTN_RES_TOPOLOGY == "split" and M > _FUSED_ATTN_RES_MAX_TOKENS + return _ATTN_RES_TOPOLOGY == "split" and M > max_fused_tokens def _apply_attn_res_fused( @@ -532,8 +551,12 @@ def _apply_attn_res_rmsnorm_fused( proj: nn.Linear, norm: KimiK3RMSNorm, output_norm: nn.Module, + max_fused_tokens: Optional[int] = None, ) -> Optional[torch.Tensor]: - """Fuse attention-residual mixing with its immediately following norm.""" + """Fuse attention-residual mixing with its immediately following norm (at most ``max_fused_tokens`` tokens, + default ``KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS``).""" + if max_fused_tokens is None: + max_fused_tokens = _FUSED_ATTN_RES_MAX_TOKENS if ( prefix_sum.dtype is not torch.bfloat16 or not prefix_sum.is_cuda @@ -543,10 +566,10 @@ def _apply_attn_res_rmsnorm_fused( M, H = prefix_sum.shape K = int(block_residual.shape[0]) N = K + 1 - # The fused path is taken for M <= _FUSED_ATTN_RES_MAX_TOKENS, H == 7168 and - # N <= 12, which is the measured window; larger token counts have not been - # measured and fall back to the unfused add + attn_res_fwd + RMSNorm path. - if _use_persistent_attn_res(M, H, N): + # The fused path is taken for M <= max_fused_tokens, H == 7168 and N <= 12, + # which is the measured window; larger token counts have not been measured + # and fall back to the unfused add + attn_res_fwd + RMSNorm path. + if _use_persistent_attn_res(M, H, N, max_fused_tokens): try: persistent_op = torch.ops.trtllm.attn_res_add_rmsnorm_persistent_fwd except (AttributeError, RuntimeError): @@ -564,7 +587,7 @@ def _apply_attn_res_rmsnorm_fused( _note_attn_res_fusion("attn_res+norm/persistent", True, M, H, N) return output.reshape(M, H) - if M > _FUSED_ATTN_RES_MAX_TOKENS or H != 7168 or N > 12: + if M > max_fused_tokens or H != 7168 or N > 12: _note_attn_res_fusion("attn_res+norm", False, M, H, N) return None try: @@ -593,8 +616,10 @@ def _apply_attn_res_add_rmsnorm_fused( proj: nn.Linear, norm: KimiK3RMSNorm, output_norm: nn.Module, + max_fused_tokens: Optional[int] = None, ) -> Optional[Tuple[torch.Tensor, torch.Tensor]]: - """Fuse ``prefix_sum + addend``, attention-residual, and trailing norm. + """Fuse ``prefix_sum + addend``, attention-residual, and trailing norm (``max_fused_tokens`` as in + ``_apply_attn_res_rmsnorm_fused``). The production residual add produces a BF16 tensor that remains live across the following MLP. The kernel therefore returns that materialized, @@ -602,6 +627,8 @@ def _apply_attn_res_add_rmsnorm_fused( while avoiding a separate add launch and a re-read of the intermediate by attention-residual selection. """ + if max_fused_tokens is None: + max_fused_tokens = _FUSED_ATTN_RES_MAX_TOKENS if ( prefix_sum.dtype is not torch.bfloat16 or addend.dtype is not torch.bfloat16 @@ -615,7 +642,7 @@ def _apply_attn_res_add_rmsnorm_fused( K = int(block_residual.shape[0]) N = K + 1 # Same measured window as _apply_attn_res_rmsnorm_fused above. - if _use_persistent_attn_res(M, H, N): + if _use_persistent_attn_res(M, H, N, max_fused_tokens): try: persistent_op = torch.ops.trtllm.attn_res_add_rmsnorm_persistent_fwd except (AttributeError, RuntimeError): @@ -633,7 +660,7 @@ def _apply_attn_res_add_rmsnorm_fused( _note_attn_res_fusion("add+attn_res+norm/persistent", True, M, H, N) return updated_prefix_sum.reshape(M, H), output.reshape(M, H) - if M > _FUSED_ATTN_RES_MAX_TOKENS or H != 7168 or N > 12: + if M > max_fused_tokens or H != 7168 or N > 12: _note_attn_res_fusion("add+attn_res+norm", False, M, H, N) return None try: @@ -691,10 +718,14 @@ def _apply_attn_res_and_rmsnorm( proj: nn.Linear, norm: KimiK3RMSNorm, output_norm: nn.Module, + max_fused_tokens: Optional[int] = None, ) -> torch.Tensor: - """Apply attention-residual selection and the next RMSNorm.""" + """Apply attention-residual selection and the next RMSNorm. ``max_fused_tokens``: the largest token count the + fused kernel takes (default ``KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS``).""" if _FUSED_ATTN_RES_ENABLED: - fused = _apply_attn_res_rmsnorm_fused(prefix_sum, block_residual, proj, norm, output_norm) + fused = _apply_attn_res_rmsnorm_fused( + prefix_sum, block_residual, proj, norm, output_norm, max_fused_tokens + ) if fused is not None: return fused return output_norm(_apply_attn_res(prefix_sum, block_residual, proj, norm)) @@ -707,20 +738,36 @@ def _apply_attn_res_add_and_rmsnorm( proj: nn.Linear, norm: KimiK3RMSNorm, output_norm: nn.Module, + max_fused_tokens: Optional[int] = None, ) -> Tuple[torch.Tensor, torch.Tensor]: - """Add an attention output to the running residual, then select and norm.""" + """Add an attention output to the running residual, then select and norm (``max_fused_tokens`` as in + ``_apply_attn_res_and_rmsnorm``).""" if _FUSED_ATTN_RES_ENABLED: fused = _apply_attn_res_add_rmsnorm_fused( - prefix_sum, addend, block_residual, proj, norm, output_norm + prefix_sum, addend, block_residual, proj, norm, output_norm, max_fused_tokens ) if fused is not None: return fused updated_prefix_sum = prefix_sum + addend return updated_prefix_sum, _apply_attn_res_and_rmsnorm( - updated_prefix_sum, block_residual, proj, norm, output_norm + updated_prefix_sum, block_residual, proj, norm, output_norm, max_fused_tokens ) +# A wide decode step's residual epilogues take the fused add + attn_res + RMSNorm kernels up to this many tokens, where +# they are faster than the add -> attn_res -> RMSNorm chain. +_WIDE_ATTN_RES_MAX_TOKENS = 32 + + +def _attn_res_max_tokens(step: Optional[DecodeStep]) -> Optional[int]: + """The most tokens of ``step`` the fused attn_res kernels take: one token tile (``DECODE_MAX_TOKENS``) on a step + ``decode_step`` classifies, ``_WIDE_ATTN_RES_MAX_TOKENS`` on a wide decode step; None on any other step (the + generic path's ``KIMI_K3_FUSED_ATTN_RES_MAX_TOKENS``).""" + if step is None: + return None + return _WIDE_ATTN_RES_MAX_TOKENS if step.wide else DECODE_MAX_TOKENS + + _K3_ROUTED_EXPERT_KEY_SUFFIXES = ("block_sparse_moe.experts", "mlp.experts") @@ -818,6 +865,7 @@ def __init__( raise ValueError("Kimi K3 runtime expects latent_moe_use_norm=True") situ_beta, situ_linear_beta = _resolve_kimi_situ_betas(cfg) + self._situ_betas = (situ_beta, situ_linear_beta) dtype = torch.bfloat16 # Routing scores stay fp32; the gate GEMM runs bf16xbf16 with fp32 @@ -942,6 +990,9 @@ def __init__( self.routed_expert_norm = RMSNorm( hidden_size=self.moe_hidden_size, eps=cfg.rms_norm_eps, dtype=dtype ) + # The decode path (decode_moe.py) and the decode GEMVs' state, set by the target's post_load_weights. + self.decode_moe: Optional[_decode_moe.K3DecodeMoeLayer] = None + self.decode_gemvs: Optional[_decode_gemv.K3DecodeGemvs] = None @staticmethod def _routed_projection(hidden_states: torch.Tensor, projection: nn.Module) -> torch.Tensor: @@ -1105,8 +1156,27 @@ def _routed_moe_model_config(model_config: ModelConfig) -> ModelConfig: routed_model_config._frozen = True return routed_model_config - def forward(self, hidden_states: torch.Tensor, all_rank_num_tokens=None) -> torch.Tensor: - """``hidden_states``: ``[num_tokens, hidden_size]`` bf16.""" + def tail_rp_eligible(self, hidden_states: torch.Tensor, step: Optional[DecodeStep]) -> bool: + """Whether this layer's forward on ``step`` can hand its output on as this rank's unreduced row-parallel + share (``partial_tail``): the decode path takes the step with that tail.""" + return self.decode_moe is not None and self.decode_moe.takes(hidden_states, step, True) + + def forward( + self, + hidden_states: torch.Tensor, + all_rank_num_tokens=None, + partial_tail: bool = False, + step: Optional[DecodeStep] = None, + ) -> Union[torch.Tensor, _decode_comm.PendingTail]: + """``hidden_states``: ``[num_tokens, hidden_size]`` bf16. ``step``: the step's classification + (``decode_step``); the decode path (``decode_moe.py``) runs the steps it takes. ``partial_tail`` (only where + ``tail_rp_eligible`` holds): return this rank's unreduced share of the output instead, a ``PendingTail`` at + most 8 tokens, a tensor on a wide decode step.""" + decode = self.decode_moe + if decode is not None and decode.takes(hidden_states, step, partial_tail): + return decode.forward(self, hidden_states, self.decode_gemvs, partial_tail) + if partial_tail: + raise RuntimeError("the row-parallel MoE tail needs the decode path to take the step") identity = hidden_states router_logits = self.gate.compute_logits(hidden_states) moe_all_reduce = self.routed_experts.all_reduce if self._reduce_routed_output else None @@ -1260,15 +1330,27 @@ def __init__( aux_stream_dict=aux_stream_dict, ) + def will_run_decode_branch( + self, attn_metadata: AttentionMetadata, step: Optional[DecodeStep] + ) -> bool: + """Whether the mixer runs ``step`` on the decode kernels (``K3DecodeMLA.will_run_decode_branch``).""" + return self.mixer.will_run_decode_branch(attn_metadata, step) + def forward( self, hidden_states: torch.Tensor, attn_metadata: AttentionMetadata, step: Optional[DecodeStep] = None, + reduce_output: bool = True, + project_output: bool = True, ) -> torch.Tensor: + """``reduce_output=False`` returns ``o_proj``'s TP partial (no all-reduce); ``project_output=False`` the + mixer's gated attention output before ``o_proj``, only where ``will_run_decode_branch`` holds.""" # MLA.forward takes position_ids first; K3 is NoPE, so pass None. - out = self.mixer(None, hidden_states, attn_metadata, step=step) - if self._o_allreduce is not None: + out = self.mixer( + None, hidden_states, attn_metadata, step=step, project_output=project_output + ) + if project_output and reduce_output and self._o_allreduce is not None: # Head-sharded TP: sum the row-sharded o_proj partials across # the head-shard group. out = self._o_allreduce(out) @@ -1382,6 +1464,10 @@ def __init__( self.mlp_res_norm = KimiK3RMSNorm(cfg.hidden_size, eps=cfg.rms_norm_eps, dtype=dtype) self.self_attention_res_proj = nn.Linear(cfg.hidden_size, 1, bias=False, dtype=dtype) self.mlp_res_proj = nn.Linear(cfg.hidden_size, 1, bias=False, dtype=dtype) + # The decode path's collectives (decode_comm.py), and whether their sandwich takes this layer's o_proj; set + # by the target's post_load_weights. + self.decode_comm: Optional[_decode_comm.K3DecodeComm] = None + self.sandwich_oproj = False def forward( self, @@ -1392,7 +1478,9 @@ def forward( capture: Optional[Tuple[Any, int]] = None, step: Optional[DecodeStep] = None, prenormed: bool = False, - ) -> Tuple[torch.Tensor, int]: + pending_moe_partial: Optional[Union[torch.Tensor, _decode_comm.PendingTail]] = None, + defer_moe_tail: bool = False, + ) -> Union[Tuple[torch.Tensor, int], Tuple[torch.Tensor, int, Any]]: """Port of HF ``KimiDecoderLayer._forward_attn_residual`` (per token). ``block_residual`` is a preallocated snapshot bank in kernel-native @@ -1412,12 +1500,62 @@ def forward( ``prenormed`` (layer 0 on a decode step): ``hidden_states`` already is this layer's input norm, and the layer's input, the step's embedding, already is in ``block_residual[0]`` (``K3DecodeGemvs.embed_norm``). + + The post-attention step of at most ``AR_ATTN_RES_MAX_TOKENS`` tokens + runs the attention's all-reduce and the residual update in one + collective once the target has built its ``decode_comm``: with o_proj + (``K3DecodeComm.sandwich_oproj``) where the sandwich takes the layer, + the step and the attention's decode branch, else on the attention's + unreduced o_proj output (``K3DecodeComm.allreduce_attn_res``). + + ``pending_moe_partial`` (only where ``accepts_moe_partial`` held): the + previous layer's MoE output, unreduced (``defer_moe_tail``); then + ``hidden_states`` is the prefix sum without it, and this layer's + pre-attention step reduces and adds it: a ``PendingTail`` in one + ``K3DecodeComm.sandwich_tail`` call, a wide decode step's tensor by its + all-reduce and the fused add + attn_res + RMSNorm. + + ``defer_moe_tail``: return ``(prefix_sum, num_snapshots, partial)`` + instead, ``partial`` this layer's MoE output unreduced, for the next + consumer to reduce. """ prefix_sum = hidden_states valid_block_residual = block_residual[:num_snapshots] + attn_res_max_tokens = _attn_res_max_tokens(step) + tail = ( + pending_moe_partial + if isinstance(pending_moe_partial, _decode_comm.PendingTail) + else None + ) + # A snapshot layer whose pre-attention step is the sandwich tail has the kernel store the running prefix sum + # straight into the bank row it snapshots. + snapshot_row = None + if tail is not None and self.layer_idx % self.attn_res_block_size == 0: + snapshot_row = block_residual[num_snapshots] if prenormed: assert num_snapshots == 0 and self.layer_idx % self.attn_res_block_size == 0 + assert pending_moe_partial is None and capture is None + elif tail is not None: + hidden_states, prefix_sum = self.decode_comm.sandwich_tail( + tail, + prefix_sum, + valid_block_residual, + self.self_attention_res_proj, + self.self_attention_res_norm, + self.input_layernorm, + updated_out=snapshot_row, + ) + elif pending_moe_partial is not None: + prefix_sum, hidden_states = _apply_attn_res_add_and_rmsnorm( + prefix_sum, + _decode_comm.wide_all_reduce(self._o_allreduce(), pending_moe_partial), + valid_block_residual, + self.self_attention_res_proj, + self.self_attention_res_norm, + self.input_layernorm, + attn_res_max_tokens, + ) elif capture is not None: # The mixture tap needs the PRE-norm value, which the fused # attn-res + RMSNorm kernel does not expose. Keep the two steps @@ -1442,43 +1580,95 @@ def forward( self.self_attention_res_proj, self.self_attention_res_norm, self.input_layernorm, + attn_res_max_tokens, ) else: hidden_states = self.input_layernorm(hidden_states) + if capture is not None and pending_moe_partial is not None: + # The tapped layer handed its MoE output on: the step above reduced it into prefix_sum. Tap that value's + # pre-norm attn_res mixture, what the split path captures. + tapped = ( + _apply_attn_res( + prefix_sum, + valid_block_residual, + self.self_attention_res_proj, + self.self_attention_res_norm, + ) + if _AUX_ATTN_RES_STREAM_ENABLED + else prefix_sum + ) + capture[0].maybe_capture_hidden_states(capture[1], tapped, None) + if self.layer_idx % self.attn_res_block_size == 0: - if not prenormed: + if not prenormed and snapshot_row is None: block_residual[num_snapshots].copy_(prefix_sum) num_snapshots += 1 valid_block_residual = block_residual[:num_snapshots] prefix_sum = None - if self.is_kda: - hidden_states = self.linear_attn(hidden_states, attn_metadata, step=step) - else: - hidden_states = self.self_attn(hidden_states, attn_metadata, step=step) - - if prefix_sum is None: - prefix_sum = hidden_states - hidden_states = _apply_attn_res_and_rmsnorm( - prefix_sum, - valid_block_residual, - self.mlp_res_proj, - self.mlp_res_norm, - self.post_attention_layernorm, - ) + attention = self.linear_attn if self.is_kda else self.self_attn + comm = self.decode_comm + if comm is not None and comm.takes_post_attention(hidden_states, step): + if ( + self.sandwich_oproj + and step is not None + and step.num_tokens <= _decode_comm.SANDWICH_MAX_TOKENS + and attention.will_run_decode_branch(attn_metadata, step) + ): + core = attention(hidden_states, attn_metadata, step=step, project_output=False) + hidden_states, prefix_sum = comm.sandwich_oproj( + core, + self._o_proj().weight, + prefix_sum, + valid_block_residual, + self.mlp_res_proj, + self.mlp_res_norm, + self.post_attention_layernorm, + ) + else: + partial = attention(hidden_states, attn_metadata, step=step, reduce_output=False) + hidden_states, prefix_sum = comm.allreduce_attn_res( + partial, + prefix_sum, + valid_block_residual, + self.mlp_res_proj, + self.mlp_res_norm, + self.post_attention_layernorm, + ) else: - prefix_sum, hidden_states = _apply_attn_res_add_and_rmsnorm( - prefix_sum, - hidden_states, - valid_block_residual, - self.mlp_res_proj, - self.mlp_res_norm, - self.post_attention_layernorm, + if comm is not None and step is not None and step.wide: + partial = attention(hidden_states, attn_metadata, step=step, reduce_output=False) + attention_output = _decode_comm.wide_all_reduce(self._o_allreduce(), partial) + else: + attention_output = attention(hidden_states, attn_metadata, step=step) + if prefix_sum is None: + prefix_sum = attention_output + hidden_states = _apply_attn_res_and_rmsnorm( + prefix_sum, + valid_block_residual, + self.mlp_res_proj, + self.mlp_res_norm, + self.post_attention_layernorm, + attn_res_max_tokens, + ) + else: + prefix_sum, hidden_states = _apply_attn_res_add_and_rmsnorm( + prefix_sum, + attention_output, + valid_block_residual, + self.mlp_res_proj, + self.mlp_res_norm, + self.post_attention_layernorm, + attn_res_max_tokens, + ) + all_rank_num_tokens = getattr(attn_metadata, "all_rank_num_tokens", None) + if self.is_moe and defer_moe_tail: + partial = self.block_sparse_moe( + hidden_states, all_rank_num_tokens, partial_tail=True, step=step ) + return prefix_sum, num_snapshots, partial if self.is_moe: - hidden_states = self.block_sparse_moe( - hidden_states, getattr(attn_metadata, "all_rank_num_tokens", None) - ) + hidden_states = self.block_sparse_moe(hidden_states, all_rank_num_tokens, step=step) else: hidden_states = self._dense_mlp(hidden_states, step) @@ -1505,6 +1695,19 @@ def _mnnvl_allreduce(self): attention = self.linear_attn if self.is_kda else self.self_attn return getattr(getattr(attention, "_o_allreduce", None), "mnnvl_allreduce", None) + def _o_proj(self) -> nn.Module: + """This layer's attention output projection (row parallel).""" + return self.linear_attn.o_proj if self.is_kda else self.self_attn.mixer.o_proj + + def _o_allreduce(self) -> AllReduce: + """The all-reduce module of this layer's attention output; it also reduces a wide step's MoE partial.""" + return (self.linear_attn if self.is_kda else self.self_attn)._o_allreduce + + def accepts_moe_partial(self, num_snapshots: int) -> bool: + """Whether this layer's pre-attention step can reduce the previous layer's unreduced MoE output: the target + built its collective state, and the snapshot bank is not empty.""" + return num_snapshots > 0 and self.decode_comm is not None + def skip_forward( self, hidden_states: torch.Tensor, @@ -1573,6 +1776,29 @@ def __init__(self, model_config: ModelConfig): # >>> route B: no per-token KDA verify states (kda_token_states): no speculative decoding # <<< route B + def _defer_moe_tail( + self, + i: int, + hidden_states: torch.Tensor, + num_snapshots: int, + spec_metadata, + capture_set, + step: Optional[DecodeStep], + ) -> bool: + """Whether layer ``i`` hands its MoE output on unreduced (the row-parallel tail, ``decode_moe.py``): its MoE + takes the step with that tail, and its consumer, the next layer's pre-attention step or the final norm after + the last layer, accepts it. A layer DSpark taps defers too: the next layer's step reduces the output, then taps + its pre-norm mixture. An unknown capture set (every layer tapped) and a tapped last layer keep the replicated + tail.""" + layer = self.layers[i] + if not (layer.is_moe and layer.block_sparse_moe.tail_rp_eligible(hidden_states, step)): + return False + if spec_metadata is not None and ( + capture_set is None or (layer.layer_idx in capture_set and i == len(self.layers) - 1) + ): + return False + return self.layers[min(i + 1, len(self.layers) - 1)].accepts_moe_partial(num_snapshots) + def forward( self, attn_metadata: AttentionMetadata, @@ -1618,6 +1844,9 @@ def forward( if spec_metadata is not None else None ) + # A MoE layer's output handed on unreduced, which the next layer's pre-attention step (or the final norm) + # reduces: a decode_comm.PendingTail, or a wide decode step's tensor. + pending_moe_partial = None for i, layer in enumerate(self.layers): # DFlash/DSpark hidden-state capture. The drafter is distilled on # the aggregated stream value -- the pre-norm softmax mixture its @@ -1636,7 +1865,10 @@ def forward( and (capture_set is None or self.layers[i - 1].layer_idx in capture_set) ): capture = (spec_metadata, self.layers[i - 1].layer_idx) - hidden_states, num_snapshots = layer( + defer_moe_tail = self._defer_moe_tail( + i, hidden_states, num_snapshots, spec_metadata, capture_set, step + ) + outputs = layer( hidden_states, block_residual, num_snapshots, @@ -1644,7 +1876,35 @@ def forward( capture=capture, step=step, prenormed=i == 0 and prenormed is not None, + pending_moe_partial=pending_moe_partial, + defer_moe_tail=defer_moe_tail, + ) + if defer_moe_tail: + hidden_states, num_snapshots, pending_moe_partial = outputs + else: + (hidden_states, num_snapshots), pending_moe_partial = outputs, None + + if isinstance(pending_moe_partial, _decode_comm.PendingTail): + normed, _ = self.layers[-1].decode_comm.sandwich_tail( + pending_moe_partial, + hidden_states, + block_residual[:num_snapshots], + self.output_attn_res_proj, + self.output_attn_res_norm, + self.norm, ) + return normed + if pending_moe_partial is not None: + _, normed = _apply_attn_res_add_and_rmsnorm( + hidden_states, + _decode_comm.wide_all_reduce(self.layers[-1]._o_allreduce(), pending_moe_partial), + block_residual[:num_snapshots], + self.output_attn_res_proj, + self.output_attn_res_norm, + self.norm, + _attn_res_max_tokens(step), + ) + return normed # The last layer has no successor, so this one recompute is # unavoidable -- output-side score weights, matching SGLang's @@ -1672,6 +1932,7 @@ def forward( self.output_attn_res_proj, self.output_attn_res_norm, self.norm, + _attn_res_max_tokens(step), ) @@ -1733,17 +1994,32 @@ def finalize_decode_weights(self) -> None: self.k3_proj_weight = fused self._qkvg_proj_weight, self._bfa_proj_weight = fused[:rows], fused[rows:] + def will_run_decode_branch( + self, attn_metadata: AttentionMetadata, step: Optional[DecodeStep] + ) -> bool: + """Whether ``forward`` runs ``step`` on the decode branch: a step ``decode_step`` classifies, outside a + breakable CUDA graph.""" + return step is not None and not is_in_breakable_cuda_graph() + def forward( self, hidden_states: torch.Tensor, attn_metadata: AttentionMetadata, step: Optional[DecodeStep] = None, + reduce_output: bool = True, + project_output: bool = True, ) -> torch.Tensor: - """The built-in forward on a step ``decode_step`` does not classify, and under a breakable CUDA graph. On the - others: the plain decode on ``ssm/k3_kda_decode_attn`` on a decode step of one token per request, else the - built-in dispatch; then ``o_proj`` on its decode GEMV site.""" - if step is None or is_in_breakable_cuda_graph(): - return super().forward(hidden_states, attn_metadata) + """The built-in forward where ``will_run_decode_branch`` does not hold. On the decode branch: the plain decode + on ``ssm/k3_kda_decode_attn`` on a decode step of one token per request, else the built-in dispatch; then + ``o_proj`` on its decode GEMV site and the TP all-reduce. + + On every path, ``reduce_output=False`` returns ``o_proj``'s TP partial (no all-reduce), and + ``project_output=False`` the post-o_norm core ``[N, H * 128]`` (no ``o_proj``).""" + if not self.will_run_decode_branch(attn_metadata, step): + if reduce_output and project_output: + return super().forward(hidden_states, attn_metadata) + core = self._builtin_core(hidden_states, attn_metadata).reshape(-1, self.proj_size) + return self.o_proj(core) if project_output else core if ( step.decode and step.tokens_per_request == 1 @@ -1753,19 +2029,37 @@ def forward( core = self._k3_decode(hidden_states[: step.num_tokens], attn_metadata) else: core = self._forward_impl(hidden_states, attn_metadata) - return self._k3_project_output(core) + if not project_output: + return core.reshape(-1, self.proj_size) + return self._k3_project_output(core, reduce_output) - def _k3_project_output(self, core: torch.Tensor) -> torch.Tensor: - """``o_proj`` on the ``o_proj`` decode GEMV site where it takes the rows (else the module), then the TP - all-reduce.""" + def _builtin_core( + self, hidden_states: torch.Tensor, attn_metadata: AttentionMetadata + ) -> torch.Tensor: + """The built-in forward's core ``[N, H, 128]``: inside a breakable CUDA graph from the eager + ``maybe_bcg_kda_core_inplace``, else from ``_forward_impl``.""" + if self.register_to_config and is_in_breakable_cuda_graph(): + core = hidden_states.new_empty( + (hidden_states.shape[0], self.num_heads, self.head_dim), dtype=torch.bfloat16 + ) + maybe_bcg_kda_core_inplace(hidden_states, self.layer_idx_str, core) + return core + return self._forward_impl(hidden_states, attn_metadata) + + def _k3_project_output(self, core: torch.Tensor, reduce_output: bool = True) -> torch.Tensor: + """``o_proj`` on the ``o_proj`` decode GEMV site where it takes the rows (else the module), then, with + ``reduce_output``, the TP all-reduce.""" + core2d = core.reshape(-1, self.proj_size) out = None if self.decode_gemvs is not None: - out = self.decode_gemvs.project( - "o_proj", core.reshape(-1, self.proj_size), self.o_proj.weight - ) + out = self.decode_gemvs.project("o_proj", core2d, self.o_proj.weight) if out is None: - return self._project_output(core) - return out if self._o_allreduce is None else self._o_allreduce(out) + if reduce_output: + return self._project_output(core) + out = self.o_proj(core2d) + if reduce_output and self._o_allreduce is not None: + out = self._o_allreduce(out) + return out def _k3_decode(self, x: torch.Tensor, attn_metadata: AttentionMetadata) -> torch.Tensor: """``ssm/k3_kda_decode_attn``: the core output ``[R, H, 128]`` of one token of each of the step's R requests; @@ -1879,6 +2173,29 @@ def _k3_layout_gaps(self) -> list: ) return [why for ok, why in checks if not ok] + def will_run_decode_branch( + self, attn_metadata: AttentionMetadata, step: Optional[DecodeStep] + ) -> bool: + """Whether ``forward`` runs ``step`` on the decode kernels (see ``_k3_step_view``). Only there does it return + the gated attention output before ``o_proj``; the KV writes forbid running the attention twice, so a caller + that needs it asks first.""" + return self._k3_step_view(attn_metadata, step) is not None + + def _k3_step_view( + self, attn_metadata: AttentionMetadata, step: Optional[DecodeStep] + ) -> Optional[dict]: + """The paged-cache view the decode kernels read on ``step``: a decode step, this layer's ``[W_a; W_g]`` and + workspace built, outside a breakable CUDA graph, and a cache ``k3_mla_decode_view`` takes. Else None.""" + if not ( + step is not None + and step.decode + and self.k3_ag_weight is not None + and self.k3_workspace is not None + and not is_in_breakable_cuda_graph() + ): + return None + return self._k3_decode_view(attn_metadata, step.num_tokens) + def forward( self, position_ids: Optional[torch.Tensor], @@ -1887,19 +2204,15 @@ def forward( all_reduce_params=None, latent_cache_gen: Optional[torch.Tensor] = None, step: Optional[DecodeStep] = None, + project_output: bool = True, ) -> torch.Tensor: - """The built-in forward, except on a decode step whose cache the decode kernels read.""" - view = None - if ( - step is not None - and step.decode - and latent_cache_gen is None - and self.k3_ag_weight is not None - and self.k3_workspace is not None - and not is_in_breakable_cuda_graph() - ): - view = self._k3_decode_view(attn_metadata, step.num_tokens) + """The built-in forward, except on a decode step whose cache the decode kernels read + (``will_run_decode_branch``). There only, ``project_output=False`` returns the gated attention output + ``[M, H * 128]``, the input of ``o_proj``.""" + view = None if latent_cache_gen is not None else self._k3_step_view(attn_metadata, step) if view is None: + if not project_output: + raise ValueError("project_output=False needs a step will_run_decode_branch takes") return super().forward( position_ids, hidden_states, attn_metadata, all_reduce_params, latent_cache_gen ) @@ -1939,6 +2252,8 @@ def forward( gate=ag, gate_col0=rows, ) + if not project_output: + return attn_output out = None if gemvs is None else gemvs.project("o_proj", attn_output, self.o_proj.weight) if out is None: out = self._project_output( @@ -2170,8 +2485,12 @@ def cache_derived_state(self) -> None: def post_load_weights(self) -> None: """The state the decode kernels share, built once per device before any CUDA-graph capture and handed to - every layer of its kind: the KDA projection's Lamport buffers and the MLA decode attention's workspace; and - the decode GEMVs' state (built by ``cache_derived_state``) handed to every attention module.""" + every layer of its kind: the KDA projection's Lamport buffers and the MLA decode attention's workspace; the + decode GEMVs' state (built by ``cache_derived_state``) handed to every attention module; and, where every + attention all-reduce runs over MNNVL, the TP group's collective state (``K3DecodeComm``, collective: every + rank builds it here) handed to every layer, with whether its sandwich takes the layer's o_proj, the decode + path's one-shot ceiling on every stock MNNVL all-reduce (``use_decode_one_shot``), and the MoE decode path + (``_build_decode_moe``).""" super().post_load_weights() kda = [layer.linear_attn for layer in self.model.layers if layer.is_kda] mla = [layer.self_attn.mixer for layer in self.model.layers if not layer.is_kda] @@ -2186,14 +2505,80 @@ def post_load_weights(self) -> None: module.k3_workspace = workspace for module in kda + mla: module.decode_gemvs = self.model.decode_gemvs + layers = self.model.layers + comm = None + if all(layer._mnnvl_allreduce() is not None for layer in layers): + for layer in layers: + layer.sandwich_oproj = _decode_comm.K3DecodeComm.takes_oproj( + layer._o_proj(), self.model.num_attn_res_snapshots + ) + oproj = next((layer._o_proj().weight for layer in layers if layer.sandwich_oproj), None) + comm = _decode_comm.K3DecodeComm.create(self.model_config.mapping, oproj) + for layer in layers: + layer.decode_comm = comm + _decode_comm.use_decode_one_shot(self) + # >>> route B: no MoE decode path yet (its build is tp16_moetp4ep4's; k3_moe's wide build does not fit 896 + # local experts). The MoE layers run the generic path until route B's own engines are wired. + moe_layers = 0 + # <<< route B logger.info( # >>> route B: no verify kernels "Kimi K3 decode kernels: KDA on k3_kda_decode_attn " # <<< route B f"({sum(m.takes_k3_kernels for m in kda)} / {len(kda)} layers take them), MLA on k3_mla_qkv and " - f"k3_mla_attn_vb_out ({len(mla)} layers)" + f"k3_mla_attn_vb_out ({len(mla)} layers); the post-attention all-reduce of at most " + f"{_decode_comm.AR_ATTN_RES_MAX_TOKENS} tokens " + + ( + "with the residual update: k3_sandwich_oproj with o_proj at most " + f"{_decode_comm.SANDWICH_MAX_TOKENS} decode tokens " + f"({sum(layer.sandwich_oproj for layer in layers)} / {len(layers)} layers), else " + "mnnvl_allreduce_attn_res" + if comm is not None + else "unfused (an attention all-reduce does not run over MNNVL)" + ) + + f"; MoE on k3_moe_front, k3_moe and the row-parallel tail ({moe_layers} layers)" ) + def _build_decode_moe(self, comm: _decode_comm.K3DecodeComm) -> int: + """The MoE decode path (``decode_moe.py``) on every MoE layer it takes: the shared state (collective: the + front's head workspace), each layer's decode weights with its latent norm folded into the latent up + projection, the decode GEMVs' state; then one call of every kernel of the path (collective: the front and + the sandwich tail). Every rank builds it here. Returns the number of MoE layers it takes.""" + mapping = self.model_config.mapping + moes = [layer.block_sparse_moe for layer in self.model.layers if layer.is_moe] + gaps = { + id(moe): _decode_moe.layout_gaps( + moe, mapping.tp_size, self.model.num_attn_res_snapshots + ) + for moe in moes + } + takes = [moe for moe in moes if not gaps[id(moe)]] + for reason in sorted({gap for moe in moes for gap in gaps[id(moe)]}): + logger.info_once( + f"Kimi K3 MoE decode path off on some layers: {reason}", + key=f"k3_decode_moe_off_{reason}", + ) + if not takes: + return 0 + backend = takes[0].routed_experts.backend + state = _decode_moe.K3DecodeMoe.create( + mapping, + backend.w3_w1_weight.device, + backend.w3_w1_weight.shape[1] // 2, + backend.expert_size_per_partition, + comm.mnnvl, + ) + for moe in takes: + _decode_moe.fold_latent_norm(moe) + moe.decode_moe = _decode_moe.K3DecodeMoeLayer.create( + moe, state, mapping.tp_rank, mapping.tp_size + ) + moe.decode_gemvs = self.model.decode_gemvs + first = takes[0].decode_moe + first.warm_up(takes[0]) + comm.compile_tail(takes[0].moe_hidden_size, first.shared_cols, first.tail_weight) + return len(takes) + def _check_step_contract(self, attn_metadata: AttentionMetadata) -> None: """First-forward checks of the engine surface and the per-engine settings.""" objects = { From ba53c843592e8c4f7cb1f95c8c1bab1c9e85e4e4 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 01:02:27 -0700 Subject: [PATCH 110/161] [None][fix] modeling_v2 Kimi K3 target: the MoE decode path hands its kernels the routing bias detached The target's post_load_weights builds the MoE decode path and runs each of its kernels once, outside inference mode with autograd on. The front was given the gate's e_score_correction_bias, an nn.Parameter that requires grad, so the op returned outputs that require grad. k3_moe's first-call compile exports its arguments through DLPack, which refuses such tensors, and every rank failed to start: "Can't export tensors that require gradient". K3DecodeMoeLayer keeps a detached view of the bias and hands that to k3_moe_front and k3_route_quant. The tp16_moetp16ep1 copy of decode_moe.py takes the same change. The 4-rank decode collectives test gains a first check that builds the decode MoE as post_load_weights does: from a MoE layer's own modules (the target's gate, the latent projections, the stock RMSNorm and the shared GatedMLP) with zeroed routed experts, outside inference mode with autograd on, then a 3-token step and a wide 16-token step against the shared expert's forward. Before this change it fails with the same error. Signed-off-by: Vasanth Sabavat --- .../decode_moe.py | 17 ++- .../decode_moe.py | 17 ++- .../comm/_kimi_k3_decode_comm_op_matrix.py | 140 +++++++++++++++++- 3 files changed, 159 insertions(+), 15 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py index 75f5b06ed66e..951c928e5e0e 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py @@ -178,14 +178,17 @@ def fold_latent_norm(moe: nn.Module) -> None: class K3DecodeMoeLayer: """One MoE layer's decode path: the front weight (this rank's head slice padded to whole tiles, then the shared gate_up re-ordered), the head slice (a view of it), the row-parallel tail weight - ``[latent up columns lo:lo+width | padding | shared down]``, and its ``k3_moe`` handles on the shared builds. - Built by `create` once the weights are final (and the latent norm folded).""" + ``[latent up columns lo:lo+width | padding | shared down]``, the routing bias, and its ``k3_moe`` handles on the + shared builds. Built by `create` once the weights are final (and the latent norm folded).""" state: K3DecodeMoe front_weight: torch.Tensor head_weight: torch.Tensor tail_weight: torch.Tensor tail_pad: Optional[torch.Tensor] + # The gate's routing bias detached (a view of the parameter): an op given a tensor that requires grad while + # autograd records, as in post_load_weights, returns outputs that require grad, which k3_moe cannot take. + bias: torch.Tensor lo: int width: int shared_cols: int @@ -224,6 +227,7 @@ def create( head_weight=front[: head.shape[0]], tail_weight=tail, tail_pad=up.new_zeros(WIDE_MAX_TOKENS, pad) if pad else None, + bias=moe.gate.e_score_correction_bias.detach(), lo=lo, width=width, shared_cols=gate_up.shape[0] // 2, @@ -242,9 +246,8 @@ def warm_up(self, moe: nn.Module) -> None: logits = torch.zeros(1, moe.num_experts, dtype=torch.float32, device=device) latent = torch.zeros(1, moe.moe_hidden_size, dtype=torch.bfloat16, device=device) ids, weights, x_fp8, x_sf = k3_route_quant( - logits, moe.gate.e_score_correction_bias, latent, float(moe.gate.routed_scaling_factor), - early_trigger=True, - ) # fmt: skip + logits, self.bias, latent, float(moe.gate.routed_scaling_factor), early_trigger=True + ) k3_moe(x_fp8, x_sf, ids, weights, offset, self.wide) torch.cuda.synchronize(device) @@ -252,7 +255,7 @@ def _front(self, moe: nn.Module, x: torch.Tensor): return k3_moe_front( x.contiguous(), self.front_weight, - moe.gate.e_score_correction_bias, + self.bias, float(moe.gate.routed_scaling_factor), self.shared_cols, *moe._situ_betas, @@ -307,7 +310,7 @@ def _wide(self, moe: nn.Module, hidden_states: torch.Tensor, gemvs) -> torch.Ten def _routed_partial(): ids, weights, x_fp8, x_sf = k3_route_quant( - router_logits, moe.gate.e_score_correction_bias, routed_in.contiguous(), + router_logits, self.bias, routed_in.contiguous(), float(moe.gate.routed_scaling_factor), early_trigger=True, ) # fmt: skip return k3_moe( diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_moe.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_moe.py index 75f5b06ed66e..951c928e5e0e 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_moe.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_moe.py @@ -178,14 +178,17 @@ def fold_latent_norm(moe: nn.Module) -> None: class K3DecodeMoeLayer: """One MoE layer's decode path: the front weight (this rank's head slice padded to whole tiles, then the shared gate_up re-ordered), the head slice (a view of it), the row-parallel tail weight - ``[latent up columns lo:lo+width | padding | shared down]``, and its ``k3_moe`` handles on the shared builds. - Built by `create` once the weights are final (and the latent norm folded).""" + ``[latent up columns lo:lo+width | padding | shared down]``, the routing bias, and its ``k3_moe`` handles on the + shared builds. Built by `create` once the weights are final (and the latent norm folded).""" state: K3DecodeMoe front_weight: torch.Tensor head_weight: torch.Tensor tail_weight: torch.Tensor tail_pad: Optional[torch.Tensor] + # The gate's routing bias detached (a view of the parameter): an op given a tensor that requires grad while + # autograd records, as in post_load_weights, returns outputs that require grad, which k3_moe cannot take. + bias: torch.Tensor lo: int width: int shared_cols: int @@ -224,6 +227,7 @@ def create( head_weight=front[: head.shape[0]], tail_weight=tail, tail_pad=up.new_zeros(WIDE_MAX_TOKENS, pad) if pad else None, + bias=moe.gate.e_score_correction_bias.detach(), lo=lo, width=width, shared_cols=gate_up.shape[0] // 2, @@ -242,9 +246,8 @@ def warm_up(self, moe: nn.Module) -> None: logits = torch.zeros(1, moe.num_experts, dtype=torch.float32, device=device) latent = torch.zeros(1, moe.moe_hidden_size, dtype=torch.bfloat16, device=device) ids, weights, x_fp8, x_sf = k3_route_quant( - logits, moe.gate.e_score_correction_bias, latent, float(moe.gate.routed_scaling_factor), - early_trigger=True, - ) # fmt: skip + logits, self.bias, latent, float(moe.gate.routed_scaling_factor), early_trigger=True + ) k3_moe(x_fp8, x_sf, ids, weights, offset, self.wide) torch.cuda.synchronize(device) @@ -252,7 +255,7 @@ def _front(self, moe: nn.Module, x: torch.Tensor): return k3_moe_front( x.contiguous(), self.front_weight, - moe.gate.e_score_correction_bias, + self.bias, float(moe.gate.routed_scaling_factor), self.shared_cols, *moe._situ_betas, @@ -307,7 +310,7 @@ def _wide(self, moe: nn.Module, hidden_states: torch.Tensor, gemvs) -> torch.Ten def _routed_partial(): ids, weights, x_fp8, x_sf = k3_route_quant( - router_logits, moe.gate.e_score_correction_bias, routed_in.contiguous(), + router_logits, self.bias, routed_in.contiguous(), float(moe.gate.routed_scaling_factor), early_trigger=True, ) # fmt: skip return k3_moe( diff --git a/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py index be6448fa7f28..786445cadfdc 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py @@ -20,6 +20,10 @@ [7168, 256 + 384]), or returns that tail reduced in torch and the stock all-reduce. Checks: + * decode MoE (first, outside inference mode with autograd on, as post_load_weights runs): the MoE decode path built + from a MoE layer's own modules (the target's gate, whose parameters require grad, the latent projections, the + stock RMSNorm and shared GatedMLP; zeroed routed experts) at this TP, every kernel warmed up, then a 3-token step + and a wide 16-token step against the shared expert's forward. * takes_oproj: the TP16 o_proj shape only, bias-free, within the kernel's candidate count. * use_decode_one_shot: every stock MNNVL all-reduce takes the decode ceiling. * fused vs built-in: for each dense layer and step (decode of 1 / 3 / 8 tokens, DSpark 1 x 6, a 5-token prefill, @@ -572,6 +576,138 @@ def load(): del graph +# The MoE decode path at this check's TP: a moetp4ep4 rank's routed experts (intermediate columns, experts), and a +# shared expert whose rank slice at TP 4 is a TP16 rank's (384 activation columns). +I_TP, E_LOCAL, NUM_EXPERTS = 768, 224, 896 +SHARED = 384 * 4 +SITU_CAPS = (4.0, 25.0) +MOE_TOL = 3e-2 # the front's fused shared gate_up + SiTU against the shared expert's own GEMMs + + +def _decode_moe_layer(): + """A MoE layer with the parts the decode path reads, from the modules KimiK3MoERuntime builds: the target's + KimiK3MoEGate, nn.Linear latent projections, the stock RMSNorm and GatedMLP (at this check's TP), each with its + own parameters (the gate's and the latent projections' require grad, as in the model), and zeroed routed-expert + buffers in TRTLLM-Gen's W4A8_MXFP4_MXFP8 layout, so the routed experts add nothing.""" + from tensorrt_llm._torch.distributed import AllReduce + from tensorrt_llm._torch.model_config import ModelConfig + from tensorrt_llm._torch.modules.gated_mlp import GatedMLP + from tensorrt_llm._torch.modules.rms_norm import RMSNorm + from tensorrt_llm._torch.modules.situ import SituAndMul + from tensorrt_llm.functional import AllReduceStrategy + + cfg = SimpleNamespace( + num_experts_per_token=16, + num_experts=NUM_EXPERTS, + routed_scaling_factor=2.827, + moe_router_activation_func="sigmoid", + num_expert_group=1, + topk_group=1, + moe_renormalize=True, + hidden_size=H, + ) + with torch.device("cuda"): + gate = T.KimiK3MoEGate(cfg, logits_gemm_dtype=torch.bfloat16) + down = nn.Linear(H, LATENT, bias=False, dtype=torch.bfloat16) + up = nn.Linear(LATENT, H, bias=False, dtype=torch.bfloat16) + norm = RMSNorm(hidden_size=LATENT, eps=1e-5, dtype=torch.bfloat16) + shared = GatedMLP( + hidden_size=H, + intermediate_size=SHARED, + bias=False, + activation=SituAndMul( + beta=SITU_CAPS[0], linear_beta=SITU_CAPS[1], use_fused_activation=True + ), + dtype=torch.bfloat16, + config=ModelConfig(mapping=R.mapping, allreduce_strategy=AllReduceStrategy.MNNVL), + reduce_output=True, + layer_idx=MOE_LAYERS[0], + is_shared_expert=True, + ) + g = _gen(5000) # replicated + gr = _gen(5100 + R.rank) # this rank's shared expert slices + with torch.no_grad(): + gate.weight.copy_(_normal(g, gate.weight.shape, 0.02)) + gate.e_score_correction_bias.copy_( + 0.05 * torch.randn(NUM_EXPERTS, generator=g, device="cuda") + ) + down.weight.copy_(_normal(g, down.weight.shape, 0.02)) + up.weight.copy_(_normal(g, up.weight.shape, 0.02)) + norm.weight.copy_(_normal(g, norm.weight.shape, 0.1, 1.0)) + shared.gate_up_proj.weight.copy_(_normal(gr, shared.gate_up_proj.weight.shape, 0.02)) + shared.down_proj.weight.copy_(_normal(gr, shared.down_proj.weight.shape, 0.02)) + + def zeros(*shape): + return nn.Parameter( + torch.zeros(shape, dtype=torch.uint8, device="cuda"), requires_grad=False + ) + + backend = SimpleNamespace( + w3_w1_weight=zeros(E_LOCAL, 2 * I_TP, LATENT // 2), + w3_w1_weight_scale=zeros(E_LOCAL, 2 * I_TP, LATENT // 32), + w2_weight=zeros(E_LOCAL, LATENT, I_TP // 2), + w2_weight_scale=zeros(E_LOCAL, LATENT, I_TP // 32), + expert_size_per_partition=E_LOCAL, + slot_start=R.rank * E_LOCAL % NUM_EXPERTS, + ) + all_reduce = AllReduce( + mapping=R.mapping, strategy=AllReduceStrategy.MNNVL, dtype=torch.bfloat16 + ) + return SimpleNamespace( + num_experts=NUM_EXPERTS, + top_k=cfg.num_experts_per_token, + moe_hidden_size=LATENT, + hidden_size=H, + gate=gate, + routed_expert_down_proj=down, + routed_expert_up_proj=up, + routed_expert_norm=norm, + shared_experts=shared, + routed_experts=SimpleNamespace(backend=backend, all_reduce=all_reduce), + _situ_betas=SITU_CAPS, + moe_main_event=torch.cuda.Event(), + moe_shared_event=torch.cuda.Event(), + shared_expert_stream=torch.cuda.Stream(), + ) + + +def check_decode_moe_from_model_parameters(): + """The MoE decode path built as post_load_weights builds it, outside inference mode with autograd on, from a MoE + layer's own modules (_decode_moe_layer): the shared state, the folded latent norm, the layer's decode weights and + one call of every kernel (the front, both k3_moe builds, k3_route_quant). Then a step of 3 tokens (the replicated + tail) and a wide step of 16 (this rank's share, reduced here) against the shared expert's forward: the routed + experts are zeros, so the two agree.""" + if R.world not in (4, 8, 16): + if R.rank == 0: + print( + f"[rank 0] decode MoE skipped at world {R.world}: the front runs 4, 8 or 16 ranks", + flush=True, + ) + return + assert torch.is_grad_enabled() and not torch.is_inference_mode_enabled() + dm = T._decode_moe + moe = _decode_moe_layer() + assert moe.gate.e_score_correction_bias.requires_grad + mnnvl = T._decode_comm.MnnvlWorkspace.create(R.mapping, T._decode_comm.MNNVL_BUFFER_BYTES) + device = torch.device("cuda", torch.cuda.current_device()) + state = dm.K3DecodeMoe.create(R.mapping, device, I_TP, E_LOCAL, mnnvl) + dm.fold_latent_norm(moe) + layer = dm.K3DecodeMoeLayer.create(moe, state, R.rank, R.world) + layer.warm_up(moe) + for rows in (3, 16): + x = _normal(_gen(5200 + rows), (rows, H), 0.5) + y = layer.forward(moe, x, None, partial_tail=False) + if rows > dm.MAX_TOKENS: # a wide step returns this rank's share + y = moe.routed_experts.all_reduce(y) + want = moe.shared_experts(x) + torch.cuda.synchronize() + err = ls.rel_err(y, want) + if R.rank == 0: + print(f"[rank 0] decode MoE, {rows} tokens: vs the shared expert {err:.2e}", flush=True) + assert torch.isfinite(y.float()).all() and err <= MOE_TOL, (rows, err) + assert R.same_on_ranks(y), (rows, "ranks differ") + + CHECKS = [ check_takes_oproj, check_decode_one_shot, @@ -595,10 +731,12 @@ def _run_one_rank(args) -> int: "tensorrt_llm._torch._experimental.modeling_v2.models.kimi_k3_vl." "kimi_k3_mxfp4__sm_100__tp16_moetp4ep4.modeling" ) + # First, before anything exists as an inference tensor: the decode MoE as post_load_weights builds it. + code = ls.run_checks(R, [check_decode_moe_from_model_parameters]) with torch.inference_mode(): LAYERS_BUILT = build_layers() COMM = make_comm(LAYERS_BUILT) - code = ls.run_checks(R, CHECKS) + code = ls.run_checks(R, CHECKS) or code if R.rank == 0: print(f"[rank 0] world {R.world}; {STATS}", flush=True) return code From 7dc60ccc4bb08dbbde8e305f30c61d435d9b1ebb Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 23:00:39 -0700 Subject: [PATCH 111/161] [None][feat] SpecDecOneEngineForCausalLM: build the drafter through an overridable method An engine change outside modeling_v2: the one-engine shell now builds its draft model with SpecDecOneEngineForCausalLM._build_draft_model(), whose default is the mode registry's builder (get_draft_model) with the same arguments as before, so behaviour is unchanged. A model that owns its drafter overrides the method instead of copying the shell's speculative setup (draft config, separate draft KV cache, epilogue, logits processor, worker). The new hw-agnostic test checks both: the default calls get_draft_model with (model_config, draft_config, lm_head, model), and __init__ builds the drafter through a subclass override, once, with the resolved draft config. Signed-off-by: Vasanth Sabavat --- .../_torch/models/modeling_speculative.py | 17 ++++- .../test_build_draft_model_hook.py | 73 +++++++++++++++++++ 2 files changed, 87 insertions(+), 3 deletions(-) create mode 100644 tests/unittest/_torch/speculative/hw_agnostic/test_build_draft_model_hook.py diff --git a/tensorrt_llm/_torch/models/modeling_speculative.py b/tensorrt_llm/_torch/models/modeling_speculative.py index ba54c65d95d1..3e6aff1bf21b 100644 --- a/tensorrt_llm/_torch/models/modeling_speculative.py +++ b/tensorrt_llm/_torch/models/modeling_speculative.py @@ -1783,9 +1783,8 @@ def __init__(self, self.use_separate_draft_kv_cache = should_use_separate_draft_kv_cache( spec_config) - self.draft_model = get_draft_model(model_config, - self.draft_config, - self.lm_head, self.model) + self.draft_model = self._build_draft_model( + model_config, self.draft_config) if spec_config.uses_replacement_heads: self.draft_config = self.draft_model.model_config if self.draft_model is not None: @@ -1808,6 +1807,18 @@ def __init__(self, self.epilogue.append(self.spec_worker) self.layer_idx = -1 + def _build_draft_model( + self, model_config: ModelConfig, + draft_config: Optional[ModelConfig]) -> Optional[nn.Module]: + """Build the draft model for ``model_config``'s speculative mode. + + The default is the mode registry's builder (``get_draft_model``). A + model that owns its drafter overrides this; ``__init__`` calls it once, + after ``draft_config`` is resolved and before the worker is built. + """ + return get_draft_model(model_config, draft_config, self.lm_head, + self.model) + def setup_aliases(self) -> None: if (self.draft_model is not None and getattr(self.draft_model, "shares_target_kv_cache", False)): diff --git a/tests/unittest/_torch/speculative/hw_agnostic/test_build_draft_model_hook.py b/tests/unittest/_torch/speculative/hw_agnostic/test_build_draft_model_hook.py new file mode 100644 index 000000000000..1d5ef8c4a97a --- /dev/null +++ b/tests/unittest/_torch/speculative/hw_agnostic/test_build_draft_model_hook.py @@ -0,0 +1,73 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""``SpecDecOneEngineForCausalLM._build_draft_model``: the default is the mode registry's builder, and a subclass +override is what ``__init__`` builds the drafter with.""" + +import pytest +from torch import nn +from transformers import PretrainedConfig + +from tensorrt_llm._torch.model_config import ModelConfig +from tensorrt_llm._torch.models import modeling_speculative +from tensorrt_llm._torch.models.modeling_speculative import SpecDecOneEngineForCausalLM +from tensorrt_llm._torch.models.modeling_utils import DecoderModelForCausalLM +from tensorrt_llm.llmapi.llm_args import MTPDecodingConfig + +pytestmark = pytest.mark.cpu_only + + +def test_default_build_draft_model_is_get_draft_model(monkeypatch): + built = nn.Module() + calls = [] + + def fake_get_draft_model(model_config, draft_config, lm_head, model): + calls.append((model_config, draft_config, lm_head, model)) + return built + + monkeypatch.setattr(modeling_speculative, "get_draft_model", fake_get_draft_model) + shell = object.__new__(SpecDecOneEngineForCausalLM) + nn.Module.__init__(shell) + shell.lm_head = nn.Linear(2, 2) + shell.model = nn.Module() + model_config, draft_config = object(), object() + + assert shell._build_draft_model(model_config, draft_config) is built + assert calls == [(model_config, draft_config, shell.lm_head, shell.model)] + + +def _minimal_decoder_init(self, model, *, config, hidden_size, vocab_size): + """``DecoderModelForCausalLM.__init__`` reduced to what the one-engine shell reads.""" + nn.Module.__init__(self) + self.model_config = config + self.model = model + self.lm_head = nn.Linear(hidden_size, vocab_size, bias=False) + self.logits_processor = object() + self.epilogue = [] + + +def test_init_builds_the_drafter_through_an_override(monkeypatch): + monkeypatch.setattr(DecoderModelForCausalLM, "__init__", _minimal_decoder_init) + monkeypatch.setattr(DecoderModelForCausalLM, "__post_init__", lambda self: None) + monkeypatch.setattr(modeling_speculative, "get_spec_worker", lambda *args, **kwargs: None) + + def registry_builder_not_called(*args, **kwargs): + raise AssertionError("the override replaces get_draft_model") + + monkeypatch.setattr(modeling_speculative, "get_draft_model", registry_builder_not_called) + drafter = nn.Linear(1, 1) + calls = [] + + class OwnDrafter(SpecDecOneEngineForCausalLM): + def _build_draft_model(self, model_config, draft_config): + calls.append((model_config, draft_config)) + return drafter + + model_config = ModelConfig( + pretrained_config=PretrainedConfig(hidden_size=8, vocab_size=16, num_hidden_layers=2), + spec_config=MTPDecodingConfig(max_draft_len=1), + ) + shell = OwnDrafter(nn.Module(), model_config) + + assert calls == [(model_config, shell.draft_config)] + assert shell.draft_model is drafter + assert shell.epilogue == [drafter] From 31d76e08aae85cf34e0b117c4795c5e620d33f09 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 23:10:41 -0700 Subject: [PATCH 112/161] [None][feat] modeling_v2 Kimi K3 target: decode GEMV sites for the DSpark drafter Four sites at the drafter's TP16 per-rank shapes, each a cell the G2 contracts already certify at 1..8 rows: drafter_qkv [512, 7168] and drafter_gate_up [1792, 7168] on k3_ctm_gemv_long (split 8, ring 6, push), drafter_o [7168, 384] on k3_ctm_gemv (split 1, push), and drafter_down [7168, 896] on k3_ctm_gemv_swiglu (split 2, push), whose rows are the [gate | up] input of the SiLU-and-mul. create() warms them with the other sites; the site test covers them, the SiLU-and-mul site against the float64 product of torch's bf16 silu(gate) * up. Signed-off-by: Vasanth Sabavat --- .../decode_gemv.py | 60 ++++++++++++++----- .../test_modeling_v2_kimi_k3_decode_gemv.py | 24 ++++++-- 2 files changed, 64 insertions(+), 20 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py index 0e26e3971d0e..33fbbac23ca7 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py @@ -4,11 +4,12 @@ * **Per-site GEMVs** (`K3DecodeGemvs.project`): a projection of a decode step runs on the kernel measured fastest at its call site's weight shape (`SITES`): at most `MAX_ROWS` rows on `gemm/k3_decode_gemv`, - `gemm/k3_ctm_gemv_wide` or `gemm/k3_ctm_gemv_long`, and, where the site lists it, more rows (up to `WIDE_ROWS`) - on `gemm/k3_ctm_gemv_wide`. The sites are the decode kernels' fused projections (MLA's [W_a; W_g] with the gate - rows through a sigmoid, KDA's [q | k | v | g | f_a | b]), the attention output projection, the built-in MLA path's - q_a / kv_a, q_b and gate projections, and the MoE decode path's projections (`decode_moe.py`), two of them with - fp32 outputs. + `gemm/k3_ctm_gemv_wide`, `gemm/k3_ctm_gemv_long`, `gemm/k3_ctm_gemv` or `gemm/k3_ctm_gemv_swiglu`, and, where the + site lists it, more rows (up to `WIDE_ROWS`) on `gemm/k3_ctm_gemv_wide`. The sites are the decode kernels' fused + projections (MLA's [W_a; W_g] with the gate rows through a sigmoid, KDA's [q | k | v | g | f_a | b]), the attention + output projection, the built-in MLA path's q_a / kv_a, q_b and gate projections, the MoE decode path's projections + (`decode_moe.py`), two of them with fp32 outputs, and the DSpark drafter's q / k / v, output, gate / up and down + projections (`K3DSparkDrafter`). * **LM head** (`K3LogitsProcessor`): at most `MAX_ROWS` rows of this rank's vocabulary shard on `gemm/k3_head_gemv` over the target's `K3HeadGemvWorkspace`, then the shards gathered (`comm/allgather`) as the stock head gathers them. It is the shell's logits processor, so the speculative worker's target logits and the @@ -41,9 +42,13 @@ from tensorrt_llm._torch._experimental.modeling_v2.catalog.activation.k3_situ_mul import k3_situ_mul from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.allgather import allgather +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv import k3_ctm_gemv from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv_long import ( k3_ctm_gemv_long, ) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv_swiglu import ( + k3_ctm_gemv_swiglu, +) from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv_wide import ( k3_ctm_gemv_wide, ) @@ -73,10 +78,11 @@ @dataclass(frozen=True) class Site: """A call site's weight shape (this target's per-rank shapes) and its kernels: ``small`` at 1..MAX_ROWS rows - ("decode", "wide" or "long"), and k3_ctm_gemv_wide at MAX_ROWS+1..``wide_rows`` rows where ``wide``. Output columns - from ``sig_col0`` on are stored through a sigmoid; ``out_fp32``: an fp32 output (k3_ctm_gemv_wide only). - ``split`` / ``ring`` / ``push``: k3_ctm_gemv_long's CTAs per 128-row weight tile, weight-ring stages, and whether - the partial sums are pushed to each row's owner.""" + ("decode", "wide", "long", "ctm" for k3_ctm_gemv, or "swiglu" for k3_ctm_gemv_swiglu, whose rows are the + ``[gate | up]`` input of the SiLU-and-mul, 2 ``k`` wide), and k3_ctm_gemv_wide at MAX_ROWS+1..``wide_rows`` rows + where ``wide``. Output columns from ``sig_col0`` on are stored through a sigmoid; ``out_fp32``: an fp32 output + (k3_ctm_gemv_wide only). ``split`` / ``ring`` / ``push``: the CTM kernels' CTAs per 128-row weight tile, + k3_ctm_gemv_long's weight-ring stages, and whether the partial sums are pushed to each row's owner.""" n: int k: int @@ -113,9 +119,21 @@ class Site: "moe_shared_gate_up": Site(768, 7168, "wide", wide=True), "moe_tail": Site(7168, 640, "wide", wide=True, wide_rows=32), "moe_up": Site(7168, 3584, "wide", out_fp32=True), + # The DSpark drafter's projections (K3DSparkDrafter; 6 query heads, 1 KV head of 64, MLP width 896): the fused + # q / k / v, the output projection (row parallel), gate_up [gate 896 | up 896], and the down projection with the + # SiLU-and-mul (row parallel). + "drafter_qkv": Site(512, 7168, "long", split=8, ring=6, push=True), + "drafter_o": Site(7168, 384, "ctm", split=1, push=True), + "drafter_gate_up": Site(1792, 7168, "long", split=8, ring=6, push=True), + "drafter_down": Site(7168, 896, "swiglu", split=2, push=True), } +def _width(spec: Site) -> int: + """The row width a site's kernel reads: ``k``, or the ``[gate | up]`` input of a SiLU-and-mul site, 2 ``k``.""" + return 2 * spec.k if spec.small == "swiglu" else spec.k + + def _wide_tile(rows: int) -> int: from tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv import k3_ctm_gemv_kernel @@ -136,6 +154,14 @@ def _run( return k3_ctm_gemv_wide(x2d, weight, sig_col0=spec.sig_col0, out_fp32=spec.out_fp32) if spec.out_fp32: return None + if kernel == "ctm": + if spec.sig_col0 >= 0 or not _ctm_op.supports(x2d, weight, spec.split): + return None + return k3_ctm_gemv(x2d, weight, trigger_early=True, split=spec.split, push=spec.push) + if kernel == "swiglu": + if spec.sig_col0 >= 0 or not _ctm_op.supports_swiglu(x2d, weight, spec.split): + return None + return k3_ctm_gemv_swiglu(x2d, weight, trigger_early=True, split=spec.split, push=spec.push) # One wave of the GPU's SMs: beyond it the long GEMV loses to the others. sms = torch.cuda.get_device_properties(x2d.device).multi_processor_count if math.ceil(spec.n / 128) * spec.split > sms or not _ctm_op.supports_long( @@ -240,7 +266,7 @@ def create( for rows in (1, 16, 32, 64) if spec.wide else (1,): if rows > spec.wide_rows: continue - state._project(site, weight.new_zeros(rows, spec.k), weight, warm=True) + state._project(site, weight.new_zeros(rows, _width(spec)), weight, warm=True) del weight if "dense_gate_up" in sites: gu = torch.zeros(1, SITES["dense_gate_up"].n, dtype=torch.bfloat16, device=device) @@ -262,25 +288,27 @@ def create( def project(self, site: str, x: torch.Tensor, weight: torch.Tensor) -> Optional[torch.Tensor]: """``x @ weight.T`` (``[..., N]``: bf16, the site's sigmoid columns through the sigmoid; fp32 at an - ``out_fp32`` site) for ``site``'s weight on its decode kernel, or None where none takes the call: more rows - than the site's kernels take, another shape or dtype, or, under capture, a kernel that has not run eagerly. - The caller then runs its GEMM.""" + ``out_fp32`` site; on a SiLU-and-mul site, ``silu(gate) * up`` of ``x = [gate | up]`` first) for ``site``'s + weight on its decode kernel, or None where none takes the call: more rows than the site's kernels take, + another shape or dtype, or, under capture, a kernel that has not run eagerly. The caller then runs its + module.""" return self._project(site, x, weight, warm=False) def _project( self, site: str, x: torch.Tensor, weight: torch.Tensor, warm: bool ) -> Optional[torch.Tensor]: spec = SITES[site] + width = _width(spec) if ( weight.dtype != torch.bfloat16 or tuple(weight.shape) != (spec.n, spec.k) or not weight.is_contiguous() or x.dtype != torch.bfloat16 or x.dim() < 1 - or x.shape[-1] != spec.k + or x.shape[-1] != width ): return None - rows = x.numel() // spec.k + rows = x.numel() // width if 0 < rows <= MAX_ROWS: kernel = spec.small elif spec.wide and MAX_ROWS < rows <= spec.wide_rows: @@ -291,7 +319,7 @@ def _project( capturing = _capturing() if capturing and not warm and key not in self._ran: return None - y = _run(spec, kernel, _dense_rows(x.reshape(rows, spec.k)), weight) + y = _run(spec, kernel, _dense_rows(x.reshape(rows, width)), weight) if y is None: return None if not capturing: diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_gemv.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_gemv.py index 135166cde8b2..195562ee38a0 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_gemv.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_gemv.py @@ -5,8 +5,8 @@ * Each GEMV site at every row count it takes (1..8, and 9..its ``wide_rows`` where k3_ctm_gemv_wide takes it): the bits of its catalog entry's call (fp32 at an ``out_fp32`` site), within 8e-3 of ``max |ref|`` of a float64 product - (the sigmoid columns within 1e-2 of the sigmoid of it); declined above its rows, at another shape or dtype, and - under capture before an eager call. + (the sigmoid columns within 1e-2 of the sigmoid of it; a SiLU-and-mul site's product taken of torch's bf16 + ``silu(gate) * up``); declined above its rows, at another shape or dtype, and under capture before an eager call. * The LM head through ``K3LogitsProcessor`` on a real ``LMHead`` (one rank): fp32 logits from ``k3_head_gemv`` within the same bound, the stock processor's rows selected, the stock path above 8 rows and before the state is built. @@ -23,9 +23,13 @@ from torch import nn from tensorrt_llm._torch._experimental.modeling_v2.catalog.activation.k3_situ_mul import k3_situ_mul +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv import k3_ctm_gemv from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv_long import ( k3_ctm_gemv_long, ) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv_swiglu import ( + k3_ctm_gemv_swiglu, +) from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv_wide import ( k3_ctm_gemv_wide, ) @@ -85,6 +89,10 @@ def _entry(spec, rows, x, w): return k3_ctm_gemv_wide(x, w, sig_col0=spec.sig_col0, out_fp32=spec.out_fp32) if spec.small == "decode": return k3_decode_gemv(x, w) + if spec.small == "ctm": + return k3_ctm_gemv(x, w, split=spec.split, push=spec.push) + if spec.small == "swiglu": + return k3_ctm_gemv_swiglu(x, w, split=spec.split, push=spec.push) return k3_ctm_gemv_long( x, w, sig_col0=spec.sig_col0, split=spec.split, ring=spec.ring, push=spec.push ) @@ -96,6 +104,14 @@ def _site_rows(spec): ) +def _product_input(spec, x): + """The rows the site's weight multiplies: ``x``, or a SiLU-and-mul site's torch bf16 ``silu(gate) * up``.""" + if spec.small != "swiglu": + return x + gate, up = x.float().chunk(2, dim=-1) + return (torch.nn.functional.silu(gate) * up).to(torch.bfloat16) + + @pytest.mark.parametrize("site", list(decode_gemv.SITES)) def test_site(site): """Every row count the site takes runs its catalog entry: its bits, within TOL of the float64 product.""" @@ -103,12 +119,12 @@ def test_site(site): w = _weight(spec.n, spec.k, seed=len(site)) gemvs = decode_gemv.K3DecodeGemvs.create(None, sites=[site]) for m in _site_rows(spec): - x = _rows(m, spec.k, seed=m) + x = _rows(m, decode_gemv._width(spec), seed=m) y = gemvs.project(site, x, w) assert y is not None and y.shape == (m, spec.n), (site, m) assert y.dtype == (torch.float32 if spec.out_fp32 else torch.bfloat16), (site, y.dtype) assert torch.equal(_bits(y), _bits(_entry(spec, m, x, w))), (site, m) - _check_product(y, x, w, spec.sig_col0) + _check_product(y, _product_input(spec, x), w, spec.sig_col0) @pytest.mark.parametrize("site", ["kv_a", "mla_ag", "moe_head"]) From a86719999a30ffe6524bb8af0d78fc875c6b2381 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 23:11:23 -0700 Subject: [PATCH 113/161] [None][feat] modeling_v2 Kimi K3 target: its DSpark drafter on the drafter entries The target builds DSpark's standalone GQA drafter as K3DSparkDrafter through the shell's _build_draft_model; the worker stays the stock DSparkWorker. K3DSparkDrafter is the stock GQADSparkForCausalLM with dflash_forward taking a decode block on the catalog's drafter entries: per layer, the q / k / v projection on the drafter_qkv GEMV site, attention/k3_drafter_attn_qknorm (the q / k RMSNorm, NeoX RoPE and the block's attention over its paged context and its own k / v in one launch, without storing the block's k / v), the output projection on drafter_o and gate / up and down (SiLU-and-mul fused) on drafter_gate_up / drafter_down, each row-parallel output then reduced by its module's all-reduce. The norms and residual adds stay the stock modules'. Every other block runs the stock dflash_forward: a split the attention entry does not certify (DRAFTER_ATTN_SPLITS mirrors its contract), another backend, head layout, RoPE base, epsilon or MLP, a cache layout the kernel does not read, or an attention compile key that has not run eagerly when under CUDA-graph capture. The stock drafter classes and the builder's checks are declared in UNCERTIFIED_GENERIC_CALLS, the new ops in REQUIRED_TRTLLM_OPS. cache_derived_state hands the decode GEMVs' state to the drafter. Signed-off-by: Vasanth Sabavat --- .../modeling.py | 302 +++++++++++++++++- 1 file changed, 300 insertions(+), 2 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index c259ab07f728..e4a5c3396d77 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -69,7 +69,10 @@ weights are a predicted non-load, `weights.py`), and a step carrying multimodal input raises. **Speculative decoding** goes through the stock one-engine shell: DSpark or DFlash with an external drafter -checkpoint, and SA. The worker and its kernels stay upstream code; this target does not own a worker. +checkpoint, and SA. The DSpark drafter is this target's `K3DSparkDrafter`, the stock GQA drafter with a decode step's +block on the drafter entries (`attention/k3_drafter_attn_qknorm` and the drafter's decode GEMV sites); the shell +builds it through `_build_draft_model`. The worker and its kernels stay upstream code; this target does not own a +worker. """ from __future__ import annotations @@ -83,6 +86,9 @@ import torch from torch import nn +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.k3_drafter_attn_qknorm import ( + k3_drafter_attn_qknorm, +) from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.k3_mla_attn_vb_out import ( k3_mla_attn_vb_out, ) @@ -103,6 +109,11 @@ from tensorrt_llm._torch.custom_ops import cute_dsl_kimi_k3_kda_mtp_ops # noqa: F401 from tensorrt_llm._torch.distributed import AllReduce from tensorrt_llm._torch.model_config import ModelConfig +from tensorrt_llm._torch.models.modeling_dflash import DFlashForCausalLM +from tensorrt_llm._torch.models.modeling_dspark import ( + GQADSparkForCausalLM, + draft_is_embedded_in_target, +) from tensorrt_llm._torch.models.modeling_kimi_linear import KimiLinearForCausalLM from tensorrt_llm._torch.models.modeling_speculative import SpecDecOneEngineForCausalLM from tensorrt_llm._torch.models.modeling_utils import DecoderModel, register_auto_model @@ -122,6 +133,7 @@ from tensorrt_llm._torch.moe.fused_moe.interface import MoESchedulerKind from tensorrt_llm._torch.moe.fused_moe.routing import DeepSeekV3MoeRoutingMethod from tensorrt_llm._torch.pyexecutor.breakable_cuda_graph import is_in_breakable_cuda_graph +from tensorrt_llm._torch.pyexecutor.config_utils import is_mla from tensorrt_llm._torch.pyexecutor.kv_cache.mamba_cache_manager import MambaHybridCacheManagerV2 from tensorrt_llm._torch.utils import AuxStreamType from tensorrt_llm.functional import AllReduceStrategy @@ -176,6 +188,10 @@ "k3_moe", "k3_route_quant", "mnnvl_allgather_split", + # The DSpark drafter's decode path (K3DSparkDrafter): its GEMV sites and its block attention. + "k3_ctm_gemv", + "k3_ctm_gemv_swiglu", + "k3_drafter_attn_qknorm", ) #: The engine surface the first forward checks before this target relies on it: per object, the attributes read. @@ -216,6 +232,13 @@ "tensorrt_llm._torch.modules.rms_norm.RMSNorm", "tensorrt_llm._torch.distributed.AllReduce", "tensorrt_llm._torch.modules.multi_stream_utils.maybe_execute_in_parallel", + # The DSpark drafter: the stock GQA drafter K3DSparkDrafter extends (its block forward where the drafter entries + # do not take a block, its context projection and context k / v, its heads and its weight load), and the stock + # builder's checks for which drafter a checkpoint gets. + "tensorrt_llm._torch.models.modeling_dspark.GQADSparkForCausalLM", + "tensorrt_llm._torch.models.modeling_dspark.draft_is_embedded_in_target", + "tensorrt_llm._torch.models.modeling_dflash.DFlashForCausalLM", + "tensorrt_llm._torch.pyexecutor.config_utils.is_mla", ) # The K3 decode kernels' bounds: the token-count kernels take one tile of DECODE_MAX_TOKENS rows; the request-aware @@ -2360,6 +2383,261 @@ def _k3_decode_view(self, attn_metadata: AttentionMetadata, num_tokens: int) -> return view +# ---------------------------------------------------------------------------------------------------------------------- +# The DSpark drafter: the stock GQA drafter, with a decode step's block on the drafter entries. +# ---------------------------------------------------------------------------------------------------------------------- + +#: The (requests, block tokens) splits attention/k3_drafter_attn_qknorm certifies. +DRAFTER_ATTN_SPLITS = frozenset({(1, 1), (1, 8), (2, 4), (4, 2), (8, 1), (3, 1), (2, 8), (8, 8)}) +# The drafter layout it certifies: 6 query heads and 1 KV head of 64 per rank, context pages of 64 rows, the q / k +# RMSNorm's epsilon and the NeoX RoPE's base. +_DRAFTER_HEADS = (6, 1) +_DRAFTER_HEAD_DIM = 64 +_DRAFTER_PAGE = 64 +_DRAFTER_EPS = 1e-5 +_DRAFTER_ROPE_BASE = 10000.0 + + +class K3DSparkDrafter(GQADSparkForCausalLM): + """Kimi K3's DSpark drafter: the stock GQA drafter, with a decode step's block forward on the drafter entries. + + A block the entries take runs, per layer: + + * the fused q / k / v projection on the ``drafter_qkv`` decode GEMV site; + * ``attention/k3_drafter_attn_qknorm``: the q / k RMSNorm, the NeoX RoPE, and each request's block attending to + its paged context and to its own k / v, in one launch. The block's k / v are not stored in the cache: the + worker writes the accepted rows' context k / v before a block reads them; + * the output projection on ``drafter_o``, then the module's all-reduce; + * gate / up on ``drafter_gate_up`` and the down projection with the SiLU-and-mul on ``drafter_down``, then the + module's all-reduce. + + The norms and the residual adds are the stock modules'; a projection whose site does not take its rows runs its + module. Every other block runs the stock ``dflash_forward``: a split `DRAFTER_ATTN_SPLITS` does not list, another + attention backend, head layout, RoPE or normalization, a cache the kernel does not read, or a compile key of the + attention that has not run eagerly, under CUDA-graph capture. The worker that calls it (the stock + ``DSparkWorker``), the context projection, the context k / v and the Markov head stay upstream code. + """ + + def __init__( + self, draft_config: ModelConfig, *, dflash_attention_backend: str = "AUTO" + ) -> None: + super().__init__(draft_config, dflash_attention_backend=dflash_attention_backend) + # The decode GEMVs' state (decode_gemv.py), shared by the target's layers; set by the target once built. + self.decode_gemvs: Optional[_decode_gemv.K3DecodeGemvs] = None + # Whether every layer is one the drafter entries take; decided on the first block, the weights loaded. + self._k3_take: Optional[bool] = None + # The attention's compile keys that ran eagerly ((more than one request, page stride)); a capture takes + # only these. + self._k3_attn_ran: set = set() + + def dflash_forward( + self, + noise_embedding: torch.Tensor, + query_positions: torch.Tensor, + num_ctx_per_req: torch.Tensor, + ctx_k_cache: torch.Tensor, + ctx_v_cache: torch.Tensor, + ctx_cache_batch_idx: torch.Tensor, + ctx_kv_cache: Optional[torch.Tensor] = None, + ctx_page_table: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + """The block's hidden states ``[B * block, hidden]``: on the drafter entries where they take it (the class + docstring), else the stock forward.""" + keys = self._k3_block_keys(noise_embedding, ctx_kv_cache, ctx_page_table) + if keys is None: + return super().dflash_forward( + noise_embedding, + query_positions, + num_ctx_per_req, + ctx_k_cache, + ctx_v_cache, + ctx_cache_batch_idx, + ctx_kv_cache, + ctx_page_table, + ) + out = self._k3_block_forward( + noise_embedding, + query_positions, + num_ctx_per_req, + ctx_cache_batch_idx, + ctx_kv_cache, + ctx_page_table, + ) + if not torch.cuda.is_current_stream_capturing(): + self._k3_attn_ran |= keys + return out + + def _k3_layers_take(self) -> bool: + """Whether every layer is one the drafter entries take: plain NeoX RoPE with the q / k RMSNorm at the head + layout, epsilon and RoPE base the attention certifies, bias-free projections, a plain SiLU-and-mul MLP, + non-causal unwindowed block attention, and no block convolution or post-attention gate.""" + if self._k3_take is None: + if self._fused_kv_weight is None: + self._build_fused_kv_buffers() + take = ( + self._use_fused_qk_norm_rope + and not self.has_block_conv + and type(self)._post_attention_gate is DFlashForCausalLM._post_attention_gate + ) + for layer_idx, layer in enumerate(self.model.layers): + attn, mlp = layer.self_attn, layer.mlp + pos = getattr(attn, "pos_embd_params", None) + rope = getattr(pos, "rope", None) + norms = (attn.q_norm, attn.k_norm) + take = take and ( + (attn.num_heads, attn.num_key_value_heads) == _DRAFTER_HEADS + and attn.head_dim == _DRAFTER_HEAD_DIM + and getattr(attn, "is_qk_norm", False) + and not getattr(attn, "use_gemma_rms_norm", True) + and rope is not None + and pos.is_neox + and getattr(pos, "mrope_section", None) is None + and rope.theta == _DRAFTER_ROPE_BASE + and all( + norm.variance_epsilon == _DRAFTER_EPS + and norm.weight.dtype == torch.bfloat16 + and tuple(norm.weight.shape) == (_DRAFTER_HEAD_DIM,) + for norm in norms + ) + and attn.qkv_proj.bias is None + and attn.o_proj.bias is None + and mlp.activation is torch.nn.functional.silu + and mlp.swiglu_limit is None + and mlp.swiglu_alpha in (None, 1.0) + and mlp.swiglu_beta in (None, 0.0) + and mlp.gate_up_proj.bias is None + and mlp.down_proj.bias is None + and self._resolve_block_attention(layer_idx) == (False, (-1, -1)) + ) + self._k3_take = bool(take) + logger.info( + f"Kimi K3 DSpark drafter: {len(self.model.layers)} layers, decode blocks " + + ( + "on k3_drafter_attn_qknorm and the drafter GEMV sites" + if self._k3_take + else "on the stock block forward (a layer the drafter entries do not take)" + ) + ) + return self._k3_take + + def _k3_block_keys( + self, + noise_embedding: torch.Tensor, + ctx_kv_cache: Optional[torch.Tensor], + ctx_page_table: Optional[torch.Tensor], + ) -> Optional[frozenset]: + """The attention's compile keys for this block when the drafter entries take it, else None.""" + if ( + self.dflash_attention_backend != "TRTLLM" + or ctx_kv_cache is None + or ctx_page_table is None + or noise_embedding.dtype != torch.bfloat16 + or tuple(noise_embedding.shape[:2]) not in DRAFTER_ATTN_SPLITS + or is_in_breakable_cuda_graph() + or not self._k3_layers_take() + ): + return None + num_kv_heads = self.model.layers[0].self_attn.num_key_value_heads + page = (2, num_kv_heads, _DRAFTER_PAGE, _DRAFTER_HEAD_DIM) + dense = ( + num_kv_heads * _DRAFTER_PAGE * _DRAFTER_HEAD_DIM, + _DRAFTER_PAGE * _DRAFTER_HEAD_DIM, + ) + strides = set() + for layer_idx in range(len(self.model.layers)): + cache = ctx_kv_cache[layer_idx] + if ( + cache.dtype != torch.bfloat16 + or cache.dim() != 5 + or tuple(cache.shape[1:]) != page + or cache.stride()[1:] != (*dense, _DRAFTER_HEAD_DIM, 1) + or cache.stride(0) % 8 + ): + return None + strides.add(cache.stride(0)) + keys = frozenset((noise_embedding.shape[0] > 1, stride) for stride in strides) + if torch.cuda.is_current_stream_capturing() and not keys <= self._k3_attn_ran: + return None + return keys + + def _k3_block_forward( + self, + noise_embedding: torch.Tensor, + query_positions: torch.Tensor, + num_ctx_per_req: torch.Tensor, + ctx_cache_batch_idx: torch.Tensor, + ctx_kv_cache: torch.Tensor, + ctx_page_table: torch.Tensor, + ) -> torch.Tensor: + batch, block = noise_embedding.shape[:2] + rows = batch * block + ctx_len = num_ctx_per_req[:batch].to(torch.int32) + page_table = ctx_page_table.index_select(0, ctx_cache_batch_idx.to(torch.long)) + positions = query_positions.reshape(-1).contiguous() + hidden = noise_embedding.reshape(rows, -1) + residual = None + for layer_idx, layer in enumerate(self.model.layers): + attn = layer.self_attn + if residual is None: + residual = hidden.clone() + normed = layer.input_layernorm(hidden) + else: + normed, residual = layer.input_layernorm(hidden, residual) + qkv = self._k3_project("drafter_qkv", normed, attn.qkv_proj) + out = torch.empty(rows, attn.q_size, dtype=torch.bfloat16, device=qkv.device) + k3_drafter_attn_qknorm( + qkv, + attn.q_norm.weight, + attn.k_norm.weight, + positions, + attn.q_norm.variance_epsilon, + attn.pos_embd_params.rope.theta, + ctx_kv_cache[layer_idx], + page_table, + ctx_len, + attn.num_heads, + attn.num_key_value_heads, + out, + ) + hidden = self._k3_project("drafter_o", out, attn.o_proj) + hidden, residual = layer.post_attention_layernorm(hidden, residual) + hidden = self._k3_mlp(layer.mlp, hidden) + out, _ = self.model.norm(hidden, residual) + return out + + def _k3_project(self, site: str, x: torch.Tensor, linear: nn.Module) -> torch.Tensor: + """``linear(x)``, its GEMM on the ``site`` decode GEMV where that takes the rows, then a row-parallel + projection's all-reduce.""" + y = None if self.decode_gemvs is None else self.decode_gemvs.project(site, x, linear.weight) + if y is None: + return linear(x) + return _row_parallel_reduce(linear, y) + + def _k3_mlp(self, mlp: nn.Module, x: torch.Tensor) -> torch.Tensor: + """``mlp(x)``: gate / up on ``drafter_gate_up`` and the SiLU-and-mul with the down projection on + ``drafter_down`` where both take the rows, then the down projection's all-reduce.""" + gemvs = self.decode_gemvs + gate_up = ( + None if gemvs is None else gemvs.project("drafter_gate_up", x, mlp.gate_up_proj.weight) + ) + down = ( + None + if gate_up is None + else gemvs.project("drafter_down", gate_up, mlp.down_proj.weight) + ) + if down is None: + return mlp(x) + return _row_parallel_reduce(mlp.down_proj, down) + + +def _row_parallel_reduce(linear: nn.Module, partial: torch.Tensor) -> torch.Tensor: + """``partial`` all-reduced as ``linear`` reduces its output: a row-parallel projection's all-reduce, if any.""" + all_reduce = getattr(linear, "all_reduce", None) + if getattr(getattr(linear, "tp_mode", None), "name", None) == "ROW" and all_reduce is not None: + return all_reduce(partial) + return partial + + # ---------------------------------------------------------------------------------------------------------------------- # The target: step classification, the construction checks and the registration shell. # ---------------------------------------------------------------------------------------------------------------------- @@ -2547,13 +2825,31 @@ def __init__(self, model_config: ModelConfig): model_config.pretrained_config = self.config model_config._frozen = True + def _build_draft_model( + self, model_config: ModelConfig, draft_config: Optional[ModelConfig] + ) -> Optional[nn.Module]: + """DSpark's drafter from a standalone GQA checkpoint is this target's `K3DSparkDrafter`; any other drafter + is the mode registry's.""" + spec_config = model_config.spec_config + if ( + spec_config.spec_dec_mode.is_dspark() + and draft_config is not None + and not draft_is_embedded_in_target(model_config) + and not is_mla(draft_config.pretrained_config) + ): + return K3DSparkDrafter( + draft_config, dflash_attention_backend=spec_config.attention_backend + ) + return super()._build_draft_model(model_config, draft_config) + def load_weights(self, weights, *args, **kwargs): _weights.load(self, weights) def cache_derived_state(self) -> None: """Build the decode GEMVs' state once the weights are final: the LM head's workspace, and one eager call of every decode GEMV kernel at its site's shape (decode_gemv.SITES), so none compiles under capture. Built once: - a later call keeps it, since CUDA graphs captured in between hold its workspace.""" + a later call keeps it, since CUDA graphs captured in between hold its workspace. The DSpark drafter's sites + are among them, and the drafter gets the state too.""" super().cache_derived_state() if self.model.decode_gemvs is not None: return @@ -2563,6 +2859,8 @@ def cache_derived_state(self) -> None: for layer in self.model.layers: if not layer.is_moe: layer.decode_gemvs = gemvs + if isinstance(self.draft_model, K3DSparkDrafter): + self.draft_model.decode_gemvs = gemvs def post_load_weights(self) -> None: """The state the decode kernels share, built once per device before any CUDA-graph capture and handed to From b79745b29b3b61af3c4ef73983ba5dcf351ddd18 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 23:14:23 -0700 Subject: [PATCH 114/161] [None][test] modeling_v2 Kimi K3 target: the DSpark drafter against the stock block forward One GPU at one TP16 rank's drafter shapes (6 / 1 heads of 64, hidden 7168, MLP width 896, 64-row context pages), two layers: every split attention/k3_drafter_attn_qknorm certifies runs the entry once per layer and matches DFlashForCausalLM.dflash_forward on the same module and context, with the torch GEMMs and with the decode GEMV sites (taken at up to 8 rows, declined above), and leaves the cache untouched; an uncertified split runs the stock forward bit for bit; under capture the entries take a block only once its compile key ran eagerly; a weight changed between the two forwards fails the comparison. Listed in l0_b200.yml next to the decode GEMV test. Signed-off-by: Vasanth Sabavat --- .../test_lists/test-db/l0_b200.yml | 2 + .../test_modeling_v2_kimi_k3_drafter.py | 247 ++++++++++++++++++ 2 files changed, 249 insertions(+) create mode 100644 tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index 1e4e2de89696..14470959f083 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -160,6 +160,8 @@ l0_b200: - unittest/_torch/modeling_v2/gemm/test_modeling_v2_k3_head_gemv.py # The Kimi K3 target's decode GEMVs, LM head and embedding on those entries. - unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_gemv.py + # The Kimi K3 target's DSpark drafter on the drafter entries, against the stock block forward. + - unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py # KDA runtime: host-derived prefill metadata and the bf16 state pool # round-trip. Both are single-device cases that use GPU 0 only. - unittest/_torch/modules/kimi_kda/test_kda_host_metadata.py diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py new file mode 100644 index 000000000000..171636aac594 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py @@ -0,0 +1,247 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""The Kimi K3 target's DSpark drafter (``K3DSparkDrafter`` of ``kimi_k3_mxfp4__sm_100__tp16_moetp4ep4``) on one GPU, +at one TP16 rank's drafter shapes: 6 query heads and 1 KV head of 64, hidden 7168, MLP width 896, context pages of +64 rows, two layers. + +* A decode block of every split ``attention/k3_drafter_attn_qknorm`` certifies runs the drafter entries (the + attention once per layer) and matches the stock block forward (``DFlashForCausalLM.dflash_forward`` on the same + module and context), with the torch GEMMs and with the decode GEMV sites (which take the blocks of at most 8 rows). + The entries leave the cache untouched. +* A split the entry does not certify runs the stock forward, bit for bit, without the entry. +* Under CUDA-graph capture, the entries take a block only once its attention compile key has run eagerly. +* Negative control: a weight changed between the two forwards fails the comparison. +""" + +import math + +import pytest +import torch +from transformers import Qwen3Config + +from tensorrt_llm._torch._experimental.modeling_v2.models.kimi_k3_vl.kimi_k3_mxfp4__sm_100__tp16_moetp4ep4 import ( # noqa: E501 + decode_gemv, +) +from tensorrt_llm._torch._experimental.modeling_v2.models.kimi_k3_vl.kimi_k3_mxfp4__sm_100__tp16_moetp4ep4 import ( # noqa: E501 + modeling as target, +) +from tensorrt_llm._torch.model_config import ModelConfig +from tensorrt_llm._torch.models.modeling_dflash import DFlashForCausalLM + +pytestmark = pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.get_device_capability() != (10, 0), + reason="the K3 drafter entries run on sm_100 only", +) + +HIDDEN, INTER, HEADS, KV_HEADS, HEAD_DIM, LAYERS, PAGE = 7168, 896, 6, 1, 64, 2, 64 +VOCAB, RANK = 1024, 16 +# Context lengths per request: inside a page, on a page boundary, across pages, empty. +CTX = (37, 64, 130, 0, 5, 200, 63, 1) +REL_L2 = 1e-2 +SITES = ("drafter_qkv", "drafter_o", "drafter_gate_up", "drafter_down") + + +def _config(): + return Qwen3Config.from_dict( + dict( + architectures=["Qwen3ForCausalLM"], + model_type="qwen3", + hidden_size=HIDDEN, + intermediate_size=INTER, + num_hidden_layers=LAYERS, + num_attention_heads=HEADS, + num_key_value_heads=KV_HEADS, + head_dim=HEAD_DIM, + hidden_act="silu", + rms_norm_eps=1e-5, + vocab_size=VOCAB, + max_position_embeddings=4096, + rope_theta=10000.0, + rope_scaling=None, + attention_bias=False, + torch_dtype="bfloat16", + tie_word_embeddings=False, + markov_rank=RANK, + markov_head_type="vanilla", + enable_confidence_head=True, + dflash_config={"mask_token_id": VOCAB - 2, "target_layer_ids": [0, 1]}, + ) + ) + + +def _weights(seed=11): + g = torch.Generator().manual_seed(seed) + + def rnd(*shape, scale=0.02): + return (torch.randn(*shape, generator=g) * scale).to(torch.bfloat16) + + w = { + "fc.weight": rnd(HIDDEN, 2 * HIDDEN), + "hidden_norm.weight": rnd(HIDDEN, scale=0.05) + 1.0, + "norm.weight": rnd(HIDDEN, scale=0.05) + 1.0, + "markov_head.markov_w1.weight": rnd(VOCAB, RANK), + "markov_head.markov_w2.weight": rnd(VOCAB, RANK), + "confidence_head.proj.weight": rnd(1, HIDDEN + RANK), + "confidence_head.proj.bias": rnd(1), + } + for i in range(LAYERS): + p = f"layers.{i}." + w[p + "self_attn.q_proj.weight"] = rnd(HEADS * HEAD_DIM, HIDDEN) + w[p + "self_attn.k_proj.weight"] = rnd(KV_HEADS * HEAD_DIM, HIDDEN) + w[p + "self_attn.v_proj.weight"] = rnd(KV_HEADS * HEAD_DIM, HIDDEN) + w[p + "self_attn.o_proj.weight"] = rnd(HIDDEN, HEADS * HEAD_DIM) + w[p + "self_attn.q_norm.weight"] = rnd(HEAD_DIM, scale=0.05) + 1.0 + w[p + "self_attn.k_norm.weight"] = rnd(HEAD_DIM, scale=0.05) + 1.0 + w[p + "input_layernorm.weight"] = rnd(HIDDEN, scale=0.05) + 1.0 + w[p + "post_attention_layernorm.weight"] = rnd(HIDDEN, scale=0.05) + 1.0 + w[p + "mlp.gate_proj.weight"] = rnd(INTER, HIDDEN) + w[p + "mlp.up_proj.weight"] = rnd(INTER, HIDDEN) + w[p + "mlp.down_proj.weight"] = rnd(HIDDEN, INTER) + return w + + +@pytest.fixture(scope="module") +def drafter(): + model_config = ModelConfig(pretrained_config=_config(), attn_backend="TRTLLM") + module = target.K3DSparkDrafter(model_config, dflash_attention_backend="TRTLLM").to("cuda") + module.load_weights(_weights()) + assert module._k3_layers_take(), "the test drafter must be one the drafter entries take" + return module + + +@pytest.fixture(scope="module") +def gemvs(): + return decode_gemv.K3DecodeGemvs.create(None, sites=SITES) + + +def _block(batch, block, seed): + """A block's inputs: noise rows, query positions after each request's context, and a paged context cache + (random rows, the requests' pages scattered) with its page table.""" + g = torch.Generator(device="cuda").manual_seed(seed) + ctx = torch.tensor(CTX[:batch], dtype=torch.int32, device="cuda") + width = math.ceil((int(ctx.max()) + block) / PAGE) + 1 + pages = batch * width + 3 + order = torch.randperm(pages, generator=g, device="cuda")[: batch * width] + caches = [ + (torch.randn(pages, 2, KV_HEADS, PAGE, HEAD_DIM, generator=g, device="cuda")).to( + torch.bfloat16 + ) + for _ in range(LAYERS) + ] + return dict( + noise_embedding=torch.randn(batch, block, HIDDEN, generator=g, device="cuda").to( + torch.bfloat16 + ), + query_positions=(ctx.long()[:, None] + torch.arange(block, device="cuda")), + num_ctx_per_req=ctx, + ctx_k_cache=None, + ctx_v_cache=None, + ctx_cache_batch_idx=torch.arange(batch, dtype=torch.int32, device="cuda"), + ctx_kv_cache=caches, + ctx_page_table=order.view(batch, width).to(torch.int32), + ) + + +def _copy(inputs): + out = dict(inputs) + out["ctx_kv_cache"] = [c.clone() for c in inputs["ctx_kv_cache"]] + return out + + +def _rel_l2(y, ref): + return ((y.double() - ref.double()).norm() / ref.double().norm()).item() + + +@pytest.fixture +def attn_calls(monkeypatch): + """Counts the drafter's calls of the attention entry.""" + calls = [] + entry = target.k3_drafter_attn_qknorm + + def counted(*args): + calls.append(args[0].shape[0]) + return entry(*args) + + monkeypatch.setattr(target, "k3_drafter_attn_qknorm", counted) + return calls + + +@pytest.mark.parametrize("use_gemvs", [False, True], ids=["torch", "gemv"]) +@pytest.mark.parametrize("split", sorted(target.DRAFTER_ATTN_SPLITS)) +def test_block_matches_the_stock_forward(drafter, gemvs, attn_calls, monkeypatch, use_gemvs, split): + batch, block = split + projected = [] + project = gemvs.project + + def recorded(site, x, weight): + y = project(site, x, weight) + projected.append((site, y is not None)) + return y + + monkeypatch.setattr(gemvs, "project", recorded) + drafter.decode_gemvs = gemvs if use_gemvs else None + inputs = _block(batch, block, seed=batch * 10 + block) + stock_inputs = _copy(inputs) + before = [c.clone() for c in inputs["ctx_kv_cache"]] + out = drafter.dflash_forward(**inputs) + ref = DFlashForCausalLM.dflash_forward(drafter, **stock_inputs) + torch.cuda.synchronize() + assert attn_calls == [batch * block] * LAYERS + assert out.shape == ref.shape == (batch * block, HIDDEN) + err = _rel_l2(out, ref) + assert err <= REL_L2, err + assert all(torch.equal(c, b) for c, b in zip(inputs["ctx_kv_cache"], before)) + if use_gemvs: + # At most 8 rows every site takes its projection; above, each declines and its module runs (the down + # projection's site is not asked once gate / up's declined). + takes = batch * block <= decode_gemv.MAX_ROWS + asked = SITES if takes else SITES[:3] + assert sorted(projected) == sorted( + (site, takes) for site in asked for _ in range(LAYERS) + ), projected + else: + assert projected == [] + + +def test_an_uncertified_split_runs_the_stock_forward(drafter, attn_calls): + drafter.decode_gemvs = None + batch, block = 2, 7 + assert (batch, block) not in target.DRAFTER_ATTN_SPLITS + inputs = _block(batch, block, seed=27) + stock_inputs = _copy(inputs) + out = drafter.dflash_forward(**inputs) + ref = DFlashForCausalLM.dflash_forward(drafter, **stock_inputs) + assert attn_calls == [] + assert torch.equal(out, ref) + + +def test_capture_takes_only_compiled_keys(drafter): + inputs = _block(1, 8, seed=18) + args = (inputs["noise_embedding"], inputs["ctx_kv_cache"], inputs["ctx_page_table"]) + drafter._k3_attn_ran.clear() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + refused = drafter._k3_block_keys(*args) + assert refused is None + drafter.dflash_forward(**inputs) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + taken = drafter._k3_block_keys(*args) + assert taken is not None and taken <= drafter._k3_attn_ran + + +def test_negative_control_a_changed_weight_fails(drafter): + drafter.decode_gemvs = None + inputs = _block(1, 8, seed=81) + stock_inputs = _copy(inputs) + weight = drafter.model.layers[0].self_attn.o_proj.weight + saved = weight.detach().clone() + try: + with torch.no_grad(): + weight.add_(torch.randn_like(weight) * 0.02) + out = drafter.dflash_forward(**inputs) + finally: + with torch.no_grad(): + weight.copy_(saved) + ref = DFlashForCausalLM.dflash_forward(drafter, **stock_inputs) + assert _rel_l2(out, ref) > REL_L2 From 592ea1a65b983a11d83b686ee4086ab94b9db52a Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 23:16:19 -0700 Subject: [PATCH 115/161] [None][test] Kimi K3 drafter attention: the R x 7 splits DSpark's draft block under shift_label is max_draft_len tokens (7 at max_draft_len 7), not max_draft_len + 1, so every DSpark decode step calls k3_drafter_attn(_qknorm) with T = 7. The op test and both catalog entries' tests now run R x 7 for R = 1..8 alongside the splits they already certify. Signed-off-by: Vasanth Sabavat --- .../kimi_k3/test_k3_drafter_attn.py | 19 +++++++++++-------- .../test_modeling_v2_k3_drafter_attn.py | 4 +++- ...test_modeling_v2_k3_drafter_attn_qknorm.py | 4 +++- 3 files changed, 17 insertions(+), 10 deletions(-) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_drafter_attn.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_drafter_attn.py index 0230d4da41ec..22f7cdb5bcc3 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_drafter_attn.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_drafter_attn.py @@ -16,10 +16,10 @@ requests of T tokens, at the in-model TP16 group (6 query heads, 1 KV head; TP4 24 / 4 for a few splits), head dim 64, HND pages of 64. -Every split R x T (R T <= 8, and R x 8 for R <= 8) with mixed per-request context lengths (blocks crossing a page, a -128-row tile and the cluster's 16-tile round, no context, several tiles per CTA), page tables as strided row views -and as dense rows, against an fp32 reference and the model's production path (flashinfer append_paged_kv_cache + -trtllm-gen batch_context_with_kv_cache, non-causal), plus: no NaN (the pool's unused rows are NaN), the pool left +Every split R x T (R T <= 8, and R x 8 and R x 7 for R <= 8) with mixed per-request context lengths (blocks crossing a +page, a 128-row tile and the cluster's 16-tile round, no context, several tiles per CTA), page tables as strided row +views and as dense rows, against an fp32 reference and the model's production path (flashinfer append_paged_kv_cache ++ trtllm-gen batch_context_with_kv_cache, non-causal), plus: no NaN (the pool's unused rows are NaN), the pool left untouched, reruns bit-identical, requests isolated (a change in one request's context or block changes only its rows), CUDA-graph replays with rewritten inputs. @@ -37,10 +37,13 @@ THETA = 10000.0 TOL = 1e-2 # max |err| / max |ref|; bf16 P and output roundings are ~4e-3 -# Every split the engine schedules for the drafter: R requests x T tokens with R T <= 8, and DSpark's R x 8. -SPLITS = [(1, 8), (2, 4), (4, 2), (8, 1), (1, 1), (2, 1), (3, 1), (4, 1), (5, 1)] + [ - (r, 8) for r in range(2, 9) -] +# Every split the engine schedules for the drafter: R requests x T tokens with R T <= 8, R x 8, and DSpark's R x 7 +# (its block under shift_label is max_draft_len tokens). +SPLITS = ( + [(1, 8), (2, 4), (4, 2), (8, 1), (1, 1), (2, 1), (3, 1), (4, 1), (5, 1)] + + [(r, 8) for r in range(2, 9)] + + [(r, 7) for r in range(1, 9)] +) # Context lengths, cycled over the requests: blocks crossing a page (60, 121), a 128-row tile (124, 127), the # cluster's 16-tile round (2040, 2041); starting a page / tile (0, 64, 128, 1024, 2048); several tiles per CTA. LENGTHS = [60, 2040, 5, 124, 1000, 0, 2041, 64, 3000, 127, 1024, 121, 128, 4000, 2048, 1500] diff --git a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn.py index f7180da1bee5..f75a33353d2e 100644 --- a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn.py +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn.py @@ -28,7 +28,9 @@ # Context lengths, cycled over the requests: none, a block crossing a page, page and 128-row tile boundaries, the # cluster's 16-tile round, several tiles per CTA. LENGTHS = (60, 2041, 0, 127, 1000, 64, 128, 5) -SPLITS = ((1, 1), (1, 8), (2, 4), (4, 2), (8, 1), (3, 1), (2, 8), (8, 8)) +SPLITS = ((1, 1), (1, 8), (2, 4), (4, 2), (8, 1), (3, 1), (2, 8), (8, 8)) + tuple( + (r, 7) for r in range(1, 9) +) # R x 7: DSpark's block under shift_label (max_draft_len 7) def _bits(t: torch.Tensor) -> torch.Tensor: diff --git a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn_qknorm.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn_qknorm.py index 3f08befe6d89..0494fe8b26a9 100644 --- a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn_qknorm.py +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn_qknorm.py @@ -24,7 +24,9 @@ # Context lengths, cycled over the requests: none, a block crossing a page, page and 128-row tile boundaries, the # cluster's 16-tile round, several tiles per CTA. LENGTHS = (60, 2041, 0, 127, 1000, 64, 128, 5) -SPLITS = ((1, 1), (1, 8), (2, 4), (4, 2), (8, 1), (3, 1), (2, 8), (8, 8)) +SPLITS = ((1, 1), (1, 8), (2, 4), (4, 2), (8, 1), (3, 1), (2, 8), (8, 8)) + tuple( + (r, 7) for r in range(1, 9) +) # R x 7: DSpark's block under shift_label (max_draft_len 7) def _bits(t: torch.Tensor) -> torch.Tensor: From 31b6275c70d1018d4ddd392ab96f9f622bd197ed Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 23:20:18 -0700 Subject: [PATCH 116/161] [None][test] Kimi K3 drafter attention tests: name the splits they run The op test's list names each split's source (R T <= 8 short blocks, DFlash's R x 8, DSpark's R x 7) instead of "every split the engine schedules", its docstring says "the splits", and the catalog tests' R x 7 comment sits above SPLITS. Signed-off-by: Vasanth Sabavat --- .../_torch/cute_dsl_kernels/kimi_k3/test_k3_drafter_attn.py | 6 +++--- .../attention/test_modeling_v2_k3_drafter_attn.py | 3 ++- .../attention/test_modeling_v2_k3_drafter_attn_qknorm.py | 3 ++- 3 files changed, 7 insertions(+), 5 deletions(-) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_drafter_attn.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_drafter_attn.py index 22f7cdb5bcc3..e31de3eb3902 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_drafter_attn.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_drafter_attn.py @@ -16,7 +16,7 @@ requests of T tokens, at the in-model TP16 group (6 query heads, 1 KV head; TP4 24 / 4 for a few splits), head dim 64, HND pages of 64. -Every split R x T (R T <= 8, and R x 8 and R x 7 for R <= 8) with mixed per-request context lengths (blocks crossing a +The splits R x T (R T <= 8, and R x 8 and R x 7 for R <= 8) with mixed per-request context lengths (blocks crossing a page, a 128-row tile and the cluster's 16-tile round, no context, several tiles per CTA), page tables as strided row views and as dense rows, against an fp32 reference and the model's production path (flashinfer append_paged_kv_cache + trtllm-gen batch_context_with_kv_cache, non-causal), plus: no NaN (the pool's unused rows are NaN), the pool left @@ -37,8 +37,8 @@ THETA = 10000.0 TOL = 1e-2 # max |err| / max |ref|; bf16 P and output roundings are ~4e-3 -# Every split the engine schedules for the drafter: R requests x T tokens with R T <= 8, R x 8, and DSpark's R x 7 -# (its block under shift_label is max_draft_len tokens). +# The splits tested, R requests x T tokens: R T <= 8 (short blocks), R x 8 (DFlash's block at max_draft_len 7: +# max_draft_len + 1 tokens) and R x 7 (DSpark's block at max_draft_len 7: max_draft_len tokens under shift_label). SPLITS = ( [(1, 8), (2, 4), (4, 2), (8, 1), (1, 1), (2, 1), (3, 1), (4, 1), (5, 1)] + [(r, 8) for r in range(2, 9)] diff --git a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn.py index f75a33353d2e..b7036f753adb 100644 --- a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn.py +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn.py @@ -28,9 +28,10 @@ # Context lengths, cycled over the requests: none, a block crossing a page, page and 128-row tile boundaries, the # cluster's 16-tile round, several tiles per CTA. LENGTHS = (60, 2041, 0, 127, 1000, 64, 128, 5) +# R x 7: DSpark's block at max_draft_len 7 (max_draft_len tokens under shift_label). SPLITS = ((1, 1), (1, 8), (2, 4), (4, 2), (8, 1), (3, 1), (2, 8), (8, 8)) + tuple( (r, 7) for r in range(1, 9) -) # R x 7: DSpark's block under shift_label (max_draft_len 7) +) def _bits(t: torch.Tensor) -> torch.Tensor: diff --git a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn_qknorm.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn_qknorm.py index 0494fe8b26a9..3a1657098d1b 100644 --- a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn_qknorm.py +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn_qknorm.py @@ -24,9 +24,10 @@ # Context lengths, cycled over the requests: none, a block crossing a page, page and 128-row tile boundaries, the # cluster's 16-tile round, several tiles per CTA. LENGTHS = (60, 2041, 0, 127, 1000, 64, 128, 5) +# R x 7: DSpark's block at max_draft_len 7 (max_draft_len tokens under shift_label). SPLITS = ((1, 1), (1, 8), (2, 4), (4, 2), (8, 1), (3, 1), (2, 8), (8, 8)) + tuple( (r, 7) for r in range(1, 9) -) # R x 7: DSpark's block under shift_label (max_draft_len 7) +) def _bits(t: torch.Tensor) -> torch.Tensor: From 18f1edd040e9f5f42e6f42375b9e0a2c913bed12 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 00:20:55 -0700 Subject: [PATCH 117/161] [None][doc] modeling_v2 catalog: the Kimi K3 drafter attention certifies R x 7 attention/k3_drafter_attn and attention/k3_drafter_attn_qknorm add 1x7 to 8x7 to their certified splits (DSpark's block at max_draft_len 7). Receipts on sm_100 (GB200, the k3rest build ed82d3435b with these tests): k3_drafter_attn 36 tests passed (the 16 splits at both head layouts, table forms, graph replay, the out-of-contract refusals), k3_drafter_attn_qknorm 16 passed; the op test 105 passed. Signed-off-by: Vasanth Sabavat --- .../modeling_v2/catalog/attention/k3_drafter_attn.md | 7 ++++--- .../catalog/attention/k3_drafter_attn_qknorm.md | 5 +++-- 2 files changed, 7 insertions(+), 5 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn.md index 8f816d76aca1..5407ce7d4a1c 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn.md @@ -1,6 +1,6 @@ --- receipts: - sm_100: {status: passed, tests: 20} + sm_100: {status: passed, tests: 36} --- # k3_drafter_attn @@ -53,8 +53,9 @@ def k3_drafter_attn( | `num_heads`, `num_kv_heads` | 6 / 1 (Kimi K3 at TP16) and 24 / 4 (TP4) certified; `num_heads` = 6 `num_kv_heads` | Python int | — | — | | `out` | `[M, num_heads * 64]` | bf16 | contiguous | CUDA | -Certified splits `R x T`: 1x1, 1x8, 2x4, 4x2, 8x1, 3x1, 2x8, 8x8, with context lengths 0 to 2041 (blocks crossing a -page, 64- and 128-row boundaries, 2041 rows spanning the kernel's 16-tile round), at both head layouts. +Certified splits `R x T`: 1x1, 1x8, 2x4, 4x2, 8x1, 3x1, 2x8, 8x8, and 1x7 to 8x7 (DSpark's block at +`max_draft_len` 7), with context lengths 0 to 2041 (blocks crossing a page, 64- and 128-row boundaries, 2041 rows +spanning the kernel's 16-tile round), at both head layouts. ## Metadata consumed diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn_qknorm.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn_qknorm.md index 45b60ecba263..e3100d69ad20 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn_qknorm.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/k3_drafter_attn_qknorm.md @@ -1,6 +1,6 @@ --- receipts: - sm_100: {status: passed, tests: 8} + sm_100: {status: passed, tests: 16} --- # k3_drafter_attn_qknorm @@ -55,7 +55,8 @@ def k3_drafter_attn_qknorm( | `cache`, `page_table`, `ctx_len`, `out` | as `k3_drafter_attn` | | | | | `num_heads`, `num_kv_heads` | 6 / 1 certified | Python int | — | — | -Certified splits `R x T`: 1x1, 1x8, 2x4, 4x2, 8x1, 3x1, 2x8, 8x8, context lengths 0 to 2041. +Certified splits `R x T`: 1x1, 1x8, 2x4, 4x2, 8x1, 3x1, 2x8, 8x8, and 1x7 to 8x7 (DSpark's block at +`max_draft_len` 7), context lengths 0 to 2041. ## Metadata consumed From 0d16d098615f0e8a20b44f3fc383edf9c103f57e Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 00:20:56 -0700 Subject: [PATCH 118/161] [None][feat] modeling_v2 Kimi K3 target: the DSpark drafter takes R x 7 blocks DRAFTER_ATTN_SPLITS gains 1x7 to 8x7, now certified by the attention entries, so a dspark7 block (7 tokens per request under shift_label) runs on the drafter entries. The drafter test's uncertified split is 4 x 4. Signed-off-by: Vasanth Sabavat --- .../kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py | 7 +++++-- .../_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py | 4 ++-- 2 files changed, 7 insertions(+), 4 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index e4a5c3396d77..36cf98f2eac2 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -2387,8 +2387,11 @@ def _k3_decode_view(self, attn_metadata: AttentionMetadata, num_tokens: int) -> # The DSpark drafter: the stock GQA drafter, with a decode step's block on the drafter entries. # ---------------------------------------------------------------------------------------------------------------------- -#: The (requests, block tokens) splits attention/k3_drafter_attn_qknorm certifies. -DRAFTER_ATTN_SPLITS = frozenset({(1, 1), (1, 8), (2, 4), (4, 2), (8, 1), (3, 1), (2, 8), (8, 8)}) +#: The (requests, block tokens) splits attention/k3_drafter_attn_qknorm certifies; R x 7 is DSpark's block at +#: max_draft_len 7. +DRAFTER_ATTN_SPLITS = frozenset( + {(1, 1), (1, 8), (2, 4), (4, 2), (8, 1), (3, 1), (2, 8), (8, 8)} | {(r, 7) for r in range(1, 9)} +) # The drafter layout it certifies: 6 query heads and 1 KV head of 64 per rank, context pages of 64 rows, the q / k # RMSNorm's epsilon and the NeoX RoPE's base. _DRAFTER_HEADS = (6, 1) diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py index 171636aac594..f2dad0161298 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py @@ -205,9 +205,9 @@ def recorded(site, x, weight): def test_an_uncertified_split_runs_the_stock_forward(drafter, attn_calls): drafter.decode_gemvs = None - batch, block = 2, 7 + batch, block = 4, 4 assert (batch, block) not in target.DRAFTER_ATTN_SPLITS - inputs = _block(batch, block, seed=27) + inputs = _block(batch, block, seed=44) stock_inputs = _copy(inputs) out = drafter.dflash_forward(**inputs) ref = DFlashForCausalLM.dflash_forward(drafter, **stock_inputs) From 706cf2d4ff98d39c9dba008d9c60e28cadd3a9b3 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 05:09:34 -0700 Subject: [PATCH 119/161] [None][test] Kimi K3 kernel lints: the KDA, MLA and drafter kernels' rows The source lints that keep the Kimi K3 kernels' cluster-scope mailbox waits and tcgen05 thread-sync fences list only the GEMV, MoE and sandwich kernels. The KDA (k3_kda_attn, k3_kda_decode, k3_kda_verify), MLA (k3_mla_attn, k3_mla_q) and drafter (k3_drafter_attn) kernels these rows were written for came in other changes, and their rows were left out: - test_k3_cluster_waits.py: their st.async mailboxes, and k3_mla_attn's no_cluster block, whose waits are on the CTA's own arrival; - test_k3_tcgen05_fences.py: the five KDA and MLA kernels. All six kernels already pass both lints; the rows keep them that way. Signed-off-by: Vasanth Sabavat --- .../kimi_k3/test_k3_cluster_waits.py | 19 ++++++++++++++++--- .../kimi_k3/test_k3_tcgen05_fences.py | 5 +++++ 2 files changed, 21 insertions(+), 3 deletions(-) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_cluster_waits.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_cluster_waits.py index e9860683aab6..6ce12ecb6ac5 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_cluster_waits.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_cluster_waits.py @@ -28,8 +28,19 @@ # Kernel module -> the barriers in it that other CTAs complete with st.async. MAILBOXES = { + "tensorrt_llm._torch.cute_dsl_kernels.k3_kda_attn.k3_kda_attn_kernel": ("mbox_bar", "ss_ready"), + "tensorrt_llm._torch.cute_dsl_kernels.k3_kda_attn.k3_kda_decode_kernel": ("ss_ready",), + "tensorrt_llm._torch.cute_dsl_kernels.k3_kda_verify.k3_kda_verify_kernel": ("ss_ready",), + "tensorrt_llm._torch.cute_dsl_kernels.k3_drafter.k3_drafter_attn_kernel": ("mail_full",), "tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv.k3_ctm_gemv_kernel": ("mail_full",), "tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.k3_moe_front": ("mail_full",), + "tensorrt_llm._torch.cute_dsl_kernels.k3_mla.k3_mla_attn_kernel": ("ml_full",), + "tensorrt_llm._torch.cute_dsl_kernels.k3_mla.k3_mla_q_kernel": ( + "qn_full", + "norm_full", + "kb_full", + "mail_full", + ), # rms_full: stage 0 only (filled by the cluster CTAs' st.async pushes); its other stages are completed by this # CTA's own arrive or bulk copy. "tensorrt_llm._torch.cute_dsl_kernels.k3_sandwich.k3_sandwich_kernel": ( @@ -43,9 +54,11 @@ "mb_sq", ), } -# Kernel module -> the header of the blocks whose mailbox waits are on this CTA's own arrival and keep CTA scope (a -# kernel that arrives on its mailbox itself and acquires the other CTAs' data through a counter). -OWN_ARRIVAL = {} +# Kernel module -> the header of the blocks whose mailbox waits are on this CTA's own arrival and keep CTA scope +# (k3_mla_attn's no_cluster mode arrives on ml_full itself and acquires the other CTAs' (m, l) through a counter). +OWN_ARRIVAL = { + "tensorrt_llm._torch.cute_dsl_kernels.k3_mla.k3_mla_attn_kernel": r"if cutlass\.const_expr\(no_cluster\):", +} CTA_WAIT = re.compile(r"mbarrier_(test|try)_wait\(") CLUSTER_WAIT = re.compile(r"_(test|try)_wait_cluster\(") diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py index 14edccd65cc6..57d843d805e8 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py @@ -35,6 +35,11 @@ "tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.k3_moe_kernel", "tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.k3_moe_m1_kernel", "tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.k3_moe_m2_kernel", + "tensorrt_llm._torch.cute_dsl_kernels.k3_kda_attn.k3_kda_attn_kernel", + "tensorrt_llm._torch.cute_dsl_kernels.k3_kda_attn.k3_kda_decode_kernel", + "tensorrt_llm._torch.cute_dsl_kernels.k3_kda_verify.k3_kda_verify_kernel", + "tensorrt_llm._torch.cute_dsl_kernels.k3_mla.k3_mla_attn_kernel", + "tensorrt_llm._torch.cute_dsl_kernels.k3_mla.k3_mla_q_kernel", "tensorrt_llm._torch.cute_dsl_kernels.k3_sandwich.k3_sandwich_kernel", ] From f0e7d4c1b62eff6a425d0a4c3921f7b3e58f3c78 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 05:10:02 -0700 Subject: [PATCH 120/161] [None][test] Kimi K3 SiTU MoE: the TP16 shard (192 values, padded to 256) test_tp8_sharded_forward_matches_whole_expert becomes test_tp_sharded_forward_matches_whole_expert at TP8 and TP16. At TP16 each rank's slice of an expert is 192 intermediate values, which the MXFP4 quant method zero-pads to 256 (the weight alignment) and runs as 192 valid. The sum of the 16 shard outputs must match the whole expert's output within the bf16 band, as the 8 shards' sum does. This is the expert layout of the Kimi K3 tp16_moetp16ep1 target. Signed-off-by: Vasanth Sabavat --- tests/unittest/_torch/moe/test_kimi_k3_situ_moe.py | 13 +++++++------ 1 file changed, 7 insertions(+), 6 deletions(-) diff --git a/tests/unittest/_torch/moe/test_kimi_k3_situ_moe.py b/tests/unittest/_torch/moe/test_kimi_k3_situ_moe.py index 7f8dddd5adde..cac54952cafb 100644 --- a/tests/unittest/_torch/moe/test_kimi_k3_situ_moe.py +++ b/tests/unittest/_torch/moe/test_kimi_k3_situ_moe.py @@ -1178,16 +1178,17 @@ def test_tp_shard_loader_matches_manual_slice(tp_size): @situ_supported +@pytest.mark.parametrize("tp_size", [8, 16], ids=lambda n: f"tp{n}") @pytest.mark.parametrize("num_tokens", [1, 16], ids=lambda n: f"tokens{n}") -def test_tp8_sharded_forward_matches_whole_expert(num_tokens): - """Sum of 8 TP-shard partial outputs == whole-expert reference. +def test_tp_sharded_forward_matches_whole_expert(num_tokens, tp_size): + """Sum of the TP-shard partial outputs == whole-expert reference. Per-element MXFP4/MXFP8 numerics are identical between the two layouts - (group-32 boundaries align: 384 % 32 == 0), so the only expected error - is bf16 rounding of the per-shard FC2 partial sums. + (group-32 boundaries align: 384 and 192 are multiples of 32), so the only + expected error is bf16 rounding of the per-shard FC2 partial sums. At TP16 + each shard's 192-wide slice is zero-padded to 256 and run as valid 192. """ - tp_size = 8 - ipp = _TP_INTERMEDIATE // tp_size # 384 — the production TP8 shard size + ipp = _TP_INTERMEDIATE // tp_size # 384 at TP8; 192 at TP16 (the no-spec layout) bank = _make_packed_expert_bank(_TP_EXPERTS, _TP_INTERMEDIATE, _TP_HIDDEN) gate = _make_test_gate() From 87c7e70ce3ad7d9ba82b7551eccf35be640fed44 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 05:25:28 -0700 Subject: [PATCH 121/161] [None][test] modeling_v2 MNNVL op matrices: run on sm_100 only The mnnvl_fusion_allreduce and mnnvl_allgather_split catalog entries are certified on sm_100 only (their contracts carry sm_100 receipts), but their collected op-matrix tests had no SM guard, so a multi-GPU list that collects the modeling_v2 comm directory (l0_gb300_multi_gpus) would run them on sm_103. They now skip off sm_100 with the module-level guard the Kimi K3 comm entries use. Signed-off-by: Vasanth Sabavat --- .../comm/test_modeling_v2_mnnvl_allgather_split_op_matrix.py | 4 ++++ .../comm/test_modeling_v2_mnnvl_fusion_allreduce_op_matrix.py | 4 ++++ 2 files changed, 8 insertions(+) diff --git a/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_allgather_split_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_allgather_split_op_matrix.py index ae0e6db01dd7..1d8186d613f9 100644 --- a/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_allgather_split_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_allgather_split_op_matrix.py @@ -12,6 +12,10 @@ assert torch.cuda.is_available(), "mnnvl_allgather_split requires CUDA devices" +if torch.cuda.get_device_capability() != (10, 0): + # The entry is certified on sm_100 (GB200) only; see its contract's receipts. + pytest.skip("mnnvl_allgather_split is certified on sm_100 only", allow_module_level=True) + # Each case starts its own W-rank mpirun over the visible devices; under xdist several workers would fight for them. @pytest.mark.no_xdist diff --git a/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_fusion_allreduce_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_fusion_allreduce_op_matrix.py index b0fd79635674..f12370896242 100644 --- a/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_fusion_allreduce_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_mnnvl_fusion_allreduce_op_matrix.py @@ -12,6 +12,10 @@ assert torch.cuda.is_available(), "mnnvl_fusion_allreduce requires CUDA devices" +if torch.cuda.get_device_capability() != (10, 0): + # The entry is certified on sm_100 (GB200) only; see its contract's receipts. + pytest.skip("mnnvl_fusion_allreduce is certified on sm_100 only", allow_module_level=True) + # Each case starts its own W-rank mpirun over the visible devices; under xdist several workers would fight for them. @pytest.mark.no_xdist From 52ca9fa8948a746d7eaae347392a1422ce19c149 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Fri, 2 Oct 2026 23:35:58 -0700 Subject: [PATCH 122/161] [None][test] Kimi K3 MLA decode view test: list it in l0_cpu, not l0_b200 Every test in test_k3_mla_decode_view.py is cpu_only and uses CPU tensors. GPU stages run unittests with -m "not cpu_only" (jenkins/L0_Test.groovy), so its l0_b200 entry collected nothing. The entry moves to l0_cpu.yml, whose stage runs -m cpu_only. Signed-off-by: Vasanth Sabavat --- tests/integration/test_lists/test-db/l0_b200.yml | 1 - tests/integration/test_lists/test-db/l0_cpu.yml | 2 ++ 2 files changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index 14470959f083..28caf920b148 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -433,7 +433,6 @@ l0_b200: - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_kda_verify.py - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_attn.py - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_q.py - - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_decode_view.py - unittest/_torch/modeling_v2/ssm - unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_mla_qkv.py - unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_mla_attn_vb_out.py diff --git a/tests/integration/test_lists/test-db/l0_cpu.yml b/tests/integration/test_lists/test-db/l0_cpu.yml index 2f1a4de55524..521187488659 100644 --- a/tests/integration/test_lists/test-db/l0_cpu.yml +++ b/tests/integration/test_lists/test-db/l0_cpu.yml @@ -39,6 +39,8 @@ l0_cpu: # runs -m cpu_only, so this entry contributes only the 3 marked files under attention/. - unittest/_torch/attention - unittest/_torch/cute_dsl/test_kimi_k3_kda_ptx_patch.py + # Kimi K3 MLA decode view from the attention metadata: every test is cpu_only (CPU tensors). + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mla_decode_view.py - unittest/_torch/compilation/test_auto_multi_stream.py - unittest/_torch/compilation/test_remove_copy_pass.py::test_remove_copy_preserves_minimax_producer_outputs - unittest/_torch/distributed From 543ebb37819367d069bd5f2df4bc9f8507ac4c50 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 06:01:22 -0700 Subject: [PATCH 123/161] [None][test] modeling_v2 Kimi K3 construction and drift tests: list them in l0_b200 The Kimi K3 targets' construction checks and the drift check of tp16_moetp16ep1's copied modules are collected only by l0_b300's unittest/_torch/modeling_v2 entry, and waives.txt skips that entry as a whole (nvbugs/6853741), so no CI stage runs them. Neither file is cpu_only, so B200's "not cpu_only" stage runs them. Signed-off-by: Vasanth Sabavat --- tests/integration/test_lists/test-db/l0_b200.yml | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index 28caf920b148..f338fb67a9f8 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -162,6 +162,10 @@ l0_b200: - unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_gemv.py # The Kimi K3 target's DSpark drafter on the drafter entries, against the stock block forward. - unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py + # The Kimi K3 targets' construction checks and the drift check of tp16_moetp16ep1's copied modules. + # l0_b300's modeling_v2 entry, the only other one that collects them, is waived (nvbugs/6853741). + - unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_construction.py + - unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drift.py # KDA runtime: host-derived prefill metadata and the bf16 state pool # round-trip. Both are single-device cases that use GPU 0 only. - unittest/_torch/modules/kimi_kda/test_kda_host_metadata.py From 871ef643e035329aea69e0ddb09665edd6d853cf Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 06:01:22 -0700 Subject: [PATCH 124/161] [None][test] modeling_v2 Kimi K3 routing and claims tests: list them in l0_b200 The Kimi K3 routing tests and the claims check that a target's generic path declares every stock import are collected only by l0_b300's unittest/_torch/modeling_v2 entry, and waives.txt skips that entry as a whole (nvbugs/6853741), so no CI stage runs them. None of them is cpu_only, so B200's "not cpu_only" stage runs them. l0_b200 selects them by -k and by node id, so the files' other tests stay under the waiver. Signed-off-by: Vasanth Sabavat --- tests/integration/test_lists/test-db/l0_b200.yml | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index f338fb67a9f8..991644a9f59e 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -144,6 +144,10 @@ l0_b200: - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_head_gemv.py - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_cluster_waits.py - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py + # The Kimi K3 targets' routing and the claims check of their stock imports. l0_b300's modeling_v2 + # entry, the only other one that collects them, is waived (nvbugs/6853741). + - unittest/_torch/modeling_v2/test_modeling_v2_routing.py -k "kimi_k3" + - unittest/_torch/modeling_v2/test_modeling_v2_claims.py::test_uncertified_generic_calls_name_every_stock_import # modeling_v2 catalog entries with sm_100 receipts (the K3 ones skip on # other architectures). - unittest/_torch/modeling_v2/norm/test_modeling_v2_k3_embed_norm.py From 8806ae3a6bf0d28fb421a5fa57a637293a64f947 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 07:55:37 -0700 Subject: [PATCH 125/161] [None][chore] modeling_v2 Kimi K3 tp16_moetp16ep1: carry C4 into the copy The tp16_moetp4ep4 target gains its DSpark drafter (K3DSparkDrafter, which _build_draft_model builds) and the drafter's four decode GEMV sites (drafter_qkv, drafter_o, drafter_gate_up and drafter_down, on k3_ctm_gemv_long, k3_ctm_gemv and k3_ctm_gemv_swiglu), with their REQUIRED_TRTLLM_OPS and UNCERTIFIED_GENERIC_CALLS entries. The copy here takes them: decode_gemv.py byte for byte, and the changes to modeling.py outside route B's blocks. This target decodes without speculation: a speculative decoding config fails construction before the one-engine shell would build a drafter, so K3DSparkDrafter and _build_draft_model are never reached here. The module docstring's block keeps route B's paragraph on that in place of tp16_moetp4ep4's and adds that the drafter stays as tp16_moetp4ep4 has it. No decode GEMV site this target takes changes. _width is a site's k except at the SiLU-and-mul site (drafter_down), and the new _run branches run only the ctm and swiglu kernels, which only the drafter sites name. K3DecodeGemvs.create warms every site, so this target now also compiles and runs the four drafter sites once on zero weights at startup, as tp16_moetp4ep4 does; nothing calls them here. Signed-off-by: Vasanth Sabavat --- .../decode_gemv.py | 60 +++- .../modeling.py | 302 +++++++++++++++++- 2 files changed, 345 insertions(+), 17 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_gemv.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_gemv.py index 0e26e3971d0e..33fbbac23ca7 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_gemv.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_gemv.py @@ -4,11 +4,12 @@ * **Per-site GEMVs** (`K3DecodeGemvs.project`): a projection of a decode step runs on the kernel measured fastest at its call site's weight shape (`SITES`): at most `MAX_ROWS` rows on `gemm/k3_decode_gemv`, - `gemm/k3_ctm_gemv_wide` or `gemm/k3_ctm_gemv_long`, and, where the site lists it, more rows (up to `WIDE_ROWS`) - on `gemm/k3_ctm_gemv_wide`. The sites are the decode kernels' fused projections (MLA's [W_a; W_g] with the gate - rows through a sigmoid, KDA's [q | k | v | g | f_a | b]), the attention output projection, the built-in MLA path's - q_a / kv_a, q_b and gate projections, and the MoE decode path's projections (`decode_moe.py`), two of them with - fp32 outputs. + `gemm/k3_ctm_gemv_wide`, `gemm/k3_ctm_gemv_long`, `gemm/k3_ctm_gemv` or `gemm/k3_ctm_gemv_swiglu`, and, where the + site lists it, more rows (up to `WIDE_ROWS`) on `gemm/k3_ctm_gemv_wide`. The sites are the decode kernels' fused + projections (MLA's [W_a; W_g] with the gate rows through a sigmoid, KDA's [q | k | v | g | f_a | b]), the attention + output projection, the built-in MLA path's q_a / kv_a, q_b and gate projections, the MoE decode path's projections + (`decode_moe.py`), two of them with fp32 outputs, and the DSpark drafter's q / k / v, output, gate / up and down + projections (`K3DSparkDrafter`). * **LM head** (`K3LogitsProcessor`): at most `MAX_ROWS` rows of this rank's vocabulary shard on `gemm/k3_head_gemv` over the target's `K3HeadGemvWorkspace`, then the shards gathered (`comm/allgather`) as the stock head gathers them. It is the shell's logits processor, so the speculative worker's target logits and the @@ -41,9 +42,13 @@ from tensorrt_llm._torch._experimental.modeling_v2.catalog.activation.k3_situ_mul import k3_situ_mul from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.allgather import allgather +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv import k3_ctm_gemv from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv_long import ( k3_ctm_gemv_long, ) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv_swiglu import ( + k3_ctm_gemv_swiglu, +) from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.k3_ctm_gemv_wide import ( k3_ctm_gemv_wide, ) @@ -73,10 +78,11 @@ @dataclass(frozen=True) class Site: """A call site's weight shape (this target's per-rank shapes) and its kernels: ``small`` at 1..MAX_ROWS rows - ("decode", "wide" or "long"), and k3_ctm_gemv_wide at MAX_ROWS+1..``wide_rows`` rows where ``wide``. Output columns - from ``sig_col0`` on are stored through a sigmoid; ``out_fp32``: an fp32 output (k3_ctm_gemv_wide only). - ``split`` / ``ring`` / ``push``: k3_ctm_gemv_long's CTAs per 128-row weight tile, weight-ring stages, and whether - the partial sums are pushed to each row's owner.""" + ("decode", "wide", "long", "ctm" for k3_ctm_gemv, or "swiglu" for k3_ctm_gemv_swiglu, whose rows are the + ``[gate | up]`` input of the SiLU-and-mul, 2 ``k`` wide), and k3_ctm_gemv_wide at MAX_ROWS+1..``wide_rows`` rows + where ``wide``. Output columns from ``sig_col0`` on are stored through a sigmoid; ``out_fp32``: an fp32 output + (k3_ctm_gemv_wide only). ``split`` / ``ring`` / ``push``: the CTM kernels' CTAs per 128-row weight tile, + k3_ctm_gemv_long's weight-ring stages, and whether the partial sums are pushed to each row's owner.""" n: int k: int @@ -113,9 +119,21 @@ class Site: "moe_shared_gate_up": Site(768, 7168, "wide", wide=True), "moe_tail": Site(7168, 640, "wide", wide=True, wide_rows=32), "moe_up": Site(7168, 3584, "wide", out_fp32=True), + # The DSpark drafter's projections (K3DSparkDrafter; 6 query heads, 1 KV head of 64, MLP width 896): the fused + # q / k / v, the output projection (row parallel), gate_up [gate 896 | up 896], and the down projection with the + # SiLU-and-mul (row parallel). + "drafter_qkv": Site(512, 7168, "long", split=8, ring=6, push=True), + "drafter_o": Site(7168, 384, "ctm", split=1, push=True), + "drafter_gate_up": Site(1792, 7168, "long", split=8, ring=6, push=True), + "drafter_down": Site(7168, 896, "swiglu", split=2, push=True), } +def _width(spec: Site) -> int: + """The row width a site's kernel reads: ``k``, or the ``[gate | up]`` input of a SiLU-and-mul site, 2 ``k``.""" + return 2 * spec.k if spec.small == "swiglu" else spec.k + + def _wide_tile(rows: int) -> int: from tensorrt_llm._torch.cute_dsl_kernels.k3_ctm_gemv import k3_ctm_gemv_kernel @@ -136,6 +154,14 @@ def _run( return k3_ctm_gemv_wide(x2d, weight, sig_col0=spec.sig_col0, out_fp32=spec.out_fp32) if spec.out_fp32: return None + if kernel == "ctm": + if spec.sig_col0 >= 0 or not _ctm_op.supports(x2d, weight, spec.split): + return None + return k3_ctm_gemv(x2d, weight, trigger_early=True, split=spec.split, push=spec.push) + if kernel == "swiglu": + if spec.sig_col0 >= 0 or not _ctm_op.supports_swiglu(x2d, weight, spec.split): + return None + return k3_ctm_gemv_swiglu(x2d, weight, trigger_early=True, split=spec.split, push=spec.push) # One wave of the GPU's SMs: beyond it the long GEMV loses to the others. sms = torch.cuda.get_device_properties(x2d.device).multi_processor_count if math.ceil(spec.n / 128) * spec.split > sms or not _ctm_op.supports_long( @@ -240,7 +266,7 @@ def create( for rows in (1, 16, 32, 64) if spec.wide else (1,): if rows > spec.wide_rows: continue - state._project(site, weight.new_zeros(rows, spec.k), weight, warm=True) + state._project(site, weight.new_zeros(rows, _width(spec)), weight, warm=True) del weight if "dense_gate_up" in sites: gu = torch.zeros(1, SITES["dense_gate_up"].n, dtype=torch.bfloat16, device=device) @@ -262,25 +288,27 @@ def create( def project(self, site: str, x: torch.Tensor, weight: torch.Tensor) -> Optional[torch.Tensor]: """``x @ weight.T`` (``[..., N]``: bf16, the site's sigmoid columns through the sigmoid; fp32 at an - ``out_fp32`` site) for ``site``'s weight on its decode kernel, or None where none takes the call: more rows - than the site's kernels take, another shape or dtype, or, under capture, a kernel that has not run eagerly. - The caller then runs its GEMM.""" + ``out_fp32`` site; on a SiLU-and-mul site, ``silu(gate) * up`` of ``x = [gate | up]`` first) for ``site``'s + weight on its decode kernel, or None where none takes the call: more rows than the site's kernels take, + another shape or dtype, or, under capture, a kernel that has not run eagerly. The caller then runs its + module.""" return self._project(site, x, weight, warm=False) def _project( self, site: str, x: torch.Tensor, weight: torch.Tensor, warm: bool ) -> Optional[torch.Tensor]: spec = SITES[site] + width = _width(spec) if ( weight.dtype != torch.bfloat16 or tuple(weight.shape) != (spec.n, spec.k) or not weight.is_contiguous() or x.dtype != torch.bfloat16 or x.dim() < 1 - or x.shape[-1] != spec.k + or x.shape[-1] != width ): return None - rows = x.numel() // spec.k + rows = x.numel() // width if 0 < rows <= MAX_ROWS: kernel = spec.small elif spec.wide and MAX_ROWS < rows <= spec.wide_rows: @@ -291,7 +319,7 @@ def _project( capturing = _capturing() if capturing and not warm and key not in self._ran: return None - y = _run(spec, kernel, _dense_rows(x.reshape(rows, spec.k)), weight) + y = _run(spec, kernel, _dense_rows(x.reshape(rows, width)), weight) if y is None: return None if not capturing: diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py index 9bf03a16eb78..75f43fcec212 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py @@ -72,6 +72,8 @@ split. The causal LM is still the stock one-engine shell the built-in model builds on, which without that config builds no drafter and no worker. This target drops `tp16_moetp4ep4`'s KDA verify kernels and the per-token verify states they keep; the text model's speculative hidden-state taps stay as `tp16_moetp4ep4` has them and never fire. +The DSpark drafter (`K3DSparkDrafter`, which `_build_draft_model` builds) also stays as `tp16_moetp4ep4` has it +and is never built here; its decode GEMV sites are warmed with the others and never called. """ # <<< route B @@ -86,6 +88,9 @@ import torch from torch import nn +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.k3_drafter_attn_qknorm import ( + k3_drafter_attn_qknorm, +) from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.k3_mla_attn_vb_out import ( k3_mla_attn_vb_out, ) @@ -105,6 +110,11 @@ # <<< route B from tensorrt_llm._torch.distributed import AllReduce from tensorrt_llm._torch.model_config import ModelConfig +from tensorrt_llm._torch.models.modeling_dflash import DFlashForCausalLM +from tensorrt_llm._torch.models.modeling_dspark import ( + GQADSparkForCausalLM, + draft_is_embedded_in_target, +) from tensorrt_llm._torch.models.modeling_kimi_linear import KimiLinearForCausalLM from tensorrt_llm._torch.models.modeling_speculative import SpecDecOneEngineForCausalLM from tensorrt_llm._torch.models.modeling_utils import DecoderModel, register_auto_model @@ -124,6 +134,7 @@ from tensorrt_llm._torch.moe.fused_moe.interface import MoESchedulerKind from tensorrt_llm._torch.moe.fused_moe.routing import DeepSeekV3MoeRoutingMethod from tensorrt_llm._torch.pyexecutor.breakable_cuda_graph import is_in_breakable_cuda_graph +from tensorrt_llm._torch.pyexecutor.config_utils import is_mla from tensorrt_llm._torch.pyexecutor.kv_cache.mamba_cache_manager import MambaHybridCacheManagerV2 from tensorrt_llm._torch.utils import AuxStreamType from tensorrt_llm.functional import AllReduceStrategy @@ -177,6 +188,10 @@ "k3_moe", "k3_route_quant", "mnnvl_allgather_split", + # The DSpark drafter's decode path (K3DSparkDrafter): its GEMV sites and its block attention. + "k3_ctm_gemv", + "k3_ctm_gemv_swiglu", + "k3_drafter_attn_qknorm", ) #: The engine surface the first forward checks before this target relies on it: per object, the attributes read. @@ -218,6 +233,13 @@ "tensorrt_llm._torch.modules.rms_norm.RMSNorm", "tensorrt_llm._torch.distributed.AllReduce", "tensorrt_llm._torch.modules.multi_stream_utils.maybe_execute_in_parallel", + # The DSpark drafter: the stock GQA drafter K3DSparkDrafter extends (its block forward where the drafter entries + # do not take a block, its context projection and context k / v, its heads and its weight load), and the stock + # builder's checks for which drafter a checkpoint gets. + "tensorrt_llm._torch.models.modeling_dspark.GQADSparkForCausalLM", + "tensorrt_llm._torch.models.modeling_dspark.draft_is_embedded_in_target", + "tensorrt_llm._torch.models.modeling_dflash.DFlashForCausalLM", + "tensorrt_llm._torch.pyexecutor.config_utils.is_mla", ) # The K3 decode kernels' bounds: the token-count kernels take one tile of DECODE_MAX_TOKENS rows; the request-aware @@ -2274,6 +2296,264 @@ def _k3_decode_view(self, attn_metadata: AttentionMetadata, num_tokens: int) -> return view +# ---------------------------------------------------------------------------------------------------------------------- +# The DSpark drafter: the stock GQA drafter, with a decode step's block on the drafter entries. +# ---------------------------------------------------------------------------------------------------------------------- + +#: The (requests, block tokens) splits attention/k3_drafter_attn_qknorm certifies; R x 7 is DSpark's block at +#: max_draft_len 7. +DRAFTER_ATTN_SPLITS = frozenset( + {(1, 1), (1, 8), (2, 4), (4, 2), (8, 1), (3, 1), (2, 8), (8, 8)} | {(r, 7) for r in range(1, 9)} +) +# The drafter layout it certifies: 6 query heads and 1 KV head of 64 per rank, context pages of 64 rows, the q / k +# RMSNorm's epsilon and the NeoX RoPE's base. +_DRAFTER_HEADS = (6, 1) +_DRAFTER_HEAD_DIM = 64 +_DRAFTER_PAGE = 64 +_DRAFTER_EPS = 1e-5 +_DRAFTER_ROPE_BASE = 10000.0 + + +class K3DSparkDrafter(GQADSparkForCausalLM): + """Kimi K3's DSpark drafter: the stock GQA drafter, with a decode step's block forward on the drafter entries. + + A block the entries take runs, per layer: + + * the fused q / k / v projection on the ``drafter_qkv`` decode GEMV site; + * ``attention/k3_drafter_attn_qknorm``: the q / k RMSNorm, the NeoX RoPE, and each request's block attending to + its paged context and to its own k / v, in one launch. The block's k / v are not stored in the cache: the + worker writes the accepted rows' context k / v before a block reads them; + * the output projection on ``drafter_o``, then the module's all-reduce; + * gate / up on ``drafter_gate_up`` and the down projection with the SiLU-and-mul on ``drafter_down``, then the + module's all-reduce. + + The norms and the residual adds are the stock modules'; a projection whose site does not take its rows runs its + module. Every other block runs the stock ``dflash_forward``: a split `DRAFTER_ATTN_SPLITS` does not list, another + attention backend, head layout, RoPE or normalization, a cache the kernel does not read, or a compile key of the + attention that has not run eagerly, under CUDA-graph capture. The worker that calls it (the stock + ``DSparkWorker``), the context projection, the context k / v and the Markov head stay upstream code. + """ + + def __init__( + self, draft_config: ModelConfig, *, dflash_attention_backend: str = "AUTO" + ) -> None: + super().__init__(draft_config, dflash_attention_backend=dflash_attention_backend) + # The decode GEMVs' state (decode_gemv.py), shared by the target's layers; set by the target once built. + self.decode_gemvs: Optional[_decode_gemv.K3DecodeGemvs] = None + # Whether every layer is one the drafter entries take; decided on the first block, the weights loaded. + self._k3_take: Optional[bool] = None + # The attention's compile keys that ran eagerly ((more than one request, page stride)); a capture takes + # only these. + self._k3_attn_ran: set = set() + + def dflash_forward( + self, + noise_embedding: torch.Tensor, + query_positions: torch.Tensor, + num_ctx_per_req: torch.Tensor, + ctx_k_cache: torch.Tensor, + ctx_v_cache: torch.Tensor, + ctx_cache_batch_idx: torch.Tensor, + ctx_kv_cache: Optional[torch.Tensor] = None, + ctx_page_table: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + """The block's hidden states ``[B * block, hidden]``: on the drafter entries where they take it (the class + docstring), else the stock forward.""" + keys = self._k3_block_keys(noise_embedding, ctx_kv_cache, ctx_page_table) + if keys is None: + return super().dflash_forward( + noise_embedding, + query_positions, + num_ctx_per_req, + ctx_k_cache, + ctx_v_cache, + ctx_cache_batch_idx, + ctx_kv_cache, + ctx_page_table, + ) + out = self._k3_block_forward( + noise_embedding, + query_positions, + num_ctx_per_req, + ctx_cache_batch_idx, + ctx_kv_cache, + ctx_page_table, + ) + if not torch.cuda.is_current_stream_capturing(): + self._k3_attn_ran |= keys + return out + + def _k3_layers_take(self) -> bool: + """Whether every layer is one the drafter entries take: plain NeoX RoPE with the q / k RMSNorm at the head + layout, epsilon and RoPE base the attention certifies, bias-free projections, a plain SiLU-and-mul MLP, + non-causal unwindowed block attention, and no block convolution or post-attention gate.""" + if self._k3_take is None: + if self._fused_kv_weight is None: + self._build_fused_kv_buffers() + take = ( + self._use_fused_qk_norm_rope + and not self.has_block_conv + and type(self)._post_attention_gate is DFlashForCausalLM._post_attention_gate + ) + for layer_idx, layer in enumerate(self.model.layers): + attn, mlp = layer.self_attn, layer.mlp + pos = getattr(attn, "pos_embd_params", None) + rope = getattr(pos, "rope", None) + norms = (attn.q_norm, attn.k_norm) + take = take and ( + (attn.num_heads, attn.num_key_value_heads) == _DRAFTER_HEADS + and attn.head_dim == _DRAFTER_HEAD_DIM + and getattr(attn, "is_qk_norm", False) + and not getattr(attn, "use_gemma_rms_norm", True) + and rope is not None + and pos.is_neox + and getattr(pos, "mrope_section", None) is None + and rope.theta == _DRAFTER_ROPE_BASE + and all( + norm.variance_epsilon == _DRAFTER_EPS + and norm.weight.dtype == torch.bfloat16 + and tuple(norm.weight.shape) == (_DRAFTER_HEAD_DIM,) + for norm in norms + ) + and attn.qkv_proj.bias is None + and attn.o_proj.bias is None + and mlp.activation is torch.nn.functional.silu + and mlp.swiglu_limit is None + and mlp.swiglu_alpha in (None, 1.0) + and mlp.swiglu_beta in (None, 0.0) + and mlp.gate_up_proj.bias is None + and mlp.down_proj.bias is None + and self._resolve_block_attention(layer_idx) == (False, (-1, -1)) + ) + self._k3_take = bool(take) + logger.info( + f"Kimi K3 DSpark drafter: {len(self.model.layers)} layers, decode blocks " + + ( + "on k3_drafter_attn_qknorm and the drafter GEMV sites" + if self._k3_take + else "on the stock block forward (a layer the drafter entries do not take)" + ) + ) + return self._k3_take + + def _k3_block_keys( + self, + noise_embedding: torch.Tensor, + ctx_kv_cache: Optional[torch.Tensor], + ctx_page_table: Optional[torch.Tensor], + ) -> Optional[frozenset]: + """The attention's compile keys for this block when the drafter entries take it, else None.""" + if ( + self.dflash_attention_backend != "TRTLLM" + or ctx_kv_cache is None + or ctx_page_table is None + or noise_embedding.dtype != torch.bfloat16 + or tuple(noise_embedding.shape[:2]) not in DRAFTER_ATTN_SPLITS + or is_in_breakable_cuda_graph() + or not self._k3_layers_take() + ): + return None + num_kv_heads = self.model.layers[0].self_attn.num_key_value_heads + page = (2, num_kv_heads, _DRAFTER_PAGE, _DRAFTER_HEAD_DIM) + dense = ( + num_kv_heads * _DRAFTER_PAGE * _DRAFTER_HEAD_DIM, + _DRAFTER_PAGE * _DRAFTER_HEAD_DIM, + ) + strides = set() + for layer_idx in range(len(self.model.layers)): + cache = ctx_kv_cache[layer_idx] + if ( + cache.dtype != torch.bfloat16 + or cache.dim() != 5 + or tuple(cache.shape[1:]) != page + or cache.stride()[1:] != (*dense, _DRAFTER_HEAD_DIM, 1) + or cache.stride(0) % 8 + ): + return None + strides.add(cache.stride(0)) + keys = frozenset((noise_embedding.shape[0] > 1, stride) for stride in strides) + if torch.cuda.is_current_stream_capturing() and not keys <= self._k3_attn_ran: + return None + return keys + + def _k3_block_forward( + self, + noise_embedding: torch.Tensor, + query_positions: torch.Tensor, + num_ctx_per_req: torch.Tensor, + ctx_cache_batch_idx: torch.Tensor, + ctx_kv_cache: torch.Tensor, + ctx_page_table: torch.Tensor, + ) -> torch.Tensor: + batch, block = noise_embedding.shape[:2] + rows = batch * block + ctx_len = num_ctx_per_req[:batch].to(torch.int32) + page_table = ctx_page_table.index_select(0, ctx_cache_batch_idx.to(torch.long)) + positions = query_positions.reshape(-1).contiguous() + hidden = noise_embedding.reshape(rows, -1) + residual = None + for layer_idx, layer in enumerate(self.model.layers): + attn = layer.self_attn + if residual is None: + residual = hidden.clone() + normed = layer.input_layernorm(hidden) + else: + normed, residual = layer.input_layernorm(hidden, residual) + qkv = self._k3_project("drafter_qkv", normed, attn.qkv_proj) + out = torch.empty(rows, attn.q_size, dtype=torch.bfloat16, device=qkv.device) + k3_drafter_attn_qknorm( + qkv, + attn.q_norm.weight, + attn.k_norm.weight, + positions, + attn.q_norm.variance_epsilon, + attn.pos_embd_params.rope.theta, + ctx_kv_cache[layer_idx], + page_table, + ctx_len, + attn.num_heads, + attn.num_key_value_heads, + out, + ) + hidden = self._k3_project("drafter_o", out, attn.o_proj) + hidden, residual = layer.post_attention_layernorm(hidden, residual) + hidden = self._k3_mlp(layer.mlp, hidden) + out, _ = self.model.norm(hidden, residual) + return out + + def _k3_project(self, site: str, x: torch.Tensor, linear: nn.Module) -> torch.Tensor: + """``linear(x)``, its GEMM on the ``site`` decode GEMV where that takes the rows, then a row-parallel + projection's all-reduce.""" + y = None if self.decode_gemvs is None else self.decode_gemvs.project(site, x, linear.weight) + if y is None: + return linear(x) + return _row_parallel_reduce(linear, y) + + def _k3_mlp(self, mlp: nn.Module, x: torch.Tensor) -> torch.Tensor: + """``mlp(x)``: gate / up on ``drafter_gate_up`` and the SiLU-and-mul with the down projection on + ``drafter_down`` where both take the rows, then the down projection's all-reduce.""" + gemvs = self.decode_gemvs + gate_up = ( + None if gemvs is None else gemvs.project("drafter_gate_up", x, mlp.gate_up_proj.weight) + ) + down = ( + None + if gate_up is None + else gemvs.project("drafter_down", gate_up, mlp.down_proj.weight) + ) + if down is None: + return mlp(x) + return _row_parallel_reduce(mlp.down_proj, down) + + +def _row_parallel_reduce(linear: nn.Module, partial: torch.Tensor) -> torch.Tensor: + """``partial`` all-reduced as ``linear`` reduces its output: a row-parallel projection's all-reduce, if any.""" + all_reduce = getattr(linear, "all_reduce", None) + if getattr(getattr(linear, "tp_mode", None), "name", None) == "ROW" and all_reduce is not None: + return all_reduce(partial) + return partial + + # ---------------------------------------------------------------------------------------------------------------------- # The target: step classification, the construction checks and the registration shell. # ---------------------------------------------------------------------------------------------------------------------- @@ -2466,13 +2746,31 @@ def __init__(self, model_config: ModelConfig): model_config.pretrained_config = self.config model_config._frozen = True + def _build_draft_model( + self, model_config: ModelConfig, draft_config: Optional[ModelConfig] + ) -> Optional[nn.Module]: + """DSpark's drafter from a standalone GQA checkpoint is this target's `K3DSparkDrafter`; any other drafter + is the mode registry's.""" + spec_config = model_config.spec_config + if ( + spec_config.spec_dec_mode.is_dspark() + and draft_config is not None + and not draft_is_embedded_in_target(model_config) + and not is_mla(draft_config.pretrained_config) + ): + return K3DSparkDrafter( + draft_config, dflash_attention_backend=spec_config.attention_backend + ) + return super()._build_draft_model(model_config, draft_config) + def load_weights(self, weights, *args, **kwargs): _weights.load(self, weights) def cache_derived_state(self) -> None: """Build the decode GEMVs' state once the weights are final: the LM head's workspace, and one eager call of every decode GEMV kernel at its site's shape (decode_gemv.SITES), so none compiles under capture. Built once: - a later call keeps it, since CUDA graphs captured in between hold its workspace.""" + a later call keeps it, since CUDA graphs captured in between hold its workspace. The DSpark drafter's sites + are among them, and the drafter gets the state too.""" super().cache_derived_state() if self.model.decode_gemvs is not None: return @@ -2482,6 +2780,8 @@ def cache_derived_state(self) -> None: for layer in self.model.layers: if not layer.is_moe: layer.decode_gemvs = gemvs + if isinstance(self.draft_model, K3DSparkDrafter): + self.draft_model.decode_gemvs = gemvs def post_load_weights(self) -> None: """The state the decode kernels share, built once per device before any CUDA-graph capture and handed to From 32a581ab11d2ae5eda96aff29bfbd3f5d8bdf6fc Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 10:29:53 -0700 Subject: [PATCH 126/161] [None][feat] Kimi K3 DSpark decode: k3_spec_accept, k3_ctx_kv and k3_markov kernels Three CuTe DSL kernels for the DSpark speculative decode step and their tests: - trtllm::k3_spec_accept: the step's acceptance against vocabulary- sharded target logits (a cross-rank top-1 exchange instead of the logits all-gather), the block-table decode, the KV lengths and the drafter's inputs in one launch; - trtllm::k3_ctx_kv: every drafter layer's context K / V for the step's accepted rows written into the paged pool in one launch; - trtllm::k3_markov: the vanilla-Markov draft chain over the vocabulary-sharded draft logits, with the per-position greedy exchange across the ranks. The worker does not call them yet. Signed-off-by: Vasanth Sabavat --- .../cute_dsl_kernels/k3_ctx_kv/__init__.py | 19 + .../k3_ctx_kv/k3_ctx_kv_kernel.py | 882 +++++++++++++++ .../_torch/cute_dsl_kernels/k3_ctx_kv/op.py | 216 ++++ .../cute_dsl_kernels/k3_markov/__init__.py | 19 + .../k3_markov/k3_markov_kernel.py | 572 ++++++++++ .../_torch/cute_dsl_kernels/k3_markov/op.py | 306 +++++ .../k3_spec_accept/__init__.py | 19 + .../k3_spec_accept/k3_spec_accept_kernel.py | 586 ++++++++++ .../cute_dsl_kernels/k3_spec_accept/op.py | 309 +++++ .../kimi_k3/test_k3_ctx_kv.py | 650 +++++++++++ .../kimi_k3/test_k3_markov.py | 1006 +++++++++++++++++ .../kimi_k3/test_k3_spec_accept.py | 681 +++++++++++ .../kimi_k3/test_k3_spec_accept_sharded.py | 568 ++++++++++ 13 files changed, 5833 insertions(+) create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_ctx_kv/__init__.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_ctx_kv/k3_ctx_kv_kernel.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_ctx_kv/op.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_markov/__init__.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_markov/k3_markov_kernel.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_markov/op.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_spec_accept/__init__.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_spec_accept/k3_spec_accept_kernel.py create mode 100644 tensorrt_llm/_torch/cute_dsl_kernels/k3_spec_accept/op.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctx_kv.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_markov.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_spec_accept.py create mode 100644 tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_spec_accept_sharded.py diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_ctx_kv/__init__.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_ctx_kv/__init__.py new file mode 100644 index 000000000000..3a7fee9a613e --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_ctx_kv/__init__.py @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""The block drafter's context K/V of a decode step in CuTe DSL (``trtllm::k3_ctx_kv``). + +Importing :mod:`.op` registers the torch op; nothing is imported eagerly here so that +the CuTe DSL dependency stays optional for every other model. +""" diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_ctx_kv/k3_ctx_kv_kernel.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_ctx_kv/k3_ctx_kv_kernel.py new file mode 100644 index 000000000000..0b04527fa29d --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_ctx_kv/k3_ctx_kv_kernel.py @@ -0,0 +1,882 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +# ============================================================================= +# Block drafter context K/V of a decode step -- CTM (prims/cute) kernel +# ============================================================================= +# +# For the N <= 64 context tokens of a step (B <= 8 requests of K1 <= 8; token t = request b = t / K1, candidate +# j = t % K1) and every drafter layer l and K/V head h of this rank: +# kv[t] = x[t] @ W^T W = the layers' stacked K / V projection rows [L0 K | L0 V | L1 K | ...] (bf16) +# K head: k = bf16(kv); k = bf16(k * rsqrt(mean(k^2) + eps)); k = bf16(k * k_norm[l]); k = bf16(NeoX RoPE(k, pos[t])) +# V head: v = bf16(kv) +# a masked token (j >= num_acc[b]) is multiplied by 0.0 (signed zeros, as the Python path), and every token's rows go +# into the paged pool (HND per layer: page, K/V, head, slot, 64; every layer a view of one allocation) at column +# col = clamp(min(ctx_len[slot_b] + j, counts[row_b] page - 1), 0), page = table[row_b, col / page]; +# then ctx_len[slot_b] = min(ctx_len + num_acc[b], max_ctx) and num_ctx[b] = min(new ctx_len, +# max(counts[row_b] page - block, 0)). +# (DFlashWorker: precompute_context_kv, the write mask, _store_context_kv_paged, the ctx_len update and num_ctx.) +# +# Rounding points are the Python path's; the GEMM adds in another order than cuBLAS (fp32 tolerance). +# +# Geometry (k3_ctm_gemv_long): one 128-row tile of W per cluster of SPLIT CTAs; rank r streams the k-tiles r, +# r + SPLIT, ... through a RING-stage ring filled before griddepcontrol.wait (the rest prefetched into L2 then), x +# resident after the wait, tcgen05 M128 x N into TMEM (N = 8, 16, 32 or 64 token columns: the smallest that holds the +# step's tokens; rows past them arrive as zeros). The tokens form chunks of 8; pair p = half h x C + chunk c (half h of +# the tile: 64 rows, one head of one layer, K or V; C = N / 8 chunks) is owned by rank p % SPLIT (C <= SPLIT, so a warp +# pair owns at most one chunk; N = 8: half o by rank o). The other ranks push their fp32 partials of a pair into slot +# [p / SPLIT][source] of its owner's mailbox (st.async, completing the owner's barrier p / SPLIT by bytes); the +# owner's two epilogue warps (TMEM lanes 64 h .. 64 h + 63; thread = head row d, the chunk's 8 token values) add the +# SPLIT partials in rank order and finish the head: the per-token sums of squares by a reduce-scatter warp tree and +# the two warps' halves through shared memory, the RoPE partner row d ^ 32 (the other warp, same lane) through shared +# memory, one bf16 store per (token, row): 128 contiguous bytes per token and head. +# +# ctx_len / num_ctx: every CTA reads what it needs of ctx_len (its epilogue threads, after the grid dependency), joins +# its epilogue threads and arrives on a counter (atom.acq_rel.gpu); the last arrival writes ctx_len and num_ctx and +# re-arms the counter to 0 (no CTA waits on it). +# +# Warps: 0 weight TMA (+ early dependent trigger), 1 x TMA after the grid dependency, 2 TMEM allocation + MMA, 3 idle, +# 4-7 epilogue (the token metadata after the grid dependency, while the MMAs run). +# ============================================================================= +"""CTM block-drafter context K/V: the stacked K/V projection, k_norm, NeoX RoPE, the write mask and the paged store +of every drafter layer for a decode step's context tokens, and the context-length update, in one launch.""" + +from __future__ import annotations + +import cuda.bindings.driver as cuda_driver +import cutlass +import cutlass.cute as cute +import cutlass.experimental.cuda as cuda +from cutlass._mlir.dialects import llvm as _llvm +from cutlass._mlir.extras import types as _T +from cutlass.cutlass_dsl import dsl_user_op +from cutlass.experimental import primitives as prims + +CTA_M = 128 # weight rows per tile = the MMA's M +HEAD = 64 # rows of one K or V head (the head dimension) +ROPE_HALF = HEAD // 2 +CHUNK = 8 # tokens per owned chunk: the token values an epilogue thread finishes +MAX_TOKENS = 64 +MAX_BATCH = 8 +CTA_K = 128 # one k-tile: two 64-element halves of the 128-byte swizzle +MMA_K = 16 +TMA_K_BOX = 64 +TMA_COPY_ITERS = CTA_K // TMA_K_BOX +K_BLOCKS_PER_HALF = TMA_K_BOX // MMA_K +THREADS = 256 +EPI_THREADS = 128 +OWNER_THREADS = 64 +ELEM_BYTES = 2 +EVICT_FIRST = 0x12F0000000000000 # createpolicy.fractional.L2::evict_first, fraction 1.0 (sm_100) +SMEM_BUDGET = 220 * 1024 +BAR_JOIN = 1 # the 128 epilogue threads (TMEM reads done, ctx_len reads done) +BAR_OWNER = 2 # the owning half's two epilogue warps (+ half when a rank owns a chunk of each half) + +# Shared-memory descriptor strides for the 128-byte swizzle, in 16-byte units. +LEADING = 16 +STRIDE = 8 * TMA_K_BOX * ELEM_BYTES +A_HALF_ELEMS = CTA_M * TMA_K_BOX +STEP = (MMA_K * ELEM_BYTES) >> 4 +A_BOX = A_HALF_ELEMS >> 3 +STAGE_A = (CTA_M * CTA_K * ELEM_BYTES) >> 4 + +io_dtype = cutlass.BFloat16 + + +def mma_n(n_tokens: int) -> int: + """The MMA's token columns: the smallest of 8, 16, 32, 64 that holds the step's tokens.""" + for n in (8, 16, 32, 64): + if n_tokens <= n: + return n + return 0 + + +def owner_slots(n_tokens: int, split: int) -> int: + """Pairs (half, chunk) a rank may own: 1 when the 2 C pairs land on distinct ranks, else 2 (one of each half).""" + return 1 if 2 * (mma_n(n_tokens) // CHUNK) <= split else 2 + + +def smem_bytes(split: int, ring: int, my_tiles: int, n_tokens: int = CHUNK) -> int: + """Shared memory of a CTA: the weight ring, the resident x k-tiles, the mailbox, the exchanges and barriers.""" + slots = owner_slots(n_tokens, split) + return ( + ring * CTA_M * CTA_K * ELEM_BYTES + + my_tiles * mma_n(n_tokens) * CTA_K * ELEM_BYTES + + slots * split * HEAD * CHUNK * 4 + + slots * HEAD * CHUNK * 4 + + slots * 2 * CHUNK * 4 + + 1024 + ) + + +def pick_ring(k_in: int, split: int, n_tokens: int = CHUNK) -> int: + """The deepest weight ring that fits next to the resident x, the mailbox and the exchanges (0: none).""" + my_tiles = (k_in // CTA_K) // split + for ring in range(min(my_tiles, 8), 0, -1): + if smem_bytes(split, ring, my_tiles, n_tokens) <= SMEM_BUDGET: + return ring + return 0 + + +def supports(n_rows: int, k_in: int, split: int, nkv: int, k1: int, n_tokens: int) -> bool: + """Shapes the kernel runs: whole 128-row tiles of whole heads, whole k-tiles split evenly over a 2-8 CTA cluster, + a ring that fits, N <= 64 tokens made of B <= 8 whole requests of K1 <= 8, at most one chunk of 8 tokens per + rank and half (N / 8 <= SPLIT).""" + k_tiles = k_in // CTA_K + return ( + nkv > 0 + and n_rows % (2 * nkv * HEAD) == 0 + and n_rows % CTA_M == 0 + and k_in % CTA_K == 0 + and split in (2, 4, 8) + and k_tiles % split == 0 + and 0 < k1 <= CHUNK + and 0 < n_tokens <= MAX_TOKENS + and n_tokens % k1 == 0 + and n_tokens // k1 <= MAX_BATCH + and mma_n(n_tokens) // CHUNK <= split + and pick_ring(k_in, split, n_tokens) > 0 + ) + + +def _plus(base, x): + """``base + x``; ``x`` itself when ``base`` is a Python 0 (the one-chunk build: no op in the IR).""" + if type(base) is int and base == 0: + return x + return base + x + + +def _slot(barriers, i): + """Barrier ``i`` of ``barriers`` (the array itself for a Python 0).""" + if type(i) is int and i == 0: + return barriers + return barriers.subview(i) + + +def _push_chunk(mailbox, mail_full, acc, c, half, chunks, split, slots_owned, rank, slot_row): + """st.async of this lane's partials of chunk ``c`` of its half (8 token values) into slot [pair / split][rank] of + the pair's owner, completing the owner's barrier pair / split by 32 bytes.""" + if chunks == 1: + dest = half + slot_grp = 0 + else: + pair = half * cutlass.Int32(chunks) + cutlass.Int32(c) + dest = pair % cutlass.Int32(split) + slot_grp = 0 if slots_owned == 1 else pair // cutlass.Int32(split) + base = _plus(slot_grp * split, rank) * cutlass.Int32(HEAD * CHUNK) + slot_row + peer_slot = _mapa_u32(mailbox.subview(base).data_ptr(), dest) + mbar_peer = _mapa_u32(_slot(mail_full, slot_grp).data_ptr(), dest) + _st_async_v4( + peer_slot, + acc[CHUNK * c], + acc[CHUNK * c + 1], + acc[CHUNK * c + 2], + acc[CHUNK * c + 3], + mbar_peer, + ) + _st_async_v4( + peer_slot + cutlass.Int32(16), acc[CHUNK * c + 4], acc[CHUNK * c + 5], acc[CHUNK * c + 6], + acc[CHUNK * c + 7], mbar_peer, + ) # fmt: skip + + +def _own(acc, t, mine): + """The rank's own partial of token ``t``: ``mine`` when chosen among chunks, else ``acc[t]`` (read in place).""" + return cutlass.Float32(acc[t]) if mine is None else mine + + +def _bf16_rn(v): + """fp32 -> bf16 precision (round to nearest even) in the integer domain, kept as fp32 (no fptrunc/fpext pair the + compiler could fold into the next multiply).""" + u = v.bitcast(cutlass.Int32) + u = u + (((u >> 16) & 1) + 0x7FFF) + u = (u >> 16) << 16 + return u.bitcast(cutlass.Float32) + + +@dsl_user_op +def _atomic_add_acq_rel(addr, value, *, loc=None, ip=None): + """atom.add.acq_rel.gpu.s32 on a global address; returns the old value.""" + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [cutlass.Int64(addr).ir_value(loc=loc, ip=ip), value.ir_value(loc=loc, ip=ip)], + "atom.add.acq_rel.gpu.s32 $0, [$1], $2;", "=r,l,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _mapa_u32(smem_ptr, peer, *, loc=None, ip=None): + """The shared::cluster address of this CTA's shared-memory location in cluster CTA ``peer``.""" + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [smem_ptr.toint(loc=loc, ip=ip).ir_value(), cutlass.Int32(peer).ir_value(loc=loc, ip=ip)], + "mapa.shared::cluster.u32 $0, $1, $2;", "=r,r,r", has_side_effects=False, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _test_wait_cluster(mbar, parity, *, loc=None, ip=None): + """mbarrier.test_wait.parity.acquire.cluster on a barrier of this CTA (``mbar``, its shared-memory pointer): whether + phase ``parity`` has completed, acquiring at cluster scope. For barriers that other CTAs' st.async complete (their + complete_tx releases at cluster scope; a CTA-scope acquire does not synchronize with it).""" + done = cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [mbar.toint(loc=loc, ip=ip).ir_value(), cutlass.Int32(parity).ir_value(loc=loc, ip=ip)], + "{\n\t.reg .pred p;\n\tmbarrier.test_wait.parity.acquire.cluster.shared::cta.b64 p, [$1], $2;\n\t" + "selp.u32 $0, 1, 0, p;\n\t}", "=r,r,r,~{memory}", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + return done != cutlass.Int32(0) + + +@dsl_user_op +def _st_async_v4(dst, a, b, c, d, mbar, *, loc=None, ip=None): + """st.async of four fp32 (16 bytes) to a shared::cluster address, completing ``mbar`` (a shared::cluster address) + by 16 bytes.""" + _llvm.inline_asm( + None, + [cutlass.Int32(dst).ir_value(loc=loc, ip=ip), cutlass.Float32(a).ir_value(loc=loc, ip=ip), + cutlass.Float32(b).ir_value(loc=loc, ip=ip), cutlass.Float32(c).ir_value(loc=loc, ip=ip), + cutlass.Float32(d).ir_value(loc=loc, ip=ip), cutlass.Int32(mbar).ir_value(loc=loc, ip=ip)], + "st.async.shared::cluster.mbarrier::complete_tx::bytes.v4.f32 [$0], {$1, $2, $3, $4}, [$5];", "r,f,f,f,f,r", + has_side_effects=True, is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) # fmt: skip + + +def _warp_token_sums(vals, lane): + """Per-token sums over the warp's 32 lanes of 8 per-lane values (a reduce-scatter over lane distances 16, 8, 4, + then the full tree at 2 and 1: 9 shuffles). Lanes 4 i .. 4 i + 3 end with token i's sum.""" + parts = vals + for offset in [16, 8, 4]: + upper = (lane & cutlass.Int32(offset)) != cutlass.Int32(0) + half = len(parts) // 2 + kept = [] + for i in range(half): + send = cutlass.Float32(cutlass.select_(upper, parts[i], parts[half + i])) + keep = cutlass.Float32(cutlass.select_(upper, parts[half + i], parts[i])) + kept.append( + keep + cute.arch.shuffle_sync_bfly(send, offset=offset, mask=-1, mask_and_clamp=31) + ) + parts = kept + total = parts[0] + for offset in [2, 1]: + total = total + cute.arch.shuffle_sync_bfly( + total, offset=offset, mask=-1, mask_and_clamp=31 + ) + return total + + +@cute.kernel +def k3_ctx_kv_kernel( + tma_desc_w: cutlass.GridConstant[ + cuda.TensorMap + ], # W [n_rows, K] bf16, 5-D, one call per k-tile + tma_desc_x: cutlass.GridConstant[cuda.TensorMap], # x [N, K] bf16, box 64 x MMA N + k_norm: cutlass.Array, # bf16 [L * 64] + cos_sin: cutlass.Array, # fp32 [max_pos * 64]: cos of the 32 pairs, then sin + cpos: cutlass.Array, # int64 [N]: RoPE positions + num_acc: cutlass.Array, # int32 [B] + ctx_len: cutlass.Array, # int64 [slots], updated + slots: cutlass.Array, # int64 [B] + rows: cutlass.Array, # int64 [B]: rows of the block table + table: cutlass.Array, # int32 [table rows * table_stride] + counts: cutlass.Array, # int64 [table rows] + pool: cutlass.Array, # bf16 elements of the allocation holding every layer's pool + layer_off: cutlass.Array, # int64 [L]: element offset of each layer's pool in it + num_ctx: cutlass.Array, # int32 [B], out + counter: cutlass.Array, # int32 [1], zero between launches + eps: cutlass.Float32, + max_ctx: cutlass.Int64, + page: cutlass.Int64, + block_size: cutlass.Int64, + table_stride: cutlass.Int64, + page_stride: cutlass.Int64, + kv_stride: cutlass.Int64, + head_stride: cutlass.Int64, + n_rows: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + split: cutlass.Constexpr[int], + ring: cutlass.Constexpr[int], + nkv: cutlass.Constexpr[int], + k1: cutlass.Constexpr[int], + n_tokens: cutlass.Constexpr[int], + trigger_early: cutlass.Constexpr[bool], +): + k_tiles = k_in // CTA_K + my_tiles = k_tiles // split + grid = (n_rows // CTA_M) * split + batch = n_tokens // k1 + n_cols = mma_n(n_tokens) + chunks = n_cols // CHUNK + slots_owned = owner_slots(n_tokens, split) + tmem_cols = max(32, n_cols) + stage_b = (n_cols * CTA_K * ELEM_BYTES) >> 4 + b_half_elems = n_cols * TMA_K_BOX + tx, _, _ = cute.arch.thread_idx() + bx, _, _ = cute.arch.block_idx() + warp_id = cute.arch.warp_idx() + rank = cute.arch.block_idx_in_cluster() + tile = bx // cutlass.Int32(split) + m_offset = tile * cutlass.Int32(CTA_M) + tma_ptr_w = tma_desc_w.get_ptr() + tma_ptr_x = tma_desc_x.get_ptr() + + smem_a = cutlass.Array( + io_dtype, ring * CTA_M * CTA_K, space=cutlass.AddressSpace.smem, alignment=1024 + ) + smem_b = cutlass.Array( + io_dtype, my_tiles * n_cols * CTA_K, space=cutlass.AddressSpace.smem, alignment=1024 + ) + tma_full = cutlass.Array(cutlass.Int64, ring, space=cutlass.AddressSpace.smem, alignment=8) + mma_done = cutlass.Array(cutlass.Int64, ring, space=cutlass.AddressSpace.smem, alignment=8) + act_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + acc_done = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + mail_full = cutlass.Array( + cutlass.Int64, slots_owned, space=cutlass.AddressSpace.smem, alignment=8 + ) + tmem_ptr_i32 = cutlass.Array(cutlass.Int32, 1, space=cutlass.AddressSpace.smem) + # [pair p / split][source rank][head row][token of the chunk] fp32 partials of the pairs this CTA owns (the own + # rank's slot stays unused). + mailbox = cutlass.Array( + cutlass.Float32, + slots_owned * split * HEAD * CHUNK, + space=cutlass.AddressSpace.smem, + alignment=16, + ) + # Per owning half (one buffer unless a rank owns a chunk of each): [owner warp][token]: the warp's sums of squares; + # [head row][token]: the RoPE partners. + s_ss = cutlass.Array( + cutlass.Float32, slots_owned * 2 * CHUNK, space=cutlass.AddressSpace.smem, alignment=16 + ) + s_rope = cutlass.Array( + cutlass.Float32, slots_owned * HEAD * CHUNK, space=cutlass.AddressSpace.smem, alignment=16 + ) + + if warp_id == 0: + prims.prefetch_tensormap(tma_ptr_w) + prims.prefetch_tensormap(tma_ptr_x) + if prims.elect_sync(): + for s in cutlass.range_constexpr(ring): + prims.mbarrier_init(tma_full.subview(s), 1) + prims.mbarrier_init(mma_done.subview(s), 1) + prims.mbarrier_init(act_full, 1) + prims.mbarrier_init(acc_done, 1) + # Owner ranks: the other ranks' partials of an owned pair arrive by st.async (16-byte stores that + # complete its barrier's transaction count); expected here, before cluster formation. + for s in cutlass.range_constexpr(slots_owned): + prims.mbarrier_init(_slot(mail_full, s), 1) + prims.mbarrier_arrive_expect_tx(_slot(mail_full, s), (split - 1) * HEAD * CHUNK * 4) + if warp_id == 2: + prims.tcgen05_alloc(tmem_ptr_i32, tmem_cols) + prims.tcgen05_relinquish_alloc_permit() + prims.fence_mbarrier_init() + # Cluster formation: the peers' shared memory and barriers are addressable from here on. + prims.barrier_cluster_arrive_relaxed() + prims.barrier_cluster_wait() + prims.barrier_cta_sync(0) + tmem_ptr = cutlass.inttoptr(tmem_ptr_i32.load(), 6, cutlass.Int32) + + if warp_id == 0: + # ===================================================================== + # Weight TMA: fill the ring and prefetch the rest into L2 before the + # grid dependency; then refill each stage when its MMAs are done. + # ===================================================================== + if prims.elect_sync(): + for i in cutlass.range_constexpr(ring): + k = rank + cutlass.Int32(i * split) + prims.mbarrier_arrive_expect_tx(tma_full.subview(i), CTA_M * CTA_K * ELEM_BYTES) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a.subview(i * CTA_M * CTA_K), + tma_ptr_w, + ( + cutlass.Int32(0), + m_offset, + k * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), + cutlass.Int32(0), + ), + tma_full.subview(i), + l2_cache_hint=EVICT_FIRST, + ) + for i in cutlass.range_constexpr(ring, my_tiles): + k = rank + cutlass.Int32(i * split) + prims.cp_async_bulk_tensor_prefetch( + tma_ptr_w, + [ + cutlass.Int32(0), + m_offset, + k * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), + cutlass.Int32(0), + ], + [], # tile mode: no im2col offsets + ) + if cutlass.const_expr(trigger_early): + # Dependents may launch now; they wait for this whole grid before reading the pool or ctx_len. + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + if prims.elect_sync(): + stage = cutlass.Int32(0) + phase = cutlass.Int32(0) + for i in cutlass.range(ring, my_tiles, unroll=1): + while not cute.arch.mbarrier_try_wait(mma_done.subview(stage).data_ptr(), phase): + pass + k = rank + i * cutlass.Int32(split) + prims.mbarrier_arrive_expect_tx(tma_full.subview(stage), CTA_M * CTA_K * ELEM_BYTES) + prims.cp_async_bulk_tensor_shared_cta_global( + smem_a.subview(stage * cutlass.Int32(CTA_M * CTA_K)), + tma_ptr_w, + ( + cutlass.Int32(0), + m_offset, + k * cutlass.Int32(TMA_COPY_ITERS), + cutlass.Int32(0), + cutlass.Int32(0), + ), + tma_full.subview(stage), + l2_cache_hint=EVICT_FIRST, + ) + stage = stage + cutlass.Int32(1) + if stage == cutlass.Int32(ring): + stage = cutlass.Int32(0) + phase = phase ^ cutlass.Int32(1) + elif warp_id == 1: + # ===================================================================== + # Activation TMA: the rank's k-tiles of x, resident, after the wait. + # ===================================================================== + prims.griddepcontrol(prims.GridDepAction.WAIT) + if prims.elect_sync(): + prims.mbarrier_arrive_expect_tx(act_full, my_tiles * n_cols * CTA_K * ELEM_BYTES) + for i in cutlass.range_constexpr(my_tiles): + k = rank + cutlass.Int32(i * split) + for half in cutlass.range_constexpr(TMA_COPY_ITERS): + prims.cp_async_bulk_tensor_shared_cta_global( + smem_b.subview(i * n_cols * CTA_K + half * b_half_elems), + tma_ptr_x, + ( + k * cutlass.Int32(CTA_K) + cutlass.Int32(half * TMA_K_BOX), + cutlass.Int32(0), + ), + act_full, + ) + elif warp_id == 2: + # ===================================================================== + # MMA: the rank's k-tiles in order, one TMEM accumulator (partial sum). + # ===================================================================== + idesc = prims.Tcgen05InstrDesc.build( + c_dtype=cutlass.Float32, a_dtype=io_dtype, b_dtype=io_dtype, n_dim=n_cols, m_dim=CTA_M + ) + desc_a_base = prims.Tcgen05SmemDesc.build( + start_address=smem_a, leading_byte_offset=LEADING, stride_byte_offset=STRIDE, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + desc_b_base = prims.Tcgen05SmemDesc.build( + start_address=smem_b, leading_byte_offset=LEADING, stride_byte_offset=STRIDE, + layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B, + ) # fmt: skip + while not cute.arch.mbarrier_try_wait(act_full.data_ptr(), 0): + pass + stage = cutlass.Int32(0) + phase = cutlass.Int32(0) + for i in cutlass.range(my_tiles, unroll=1): + while not cute.arch.mbarrier_try_wait(tma_full.subview(stage).data_ptr(), phase): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + for kb in cutlass.range_constexpr(TMA_COPY_ITERS * K_BLOCKS_PER_HALF): + box = kb // K_BLOCKS_PER_HALF + within = kb % K_BLOCKS_PER_HALF + desc_a = desc_a_base + ( + stage * cutlass.Int32(STAGE_A) + cutlass.Int32(box * A_BOX + within * STEP) + ) + desc_b = desc_b_base + ( + i * cutlass.Int32(stage_b) + + cutlass.Int32(box * (b_half_elems >> 3) + within * STEP) + ) + accumulate = cutlass.Boolean(True) + if cutlass.const_expr(kb == 0): + accumulate = i > cutlass.Int32(0) + if prims.elect_sync(): + prims.tcgen05_mma( + prims.Tcgen05MMAKind.F16, + prims.CTAGroup.CTA_1, + tmem_ptr, + desc_a, + desc_b, + idesc, + accumulate, + ) + if prims.elect_sync(): + prims.tcgen05_commit(mma_done.subview(stage)) + stage = stage + cutlass.Int32(1) + if stage == cutlass.Int32(ring): + stage = cutlass.Int32(0) + phase = phase ^ cutlass.Int32(1) + if prims.elect_sync(): + prims.tcgen05_commit(acc_done) + elif warp_id >= 4: + # ===================================================================== + # Epilogue: token metadata, TMEM -> registers, push to / reduce at the + # pair's owner, the head's epilogue and the paged store. + # ===================================================================== + lane = tx % cutlass.Int32(32) + w = warp_id - cutlass.Int32(4) # TMEM lanes 32 w .. 32 w + 31: tile rows 32 w + lane + half = w // cutlass.Int32(2) + d = (w % cutlass.Int32(2)) * cutlass.Int32(32) + lane # row within the head + # The half's 64-row block of W: (layer, K or V, head). + q = tile * cutlass.Int32(2) + half + layer = q // cutlass.Int32(2 * nkv) + kv = (q % cutlass.Int32(2 * nkv)) // cutlass.Int32(nkv) + head = q % cutlass.Int32(nkv) + row_base = ( + cutlass.Int64(layer_off.load(idx=layer)) + cutlass.Int64(kv) * kv_stride + cutlass.Int64(head) * head_stride + + cutlass.Int64(d) + ) # fmt: skip + k_weight = cutlass.BFloat16(k_norm.load(idx=layer * cutlass.Int32(HEAD) + d)).to( + cutlass.Float32 + ) + # The chunk c of this half that this rank owns (pair half C + c on rank (half C + c) % split), its mailbox + # slot and its exchange buffer (the half's own when a rank owns a chunk of each half). + if cutlass.const_expr(chunks == 1): + c_own = 0 # half h is owned by rank h + else: + c_own = ( + rank + cutlass.Int32(2 * split) - half * cutlass.Int32(chunks) + ) % cutlass.Int32(split) # >= chunks: this rank owns none of the half's chunks + if cutlass.const_expr(slots_owned == 1): + own_slot = buf = 0 # Python zeros: no slot / buffer arithmetic (see _plus) + owner_bar = BAR_OWNER + else: + own_slot = (half * cutlass.Int32(chunks) + c_own) // cutlass.Int32(split) + buf = half + owner_bar = cutlass.Int32(BAR_OWNER) + half + + # Token metadata (predecessor outputs: after the grid dependency), while the MMAs run: each owned token's pool + # element offset, its mask factor and its RoPE cos / sin; per request the slot, length, accepted count and + # allocation (for the last CTA's update). + prims.griddepcontrol(prims.GridDepAction.WAIT) + dst = [] + keep = [] + cos_t = [] + sin_t = [] + req_slot = [] + req_len = [] + req_acc = [] + req_cap = [] + pair = d & cutlass.Int32(ROPE_HALF - 1) + for b in cutlass.range_constexpr(batch): + slot = cutlass.Int32(slots.load(idx=b)) + trow = cutlass.Int32(rows.load(idx=b)) + c = cutlass.Int64(ctx_len.load(idx=slot)) + cap = cutlass.Int64(counts.load(idx=trow)) * page + n_acc = cutlass.Int32(num_acc.load(idx=b)) + req_slot.append(slot) + req_len.append(c) + req_acc.append(n_acc) + req_cap.append(cap) + if cutlass.const_expr(chunks == 1): + for j in cutlass.range_constexpr(k1): + col = c + cutlass.Int64(j) + col = cutlass.Int64( + cutlass.select_(col > cap - cutlass.Int64(1), cap - cutlass.Int64(1), col) + ) + col = cutlass.Int64( + cutlass.select_(col < cutlass.Int64(0), cutlass.Int64(0), col) + ) + pg = cutlass.Int64( + table.load(idx=cutlass.Int64(trow) * table_stride + col // page) + ) + dst.append(pg * page_stride + (col % page) * cutlass.Int64(HEAD) + row_base) + keep.append(cutlass.Float32(cutlass.select_(cutlass.Int32(j) < n_acc, cutlass.Float32(1.0), + cutlass.Float32(0.0)))) # fmt: skip + pos = cutlass.Int64(cpos.load(idx=b * k1 + j)) + cos_t.append( + cutlass.Float32( + cos_sin.load(idx=pos * cutlass.Int64(HEAD) + cutlass.Int64(pair)) + ) + ) + sin_t.append( + cutlass.Float32( + cos_sin.load( + idx=pos * cutlass.Int64(HEAD) + + cutlass.Int64(pair + cutlass.Int32(ROPE_HALF)) + ) + ) + ) + if cutlass.const_expr(chunks > 1): + # The owned chunk's tokens CHUNK c_own + t (a token past the step's reads token 0's, never stored). + for t in cutlass.range_constexpr(CHUNK): + tok = c_own * cutlass.Int32(CHUNK) + cutlass.Int32(t) + tok = cutlass.Int32( + cutlass.select_(tok < cutlass.Int32(n_tokens), tok, cutlass.Int32(0)) + ) + b = tok // cutlass.Int32(k1) + j = tok - b * cutlass.Int32(k1) + slot = cutlass.Int32(slots.load(idx=b)) + trow = cutlass.Int32(rows.load(idx=b)) + c = cutlass.Int64(ctx_len.load(idx=slot)) + cap = cutlass.Int64(counts.load(idx=trow)) * page + n_acc = cutlass.Int32(num_acc.load(idx=b)) + col = c + cutlass.Int64(j) + col = cutlass.Int64( + cutlass.select_(col > cap - cutlass.Int64(1), cap - cutlass.Int64(1), col) + ) + col = cutlass.Int64(cutlass.select_(col < cutlass.Int64(0), cutlass.Int64(0), col)) + pg = cutlass.Int64(table.load(idx=cutlass.Int64(trow) * table_stride + col // page)) + dst.append(pg * page_stride + (col % page) * cutlass.Int64(HEAD) + row_base) + keep.append(cutlass.Float32(cutlass.select_(j < n_acc, cutlass.Float32(1.0), + cutlass.Float32(0.0)))) # fmt: skip + pos = cutlass.Int64(cpos.load(idx=tok)) + cos_t.append( + cutlass.Float32( + cos_sin.load(idx=pos * cutlass.Int64(HEAD) + cutlass.Int64(pair)) + ) + ) + sin_t.append( + cutlass.Float32( + cos_sin.load( + idx=pos * cutlass.Int64(HEAD) + + cutlass.Int64(pair + cutlass.Int32(ROPE_HALF)) + ) + ) + ) + + while not cute.arch.mbarrier_try_wait(acc_done.data_ptr(), 0): + pass + prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) + acc = prims.tcgen05_ld( + "32x32b", cutlass.inttoptr(tmem_ptr_i32.load(), 6, cutlass.Float32), num=n_cols + ) + prims.tcgen05_wait(prims.Tcgen05Wait.LOAD) + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + # Every epilogue thread has read TMEM and ctx_len. + prims.barrier_cta_sync(BAR_JOIN, thread_count=EPI_THREADS) + if warp_id == 4: + prims.tcgen05_dealloc(tmem_ptr, tmem_cols) + slot_row = d * cutlass.Int32(CHUNK) + # Push every chunk of this half that another rank owns: two 16-byte st.async per lane into slot + # [pair / split][rank]; the owner's barrier completes when all (split - 1) x 64 lanes' land. A rank that owns + # none of the half's chunks only pushes; an owner pushes the others, then finishes its own. + if cutlass.const_expr(chunks == 1): + push_only = rank != half + else: + push_only = c_own >= cutlass.Int32(chunks) + if push_only: + for ch in cutlass.range_constexpr(chunks): + _push_chunk( + mailbox, mail_full, acc, ch, half, chunks, split, slots_owned, rank, slot_row + ) + else: + if cutlass.const_expr(chunks > 1): + for ch in cutlass.range_constexpr(chunks): + pushed = half * cutlass.Int32(chunks) + cutlass.Int32(ch) # the pair index + if pushed % cutlass.Int32(split) != rank: + _push_chunk( + mailbox, + mail_full, + acc, + ch, + half, + chunks, + split, + slots_owned, + rank, + slot_row, + ) + # test_wait spin: a warp suspended in try_wait on a barrier completed by remote st.async wakes late. + while not _test_wait_cluster(_slot(mail_full, own_slot).data_ptr(), 0): + pass + slot_base = _plus(own_slot * (split * HEAD * CHUNK), slot_row) + y = [] + for t in cutlass.range_constexpr(CHUNK): + # This rank's own partial of token t of its chunk (several chunks: picked by c_own). + mine = None + if cutlass.const_expr(chunks > 1): + mine = cutlass.Float32(acc[t]) + for ch in cutlass.range_constexpr(1, chunks): + mine = cutlass.Float32( + cutlass.select_( + c_own == cutlass.Int32(ch), + cutlass.Float32(acc[CHUNK * ch + t]), + mine, + ) + ) + total = cutlass.Float32(0.0) + for src in cutlass.range_constexpr(split): + part = cutlass.Float32( + mailbox.load(idx=cutlass.Int32(src * HEAD * CHUNK + t) + slot_base) + ) + total = total + cutlass.Float32( + cutlass.select_(rank == cutlass.Int32(src), _own(acc, t, mine), part) + ) + y.append(_bf16_rn(total)) + # Tokens of the chunk that the step has (the rest are the zero rows past N): static for one chunk. + stored = [] + if cutlass.const_expr(chunks > 1): + for t in cutlass.range_constexpr(CHUNK): + stored.append( + c_own * cutlass.Int32(CHUNK) + cutlass.Int32(t) < cutlass.Int32(n_tokens) + ) + if kv == cutlass.Int32(0): + # K head: RMSNorm over the 64 rows of each token, k_norm, NeoX RoPE. + token_ss = _warp_token_sums([y[t] * y[t] for t in range(CHUNK)], lane) + ss_base = buf * (2 * CHUNK) + if (lane & cutlass.Int32(3)) == cutlass.Int32(0): + s_ss.store( + token_ss, + idx=_plus( + ss_base, + (w % cutlass.Int32(2)) * cutlass.Int32(CHUNK) + + (lane >> cutlass.Int32(2)), + ), + ) + prims.barrier_cta_sync(owner_bar, thread_count=OWNER_THREADS) + kn = [] + rope_base = buf * (HEAD * CHUNK) + for t in cutlass.range_constexpr(CHUNK): + ss = cutlass.Float32(s_ss.load(idx=_plus(ss_base, t))) + cutlass.Float32( + s_ss.load(idx=_plus(ss_base, CHUNK + t)) + ) + r = cute.math.rsqrt(ss * cutlass.Float32(1.0 / HEAD) + eps, fastmath=True) + kn.append(_bf16_rn(_bf16_rn(y[t] * r) * k_weight)) + s_rope.store(kn[t], idx=_plus(rope_base, slot_row + cutlass.Int32(t))) + prims.barrier_cta_sync(owner_bar, thread_count=OWNER_THREADS) + partner_row = _plus( + rope_base, (d ^ cutlass.Int32(ROPE_HALF)) * cutlass.Int32(CHUNK) + ) + sign = cutlass.Float32(cutlass.select_(d < cutlass.Int32(ROPE_HALF), cutlass.Float32(-1.0), + cutlass.Float32(1.0))) # fmt: skip + if cutlass.const_expr(chunks == 1): + for t in cutlass.range_constexpr(n_tokens): + partner = cutlass.Float32(s_rope.load(idx=partner_row + cutlass.Int32(t))) + out = _bf16_rn(kn[t] * cos_t[t] + sign * partner * sin_t[t]) + pool.store((out * keep[t]).to(io_dtype), idx=dst[t]) + else: + for t in cutlass.range_constexpr(CHUNK): + if stored[t]: + partner = cutlass.Float32( + s_rope.load(idx=partner_row + cutlass.Int32(t)) + ) + out = _bf16_rn(kn[t] * cos_t[t] + sign * partner * sin_t[t]) + pool.store((out * keep[t]).to(io_dtype), idx=dst[t]) + else: + if cutlass.const_expr(chunks == 1): + for t in cutlass.range_constexpr(n_tokens): + pool.store((y[t] * keep[t]).to(io_dtype), idx=dst[t]) + else: + for t in cutlass.range_constexpr(CHUNK): + if stored[t]: + pool.store((y[t] * keep[t]).to(io_dtype), idx=dst[t]) + if tx == cutlass.Int32(4 * 32): + # After this CTA's push or stores (off their critical path); every epilogue thread read ctx_len before the + # join above. The last CTA of the grid updates the lengths. + arrived = _atomic_add_acq_rel(counter.data_ptr(0).toint(), cutlass.Int32(1)) + if arrived == cutlass.Int32(grid - 1): + # From this thread's own metadata (read before its arrival, like every other CTA's). + for b in cutlass.range_constexpr(batch): + grown = req_len[b] + cutlass.Int64(req_acc[b]) + grown = cutlass.Int64(cutlass.select_(grown > max_ctx, max_ctx, grown)) + ctx_len.store(grown, idx=req_slot[b]) + allocated = req_cap[b] - block_size + allocated = cutlass.Int64( + cutlass.select_(allocated < cutlass.Int64(0), cutlass.Int64(0), allocated) + ) + num_ctx.store( + cutlass.Int32(cutlass.select_(grown > allocated, allocated, grown)), idx=b + ) + counter.store(cutlass.Int32(0), idx=0) + + +def _weight_tensor_map(w, n_rows, k_in): + """W as five TMA dimensions (64-element column chunk, row, 64-element chunk index, 1, 1) so one call per k-tile + lands both 128-byte-swizzled halves; strides in 16-byte units.""" + return cuda.create_tensor_map_tiled( + global_address=w.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[TMA_K_BOX, n_rows, k_in // TMA_K_BOX, 1, 1], + global_strides=[ + (k_in * ELEM_BYTES) // 16, + (TMA_K_BOX * ELEM_BYTES) // 16, + (n_rows * k_in * ELEM_BYTES) // 16, + (n_rows * k_in * ELEM_BYTES) // 16, + ], + box_dims=[TMA_K_BOX, CTA_M, TMA_COPY_ITERS, 1, 1], + swizzle=cuda.TensorMapSwizzle.s128b, + ) + + +def _activation_tensor_map(x, cols, num_tokens): + """x [N, cols] (rows dense) as (cols, N) with a box of the MMA's token columns: rows past N arrive as zeros.""" + return cuda.create_tensor_map_tiled( + global_address=x.iterator.toint(), + dtype=cutlass.BFloat16, + global_dims=[cols, num_tokens], + global_strides=[(cols * ELEM_BYTES) // 16], + box_dims=[TMA_K_BOX, mma_n(num_tokens)], + swizzle=cuda.TensorMapSwizzle.s128b, + ) + + +@cute.jit +def k3_ctx_kv( + w: cute.Tensor, # [n_rows, K] bf16, K contiguous + x: cute.Tensor, # [N, K] bf16, K contiguous + k_norm: cute.Tensor, + cos_sin: cute.Tensor, + cpos: cute.Tensor, + num_acc: cute.Tensor, + ctx_len: cute.Tensor, + slots: cute.Tensor, + rows: cute.Tensor, + table: cute.Tensor, + counts: cute.Tensor, + pool: cute.Tensor, + layer_off: cute.Tensor, + num_ctx: cute.Tensor, + counter: cute.Tensor, + eps: cutlass.Float32, + max_ctx: cutlass.Int64, + page: cutlass.Int64, + block_size: cutlass.Int64, + table_stride: cutlass.Int64, + page_stride: cutlass.Int64, + kv_stride: cutlass.Int64, + head_stride: cutlass.Int64, + n_rows: cutlass.Constexpr[int], + k_in: cutlass.Constexpr[int], + split: cutlass.Constexpr[int], + ring: cutlass.Constexpr[int], + nkv: cutlass.Constexpr[int], + k1: cutlass.Constexpr[int], + n_tokens: cutlass.Constexpr[int], + trigger_early: cutlass.Constexpr[bool], + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + tma_desc_w = _weight_tensor_map(w, n_rows, k_in) + tma_desc_x = _activation_tensor_map(x, k_in, n_tokens) + k3_ctx_kv_kernel( + tma_desc_w, tma_desc_x, k_norm, cos_sin, cpos, num_acc, ctx_len, slots, rows, table, counts, pool, layer_off, + num_ctx, counter, eps, max_ctx, page, block_size, table_stride, page_stride, kv_stride, head_stride, + n_rows, k_in, split, ring, nkv, k1, n_tokens, trigger_early, + ).launch( + grid=((n_rows // CTA_M) * split, 1, 1), + block=(THREADS, 1, 1), + cluster=(split, 1, 1), + stream=stream, + use_pdl=use_pdl, + ) # fmt: skip diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_ctx_kv/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_ctx_kv/op.py new file mode 100644 index 000000000000..26c330ba8152 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_ctx_kv/op.py @@ -0,0 +1,216 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Torch op of the block drafter's context K/V (``trtllm::k3_ctx_kv``). + +For the N = B (K + 1) <= 64 context tokens of a decode step (B <= 8 requests, K + 1 <= 8): the stacked K/V projection +of every drafter layer, k_norm, NeoX RoPE, the write mask and the paged store into the drafter's context pool, then +``ctx_len += num_accepted`` (clamped) and the context length each request may advertise, in one launch (see +``k3_ctx_kv_kernel``). Compiled on the first call for its shape (B, K + 1), which must happen outside CUDA-graph +capture. +""" + +from __future__ import annotations + +import os +import threading +from typing import Dict, List + +import torch + +_lock = threading.Lock() +_compiled: Dict[tuple, object] = {} +_counters: Dict[torch.device, torch.Tensor] = {} +_sm_counts: Dict[torch.device, int] = {} + + +def _kernel_module(): + from . import k3_ctx_kv_kernel + + return k3_ctx_kv_kernel + + +def _arg(t: torch.Tensor, align: int = 16): + from cutlass.cute.runtime import from_dlpack + + # detach(): DLPack refuses tensors that require grad (weights are parameters). + return from_dlpack(t.detach(), assumed_align=align).mark_layout_dynamic(leading_dim=t.dim() - 1) + + +def _index_arg(t: torch.Tensor): + """An int index / length / table tensor the kernel reads element by element, declared at its element's alignment: + the drafter passes per-request slices such as ``num_accepted_tokens[num_contexts:]``, which start on any element + boundary.""" + return _arg(t, t.element_size()) + + +def _sm_count(device: torch.device) -> int: + n = _sm_counts.get(device) + if n is None: + n = _sm_counts[device] = torch.cuda.get_device_properties(device).multi_processor_count + return n + + +def pick_split( + n_rows: int, k_in: int, nkv: int, k1: int, n_tokens: int, device: torch.device +) -> int: + """The widest cluster (8, 4, 2) whose grid fits the SMs at once and whose shared memory holds the step's tokens + (0: the shape does not run).""" + kern = _kernel_module() + for split in (8, 4, 2): + if (n_rows // kern.CTA_M) * split <= _sm_count(device) and kern.supports( + n_rows, k_in, split, nkv, k1, n_tokens + ): + return split + return 0 + + +def pool_view(layers: List[torch.Tensor]): + """``(base, layer_off, page_stride, kv_stride, head_stride)`` for per-layer HND views [pages, 2, heads, slot, 64] + with one set of strides and dense rows, or None. The views may be separate tensors (a V2 manager wraps each layer's + base address in its own tensor) of one pool allocation: ``base`` is a one-element view at the lowest layer base and + ``layer_off`` the element offsets of the layers from it; the kernel addresses the pool from there.""" + first = layers[0] + if ( + first.dtype != torch.bfloat16 + or first.dim() != 5 + or first.stride(4) != 1 + or first.stride(3) != first.size(4) + ): + return None + ptrs = [t.data_ptr() for t in layers] + lowest = min(ptrs) + for t, p in zip(layers, ptrs): + if ( + t.stride() != first.stride() + or t.dtype != first.dtype + or t.device != first.device + or (p - lowest) % 2 + ): + return None + anchor = layers[ptrs.index(lowest)] + base = anchor.as_strided((1,), (1,), anchor.storage_offset()) + layer_off = torch.tensor( + [(p - lowest) // 2 for p in ptrs], dtype=torch.int64, device=first.device + ) + return base, layer_off, first.stride(0), first.stride(1), first.stride(2) + + +@torch.library.custom_op("trtllm::k3_ctx_kv", mutates_args=("ctx_len", "pool")) +def k3_ctx_kv( + x: torch.Tensor, + weight: torch.Tensor, + k_norm: torch.Tensor, + cos_sin: torch.Tensor, + cpos: torch.Tensor, + num_acc: torch.Tensor, + ctx_len: torch.Tensor, + slots: torch.Tensor, + rows: torch.Tensor, + table: torch.Tensor, + counts: torch.Tensor, + pool: torch.Tensor, + layer_off: torch.Tensor, + page_stride: int, + kv_stride: int, + head_stride: int, + eps: float, + max_ctx: int, + page: int, + block_size: int, + nkv: int, +) -> torch.Tensor: + """Writes the context K/V of every drafter layer for ``x`` [N = B (K + 1), H] bf16 into the pool addressed from + ``pool`` (bf16; :func:`pool_view`'s base) at ``layer_off`` [L] int64 element offsets (page / K-V / head strides + in elements) and updates ``ctx_len`` [slots] int64 in place; returns ``num_ctx`` [B] int32. ``weight``: + ``_fused_kv_weight`` [L 2 nkv 64, H] bf16; ``k_norm`` [L, 64] bf16; ``cos_sin`` [max_pos, 64] fp32 (cos | sin); + ``cpos`` [B, K + 1] int64 RoPE positions; ``num_acc`` [B] int32; ``slots`` / ``rows`` [B] int64; ``table`` + [rows, width] int32 and ``counts`` [rows] int64 (the draft pool's block table).""" + import cuda.bindings.driver as cuda_driver + + kern = _kernel_module() + batch, k1 = cpos.shape + n_tokens, k_in = x.shape + n_rows = weight.shape[0] + split = pick_split(n_rows, k_in, nkv, k1, n_tokens, x.device) + if ( + split == 0 + or n_tokens != batch * k1 + or x.dtype != torch.bfloat16 + or weight.dtype != torch.bfloat16 + or weight.shape[1] != k_in + or not x.is_contiguous() + or not weight.is_contiguous() + or k_norm.dtype != torch.bfloat16 + or k_norm.numel() * 2 * nkv != n_rows + or cos_sin.dtype != torch.float32 + or cos_sin.shape[-1] != kern.HEAD + or cpos.dtype != torch.int64 + or num_acc.dtype != torch.int32 + or ctx_len.dtype != torch.int64 + or slots.dtype != torch.int64 + or rows.dtype != torch.int64 + or table.dtype != torch.int32 + or counts.dtype != torch.int64 + or pool.dtype != torch.bfloat16 + or layer_off.dtype != torch.int64 + or layer_off.numel() * 2 * nkv * kern.HEAD != n_rows + ): + raise ValueError( + f"k3_ctx_kv: unsupported call: x {tuple(x.shape)} {x.dtype}, weight {tuple(weight.shape)} {weight.dtype}, " + f"nkv {nkv}, cpos {tuple(cpos.shape)} (N = B (K + 1) <= {kern.MAX_TOKENS} tokens, B <= {kern.MAX_BATCH}, " + f"K + 1 <= {kern.CHUNK}, H % {kern.CTA_K} == 0, whole heads of 64, bf16 weights, int64 positions / " + f"lengths / slots / rows, int32 table)" + ) # fmt: skip + device = x.device + counter = _counters.get(device) + if counter is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError("trtllm::k3_ctx_kv must run once outside CUDA-graph capture first") + counter = _counters[device] = torch.zeros(1, dtype=torch.int32, device=device) + num_ctx = torch.empty(batch, dtype=torch.int32, device=device) + ring = kern.pick_ring(k_in, split, n_tokens) + args = ( + _arg(weight), _arg(x), _arg(k_norm.reshape(-1)), _arg(cos_sin.reshape(-1)), _index_arg(cpos.reshape(-1)), + _index_arg(num_acc), _index_arg(ctx_len), _index_arg(slots), _index_arg(rows), _index_arg(table.reshape(-1)), + _index_arg(counts), _arg(pool), _index_arg(layer_off), _arg(num_ctx), _arg(counter), + ) # fmt: skip + scalars = (float(eps), int(max_ctx), int(page), int(block_size), int(table.stride(0)), int(page_stride), + int(kv_stride), int(head_stride)) # fmt: skip + use_pdl = os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" + consts = (n_rows, k_in, split, ring, nkv, k1, n_tokens, True) + stream = cuda_driver.CUstream(torch.cuda.current_stream(device).cuda_stream) + key = consts + (use_pdl,) + fn = _compiled.get(key) + if fn is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "trtllm::k3_ctx_kv must run once per shape outside CUDA-graph capture first" + ) + import cutlass.cute as cute + + with _lock: + fn = _compiled.get(key) + if fn is None: + fn = _compiled[key] = cute.compile( + kern.k3_ctx_kv, *args, *scalars, *consts, use_pdl, stream + ) + fn(*args, *scalars, stream) + return num_ctx + + +@k3_ctx_kv.register_fake +def _(x, weight, k_norm, cos_sin, cpos, num_acc, ctx_len, slots, rows, table, counts, pool, layer_off, page_stride, + kv_stride, head_stride, eps, max_ctx, page, block_size, nkv): # fmt: skip + return ctx_len.new_empty((cpos.shape[0],), dtype=torch.int32) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_markov/__init__.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_markov/__init__.py new file mode 100644 index 000000000000..b51a8c03554f --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_markov/__init__.py @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""DSpark vanilla-Markov draft chain over a vocab-sharded draft head in CuTe DSL (``trtllm::k3_markov``). + +Importing :mod:`.op` registers the torch op; nothing is imported eagerly here so that +the CuTe DSL dependency stays optional for every other model. +""" diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_markov/k3_markov_kernel.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_markov/k3_markov_kernel.py new file mode 100644 index 000000000000..22b9d5fa8864 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_markov/k3_markov_kernel.py @@ -0,0 +1,572 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +# ============================================================================= +# DSpark vanilla-Markov draft chain for a vocab-sharded draft head -- CTM (prims/cute) kernel +# ============================================================================= +# +# For every request b and block position k = 0 .. K-1 (greedy, temperature 0): +# bias[v] = bf16(markov_w2[v] . markov_w1[prev]) v in this rank's vocab shard +# corrected[v] = base[b, k, v] + bias[v] (fp32) +# token[b, k] = the first maximum of corrected over the whole vocabulary (all ranks) +# prev = token[b, k] (prev for k = 0 is the anchor, first_prev[b]) +# so the K positions are a chain of dependent global argmaxes across the tensor-parallel ranks. +# +# Numerics are dsparkMarkovChainKernel's (the C++ kernel this replaces): lane l of a warp sums the 8 columns +# 8l .. 8l+7 of a row by FMAs from 0, the 32 lane sums are added by the xor tree 16, 8, 4, 2, 1, the sum is rounded +# to bf16 (cvt.rn), and added to the fp32 base. Argmax order: larger value, then lower vocabulary index; NaN never +# wins. The order is total, so every split of the reduction gives the same token. +# +# Geometry: G CTAs of 256 threads, CTA c owning the RC = S / G rows [c RC, c RC + RC) of the shard (RC a multiple +# of 64). Warp w owns the RC / 8 contiguous rows [w RC / 8, ...), in blocks of 8 (row j's sum ends on lane 4j after +# the reduce-scatter form of the xor tree). The CTA's rows of markov_w2 are bulk-copied into shared memory before the +# grid-dependency wait (a weight) and stay there for all K positions. markov_w1[prev] (512 bytes per request) is one +# bulk copy per CTA and position into shared memory, issued by the thread that reduces request b's global argmax, on +# one phase of an mbarrier per position. +# +# Exchange (one NVLink hop per position): every CTA reduces its rows to one (value, index) entry per request and +# multicast-stores it into slot [k][b][rank][c] of every rank's Lamport buffer; every CTA then polls all +# W x G entries of [k][b] in its own rank's copy (the data is the ready signal: 0x80000000 = empty; values carry +# -0.0 as +0.0 and indices are non-negative, so no entry is the sentinel) and reduces them with the same order. +# Every CTA of every rank computes the same global argmax: no leader, no grid barrier, no publish word. +# +# Buffers: 3 rotating Lamport buffers. Call n uses buffer n % 3 and re-arms buffer (n - 1) % 3 (its readers, all +# CTAs of call n - 1, finished before this call's grid wait returned; its next writer is a remote rank's call n + 2, +# which needs this rank's call n + 1 pushes). flags: [0] buffer of this call, [1] CTAs arrived, [2] words the +# previous call used (to re-arm), [3] unused. The last CTA to arrive advances them. +# +# Roles: warp 0's elected lane issues the weight bulk copy (and the early dependent trigger); thread b < B pushes +# request b's CTA entry, reduces its global argmax and requests the next markov_w1 row; all 8 warps compute and poll +# (the chain inside a position is serial, and the poll is spread over all 256 threads: one load wave). Per position: +# rows → barrier → push → poll → barrier → (thread b) argmax, next row; no barrier between the argmax and the next +# position (its readers wait for the row's phase). +# ============================================================================= +"""CTM DSpark Markov chain: the draft head's greedy intra-block Markov bias, the per-position global argmax +across the tensor-parallel vocab shards (the draft tokens), next_new_tokens and the target's KV-length rewind after +the draft forward, in one launch.""" + +from __future__ import annotations + +import cuda.bindings.driver as cuda_driver +import cutlass +import cutlass.cute as cute +from cutlass._mlir.dialects import llvm as _llvm +from cutlass._mlir.extras import types as _T +from cutlass.cutlass_dsl import dsl_user_op +from cutlass.experimental import primitives as prims + +THREADS = 256 +WARPS = THREADS // 32 +MARKOV_RANK = 256 # 32 lanes x 8 columns +ROW_WORDS = MARKOV_RANK // 2 # int32 words (bf16 pairs) per markov row +ROW_BYTES = MARKOV_RANK * 2 +BLOCK_ROWS = 8 # rows per warp block: one per epilogue lane +MAX_BATCH = 16 +MAX_BLOCK = 16 +BUFFERS = 3 +FLAG_WORDS = 4 +EMPTY_WORD = -(2**31) # 0x80000000: fp32 -0.0, never a pushed word +INT_MAX = 2**31 - 1 +NEG_INF_BITS = -8388608 # 0xFF800000 +COPY_ROWS = 32 # rows per weight bulk copy (16 KB) +SMEM_ROW_BUDGET = 400 # 200 KB of markov_w2 rows per CTA +SMEM_BYTES = 227 * 1024 + + +def rows_per_cta(shard: int, grid: int) -> int: + return shard // grid + + +def supports(shard: int, grid: int, block: int, batch: int) -> bool: + """Shapes the kernel runs: whole 8-row blocks for every warp (RC a multiple of 64), the CTA's markov_w2 rows + within the shared-memory budget, block <= 16, batch <= 16, an even entry count per exchange.""" + if grid <= 0 or shard % grid != 0: + return False + rc = rows_per_cta(shard, grid) + return ( + rc % (WARPS * BLOCK_ROWS) == 0 + and rc <= SMEM_ROW_BUDGET + and smem_bytes(rc, batch) <= SMEM_BYTES + and grid % 2 == 0 + and 0 < block <= MAX_BLOCK + and 0 < batch <= MAX_BATCH + ) + + +def smem_bytes(rows: int, batch: int) -> int: + """Shared memory of a CTA: its markov_w2 rows, the markov_w1[prev] rows, the two [b][warp] reductions, the + barriers, flags and prev (with the allocator's alignment padding).""" + return rows * ROW_BYTES + batch * ROW_BYTES + 2 * batch * WARPS * 8 + 1024 + + +def buffer_words(block: int, batch: int, slots: int, grid: int) -> int: + """Int32 words of one Lamport buffer: one (value, index) entry per position, request, rank slot and CTA.""" + return block * batch * slots * grid * 2 + + +@dsl_user_op +def _atomic_add_acq_rel(addr_i64, val, *, loc=None, ip=None): + """atom.acq_rel.gpu.global.add.u32, returning the old value.""" + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [addr_i64.ir_value(loc=loc, ip=ip), val.ir_value(loc=loc, ip=ip)], + "atom.acq_rel.gpu.global.add.u32 $0, [$1], $2;", "=r,l,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _fma_rn(a, b, c, *, loc=None, ip=None): + """fma.rn.f32 (fmaf).""" + return cutlass.Float32( + _llvm.inline_asm( + _T.f32(), [a.ir_value(loc=loc, ip=ip), b.ir_value(loc=loc, ip=ip), c.ir_value(loc=loc, ip=ip)], + "fma.rn.f32 $0, $1, $2, $3;", "=f,f,f,f", has_side_effects=False, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _add_rn(a, b, *, loc=None, ip=None): + """add.rn.f32: an fp32 add the compiler cannot contract with a neighbouring multiply.""" + return cutlass.Float32( + _llvm.inline_asm( + _T.f32(), [a.ir_value(loc=loc, ip=ip), b.ir_value(loc=loc, ip=ip)], + "add.rn.f32 $0, $1, $2;", "=f,f,f", has_side_effects=False, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +@dsl_user_op +def _pack_bf16x2(hi, lo, *, loc=None, ip=None): + """(bf16(hi) << 16) | bf16(lo), round to nearest even (cvt.rn, as __float2bfloat16_rn).""" + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [hi.ir_value(loc=loc, ip=ip), lo.ir_value(loc=loc, ip=ip)], + "cvt.rn.bf16x2.f32 $0, $1, $2;", "=r,f,f", has_side_effects=False, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +def _lo_f32(word): + """The bf16 in the low half of a word, as fp32 (exact).""" + return cutlass.Int32(word << cutlass.Int32(16)).bitcast(cutlass.Float32) + + +def _hi_f32(word): + """The bf16 in the high half of a word, as fp32 (exact).""" + return cutlass.Int32(word & cutlass.Int32(-65536)).bitcast(cutlass.Float32) + + +def _better(value, index, best_value, best_index): + """dsparkMarkovChainKernel's keepBetter order: larger value, or an equal value at a lower index. NaN never + compares better, so it never wins.""" + return (value > best_value) | ((value == best_value) & (index < best_index)) + + +def _keep_better(best_value, best_index, value, index): + take = _better(value, index, best_value, best_index) + return cutlass.Float32(cutlass.select_(take, value, best_value)), cutlass.Int32( + cutlass.select_(take, index, best_index) + ) + + +def _warp_best(value, index): + """Every lane gets the warp's best (value, index) (xor tree: the order is total, so all lanes agree).""" + for offset in [16, 8, 4, 2, 1]: + other_value = cute.arch.shuffle_sync_bfly(value, offset=offset, mask=-1, mask_and_clamp=31) + other_index = cute.arch.shuffle_sync_bfly(index, offset=offset, mask=-1, mask_and_clamp=31) + value, index = _keep_better(value, index, other_value, other_index) + return value, index + + +def _block_dot(s_w2, first_row, lane, w1f): + """markov_w2[first_row + j] . markov_w1[prev] for the block's 8 rows, in the C++ kernel's order: lane l sums the + 8 columns 8l .. 8l+7 by an FMA chain from 0, then the 32 lane sums are added in pairs at lane distances 16, 8, 4, 2, + 1 (the xor tree). Done as a reduce-scatter: at distances 16, 8 and 4 a lane keeps half of its rows and adds its + partner's partial of each kept row (the pairs of the full tree, so the same sums: 9 shuffles instead of 40), then + the full tree at 2 and 1. Returns row (lane >> 2)'s sum.""" + parts = [] + for r in range(BLOCK_ROWS): + w2v = s_w2.load(idx=(first_row + cutlass.Int32(r)) * cutlass.Int32(ROW_WORDS) + lane * cutlass.Int32(4), + vector_size=4, alignment=16) # fmt: skip + acc = cutlass.Float32(0.0) + for q in range(4): + word = cutlass.Int32(w2v[q]) + acc = _fma_rn(w1f[2 * q], _lo_f32(word), acc) + acc = _fma_rn(w1f[2 * q + 1], _hi_f32(word), acc) + parts.append(acc) + for offset in [16, 8, 4]: + upper = (lane & cutlass.Int32(offset)) != cutlass.Int32(0) + half = len(parts) // 2 + kept = [] + for i in range(half): + send = cutlass.Float32(cutlass.select_(upper, parts[i], parts[half + i])) + keep = cutlass.Float32(cutlass.select_(upper, parts[half + i], parts[i])) + kept.append( + _add_rn( + keep, + cute.arch.shuffle_sync_bfly(send, offset=offset, mask=-1, mask_and_clamp=31), + ) + ) + parts = kept + dot = parts[0] + for offset in [2, 1]: + dot = _add_rn( + dot, cute.arch.shuffle_sync_bfly(dot, offset=offset, mask=-1, mask_and_clamp=31) + ) + return dot + + +def _fetch_w1_row(s_w1, s_valid, w1, mbar, b, prev, vocab): + """Request b's markov_w1[prev] row into shared memory: one 512-byte bulk copy on the position's barrier phase (a + prev outside the vocabulary copies row 0 and clears s_valid[b]: the readers zero its bias). s_valid[b] is stored + before the arrive, whose release the readers' phase wait acquires.""" + valid = (prev >= cutlass.Int32(0)) & (prev < vocab) + row = cutlass.Int32(cutlass.select_(valid, prev, cutlass.Int32(0))) + s_valid.store(cutlass.Int32(cutlass.select_(valid, cutlass.Int32(1), cutlass.Int32(0))), idx=b) + prims.mbarrier_arrive_expect_tx(mbar, ROW_BYTES) + prims.cp_async_bulk_shared_cluster_global( + s_w1.subview(b * cutlass.Int32(ROW_WORDS)), + w1.subview(row * cutlass.Int32(ROW_WORDS)), + mbar, + ROW_BYTES, + ) + + +def _reduce_slots(slots_arr, first_word): + """The best of 8 (value bits, index) pairs at first_word .. in warp order.""" + value = cutlass.Int32(NEG_INF_BITS).bitcast(cutlass.Float32) + index = cutlass.Int32(INT_MAX) + for w in range(WARPS): + pair = slots_arr.load(idx=first_word + cutlass.Int32(w * 2), vector_size=2, alignment=8) + value, index = _keep_better( + value, index, cutlass.Int32(pair[0]).bitcast(cutlass.Float32), cutlass.Int32(pair[1]) + ) + return value, index + + +@cute.kernel +def k3_markov_kernel( + base: cutlass.Array, # fp32 [B * K * S] (or its int16 bf16 view with base_bf16): the draft head's shard logits + first_prev: cutlass.Array, # int64 [B]: the anchor token of each request + w1: cutlass.Array, # int32 words of bf16 markov_w1 [V, 256] + w2: cutlass.Array, # int32 words of bf16 markov_w2[shard rows] [S, 256] + corrected: cutlass.Array, # fp32 [B * K * S], out + tokens: cutlass.Array, # int32 [B * K], out + next_new: cutlass.Array, # int32 [B * (K + 1)], out + accepted: cutlass.Array, # int32 [*, accepted_stride]: accepted tokens + num_accepted: cutlass.Array, # int32 [B] + accepted_rows: cutlass.Array, # int32 [B]: the row of accepted of each request + kv_lens: cutlass.Array, # int32: the target's KV lengths, rewound here after the draft forward + rewind: cutlass.Array, # int32 [rewind_count]: subtracted from kv_lens[rewind_first ...], clamped at 0 + buf_uc: cutlass.Array, # int32 words: this rank's Lamport buffers + buf_mc: cutlass.Array, # int32 words: their multicast mapping + flags: cutlass.Array, # int32 [4] + vocab: cutlass.Int32, # rows of markov_w1: a prev outside [0, vocab) gets a zero bias + vocab_offset: cutlass.Int32, # first vocabulary index of this rank's shard + rank: cutlass.Int32, + accepted_stride: cutlass.Int32, + rewind_first: cutlass.Int32, + rewind_count: cutlass.Int32, # 0: no rewind + grid: cutlass.Constexpr[int], + rows: cutlass.Constexpr[int], # RC + shard: cutlass.Constexpr[int], # S + block: cutlass.Constexpr[int], # K + batch: cutlass.Constexpr[int], # B + slots: cutlass.Constexpr[int], # rank slots per (k, b): world x push_copies + push_copies: cutlass.Constexpr[ + int + ], # test only: each rank's entry goes into push_copies consecutive slots + buf_words: cutlass.Constexpr[int], + base_bf16: cutlass.Constexpr[bool], +): + rows_per_warp = rows // WARPS + blocks_per_warp = rows_per_warp // BLOCK_ROWS + entries = slots * grid # (value, index) entries per (k, b) + vectors = entries // 2 # 16-byte loads per (k, b) + vec_per_thread = (vectors + THREADS - 1) // THREADS + used_words = block * batch * entries * 2 + + tx, _, _ = cute.arch.thread_idx() + cta, _, _ = cute.arch.block_idx() + warp = cute.arch.make_warp_uniform(cute.arch.warp_idx()) + lane = tx % cutlass.Int32(32) + + s_w2 = cutlass.Array( + cutlass.Int32, rows * ROW_WORDS, space=cutlass.AddressSpace.smem, alignment=128 + ) + # [b]: markov_w1[prev] of the position (one bulk copy per request, issued by the thread that learns prev). + s_w1 = cutlass.Array( + cutlass.Int32, batch * ROW_WORDS, space=cutlass.AddressSpace.smem, alignment=128 + ) + s_mbar = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + s_mbar_w1 = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8) + # [b][warp] (value bits, index): the CTA reduction of the rows, then of the polled entries. + s_row_best = cutlass.Array( + cutlass.Int32, batch * WARPS * 2, space=cutlass.AddressSpace.smem, alignment=16 + ) + s_poll_best = cutlass.Array( + cutlass.Int32, batch * WARPS * 2, space=cutlass.AddressSpace.smem, alignment=16 + ) + s_flags = cutlass.Array(cutlass.Int32, 4, space=cutlass.AddressSpace.smem, alignment=16) + s_valid = cutlass.Array(cutlass.Int32, MAX_BATCH, space=cutlass.AddressSpace.smem, alignment=16) + + if warp == 0: + if prims.elect_sync(): + prims.mbarrier_init(s_mbar, 1) + prims.mbarrier_init(s_mbar_w1, batch) + prims.fence_mbarrier_init() + prims.barrier_cta_sync(0) + + # The CTA's markov_w2 rows: a weight, loaded before the grid dependency. + if warp == 0: + if prims.elect_sync(): + prims.mbarrier_arrive_expect_tx(s_mbar, rows * ROW_BYTES) + for c in cutlass.range_constexpr((rows + COPY_ROWS - 1) // COPY_ROWS): + n = min(COPY_ROWS, rows - c * COPY_ROWS) + prims.cp_async_bulk_shared_cluster_global( + s_w2.subview(c * COPY_ROWS * ROW_WORDS), + w2.subview( + cta * cutlass.Int32(rows * ROW_WORDS) + + cutlass.Int32(c * COPY_ROWS * ROW_WORDS) + ), + s_mbar, + n * ROW_BYTES, + ) + # Dependents wait for this whole grid before reading its outputs. + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + + prims.griddepcontrol(prims.GridDepAction.WAIT) + if tx == 0: + s_flags.store(flags.load(idx=0, is_volatile=True), idx=0) + s_flags.store(flags.load(idx=2, is_volatile=True), idx=2) + if tx < cutlass.Int32(batch): + _fetch_w1_row( + s_w1, s_valid, w1, s_mbar_w1, tx, cutlass.Int32(first_prev.load(idx=tx)), vocab + ) + prims.barrier_cta_sync(0) + cur = s_flags.load(idx=0) + cur_base = cur * cutlass.Int32(buf_words) + + # Re-arm the previous call's buffer (this CTA's 16-byte chunks of the words it used). + dirty_base = ((cur + cutlass.Int32(2)) % cutlass.Int32(BUFFERS)) * cutlass.Int32(buf_words) + dirty_words = s_flags.load(idx=2) + empty = cutlass.Int32(EMPTY_WORD) + w_clear = (cta * cutlass.Int32(THREADS) + tx) * cutlass.Int32(4) + while w_clear < dirty_words: + buf_uc.store((empty, empty, empty, empty), idx=dirty_base + w_clear, alignment=16) + w_clear += cutlass.Int32(grid * THREADS * 4) + + # next_new_tokens[b] = [accepted[row_b, num_accepted[b] - 1], tokens[b, :]] (-1 wraps, as torch indexing). + if cta == 0: + if tx < cutlass.Int32(batch): + col = num_accepted.load(idx=tx) - cutlass.Int32(1) + col = cutlass.Int32(cutlass.select_(col < cutlass.Int32(0), col + accepted_stride, col)) + row = accepted_rows.load(idx=tx) + next_new.store( + accepted.load(idx=row * accepted_stride + col), idx=tx * cutlass.Int32(block + 1) + ) + # The KV-length rewind after the draft forward (whose kernels, the readers of kv_lens, all finished before + # this grid's wait): kv_lens[first + b] = max(kv_lens[first + b] - rewind[b], 0). + if tx < rewind_count: + left = kv_lens.load(idx=rewind_first + tx) - rewind.load(idx=tx) + kv_lens.store(cutlass.Int32(cutlass.select_(left < cutlass.Int32(0), cutlass.Int32(0), left)), + idx=rewind_first + tx) # fmt: skip + + # The weights have landed (issued before the wait). + while not cute.arch.mbarrier_test_wait(s_mbar.data_ptr(), 0): + pass + + row0 = warp * cutlass.Int32(rows_per_warp) # first row of the warp within the CTA + shard_row0 = cta * cutlass.Int32(rows) + row0 # ... within the shard + for k in cutlass.range(block, unroll=1): + for b in cutlass.range(batch, unroll=1): + base_row = (b * cutlass.Int32(block) + k) * cutlass.Int32(shard) + # Lane 4j finishes row j of each 8-row block (the reduce-scatter leaves row j's sum on lanes 4j .. 4j+3): + # the warp's base logits, in flight during the wait below. + has_row = (lane & cutlass.Int32(3)) == cutlass.Int32(0) + lane_row = shard_row0 + (lane >> cutlass.Int32(2)) + base_vals = [] + for g in cutlass.range_constexpr(blocks_per_warp): + if cutlass.const_expr(base_bf16): + base_vals.append( + cutlass.Int32( + cutlass.Int32( + base.load(idx=base_row + lane_row + cutlass.Int32(g * BLOCK_ROWS)) + ) + << cutlass.Int32(16) + ).bitcast(cutlass.Float32) + ) + else: + base_vals.append( + base.load(idx=base_row + lane_row + cutlass.Int32(g * BLOCK_ROWS)) + ) + # markov_w1[prev] landed (the position's phase of the W1 barrier). + while not cute.arch.mbarrier_test_wait(s_mbar_w1.data_ptr(), k & cutlass.Int32(1)): + pass + valid = s_valid.load(idx=b) != cutlass.Int32(0) + # markov_w1[prev], columns 8 lane .. 8 lane + 7, as fp32 (zeros for an invalid prev). + w1v = s_w1.load( + idx=b * cutlass.Int32(ROW_WORDS) + lane * 4, vector_size=4, alignment=16 + ) + w1f = [] + for q in cutlass.range_constexpr(4): + word = cutlass.Int32( + cutlass.select_(valid, cutlass.Int32(w1v[q]), cutlass.Int32(0)) + ) + w1f.append(_lo_f32(word)) + w1f.append(_hi_f32(word)) + best_value = cutlass.Int32(NEG_INF_BITS).bitcast(cutlass.Float32) + best_index = cutlass.Int32(INT_MAX) + for g in cutlass.range_constexpr(blocks_per_warp): + dot = _block_dot(s_w2, row0 + cutlass.Int32(g * BLOCK_ROWS), lane, w1f) + bias = _lo_f32(_pack_bf16x2(dot, dot)) + value = _add_rn(base_vals[g], bias) + if has_row: + my_row = lane_row + cutlass.Int32(g * BLOCK_ROWS) + corrected.store(value, idx=base_row + my_row) + best_value, best_index = _keep_better( + best_value, best_index, value, vocab_offset + my_row + ) + best_value, best_index = _warp_best(best_value, best_index) + if lane == 0: + s_row_best.store((best_value.bitcast(cutlass.Int32), best_index), + idx=(b * cutlass.Int32(WARPS) + warp) * cutlass.Int32(2), alignment=8) # fmt: skip + prims.barrier_cta_sync(0) + # Thread b pushes request b's CTA maximum into every rank's buffer. + if tx < cutlass.Int32(batch): + cta_value, cta_index = _reduce_slots(s_row_best, tx * cutlass.Int32(WARPS * 2)) + bits = cta_value.bitcast(cutlass.Int32) + bits = cutlass.Int32( + cutlass.select_(bits == cutlass.Int32(EMPTY_WORD), cutlass.Int32(0), bits) + ) + for c in cutlass.range_constexpr(push_copies): + s_idx = rank * cutlass.Int32(push_copies) + cutlass.Int32(c) + word = cur_base + (((k * cutlass.Int32(batch) + tx) * cutlass.Int32(slots) + s_idx) + * cutlass.Int32(grid) + cta) * cutlass.Int32(2) # fmt: skip + buf_mc.store((bits, cta_index), idx=word, alignment=8) + # Drain the posted multicast store now (the polls below issue no release that would). + prims.fence_acq_rel(prims.MemScope.CLUSTER) + # Poll every rank's entries of [k][b] (this thread's 16-byte vectors, all loads in flight together) until + # none is empty; the pass that finds them all complete also reduces them. + for b in cutlass.range(batch, unroll=1): + region = cur_base + (k * cutlass.Int32(batch) + b) * cutlass.Int32(entries * 2) + best_value = cutlass.Int32(NEG_INF_BITS).bitcast(cutlass.Float32) + best_index = cutlass.Int32(INT_MAX) + pending = cutlass.Boolean(True) + while pending: + pass_value = cutlass.Int32(NEG_INF_BITS).bitcast(cutlass.Float32) + pass_index = cutlass.Int32(INT_MAX) + pass_pending = cutlass.Boolean(False) + for i in cutlass.range_constexpr(vec_per_thread): + v_idx = tx + cutlass.Int32(i * THREADS) + live = v_idx < cutlass.Int32(vectors) + got = buf_uc.load( + idx=region + + cutlass.Int32(cutlass.select_(live, v_idx, cutlass.Int32(0))) + * cutlass.Int32(4), + vector_size=4, + alignment=16, + is_volatile=True, + ) + for e in cutlass.range_constexpr(2): + vw = cutlass.Int32(got[2 * e]) + iw = cutlass.Int32(got[2 * e + 1]) + pass_pending = pass_pending | (live & ((vw == empty) | (iw == empty))) + take = live & _better( + vw.bitcast(cutlass.Float32), iw, pass_value, pass_index + ) + pass_value = cutlass.Float32( + cutlass.select_(take, vw.bitcast(cutlass.Float32), pass_value) + ) + pass_index = cutlass.Int32(cutlass.select_(take, iw, pass_index)) + best_value = pass_value + best_index = pass_index + pending = pass_pending + best_value, best_index = _warp_best(best_value, best_index) + if lane == 0: + s_poll_best.store((best_value.bitcast(cutlass.Int32), best_index), + idx=(b * cutlass.Int32(WARPS) + warp) * cutlass.Int32(2), alignment=8) # fmt: skip + prims.barrier_cta_sync(0) + if tx < cutlass.Int32(batch): + g_value, g_index = _reduce_slots(s_poll_best, tx * cutlass.Int32(WARPS * 2)) + # The next position's markov_w1 row, requested as soon as prev is known: every reader of s_w1 and + # s_valid for this position passed the barrier before the push. No barrier follows: the other threads + # go on to the next position's base logits and wait for the row's phase. s_poll_best is next written + # after the next push barrier, which this thread reaches only after these reads. + if k + cutlass.Int32(1) < cutlass.Int32(block): + _fetch_w1_row(s_w1, s_valid, w1, s_mbar_w1, tx, g_index, vocab) + if cta == 0: + tokens.store(g_index, idx=tx * cutlass.Int32(block) + k) + next_new.store(g_index, idx=tx * cutlass.Int32(block + 1) + k + cutlass.Int32(1)) + + # Every CTA read the flags before arriving; the last one advances them for the next call. + if tx == 0: + arrived = _atomic_add_acq_rel(flags.data_ptr(1).toint(), cutlass.Int32(1)) + if arrived == cutlass.Int32(grid - 1): + flags.store(cutlass.Int32(0), idx=1) + flags.store(cutlass.Int32(used_words), idx=2) + flags.store((cur + cutlass.Int32(1)) % cutlass.Int32(BUFFERS), idx=0) + + +@cute.jit +def k3_markov( + base: cute.Tensor, + first_prev: cute.Tensor, + w1: cute.Tensor, + w2: cute.Tensor, + corrected: cute.Tensor, + tokens: cute.Tensor, + next_new: cute.Tensor, + accepted: cute.Tensor, + num_accepted: cute.Tensor, + accepted_rows: cute.Tensor, + kv_lens: cute.Tensor, + rewind: cute.Tensor, + buf_uc: cute.Tensor, + buf_mc: cute.Tensor, + flags: cute.Tensor, + vocab: cutlass.Int32, + vocab_offset: cutlass.Int32, + rank: cutlass.Int32, + accepted_stride: cutlass.Int32, + rewind_first: cutlass.Int32, + rewind_count: cutlass.Int32, + grid: cutlass.Constexpr[int], + rows: cutlass.Constexpr[int], + shard: cutlass.Constexpr[int], + block: cutlass.Constexpr[int], + batch: cutlass.Constexpr[int], + slots: cutlass.Constexpr[int], + push_copies: cutlass.Constexpr[int], + buf_words: cutlass.Constexpr[int], + base_bf16: cutlass.Constexpr[bool], + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + k3_markov_kernel( + base, first_prev, w1, w2, corrected, tokens, next_new, accepted, num_accepted, accepted_rows, kv_lens, + rewind, buf_uc, buf_mc, flags, vocab, vocab_offset, rank, accepted_stride, rewind_first, rewind_count, + grid, rows, shard, block, batch, slots, push_copies, buf_words, base_bf16, + ).launch( + grid=[grid, 1, 1], + block=[THREADS, 1, 1], + stream=stream, + use_pdl=use_pdl, + ) # fmt: skip diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_markov/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_markov/op.py new file mode 100644 index 000000000000..cf7cb969f923 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_markov/op.py @@ -0,0 +1,306 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``trtllm::k3_markov``: the DSpark vanilla-Markov draft chain over a vocab-sharded draft head in one CTM kernel. + +For the block logits ``base`` [B, K, S] of this rank's vocab shard it returns the Markov-corrected logits (fp32, as +``dspark_markov_chain`` with fp32 logits), every position's global greedy token (first maximum, lowest vocabulary +index, across the tensor-parallel ranks) and the next step's input tokens +``[accepted[row_b, num_accepted[b] - 1], tokens[b, :]]``: what ``dspark_markov_chain`` over the gathered logits +followed by the TP-gathered greedy sampler computes. + +The ranks exchange one (value, index) entry per CTA and position through this module's own MNNVL multicast buffers +(``markov_workspace``), allocated collectively on the first call of a TP group, which must happen outside CUDA-graph +capture (the kernel also compiles there). +""" + +from __future__ import annotations + +import importlib.util +import os +import sys +import threading +from typing import Dict, Optional, Tuple + +import torch + +_lock = threading.Lock() +_compiled: Dict[tuple, object] = {} +_workspaces: Dict[object, dict] = {} +_modules: Dict[str, object] = {} + +WORKSPACE_MAX_BLOCK = 8 +WORKSPACE_MAX_GRID = 152 + + +def _kernel_module(): + """k3_markov_kernel.py next to this file (also when op.py is loaded outside the package).""" + mod = _modules.get("kernel") + if mod is None: + path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "k3_markov_kernel.py") + spec = importlib.util.spec_from_file_location(f"{__name__}_k3_markov_kernel", path) + mod = importlib.util.module_from_spec(spec) + sys.modules[spec.name] = mod + spec.loader.exec_module(mod) + _modules["kernel"] = mod + return mod + + +def _sm_count(device) -> int: + return torch.cuda.get_device_properties(device).multi_processor_count + + +def pick_grid(shard: int, block: int, batch: int, device=None) -> int: + """CTAs for a shard: 128 when the kernel supports that split (see ``k3_markov_kernel.supports``), else the + largest supported count up to the SM count; 0 if none.""" + kern = _kernel_module() + limit = min( + WORKSPACE_MAX_GRID, _sm_count(device if device is not None else torch.cuda.current_device()) + ) + if 128 <= limit and kern.supports(shard, 128, block, batch): + return 128 + for grid in range(limit, 0, -1): + if kern.supports(shard, grid, block, batch): + return grid + return 0 + + +def supports(base: torch.Tensor, markov_w1: torch.Tensor, markov_w2_shard: torch.Tensor) -> bool: + """Whether ``k3_markov`` runs these operands (the TP group also needs MNNVL, see ``markov_workspace``).""" + kern = _kernel_module() + if base.dim() != 3 or base.dtype not in (torch.float32, torch.bfloat16) or not base.is_cuda: + return False + batch, block, shard = base.shape + return ( + markov_w1.dtype == torch.bfloat16 + and markov_w2_shard.dtype == torch.bfloat16 + and markov_w1.dim() == 2 + and markov_w1.shape[1] == kern.MARKOV_RANK + and tuple(markov_w2_shard.shape) == (shard, kern.MARKOV_RANK) + and block <= WORKSPACE_MAX_BLOCK + and pick_grid(shard, block, batch, base.device) > 0 + ) + + +def markov_workspace(mapping, push_copies: int = 1) -> dict: + """This TP group's Lamport buffers (3 rotating buffers, every word 0x80000000 when armed) behind one multicast + mapping, and its flag words. Collective on first use: every rank of the group must make the first call at the + same point, outside CUDA-graph capture. ``push_copies`` > 1 (tests only) gives every rank that many slots, to + emulate a larger group's exchange volume.""" + key = (mapping, push_copies) + ws = _workspaces.get(key) + if ws is not None: + return ws + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError("k3_markov: the workspace must be allocated outside CUDA-graph capture") + from tensorrt_llm._torch.distributed.ops import ( + _get_mnnvl_workspace_comm, + _make_mnnvl_mcast_buffer, + _mnnvl_workspace_all_succeeded, + ) + + kern = _kernel_module() + world = mapping.tp_size + slots = world * push_copies + buf_words = kern.buffer_words(WORKSPACE_MAX_BLOCK, kern.MAX_BATCH, slots, WORKSPACE_MAX_GRID) + comm = _get_mnnvl_workspace_comm(mapping) + use_fabric_handle = ( + os.environ.get("TRTLLM_FORCE_MNNVL_AR", "0") == "1" or mapping.is_multi_node() + ) + error: Optional[Exception] = None + try: + words = kern.BUFFERS * buf_words + handle = _make_mnnvl_mcast_buffer(comm, words * 4, mapping, use_fabric_handle) + uc = handle.get_uc_buffer(mapping.tp_rank, (words,), torch.int32, 0) + mc = handle.get_mc_buffer((words,), torch.int32, 0) + with torch.inference_mode(): + uc.fill_(kern.EMPTY_WORD) + flags = torch.zeros(kern.FLAG_WORDS, dtype=torch.int32, device=uc.device) + torch.cuda.synchronize() + ws = dict(handle=handle, comm=comm, uc=uc, mc=mc, flags=flags, rank=mapping.tp_rank, world=world, + slots=slots, push_copies=push_copies, buf_words=buf_words) # fmt: skip + except Exception as exc: # noqa: BLE001 -- reported to every rank below, then re-raised + error = exc + # Also the barrier that keeps any rank from pushing into a peer's buffer before the peer has armed it. + if not _mnnvl_workspace_all_succeeded(comm, error is None): + raise RuntimeError("k3_markov Lamport buffers failed on at least one rank") from error + _workspaces[key] = ws + return ws + + +def _arg(t: torch.Tensor, align: int = 16): + from cutlass.cute.runtime import from_dlpack + + return from_dlpack(t.detach(), assumed_align=align).mark_layout_dynamic(leading_dim=0) + + +@torch.library.custom_op("trtllm::k3_markov", mutates_args=("kv_lens",)) +def k3_markov( + base: torch.Tensor, + first_prev: torch.Tensor, + markov_w1: torch.Tensor, + markov_w2_shard: torch.Tensor, + shard_offset: int, + ws_uc: torch.Tensor, + ws_mc: torch.Tensor, + ws_flags: torch.Tensor, + rank: int, + slots: int, + push_copies: int, + buf_words: int, + accepted: torch.Tensor, + num_accepted: torch.Tensor, + accepted_rows: torch.Tensor, + kv_lens: torch.Tensor, + rewind: torch.Tensor, + rewind_first: int, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """``(corrected [B, K, S] fp32, tokens [B, K] int32, next_new [B, K + 1] int32)`` for the block logits ``base`` + [B, K, S] (fp32, or bf16 converted exactly) of this rank's shard, which starts at vocabulary index + ``shard_offset``. ``ws_*``, ``slots``, ``push_copies``, ``buf_words``: ``markov_workspace``'s. ``accepted`` + [rows, K + 1] int32, ``num_accepted`` [B] int32 and ``accepted_rows`` [B] int32 give next_new's first column + ``accepted[accepted_rows[b], num_accepted[b] - 1]``. ``kv_lens`` (int32) is rewound in place: + ``kv_lens[rewind_first + b] = max(kv_lens[rewind_first + b] - rewind[b], 0)`` for b < ``rewind.numel()`` (<= B; + empty: no rewind).""" + import cuda.bindings.driver as cuda_driver + import cutlass.cute as cute + + kern = _kernel_module() + if not supports(base, markov_w1, markov_w2_shard): + raise ValueError( + f"k3_markov: unsupported operands: base {tuple(base.shape)} {base.dtype} (fp32/bf16 [B <= " + f"{kern.MAX_BATCH}, K <= {WORKSPACE_MAX_BLOCK}, S]), markov_w1 {tuple(markov_w1.shape)} " + f"{markov_w1.dtype} (bf16 [V, {kern.MARKOV_RANK}]), markov_w2_shard {tuple(markov_w2_shard.shape)} " + f"{markov_w2_shard.dtype} (bf16 [S, {kern.MARKOV_RANK}]); S must split over an even number of CTAs " + f"(at most the SM count) into a multiple of 64 rows each, at most {kern.SMEM_ROW_BUDGET}" + ) + batch, block, shard = base.shape + grid = pick_grid(shard, block, batch, base.device) + rows = kern.rows_per_cta(shard, grid) + if kern.buffer_words(block, batch, slots, grid) > buf_words: + raise ValueError("k3_markov: the exchange does not fit the workspace buffers") + if grid > _sm_count(base.device): + raise ValueError(f"k3_markov: {grid} CTAs must all be resident; the GPU has fewer SMs") + if first_prev.dtype != torch.int64 or first_prev.numel() != batch: + raise ValueError("k3_markov: first_prev must be int64 [B]") + if ( + accepted.dtype != torch.int32 + or accepted.dim() != 2 + or num_accepted.dtype != torch.int32 + or accepted_rows.dtype != torch.int32 + or num_accepted.numel() < batch + or accepted_rows.numel() < batch + ): + raise ValueError( + "k3_markov: accepted must be int32 [rows, K + 1], num_accepted / accepted_rows int32 [>= B]" + ) + rewind_count = rewind.numel() + if ( + kv_lens.dtype != torch.int32 + or rewind.dtype != torch.int32 + or rewind_count > batch + or ( + rewind_count > 0 and (rewind_first < 0 or rewind_first + rewind_count > kv_lens.numel()) + ) + ): + raise ValueError( + "k3_markov: kv_lens / rewind must be int32, the rewind at most B rows inside kv_lens" + ) + device = base.device + corrected = torch.empty(batch, block, shard, dtype=torch.float32, device=device) + tokens = torch.empty(batch, block, dtype=torch.int32, device=device) + next_new = torch.empty(batch, block + 1, dtype=torch.int32, device=device) + base_bf16 = base.dtype == torch.bfloat16 + base_view = base.contiguous().view(-1) + if base_bf16: + base_view = base_view.view(torch.int16) + # The kernel reads first_prev, accepted, num_accepted, accepted_rows, kv_lens and rewind one element at a time: + # declared at their element alignment, they may be views at any offset (e.g. the generation requests' rows after + # the context requests'). + args = ( + _arg(base_view), + _arg(first_prev.contiguous().view(-1), align=8), + _arg(markov_w1.contiguous().view(-1).view(torch.int32)), + _arg(markov_w2_shard.contiguous().view(-1).view(torch.int32)), + _arg(corrected.view(-1)), + _arg(tokens.view(-1)), + _arg(next_new.view(-1)), + _arg(accepted.contiguous().view(-1), align=4), + _arg(num_accepted.contiguous().view(-1), align=4), + _arg(accepted_rows.contiguous().view(-1), align=4), + _arg(kv_lens.view(-1), align=4), + _arg(rewind.contiguous().view(-1) if rewind_count > 0 else kv_lens.view(-1), align=4), + _arg(ws_uc.view(-1)), + _arg(ws_mc.view(-1)), + _arg(ws_flags.view(-1)), + ) + scalars = (int(markov_w1.shape[0]), int(shard_offset), int(rank), int(accepted.shape[1]), int(rewind_first), + int(rewind_count)) # fmt: skip + use_pdl = os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" + stream = cuda_driver.CUstream(torch.cuda.current_stream(device).cuda_stream) + consts = (grid, rows, shard, block, batch, slots, push_copies, buf_words, base_bf16) + key = consts + (use_pdl,) + fn = _compiled.get(key) + if fn is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "trtllm::k3_markov must run once per shape outside CUDA-graph capture first" + ) + with _lock: + fn = _compiled.get(key) + if fn is None: + fn = _compiled[key] = cute.compile( + kern.k3_markov, *args, *scalars, *consts, use_pdl, stream + ) + fn(*args, *scalars, stream) + return corrected, tokens, next_new + + +@k3_markov.register_fake +def _(base, first_prev, markov_w1, markov_w2_shard, shard_offset, ws_uc, ws_mc, ws_flags, rank, slots, push_copies, + buf_words, accepted, num_accepted, accepted_rows, kv_lens, rewind, rewind_first): # fmt: skip + batch, block, shard = base.shape + return ( + base.new_empty((batch, block, shard), dtype=torch.float32), + base.new_empty((batch, block), dtype=torch.int32), + base.new_empty((batch, block + 1), dtype=torch.int32), + ) + + +def markov_chain( + mapping, + base: torch.Tensor, + first_prev: torch.Tensor, + markov_w1: torch.Tensor, + markov_w2_shard: torch.Tensor, + shard_offset: int, + accepted: torch.Tensor, + num_accepted: torch.Tensor, + accepted_rows: torch.Tensor, + push_copies: int = 1, + kv_lens: Optional[torch.Tensor] = None, + rewind: Optional[torch.Tensor] = None, + rewind_first: int = 0, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """``trtllm::k3_markov`` on ``mapping``'s TP group workspace; with ``kv_lens`` and ``rewind`` it also applies the + KV-length rewind (see ``k3_markov``).""" + ws = markov_workspace(mapping, push_copies) + if kv_lens is None or rewind is None: + # No rewind: an empty rewind leaves the kernel's kv_lens argument (the flags words) untouched. + kv_lens, rewind = ws["flags"], ws["flags"][:0] + return torch.ops.trtllm.k3_markov( + base, first_prev, markov_w1, markov_w2_shard, int(shard_offset), ws["uc"], ws["mc"], ws["flags"], + ws["rank"], ws["slots"], ws["push_copies"], ws["buf_words"], accepted, num_accepted, accepted_rows, kv_lens, + rewind, int(rewind_first), + ) # fmt: skip diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_spec_accept/__init__.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_spec_accept/__init__.py new file mode 100644 index 000000000000..24af92675a48 --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_spec_accept/__init__.py @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""Speculative-decode acceptance + block-drafter input prep in CuTe DSL (``trtllm::k3_spec_accept``). + +Importing :mod:`.op` registers the torch op; nothing is imported eagerly here so that +the CuTe DSL dependency stays optional for every other model. +""" diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_spec_accept/k3_spec_accept_kernel.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_spec_accept/k3_spec_accept_kernel.py new file mode 100644 index 000000000000..85340fbf648d --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_spec_accept/k3_spec_accept_kernel.py @@ -0,0 +1,586 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +# ============================================================================= +# Speculative-decode acceptance + block-drafter input prep for one decode step -- CTM (prims/cute) kernel +# ============================================================================= +# +# For B generation requests (no context requests), K drafts each, greedy strict acceptance (DFlash / DSpark): +# t[r] = argmax(logits[r, :]) r < B (K + 1); torch's order: NaN first, then larger, lower index +# accepted[b, j] = t[b (K + 1) + j] +# n[b] = 1 + #leading j < K with draft[b, j] == t[b (K + 1) + j] (or the forced count) +# prev_acc[si[b]] = dummy[b] ? prev_acc[si[b]] : max(n[b] - 1, 0) (the KDA replay record) +# rewind[b] = 1 - n[b]; kv_lens[b] += 1 +# bonus[b] = accepted[b, max(n[b] - 1, 0)] +# c = ctx_len[slot[b]]; qpos[b, j] = min(c + n[b], max_ctx) + j (j < block); cpos[b, j] = c + j (j <= K) +# noise[b, 0, :] = embed[bonus[b], :]; noise[b, j > 0, :] = mask_row +# tables[b, i] = max(off[b, i], 0) // divisor; counts[b] = #(off[b, i] >= 0) (the draft pool's block table) +# +# Geometry: G CTAs of 256 threads; CTA c reduces columns [1024 c, 1024 c + 1024) of every logits row (one 16-byte +# load per thread per row) to a (value, index) partial, writes its share of the constant mask rows, and (CTA 0) the +# block table. Each CTA publishes its partials (fence.acq_rel.gpu, then an acq_rel arrival); the last CTA to arrive +# reduces all partials in the same total order, runs the per-request bookkeeping on lanes 0 .. B-1 of warp 0, copies +# the bonus embedding rows, and zeroes the arrival counter for the next call. No CTA waits on another. +# +# Vocabulary-sharded logits (SLOTS > 1): the rank holds bf16 logits of columns [shard_offset, shard_offset + V) of a +# TP column-parallel head. The last CTA's per-row (value, global index) goes to slot [rank] of every rank's Lamport +# buffer through its multicast mapping (fp32 bits, -0.0 sent as +0.0, so no pushed word is the empty word +# 0x80000000); the CTA then polls the SLOTS entries of each row and reduces them in rank order (= index order) with +# the same total order, so every rank gets the argmax of the gathered logits. Call n uses buffer n % 3 (flags[0]) +# and first re-arms all MAX_ROWS rows of buffer (n + 2) % 3 (an earlier call may have pushed more rows than this one), +# whose last readers (this rank's call n - 1) are done and whose next writers (the peers' call n + 2) need this rank's +# call n + 1 pushes; only the last CTA pushes, and its push depends on this rank's arrivals alone, so no rank waits on +# a push that waits on it. +# ============================================================================= +"""CTM speculative-decode acceptance + block-drafter input prep (``trtllm::k3_spec_accept``).""" + +from __future__ import annotations + +import cuda.bindings.driver as cuda_driver +import cutlass +import cutlass.cute as cute +from cutlass._mlir.dialects import llvm as _llvm +from cutlass._mlir.extras import types as _T +from cutlass.cutlass_dsl import dsl_user_op +from cutlass.experimental import primitives as prims + +THREADS = 256 +WARPS = THREADS // 32 +COLS_PER_CTA = THREADS * 4 # one fp32x4 per thread per row +MAX_BATCH = 8 +MAX_ROWS = 64 # B (K + 1) +INT_MAX = 2**31 - 1 +NEG_INF_BITS = -8388608 # 0xFF800000 +FORCE_OFF = 0 # natural acceptance +FORCE_INT = 1 # n = base_total +FORCE_FRAC = 2 # n = min(base_total + (pool[idx] < frac), K + 1) +RNG_POOL_MASK = (1 << 16) - 1 +BUFFERS = 3 +FLAG_WORDS = 4 +MAX_SLOTS = 16 +BUF_WORDS = ( + MAX_ROWS * MAX_SLOTS * 2 +) # int32 words per Lamport buffer: [row][slot] (value bits, index) +EMPTY_WORD = -(2**31) # 0x80000000: fp32 -0.0, never a pushed word +RNG_COUNTER_STRIDE = 6007 +RNG_SLOT_STRIDE = 1009 + + +def supports_columns(vocab: int, slots: int = 1) -> bool: + """``vocab``: the columns this rank reduces (the whole vocabulary, or the rank's shard when ``slots`` > 1, one + exchange slot per rank; an even count, so every row's slots are whole 16-byte vectors).""" + return vocab % COLS_PER_CTA == 0 and (slots == 1 or (slots % 2 == 0 and slots <= MAX_SLOTS)) + + +def supports(vocab: int, batch: int, block: int, drafts: int, hidden: int, slots: int = 1) -> bool: + """``vocab`` and ``slots``: see :func:`supports_columns`.""" + return ( + supports_columns(vocab, slots) + and 0 < batch <= MAX_BATCH + and batch * (drafts + 1) <= MAX_ROWS + and 0 < block <= 16 + and drafts + 1 <= 16 + and hidden % 8 == 0 + ) + + +@dsl_user_op +def _atomic_add_acq_rel(addr_i64, val, *, loc=None, ip=None): + """atom.acq_rel.gpu.global.add.u32, returning the old value.""" + return cutlass.Int32( + _llvm.inline_asm( + _T.i32(), [addr_i64.ir_value(loc=loc, ip=ip), val.ir_value(loc=loc, ip=ip)], + "atom.acq_rel.gpu.global.add.u32 $0, [$1], $2;", "=r,l,r", has_side_effects=True, + is_align_stack=False, asm_dialect=_llvm.AsmDialect.AD_ATT, loc=loc, ip=ip, + ) + ) # fmt: skip + + +def _is_nan(x): + return x != x + + +def _torch_better(value, index, best_value, best_index): + """torch's argmax order (GreaterOrNan): a NaN beats any number, equal values (or two NaNs) go to the lower index, + else the larger value.""" + v_nan = _is_nan(value) + b_nan = _is_nan(best_value) + lower = index < best_index + numbers = (value > best_value) | ((value == best_value) & lower) + if_nan = cutlass.Boolean(cutlass.select_(b_nan, lower, cutlass.Boolean(True))) + if_number = cutlass.Boolean(cutlass.select_(b_nan, cutlass.Boolean(False), numbers)) + return cutlass.Boolean(cutlass.select_(v_nan, if_nan, if_number)) + + +def _keep(best_value, best_index, value, index): + take = _torch_better(value, index, best_value, best_index) + return cutlass.Float32(cutlass.select_(take, value, best_value)), cutlass.Int32( + cutlass.select_(take, index, best_index) + ) + + +def _warp_best(value, index): + for offset in [16, 8, 4, 2, 1]: + other_value = cute.arch.shuffle_sync_bfly(value, offset=offset, mask=-1, mask_and_clamp=31) + other_index = cute.arch.shuffle_sync_bfly(index, offset=offset, mask=-1, mask_and_clamp=31) + value, index = _keep(value, index, other_value, other_index) + return value, index + + +@cute.kernel +def k3_spec_accept_kernel( + logits: cutlass.Array, # fp32 [R * V], R = B (K + 1) + draft: cutlass.Array, # int32 [B * K] + block_off: cutlass.Array, # int32: the draft pool's block offsets (row b at off_base + b * off_stride) + block_counts: cutlass.Array, # int64 [>= B], out + block_tables: cutlass.Array, # int32 [>= B, max_blocks], out + prev_acc: cutlass.Array, # int32 [state slots], in / out + state_idx: cutlass.Array, # int32 [>= B] + dummy: cutlass.Array, # uint8 [>= B] + kv_lens: cutlass.Array, # int32 [>= B], in / out (+= 1) + batch_to_slot: cutlass.Array, # int64 [>= B] + ctx_len: cutlass.Array, # int64 [slots] + embed: cutlass.Array, # int32 words of the bf16 embedding [V_embed, H] + mask_row: cutlass.Array, # int32 words of the bf16 mask embedding [H] + rng_pool: cutlass.Array, # fp32 [65536] (FORCE_FRAC) + rng_counter: cutlass.Array, # int64 [1], in / out (FORCE_FRAC) + accepted: cutlass.Array, # int32 [B * (K + 1)], out + num_acc: cutlass.Array, # int32 [B], out + rewind: cutlass.Array, # int32 [B], out + bonus: cutlass.Array, # int64 [B], out + qpos: cutlass.Array, # int64 [B * block], out + cpos: cutlass.Array, # int64 [B * (K + 1)], out + noise: cutlass.Array, # int32 words of bf16 [B, block, H], out + partials: cutlass.Array, # int32 [R * G * 2] scratch + counter: cutlass.Array, # int32 [1]: CTAs arrived (0 between calls) + buf_uc: cutlass.Array, # sharded: int32 [BUFFERS * BUF_WORDS], this rank's Lamport words + buf_mc: cutlass.Array, # sharded: the same words through the multicast mapping (every rank's) + flags: cutlass.Array, # sharded: int32 [FLAG_WORDS]; [0] the buffer of this call + off_base: cutlass.Int32, + off_stride: cutlass.Int32, + divisor: cutlass.Int32, + max_ctx: cutlass.Int64, + force_total: cutlass.Int32, # min(int(f) + 1, K + 1) + force_frac: cutlass.Float32, + force_mode: cutlass.Int32, # a runtime value: warmup (off) and the captured step (forced) share one build + rank: cutlass.Int32, # sharded: this rank (its slots are rank * push_copies + c) + shard_offset: cutlass.Int32, # sharded: the global index of this rank's first column + grid: cutlass.Constexpr[int], + vocab: cutlass.Constexpr[int], + batch: cutlass.Constexpr[int], + drafts: cutlass.Constexpr[int], # K + block: cutlass.Constexpr[int], # query block width + hidden: cutlass.Constexpr[int], + max_blocks: cutlass.Constexpr[int], + slots: cutlass.Constexpr[ + int + ], # 1: the whole vocabulary is here (fp32); > 1: bf16 shards, one slot per rank + push_copies: cutlass.Constexpr[ + int + ], # slots each rank fills (tests emulate a larger group with > 1) +): + rows = batch * (drafts + 1) + row_words = hidden // 2 # int32 words per embedding row + row_vecs = row_words // 4 # 16-byte vectors per embedding row + tx, _, _ = cute.arch.thread_idx() + cta, _, _ = cute.arch.block_idx() + warp = cute.arch.make_warp_uniform(cute.arch.warp_idx()) + lane = tx % cutlass.Int32(32) + + s_best = cutlass.Array( + cutlass.Int32, rows * WARPS * 2, space=cutlass.AddressSpace.smem, alignment=16 + ) + s_tok = cutlass.Array(cutlass.Int32, MAX_ROWS, space=cutlass.AddressSpace.smem, alignment=16) + s_bonus = cutlass.Array(cutlass.Int32, MAX_BATCH, space=cutlass.AddressSpace.smem, alignment=16) + # [b]: the first query position and the context length of each request (for the per-element stores). + s_now = cutlass.Array(cutlass.Int64, MAX_BATCH, space=cutlass.AddressSpace.smem, alignment=16) + s_ctx = cutlass.Array(cutlass.Int64, MAX_BATCH, space=cutlass.AddressSpace.smem, alignment=16) + s_last = cutlass.Array(cutlass.Int32, 4, space=cutlass.AddressSpace.smem, alignment=16) + # The last CTA's inputs of the per-request bookkeeping, read while the rows reduce: [b] kv_len, state index, + # prev_acc[state], dummy; the drafts [b][j]. + s_in = cutlass.Array( + cutlass.Int32, MAX_BATCH * 4, space=cutlass.AddressSpace.smem, alignment=16 + ) + s_draft = cutlass.Array(cutlass.Int32, MAX_ROWS, space=cutlass.AddressSpace.smem, alignment=16) + # [row]: the (value bits, index) of the row over this rank's columns (the exchange's payload). + s_row = cutlass.Array( + cutlass.Int32, MAX_ROWS * 2, space=cutlass.AddressSpace.smem, alignment=16 + ) + + # Nothing this kernel reads is a weight: all of it comes after the grid dependency. + if warp == 0: + prims.griddepcontrol(prims.GridDepAction.LAUNCH_DEPENDENTS) + prims.griddepcontrol(prims.GridDepAction.WAIT) + + # Per-row partial argmax over this CTA's 1024 columns. + col0 = cta * cutlass.Int32(COLS_PER_CTA) + tx * cutlass.Int32(4) + for r in cutlass.range_constexpr(rows): + vals = [] + if cutlass.const_expr(slots > 1): + # bf16 pairs as int32 words; bf16 -> fp32 is exact (the high half of the word). + lw = logits.load( + idx=(cutlass.Int32(r * vocab) + col0) // cutlass.Int32(2), + vector_size=2, + alignment=8, + ) + for q in cutlass.range_constexpr(2): + word = cutlass.Int32(lw[q]) + vals.append(cutlass.Int32(word << cutlass.Int32(16)).bitcast(cutlass.Float32)) + vals.append(cutlass.Int32(word & cutlass.Int32(-65536)).bitcast(cutlass.Float32)) + else: + v = logits.load(idx=cutlass.Int32(r * vocab) + col0, vector_size=4, alignment=16) + for q in cutlass.range_constexpr(4): + vals.append(cutlass.Float32(v[q])) + best_value = vals[0] + best_index = col0 + for q in cutlass.range_constexpr(1, 4): + best_value, best_index = _keep(best_value, best_index, vals[q], col0 + cutlass.Int32(q)) + best_value, best_index = _warp_best(best_value, best_index) + if lane == 0: + s_best.store((best_value.bitcast(cutlass.Int32), best_index), idx=cutlass.Int32(r * WARPS * 2) + warp * 2, + alignment=8) # fmt: skip + + # The constant mask rows of the block (j >= 1), this CTA's share of their 16-byte vectors. + n_mask_vecs = batch * (block - 1) * row_vecs + v_id = cta * cutlass.Int32(THREADS) + tx + while v_id < cutlass.Int32(n_mask_vecs): + b_j = v_id // cutlass.Int32(row_vecs) # (b, j - 1) row + within = v_id - b_j * cutlass.Int32(row_vecs) + b = b_j // cutlass.Int32(block - 1) + j = b_j - b * cutlass.Int32(block - 1) + cutlass.Int32(1) + m = mask_row.load(idx=within * cutlass.Int32(4), vector_size=4, alignment=16) + noise.store((cutlass.Int32(m[0]), cutlass.Int32(m[1]), cutlass.Int32(m[2]), cutlass.Int32(m[3])), + idx=((b * cutlass.Int32(block) + j) * cutlass.Int32(row_vecs) + within) * cutlass.Int32(4), + alignment=16) # fmt: skip + v_id += cutlass.Int32(grid * THREADS) + + # The draft pool's block table (CTA 0): tables[b, i] = max(off, 0) // divisor, counts[b] = #(off >= 0). + if cta == 0: + for b in cutlass.range_constexpr(batch): + present = cutlass.Int32(0) + i = tx + while i < cutlass.Int32(max_blocks): + e = block_off.load(idx=off_base + cutlass.Int32(b) * off_stride + i) + block_tables.store(cutlass.Int32(cutlass.select_(e > cutlass.Int32(0), e, cutlass.Int32(0))) // divisor, + idx=cutlass.Int32(b * max_blocks) + i) # fmt: skip + present += cutlass.Int32( + cutlass.select_(e >= cutlass.Int32(0), cutlass.Int32(1), cutlass.Int32(0)) + ) + i += cutlass.Int32(THREADS) + for offset in [16, 8, 4, 2, 1]: + present += cute.arch.shuffle_sync_bfly( + present, offset=offset, mask=-1, mask_and_clamp=31 + ) + if lane == 0: + s_tok.store(present, idx=cutlass.Int32(b * WARPS) + warp) + prims.barrier_cta_sync(0) + if tx < cutlass.Int32(batch): + total = cutlass.Int32(0) + for w in cutlass.range_constexpr(WARPS): + total += s_tok.load(idx=tx * cutlass.Int32(WARPS) + cutlass.Int32(w)) + block_counts.store(cutlass.Int64(total), idx=tx) + + prims.barrier_cta_sync(0) + # This CTA's partial per row, then its arrival (the partials first: fence.acq_rel.gpu + the acq_rel atom). + if tx < cutlass.Int32(rows): + value = cutlass.Int32(NEG_INF_BITS).bitcast(cutlass.Float32) + index = cutlass.Int32(INT_MAX) + for w in cutlass.range_constexpr(WARPS): + pair = s_best.load( + idx=tx * cutlass.Int32(WARPS * 2) + cutlass.Int32(w * 2), vector_size=2, alignment=8 + ) + value, index = _keep( + value, + index, + cutlass.Int32(pair[0]).bitcast(cutlass.Float32), + cutlass.Int32(pair[1]), + ) + partials.store((value.bitcast(cutlass.Int32), index), idx=(tx * cutlass.Int32(grid) + cta) * cutlass.Int32(2), + alignment=8) # fmt: skip + cute.arch.fence_acq_rel_gpu() + prims.barrier_cta_sync(0) + if tx == 0: + arrived = _atomic_add_acq_rel(counter.data_ptr(0).toint(), cutlass.Int32(1)) + s_last.store( + cutlass.Int32( + cutlass.select_( + arrived == cutlass.Int32(grid - 1), cutlass.Int32(1), cutlass.Int32(0) + ) + ), + idx=0, + ) + prims.barrier_cta_sync(0) + if s_last.load(idx=0) != cutlass.Int32(0): + cute.arch.fence_acq_rel_gpu() + # The bookkeeping's inputs (none depends on the argmax), in flight while the rows reduce: one element per + # thread for the drafts; lane b of warp 0 for request b's state. + if tx < cutlass.Int32(batch * drafts): + s_draft.store(cutlass.Int32(draft.load(idx=tx)), idx=tx) + if tx < cutlass.Int32(batch): + si = cutlass.Int32(state_idx.load(idx=tx)) + s_in.store(cutlass.Int32(kv_lens.load(idx=tx)), idx=tx * cutlass.Int32(4)) + s_in.store(si, idx=tx * cutlass.Int32(4) + cutlass.Int32(1)) + s_in.store( + cutlass.Int32(prev_acc.load(idx=si)), idx=tx * cutlass.Int32(4) + cutlass.Int32(2) + ) + s_in.store( + cutlass.Int32(dummy.load(idx=tx)), idx=tx * cutlass.Int32(4) + cutlass.Int32(3) + ) + s_ctx.store( + cutlass.Int64(ctx_len.load(idx=cutlass.Int32(batch_to_slot.load(idx=tx)))), idx=tx + ) + # The target token of each row: warp w reduces rows w, w + 8, ... over the G partials (every partial load of + # the lane in flight at once). + r_w = warp + while r_w < cutlass.Int32(rows): + pairs = [] + for i in cutlass.range_constexpr((grid + 31) // 32): + p = lane + cutlass.Int32(i * 32) + p = cutlass.Int32(cutlass.select_(p < cutlass.Int32(grid), p, cutlass.Int32(0))) + pairs.append(partials.load(idx=(r_w * cutlass.Int32(grid) + p) * cutlass.Int32(2), vector_size=2, + alignment=8, is_volatile=True)) # fmt: skip + value = cutlass.Int32(NEG_INF_BITS).bitcast(cutlass.Float32) + index = cutlass.Int32(INT_MAX) + for i in cutlass.range_constexpr((grid + 31) // 32): + live = lane + cutlass.Int32(i * 32) < cutlass.Int32(grid) + v_i, i_i = _keep( + value, + index, + cutlass.Int32(pairs[i][0]).bitcast(cutlass.Float32), + cutlass.Int32(pairs[i][1]), + ) + value = cutlass.Float32(cutlass.select_(live, v_i, value)) + index = cutlass.Int32(cutlass.select_(live, i_i, index)) + value, index = _warp_best(value, index) + if lane == 0: + s_tok.store(index, idx=r_w) + if cutlass.const_expr(slots > 1): + s_row.store((value.bitcast(cutlass.Int32), index + shard_offset), idx=r_w * cutlass.Int32(2), + alignment=8) # fmt: skip + r_w += cutlass.Int32(WARPS) + prims.barrier_cta_sync(0) + if cutlass.const_expr(slots > 1): + # The vocabulary shards' exchange (see the header): re-arm the buffer before last, push, poll, reduce. + cur = flags.load(idx=0, is_volatile=True) + cur_base = cur * cutlass.Int32(BUF_WORDS) + dirty_base = ((cur + cutlass.Int32(2)) % cutlass.Int32(BUFFERS)) * cutlass.Int32( + BUF_WORDS + ) + empty = cutlass.Int32(EMPTY_WORD) + w_clear = tx * cutlass.Int32(4) + while w_clear < cutlass.Int32(MAX_ROWS * slots * 2): + buf_uc.store((empty, empty, empty, empty), idx=dirty_base + w_clear, alignment=16) + w_clear += cutlass.Int32(THREADS * 4) + if tx < cutlass.Int32(rows): + mine = s_row.load(idx=tx * cutlass.Int32(2), vector_size=2, alignment=8) + mine_bits = cutlass.Int32(mine[0]) + mine_bits = cutlass.Int32( + cutlass.select_(mine_bits == empty, cutlass.Int32(0), mine_bits) + ) + for cp in cutlass.range_constexpr(push_copies): + slot = rank * cutlass.Int32(push_copies) + cutlass.Int32(cp) + push_at = cur_base + (tx * cutlass.Int32(slots) + slot) * cutlass.Int32(2) + buf_mc.store((mine_bits, cutlass.Int32(mine[1])), idx=push_at, alignment=8) + # Drain the posted multicast store now (the polls below issue no release that would). + prims.fence_acq_rel(prims.MemScope.CLUSTER) + region = cur_base + tx * cutlass.Int32(slots * 2) + g_index = cutlass.Int32(INT_MAX) + pending = cutlass.Boolean(True) + while pending: + p_value = cutlass.Int32(NEG_INF_BITS).bitcast(cutlass.Float32) + p_index = cutlass.Int32(INT_MAX) + p_pending = cutlass.Boolean(False) + got = [] + for v4 in cutlass.range_constexpr((slots * 2 + 3) // 4): + got.append(buf_uc.load(idx=region + cutlass.Int32(v4 * 4), vector_size=4, alignment=16, + is_volatile=True)) # fmt: skip + for sl in cutlass.range_constexpr(slots): + vw = cutlass.Int32(got[(2 * sl) // 4][(2 * sl) % 4]) + iw = cutlass.Int32(got[(2 * sl + 1) // 4][(2 * sl + 1) % 4]) + p_pending = p_pending | (vw == empty) | (iw == empty) + p_value, p_index = _keep(p_value, p_index, vw.bitcast(cutlass.Float32), iw) + g_index = p_index + pending = p_pending + s_tok.store(g_index, idx=tx) + prims.barrier_cta_sync(0) + if tx == 0: + flags.store((cur + cutlass.Int32(1)) % cutlass.Int32(BUFFERS), idx=0) + # Per-request bookkeeping on lane b of warp 0. + if tx < cutlass.Int32(batch): + b = tx + base_row = b * cutlass.Int32(drafts + 1) + n = cutlass.Int32(1) + run = cutlass.Boolean(True) + # Volatile: a run of adjacent loads at a dynamic offset is otherwise merged into 16-byte vectors, which + # the rows of b >= 1 misalign. + for j in cutlass.range_constexpr(drafts + 1): + t = s_tok.load(idx=base_row + cutlass.Int32(j), is_volatile=True) + if cutlass.const_expr(j < drafts): + run = run & ( + s_draft.load( + idx=b * cutlass.Int32(drafts) + cutlass.Int32(j), is_volatile=True + ) + == t + ) + n += cutlass.Int32(cutlass.select_(run, cutlass.Int32(1), cutlass.Int32(0))) + if force_mode == cutlass.Int32(FORCE_INT): + n = force_total + if force_mode == cutlass.Int32(FORCE_FRAC): + step = rng_counter.load(idx=0) + cutlass.Int64(1) + pool_idx = ( + step * cutlass.Int64(RNG_COUNTER_STRIDE) + + cutlass.Int64(b) * cutlass.Int64(RNG_SLOT_STRIDE) + ) & cutlass.Int64(RNG_POOL_MASK) + extra = cutlass.Int32(cutlass.select_(rng_pool.load(idx=cutlass.Int32(pool_idx)) < force_frac, + cutlass.Int32(1), cutlass.Int32(0))) # fmt: skip + n = force_total + extra + n = cutlass.Int32( + cutlass.select_(n > cutlass.Int32(drafts + 1), cutlass.Int32(drafts + 1), n) + ) + num_acc.store(n, idx=b) + rewind.store(cutlass.Int32(1) - n, idx=b) + kv_lens.store( + cutlass.Int32(s_in.load(idx=b * cutlass.Int32(4))) + cutlass.Int32(1), idx=b + ) + # The KDA replay record of the accepted drafts. + si = cutlass.Int32(s_in.load(idx=b * cutlass.Int32(4) + cutlass.Int32(1))) + keep_prev = cutlass.Int32( + s_in.load(idx=b * cutlass.Int32(4) + cutlass.Int32(3)) + ) != cutlass.Int32(0) + drafted = n - cutlass.Int32(1) + drafted = cutlass.Int32( + cutlass.select_(drafted < cutlass.Int32(0), cutlass.Int32(0), drafted) + ) + prev_acc.store( + cutlass.Int32( + cutlass.select_( + keep_prev, + cutlass.Int32(s_in.load(idx=b * cutlass.Int32(4) + cutlass.Int32(2))), + drafted, + ) + ), + idx=si, + ) + # The drafter's inputs. + last = n - cutlass.Int32(1) + last = cutlass.Int32(cutlass.select_(last < cutlass.Int32(0), cutlass.Int32(0), last)) + tok = s_tok.load(idx=base_row + last) + bonus.store(cutlass.Int64(tok), idx=b) + s_bonus.store(tok, idx=b) + c = cutlass.Int64(s_ctx.load(idx=b)) + now = c + cutlass.Int64(n) + now = cutlass.Int64(cutlass.select_(now > max_ctx, max_ctx, now)) + s_now.store(now, idx=b) + prims.barrier_cta_sync(0) + # After the barrier: every lane above has read the counter (the lanes of warp 0 need not run in step). + if force_mode == cutlass.Int32(FORCE_FRAC): + if tx == 0: + rng_counter.store(rng_counter.load(idx=0) + cutlass.Int64(1), idx=0) + # One element per thread of accepted, qpos and cpos (a thread storing a run of adjacent elements at a dynamic + # offset gets them merged into wider stores that the odd rows misalign). + if tx < cutlass.Int32(rows): + accepted.store(s_tok.load(idx=tx), idx=tx) + b_c = tx // cutlass.Int32(drafts + 1) + cpos.store( + s_ctx.load(idx=b_c) + cutlass.Int64(tx - b_c * cutlass.Int32(drafts + 1)), idx=tx + ) + if tx < cutlass.Int32(batch * block): + b_q = tx // cutlass.Int32(block) + qpos.store(s_now.load(idx=b_q) + cutlass.Int64(tx - b_q * cutlass.Int32(block)), idx=tx) + # noise[b, 0, :] = embed[bonus[b], :]: every load of the thread in flight, then the stores. + embed_vecs = [] + for i in cutlass.range_constexpr((batch * row_vecs + THREADS - 1) // THREADS): + v_id = tx + cutlass.Int32(i * THREADS) + v_id = cutlass.Int32( + cutlass.select_(v_id < cutlass.Int32(batch * row_vecs), v_id, cutlass.Int32(0)) + ) + b = v_id // cutlass.Int32(row_vecs) + within = v_id - b * cutlass.Int32(row_vecs) + tok = cutlass.Int32(s_bonus.load(idx=b)) + embed_vecs.append(embed.load(idx=(tok * cutlass.Int32(row_vecs) + within) * cutlass.Int32(4), vector_size=4, + alignment=16)) # fmt: skip + for i in cutlass.range_constexpr((batch * row_vecs + THREADS - 1) // THREADS): + v_id = tx + cutlass.Int32(i * THREADS) + if v_id < cutlass.Int32(batch * row_vecs): + b = v_id // cutlass.Int32(row_vecs) + within = v_id - b * cutlass.Int32(row_vecs) + e = embed_vecs[i] + noise.store((cutlass.Int32(e[0]), cutlass.Int32(e[1]), cutlass.Int32(e[2]), cutlass.Int32(e[3])), + idx=((b * cutlass.Int32(block)) * cutlass.Int32(row_vecs) + within) * cutlass.Int32(4), + alignment=16) # fmt: skip + if tx == 0: + counter.store(cutlass.Int32(0), idx=0) + + +@cute.jit +def k3_spec_accept( + logits: cute.Tensor, + draft: cute.Tensor, + block_off: cute.Tensor, + block_counts: cute.Tensor, + block_tables: cute.Tensor, + prev_acc: cute.Tensor, + state_idx: cute.Tensor, + dummy: cute.Tensor, + kv_lens: cute.Tensor, + batch_to_slot: cute.Tensor, + ctx_len: cute.Tensor, + embed: cute.Tensor, + mask_row: cute.Tensor, + rng_pool: cute.Tensor, + rng_counter: cute.Tensor, + accepted: cute.Tensor, + num_acc: cute.Tensor, + rewind: cute.Tensor, + bonus: cute.Tensor, + qpos: cute.Tensor, + cpos: cute.Tensor, + noise: cute.Tensor, + partials: cute.Tensor, + counter: cute.Tensor, + buf_uc: cute.Tensor, + buf_mc: cute.Tensor, + flags: cute.Tensor, + off_base: cutlass.Int32, + off_stride: cutlass.Int32, + divisor: cutlass.Int32, + max_ctx: cutlass.Int64, + force_total: cutlass.Int32, + force_frac: cutlass.Float32, + force_mode: cutlass.Int32, + rank: cutlass.Int32, + shard_offset: cutlass.Int32, + grid: cutlass.Constexpr[int], + vocab: cutlass.Constexpr[int], + batch: cutlass.Constexpr[int], + drafts: cutlass.Constexpr[int], + block: cutlass.Constexpr[int], + hidden: cutlass.Constexpr[int], + max_blocks: cutlass.Constexpr[int], + slots: cutlass.Constexpr[int], + push_copies: cutlass.Constexpr[int], + use_pdl: cutlass.Constexpr[bool], + stream: cuda_driver.CUstream, +) -> None: + k3_spec_accept_kernel( + logits, draft, block_off, block_counts, block_tables, prev_acc, state_idx, dummy, kv_lens, batch_to_slot, + ctx_len, embed, mask_row, rng_pool, rng_counter, accepted, num_acc, rewind, bonus, qpos, cpos, noise, partials, + counter, buf_uc, buf_mc, flags, off_base, off_stride, divisor, max_ctx, force_total, force_frac, force_mode, + rank, shard_offset, grid, vocab, batch, drafts, block, hidden, max_blocks, slots, push_copies, + ).launch( + grid=[grid, 1, 1], + block=[THREADS, 1, 1], + stream=stream, + use_pdl=use_pdl, + ) # fmt: skip diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_spec_accept/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_spec_accept/op.py new file mode 100644 index 000000000000..82bea70f884f --- /dev/null +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_spec_accept/op.py @@ -0,0 +1,309 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``trtllm::k3_spec_accept``: one decode step's speculative acceptance and the block drafter's inputs in one CTM +kernel (greedy strict acceptance of B generation requests, K drafts each; see k3_spec_accept_kernel.py). + +Bit-identical to the Python path it replaces: ``_refresh_ctx_block_tables``, ``_sample_and_accept_draft_tokens_base`` +(with the forced-acceptance override), the KDA replay record of ``update_mamba_states``, the kv_lens update of +``_prepare_kv_for_draft_forward`` and ``prepare_1st_drafter_inputs`` up to the fc (bonus, positions, noise embedding). +With a :func:`workspace`, the target logits may stay vocabulary-sharded (this rank's bf16 columns of a TP +column-parallel head): the ranks exchange their row maxima over a multicast Lamport buffer instead of all-gathering +the logits, with the same argmax. Compiled on the first call per shape, which must happen outside CUDA-graph capture. +""" + +from __future__ import annotations + +import importlib.util +import os +import sys +import threading +from typing import Dict, List, Optional + +import torch + +_lock = threading.Lock() +_compiled: Dict[tuple, object] = {} +_scratch: Dict[tuple, tuple] = {} +_modules: Dict[str, object] = {} +_workspaces: Dict[object, dict] = {} + + +def _kernel_module(): + mod = _modules.get("kernel") + if mod is None: + path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "k3_spec_accept_kernel.py") + spec = importlib.util.spec_from_file_location(f"{__name__}_k3_spec_accept_kernel", path) + mod = importlib.util.module_from_spec(spec) + sys.modules[spec.name] = mod + spec.loader.exec_module(mod) + _modules["kernel"] = mod + return mod + + +def supports(vocab: int, batch: int, block: int, drafts: int, hidden: int, slots: int = 1) -> bool: + """``vocab``: the columns of the logits this rank holds (its shard when ``slots`` > 1).""" + return _kernel_module().supports(vocab, batch, block, drafts, hidden, slots) + + +def supports_columns(vocab: int, slots: int = 1) -> bool: + """Whether the kernel splits ``vocab`` columns (a rank's shard of a ``slots``-rank group when ``slots`` > 1).""" + return _kernel_module().supports_columns(vocab, slots) + + +def existing_workspace(mapping, push_copies: int = 1) -> Optional[dict]: + """The group's :func:`workspace` if it has been allocated (inside CUDA-graph capture it cannot be).""" + return _workspaces.get((mapping, push_copies)) + + +def workspace(mapping, push_copies: int = 1) -> dict: + """This TP group's Lamport buffers for the sharded argmax exchange (3 rotating buffers, every word 0x80000000 + when armed) behind one multicast mapping, and the flag words. Collective on first use: every rank of the group + must make the first call at the same point, outside CUDA-graph capture. ``push_copies`` > 1 (tests only) gives + every rank that many slots, to emulate a larger group.""" + key = (mapping, push_copies) + ws = _workspaces.get(key) + if ws is not None: + return ws + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "k3_spec_accept: the workspace must be allocated outside CUDA-graph capture" + ) + from tensorrt_llm._torch.distributed.ops import ( + _get_mnnvl_workspace_comm, + _make_mnnvl_mcast_buffer, + _mnnvl_workspace_all_succeeded, + ) + + kern = _kernel_module() + slots = mapping.tp_size * push_copies + if not (slots % 2 == 0 and slots <= kern.MAX_SLOTS): + raise ValueError( + f"k3_spec_accept: {slots} exchange slots (an even count up to {kern.MAX_SLOTS})" + ) + comm = _get_mnnvl_workspace_comm(mapping) + use_fabric_handle = ( + os.environ.get("TRTLLM_FORCE_MNNVL_AR", "0") == "1" or mapping.is_multi_node() + ) + error: Optional[Exception] = None + try: + words = kern.BUFFERS * kern.BUF_WORDS + handle = _make_mnnvl_mcast_buffer(comm, words * 4, mapping, use_fabric_handle) + uc = handle.get_uc_buffer(mapping.tp_rank, (words,), torch.int32, 0) + mc = handle.get_mc_buffer((words,), torch.int32, 0) + with torch.inference_mode(): + uc.fill_(kern.EMPTY_WORD) + flags = torch.zeros(kern.FLAG_WORDS, dtype=torch.int32, device=uc.device) + torch.cuda.synchronize() + ws = dict(handle=handle, comm=comm, uc=uc, mc=mc, flags=flags, rank=mapping.tp_rank, slots=slots, + push_copies=push_copies) # fmt: skip + except Exception as exc: # noqa: BLE001 -- reported to every rank below, then re-raised + error = exc + # Also the barrier that keeps any rank from pushing into a peer's buffer before the peer has armed it. + if not _mnnvl_workspace_all_succeeded(comm, error is None): + raise RuntimeError("k3_spec_accept Lamport buffers failed on at least one rank") from error + _workspaces[key] = ws + return ws + + +def force_mode(force: float, drafts: int): + """(mode, base total, fraction) of ``_apply_force_accepted_tokens`` for the forced value ``force`` (0: off).""" + kern = _kernel_module() + if force == 0.0: + return kern.FORCE_OFF, 0, 0.0 + int_part = int(force) + frac = force - int_part + total = min(int_part + 1, drafts + 1) + if frac > 0.0 and total < drafts + 1: + return kern.FORCE_FRAC, total, frac + return kern.FORCE_INT, total, 0.0 + + +def _arg(t: torch.Tensor): + from cutlass.cute.runtime import from_dlpack + + return from_dlpack(t.detach(), assumed_align=16).mark_layout_dynamic(leading_dim=0) + + +@torch.library.custom_op( + "trtllm::k3_spec_accept", + mutates_args=("block_counts", "block_tables", "prev_acc", "kv_lens", "rng_counter"), +) +def k3_spec_accept( + logits: torch.Tensor, + draft: torch.Tensor, + block_off: torch.Tensor, + pool_idx: int, + divisor: int, + block_counts: torch.Tensor, + block_tables: torch.Tensor, + prev_acc: torch.Tensor, + state_idx: torch.Tensor, + dummy: torch.Tensor, + kv_lens: torch.Tensor, + batch_to_slot: torch.Tensor, + ctx_len: torch.Tensor, + max_ctx: int, + embed: torch.Tensor, + mask_row: torch.Tensor, + rng_pool: torch.Tensor, + rng_counter: torch.Tensor, + force: float, + block: int, + ws_uc: Optional[torch.Tensor] = None, + ws_mc: Optional[torch.Tensor] = None, + ws_flags: Optional[torch.Tensor] = None, + rank: int = 0, + slots: int = 1, + push_copies: int = 1, + shard_offset: int = 0, +) -> List[torch.Tensor]: + """Returns ``[accepted [B, K + 1] int32, num_accepted [B] int32, rewind [B] int32, bonus [B] int64, + query_positions [B, block] int64, ctx_positions [B, K + 1] int64, noise [B, block, H] bf16]`` and updates + ``block_counts[:B]`` / ``block_tables[:B]`` (the draft pool's block table decoded from ``block_off[pool_idx, :B, + 0]`` / ``divisor``), ``prev_acc`` (the KDA replay record), ``kv_lens[:B]`` (+ 1) and ``rng_counter`` (forced + fractional acceptance). ``logits``: fp32 [B (K + 1), V]; or, with ``slots`` > 1 and a :func:`workspace` + (``ws_*``, ``rank`` and ``push_copies`` from it), this rank's bf16 columns [B (K + 1), V / TP] starting at + ``shard_offset``. ``draft``: int32 [B, K].""" + import cuda.bindings.driver as cuda_driver + import cutlass.cute as cute + + kern = _kernel_module() + batch, drafts = draft.shape + rows, vocab = logits.shape + hidden = embed.shape[1] + sharded = slots > 1 + logits_dtype = torch.bfloat16 if sharded else torch.float32 + if ( + logits.dtype != logits_dtype + or rows != batch * (drafts + 1) + or draft.dtype != torch.int32 + or not kern.supports(vocab, batch, block, drafts, hidden, slots) + or (sharded and (ws_uc is None or ws_mc is None or ws_flags is None)) + ): + raise ValueError( + f"k3_spec_accept: unsupported call: logits {tuple(logits.shape)} {logits.dtype}, draft " + f"{tuple(draft.shape)} {draft.dtype}, block {block}, hidden {hidden}, slots {slots} (fp32 [B (K + 1), V], " + f"or bf16 shards with a workspace; V % {kern.COLS_PER_CTA} == 0, int32 [B <= {kern.MAX_BATCH}, K], " + f"block <= 16)" + ) # fmt: skip + if ( + block_off.dtype != torch.int32 + or block_off.dim() != 4 + or not block_off.is_contiguous() + or block_tables.dtype != torch.int32 + or block_tables.shape[1] != block_off.shape[3] + or block_counts.dtype != torch.int64 + or prev_acc.dtype != torch.int32 + or state_idx.dtype != torch.int32 + or kv_lens.dtype != torch.int32 + or batch_to_slot.dtype != torch.int64 + or ctx_len.dtype != torch.int64 + or embed.dtype != torch.bfloat16 + or mask_row.dtype != torch.bfloat16 + or mask_row.numel() != hidden + ): + raise ValueError("k3_spec_accept: unexpected state tensor dtypes or layouts") # fmt: skip + device = logits.device + grid = vocab // kern.COLS_PER_CTA + scratch = _scratch.get((device, grid)) + if scratch is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "trtllm::k3_spec_accept must run once outside CUDA-graph capture first" + ) + scratch = _scratch[(device, grid)] = ( + torch.empty(kern.MAX_ROWS * grid * 2, dtype=torch.int32, device=device), + torch.zeros(1, dtype=torch.int32, device=device), + ) + mode, total, frac = force_mode(force, drafts) + accepted = torch.empty(batch, drafts + 1, dtype=torch.int32, device=device) + num_acc = torch.empty(batch, dtype=torch.int32, device=device) + rewind = torch.empty(batch, dtype=torch.int32, device=device) + bonus = torch.empty(batch, dtype=torch.int64, device=device) + qpos = torch.empty(batch, block, dtype=torch.int64, device=device) + cpos = torch.empty(batch, drafts + 1, dtype=torch.int64, device=device) + noise = torch.empty(batch, block, hidden, dtype=torch.bfloat16, device=device) + logits_arg = logits.contiguous().view(-1) + if sharded: + logits_arg = logits_arg.view(torch.int32) + else: + ws_uc = ws_mc = ws_flags = scratch[1] + args = ( + _arg(logits_arg), + _arg(draft.contiguous().view(-1)), + _arg(block_off.view(-1)), + _arg(block_counts.view(-1)), + _arg(block_tables.view(-1)), + _arg(prev_acc.view(-1)), + _arg(state_idx.view(-1)), + _arg(dummy.view(-1).view(torch.uint8)), + _arg(kv_lens.view(-1)), + _arg(batch_to_slot.view(-1)), + _arg(ctx_len.view(-1)), + _arg(embed.view(-1).view(torch.int32)), + _arg(mask_row.contiguous().view(-1).view(torch.int32)), + _arg(rng_pool.view(-1)), + _arg(rng_counter.view(-1)), + _arg(accepted.view(-1)), + _arg(num_acc.view(-1)), + _arg(rewind.view(-1)), + _arg(bonus.view(-1)), + _arg(qpos.view(-1)), + _arg(cpos.view(-1)), + _arg(noise.view(-1).view(torch.int32)), + _arg(scratch[0]), + _arg(scratch[1]), + _arg(ws_uc), + _arg(ws_mc), + _arg(ws_flags), + ) + n_seq, max_blocks = block_off.shape[1], block_off.shape[3] + scalars = (int(pool_idx) * n_seq * 2 * max_blocks, 2 * max_blocks, int(divisor), int(max_ctx), int(total), + float(frac), int(mode), int(rank), int(shard_offset)) # fmt: skip + use_pdl = os.environ.get("TRTLLM_ENABLE_PDL", "1") == "1" + consts = (grid, vocab, batch, drafts, block, hidden, max_blocks, int(slots), int(push_copies)) + key = consts + (use_pdl,) + fn = _compiled.get(key) + if fn is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "trtllm::k3_spec_accept must run once per shape outside CUDA-graph capture first" + ) + with _lock: + fn = _compiled.get(key) + if fn is None: + fn = _compiled[key] = cute.compile( + kern.k3_spec_accept, *args, *scalars, *consts, use_pdl, + cuda_driver.CUstream(torch.cuda.current_stream(device).cuda_stream), + ) # fmt: skip + fn(*args, *scalars, cuda_driver.CUstream(torch.cuda.current_stream(device).cuda_stream)) + return [accepted, num_acc, rewind, bonus, qpos, cpos, noise] + + +@k3_spec_accept.register_fake +def _(logits, draft, block_off, pool_idx, divisor, block_counts, block_tables, prev_acc, state_idx, dummy, kv_lens, + batch_to_slot, ctx_len, max_ctx, embed, mask_row, rng_pool, rng_counter, force, block, ws_uc=None, ws_mc=None, + ws_flags=None, rank=0, slots=1, push_copies=1, shard_offset=0): # fmt: skip + batch, drafts = draft.shape + hidden = embed.shape[1] + return [ + draft.new_empty((batch, drafts + 1)), + draft.new_empty((batch,)), + draft.new_empty((batch,)), + draft.new_empty((batch,), dtype=torch.int64), + draft.new_empty((batch, block), dtype=torch.int64), + draft.new_empty((batch, drafts + 1), dtype=torch.int64), + embed.new_empty((batch, block, hidden)), + ] diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctx_kv.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctx_kv.py new file mode 100644 index 000000000000..177009b19ba4 --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctx_kv.py @@ -0,0 +1,650 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``trtllm::k3_ctx_kv`` (the block drafter's context K/V of a decode step) against the DFlashWorker Python path, for +B requests of K + 1 tokens (N = B (K + 1) <= 64), at the in-model TP16 shape (5 drafter layers, 1 KV head, hidden +7168; TP4's 4 KV heads where the kernel takes it). + +The Python path, as DFlashWorker runs it (modeling_dflash precompute_context_kv, dflash _store_context_kv_paged): +F.linear with the stacked K/V weight, K/V .contiguous(), F.rms_norm + k_norm, flashinfer NeoX RoPE in place, the write +mask, column clamps, one flashinfer paged append per layer, ctx_len += num_accepted (clamped), num_ctx = min(ctx_len, +counts page - block). + +Every split with mixed per-request lengths (columns crossing a page), accepted counts and slots / table rows, a column +clamp at the allocation, page-major and arena pools: every written pool row against the Python path and an fp32 +reference (the kernel's error must be the Python path's), masked rows zero, the pool's other elements untouched, +ctx_len and num_ctx exact, reruns bit-identical; CUDA-graph replays with rewritten inputs. + +Batch-1 identity against the unmodified kernel (N <= 8): set ``K3_BASE_TRTLLM`` to an unmodified ``tensorrt_llm`` +package directory. Timing: ``python3 test_k3_ctx_kv.py time [--base]``; error table: ``python3 test_k3_ctx_kv.py +report``. +""" + +import importlib.util +import os +import statistics +import sys + +import pytest +import torch +import torch.nn.functional as F + +HIDDEN = 7168 +HEAD = 64 +LAYERS = 5 +PAGE = 64 +EPS = 1e-6 +MAX_POS = 8192 +MAX_CTX = 4096 +THETA = 1.0e6 +BLOCK = 8 + +# (B requests, K + 1 tokens each): every split with B (K + 1) <= 8, DSpark's B x 8, and other draft lengths. +SPLITS = ( + [(1, 8), (2, 4), (4, 2), (8, 1), (1, 1), (2, 1), (3, 1), (4, 1), (5, 1)] + + [(b, 8) for b in range(2, 9)] + + [(8, 7), (2, 7), (8, 4), (3, 4)] +) + + +def _sm100() -> bool: + return torch.cuda.is_available() and torch.cuda.get_device_capability()[0] == 10 + + +pytestmark = pytest.mark.skipif(not _sm100(), reason="needs SM100 (tcgen05, TMA, clusters)") + + +def _op(): + import tensorrt_llm # noqa: F401 (registers the trtllm ops) + from tensorrt_llm._torch.cute_dsl_kernels.k3_ctx_kv import op + + return op + + +def cos_sin_cache(device): + inv = 1.0 / (THETA ** (torch.arange(0, HEAD, 2, dtype=torch.float64, device=device) / HEAD)) + ang = torch.arange(MAX_POS, dtype=torch.float64, device=device)[:, None] * inv[None, :] + return torch.cat([ang.cos(), ang.sin()], dim=1).float().contiguous() + + +def make_pool(style, pages, nkv, gen): + """Per-layer HND views [pages, 2, nkv, PAGE, 64] of one allocation, filled with a pattern.""" + per_page = 2 * nkv * PAGE * HEAD + if style == "arena": + buf = torch.randn( + LAYERS, pages, 2, nkv, PAGE, HEAD, generator=gen, device="cuda" + ).bfloat16() + return buf, [buf[layer] for layer in range(LAYERS)] + # Page-major: each layer at offset layer * per_page inside every page group. + buf = torch.randn(pages * LAYERS * per_page, generator=gen, device="cuda").bfloat16() + views = [ + buf.as_strided( + (pages, 2, nkv, PAGE, HEAD), + (LAYERS * per_page, nkv * PAGE * HEAD, PAGE * HEAD, HEAD, 1), + layer * per_page, + ) # fmt: skip + for layer in range(LAYERS) + ] + return buf, views + + +def views_of(buf, like, style): + if style == "arena": + return [buf[layer] for layer in range(LAYERS)] + return [buf.as_strided(t.shape, t.stride(), t.storage_offset()) for t in like] + + +def python_path( + x, w, k_norm, cs, cpos, num_acc, ctx_len, slots, rows, table, counts, layers, block_size +): + """DFlashWorker's ops for the step (the model's functions).""" + from tensorrt_llm._torch.custom_ops import ( + flashinfer_apply_rope_with_cos_sin_cache_inplace as rope, + ) + from tensorrt_llm._torch.speculative.dflash_attention import get_dflash_paged_append + + n, k1 = x.shape[0], cpos.shape[1] + nkv = w.shape[0] // (LAYERS * 2 * HEAD) + kv = F.linear(x, w).view(n, LAYERS, 2, nkv, HEAD) + k = kv[:, :, 0].contiguous() + v = kv[:, :, 1].contiguous() + k = F.rms_norm(k, (HEAD,), eps=EPS) + k = k * k_norm.view(1, LAYERS, 1, HEAD) + pos = cpos.reshape(-1).to(torch.int32).repeat_interleave(LAYERS) + dummy_q = k.new_empty(n * LAYERS, HEAD) + rope(pos, dummy_q, k.view(n * LAYERS, nkv * HEAD), HEAD, cs, True) + offs = torch.arange(k1, device="cuda") + mask = (offs[None, :] < num_acc.long()[:, None]).reshape(-1).view(-1, 1, 1, 1).to(k.dtype) + k.mul_(mask) + v.mul_(mask) + col = ctx_len[slots][:, None] + offs[None, :] + cap = counts[rows] * PAGE + col = torch.minimum(col, (cap - 1)[:, None]).clamp_(min=0) + rows_i32 = rows[:, None].expand(-1, k1).reshape(-1).to(torch.int32) + col_i32 = col.reshape(-1).to(torch.int32) + width = table.shape[1] + indptr = torch.arange(0, (table.shape[0] + 1) * width, width, dtype=torch.int32, device="cuda") + last = torch.full((table.shape[0],), PAGE, dtype=torch.int32, device="cuda") + append = get_dflash_paged_append() + for layer in range(LAYERS): + append(append_key=k[:, layer].contiguous(), append_value=v[:, layer].contiguous(), batch_indices=rows_i32, + positions=col_i32, paged_kv_cache=layers[layer], kv_indices=table.flatten().contiguous(), + kv_indptr=indptr, kv_last_page_len=last, kv_layout="HND") # fmt: skip + ctx_len[slots] += num_acc.long() + ctx_len.clamp_(max=MAX_CTX) + allocated = (counts * PAGE - block_size).clamp(min=0) + return torch.minimum(ctx_len[slots], allocated[rows]) + + +def fp32_rows(x, w, k_norm, cs, cpos, nkv): + """High-precision reference of every (token, layer, K/V, head) row: [N, L, 2, nkv, 64] fp32, no bf16 rounding + (computed in float64: no TF32 even where cuBLAS is told to use it).""" + n = x.shape[0] + kv = (x.double() @ w.double().t()).view(n, LAYERS, 2, nkv, HEAD) + k = kv[:, :, 0] + k = ( + k + * torch.rsqrt(k.pow(2).mean(-1, keepdim=True) + EPS) + * k_norm.double().view(1, LAYERS, 1, HEAD) + ) + c = cs.double()[cpos.reshape(-1)][:, None, None, : HEAD // 2] + s = cs.double()[cpos.reshape(-1)][:, None, None, HEAD // 2 :] + k1_, k2_ = k[..., : HEAD // 2], k[..., HEAD // 2 :] + k = torch.cat([k1_ * c - k2_ * s, k2_ * c + k1_ * s], dim=-1) + return torch.stack([k, kv[:, :, 1]], dim=2).float() + + +_views = {} + + +def pool_view(layers): + """The kernel's view of a pool, built once per pool outside capture (as the worker does); a few kept.""" + key = tuple(t.data_ptr() for t in layers) + view = _views.get(key) + if view is None: + if len(_views) >= 4: + _views.clear() + view = _views[key] = _op().pool_view(layers) + return view + + +def kernel_call( + x, w, k_norm, cs, cpos, num_acc, ctx_len, slots, rows, table, counts, layers, block_size, nkv +): + flat, layer_off, ps, kvs, hs = pool_view(layers) + return torch.ops.trtllm.k3_ctx_kv(x, w, k_norm, cs, cpos, num_acc, ctx_len, slots, rows, table, counts, flat, + layer_off, ps, kvs, hs, EPS, MAX_CTX, PAGE, block_size, nkv) # fmt: skip + + +def gather_rows(layers, table, rows, col): + """The pool rows of each (token, layer, K/V, head): [N, L, 2, nkv, 64].""" + page = table[rows[:, None].expand_as(col).reshape(-1), col.reshape(-1) // PAGE].long() + off = col.reshape(-1) % PAGE + return torch.stack([layers[layer][page, :, :, off, :] for layer in range(LAYERS)], dim=1) + + +class Step: + """One decode step's inputs: B requests (distinct slots and table rows, lengths whose columns cross pages, mixed + accepted counts), the weights, the table and the pools.""" + + def __init__(self, gen, nkv, batch, k1, style="v1", clamp=False, seed=0): + self.nkv, self.batch, self.k1, self.style = nkv, batch, k1, style + self.n = batch * k1 + n_rows = LAYERS * 2 * nkv * HEAD + self.w = (torch.randn(n_rows, HIDDEN, generator=gen, device="cuda") * 0.02).bfloat16() + self.k_norm = ( + 1.0 + 0.1 * torch.randn(LAYERS, HEAD, generator=gen, device="cuda") + ).bfloat16() + self.cs = cos_sin_cache("cuda") + width, table_rows, slots_total, self.pages = 64, batch + 3, batch + 4, 64 * (batch + 3) + 16 + self.table = ( + torch.randperm(self.pages, generator=torch.Generator().manual_seed(seed))[: table_rows * width] + .view(table_rows, width).to(torch.int32).cuda() + ) # fmt: skip + self.counts = torch.full((table_rows,), width, dtype=torch.int64, device="cuda") + g = torch.Generator().manual_seed(seed + 1) + self.slots = torch.randperm(slots_total, generator=g)[:batch].cuda() + self.rows = torch.randperm(table_rows, generator=g)[:batch].cuda() + # Lengths: a page end inside the block for every other request (64 m - 3), else random. + ctx0 = torch.randint(100, 2000, (slots_total,), generator=g) + for b in range(0, batch, 2): + ctx0[int(self.slots[b])] = 64 * int(torch.randint(2, 30, (1,), generator=g)) - 3 + if clamp: # the first request's last columns clamp to its allocation's end + self.counts[int(self.rows[0])] = 20 + ctx0[int(self.slots[0])] = 20 * PAGE - 3 + self.ctx0 = ctx0.cuda() + self.num_acc = (torch.arange(batch) * 3 + seed) % k1 + 1 + self.num_acc = self.num_acc.to(torch.int32).cuda() + self.cpos = ( + self.ctx0[self.slots][:, None] + torch.arange(k1, device="cuda")[None, :] + ).contiguous() + self.x = (torch.randn(self.n, HIDDEN, generator=gen, device="cuda") * 0.5).bfloat16() + self.buf, self.layers = make_pool(style, self.pages, nkv, gen) + + def cols(self): + return torch.minimum(self.ctx0[self.slots][:, None] + torch.arange(self.k1, device="cuda")[None, :], + (self.counts[self.rows] * PAGE - 1)[:, None]).clamp(min=0) # fmt: skip + + def run_kernel(self, buf=None, ctx=None, fn=None): + buf = self.buf.clone() if buf is None else buf + ctx = self.ctx0.clone() if ctx is None else ctx + layers = views_of(buf, self.layers, self.style) + call = kernel_call if fn is None else fn + num_ctx = call(self.x, self.w, self.k_norm, self.cs, self.cpos, self.num_acc, ctx, self.slots, self.rows, + self.table, self.counts, layers, BLOCK, self.nkv) # fmt: skip + return buf, ctx, num_ctx.long(), layers + + def run_python(self): + buf = self.buf.clone() + ctx = self.ctx0.clone() + layers = views_of(buf, self.layers, self.style) + num_ctx = python_path(self.x, self.w, self.k_norm, self.cs, self.cpos, self.num_acc, ctx, self.slots, + self.rows, self.table, self.counts, layers, BLOCK) # fmt: skip + return buf, ctx, num_ctx.long(), layers + + +def measure(nkv, batch, k1, style="v1", clamp=False, seed=0): + gen = torch.Generator(device="cuda").manual_seed(20260929 + 100 * batch + k1 + seed) + st = Step(gen, nkv, batch, k1, style, clamp, seed) + buf_p, ctx_p, nc_p, lay_p = st.run_python() + buf_k, ctx_k, nc_k, lay_k = st.run_kernel() + torch.cuda.synchronize() + col = st.cols() + ref32 = fp32_rows(st.x, st.w, st.k_norm, st.cs, st.cpos, nkv) + got = gather_rows(lay_k, st.table, st.rows, col).float() + want = gather_rows(lay_p, st.table, st.rows, col).float() + n = st.n + live = (torch.arange(k1, device="cuda")[None, :] < st.num_acc.long()[:, None]).reshape(n) + # A clamped column is written by several tokens; compare only the tokens whose column is their own. + cflat = col.reshape(-1) + full = st.rows[:, None].expand_as(col).reshape(-1) * 1_000_000 + cflat + own = torch.tensor([int((full == full[t]).sum()) == 1 for t in range(n)], device="cuda") + cmp = (live & own).view(n, 1, 1, 1, 1).expand_as(got) + scale = ref32.abs().amax(dim=-1, keepdim=True).clamp(min=1e-6).expand_as(got) + e_k = ((got - ref32).abs() / scale)[cmp].max().item() if cmp.any() else 0.0 + e_p = ((want - ref32).abs() / scale)[cmp].max().item() if cmp.any() else 0.0 + e_kp = ((got - want).abs() / scale)[cmp].max().item() if cmp.any() else 0.0 + masked = (~live & own).view(n, 1, 1, 1, 1).expand_as(got) + masked_zero = bool((got[masked] == 0).all()) if masked.any() else True + # Every pool element the Python path left alone is left alone (and the kernel wrote no others). + before = st.buf.view(torch.int16) + changed_k = buf_k.view(torch.int16) != before + changed_p = buf_p.view(torch.int16) != before + written = torch.zeros(st.buf.numel(), dtype=torch.bool, device="cuda") + idx = torch.arange(st.buf.numel(), device="cuda") + pages = st.table[st.rows[:, None].expand_as(col).reshape(-1), cflat // PAGE].long() + for layer in range(LAYERS): + idx_l = idx.view(st.buf.shape)[layer] if st.style == "arena" else idx.as_strided( + lay_p[layer].shape, lay_p[layer].stride(), lay_p[layer].storage_offset()) # fmt: skip + written[idx_l[pages, :, :, cflat % PAGE, :].reshape(-1)] = True + written = written.view(st.buf.shape) + untouched = bool((~((changed_k | changed_p) & ~written)).all()) + buf_r, ctx_r, nc_r, _ = st.run_kernel() + rerun = torch.equal(buf_r.view(torch.int16), buf_k.view(torch.int16)) and torch.equal(ctx_r, ctx_k) and \ + torch.equal(nc_r, nc_k) # fmt: skip + res = dict(split=f"{batch}x{k1}", nkv=nkv, pool=style, clamp=clamp, num_acc=st.num_acc.tolist(), + ctx=st.ctx0[st.slots].tolist(), e_kernel=e_k, e_python=e_p, e_kernel_vs_python=e_kp, + ctx_len=torch.equal(ctx_p, ctx_k), num_ctx=torch.equal(nc_p, nc_k), masked_zero=masked_zero, + untouched=untouched, rerun=rerun) # fmt: skip + res["ok"] = (res["ctx_len"] and res["num_ctx"] and masked_zero and untouched and rerun + and e_k <= max(2.0 * e_p, 1.6e-2)) # fmt: skip + return res + + +@pytest.mark.parametrize("batch,k1", SPLITS, ids=[f"{b}x{k}" for b, k in SPLITS]) +@pytest.mark.parametrize("style", ["v1", "arena"]) +def test_split(batch, k1, style): + with torch.inference_mode(): + res = measure(1, batch, k1, style, seed=1 if style == "arena" else 0) + assert res["ok"], res + + +@pytest.mark.parametrize("batch,k1", [(1, 8), (2, 8), (8, 8), (8, 1), (4, 4)]) +def test_clamp(batch, k1): + with torch.inference_mode(): + res = measure(1, batch, k1, "v1", clamp=True, seed=2) + assert res["ok"], res + + +TP4_SPLITS = [(1, 8), (2, 4), (8, 1), (2, 8), (4, 8)] + + +@pytest.mark.parametrize("batch,k1", TP4_SPLITS, ids=[f"{b}x{k}" for b, k in TP4_SPLITS]) +def test_split_tp4(batch, k1): + with torch.inference_mode(): + res = measure(4, batch, k1, "v1", seed=3) + assert res["ok"], res + + +def test_supported_shapes(): + """TP16 takes every step up to 8 x 8; TP4 (shared memory for the resident tokens) up to 32 tokens.""" + op = _op() + dev = torch.device("cuda") + for b, k1 in SPLITS: + assert op.pick_split(LAYERS * 2 * HEAD, HIDDEN, 1, k1, b * k1, dev) == 8, (b, k1) + assert op.pick_split(LAYERS * 2 * 4 * HEAD, HIDDEN, 4, 8, 32, dev) == 4 + assert op.pick_split(LAYERS * 2 * 4 * HEAD, HIDDEN, 4, 8, 64, dev) == 0 + assert op.pick_split(LAYERS * 2 * HEAD, HIDDEN, 1, 8, 72, dev) == 0 # 9 requests + assert op.pick_split(LAYERS * 2 * HEAD, HIDDEN, 1, 1, 9, dev) == 0 # 9 requests of 1 + + +@pytest.mark.parametrize("batch,k1", [(1, 8), (4, 8), (8, 8), (8, 1)]) +def test_graph_replay(batch, k1): + """20 replays of a captured call with the inputs rewritten in place (x, ctx_len, num_acc, cpos).""" + with torch.inference_mode(): + gen = torch.Generator(device="cuda").manual_seed(7 + batch) + st = Step(gen, 1, batch, k1, "v1", seed=4) + x = st.x.clone() + ctx = st.ctx0.clone() + num_acc = st.num_acc.clone() + cpos = st.cpos.clone() + buf = st.buf.clone() + lay = views_of(buf, st.layers, "v1") + out = {} + + def body(): + out["nc"] = kernel_call(x, st.w, st.k_norm, st.cs, cpos, num_acc, ctx, st.slots, st.rows, st.table, + st.counts, lay, BLOCK, 1) # fmt: skip + + body() + torch.cuda.synchronize() + stream = torch.cuda.Stream() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + body() + for rep in range(20): + x.copy_((torch.randn(x.shape, generator=gen, device="cuda") * 0.5).bfloat16()) + ctx.copy_(torch.randint(100, 2000, ctx.shape, generator=gen, device="cuda")) + num_acc.copy_(((torch.arange(batch, device="cuda") + rep) % k1 + 1).to(torch.int32)) + cpos.copy_(ctx[st.slots][:, None] + torch.arange(k1, device="cuda")[None, :]) + ref_buf = buf.clone() + ref_lay = views_of(ref_buf, st.layers, "v1") + ref_ctx = ctx.clone() + want_nc = python_path(x, st.w, st.k_norm, st.cs, cpos, num_acc, ref_ctx, st.slots, st.rows, st.table, + st.counts, ref_lay, BLOCK) # fmt: skip + graph.replay() + torch.cuda.synchronize() + col = torch.minimum(cpos, (st.counts[st.rows] * PAGE - 1)[:, None]).clamp(min=0) + got = gather_rows(lay, st.table, st.rows, col).float() + want = gather_rows(ref_lay, st.table, st.rows, col).float() + ref32 = fp32_rows(x, st.w, st.k_norm, st.cs, cpos, 1) + live = (torch.arange(k1, device="cuda")[None, :] < num_acc.long()[:, None]).reshape(-1) + live = live.view(-1, 1, 1, 1, 1).expand_as(got) + scale = ref32.abs().amax(dim=-1, keepdim=True).clamp(min=1e-6).expand_as(got) + e_k = ((got - ref32).abs() / scale)[live].max().item() + e_p = ((want - ref32).abs() / scale)[live].max().item() + assert torch.equal(ctx, ref_ctx), f"replay {rep}: ctx_len" + assert torch.equal(out["nc"].long(), want_nc.long()), f"replay {rep}: num_ctx" + assert e_k <= max(2.0 * e_p, 1.6e-2), f"replay {rep}: {e_k} vs {e_p}" + + +# The elements before an offset view: no index, length or table entry of the tests holds it. +PAD = -1 + + +def _offset(t: torch.Tensor, k: int, pads: list) -> torch.Tensor: + """``t``'s values ``k`` elements into a buffer of their own, after ``k`` PAD elements, as a view: where the + drafter's per-request slices (``num_accepted_tokens[num_contexts:]``, ``_batch_to_slot[num_contexts:]``) can + start. The buffer's pad goes to ``pads`` for the caller to check.""" + flat = torch.full((k + t.numel(),), PAD, dtype=t.dtype, device=t.device) + flat[k:] = t.reshape(-1) + pads.append(flat[:k]) + return flat[k:].view(t.shape) + + +def _kernel_call_layer_off_at(k, pads): + def call( + x, + w, + k_norm, + cs, + cpos, + num_acc, + ctx_len, + slots, + rows, + table, + counts, + layers, + block_size, + nkv, + ): + flat, layer_off, ps, kvs, hs = pool_view(layers) + return torch.ops.trtllm.k3_ctx_kv(x, w, k_norm, cs, cpos, num_acc, ctx_len, slots, rows, table, counts, flat, + _offset(layer_off, k, pads), ps, kvs, hs, EPS, MAX_CTX, PAGE, block_size, + nkv) # fmt: skip + + return call + + +@pytest.mark.parametrize( + "arg", ["cpos", "num_acc", "ctx_len", "slots", "rows", "table", "counts", "layer_off"] +) +def test_index_offset(arg): + """Each int index / length / table argument as a view 1, 2 and 3 elements into its buffer: the pool, ctx_len and + num_ctx equal those of the call on tensors of their own, bit for bit, and the elements before the view keep their + PAD values (no read or write outside it).""" + with torch.inference_mode(): + gen = torch.Generator(device="cuda").manual_seed(20260929 + 37) + st = Step(gen, 1, 3, 4, "v1", seed=5) + want_buf, want_ctx, want_nc, _ = st.run_kernel() + for k in (1, 2, 3): + ctx, fn, saved, pads = None, None, None, [] + if arg == "ctx_len": + ctx = _offset(st.ctx0, k, pads) + elif arg == "layer_off": + fn = _kernel_call_layer_off_at(k, pads) + else: + saved = getattr(st, arg) + setattr(st, arg, _offset(saved, k, pads)) + try: + buf, got_ctx, got_nc, _ = st.run_kernel(ctx=ctx, fn=fn) + torch.cuda.synchronize() + finally: + if saved is not None: + setattr(st, arg, saved) + assert torch.equal(buf.view(torch.int16), want_buf.view(torch.int16)), ( + f"{arg} at {k}: pool" + ) + assert torch.equal(got_ctx, want_ctx), f"{arg} at {k}: ctx_len" + assert torch.equal(got_nc, want_nc), f"{arg} at {k}: num_ctx" + assert len(pads) == 1 and bool((pads[0] == PAD).all()), f"{arg} at {k}: pad" + + +# ---------------------------------------------------------------------------------------------------------------- +# Batch-1 identity against the unmodified kernel (K3_BASE_TRTLLM: an unmodified tensorrt_llm package directory). +# ---------------------------------------------------------------------------------------------------------------- + +_base = {} + + +def base_kernel(): + root = os.environ.get("K3_BASE_TRTLLM") + if not root: + return None + if "mod" not in _base: + path = os.path.join(root, "_torch", "cute_dsl_kernels", "k3_ctx_kv", "k3_ctx_kv_kernel.py") + spec = importlib.util.spec_from_file_location("k3_ctx_kv_kernel_base", path) + mod = importlib.util.module_from_spec(spec) + spec.loader.exec_module(mod) + _base["mod"] = mod + return _base["mod"] + + +def base_call( + x, w, k_norm, cs, cpos, num_acc, ctx_len, slots, rows, table, counts, layers, block_size, nkv +): + """The unmodified op (N <= 8) on the same arguments.""" + import cuda.bindings.driver as cuda_driver + import cutlass.cute as cute + from cutlass.cute.runtime import from_dlpack + + kern = base_kernel() + + def arg(t): + return from_dlpack(t.detach(), assumed_align=16).mark_layout_dynamic( + leading_dim=t.dim() - 1 + ) + + flat, layer_off, ps, kvs, hs = pool_view(layers) + batch, k1 = cpos.shape + n_tokens, k_in = x.shape + n_rows = w.shape[0] + sms = torch.cuda.get_device_properties(x.device).multi_processor_count + split = next(s for s in (8, 4, 2) if (n_rows // kern.CTA_M) * s <= sms + and kern.supports(n_rows, k_in, s, nkv, k1, n_tokens)) # fmt: skip + ring = kern.pick_ring(k_in, split) + if "counter" not in _base: + _base["counter"] = torch.zeros(1, dtype=torch.int32, device="cuda") + num_ctx = torch.empty(batch, dtype=torch.int32, device="cuda") + args = (arg(w), arg(x), arg(k_norm.reshape(-1)), arg(cs.reshape(-1)), arg(cpos.reshape(-1)), arg(num_acc), + arg(ctx_len), arg(slots), arg(rows), arg(table.reshape(-1)), arg(counts), arg(flat), arg(layer_off), + arg(num_ctx), arg(_base["counter"])) # fmt: skip + scalars = ( + float(EPS), + int(MAX_CTX), + int(PAGE), + int(block_size), + int(table.stride(0)), + int(ps), + int(kvs), + int(hs), + ) + consts = (n_rows, k_in, split, ring, nkv, k1, n_tokens, True) + stream = cuda_driver.CUstream(torch.cuda.current_stream().cuda_stream) + fn = _base.get(consts) + if fn is None: + fn = _base[consts] = cute.compile(kern.k3_ctx_kv, *args, *scalars, *consts, True, stream) + fn(*args, *scalars, stream) + return num_ctx + + +@pytest.mark.skipif( + not os.environ.get("K3_BASE_TRTLLM"), reason="K3_BASE_TRTLLM (unmodified package) not set" +) +@pytest.mark.parametrize("batch,k1", [s for s in SPLITS if s[0] * s[1] <= 8], ids=[f"{b}x{k}" for b, k in SPLITS + if b * k <= 8]) # fmt: skip +def test_batch1_identity(batch, k1): + with torch.inference_mode(): + for seed, (style, clamp) in enumerate((("v1", False), ("arena", False), ("v1", True))): + gen = torch.Generator(device="cuda").manual_seed(55 + 10 * batch + k1 + seed) + st = Step(gen, 1, batch, k1, style, clamp, seed) + buf_k, ctx_k, nc_k, _ = st.run_kernel() + buf_b, ctx_b, nc_b, _ = st.run_kernel(fn=base_call) + torch.cuda.synchronize() + assert torch.equal(buf_k.view(torch.int16), buf_b.view(torch.int16)), (style, clamp) + assert torch.equal(ctx_k, ctx_b) and torch.equal(nc_k, nc_b), (style, clamp) + + +# ---------------------------------------------------------------------------------------------------------------- +# Timing (python3 test_k3_ctx_kv.py time [--base]) and the error table (report) +# ---------------------------------------------------------------------------------------------------------------- + + +def time_graph(body, calls, replays=15): + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(stream): + body(0) + torch.cuda.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + for i in range(calls): + body(i) + torch.cuda.synchronize() + for _ in range(3): + graph.replay() + torch.cuda.synchronize() + per_call = [] + for _ in range(replays): + start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + start.record() + graph.replay() + end.record() + torch.cuda.synchronize() + per_call.append(start.elapsed_time(end) * 1e3 / calls) + return statistics.median(per_call), min(per_call), max(per_call) + + +def timing(with_base: bool) -> None: + """Graphs of back-to-back calls with the weight rotating over 160 MB of copies (HBM-cold), TP16.""" + _op() + gen = torch.Generator(device="cuda").manual_seed(11) + n_rows = LAYERS * 2 * HEAD + copies = max(2, -(-(160 << 20) // (n_rows * HIDDEN * 2))) + ws = [ + (torch.randn(n_rows, HIDDEN, generator=gen, device="cuda") * 0.02).bfloat16() + for _ in range(copies) + ] + calls = 2 * copies + print(f"{torch.cuda.get_device_name()}; graphs of {calls} calls, weights rotating over {copies} copies, " + "15 replays: median (min-max) us per call") # fmt: skip + print("| split | N | k3_ctx_kv | base |") + print("| :-- | --: | --: | --: |") + with torch.inference_mode(): + for b, k1 in SPLITS: + st = Step(gen, 1, b, k1, "v1", seed=5) + st.num_acc.fill_(max(1, k1 - 2)) + ctx = st.ctx0.clone() + arms = [lambda i, fn=fn: fn(st.x, ws[i % copies], st.k_norm, st.cs, st.cpos, st.num_acc, ctx, st.slots, + st.rows, st.table, st.counts, st.layers, BLOCK, 1) + for fn in ([kernel_call, base_call] if with_base and b * k1 <= 8 else [kernel_call])] # fmt: skip + res = [[] for _ in arms] + for rep in range(3): + for a in range(len(arms)) if rep % 2 == 0 else reversed(range(len(arms))): + ctx.copy_(st.ctx0) + res[a].append(time_graph(arms[a], calls)) + cells = [] + for a in range(2): + if a < len(arms): + meds = sorted(x[0] for x in res[a]) + cells.append( + f"{meds[1]:.2f} ({min(x[1] for x in res[a]):.2f}-{max(x[2] for x in res[a]):.2f})" + ) + else: + cells.append("") + print(f"| {b}x{k1} | {b * k1} | " + " | ".join(cells) + " |", flush=True) + + +def report() -> int: + print(f"{torch.cuda.get_device_name()}") + print( + "| nkv | split | pool | clamp | kernel vs fp32 | Python vs fp32 | kernel vs Python | ctx_len | num_ctx | " + "masked 0 | untouched | rerun | result |" + ) + print("| --: | :-- | :-- | :-- | --: | --: | --: | :-- | :-- | :-- | :-- | :-- | :-- |") + ok_all = True + cases = [(1, b, k, s, False) for b, k in SPLITS for s in ("v1", "arena")] + cases += [(1, b, k, "v1", True) for b, k in [(1, 8), (2, 8), (8, 8), (8, 1), (4, 4)]] + cases += [(4, b, k, "v1", False) for b, k in TP4_SPLITS] + with torch.inference_mode(): + for nkv, b, k, style, clamp in cases: + r = measure( + nkv, b, k, style, clamp, seed=1 if style == "arena" else (2 if clamp else 0) + ) + ok_all &= r["ok"] + print(f"| {nkv} | {b}x{k} | {style} | {clamp} | {r['e_kernel']:.2e} | {r['e_python']:.2e} | " + f"{r['e_kernel_vs_python']:.2e} | {r['ctx_len']} | {r['num_ctx']} | {r['masked_zero']} | " + f"{r['untouched']} | {r['rerun']} | {'PASS' if r['ok'] else 'FAIL'} |", flush=True) # fmt: skip + print("ALL PASS" if ok_all else "FAIL") + return 0 if ok_all else 1 + + +if __name__ == "__main__": + if len(sys.argv) > 1 and sys.argv[1] == "time": + timing("--base" in sys.argv) + elif len(sys.argv) > 1 and sys.argv[1] == "report": + sys.exit(report()) + else: + sys.exit(pytest.main([__file__, "-q", "-p", "no:cacheprovider", *sys.argv[1:]])) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_markov.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_markov.py new file mode 100644 index 000000000000..c276699af735 --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_markov.py @@ -0,0 +1,1006 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``trtllm::k3_markov`` (the Kimi K3 DSpark vanilla-Markov draft chain over a vocab-sharded draft head) on W GPUs of +one NVLink domain (MPI, one rank per GPU), against two references: + +- exact: the chain in the arithmetic of the C++ ``trtllm::dspark_markov_chain`` it replaces (the TeKit port's op when + the build has it, else the same arithmetic in torch, ``Group.emulated_chain``; with both, the 'emulation' column + compares them) + the TP-gathered greedy sampler (``SpecWorkerBase.greedy_sample_draft_with_tp_gather``: each rank's + first maximum, an all-gather of (index, value), the first maximum over the ranks) + the next_new_tokens assembly + (``SpecWorkerBase._prepare_next_new_tokens``). The corrected logits (fp32 bits), the tokens and next_new must be + bit-identical to it on every rank, and identical across the ranks; +- main's torch chain (``dspark_markov_step_bias``: an F.linear of the bf16 weights, rounded to bf16 once, whose + summation order differs), fed k3_markov's tokens as the anchors: the corrected logits within one bf16 ulp of the bias + and every token a global maximum within that tolerance (the 'torch chain' column; the tolerance adds the fp32 + accumulation bound, 2^-16 of the sum of the products' magnitudes, for sums near zero). + +Checks (``report``); every case is also rerun (bit-identical) and run with the folded KV-length rewind +(``kv_lens[first + b] = max(kv_lens[first + b] - rewind[b], 0)``: positive, negative and clamped rewinds, every other +output unchanged): + +- every split, B = 1 .. 8 requests x K in (1, 3, 7) block positions, fp32 and bf16 base logits [B, K, S] (bf16 is + compared with the reference fed ``base.float()``), on two vocab shards: S = 10240 (the TP16 slice; every rank + pushes each exchange entry into ``--copies`` = 16 / W slots, the exchange volume of 16 ranks) and S = 163840 / W. + ``pick_grid`` must take every split of the TP16 slice; a shard it rejects is reported (not run), and ``supports`` + and the op must reject it too; +- crafted, at several (B, K): exact ties within an 8-row block, across blocks, warps, CTAs and ranks (identical + markov_w2 rows and base values: the lowest vocabulary index wins), anchors outside the vocabulary (zero bias), NaN + logits (never win: compared with a NaN-ignoring sampler), rewind and acceptance edges (a partial rewind from row 0, + an empty one, num_accepted 0, a wider accepted, num_accepted and the accepted rows as the model's [num_contexts:] + slices); +- mixed steps: G generation requests after num_contexts context requests, num_accepted / accepted_rows the model's + int32 slices at element num_contexts and the rewind from rewind_first = num_contexts, also against the same call on + contiguous copies (bit for bit); +- 60 eager calls with B, K and the dtype varying, interleaved with MNNVL all-reduces (and C++ chain calls when the build + has the op: the Lamport rotations of k3_markov's workspaces and of the all-reduce's, which the C++ chain shares); +- CUDA graphs of [all-reduce, k3_markov, k3_markov with the rewind], B in (1, 8) x K in (1, 7), fp32 and bf16, + replayed 20 times with rewritten inputs. + +Timing (``time``): per split at S = 10240 with the copies, CUDA graphs of back-to-back calls, each call on its own +markov_w2 shard copy (> 200 MB between two reads of one: HBM-cold), ABBA rounds behind an MPI barrier; the max over the +ranks of the median us per call of the production path (the fp32 cast of the bf16 head logits, the C++ chain or else +main's torch chain with the TP-gathered argmax per position, the gathered sampler, the next_new assembly; it exchanges +over the W ranks only) and of k3_markov on the bf16 logits. + + srun -n W --mpi=pmix python3 test_k3_markov.py [report | time] [--copies C] [--skip-perf] [--rounds R] + +W = 2 or 4 GPUs of one NVLink domain, every GPU of the node visible to every rank (rank r runs on GPU r % count). +Without a mode the checks run, then the timing (unless ``--skip-perf``); the exit code is nonzero when a check fails. +The checks also run under pytest on W >= 2 ranks (``srun -n W --mpi=pmix python3 -m pytest -p no:cacheprovider +test_k3_markov.py``); with fewer ranks the module is skipped. +""" + +import argparse +import contextlib +import os +import statistics +import sys +import traceback +import zlib +from typing import Optional + +import pytest +import torch + +VOCAB = 163840 # the draft vocabulary: rows of markov_w1 and markov_w2 +MARKOV_RANK = 256 +TP16_SHARD = VOCAB // 16 +HIDDEN = 7168 # the all-reduce's row width +BATCHES = tuple(range(1, 9)) +BLOCKS = (1, 3, 7) +DTYPES = (torch.float32, torch.bfloat16) +DTYPE_NAMES = {torch.float32: "fp32", torch.bfloat16: "bf16"} +CRAFTED = { + "ties": "ties", + "anchors": "anchors outside the vocabulary", + "nan": "NaN logits", + "edges": "rewind / acceptance edges", +} +CRAFTED_SPLITS = ((1, 1), (1, 7), (3, 3), (8, 1), (8, 7)) +# Mixed steps: G generation requests after num_contexts context requests (DSpark K = 7, bf16 logits, the TP16 slice). +MIXED_CONTEXTS = (0, 1, 3) +MIXED_GENS = (1, 2, 4, 8) +GRAPH_SPLITS = ((1, 1), (1, 7), (8, 1), (8, 7)) +# Local rows tied on every rank (with S / 2 + 3 and S - 5): 3 and 5 share an 8-row block, 11 is the warp's next block, +# then the next warps and CTAs (a CTA owns 128 rows of the TP16 slice, 320 of 163840 / 4). +TIE_ROWS = (3, 5, 11, 19, 67, 131, 323) +HBM_COLD_BYTES = 200 << 20 +FIELDS = ("corrected", "tokens", "next_new", "rerun", "ranks", "rewind", "torch", "emulation") +HEADER = ( + "| case | B | K | S | dtype | corrected | tokens | next_new | rerun | ranks identical | rewind | torch chain " + "| emulation | result |\n" + "| :-- | --: | --: | --: | :-- | :-- | :-- | :-- | :-- | :-- | :-- | :-- | :-- | :-- |" +) + + +def _sm100() -> bool: + return torch.cuda.is_available() and torch.cuda.get_device_capability()[0] == 10 + + +def _world_size() -> int: + """MPI ranks of this launch (1 without mpi4py).""" + try: + from mpi4py import MPI + except ImportError: + return 1 + return MPI.COMM_WORLD.Get_size() + + +pytestmark = [ + pytest.mark.skipif(not _sm100(), reason="needs SM100 (MNNVL multicast, TMA bulk copies, PDL)"), + pytest.mark.skipif( + _world_size() < 2, reason="needs >= 2 MPI ranks: srun -n W --mpi=pmix python3 -m pytest" + ), +] + + +def case_seed(*key) -> int: + """A seed from a case's parameters, the same in every process (unlike hash()).""" + return zlib.crc32(repr(key).encode()) + + +def bits(t: torch.Tensor) -> torch.Tensor: + return t.contiguous().view(torch.int32) + + +def identical(a: torch.Tensor, b: torch.Tensor) -> bool: + """Same shape, dtype and bits (NaN payloads included).""" + return a.shape == b.shape and a.dtype == b.dtype and torch.equal(bits(a), bits(b)) + + +def same_outputs(a, b) -> bool: + return all(identical(x, y) for x, y in zip(a, b)) + + +def differing_words(a: torch.Tensor, b: torch.Tensor) -> int: + """32-bit words of ``a`` that differ from ``b`` (all of them when the shapes or dtypes differ).""" + if a.shape != b.shape or a.dtype != b.dtype: + return max(a.numel(), 1) + return int((bits(a) != bits(b)).sum()) + + +def acceptance( + batch: int, block: int, gen, extra_rows: int = 1, extra_width: int = 0, min_accepted: int = 1 +): + """(accepted [B + extra_rows, K + 1 + extra_width], num_accepted [B] in [min_accepted, K + 1], distinct accepted + rows [B]), int32 and the same on every rank (``gen`` is a shared generator).""" + rows = batch + extra_rows + width = block + 1 + extra_width + accepted = torch.randint( + 0, VOCAB, (rows, width), generator=gen, device="cuda", dtype=torch.int32 + ) + num_accepted = torch.randint( + min_accepted, block + 2, (batch,), generator=gen, device="cuda", dtype=torch.int32 + ) + accepted_rows = torch.randperm(rows, generator=gen, device="cuda")[:batch].to(torch.int32) + return accepted, num_accepted, accepted_rows + + +def next_new_reference(acc, tokens: torch.Tensor) -> torch.Tensor: + """SpecWorkerBase._prepare_next_new_tokens: [accepted[rows[b], num_accepted[b] - 1], tokens[b, :]].""" + accepted, num_accepted, rows = acc + first = accepted[rows.long(), num_accepted.long() - 1].unsqueeze(1) + return torch.cat([first, tokens], dim=1) + + +def rewind_spec(batch: int, seed: int, count: Optional[int] = None, first: int = 2, after: int = 1): + """(kv_lens, rewind, rewind_first, the kv_lens it must leave) for the folded KV-length rewind of ``count`` (B by + default) rows from ``first``: a rewind past the length (clamped to 0), a negative one (the length grows), one to + exactly 0, random others; ``after`` untouched rows follow.""" + count = batch if count is None else count + gen = torch.Generator().manual_seed(seed) + kv_lens = torch.randint(0, 50, (first + count + after,), generator=gen, dtype=torch.int32) + amounts = torch.randint(-6, 9, (count,), generator=gen, dtype=torch.int32) + amounts[0] = kv_lens[first] + 5 + if count > 1: + amounts[1] = -1 - amounts[1].abs() + if count > 2: + amounts[2] = kv_lens[first + 2] + want = kv_lens.clone() + want[first : first + count] = (kv_lens[first : first + count] - amounts).clamp_min(0) + return kv_lens.cuda(), amounts.cuda(), first, want.cuda() + + +class Group: + """This rank of the W-rank TP group: the mapping, the MNNVL all-reduce (whose workspace the C++ chain shares), the + Markov weights (the same on every rank), the C++ chain's scratch, and the reference path.""" + + def __init__(self, copies: Optional[int] = None): + from mpi4py import MPI + + self.comm = MPI.COMM_WORLD + self.rank, self.world = self.comm.Get_rank(), self.comm.Get_size() + gpus = torch.cuda.device_count() + torch.cuda.set_device(self.rank % gpus) + import tensorrt_llm # noqa: F401 (registers the trtllm ops) + from tensorrt_llm._torch.cute_dsl_kernels.k3_markov import op + from tensorrt_llm._torch.distributed import AllReduceFusionOp, AllReduceParams, allgather + from tensorrt_llm._torch.distributed.ops import MNNVLAllReduce + from tensorrt_llm.mapping import Mapping + + self.op, self.kernel, self.allgather = op, op._kernel_module(), allgather + # The C++ chain k3_markov replaces: the TeKit port registers it, main does not. + try: + self.cpp = torch.ops.trtllm.dspark_markov_chain + except (AttributeError, RuntimeError): + self.cpp = None + self.copies = copies or max(1, 16 // self.world) + self.mapping = Mapping( + world_size=self.world, rank=self.rank, gpus_per_node=gpus, tp_size=self.world + ) + # The model's all-reduces own this workspace; the C++ chain exchanges through its Lamport buffers too. + self.all_reduce_module = MNNVLAllReduce(self.mapping, torch.bfloat16) + self.mnnvl_workspaces = MNNVLAllReduce.allreduce_mnnvl_workspaces + weights = torch.Generator(device="cuda").manual_seed(7) + self.w1 = ( + torch.randn(VOCAB, MARKOV_RANK, generator=weights, device="cuda") * 0.25 + ).bfloat16() + self.w2 = ( + torch.randn(VOCAB, MARKOV_RANK, generator=weights, device="cuda") * 0.25 + ).bfloat16() + local = torch.Generator(device="cuda").manual_seed(100 + self.rank) + self.ar_input = torch.randn(8, HIDDEN, generator=local, device="cuda").bfloat16() + self.ar_params = AllReduceParams( + fusion_op=AllReduceFusionOp.RESIDUAL_RMS_NORM, + residual=torch.zeros(8, HIDDEN, dtype=torch.bfloat16, device="cuda"), + norm_weight=torch.ones(HIDDEN, dtype=torch.bfloat16, device="cuda"), + eps=1e-5, + ) + # The C++ chain's grid-sync words (zero; every launch leaves them zero) and per-CTA partial maxima, for the + # largest split: dsparkMarkovSyncWords = K + K B + 1, dsparkMarkovPartials = K B grid pairs, grid <= the SMs. + self.sms = torch.cuda.get_device_properties( + torch.cuda.current_device() + ).multi_processor_count + block, batch = max(BLOCKS), max(BATCHES) + self.sync_words = torch.zeros(block + block * batch + 1, dtype=torch.int32, device="cuda") + self.partials = torch.zeros( + 2 * block * batch * self.sms, dtype=torch.float32, device="cuda" + ) + + def say(self, *args) -> None: + if self.rank == 0: + print(*args, flush=True) + + def say_row(self, row: dict) -> None: + cells = [row[c] for c in ("case", "B", "K", "S", "dtype", *FIELDS)] + self.say( + "| " + " | ".join(str(c) for c in cells) + f" | {'PASS' if row['ok'] else 'FAIL'} |" + ) + + def all_ranks(self, flag: bool) -> bool: + return all(self.comm.allgather(bool(flag))) + + def shards(self): + """(S, push copies): the TP16 slice with the copies, then this group's full-vocabulary shard.""" + shards = [(TP16_SHARD, self.copies)] + if VOCAB // self.world != TP16_SHARD: + shards.append((VOCAB // self.world, 1)) + return shards + + def shard(self, kind: str): + return (TP16_SHARD, self.copies) if kind == "tp16" else (VOCAB // self.world, 1) + + def shard_weights(self, shard: int) -> torch.Tensor: + return self.w2[self.rank * shard : (self.rank + 1) * shard] + + def generators(self, seed: int): + """(shared, local) CUDA generators: shared draws are the same on every rank (anchors, acceptance), local ones + are this rank's (its logits).""" + shared = torch.Generator(device="cuda").manual_seed(seed) + local = torch.Generator(device="cuda").manual_seed(seed + 7919 * (self.rank + 1)) + return shared, local + + def all_reduce(self): + return self.all_reduce_module(self.ar_input, all_reduce_params=self.ar_params) + + def emulated_chain(self, base, first, w2s, shard: int) -> torch.Tensor: + """The C++ chain's corrected logits (fp32 [B, K, S]) in torch, position by position with the TP-gathered greedy + token (NaN never wins) as the next anchor: lane l's FMA chain from +0 over columns 8 l .. 8 l + 7 (the products + of bf16 values are exact in fp32, so fp32 adds), the 32 lane sums added in pairs at lane distances 16, 8, 4, 2, + 1, the sum rounded to bf16 (nearest even) and added to the fp32 base; an anchor outside the vocabulary adds a + zero bias.""" + batch, block = base.shape[:2] + w2f = w2s.float().view(1, shard, 32, 8) + base32 = base.float() + prev = first.long() + out = [] + for k in range(block): + valid = (prev >= 0) & (prev < VOCAB) + w1r = self.w1[prev.clamp(0, VOCAB - 1)].float().view(batch, 1, 32, 8) + products = w2f * w1r + lanes = torch.zeros(batch, shard, 32, device="cuda") + for j in range(8): + lanes = lanes + products[..., j] + for half in (16, 8, 4, 2, 1): + lanes = lanes[..., :half] + lanes[..., half : 2 * half] + bias = lanes[..., 0].bfloat16().float() + bias = torch.where(valid[:, None], bias, torch.zeros_like(bias)) + corrected = base32[:, k] + bias + out.append(corrected) + prev = self.sampler(corrected.unsqueeze(1), shard, ignore_nan=True)[:, 0].long() + return torch.stack(out, dim=1) + + def chain(self, base, first, w2s, shard: int) -> torch.Tensor: + """The exact reference's corrected logits: the C++ chain when the build has it, else its torch emulation.""" + if self.cpp is not None: + return self.cpp_chain(base, first, w2s, shard) + return self.emulated_chain(base, first, w2s, shard) + + def torch_chain_failures( + self, base, first, w2s, shard: int, out, ignore_nan: bool = False + ) -> int: + """Main's torch chain at every position, fed k3_markov's tokens as the anchors: ``dspark_markov_step_bias`` + (an F.linear of the bf16 weights, one bf16 rounding) added to the fp32 base. The dot products' summation + orders differ, so k3_markov's corrected logits must be within one bf16 ulp of the bias plus the fp32 + accumulation bound (2^-16 of the sum of the products' magnitudes, which covers a sum near zero) of it, and its + token a global maximum of it within that tolerance. Returns this rank's failing (b, k) positions.""" + from tensorrt_llm._torch.models.modeling_speculative import ( + dspark_markov_step_bias, + markov_prev_embeddings, + ) + + corrected, tokens = out[0], out[1] + base32, lo = base.float(), self.rank * shard + bad = 0 + for k in range(base.shape[1]): + prev = first.long() if k == 0 else tokens[:, k - 1].long() + bias = dspark_markov_step_bias(prev, self.w1, w2s).float() + ref = base32[:, k] + bias + magnitude = torch.nn.functional.linear( + markov_prev_embeddings(prev, self.w1).float().abs(), w2s.float().abs() + ) + tol = bias.abs() * 2.0**-7 + magnitude * 2.0**-16 + ref.abs().nan_to_num(0.0) * 2.0**-22 + got = corrected[:, k] + close = ((got - ref).abs() <= tol) | (got.isnan() & ref.isnan()) + masked = ref.masked_fill(ref.isnan(), float("-inf")) + token = tokens[:, k].long() + mine = (token >= lo) & (token < lo + shard) + at = masked.gather(1, (token - lo).clamp(0, shard - 1).unsqueeze(1)).squeeze(1) + at = torch.where(mine, at, torch.full_like(at, float("-inf"))) + stats = torch.stack( + [masked.max(dim=-1).values, at, tol.nan_to_num(0.0).max(dim=-1).values] + ) + every = torch.stack(self.comm.allgather(stats.cpu())) # [ranks, 3, B] + best, value, slack = ( + every[:, 0].max(0).values, + every[:, 1].max(0).values, + every[:, 2].max(0).values, + ) + token_ok = value >= best - 2 * slack + # A NaN wins the plain sampler's row; k3_markov's NaN-free choice is checked by the exact reference. + if not ignore_nan: + token_ok |= torch.stack(self.comm.allgather(ref.isnan().any(dim=-1).cpu())).any(0) + rows_ok = close.all(dim=-1).cpu() + bad += int((~(rows_ok & token_ok)).sum()) + for b in ( + (~rows_ok).nonzero().flatten().tolist()[:2] + ): # the worst element of a failing row, for the log + excess = ((got[b] - ref[b]).abs() - tol[b]).nan_to_num(float("-inf")) + v = int(excess.argmax()) + print(f"torch chain: rank {self.rank} position {k} request {b} column {v}: k3 {got[b, v].item():.9g}, " + f"main {ref[b, v].item():.9g}, bias {bias[b, v].item():.9g}, sum |products| " + f"{magnitude[b, v].item():.6g}, tolerance {tol[b, v].item():.3g}", flush=True) # fmt: skip + return bad + + def main_reference(self, base, first, w2s, shard: int, acc): + """Main's path: its torch chain (``dspark_markov_step_bias``) with the TP-gathered greedy token per position, + then next_new: (corrected fp32 [B, K, S], tokens int32 [B, K], next_new int32 [B, K + 1]).""" + from tensorrt_llm._torch.models.modeling_speculative import dspark_markov_step_bias + + base32, prev = base.float(), first.long() + corrected, tokens = [], [] + for k in range(base.shape[1]): + step = base32[:, k] + dspark_markov_step_bias(prev, self.w1, w2s).float() + token = self.sampler(step.unsqueeze(1), shard)[:, 0] + corrected.append(step) + tokens.append(token) + prev = token.long() + tokens = torch.stack(tokens, dim=1) + return torch.stack(corrected, dim=1), tokens, next_new_reference(acc, tokens) + + def cpp_chain(self, base, first, w2s, shard: int) -> torch.Tensor: + """trtllm::dspark_markov_chain on fp32 logits (bf16 ones cast with .float(), as the model's head did).""" + ws = self.mnnvl_workspaces[self.mapping] + return self.cpp( + base.float(), first, self.w1, w2s, self.rank * shard, self.sync_words, self.partials, + ws["uc_buffer"].view(torch.bfloat16).view(3, -1), ws["buffer_flags"], + ) # fmt: skip + + def sampler( + self, corrected: torch.Tensor, shard: int, ignore_nan: bool = False + ) -> torch.Tensor: + """SpecWorkerBase.greedy_sample_draft_with_tp_gather on [B, K, S] corrected logits: each rank's first maximum + (NaN as -inf with ``ignore_nan``), the all-gather of (index, value), the first maximum over the ranks.""" + flat = corrected.reshape(-1, shard) + if ignore_nan: + flat = flat.masked_fill(flat.isnan(), float("-inf")) + values, argmax = torch.max(flat, dim=-1, keepdim=True) + index = (argmax.to(torch.int32) + self.rank * shard).float() + combined = torch.stack([index, values.float()], dim=-1).flatten(-2) + gathered = self.allgather(combined, self.mapping, dim=-1) + best = torch.argmax(gathered[..., 1::2], dim=-1, keepdim=True) + tokens = torch.gather(gathered[..., 0::2], -1, best).squeeze(-1).to(torch.int32) + return tokens.view(corrected.shape[:-1]) + + def reference(self, base, first, w2s, shard: int, acc, ignore_nan: bool = False): + """The production path: (corrected fp32 [B, K, S], tokens int32 [B, K], next_new int32 [B, K + 1]).""" + corrected = self.chain(base, first, w2s, shard) + tokens = self.sampler(corrected, shard, ignore_nan) + return corrected, tokens, next_new_reference(acc, tokens) + + def k3(self, base, first, w2s, shard: int, acc, copies: int, rewind=None): + """k3_markov as the model calls it (op.markov_chain); ``rewind``: (kv_lens, rewind, rewind_first, ...).""" + kv_lens, amounts, rewind_first = (None, None, 0) if rewind is None else rewind[:3] + return self.op.markov_chain( + self.mapping, base, first, self.w1, w2s, self.rank * shard, *acc, push_copies=copies, + kv_lens=kv_lens, rewind=amounts, rewind_first=rewind_first, + ) # fmt: skip + + def same_on_ranks(self, *tensors) -> bool: + every = self.comm.allgather([t.cpu() for t in tensors]) + return all(identical(a, b) for other in every[1:] for a, b in zip(every[0], other)) + + def summarize(self, failures: dict, **cells) -> dict: + """A table row from this rank's failure counts, summed over the ranks (a field without a count: '-').""" + every = self.comm.allgather(failures) + row = dict(cells, ok=True) + for field in FIELDS: + if field not in failures: + row[field] = "-" + continue + count = sum(r[field] for r in every) + row[field] = "yes" if count == 0 else f"no ({count})" + row["ok"] = row["ok"] and count == 0 + return row + + +def shard_support(g: Group, shard: int, copies: int): + """(rejected splits, consistent, description) of a shard: ``pick_grid`` for every split, ``supports`` agreeing + with it, and k3_markov raising ValueError on a rejected split.""" + grids = {(b, k): g.op.pick_grid(shard, k, b) for b in BATCHES for k in BLOCKS} + w2s = g.shard_weights(shard) + agree = all( + g.op.supports(torch.empty(b, k, shard, dtype=dtype, device="cuda"), g.w1, w2s) == (grid > 0) + for (b, k), grid in grids.items() + for dtype in DTYPES + ) + rejected = {split for split, grid in grids.items() if grid == 0} + used = sorted({grid for grid in grids.values() if grid > 0}) + line = f"S = {shard} ({g.world * copies} slots per exchange): " + if used: + rows = "/".join(str(shard // grid) for grid in used) + line += f"pick_grid {'/'.join(map(str, used))} CTAs of {rows} rows; " + raises = True + if rejected: + b, k = min(rejected) + shared, _ = g.generators(case_seed("rejected", shard)) + base = torch.zeros(b, k, shard, device="cuda") + first = torch.zeros(b, dtype=torch.long, device="cuda") + try: + g.k3(base, first, w2s, shard, acceptance(b, k, shared), copies) + raises = False + except ValueError: + pass + budget, multiple = g.kernel.SMEM_ROW_BUDGET, g.kernel.WARPS * g.kernel.BLOCK_ROWS + max_grid = min(g.op.WORKSPACE_MAX_GRID, g.sms) // 2 * 2 + line += ( + f"REJECTED for {len(rejected)} of {len(grids)} splits (rows per CTA a multiple of {multiple}, " + f"at most SMEM_ROW_BUDGET = {budget}; an even grid of at most min(WORKSPACE_MAX_GRID = " + f"{g.op.WORKSPACE_MAX_GRID}, {g.sms} SMs): S <= {max_grid * (budget // multiple * multiple)}); " + f"k3_markov {'raises ValueError' if raises else 'DOES NOT RAISE'}; " + ) + line += f"supports() {'agrees' if agree else 'DISAGREES'}" + return rejected, g.all_ranks(agree and raises), line + + +def check( + g: Group, + case: str, + shard: int, + copies: int, + base, + first, + w2s, + acc, + rewind, + ignore_nan: bool = False, + expect: Optional[torch.Tensor] = None, + empty_rewind: bool = False, +) -> dict: + """One case: the reference, k3_markov, its rerun and its call with the KV-length rewind (and with an empty one).""" + batch, block = base.shape[:2] + ref = g.reference(base, first, w2s, shard, acc, ignore_nan) + out = g.k3(base, first, w2s, shard, acc, copies) + again = g.k3(base, first, w2s, shard, acc, copies) + kv_lens, amounts, _, kv_want = rewind + rewind_ok = True + if empty_rewind: # kv_lens given, nothing to rewind: kv_lens untouched + spare = kv_lens.clone() + empty = g.k3(base, first, w2s, shard, acc, copies, rewind=(spare, amounts[:0], 0)) + rewind_ok = identical(spare, kv_lens) and same_outputs(empty, out) + rewound = g.k3(base, first, w2s, shard, acc, copies, rewind=rewind) + rewind_ok = rewind_ok and identical(kv_lens, kv_want) and same_outputs(rewound, out) + tokens_ok = identical(out[1], ref[1]) and (expect is None or identical(out[1], expect)) + failures = dict( + corrected=differing_words(out[0], ref[0]), + tokens=int(not tokens_ok), + next_new=int(not identical(out[2], ref[2])), + rerun=int(not same_outputs(again, out)), + ranks=int(not g.same_on_ranks(out[1], out[2])), + rewind=int(not rewind_ok), + torch=g.torch_chain_failures(base, first, w2s, shard, out, ignore_nan), + ) + if g.cpp is not None: + failures["emulation"] = differing_words(g.emulated_chain(base, first, w2s, shard), ref[0]) + return g.summarize( + failures, case=case, B=batch, K=block, S=shard, dtype=DTYPE_NAMES[base.dtype] + ) + + +def random_case(g: Group, shard: int, copies: int, batch: int, block: int, dtype) -> dict: + seed = case_seed("random", shard, batch, block, dtype) + shared, local = g.generators(seed) + base = (torch.randn(batch, block, shard, generator=local, device="cuda") * 3.0).to(dtype) + first = torch.randint(0, VOCAB, (batch,), generator=shared, device="cuda") + acc, rewind = acceptance(batch, block, shared), rewind_spec(batch, seed) + return check(g, "random", shard, copies, base, first, g.shard_weights(shard), acc, rewind) + + +def crafted_case( + g: Group, kind: str, shard: int, copies: int, batch: int, block: int, dtype +) -> dict: + seed = case_seed(kind, shard, batch, block, dtype) + shared, local = g.generators(seed) + base = torch.randn(batch, block, shard, generator=local, device="cuda") * 3.0 + first = torch.randint(0, VOCAB, (batch,), generator=shared, device="cuda") + w2s = g.shard_weights(shard) + acc = acceptance(batch, block, shared) + rewind = rewind_spec(batch, seed) + extra = {} + if kind == "ties": + rows = [r for r in TIE_ROWS if r < shard] + [shard // 2 + 3, shard - 5] + w2s = w2s.clone() + w2s[rows] = g.w2[11] # one markov_w2 row, so one bias, at every tie row of every rank + expect = torch.empty(batch, block, dtype=torch.int32) + tied = [] + for b in range(batch): + for k in range(block): + # The tie set of (b, k): every rank from the lead one, minus the lead's first `drop` rows, so the winner + # (its first kept row) moves over the levels of the reduction and over the ranks. + lead, drop = (b + k) % g.world, (b + 2 * k) % (len(rows) - 1) + expect[b, k] = lead * shard + rows[drop] + if g.rank >= lead: + tied += [(b, k, r) for r in (rows[drop:] if g.rank == lead else rows)] + if tied: + base[tuple(torch.tensor(tied, device="cuda").t())] = 40.0 + extra["expect"] = expect.cuda() + elif kind == "anchors": + outside = torch.tensor([-1, VOCAB, VOCAB + 5, -100], device="cuda") + first[0::2] = outside[: (batch + 1) // 2] + elif kind == "nan": + nan = float("nan") + b_idx = torch.arange(batch, device="cuda").repeat_interleave(block) + k_idx = torch.arange(block, device="cuda").repeat(batch) + row = 9 + 64 * ((b_idx + 3 * k_idx) % 5) + base[:, :, 7] = nan # one row of every rank at every position + # A dominant value on one rank per (b, k), between NaN rows of its 8-row block on every rank. + base[b_idx, k_idx, row - 1] = nan + base[b_idx, k_idx, row + 1] = nan + mine = (b_idx + k_idx + 1) % g.world == g.rank + base[b_idx[mine], k_idx[mine], row[mine]] = 50.0 + if g.rank == 1: + base[0, block // 2] = nan # a whole position of request 0 on one rank + extra["ignore_nan"] = True + else: # rewind and acceptance edges + accepted, num_accepted, rows = acceptance( + batch, block, shared, extra_rows=3, extra_width=2, min_accepted=0 + ) + num_accepted[0] = 0 # accepted column -1: the last one, as torch indexing + # num_accepted and the accepted rows as the model passes them: [num_contexts:] slices (a 4-byte offset). + pad = torch.zeros(1, dtype=torch.int32, device="cuda") + acc = (accepted, torch.cat([pad, num_accepted])[1:], torch.cat([pad, rows])[1:]) + rewind = rewind_spec(batch, seed, count=max(1, batch // 2), first=0, after=3) + extra["empty_rewind"] = True + return check(g, CRAFTED[kind], shard, copies, base.to(dtype), first, w2s, acc, rewind, **extra) + + +def mixed_step_case(g: Group, contexts: int, batch: int, block: int = 7) -> dict: + """A mixed step as DSparkWorker passes it: ``batch`` generation requests after ``contexts`` context requests, + num_accepted / accepted_rows the int32 slices [contexts:] of the step's per-request tensors and the KV-length + rewind on the whole kv_lens from rewind_first = ``contexts``. Against the reference, and bit for bit against the + same call on contiguous copies (kv_lens[contexts:] from rewind_first 0) in the 'rerun' column.""" + shard, copies = TP16_SHARD, g.copies + seed = case_seed("mixed", contexts, batch, block) + shared, local = g.generators(seed) + base = (torch.randn(batch, block, shard, generator=local, device="cuda") * 3.0).bfloat16() + first = torch.randint(0, VOCAB, (batch,), generator=shared, device="cuda") + w2s = g.shard_weights(shard) + accepted, num_accepted, rows = acceptance(contexts + batch, block, shared) + sliced = (accepted, num_accepted[contexts:], rows[contexts:]) + copied = (accepted, num_accepted[contexts:].clone(), rows[contexts:].clone()) + kv_lens, amounts, _, kv_want = rewind_spec(batch, seed, first=contexts) + kv_copy = kv_lens[contexts:].clone() + ref = g.reference(base, first, w2s, shard, copied) + out = g.k3(base, first, w2s, shard, sliced, copies, rewind=(kv_lens, amounts, contexts)) + out_copies = g.k3(base, first, w2s, shard, copied, copies, rewind=(kv_copy, amounts, 0)) + failures = dict( + corrected=differing_words(out[0], ref[0]), + tokens=int(not identical(out[1], ref[1])), + next_new=int(not identical(out[2], ref[2])), + rerun=int(not same_outputs(out, out_copies)), + ranks=int(not g.same_on_ranks(out[1], out[2])), + rewind=int(not (identical(kv_lens, kv_want) and identical(kv_copy, kv_want[contexts:]))), + torch=g.torch_chain_failures(base, first, w2s, shard, out), + ) + return g.summarize(failures, case=f"mixed step, {contexts} context requests first", B=batch, K=block, + S=shard, dtype="bf16") # fmt: skip + + +def eager_interleave(g: Group, calls: int = 60) -> dict: + """``calls`` eager k3_markov calls with B, K and the dtype varying (every fifth on S = 163840 / W when the kernel + splits it), every other one after an MNNVL all-reduce and every third followed by an extra C++ chain call, each + against the reference.""" + batches = (1, 8, 2, 5, 3, 7, 4, 6) + full = VOCAB // g.world + failures = dict(corrected=0, tokens=0, next_new=0) + shards, outputs = set(), [] + for i in range(calls): + batch, block = batches[i % len(batches)], BLOCKS[i % len(BLOCKS)] + dtype = DTYPES[i % len(DTYPES)] + on_full = i % 5 == 4 and g.op.pick_grid(full, block, batch) > 0 + shard, copies = (full, 1) if on_full else (TP16_SHARD, g.copies) + shards.add(shard) + shared, local = g.generators(case_seed("eager", i)) + base = (torch.randn(batch, block, shard, generator=local, device="cuda") * 3.0).to(dtype) + first = torch.randint(0, VOCAB, (batch,), generator=shared, device="cuda") + acc = acceptance(batch, block, shared) + w2s = g.shard_weights(shard) + if i % 2 == 0: + g.all_reduce() + out = g.k3(base, first, w2s, shard, acc, copies) + if i % 3 == 0 and g.cpp is not None: + g.cpp_chain(base, first, w2s, shard) + ref = g.reference(base, first, w2s, shard, acc) + failures["corrected"] += int(differing_words(out[0], ref[0]) > 0) + failures["tokens"] += int(not identical(out[1], ref[1])) + failures["next_new"] += int(not identical(out[2], ref[2])) + outputs += [out[1], out[2]] + failures["ranks"] = int(not g.same_on_ranks(*outputs)) + case = ( + f"{calls} eager calls, all-reduces{' and C++ chains' if g.cpp is not None else ''} between" + ) + shards = ", ".join(map(str, sorted(shards))) + return g.summarize(failures, case=case, B="1-8", K="1, 3, 7", S=shards, dtype="fp32, bf16") + + +def capture(fn, calls: int, stream): + """A CUDA graph of fn(0) .. fn(calls - 1), after two eager calls (compilation, NCCL and workspace setup).""" + stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(stream): + fn(0) + fn(1) + torch.cuda.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + for i in range(calls): + fn(i) + torch.cuda.synchronize() + return graph + + +def graph_replays(g: Group, batch: int, block: int, dtype, replays: int = 20) -> dict: + """A CUDA graph of [MNNVL all-reduce, k3_markov, k3_markov with the KV-length rewind] at S = 10240 with the + copies, replayed with every input rewritten in place, each replay against the reference.""" + shard, copies = TP16_SHARD, g.copies + w2s = g.shard_weights(shard) + shared, _ = g.generators(case_seed("graph", batch, block, dtype)) + base = torch.zeros(batch, block, shard, dtype=dtype, device="cuda") + first = torch.zeros(batch, dtype=torch.long, device="cuda") + acc = acceptance(batch, block, shared) + kv_lens, amounts, rewind_first, _ = rewind_spec(batch, 0) + outs = {} + + def rewrite(rep: int) -> torch.Tensor: + """New inputs, in place; returns the kv_lens the replay must leave.""" + seed = case_seed("graph", batch, block, dtype, rep) + shared, local = g.generators(seed) + base.copy_(torch.randn(batch, block, shard, generator=local, device="cuda") * 3.0) + first.copy_(torch.randint(0, VOCAB, (batch,), generator=shared, device="cuda")) + for dst, src in zip(acc, acceptance(batch, block, shared)): + dst.copy_(src) + new_kv_lens, new_amounts, _, want = rewind_spec(batch, seed) + kv_lens.copy_(new_kv_lens) + amounts.copy_(new_amounts) + return want + + def body(_): + # Held with the graph, so that no allocation inside the graph reuses the all-reduce output's address. + outs["all_reduce"] = g.all_reduce() + outs["plain"] = g.k3(base, first, w2s, shard, acc, copies) + outs["rewind"] = g.k3( + base, first, w2s, shard, acc, copies, rewind=(kv_lens, amounts, rewind_first) + ) + + rewrite(-1) + graph = capture(body, 1, torch.cuda.Stream()) + failures = dict(corrected=0, tokens=0, next_new=0, rerun=0, rewind=0) + outputs = [] + for rep in range(replays): + want = rewrite(rep) + torch.cuda.synchronize() + graph.replay() + torch.cuda.synchronize() + ref = g.reference(base, first, w2s, shard, acc) + plain, rewound = outs["plain"], outs["rewind"] + failures["corrected"] += int(differing_words(plain[0], ref[0]) > 0) + failures["tokens"] += int(not identical(plain[1], ref[1])) + failures["next_new"] += int(not identical(plain[2], ref[2])) + failures["rerun"] += int(not same_outputs(rewound, plain)) + failures["rewind"] += int(not identical(kv_lens, want)) + outputs += [plain[1].clone(), plain[2].clone()] + outs.clear() + del graph + failures["ranks"] = int(not g.same_on_ranks(*outputs)) + case = f"graph [all-reduce, k3_markov, k3_markov + rewind] x {replays} replays" + return g.summarize(failures, case=case, B=batch, K=block, S=shard, dtype=DTYPE_NAMES[dtype]) + + +def report(g: Group) -> int: + """The checks as a markdown table (rank 0); 0 when every check passes on every rank.""" + hosts = g.comm.allgather(f"{os.uname().nodename}:{torch.cuda.current_device()}") + g.say( + f"{torch.cuda.get_device_name()} x {g.world} ({', '.join(hosts)}); k3_markov: {g.op.__file__}" + ) + ok = True + runs = [] + for shard, copies in g.shards(): + rejected, consistent, line = shard_support(g, shard, copies) + g.say(line) + # The TP16 slice must take every split; a larger shard may be beyond the kernel (the model then falls back). + ok = ok and consistent and not (shard == TP16_SHARD and rejected) + runs.append((shard, copies, rejected)) + g.say( + "A failing cell shows its count over the ranks: differing words of the corrected logits, " + "else failing ranks or calls." + ) + g.say( + "Exact reference: " + + ("trtllm::dspark_markov_chain (C++); 'emulation': its torch emulation against it" + if g.cpp is not None else "the C++ chain's arithmetic in torch (this build has no C++ chain)") + + "; 'torch chain': main's torch chain fed k3_markov's tokens, positions beyond one bf16 ulp of the bias plus" + " the fp32 accumulation bound" + ) # fmt: skip + g.say(HEADER) + rows = [] + for shard, copies, rejected in runs: + for batch in BATCHES: + for block in BLOCKS: + if (batch, block) in rejected: + continue + for dtype in DTYPES: + rows.append(random_case(g, shard, copies, batch, block, dtype)) + g.say_row(rows[-1]) + for kind in CRAFTED: + for index, (batch, block) in enumerate(CRAFTED_SPLITS): + if (batch, block) in rejected: + continue + dtype = DTYPES[index % len(DTYPES)] + rows.append(crafted_case(g, kind, shard, copies, batch, block, dtype)) + g.say_row(rows[-1]) + for contexts in MIXED_CONTEXTS: + for batch in MIXED_GENS: + rows.append(mixed_step_case(g, contexts, batch)) + g.say_row(rows[-1]) + rows.append(eager_interleave(g)) + g.say_row(rows[-1]) + for batch, block in GRAPH_SPLITS: + for dtype in DTYPES: + rows.append(graph_replays(g, batch, block, dtype)) + g.say_row(rows[-1]) + not_run = [f"S = {shard}: {len(rejected)} splits" for shard, _, rejected in runs if rejected] + if not_run: + g.say("Not run (rejected by pick_grid, see above): " + ", ".join(not_run)) + ok = g.all_ranks(ok and all(row["ok"] for row in rows)) + g.say("ALL PASS" if ok else "FAIL") + return 0 if ok else 1 + + +def replay_us(g: Group, graph, stream) -> float: + """One replay, started behind an MPI barrier, in us.""" + torch.cuda.synchronize() + g.comm.Barrier() + start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + with torch.cuda.stream(stream): + start.record() + graph.replay() + end.record() + end.synchronize() + return start.elapsed_time(end) * 1e3 + + +def timing(g: Group, rounds: int = 10) -> None: + """Per split at S = 10240 with the copies: graphs of back-to-back calls of the production path and of k3_markov, + each call on its own markov_w2 shard copy; the max over the ranks of the median us per call.""" + shard, copies = TP16_SHARD, g.copies + shard_bytes = shard * MARKOV_RANK * 2 + # One shard copy per call: > 200 MB of other copies between two reads of one. + calls = -(-HBM_COLD_BYTES // shard_bytes) + 1 + gen = torch.Generator(device="cuda").manual_seed(31 + g.rank) + w2_copies = [ + (torch.randn(shard, MARKOV_RANK, generator=gen, device="cuda") * 0.25).bfloat16() + for _ in range(calls) + ] + stream = torch.cuda.Stream() + g.say( + f"{torch.cuda.get_device_name()} x {g.world}; S = {shard}, {g.world * copies} slots per k3_markov " + f"exchange; graphs of {calls} calls, each on its own markov_w2 shard copy ({calls * shard_bytes >> 20} " + f"MB: HBM-cold); {rounds} ABBA rounds; max over the ranks of the median us per call" + ) + chain = "C++ chain" if g.cpp is not None else "main's torch chain" + g.say(f"| B | K | grid | production (cast, {chain}, sampler, next_new) | k3_markov | speedup |") + g.say("| --: | --: | --: | --: | --: | --: |") + for batch in BATCHES: + for block in BLOCKS: + shared, local = g.generators(case_seed("time", batch, block)) + inputs = [] + for _ in range(4): + base = torch.randn(batch, block, shard, generator=local, device="cuda") * 3.0 + first = torch.randint(0, VOCAB, (batch,), generator=shared, device="cuda") + inputs.append((base.bfloat16(), first, acceptance(batch, block, shared))) + + def production(i): + base, first, acc = inputs[i % len(inputs)] + if g.cpp is not None: + return g.reference(base, first, w2_copies[i], shard, acc) + return g.main_reference(base, first, w2_copies[i], shard, acc) + + def fused(i): + base, first, acc = inputs[i % len(inputs)] + return g.k3(base, first, w2_copies[i], shard, acc, copies) + + graphs = [capture(production, calls, stream), capture(fused, calls, stream)] + for graph in graphs: + replay_us(g, graph, stream) + times = [[], []] + for rnd in range(rounds): + for arm in (0, 1) if rnd % 2 == 0 else (1, 0): + times[arm].append(replay_us(g, graphs[arm], stream) / calls) + medians = g.comm.allgather([statistics.median(t) for t in times]) + prod_us, k3_us = (max(m[arm] for m in medians) for arm in (0, 1)) + grid = g.op.pick_grid(shard, block, batch) + g.say( + f"| {batch} | {block} | {grid} | {prod_us:.2f} | {k3_us:.2f} | {prod_us / k3_us:.2f}x |" + ) + del graphs + + +# ---------------------------------------------------------------------------------------------------------------- +# pytest (srun -n W --mpi=pmix python3 -m pytest -p no:cacheprovider test_k3_markov.py): every rank runs the same +# tests in the same order, and every collective of a test happens before its assert. +# ---------------------------------------------------------------------------------------------------------------- + +_state = {} + + +def group() -> Group: + """This process' rank of the group, built by the first test (its workspaces are allocated collectively).""" + if "group" not in _state: + _state["group"] = Group() + return _state["group"] + + +@contextlib.contextmanager +def collective(): + """Inference mode; an exception on one rank aborts the job (its peers would wait for it in a collective).""" + with torch.inference_mode(): + try: + yield + except Exception: + traceback.print_exc() + from mpi4py import MPI + + MPI.COMM_WORLD.Abort(1) + raise + + +def test_shard_support(): + with collective(): + g = group() + results = [(shard, *shard_support(g, shard, copies)) for shard, copies in g.shards()] + for shard, rejected, consistent, line in results: + assert consistent, line + assert not (shard == TP16_SHARD and rejected), line + + +@pytest.mark.parametrize("dtype", DTYPES, ids=[DTYPE_NAMES[d] for d in DTYPES]) +@pytest.mark.parametrize("block", BLOCKS) +@pytest.mark.parametrize("batch", BATCHES) +@pytest.mark.parametrize("shard_kind", ["tp16", "full"]) +def test_split(shard_kind, batch, block, dtype): + with collective(): + g = group() + shard, copies = g.shard(shard_kind) + if g.op.pick_grid(shard, block, batch) == 0: + pytest.skip(f"pick_grid rejects S = {shard} (see test_shard_support)") + row = random_case(g, shard, copies, batch, block, dtype) + assert row["ok"], row + + +@pytest.mark.parametrize( + "index", range(len(CRAFTED_SPLITS)), ids=[f"{b}x{k}" for b, k in CRAFTED_SPLITS] +) +@pytest.mark.parametrize("kind", list(CRAFTED)) +@pytest.mark.parametrize("shard_kind", ["tp16", "full"]) +def test_crafted(shard_kind, kind, index): + batch, block = CRAFTED_SPLITS[index] + with collective(): + g = group() + shard, copies = g.shard(shard_kind) + if g.op.pick_grid(shard, block, batch) == 0: + pytest.skip(f"pick_grid rejects S = {shard} (see test_shard_support)") + row = crafted_case(g, kind, shard, copies, batch, block, DTYPES[index % len(DTYPES)]) + assert row["ok"], row + + +@pytest.mark.parametrize("batch", MIXED_GENS) +@pytest.mark.parametrize("contexts", MIXED_CONTEXTS) +def test_mixed_step(contexts, batch): + with collective(): + row = mixed_step_case(group(), contexts, batch) + assert row["ok"], row + + +def test_eager_interleave(): + with collective(): + row = eager_interleave(group()) + assert row["ok"], row + + +@pytest.mark.parametrize("dtype", DTYPES, ids=[DTYPE_NAMES[d] for d in DTYPES]) +@pytest.mark.parametrize("batch,block", GRAPH_SPLITS, ids=[f"{b}x{k}" for b, k in GRAPH_SPLITS]) +def test_graph_replay(batch, block, dtype): + with collective(): + row = graph_replays(group(), batch, block, dtype) + assert row["ok"], row + + +# ---------------------------------------------------------------------------------------------------------------- +# Script (srun -n W --mpi=pmix python3 test_k3_markov.py [report | time] ...) +# ---------------------------------------------------------------------------------------------------------------- + + +def main() -> int: + parser = argparse.ArgumentParser( + description="trtllm::k3_markov op check (see the module docstring)" + ) + parser.add_argument( + "mode", + nargs="?", + choices=("report", "time"), + help="the checks only, or the timing only (default: the checks, then the timing)", + ) + parser.add_argument( + "--copies", + type=int, + default=None, + help="push copies per rank at S = 10240 (default 16 / W: the exchange volume of 16 ranks)", + ) + parser.add_argument("--skip-perf", action="store_true", help="no timing after the checks") + parser.add_argument("--rounds", type=int, default=10, help="ABBA timing rounds") + args = parser.parse_args() + if args.copies is not None and args.copies < 1: + parser.error("--copies must be >= 1") + if not _sm100() or _world_size() < 2: + print( + "needs SM100 GPUs and >= 2 MPI ranks: srun -n W --mpi=pmix python3 test_k3_markov.py", + flush=True, + ) + return 2 + with torch.inference_mode(): + g = Group(args.copies) + status = 0 if args.mode == "time" else report(g) + if args.mode == "time" or (args.mode is None and not args.skip_perf): + timing(g, args.rounds) + return status + + +if __name__ == "__main__": + try: + sys.exit(main()) + except Exception: # a rank stopping here would leave its peers waiting in a collective + traceback.print_exc() + from mpi4py import MPI + + MPI.COMM_WORLD.Abort(1) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_spec_accept.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_spec_accept.py new file mode 100644 index 000000000000..8b231a8d2f77 --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_spec_accept.py @@ -0,0 +1,681 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``trtllm::k3_spec_accept`` (one decode step's speculative acceptance and the block drafter's inputs) against the +torch op sequence it replaces, bit for bit, on one GPU at the in-model shapes (V = 163840 target logits, hidden 7168). + +The reference, as the model runs it for a decode step of B generation requests with K drafts each (no context requests, +greedy strict acceptance): DFlashWorker._refresh_ctx_block_tables, SpecWorkerBase._sample_and_accept_draft_tokens_base +with _apply_force_accepted_tokens (the model's RNG pool and strides), the KDA replay record of +MambaHybridCacheManagerV2.update_mamba_states, the kv_lens update of DFlashWorker._prepare_kv_for_draft_forward and +DFlashWorker.prepare_1st_drafter_inputs up to the fc (bonus tokens, positions, noise embedding). + +Every split B = 1 .. 8 requests x K + 1 in {2, 4, 8} tokens at the drafter block widths the model passes (K + 1: DFlash; +K: DSpark's shift_label; 8: a draft-length schedule running K < max_draft_len = 7), 20 consecutive steps per case with +the state carried: natural acceptance with every accepted-prefix length 0 .. K (mixed over the requests), ties across +CTAs (the lowest index wins), NaN rows (the first NaN wins, torch's rule, also over +inf), +-inf and +-0 maxima, forced +acceptance (an integer, a fractional and an above-K value), dummy requests (and CUDA-graph padding sharing one state and +one context slot), placeholder (-1) and zero block offsets, block tables of 72 / 300 / 4100 columns, context lengths at +the max_ctx clamp. At every step every output and every in-place state tensor is bit-identical to the reference, the +inputs are untouched and a second run on a clone of the state gives the same bits (determinism). CUDA graphs: one +captured call per split family, replayed 20 times with the logits, drafts, dummy mask, block offsets and context lengths +rewritten in place and the state carried. + +Table: ``python3 test_k3_spec_accept.py report``. Timing: ``python3 test_k3_spec_accept.py time`` (batch 1, then every +split: CUDA graphs of 20 back-to-back calls of the reference ops and of the kernel, median over 12 replays in +alternating order). The vocabulary-sharded variant: test_k3_spec_accept_sharded.py (MPI). +""" + +import itertools +import statistics +import sys + +import pytest +import torch +import torch.nn.functional as F + +V = 163840 # target vocabulary +HIDDEN = 7168 # drafter hidden size (the noise embedding rows) +CTA_COLS = 1024 # logits columns one CTA of the kernel reduces (ties straddle these boundaries) +TP16_COLUMNS = 10240 # a TP16 rank's vocabulary shard +SEQS = 16 # rows of the per-request state; the rows past the batch must stay untouched +SLOTS = 20 # KDA replay / context slots +POOLS, POOL_IDX = 2, 1 # KV pools in the block offsets; the draft pool's index +DIVISOR = 10 # encoded block offset -> pool block (5 drafter layers x K / V) +MAX_CTX = 4100 +# Block-table width by batch % 3: one, two and 17 passes of the kernel's 256 threads. +MAX_BLOCKS = (72, 300, 4100) +PADDING_FROM = 5 # batches of 5 or more end with two CUDA-graph padding requests +STEPS = 20 +# Forced acceptance per K: (an integer value, a fractional value) of TLLM_SPEC_DECODE_FORCE_NUM_ACCEPTED_TOKENS. +FORCES = {1: (1.0, 0.4), 3: (2.0, 1.6), 7: (5.0, 5.4)} + +# (B requests, K + 1 tokens each), batch 1 first. +SPLITS = [(b, t) for b in range(1, 9) for t in (2, 4, 8)] +GRAPH_SPLITS = [(b, t) for b in (1, 4, 8) for t in (2, 4, 8)] +WIDE_SPLITS = [(b, t) for b in (1, 4, 8) for t in (2, 4)] +# (name, logits kind, forced acceptance) +CASES = ( + ("natural", "plain", None), + ("ties", "ties", None), + ("nan", "nan", None), + ("inf / +-0", "signed", None), + ("forced int", "plain", "int"), + ("forced frac", "plain", "frac"), + ("forced > K", "plain", "over"), +) +GRAPH_KINDS = ("plain", "ties", "nan", "signed") + +OUTPUTS = ("accepted", "num_acc", "rewind", "bonus", "qpos", "cpos", "noise") +STATE_KEYS = ("block_counts", "block_tables", "prev_acc", "kv_lens", "rng_counter") +INPUT_KEYS = ("block_off", "state_idx", "dummy", "batch_to_slot", "ctx_len", "mask_row") +SHARED_KEYS = ("embed", "rng_pool") # read-only and large: shared by the clones of a state + + +def _sm100() -> bool: + return torch.cuda.is_available() and torch.cuda.get_device_capability()[0] == 10 + + +pytestmark = pytest.mark.skipif(not _sm100(), reason="needs SM100 (CTM kernel)") + + +def _op(): + import tensorrt_llm # noqa: F401 (registers the trtllm ops) + from tensorrt_llm._torch.cute_dsl_kernels.k3_spec_accept import op + + return op + + +def _spec(): + """The model's forced-acceptance constants live here.""" + from tensorrt_llm._torch.speculative import interface + + return interface + + +_cache = {} + + +def embedding() -> torch.Tensor: + """The drafter embedding, bf16 [V, HIDDEN] (2.3 GB), built once: the kernel only reads it.""" + embed = _cache.get("embed") + if embed is None: + gen = torch.Generator(device="cuda").manual_seed(7) + embed = torch.randn(V, HIDDEN, generator=gen, device="cuda", dtype=torch.bfloat16) + _cache["embed"] = embed + return embed + + +def rng_pool() -> torch.Tensor: + """The model's fixed forced-acceptance pool (SpecWorkerBase._ensure_force_accept_rng_state), built once.""" + pool = _cache.get("rng_pool") + if pool is None: + spec = _spec() + gen = torch.Generator(device="cpu").manual_seed(spec._FORCE_ACCEPT_RNG_SEED) + pool = torch.rand(spec._FORCE_ACCEPT_RNG_POOL_SIZE, generator=gen).cuda() + _cache["rng_pool"] = pool + return pool + + +def force_value(drafts: int, forced) -> float: + if forced is None: + return 0.0 + if forced == "over": # min(int(f) + 1, K + 1) = K + 1: the fraction is ignored + return drafts + 2.5 + return FORCES[drafts][0 if forced == "int" else 1] + + +# ---------------------------------------------------------------------------------------------------------------- +# The decode step's state and inputs +# ---------------------------------------------------------------------------------------------------------------- + + +def block_offsets(gen, max_blocks: int, salt: int = 0) -> torch.Tensor: + """The draft KV manager's encoded block offsets [pools, seqs, K / V, max_blocks]: request b holds 5 + 9 b (+ 3 salt) + blocks for even b, max_blocks + 1 - b for odd b (request 1 the whole row), placeholders (-1) after them, a + placeholder hole and a zero offset in request 0; the other pool and the V plane hold other values.""" + off = torch.randint( + 0, 1 << 22, (POOLS, SEQS, 2, max_blocks), generator=gen, device="cuda", dtype=torch.int32 + ) + seq = torch.arange(SEQS, device="cuda") + used = torch.where(seq % 2 == 0, 5 + 9 * seq + 3 * salt, max_blocks + 1 - seq) + tail = torch.arange(max_blocks, device="cuda")[None, :] >= used[:, None] + off.masked_fill_(tail[None, :, None, :], -1) + off[POOL_IDX, 0, 0, 1] = -1 + off[POOL_IDX, 0, 0, 3] = 0 + return off + + +def context_lengths(gen, batch_to_slot: torch.Tensor, batch: int) -> torch.Tensor: + """Context lengths [SLOTS] (int64): request 0's one below max_ctx (its query positions clamp once it accepts 2 + tokens), request 1's at max_ctx, the others random.""" + ctx = torch.randint(100, MAX_CTX - 16, (SLOTS,), generator=gen, device="cuda") + ctx[batch_to_slot[0]] = MAX_CTX - 1 + if batch > 1: + ctx[batch_to_slot[1]] = MAX_CTX + return ctx + + +def make_state(gen, batch: int) -> dict: + """One decode step's state with the model's dtypes: distinct KDA slots and context slots per request, except the + two CUDA-graph padding requests of a batch of 5 or more (one state slot, one context slot); stale block counts and + tables the kernel overwrites for the batch's rows only.""" + max_blocks = MAX_BLOCKS[batch % len(MAX_BLOCKS)] + state_idx = torch.randperm(SLOTS, generator=gen, device="cuda")[:SEQS].to(torch.int32) + batch_to_slot = torch.randperm(SLOTS, generator=gen, device="cuda")[:SEQS] + if batch >= PADDING_FROM: + state_idx[batch - 1] = state_idx[batch - 2] + batch_to_slot[batch - 1] = batch_to_slot[batch - 2] + return dict( + pool_idx=POOL_IDX, + divisor=DIVISOR, + block_off=block_offsets(gen, max_blocks), + block_counts=torch.randint(0, 99, (SEQS,), generator=gen, device="cuda"), + block_tables=torch.randint(0, 99, (SEQS, max_blocks), generator=gen, device="cuda", dtype=torch.int32), + prev_acc=torch.randint(0, 8, (SLOTS,), generator=gen, device="cuda", dtype=torch.int32), + state_idx=state_idx, + dummy=torch.zeros(SEQS, dtype=torch.bool, device="cuda"), + kv_lens=torch.randint(100, 2000, (SEQS,), generator=gen, device="cuda", dtype=torch.int32), + batch_to_slot=batch_to_slot, + ctx_len=context_lengths(gen, batch_to_slot, batch), + embed=embedding(), + mask_row=torch.randn(HIDDEN, generator=gen, device="cuda").bfloat16(), + rng_pool=rng_pool(), + rng_counter=torch.full((1,), 41, dtype=torch.int64, device="cuda"), + ) # fmt: skip + + +def clone_state(st: dict) -> dict: + return { + k: v.clone() if torch.is_tensor(v) and k not in SHARED_KEYS else v for k, v in st.items() + } + + +def set_dummies(states, batch: int, step: int) -> None: + """A step's dummy-request mask: request b when (b + step) % 4 == 3, and the CUDA-graph padding (the last two + requests of a batch of 5 or more) at every step.""" + seq = torch.arange(SEQS, device="cuda") + mask = (seq + step) % 4 == 3 + if batch >= PADDING_FROM: + mask |= (seq == batch - 1) | (seq == batch - 2) + for st in states: + st["dummy"].copy_(mask) + + +def shape_logits(logits: torch.Tensor, gen, kind: str, step: int, shards: int = 1) -> None: + """Shapes a step's logits [rows, vocab] (fp32 or bf16: the values are exact in both) in place for ``kind``. + + ties: 9.0 at 3 random columns, the column after the first, both sides of a CTA boundary (and of a shard boundary). + rank_ties: 7.0 at one column of every shard (rank 0's wins). nan: a whole NaN row; rows with +inf at column 2 before + NaNs in both halves, rows with +inf before two random NaNs, rows with one NaN in the last CTA (shard). signed: two + +inf, a -inf row, -0.0 / +0.0 maxima in either order. negzero: -1.0 rows with -0.0 / +0.0 maxima across shards, a + -0.0 in the last shard only, and none (every column ties).""" + rows, vocab = logits.shape + dev = logits.device + r = torch.arange(rows, device=dev) + half = vocab // 2 + shard = vocab // shards + nan, inf = float("nan"), float("inf") + if kind == "ties": + cols = torch.randint(0, vocab, (rows, 3), generator=gen, device=dev) + edge = torch.randint(1, vocab // CTA_COLS, (rows, 1), generator=gen, device=dev) * CTA_COLS + extra = [(cols[:, :1] + 1) % vocab, edge - 1, edge] + if shards > 1: + k = 1 + (r[:, None] + step) % (shards - 1) + extra += [k * shard - 1, k * shard] + logits.scatter_(1, torch.cat([cols] + extra, dim=1), 9.0) + elif kind == "rank_ties": + q = torch.arange(shards, device=dev)[None, :] + logits.scatter_(1, q * shard + (r[:, None] * 37 + q * 11 + step) % shard, 7.0) + elif kind == "nan": + first = torch.randint(1, half, (rows,), generator=gen, device=dev) + later = first + torch.randint(1, half, (rows,), generator=gen, device=dev) + pattern = (r + step) % 3 + a, b, c = r[pattern == 0], r[pattern == 1], r[pattern == 2] + logits[a, 2] = inf + logits[a, half + 3] = nan + logits[a, vocab - 1] = nan + logits[b, first[b] // 2] = inf + logits[b, first[b]] = nan + logits[b, later[b]] = nan + logits[c, vocab - 5] = nan + logits[step % rows] = nan + elif kind == "signed": + lo = torch.randint(0, half, (rows,), generator=gen, device=dev) + hi = lo + torch.randint(1, half, (rows,), generator=gen, device=dev) + pattern = (r + step) % 4 + p0, p1, p2, p3 = (r[pattern == p] for p in range(4)) + logits[p0, lo[p0]] = inf + logits[p0, hi[p0]] = inf + logits[p1] = -inf + logits[p2] = -1.0 + logits[p2, lo[p2]] = -0.0 + logits[p2, hi[p2]] = 0.0 + logits[p3] = -1.0 + logits[p3, lo[p3]] = 0.0 + logits[p3, hi[p3]] = -0.0 + elif kind == "negzero": + pattern = (r + step) % 4 + p0, p1, p2 = (r[pattern == p] for p in range(3)) + logits.fill_(-1.0) + logits[p0, 12] = -0.0 + logits[p0, half + 5] = 0.0 + logits[p1, vocab // max(shards, 2) + 7] = 0.0 + logits[p1, vocab - 1] = -0.0 + logits[p2, vocab - 1] = -0.0 + elif kind != "plain": + raise ValueError(kind) + + +def drafts_for(gen, logits: torch.Tensor, batch: int, drafts: int, step: int) -> torch.Tensor: + """Drafts [B, K] (int32) whose accepted prefix under the target tokens of ``logits`` is (step + b) % (K + 1) for + request b: the drafts after the prefix are random, the first of them a mismatch.""" + vocab = logits.shape[1] + dev = logits.device + target = torch.argmax(logits.float(), dim=-1).view(batch, drafts + 1)[:, :drafts] + target = target.to(torch.int32) + draft = torch.randint(0, vocab, (batch, drafts), generator=gen, device=dev, dtype=torch.int32) + prefix = ((torch.arange(batch, device=dev) + step) % (drafts + 1))[:, None] + j = torch.arange(drafts, device=dev)[None, :] + draft = torch.where(j < prefix, target, draft) + return torch.where(j == prefix, (target + 1) % vocab, draft).contiguous() + + +def step_inputs(gen, batch: int, drafts: int, kind: str, step: int): + """A step's fp32 target logits [B (K + 1), V] of ``kind`` and its drafts.""" + logits = torch.randn(batch * (drafts + 1), V, generator=gen, device="cuda") + shape_logits(logits, gen, kind, step) + return logits, drafts_for(gen, logits, batch, drafts, step) + + +# ---------------------------------------------------------------------------------------------------------------- +# The reference and the kernel +# ---------------------------------------------------------------------------------------------------------------- + + +def reference(st: dict, logits, draft, force: float, block: int) -> list: + """The model's torch op sequence for the step, K = draft.shape[1]; updates the state tensors of ``st`` in place.""" + spec = _spec() + batch, drafts = draft.shape + dev = logits.device + # DFlashWorker._refresh_ctx_block_tables + encoded = st["block_off"][st["pool_idx"], :batch, 0].to(torch.int64).clone() + st["block_counts"][:batch].copy_((encoded >= 0).sum(dim=1)) + decoded = encoded.clamp_(min=0).div_(st["divisor"], rounding_mode="floor") + st["block_tables"][:batch].copy_(decoded.to(torch.int32)) + # SpecWorkerBase._sample_and_accept_draft_tokens_base (greedy: argmax) + accepted = torch.empty((batch, drafts + 1), dtype=torch.int, device=dev) + num_acc = torch.ones(batch, dtype=torch.int, device=dev) + gen_target = torch.argmax(logits, dim=-1).reshape(batch, drafts + 1) + accepted[:, : drafts + 1] = gen_target + num_acc[0:] += torch.cumprod((draft == gen_target[:, :drafts]).int(), dim=-1).sum(1) + # SpecWorkerBase._apply_force_accepted_tokens (no context requests) + if force != 0.0: + int_part = int(force) + frac = force - int_part + max_total = drafts + 1 + base_total = min(int_part + 1, max_total) + if frac > 0.0 and base_total < max_total: + st["rng_counter"] += 1 + slot_ids = torch.arange(batch, device=dev, dtype=torch.int64) + hashed = st["rng_counter"] * spec._FORCE_ACCEPT_RNG_COUNTER_STRIDE + hashed = hashed + slot_ids * spec._FORCE_ACCEPT_RNG_SLOT_STRIDE + indices = hashed & (spec._FORCE_ACCEPT_RNG_POOL_SIZE - 1) + extra = (st["rng_pool"][indices] < frac).to(num_acc.dtype) + num_acc[0:] = (base_total + extra).clamp_(max=max_total) + else: + num_acc[0:] = base_total + # MambaHybridCacheManagerV2.update_mamba_states -> _record_kda_replay_acceptance + nad = (num_acc[0:batch] - 1).to(torch.int32) + slots = st["state_idx"][0:batch].to(torch.int32).to(torch.long) + acc = nad.clamp(min=0) + current = st["prev_acc"][slots] + acc = torch.where(st["dummy"][0:batch], current, acc) + st["prev_acc"][slots] = acc + # DFlashWorker._prepare_kv_for_draft_forward + rewind = 1 - num_acc[0:batch] + st["kv_lens"][0:batch] += 1 + # DFlashWorker.prepare_1st_drafter_inputs up to the fc + bonus_idx = (num_acc - 1).clamp_min(0).long().unsqueeze(1) + bonus = accepted.gather(1, bonus_idx).squeeze(1).long() + ctx_len_gen = st["ctx_len"][st["batch_to_slot"][0:batch]] + j_block = torch.arange(block, dtype=torch.long, device=dev) + offsets = torch.arange(drafts + 1, dtype=torch.long, device=dev) + now = (ctx_len_gen + num_acc.long()).clamp_(max=MAX_CTX) + qpos = now.unsqueeze(1) + j_block.unsqueeze(0) + cpos = ctx_len_gen.unsqueeze(1) + offsets.unsqueeze(0) + noise = st["mask_row"].expand(batch, block, -1).clone() + noise[:, 0, :] = F.embedding(bonus, st["embed"]) + return [accepted, num_acc, rewind, bonus, qpos, cpos, noise] + + +def fused(st: dict, logits, draft, force: float, block: int, shard=None) -> list: + """``trtllm::k3_spec_accept`` on the state ``st``; ``shard``: (workspace, first column) when ``logits`` are this + rank's bf16 vocabulary shard.""" + ws_args = () + if shard is not None: + ws, first = shard + ws_args = (ws["uc"], ws["mc"], ws["flags"]) + ws_args += (ws["rank"], ws["slots"], ws["push_copies"], first) + return torch.ops.trtllm.k3_spec_accept( + logits, draft, st["block_off"], st["pool_idx"], st["divisor"], st["block_counts"], st["block_tables"], + st["prev_acc"], st["state_idx"], st["dummy"], st["kv_lens"], st["batch_to_slot"], st["ctx_len"], MAX_CTX, + st["embed"], st["mask_row"], st["rng_pool"], st["rng_counter"], force, block, *ws_args, + ) # fmt: skip + + +def state_of(st: dict) -> list: + return [st[k] for k in STATE_KEYS] + + +def same(xs, ys) -> list: + """Per pair: the same dtype, shape and bits.""" + return [ + x.dtype == y.dtype + and x.shape == y.shape + and torch.equal(x.contiguous().view(torch.uint8), y.contiguous().view(torch.uint8)) + for x, y in zip(xs, ys) + ] + + +def new_result(**fields) -> dict: + return dict(fields, checks={}, bad=[]) + + +def tally(res: dict, check: str, ok: bool, step: int, names=()) -> None: + """Counts one step of a check; the first failures are kept with the names of what differed.""" + passed, total = res["checks"].get(check, (0, 0)) + res["checks"][check] = (passed + bool(ok), total + 1) + if not ok and len(res["bad"]) < 4: + res["bad"].append((step, check, list(names))) + + +def record(res: dict, step: int, pairs: dict) -> None: + """Tallies one step's comparisons; ``pairs``: check -> (names, got, want).""" + for check, (names, got, want) in pairs.items(): + flags = same(got, want) + tally(res, check, all(flags), step, [n for n, f in zip(names, flags) if not f]) + + +def finish(res: dict) -> dict: + res["ok"] = all(passed == total for passed, total in res["checks"].values()) + return res + + +def cell(res: dict, check: str) -> str: + passed, total = res["checks"].get(check, (None, None)) + return "-" if total is None else f"{passed}/{total}" + + +# ---------------------------------------------------------------------------------------------------------------- +# Checks +# ---------------------------------------------------------------------------------------------------------------- + + +def measure(batch: int, drafts: int, block: int, case: int, steps: int = STEPS) -> dict: + """One case at one split: ``steps`` consecutive steps with the state carried, the kernel against the reference + (every output and state tensor), its inputs untouched and a rerun on a clone of its state.""" + _op() + name, kind, forced = CASES[case] + force = force_value(drafts, forced) + seed = 20260930 + 1000 * batch + 100 * drafts + 10 * block + case + gen = torch.Generator(device="cuda").manual_seed(seed) + st_ref = make_state(gen, batch) + st_new = clone_state(st_ref) + res = new_result(split=f"{batch}x{drafts + 1}", block=block, case=name, force=force) + for step in range(steps): + set_dummies((st_ref, st_new), batch, step) + logits, draft = step_inputs(gen, batch, drafts, kind, step) + logits_in, draft_in = logits.clone(), draft.clone() + st_again = clone_state(st_new) + want = reference(st_ref, logits, draft, force, block) + got = fused(st_new, logits, draft, force, block) + again = fused(st_again, logits, draft, force, block) + record(res, step, { + "outputs": (OUTPUTS, got, want), + "state": (STATE_KEYS, state_of(st_new), state_of(st_ref)), + "inputs": (("logits", "draft") + INPUT_KEYS, [logits, draft] + [st_new[k] for k in INPUT_KEYS], + [logits_in, draft_in] + [st_ref[k] for k in INPUT_KEYS]), + "rerun": (OUTPUTS + STATE_KEYS, again + state_of(st_again), got + state_of(st_new)), + }) # fmt: skip + return finish(res) + + +def measure_graph(batch: int, drafts: int, block: int, forced, replays: int = STEPS) -> dict: + """One call captured in a CUDA graph and replayed ``replays`` times with the logits (every kind), drafts, dummy + mask, block offsets and context lengths rewritten in place and the state carried: every replay against the + reference.""" + _op() + force = force_value(drafts, forced) + gen = torch.Generator(device="cuda").manual_seed(777 + 100 * batch + 10 * drafts + block) + st_ref = make_state(gen, batch) + st_new = clone_state(st_ref) + logits, draft = step_inputs(gen, batch, drafts, "plain", 0) + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(stream): + fused(clone_state(st_new), logits, draft, force, block) # compiled outside capture + torch.cuda.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + outs = fused(st_new, logits, draft, force, block) + torch.cuda.synchronize() + name = f"graph x{replays} ({'natural' if forced is None else 'forced ' + forced})" + res = new_result(split=f"{batch}x{drafts + 1}", block=block, case=name, force=force) + max_blocks = st_ref["block_off"].shape[3] + for rep in range(replays): + kind = GRAPH_KINDS[rep % len(GRAPH_KINDS)] + new_logits, new_draft = step_inputs(gen, batch, drafts, kind, rep) + logits.copy_(new_logits) + draft.copy_(new_draft) + set_dummies((st_ref, st_new), batch, rep) + offsets = block_offsets(gen, max_blocks, salt=rep) + ctx = context_lengths(gen, st_ref["batch_to_slot"], batch) + for st in (st_ref, st_new): + st["block_off"].copy_(offsets) + st["ctx_len"].copy_(ctx) + graph.replay() + want = reference(st_ref, logits, draft, force, block) + record(res, rep, { + "outputs": (OUTPUTS, outs, want), + "state": (STATE_KEYS, state_of(st_new), state_of(st_ref)), + "inputs": (INPUT_KEYS, [st_new[k] for k in INPUT_KEYS], [st_ref[k] for k in INPUT_KEYS]), + }) # fmt: skip + del graph + return finish(res) + + +def failures(results: list) -> list: + return [ + {k: r[k] for k in ("split", "block", "case", "checks", "bad")} + for r in results + if not r["ok"] + ] + + +@pytest.mark.parametrize("block_delta", [1, 0], ids=["block=K+1", "block=K"]) +@pytest.mark.parametrize("batch,tokens", SPLITS, ids=[f"{b}x{t}" for b, t in SPLITS]) +def test_split(batch, tokens, block_delta): + """Every case at one split and block width (DFlash's K + 1, DSpark's K), 20 steps each.""" + with torch.inference_mode(): + results = [ + measure(batch, tokens - 1, tokens - 1 + block_delta, c) for c in range(len(CASES)) + ] + assert not failures(results), failures(results) + + +@pytest.mark.parametrize("batch,tokens", WIDE_SPLITS, ids=[f"{b}x{t}" for b, t in WIDE_SPLITS]) +def test_wide_block(batch, tokens): + """Block 8 at K = 1 and 3: DFlash's block under max_draft_len 7 when a draft-length schedule runs fewer drafts.""" + with torch.inference_mode(): + results = [measure(batch, tokens - 1, 8, c) for c in range(len(CASES))] + assert not failures(results), failures(results) + + +@pytest.mark.parametrize("forced", [None, "frac"], ids=["natural", "forced-frac"]) +@pytest.mark.parametrize("block_delta", [1, 0], ids=["block=K+1", "block=K"]) +@pytest.mark.parametrize("batch,tokens", GRAPH_SPLITS, ids=[f"{b}x{t}" for b, t in GRAPH_SPLITS]) +def test_graph_replay(batch, tokens, block_delta, forced): + """A captured call replayed 20 times with its inputs rewritten in place, the state carried.""" + with torch.inference_mode(): + res = measure_graph(batch, tokens - 1, tokens - 1 + block_delta, forced) + assert res["ok"], failures([res]) + + +def test_force_mode(): + """force_mode follows _apply_force_accepted_tokens: (mode, base total, fraction) of a forced value at K drafts.""" + op = _op() + kern = op._kernel_module() + for drafts, (integer, fractional) in FORCES.items(): + assert op.force_mode(0.0, drafts) == (kern.FORCE_OFF, 0, 0.0) + want = (kern.FORCE_INT, min(int(integer) + 1, drafts + 1), 0.0) + assert op.force_mode(integer, drafts) == want, drafts + mode, total, frac = op.force_mode(fractional, drafts) + assert (mode, total) == (kern.FORCE_FRAC, int(fractional) + 1) and 0.0 < frac < 1.0, drafts + # Above K the count clamps to K + 1 and the fraction is ignored. + want = (kern.FORCE_INT, drafts + 1, 0.0) + assert op.force_mode(force_value(drafts, "over"), drafts) == want, drafts + + +def test_supports(): + """Every split the engine runs is supported (block K .. 8, the whole vocabulary or a TP shard); the limits are + rejected, and unsupported calls raise before launching.""" + op = _op() + for batch, tokens in SPLITS: + drafts = tokens - 1 + for block in (drafts, tokens, 8): + assert op.supports(V, batch, block, drafts, HIDDEN), (batch, tokens, block) + for ranks in (2, 4, 8, 16): + assert op.supports(V // ranks, batch, tokens, drafts, HIDDEN, ranks), (batch, ranks) + assert op.supports(TP16_COLUMNS, batch, tokens, drafts, HIDDEN, 16), (batch, tokens) + assert not op.supports(V, 9, 8, 7, HIDDEN) # 9 requests + assert not op.supports(V, 5, 16, 15, HIDDEN) # 80 rows + assert not op.supports(V, 1, 17, 7, HIDDEN) # block 17 + assert not op.supports(V, 1, 0, 7, HIDDEN) # no block + assert not op.supports(V, 1, 8, 16, HIDDEN) # K + 1 = 17 + assert not op.supports(V + CTA_COLS // 2, 1, 8, 7, HIDDEN) # a partial CTA of columns + assert not op.supports(V, 1, 8, 7, HIDDEN + 4) # hidden % 8 + assert not op.supports(TP16_COLUMNS, 1, 8, 7, HIDDEN, 3) # an odd slot count + assert not op.supports(TP16_COLUMNS, 1, 8, 7, HIDDEN, 32) # more than 16 slots + with torch.inference_mode(): + gen = torch.Generator(device="cuda").manual_seed(3) + st = make_state(gen, 1) + logits, draft = step_inputs(gen, 1, 7, "plain", 0) + # 7 rows for 1 x 8, bf16 logits without a workspace, int64 drafts. + bad_calls = ((logits[:7], draft), (logits.bfloat16(), draft), (logits, draft.long())) + for bad_logits, bad_draft in bad_calls: + with pytest.raises(ValueError): + fused(st, bad_logits, bad_draft, 0.0, 8) + with pytest.raises(ValueError): # a state tensor of another dtype + fused(dict(st, kv_lens=st["kv_lens"].long()), logits, draft, 0.0, 8) + with pytest.raises(ValueError): # bf16 shards without a workspace + torch.ops.trtllm.k3_spec_accept( + logits.bfloat16(), draft, st["block_off"], POOL_IDX, DIVISOR, st["block_counts"], st["block_tables"], + st["prev_acc"], st["state_idx"], st["dummy"], st["kv_lens"], st["batch_to_slot"], st["ctx_len"], + MAX_CTX, st["embed"], st["mask_row"], st["rng_pool"], st["rng_counter"], 0.0, 8, None, None, None, 0, + 2, 1, 0, + ) # fmt: skip + + +# ---------------------------------------------------------------------------------------------------------------- +# Report (python3 test_k3_spec_accept.py report) and timing (python3 test_k3_spec_accept.py time) +# ---------------------------------------------------------------------------------------------------------------- + + +def report() -> int: + """Every check as a markdown table.""" + _op() + print(f"{torch.cuda.get_device_name()}; V {V}, hidden {HIDDEN}; {STEPS} consecutive steps per case, " + "state carried") # fmt: skip + print("| split | block | case | force | outputs identical | state identical | inputs untouched | rerun identical " + "| result |") # fmt: skip + print("| :-- | --: | :-- | --: | --: | --: | --: | --: | :-- |") + ok_all = True + runs = [(b, t, t - 1 + d, c) for b, t in SPLITS for d in (1, 0) for c in range(len(CASES))] + runs += [(b, t, 8, c) for b, t in WIDE_SPLITS for c in range(len(CASES))] + with torch.inference_mode(): + results = (measure(b, t - 1, block, c) for b, t, block, c in runs) + graphs = (measure_graph(b, t - 1, t - 1 + d, forced) for b, t in GRAPH_SPLITS for d in (1, 0) + for forced in (None, "frac")) # fmt: skip + for res in itertools.chain(results, graphs): + ok_all &= res["ok"] + print(f"| {res['split']} | {res['block']} | {res['case']} | {res['force']:g} | {cell(res, 'outputs')} | " + f"{cell(res, 'state')} | {cell(res, 'inputs')} | {cell(res, 'rerun')} | " + f"{'PASS' if res['ok'] else 'FAIL ' + str(res['bad'])} |", flush=True) # fmt: skip + print("ALL PASS" if ok_all else "FAIL") + return 0 if ok_all else 1 + + +def capture(body, calls: int) -> torch.cuda.CUDAGraph: + """A CUDA graph of ``calls`` back-to-back bodies (the first body, which compiles, runs outside capture).""" + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(stream): + body() + torch.cuda.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + for _ in range(calls): + body() + torch.cuda.synchronize() + for _ in range(3): + graph.replay() + torch.cuda.synchronize() + return graph + + +def replay_us(graph: torch.cuda.CUDAGraph, calls: int) -> float: + start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + start.record() + graph.replay() + end.record() + torch.cuda.synchronize() + return start.elapsed_time(end) * 1e3 / calls + + +def timing(calls: int = 20, rounds: int = 12) -> None: + """us per call of the reference ops and of the kernel (natural acceptance, block K + 1), batch 1 first.""" + _op() + print(f"{torch.cuda.get_device_name()}; V {V}, hidden {HIDDEN}; CUDA graphs of {calls} back-to-back calls, " + f"median over {rounds} replays in alternating order: us per call") # fmt: skip + print("| split | block | reference (torch ops) | k3_spec_accept | speedup |") + print("| :-- | --: | --: | --: | --: |") + with torch.inference_mode(): + for batch, tokens in SPLITS: + drafts = tokens - 1 + gen = torch.Generator(device="cuda").manual_seed(31 + 10 * batch + tokens) + st_ref = make_state(gen, batch) + st_new = clone_state(st_ref) + logits, draft = step_inputs(gen, batch, drafts, "plain", 3) + graphs = [ + capture(lambda: reference(st_ref, logits, draft, 0.0, tokens), calls), + capture(lambda: fused(st_new, logits, draft, 0.0, tokens), calls), + ] + per_call = [[], []] + for rd in range(rounds): + for arm in (0, 1) if rd % 2 == 0 else (1, 0): + per_call[arm].append(replay_us(graphs[arm], calls)) + ref_us, kernel_us = (statistics.median(t) for t in per_call) + print(f"| {batch}x{tokens} | {tokens} | {ref_us:.2f} | {kernel_us:.2f} | {ref_us / kernel_us:.1f}x |", + flush=True) # fmt: skip + del graphs + + +if __name__ == "__main__": + if len(sys.argv) > 1 and sys.argv[1] == "time": + timing() + elif len(sys.argv) > 1 and sys.argv[1] == "report": + sys.exit(report()) + else: + sys.exit(pytest.main([__file__, "-q", "-p", "no:cacheprovider", *sys.argv[1:]])) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_spec_accept_sharded.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_spec_accept_sharded.py new file mode 100644 index 000000000000..16ea213c4c56 --- /dev/null +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_spec_accept_sharded.py @@ -0,0 +1,568 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``trtllm::k3_spec_accept`` with vocabulary-sharded bf16 target logits (the ranks exchange their row maxima over the +multicast Lamport buffer of ``op.workspace(mapping, push_copies)``) against the same kernel on the gathered fp32 logits, +on W GPUs of one NVLink domain (MPI, one rank per GPU, the W ranks one TP group). + +The gathered logits are the bf16 values in fp32 (what the all-gather and ``.float()`` give), so both must agree bit for +bit: every output and every in-place state tensor, on every rank; the gathered kernel is also checked against the torch +op sequence it replaces (``test_k3_spec_accept.reference``). Every rank draws the same logits (checked) and passes its +shard. + +Every split B = 1 .. 8 requests x K + 1 in {2, 4, 8} tokens (block K + 1), 6 consecutive calls per case with the state +carried (the exchange's three buffers rotate twice): normal (accepted prefixes 0 .. K mixed over the requests), ties +(across CTAs and shard boundaries), rank_ties (the same maximum in every rank's shard: rank 0's index wins), nan (a +whole row, NaNs in two shards, a NaN in the last shard only: the first NaN wins), negzero (+-0 maxima: -0.0 travels as ++0.0) and forced (integer and fractional values alternating). Then sharded calls of 64, 2, 2 and 64 rows with every +rank but rank 0 launching the last one late (the last call reads the buffer of the first after a 2-row call re-armed +it), a batch that dips and grows back (64, 64, 64, 16, 56 and 64 rows, then 50 calls of a random B x 8 with one random +rank late each call), and CUDA graphs of 3 sharded calls replayed 10 times with the logits, drafts and dummy mask +rewritten in place (the buffer rotation inside a graph). + +Shapes: the rank's shard V / W of V = 163840 (W exchange slots), and TP16's 10240-column shard with every rank filling +16 / W slots (the exchange of 16 ranks; ``--copies`` overrides). + + srun -n W --mpi=pmix python3 test_k3_spec_accept_sharded.py [--copies C] [--time [--base]] + srun -n W --mpi=pmix python3 -m pytest test_k3_spec_accept_sharded.py + +``--time`` also reports us per call at batch 1 and every split of the sharded kernel and of the kernel on the gathered +fp32 logits (CUDA graphs of 20 back-to-back calls, the slowest rank of each replay, median over 12 replays in +alternating order; the all-gather, cat and cast that the sharded path removes are not in the second number); +``--base`` adds the installed base package's kernel (``$K3_BASE_TRTLLM``) on the same sharded calls and workspace. +Without an MPI job of 2 or more ranks the module skips. +""" + +import argparse +import contextlib +import importlib.util +import os +import random +import statistics +import sys +import time +import traceback + +import pytest +import torch + +__extra_import_path__ = ["."] +from test_k3_spec_accept import ( # noqa: E402 (the single-GPU test's state, inputs, reference and tallies) + FORCES, + OUTPUTS, + SPLITS, + STATE_KEYS, + TP16_COLUMNS, + V, + _op, + _sm100, + cell, + clone_state, + drafts_for, + embedding, + finish, + fused, + make_state, + new_result, + record, + reference, + set_dummies, + shape_logits, + state_of, + tally, +) + +STEPS = 6 +# (name, logits kind) +CASES = ( + ("normal", "plain"), + ("ties", "ties"), + ("rank_ties", "rank_ties"), + ("nan", "nan"), + ("negzero", "negzero"), + ("forced", "plain"), +) +GRAPH_SPLITS = [(1, 2), (2, 8), (4, 4), (8, 8)] +GRAPH_KINDS = ("plain", "ties", "rank_ties", "nan", "negzero") + + +def launcher_world_size() -> int: + """The world size an MPI launcher (mpirun, MPICH, srun) gave this process, from its environment: whether to skip + is decided without initializing MPI or CUDA.""" + for name in ("OMPI_COMM_WORLD_SIZE", "PMI_SIZE", "SLURM_STEP_NUM_TASKS"): + value = os.environ.get(name, "") + if value.isdigit(): + return int(value) + return 1 + + +pytestmark = pytest.mark.skipif( + launcher_world_size() < 2 or not _sm100(), + reason="needs an MPI job of 2 or more ranks on SM100 GPUs (srun -n W --mpi=pmix python3 -m pytest ...)", +) + +_env = {} + + +def mpi_env() -> dict: + """The job's communicator, rank and world, the op and a TP mapping of every rank, set up once (the device: the + rank's index on its node).""" + if not _env: + from mpi4py import MPI + + comm = MPI.COMM_WORLD + rank, world = comm.Get_rank(), comm.Get_size() + torch.cuda.set_device(rank % torch.cuda.device_count()) + op = _op() + from tensorrt_llm.mapping import Mapping + + mapping = Mapping( + world_size=world, rank=rank, gpus_per_node=torch.cuda.device_count(), tp_size=world + ) + _env.update(comm=comm, rank=rank, world=world, op=op, mapping=mapping) + return _env + + +def on_every_rank(env: dict, fn, *args): + """``fn(*args)``, run by every rank; an exception on one rank aborts the job, whose other ranks would otherwise + wait forever in the exchange or in the next collective.""" + try: + return fn(*args) + # Whatever failed: report it, then take the whole job down rather than leave it hanging. + except Exception: + traceback.print_exc() + sys.stderr.flush() + env["comm"].Abort(1) + raise + + +def configs(world: int, copies=None) -> list: + """(name, total columns, push copies): every rank's shard of V, and TP16's shard with 16 / W slots per rank.""" + return [("V / W", V, 1), ("TP16 shard", TP16_COLUMNS * world, copies or max(1, 16 // world))] + + +def shard_of(env: dict, vocab: int): + """(columns per rank, this rank's first column).""" + shard = vocab // env["world"] + return shard, env["rank"] * shard + + +def same_logits_everywhere(env: dict, x: torch.Tensor) -> bool: + """Whether every rank drew the same logits (the shards are slices of one tensor only then): row checksums.""" + sums = x.view(torch.int16).sum(dim=1, dtype=torch.int64).tolist() + return all(other == sums for other in env["comm"].allgather(sums)) + + +def sharded_state(gen, batch: int, vocab: int) -> dict: + st = make_state(gen, batch) + st["embed"] = embedding()[:vocab] # the drafter embedding covers the vocabulary + return st + + +def base_kernel(): + """The unmodified kernel module of the installed base package (``$K3_BASE_TRTLLM``), for before/after timing.""" + path = os.path.join(os.environ["K3_BASE_TRTLLM"], "_torch", "cute_dsl_kernels", "k3_spec_accept", + "k3_spec_accept_kernel.py") # fmt: skip + spec = importlib.util.spec_from_file_location("k3_spec_accept_kernel_base", path) + module = importlib.util.module_from_spec(spec) + sys.modules[spec.name] = module + spec.loader.exec_module(module) + return module + + +def rearm(env: dict, ws: dict) -> None: + """Every rank's Lamport words empty again, with no call in flight on any rank (between the base and the new + kernel on one workspace: the base kernel re-arms only its own call's rows).""" + torch.cuda.synchronize() + env["comm"].Barrier() + ws["uc"].fill_(env["op"]._kernel_module().EMPTY_WORD) + torch.cuda.synchronize() + env["comm"].Barrier() + + +@contextlib.contextmanager +def kernel_module(op, module, compiled: dict): + """The op's calls inside run ``module``'s kernel, compiled into ``compiled``.""" + saved = op._modules.get("kernel"), op._compiled + op._modules["kernel"], op._compiled = module, compiled + try: + yield + finally: + op._modules["kernel"], op._compiled = saved + + +# ---------------------------------------------------------------------------------------------------------------- +# Checks (every rank runs the same calls in the same order: the exchange and the collectives need it) +# ---------------------------------------------------------------------------------------------------------------- + + +def check_cases(env: dict, ws: dict, vocab: int, config: str, steps: int = STEPS) -> list: + """Every split x case: ``steps`` calls with the state carried; the sharded kernel against the gathered one, and the + gathered one against the torch reference.""" + shard, lo = shard_of(env, vocab) + results = [] + for batch, tokens in SPLITS: + drafts = tokens - 1 + for index, (name, kind) in enumerate(CASES): + # One seed on every rank: the same logits, drafts and state everywhere. + seed = 4000 + 100 * batch + 10 * tokens + index + gen = torch.Generator(device="cuda").manual_seed(seed) + st_sh = sharded_state(gen, batch, vocab) + st_g, st_t = clone_state(st_sh), clone_state(st_sh) + res = new_result(config=config, split=f"{batch}x{tokens}", case=name) + for step in range(steps): + x = torch.randn(batch * tokens, vocab, generator=gen, device="cuda").bfloat16() + shape_logits(x, gen, kind, step, env["world"]) + draft = drafts_for(gen, x, batch, drafts, step) + force = FORCES[drafts][step % 2] if name == "forced" else 0.0 + set_dummies((st_sh, st_g, st_t), batch, step) + tally(res, "logits", same_logits_everywhere(env, x), step) + x_shard, x32 = x[:, lo : lo + shard].contiguous(), x.float() + got = fused(st_sh, x_shard, draft, force, tokens, (ws, lo)) + ref = fused(st_g, x32, draft, force, tokens) + want = reference(st_t, x32, draft, force, tokens) + record(res, step, { + "outputs": (OUTPUTS, got, ref), + "state": (STATE_KEYS, state_of(st_sh), state_of(st_g)), + "torch": (OUTPUTS + STATE_KEYS, ref + state_of(st_g), want + state_of(st_t)), + }) # fmt: skip + results.append(finish(res)) + return results + + +def check_transition(env: dict, ws: dict, vocab: int, config: str) -> dict: + """Sharded calls of 64, 2, 2 and 64 rows. The 4th call uses the buffer the 1st filled (row maxima of 50.0), which + the 2nd re-armed for its own 2 rows; every rank but rank 0 launches the 4th call 50 ms late, so rank 0 must wait for + its peers' pushes in every row rather than take what the 1st call left.""" + comm = env["comm"] + shard, lo = shard_of(env, vocab) + gen = torch.Generator(device="cuda").manual_seed(4999) + states = {} + for batch in (8, 1): + st = sharded_state(gen, batch, vocab) + states[batch] = (st, clone_state(st)) + res = new_result(config=config, split="8x8, 1x2", case="rows 64, 2, 2, 64 (peers late)") + for call, (batch, tokens) in enumerate(((8, 8), (1, 2), (1, 2), (8, 8))): + rows = batch * tokens + x = torch.randn(rows, vocab, generator=gen, device="cuda").bfloat16() + if call == 0: + r = torch.arange(rows, device="cuda") + x[r, vocab - 1 - r] = 50.0 + draft = drafts_for(gen, x, batch, tokens - 1, call) + tally(res, "logits", same_logits_everywhere(env, x), call) + st_sh, st_g = states[batch] + torch.cuda.synchronize() + comm.Barrier() + if call == 3 and env["rank"] > 0: + time.sleep(0.05) + got = fused(st_sh, x[:, lo : lo + shard].contiguous(), draft, 0.0, tokens, (ws, lo)) + ref = fused(st_g, x.float(), draft, 0.0, tokens) + record(res, call, { + "outputs": (OUTPUTS, got, ref), + "state": (STATE_KEYS, state_of(st_sh), state_of(st_g)), + }) # fmt: skip + torch.cuda.synchronize() + comm.Barrier() + return finish(res) + + +def check_dip_regrow(env: dict, ws: dict, vocab: int, config: str, calls: int = 50) -> dict: + """A batch that dips and grows back: sharded calls of 8, 8, 8, 2, 7 and 8 requests x 8 tokens, then ``calls`` + calls of a random B x 8, one random rank launching each of those 20 ms late (the same draws on every rank). Every + row of call c peaks at 100 - c in a random column, so a word left from an earlier call outranks the fresh ones.""" + comm = env["comm"] + shard, lo = shard_of(env, vocab) + gen = torch.Generator(device="cuda").manual_seed(5999) + draws = random.Random(5999) + tokens = 8 + states = {} + for batch in range(1, 9): + st = sharded_state(gen, batch, vocab) + states[batch] = (st, clone_state(st)) + batches = [8, 8, 8, 2, 7, 8] + [draws.randint(1, 8) for _ in range(calls)] + late = [None] * 6 + [draws.randrange(env["world"]) for _ in range(calls)] + res = new_result(config=config, split="B x 8, B varying", + case=f"rows 64, 64, 64, 16, 56, 64 + {calls} random (one rank late)") # fmt: skip + for call, batch in enumerate(batches): + rows = batch * tokens + x = torch.randn(rows, vocab, generator=gen, device="cuda").bfloat16() + peak = torch.randint(0, vocab, (rows,), generator=gen, device="cuda") + x[torch.arange(rows, device="cuda"), peak] = 100.0 - call + draft = drafts_for(gen, x, batch, tokens - 1, call) + tally(res, "logits", same_logits_everywhere(env, x), call) + st_sh, st_g = states[batch] + torch.cuda.synchronize() + comm.Barrier() + if late[call] == env["rank"]: + time.sleep(0.02) + got = fused(st_sh, x[:, lo : lo + shard].contiguous(), draft, 0.0, tokens, (ws, lo)) + ref = fused(st_g, x.float(), draft, 0.0, tokens) + record(res, call, { + "outputs": (OUTPUTS, got, ref), + "state": (STATE_KEYS, state_of(st_sh), state_of(st_g)), + }) # fmt: skip + torch.cuda.synchronize() + comm.Barrier() + return finish(res) + + +def check_graph(env: dict, ws: dict, vocab: int, config: str, batch: int, tokens: int, calls: int = 3, + replays: int = 10) -> dict: # fmt: skip + """``calls`` sharded calls captured in one CUDA graph, replayed ``replays`` times with the logits (every kind), + drafts and dummy mask rewritten in place: every call of every replay against the gathered kernel.""" + comm = env["comm"] + shard, lo = shard_of(env, vocab) + drafts, rows = tokens - 1, batch * tokens + gen = torch.Generator(device="cuda").manual_seed(5000 + 10 * batch + tokens) + st_sh = sharded_state(gen, batch, vocab) + st_g = clone_state(st_sh) + x_buf = torch.zeros(rows, shard, dtype=torch.bfloat16, device="cuda") + d_buf = torch.zeros(batch, drafts, dtype=torch.int32, device="cuda") + # Compiled outside capture (an exchange: every rank makes this call). + fused(clone_state(st_sh), x_buf, d_buf, 0.0, tokens, (ws, lo)) + torch.cuda.synchronize() + comm.Barrier() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + outs = [fused(st_sh, x_buf, d_buf, 0.0, tokens, (ws, lo)) for _ in range(calls)] + torch.cuda.synchronize() + res = new_result( + config=config, split=f"{batch}x{tokens}", case=f"graph of {calls} calls x {replays}" + ) + for rep in range(replays): + x = torch.randn(rows, vocab, generator=gen, device="cuda").bfloat16() + shape_logits(x, gen, GRAPH_KINDS[rep % len(GRAPH_KINDS)], rep, env["world"]) + draft = drafts_for(gen, x, batch, drafts, rep) + set_dummies((st_sh, st_g), batch, rep) + tally(res, "logits", same_logits_everywhere(env, x), rep) + x_buf.copy_(x[:, lo : lo + shard]) + d_buf.copy_(draft) + graph.replay() + x32 = x.float() + for c in range(calls): + record(res, rep, {"outputs": (OUTPUTS, outs[c], fused(st_g, x32, draft, 0.0, tokens))}) + record(res, rep, {"state": (STATE_KEYS, state_of(st_sh), state_of(st_g))}) + torch.cuda.synchronize() + comm.Barrier() + del graph + return finish(res) + + +def time_split(env: dict, ws: dict, vocab: int, config: str, batch: int, tokens: int, calls: int = 20, + rounds: int = 12, base=None) -> dict: # fmt: skip + """us per call at one split of the sharded kernel and of the kernel on the gathered fp32 logits (and of the + ``base`` kernel module on the sharded logits, twice: the second graph is the noise control): CUDA graphs of + ``calls`` back-to-back calls, the slowest rank of each replay, median over ``rounds`` replays in alternating + order.""" + comm = env["comm"] + shard, lo = shard_of(env, vocab) + gen = torch.Generator(device="cuda").manual_seed(6000 + 10 * batch + tokens) + st_sh = sharded_state(gen, batch, vocab) + # The CUDA-graph padding requests (one KDA slot) are dummies, as in the model: no two lanes write one slot. + set_dummies((st_sh,), batch, 0) + st_g = clone_state(st_sh) + x = torch.randn(batch * tokens, vocab, generator=gen, device="cuda").bfloat16() + x_shard, x32 = x[:, lo : lo + shard].contiguous(), x.float() + draft = drafts_for(gen, x, batch, tokens - 1, 3) + arms = [ + lambda: fused(st_sh, x_shard, draft, 0.0, tokens, (ws, lo)), + lambda: fused(st_g, x32, draft, 0.0, tokens), + ] + identical = None + if base is not None: + rearm(env, ws) + compiled = env.setdefault("base_compiled", {}) + st_b, st_b2 = clone_state(st_sh), clone_state(st_sh) + + def base_arm(st=st_b): + with kernel_module(env["op"], base, compiled): + return fused(st, x_shard, draft, 0.0, tokens, (ws, lo)) + + # Identity: one call of each kernel on the same state and logits, every output and state tensor. + st_new, st_old = clone_state(st_sh), clone_state(st_sh) + new = fused(st_new, x_shard, draft, 0.0, tokens, (ws, lo)) + with kernel_module(env["op"], base, compiled): + old = fused(st_old, x_shard, draft, 0.0, tokens, (ws, lo)) + identical = all( + torch.equal(a, b) for a, b in zip(new + state_of(st_new), old + state_of(st_old)) + ) + identical = all(comm.allgather(identical)) + arms += [base_arm, lambda: base_arm(st_b2)] + graphs = [] + for body in arms: + body() + torch.cuda.synchronize() + comm.Barrier() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + for _ in range(calls): + body() + for _ in range(3): + graph.replay() + torch.cuda.synchronize() + graphs.append(graph) + per_call = [[] for _ in arms] + for rd in range(rounds): + for arm in range(len(arms)) if rd % 2 == 0 else reversed(range(len(arms))): + comm.Barrier() + start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + start.record() + graphs[arm].replay() + end.record() + torch.cuda.synchronize() + per_call[arm].append(max(comm.allgather(start.elapsed_time(end) * 1e3 / calls))) + comm.Barrier() + return dict(config=config, split=f"{batch}x{tokens}", sharded=statistics.median(per_call[0]), + gathered=statistics.median(per_call[1]), + base=statistics.median(per_call[2]) if base is not None else None, + base2=statistics.median(per_call[3]) if base is not None else None, identical=identical) # fmt: skip + + +def run_config(env: dict, name: str, vocab: int, copies: int, time_it: bool = False, base=None, rounds: int = 12, + time_splits=None, checks: bool = True): # fmt: skip + """Every check of one shape (and its timing with ``time_it``): (skip reason or None, results, timings).""" + world = env["world"] + shard, slots = vocab // world, world * copies + config = f"{name}: {shard} columns x {world} ranks, {slots} slots" + if vocab % world or not env["op"].supports_columns(shard, slots): + reason = f"the kernel does not split {vocab} columns over {world} ranks of {copies} slots" + return reason, [], [] + try: + ws = env["op"].workspace(env["mapping"], copies) + # Raised on every rank: the ranks agree on the allocation's outcome. + except RuntimeError as exc: + return f"no multicast workspace ({exc})", [], [] + results = [] + if checks: + results = check_cases(env, ws, vocab, config) + results.append(check_transition(env, ws, vocab, config)) + results.append(check_dip_regrow(env, ws, vocab, config)) + results += [check_graph(env, ws, vocab, config, b, t) for b, t in GRAPH_SPLITS] + timings = [] + if time_it: + for b, t in time_splits or SPLITS: + timings.append(time_split(env, ws, vocab, config, b, t, rounds=rounds, base=base)) + return None, results, timings + + +@pytest.mark.parametrize("config", [0, 1], ids=["V-over-W", "TP16-shard"]) +def test_sharded(config): + env = mpi_env() + if env["world"] < 2: + pytest.skip("needs an MPI job of 2 or more ranks") + name, vocab, copies = configs(env["world"])[config] + with torch.inference_mode(): + skip, results, _ = on_every_rank(env, run_config, env, name, vocab, copies) + if skip: + pytest.skip(skip) + bad = [ + {k: r[k] for k in ("config", "split", "case", "checks", "bad")} + for r in results + if not r["ok"] + ] + oks = env["comm"].allgather(not bad) + assert all(oks), (f"ranks ok: {oks}", bad) + + +# ---------------------------------------------------------------------------------------------------------------- +# srun -n W --mpi=pmix python3 test_k3_spec_accept_sharded.py [--copies C] [--time] +# ---------------------------------------------------------------------------------------------------------------- + + +def print_report(per_rank: list, timings: list, skips: list, world: int) -> None: + """The checks as a markdown table (every count the fewest over the ranks), then the timing.""" + print(f"{torch.cuda.get_device_name()} x {world} ranks; {STEPS} consecutive calls per case, state carried; " + "counts: the fewest over the ranks") # fmt: skip + print("| config | split | case | outputs identical (sharded = gathered) | state identical | gathered = torch " + "| same logits on every rank | result |") # fmt: skip + print("| :-- | :-- | :-- | --: | --: | --: | --: | :-- |") + for rows in zip(*per_rank): + checks = { + k: (min(r["checks"][k][0] for r in rows), total) + for k, (_, total) in rows[0]["checks"].items() + } + res = dict(rows[0], checks=checks) + failed = [(rank, r["bad"]) for rank, r in enumerate(rows) if not r["ok"]] + print(f"| {res['config']} | {res['split']} | {res['case']} | {cell(res, 'outputs')} | {cell(res, 'state')} | " + f"{cell(res, 'torch')} | {cell(res, 'logits')} | {'FAIL ' + str(failed[:2]) if failed else 'PASS'} |", + flush=True) # fmt: skip + for skip in skips: + print(f"skipped: {skip}") + if timings: + print( + "\nus per call: the slowest rank, median over the replays of a CUDA graph of 20 calls (alternating order)" + ) + with_base = timings[0]["base"] is not None + print("| config | split | sharded (bf16 shard + exchange) | gathered fp32 logits |" + + (" sharded, base kernel | sharded - base | base2 - base | outputs = base |" + if with_base else "")) # fmt: skip + print("| :-- | :-- | --: | --: |" + (" --: | --: | --: | :-- |" if with_base else "")) + for t in timings: + extra = (f" {t['base']:.2f} | {t['sharded'] - t['base']:+.2f} | {t['base2'] - t['base']:+.2f} | " + f"{t['identical']} |" if with_base else "") # fmt: skip + print( + f"| {t['config']} | {t['split']} | {t['sharded']:.2f} | {t['gathered']:.2f} |{extra}", + flush=True, + ) + + +def main() -> int: + parser = argparse.ArgumentParser( + description="trtllm::k3_spec_accept: vocabulary-sharded vs gathered logits" + ) + parser.add_argument("--copies", type=int, default=None, + help="exchange slots every rank fills in the TP16 shape (default 16 / W)") # fmt: skip + parser.add_argument("--time", action="store_true", + help="also time the sharded kernel and the kernel on the gathered fp32 logits") # fmt: skip + parser.add_argument("--base", action="store_true", + help="with --time: also the base package's kernel ($K3_BASE_TRTLLM), sharded") # fmt: skip + parser.add_argument("--rounds", type=int, default=12, help="with --time: replays per arm") + parser.add_argument( + "--splits", default=None, help="with --time: only these splits, e.g. 1x2,1x8" + ) + parser.add_argument("--time-only", action="store_true", help="the timing without the checks") + args = parser.parse_args() + time_splits = None + if args.splits: + time_splits = [tuple(int(v) for v in sp.split("x")) for sp in args.splits.split(",")] + env = mpi_env() + comm, world = env["comm"], env["world"] + if world < 2: + print( + "test_k3_spec_accept_sharded: needs an MPI job of 2 or more ranks (srun -n W --mpi=pmix ...)" + ) + return 1 + results, timings, skips = [], [], [] + with torch.inference_mode(): + for name, vocab, copies in configs(world, args.copies): + skip, res, tim = on_every_rank(env, run_config, env, name, vocab, copies, args.time or args.time_only, + base_kernel() if args.base else None, args.rounds, time_splits, + not args.time_only) # fmt: skip + if skip: + skips.append(f"{name}: {skip}") + results += res + timings += tim + per_rank = comm.gather(results, root=0) + oks = comm.allgather(all(r["ok"] for r in results)) + passed = all(oks) and not skips + if env["rank"] == 0: + print_report(per_rank, timings, skips, world) + print( + "ALL PASS" if passed else f"FAIL (ranks ok: {oks}; skipped: {len(skips)})", flush=True + ) + return 0 if passed else 1 + + +if __name__ == "__main__": + sys.exit(main()) From aaef4b714b3f9ee4d6df642d4ce166cf7452c6ef Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 10:30:11 -0700 Subject: [PATCH 127/161] [None][doc] k3_markov: state the co-residency assumption The op's docstring now says that every CTA of the grid must be resident at once: no concurrent kernel holding SMs while it waits on the grid, and no SM cap below the grid. Signed-off-by: Vasanth Sabavat --- tensorrt_llm/_torch/cute_dsl_kernels/k3_markov/op.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_markov/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_markov/op.py index cf7cb969f923..139da96c479b 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_markov/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_markov/op.py @@ -22,7 +22,9 @@ The ranks exchange one (value, index) entry per CTA and position through this module's own MNNVL multicast buffers (``markov_workspace``), allocated collectively on the first call of a TP group, which must happen outside CUDA-graph -capture (the kernel also compiles there). +capture (the kernel also compiles there). Every CTA spins on the other ranks' entries, so all of the grid's CTAs (at +most the SM count) must be resident at once: no concurrent kernel may hold SMs while it waits on this grid, and no SM +cap (MPS or green contexts) may sit below the grid. """ from __future__ import annotations From a1d336dbb35ee8a91590176b16e202084ffd75a5 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 11:03:03 -0700 Subject: [PATCH 128/161] [None][feat] DFlash / DSpark worker: Kimi K3 decode kernels behind k3_decode The worker runs the k3_spec_accept, k3_ctx_kv and k3_markov kernels on the steps they take when a Kimi K3 target sets its k3_decode attribute. The attribute is False by default, and then every step takes the existing path. - trtllm::k3_spec_accept: a decode step (no context requests, at most 8 generation requests) with greedy strict acceptance (no rejection sampling, penalties or guided decoding), the V2 hybrid manager's KDA replay record, the draft pool's block table and an unsharded bf16 draft embedding runs the acceptance, the block-table decode, the replay record, kv_lens + 1 and the drafter's inputs (bonus tokens, positions, noise embedding) in one launch (_k3_accept_applies, _k3_accept). - target_logits keeps such a step's target logits as this rank's bf16 vocabulary shard: plain TP whose all-reduces own an MNNVL workspace, and an unquantized, bias-free, unpadded column-parallel bf16 head. The shard comes from the logits processor's own head kernel where it has one that takes the rows (lm_head_shard), else from the head's GEMM. The kernel exchanges the row maxima; the step returns "logits_vocab_shard", so the engine gathers the logits for any logits post-processor. - trtllm::k3_ctx_kv writes the drafter's context K / V of such a step into the manager-bound paged pool (fused K / V weight without bias or context norm, k_norm, NeoX RoPE from flashinfer's fp32 cache, up to 64 context tokens), checked once per token count and bound pool. - DSpark keeps the draft logits vocab-sharded and runs the Markov chain as trtllm::k3_markov (greedy drafting, plain TP over MNNVL, a bf16 column-parallel head whose shards tile the Markov vocabulary, bf16 Markov weights of the kernel's rank, a block, shard and batch the kernel splits). The kernel's tokens and next_new_tokens are the step's drafts and next inputs, and it applies the pending KV-length rewind. _draft_block_logits is the new hook for the drafter's block logits; _mask_token_id and _trained_mask_embedding resolve the mask row that both the torch path and the kernel use. test_k3_decode_worker.py (CPU, fakes): k3_decode is off by default and the logits come from the logits processors; each predicate takes an eligible step and declines a step that fails one of its conditions; the DSpark chain's drafts, next_new_tokens and rewind come from k3_markov. Signed-off-by: Vasanth Sabavat --- tensorrt_llm/_torch/speculative/dflash.py | 633 +++++++++++++-- tensorrt_llm/_torch/speculative/dspark.py | 223 +++++- .../hw_agnostic/test_k3_decode_worker.py | 754 ++++++++++++++++++ 3 files changed, 1541 insertions(+), 69 deletions(-) create mode 100644 tests/unittest/_torch/speculative/hw_agnostic/test_k3_decode_worker.py diff --git a/tensorrt_llm/_torch/speculative/dflash.py b/tensorrt_llm/_torch/speculative/dflash.py index eb1204d31b26..746c5d36f0f2 100644 --- a/tensorrt_llm/_torch/speculative/dflash.py +++ b/tensorrt_llm/_torch/speculative/dflash.py @@ -678,6 +678,11 @@ class DFlashWorker(SpecWorkerBase): Reference: https://arxiv.org/pdf/2602.06036 """ + # Set by a Kimi K3 target running its decode kernels: decode steps then run the acceptance and the drafter's + # inputs as trtllm::k3_spec_accept, keep the target logits vocabulary-sharded and write the drafter's context + # K/V with trtllm::k3_ctx_kv. + k3_decode = False + def __init__( self, spec_config: "DFlashDecodingConfig", @@ -1356,6 +1361,432 @@ def _prepare_kv_for_draft_forward( attn_metadata.update_for_spec_dec() + def _mask_token_id(self, draft_model) -> int: + """The drafter's mask token id, resolved once.""" + if self._resolved_mask_token_id is None: + if ( + hasattr(self.spec_config, "mask_token_id") + and self.spec_config.mask_token_id is not None + ): + self._resolved_mask_token_id = self.spec_config.mask_token_id + elif hasattr(draft_model, "mask_token_id"): + self._resolved_mask_token_id = draft_model.mask_token_id + elif hasattr(draft_model.model, "mask_token_id"): + self._resolved_mask_token_id = draft_model.model.mask_token_id + else: + raise ValueError( + "DFlash requires mask_token_id to be set. Please set it in DFlashDecodingConfig " + "or ensure the draft model config has 'dflash_config.mask_token_id' or 'mask_token_id'." + ) + return self._resolved_mask_token_id + + @staticmethod + def _trained_mask_embedding(draft_model, mask_token_id: int) -> Optional[torch.Tensor]: + """The drafter's own trained mask row, if it kept one for the mask id in use.""" + if mask_token_id != getattr(draft_model, "mask_token_id", None): + return None + return getattr(draft_model, "mask_token_embedding", None) + + def _k3_accept_applies( + self, logits, attn_metadata, spec_metadata, draft_model, num_contexts: int, num_gens: int + ) -> bool: + """Whether this step's acceptance and the drafter's inputs run as ``trtllm::k3_spec_accept``. + + That mode: a decode step (no context requests, at most 8 gen requests) with greedy strict acceptance (no + rejection sampling, penalties or guided decoding, the base draft-token and logits layouts), fp32 target + logits of a vocabulary the kernel splits, the V2 Mamba manager's KDA replay record, the draft pool's block + table, and an unsharded bf16 draft embedding. Any other step keeps the torch path. + """ + if not self._k3_accept_step_applies( + attn_metadata, spec_metadata, draft_model, num_contexts, num_gens + ): + return False + K = spec_metadata.runtime_draft_len + if ( + logits.dim() != 2 + or logits.dtype != torch.float32 + or logits.shape[0] != num_gens * (K + 1) + ): + return False + from ..cute_dsl_kernels.k3_spec_accept import op as accept_op + + embed = draft_model.draft_model_full.model.embed_tokens + return accept_op.supports( + logits.shape[1], num_gens, self._compute_block_size, K, embed.weight.shape[1] + ) + + def _k3_accept_step_applies( + self, attn_metadata, spec_metadata, draft_model, num_contexts: int, num_gens: int + ) -> bool: + """Everything ``_k3_accept_applies`` checks but the target logits.""" + if not self.k3_decode: + return False + from ..pyexecutor.kv_cache.mamba_cache_manager import MambaHybridCacheManagerV2 + + if num_contexts != 0 or not 0 < num_gens <= 8 or self.guided_decoder is not None: + return False + if ( + not spec_metadata.is_all_greedy_sample + or self._can_use_rejection_sampling(spec_metadata) + or getattr(spec_metadata, "enable_penalty", False) + or type(self)._reshape_draft_tokens_for_accept + is not SpecWorkerBase._reshape_draft_tokens_for_accept + or type(self)._reshape_logits_for_accept + is not SpecWorkerBase._reshape_logits_for_accept + ): + return False + K = spec_metadata.runtime_draft_len + draft_tokens = spec_metadata.draft_tokens + if ( + K <= 0 + or draft_tokens is None + or draft_tokens.dtype != torch.int32 + or draft_tokens.numel() != num_gens * K + ): + return False + mgr = attn_metadata.kv_cache_manager + mamba_metadata = getattr(attn_metadata, "mamba_metadata", None) + if ( + not isinstance(mgr, MambaHybridCacheManagerV2) + or type(mgr).update_mamba_states is not MambaHybridCacheManagerV2.update_mamba_states + or not mgr.use_kda_replay_update + or getattr(mgr, "prev_num_accepted_tokens", None) is None + or getattr(mgr, "_dummy_request_mask", None) is None + or mamba_metadata is None + or mamba_metadata.state_indices.dtype != torch.int32 + ): + return False + if ( + self._ctx_block_tables is None + or getattr(attn_metadata, "draft_kv_cache_block_offsets", None) is None + or getattr(attn_metadata, "kv_lens_cuda", None) is None + ): + return False + embed = draft_model.draft_model_full.model.embed_tokens + return getattr(embed, "tp_size", 1) == 1 and embed.weight.dtype == torch.bfloat16 + + def target_logits( + self, hidden_states, lm_head, logits_processor, attn_metadata, spec_metadata, draft_model + ) -> torch.Tensor: + """This rank's bf16 vocabulary shard of the target logits, [rows, vocab / TP] (``_k3_head_shard``: the head + GEMM without its all-gather and fp32 cast), for a step whose acceptance runs as the sharded + ``trtllm::k3_spec_accept`` (see ``_k3_logits_shard``); the gathered fp32 logits otherwise. At such a step the + acceptance is their only reader: the one-model spec sampler takes the worker's tokens and rejects requests for + logits or log probabilities, so the shard is also what the step returns as ``"logits"``.""" + self._k3_step_shard = None + shard = self._k3_logits_shard( + hidden_states, lm_head, attn_metadata, spec_metadata, draft_model + ) + if shard is None: + return super().target_logits( + hidden_states, lm_head, logits_processor, attn_metadata, spec_metadata, draft_model + ) + logits = self._k3_head_shard(logits_processor, lm_head, hidden_states) + self._k3_step_shard = (logits, shard) + return logits + + @staticmethod + def _k3_head_shard(logits_processor, lm_head, rows: torch.Tensor) -> torch.Tensor: + """This rank's bf16 vocabulary shard of ``lm_head(rows)``, without the head's all-gather: from the logits + processor's own head kernel where it has one that takes the rows (``lm_head_shard``; the Kimi K3 target's + ``gemm/k3_head_gemv``), so the shard holds the values of the processor's gathered logits; else from the head's + own GEMM.""" + head_shard = getattr(logits_processor, "lm_head_shard", None) + logits = None if head_shard is None else head_shard(rows, lm_head) + if logits is None: + logits = lm_head.apply_linear(rows, lm_head.bias) + return logits + + def _k3_logits_shard(self, hidden_states, lm_head, attn_metadata, spec_metadata, draft_model): + """``(workspace, first column)`` when this step's target logits stay vocabulary-sharded, else None. + + That needs a step ``_k3_accept_step_applies`` takes, plain TP whose all-reduces own an MNNVL workspace for + this mapping, and an unquantized, bias-free, unpadded, evenly split column-parallel bf16 lm_head with shards + the kernel splits. The exchange is collective, so every rank must decide alike: every condition is the + batch's or the configuration's. The workspace is allocated (collectively) and the kernel compiled outside + CUDA-graph capture only, so a graph whose warmup allocated no workspace is captured on the gathered path. + """ + if not self.k3_decode: + return None + from ..cute_dsl_kernels.k3_spec_accept import op as accept_op + from ..distributed.ops import MNNVLAllReduce + + mapping = self.mapping + if getattr(self, "_k3_shard_head", None) is not lm_head: + reason = self._k3_logits_shard_reason(lm_head) + if reason is None and MNNVLAllReduce.allreduce_mnnvl_workspaces.get(mapping) is None: + # Not kept: re-checked on later steps, since the model's MNNVL workspace may not exist yet. + logger.info_once( + "DFlash: the target logits are all-gathered (the TP all-reduces are not MNNVL)", + key="dflash_k3_logits_no_mnnvl", + ) + return None + if mapping is not None and mapping.tp_size > 1: + if reason is None: + logger.info( + f"DFlash: decode-step target logits stay vocabulary-sharded ({lm_head.weight.shape[0]} " + f"columns per rank); trtllm::k3_spec_accept exchanges the row maxima, and a request's " + f"logits post-processors get them gathered" + ) + else: + logger.info_once( + f"DFlash: the target logits are all-gathered ({reason})", key=reason + ) + self._k3_shard_head, self._k3_shard_ok = lm_head, reason is None + if not self._k3_shard_ok: + return None + num_contexts = attn_metadata.num_contexts + num_gens = attn_metadata.num_seqs - num_contexts + K = spec_metadata.runtime_draft_len + if hidden_states.shape[0] != num_gens * (K + 1) or not self._k3_accept_step_applies( + attn_metadata, spec_metadata, draft_model, num_contexts, num_gens + ): + return None + embed = draft_model.draft_model_full.model.embed_tokens + columns = lm_head.weight.shape[0] + if not accept_op.supports( + columns, num_gens, self._compute_block_size, K, embed.weight.shape[1], mapping.tp_size + ): + return None + if torch.cuda.is_current_stream_capturing(): + ws = accept_op.existing_workspace(mapping) + if ws is None: + return None + else: + ws = accept_op.workspace(mapping) + return ws, mapping.tp_rank * columns + + def _k3_logits_shard_reason(self, lm_head): + """Why the target logits cannot stay vocabulary-sharded under this configuration and head (None if they can, + given an MNNVL workspace for the mapping, which ``_k3_logits_shard`` checks).""" + from ..cute_dsl_kernels.k3_spec_accept import op as accept_op + from ..modules.linear import TensorParallelMode + + mapping = self.mapping + if mapping is None or mapping.tp_size <= 1: + return "no TP" + if mapping.enable_attention_dp or mapping.has_cp() or mapping.has_pp(): + return "attention DP, CP or PP" + if ( + getattr(lm_head, "tp_mode", None) != TensorParallelMode.COLUMN + or not getattr(lm_head, "gather_output", False) + or getattr(lm_head, "padding_size", 0) != 0 + or getattr(lm_head, "gather_output_sizes", None) is not None + ): + return "the lm_head is not an evenly split, unpadded column-parallel head" + if ( + lm_head.bias is not None + or lm_head.has_any_quant + or lm_head.weight.dtype != torch.bfloat16 + ): + return "the lm_head has a bias, is quantized or is not bf16" + if not accept_op.supports_columns(lm_head.weight.shape[0], mapping.tp_size): + columns = lm_head.weight.shape[0] + return f"trtllm::k3_spec_accept cannot split {columns} columns over {mapping.tp_size} ranks" + return None + + def _k3_static_arange(self, lo: int, hi: int) -> torch.Tensor: + """``arange(lo, hi)`` (int64), built once outside CUDA-graph capture.""" + cache = self.__dict__.setdefault("_k3_aranges", {}) + t = cache.get((lo, hi)) + if t is None: + t = torch.arange(lo, hi, dtype=torch.long, device="cuda") + if not torch.cuda.is_current_stream_capturing(): + cache[(lo, hi)] = t + return t + + def _k3_mask_row(self, draft_model) -> torch.Tensor: + """The noise block's mask embedding, bf16 [hidden], cached: the drafter's own trained mask row if it kept one + for the mask id in use, else the embedding's row (``dflash_noise_block_embedding`` reads the same row).""" + row = getattr(self, "_k3_mask_embedding", None) + if row is None: + embed = draft_model.draft_model_full.model.embed_tokens + mask_token_id = self._mask_token_id(draft_model) + trained = self._trained_mask_embedding(draft_model, mask_token_id) + if trained is not None: + row = trained.to(embed.weight.dtype).reshape(-1).contiguous() + else: + row = embed.weight[mask_token_id].contiguous() + if torch.cuda.is_current_stream_capturing(): + return row + self._k3_mask_embedding = row + return row + + def _k3_ctx_pool_key(self): + """The identity of the pool the ctx cache is bound to now (KV cache estimation rebinds it to a new one).""" + return ( + tuple(t.data_ptr() for t in self._ctx_kv_buf) if self._ctx_kv_buf is not None else None + ) + + def _k3_ctx_kv_applies(self, draft_model, projected: torch.Tensor) -> bool: + """Whether ``trtllm::k3_ctx_kv`` writes this step's context K/V: the manager-bound paged pool with K and V, + bf16, the fused K/V weight without bias or context input norm, k_norm, NeoX RoPE from flashinfer's fp32 + cache, a context the op can split (up to 64 tokens: B <= 8 requests of K + 1 <= 8, ``ctx_op.pick_split``), one + allocation for every layer's pool. Checked once per token count and + bound pool; the pool view is kept for the bound pool only, so a rebind never leaves the kernel writing the + replaced pool.""" + pool_key = self._k3_ctx_pool_key() + if getattr(self, "_k3_ctx_pool_bound", None) != pool_key: + self._k3_ctx_pool_bound = pool_key + self._k3_ctx_pool = None + self._k3_ctx_kv_ok = {} + cache = self._k3_ctx_kv_ok + n = projected.shape[0] + ok = cache.get(n) + if ok is not None: + return ok + from ..cute_dsl_kernels.k3_ctx_kv import op as ctx_op + from ..models import modeling_dflash + + reason = None + if not (self._ctx_paged and self._ctx_block_tables is not None): + reason = "the context pool is not the manager-bound paged pool" + elif projected.dtype != torch.bfloat16 or self._ctx_kv_buf[0].dtype != torch.bfloat16: + reason = "not bf16" + else: + if draft_model._fused_kv_weight is None: + draft_model._build_fused_kv_buffers() + view = ctx_op.pool_view(list(self._ctx_kv_buf)) + if ( + draft_model._fused_kv_bias is not None + or getattr(draft_model, "_input_ln_eps", None) is not None + ): + reason = "K/V bias or context input norm" + elif draft_model._k_norm_stacked is None or not draft_model._is_neox: + reason = "no k_norm or not NeoX RoPE" + elif ( + modeling_dflash._flashinfer_rope is None + or draft_model._get_cos_sin_cache().shape[1] != 64 + ): + reason = "RoPE is not flashinfer's full-rotary fp32 cache" + elif view is None or self._ctx_kv_buf[0].size(1) != 2: + reason = "the layers' pools are not K/V views of one allocation" + elif not ctx_op.pick_split( + draft_model._fused_kv_weight.shape[0], + projected.shape[1], + draft_model._num_kv_heads, + self._draft_tokens_per_req, + n, + projected.device, + ): + reason = f"{n} context tokens (at most 64) or the K/V weight shape" + else: + self._k3_ctx_pool = view + ok = reason is None + if not ok: + logger.info(f"DFlash: context K/V on the Python path ({reason})") + if not torch.cuda.is_current_stream_capturing(): + cache[n] = ok + return ok + + def _k3_ctx_kv( + self, draft_model, projected, ctx_positions, num_accepted, slots, rows + ) -> torch.Tensor: + """``trtllm::k3_ctx_kv``: every drafter layer's context K/V of this step into the paged pool and + ``_ctx_len += num_accepted`` (clamped); returns the context length each request may advertise.""" + if self._k3_ctx_pool is None or self._k3_ctx_pool_bound != self._k3_ctx_pool_key(): + raise RuntimeError("k3_ctx_kv: the context pool was rebound after the path was chosen") + flat, layer_off, page_stride, kv_stride, head_stride = self._k3_ctx_pool + return torch.ops.trtllm.k3_ctx_kv( + projected.contiguous(), + draft_model._fused_kv_weight, + draft_model._k_norm_stacked, + draft_model._get_cos_sin_cache(), + ctx_positions, + num_accepted, + self._ctx_len, + slots, + rows, + self._ctx_block_tables, + self._ctx_block_counts, + flat, + layer_off, + page_stride, + kv_stride, + head_stride, + draft_model._k_norm_eps, + self._max_ctx, + self._ctx_page_size, + self._compute_block_size, + draft_model._num_kv_heads, + ) + + def _on_acceptance( + self, accepted_tokens, num_accepted_tokens, attn_metadata, spec_metadata + ) -> None: + """Called with this step's acceptance when ``k3_spec_accept`` produced it (a hook for drafter families).""" + + def _k3_accept(self, logits, attn_metadata, spec_metadata, draft_model, shard=None): + """``trtllm::k3_spec_accept`` for a step ``_k3_accept_applies`` accepted: the acceptance, the block table, + the KDA replay record, kv_lens + 1 and the drafter's inputs in one launch. ``shard``: the exchange + workspace and first column when ``logits`` are this rank's vocabulary shard (``_k3_logits_shard``). + Returns (accepted tokens, num accepted, {rewind, bonus, qpos, cpos, noise}).""" + from ..cute_dsl_kernels.k3_spec_accept import op as accept_op + + num_gens = attn_metadata.num_seqs + K = spec_metadata.runtime_draft_len + mgr = attn_metadata.kv_cache_manager + is_warmup = spec_metadata.is_cuda_graph and not torch.cuda.is_current_stream_capturing() + force = 0.0 if is_warmup else float(self.force_num_accepted_tokens) + mode, _, _ = accept_op.force_mode(force, K) + if mode == accept_op._kernel_module().FORCE_FRAC: + self._ensure_force_accept_rng_state(logits.device) + pool, counter = self._force_accept_rng_pool, self._force_accept_rng_counter + else: + dummy_rng = getattr(self, "_k3_dummy_rng", None) + if dummy_rng is None: + dummy_rng = self._k3_dummy_rng = ( + torch.zeros(4, dtype=torch.float32, device=logits.device), + torch.zeros(2, dtype=torch.int64, device=logits.device), + ) + pool, counter = dummy_rng + embed = draft_model.draft_model_full.model.embed_tokens + accepted, num_acc, rewind, bonus, qpos, cpos, noise = torch.ops.trtllm.k3_spec_accept( + logits, + spec_metadata.draft_tokens.reshape(num_gens, K), + attn_metadata.draft_kv_cache_block_offsets, + self._ctx_pool_idx, + self._ctx_block_divisor, + self._ctx_block_counts, + self._ctx_block_tables, + mgr.prev_num_accepted_tokens, + attn_metadata.mamba_metadata.state_indices, + mgr._dummy_request_mask, + attn_metadata.kv_lens_cuda, + self._batch_to_slot, + self._ctx_len, + self._max_ctx, + embed.weight, + self._k3_mask_row(draft_model), + pool, + counter, + force, + self._compute_block_size, + *self._k3_shard_args(shard), + ) + self._on_acceptance(accepted, num_acc, attn_metadata, spec_metadata) + return ( + accepted, + num_acc, + dict(rewind=rewind, bonus=bonus, qpos=qpos, cpos=cpos, noise=noise), + ) + + @staticmethod + def _k3_shard_args(shard) -> tuple: + """``trtllm::k3_spec_accept``'s exchange arguments for a vocabulary-sharded step (none otherwise).""" + if shard is None: + return () + ws, first_column = shard + return ( + ws["uc"], + ws["mc"], + ws["flags"], + ws["rank"], + ws["slots"], + ws["push_copies"], + first_column, + ) + def _apply_kv_rewind_after_draft(self, attn_metadata, spec_metadata): """Apply the deferred kv_lens rewind after the draft forward.""" self._kv_rewind_pending = False @@ -1573,6 +2004,15 @@ def _forward_impl( batch_size = attn_metadata.num_seqs num_contexts = attn_metadata.num_contexts num_gens = batch_size - num_contexts + # Set when target_logits kept this step's logits vocabulary-sharded (only k3_spec_accept reads them). + step_shard, self._k3_step_shard = getattr(self, "_k3_step_shard", None), None + shard = None + if step_shard is not None: + if step_shard[0] is not logits: + raise RuntimeError( + "DFlash: the target logits are not the vocabulary shard target_logits returned" + ) + shard = step_shard[1] raw_logits = logits K = spec_metadata.runtime_draft_len @@ -1594,10 +2034,25 @@ def _forward_impl( draft_model, spec_metadata, attn_metadata, draft_kv_cache_manager ) spec_metadata._dflash_worker = self + # A decode step of the supported mode runs the acceptance and the drafter's inputs as one kernel + # (trtllm::k3_spec_accept), which also decodes the block table below. + if shard is None: + fused_accept = self._k3_accept_applies( + logits, attn_metadata, spec_metadata, draft_model, num_contexts, num_gens + ) + elif self._k3_accept_step_applies( + attn_metadata, spec_metadata, draft_model, num_contexts, num_gens + ): + fused_accept = True + else: + raise RuntimeError( + "DFlash: the target logits are vocabulary-sharded, but the fused acceptance no longer applies" + ) # Before any store: prefill and decode both address pages through it. # Returning False here means an empty batch -- the missing-offsets case # raises inside, with the metadata type in the message. - self._refresh_ctx_block_tables(attn_metadata, batch_size) + if not fused_accept: + self._refresh_ctx_block_tables(attn_metadata, batch_size) # Save context lengths so both warmup and a failed forward can roll # back the in-place _ctx_len updates made during drafting. @@ -1616,9 +2071,17 @@ def _forward_impl( self._execute_guided_decoder_if_present(logits) - accepted_tokens, num_accepted_tokens = self.sample_and_accept_draft_tokens( - logits, attn_metadata, spec_metadata - ) + prep = None + if fused_accept: + # Saved before the kernel updates kv_lens (warmup also snapshots kv_lens_cuda here). + self._prepare_attn_metadata_for_dflash(attn_metadata, spec_metadata) + accepted_tokens, num_accepted_tokens, prep = self._k3_accept( + logits, attn_metadata, spec_metadata, draft_model, shard + ) + else: + accepted_tokens, num_accepted_tokens = self.sample_and_accept_draft_tokens( + logits, attn_metadata, spec_metadata + ) # Opt-in acceptance recording (env-gated; eager-mode measurement # runs only). Skipped for CUDA-graph batches (capture/replay/warmup @@ -1635,18 +2098,26 @@ def _forward_impl( num_accepted_tokens[num_contexts:batch_size].tolist(), ) - # Update GDN/Mamba recurrent states to the accepted token's state. - if num_gens > 0 and isinstance(attn_metadata.kv_cache_manager, MambaHybridCacheManager): - attn_metadata.kv_cache_manager.update_mamba_states( - attn_metadata=attn_metadata, - num_accepted_tokens=num_accepted_tokens, - state_indices=attn_metadata.mamba_metadata.state_indices, - ) + if fused_accept: + # The kernel recorded the KDA replay acceptance and advanced kv_lens_cuda. + self._kv_rewind_amount = prep["rewind"] + self._kv_rewind_nc = num_contexts + self._kv_rewind_bs = batch_size + self._kv_rewind_pending = True + attn_metadata.update_for_spec_dec() + else: + # Update GDN/Mamba recurrent states to the accepted token's state. + if num_gens > 0 and isinstance(attn_metadata.kv_cache_manager, MambaHybridCacheManager): + attn_metadata.kv_cache_manager.update_mamba_states( + attn_metadata=attn_metadata, + num_accepted_tokens=num_accepted_tokens, + state_indices=attn_metadata.mamba_metadata.state_indices, + ) - self._prepare_attn_metadata_for_dflash(attn_metadata, spec_metadata) - self._prepare_kv_for_draft_forward( - attn_metadata, num_accepted_tokens, num_contexts, batch_size - ) + self._prepare_attn_metadata_for_dflash(attn_metadata, spec_metadata) + self._prepare_kv_for_draft_forward( + attn_metadata, num_accepted_tokens, num_contexts, batch_size + ) # Collapse mrope [3, 1, N] to 1D by taking the first (temporal) dimension. # The draft model uses standard 1D RoPE, so only scalar positions are needed. @@ -1688,6 +2159,7 @@ def _forward_impl( spec_metadata=spec_metadata, draft_model=draft_model, total_target_tokens=total_target_tokens, + prep=prep, ) if num_gens > 0: @@ -1715,8 +2187,8 @@ def _forward_impl( gen_gather_ids = gen_gather_ids.clamp(max=hidden_states_out.shape[0] - 1) gen_hidden_states = hidden_states_out[gen_gather_ids] - gen_logits = draft_model.logits_processor( - gen_hidden_states, draft_model.lm_head, attn_metadata, True + gen_logits = self._draft_block_logits( + draft_model, gen_hidden_states, attn_metadata, spec_metadata ) vocab_size = gen_logits.shape[-1] @@ -1803,13 +2275,17 @@ def _forward_impl( self._restore_ctx_len_host() self._ctx_len_restore_pending = False - return { + outputs = { "logits": raw_logits, "new_tokens": accepted_tokens, "new_tokens_lens": num_accepted_tokens, "next_draft_tokens": next_draft_tokens, "next_new_tokens": next_new_tokens, } + if shard is not None: + # "logits" is this rank's vocabulary shard: the engine gathers it for any logits post-processor. + outputs["logits_vocab_shard"] = True + return outputs def _draft_block_width(self, draft_model) -> int: """Block slots the draft forward must compute for K draft tokens. @@ -1849,6 +2325,20 @@ def _refine_block_logits( """ return gen_logits + def _draft_block_logits( + self, + draft_model, + gen_hidden_states: torch.Tensor, + attn_metadata: AttentionMetadata, + spec_metadata, + ) -> torch.Tensor: + """fp32 logits of the block positions over the full vocabulary (the TP lm_head all-gathers its + shards). A drafter family whose refinement and draft sampler reduce across ranks themselves may + return this rank's vocabulary shard instead.""" + return draft_model.logits_processor( + gen_hidden_states, draft_model.lm_head, attn_metadata, True + ) + def _apply_dflash2_selector( self, draft_model, @@ -1944,9 +2434,13 @@ def prepare_1st_drafter_inputs( spec_metadata: DFlashSpecMetadata, draft_model: nn.Module, total_target_tokens: int = 0, + prep: Optional[dict] = None, ): """Prepare inputs for DFlash's draft forward. + ``prep``: ``trtllm::k3_spec_accept``'s bonus tokens, positions and noise embedding for a decode step it + took (see ``_k3_accept``); computed here otherwise. + For gen requests, builds: - noise_embedding: token embeddings for [accepted + mask] tokens - query_positions: position IDs for the query tokens @@ -1958,23 +2452,7 @@ def prepare_1st_drafter_inputs( batch_size = attn_metadata.num_seqs num_gens = batch_size - num_contexts - # Resolve mask_token_id and block_size once, cache for subsequent calls - if self._resolved_mask_token_id is None: - if ( - hasattr(self.spec_config, "mask_token_id") - and self.spec_config.mask_token_id is not None - ): - self._resolved_mask_token_id = self.spec_config.mask_token_id - elif hasattr(draft_model, "mask_token_id"): - self._resolved_mask_token_id = draft_model.mask_token_id - elif hasattr(draft_model.model, "mask_token_id"): - self._resolved_mask_token_id = draft_model.model.mask_token_id - else: - raise ValueError( - "DFlash requires mask_token_id to be set. Please set it in DFlashDecodingConfig " - "or ensure the draft model config has 'dflash_config.mask_token_id' or 'mask_token_id'." - ) - mask_token_id = self._resolved_mask_token_id + mask_token_id = self._mask_token_id(draft_model) # Get the embed_tokens layer from the draft model embed_tokens = draft_model.draft_model_full.model.embed_tokens @@ -2003,39 +2481,46 @@ def prepare_1st_drafter_inputs( # Get slots for gen requests from pre-computed mapping slots = self._batch_to_slot[num_contexts : num_contexts + num_gens] - gen_rows_out = torch.arange( - num_contexts, num_contexts + num_gens, dtype=torch.long, device="cuda" - ) - K_plus_1 = K + 1 - bonus_idx = (gen_num_accepted - 1).clamp_min(0).long().unsqueeze(1) - bonus = gen_accepted_tokens.gather(1, bonus_idx).squeeze(1).long() - - ctx_len_gen = self._ctx_len[slots] - j_block = torch.arange(query_tokens_per_req, dtype=torch.long, device="cuda") - offsets_kp1 = torch.arange(K_plus_1, dtype=torch.long, device="cuda") - - # _ctx_len is clamped to _max_ctx only AFTER this step's accepted - # tokens are folded in (see the update below), so the running length - # used here has to be clamped on its own -- otherwise a request that - # already sits at the ceiling indexes num_accepted positions past - # any position the sequence can legitimately reach. - ctx_len_now = (ctx_len_gen + gen_num_accepted.long()).clamp_(max=self._max_ctx) - query_position_ids = ctx_len_now.unsqueeze(1) + j_block.unsqueeze(0) - ctx_position_ids = ctx_len_gen.unsqueeze(1) + offsets_kp1.unsqueeze(0) - - # The drafter's own trained mask row, if it kept one for the mask id in use. - trained_mask_embedding = ( - getattr(draft_model, "mask_token_embedding", None) - if mask_token_id == getattr(draft_model, "mask_token_id", None) - else None - ) - noise_embed_2d = dflash_noise_block_embedding( - embed_tokens, bonus, mask_token_id, query_tokens_per_req, trained_mask_embedding - ) + if prep is not None: + gen_rows_out = self._k3_static_arange(num_contexts, num_contexts + num_gens) + offsets_kp1 = self._k3_static_arange(0, K_plus_1) + bonus = prep["bonus"] + query_position_ids = prep["qpos"] + ctx_position_ids = prep["cpos"] + noise_embed_2d = prep["noise"] + else: + gen_rows_out = torch.arange( + num_contexts, num_contexts + num_gens, dtype=torch.long, device="cuda" + ) + + bonus_idx = (gen_num_accepted - 1).clamp_min(0).long().unsqueeze(1) + bonus = gen_accepted_tokens.gather(1, bonus_idx).squeeze(1).long() + + ctx_len_gen = self._ctx_len[slots] + j_block = torch.arange(query_tokens_per_req, dtype=torch.long, device="cuda") + offsets_kp1 = torch.arange(K_plus_1, dtype=torch.long, device="cuda") + + # _ctx_len is clamped to _max_ctx only AFTER this step's accepted + # tokens are folded in (see the update below), so the running length + # used here has to be clamped on its own -- otherwise a request that + # already sits at the ceiling indexes num_accepted positions past + # any position the sequence can legitimately reach. + ctx_len_now = (ctx_len_gen + gen_num_accepted.long()).clamp_(max=self._max_ctx) + query_position_ids = ctx_len_now.unsqueeze(1) + j_block.unsqueeze(0) + ctx_position_ids = ctx_len_gen.unsqueeze(1) + offsets_kp1.unsqueeze(0) + + noise_embed_2d = dflash_noise_block_embedding( + embed_tokens, + bonus, + mask_token_id, + query_tokens_per_req, + self._trained_mask_embedding(draft_model, mask_token_id), + ) # Accumulate new accepted features into context buffers + fused_ctx_kv = False if has_target_features: gen_start = attn_metadata.num_ctx_tokens # Target now processes exactly K+1 tokens per gen req, so the @@ -2043,6 +2528,19 @@ def prepare_1st_drafter_inputs( gen_hs = captured_hs[gen_start : gen_start + num_gens * total_tokens_per_req] gen_hs_to_project = gen_hs.reshape(-1, gen_hs.shape[-1]) projected_to_store = draft_model.project_target_hidden(gen_hs_to_project) + fused_ctx_kv = prep is not None and self._k3_ctx_kv_applies( + draft_model, projected_to_store + ) + if fused_ctx_kv: + num_ctx_fused = self._k3_ctx_kv( + draft_model, + projected_to_store, + ctx_position_ids, + gen_num_accepted, + slots, + gen_rows_out, + ) + if has_target_features and not fused_ctx_kv: gen_num_accepted_long = gen_num_accepted.long() col_idx = self._ctx_len[slots].unsqueeze(1) + offsets_kp1.unsqueeze(0) write_mask = offsets_kp1.unsqueeze(0) < gen_num_accepted_long.unsqueeze(1) @@ -2095,8 +2593,11 @@ def prepare_1st_drafter_inputs( self._ctx_len[slots] += gen_num_accepted_long self._ctx_len.clamp_(max=self._max_ctx) - num_ctx_per_req_t = self._ctx_len[slots] - if self._ctx_block_tables is not None: + if fused_ctx_kv: + num_ctx_per_req_t = num_ctx_fused + else: + num_ctx_per_req_t = self._ctx_len[slots] + if not fused_ctx_kv and self._ctx_block_tables is not None: # The write above clamps columns to the request's allocation, so # a context that outruns it has its tail written over the last # valid slot. Truncate what is advertised to match, or the read diff --git a/tensorrt_llm/_torch/speculative/dspark.py b/tensorrt_llm/_torch/speculative/dspark.py index b8dff7490d4c..7422b03471b0 100644 --- a/tensorrt_llm/_torch/speculative/dspark.py +++ b/tensorrt_llm/_torch/speculative/dspark.py @@ -865,6 +865,9 @@ class DSparkWorker(DFlashWorker): :class:`DSv4DSparkWorker`. """ + # Whether this step's block logits are the Kimi K3 decode path's vocabulary shard (see ``_draft_block_logits``). + _k3_sharded_block_logits = False + def set_draft_model(self, draft_model) -> None: """Reject an unsupported vocab mapping here rather than mid-decode. @@ -942,9 +945,10 @@ def _apply_dspark_markov_bias( sampled chain; the rejection-sampling path samples from the same biased distributions (proposal conditioned on the greedy chain). - Handles a TP vocab-sharded draft lm_head by slicing markov_w2's rows - to this rank's contiguous shard and chaining through the TP-aware - global argmax. + Block logits the Kimi K3 decode path kept vocab-sharded (see ``_draft_block_logits``) run the + whole chain as ``trtllm::k3_markov``, which also returns the greedy draft tokens and + next_new_tokens. Other TP-sharded logits slice markov_w2's rows to this rank's contiguous + shard and chain through the TP-aware global argmax. """ # The d2t guard lives in set_draft_model: it is model-static, so raising # it here would surface a load-time config error per decode step. @@ -953,6 +957,8 @@ def _apply_dspark_markov_bias( # duplicate the draft head's sharding rules. A standalone drafter # borrows the target lm_head, whose gather_output defaults to True, so # the logits normally arrive full-vocab and this branch is skipped. + self._k3_markov = None + k3_sharded, self._k3_sharded_block_logits = self._k3_sharded_block_logits, False full_vocab = draft_model.markov_w2.shape[0] shard = gen_logits.shape[-1] vocab_slice = None @@ -969,6 +975,10 @@ def _apply_dspark_markov_bias( "TP column shard of it." ) vocab_slice = slice(mapping.tp_rank * shard, (mapping.tp_rank + 1) * shard) + if k3_sharded: + return self._k3_markov_chain( + draft_model, gen_logits, first_prev_tokens, vocab_slice + ) def argmax_fn(step_logits): # Full-vocab token ids (TP-aware when sharded); tokens stay in @@ -981,3 +991,210 @@ def argmax_fn(step_logits): argmax_fn=argmax_fn, vocab_slice=vocab_slice, ) + + def _keep_draft_logits_sharded(self, draft_model, spec_metadata, num_gens: int) -> bool: + """Keep the draft logits vocab-sharded: ``trtllm::k3_markov`` reduces across the ranks itself. + + The chain's global argmaxes are the greedy draft tokens, so the lm_head's all-gather of the + full logits and the draft sampler's gather are both skipped. Needs a Kimi K3 target running its + decode kernels (``k3_decode``), plain TP, greedy drafting (rejection sampling reads + full-vocabulary probabilities), an unquantized bias-free column-parallel bf16 lm_head whose + shards tile the Markov vocabulary, bf16 Markov weights of the kernel's rank, MNNVL (the model's + all-reduces own an MNNVL workspace for this mapping), and a block, shard and batch of + ``num_gens`` requests the kernel can split. + """ + if not self.k3_decode: + return False + from ..cute_dsl_kernels.k3_markov import op as k3_markov_op + from ..distributed.ops import MNNVLAllReduce + from ..modules.linear import TensorParallelMode + + mapping = self.mapping + lm_head = getattr(draft_model, "lm_head", None) + if ( + lm_head is None + or not getattr(draft_model, "has_markov_head", False) + or mapping is None + or mapping.tp_size <= 1 + or mapping.enable_attention_dp + or spec_metadata.wants_advanced_draft_sampling + or getattr(lm_head, "tp_mode", None) != TensorParallelMode.COLUMN + or not getattr(lm_head, "gather_output", False) + or getattr(lm_head, "bias", None) is not None + or lm_head.weight.dtype != torch.bfloat16 + or draft_model.markov_w1.dtype != torch.bfloat16 + or draft_model.markov_w2.dtype != torch.bfloat16 + or lm_head.weight.shape[0] * mapping.tp_size != draft_model.markov_w2.shape[0] + or MNNVLAllReduce.allreduce_mnnvl_workspaces.get(mapping) is None + ): + return False + shard = lm_head.weight.shape[0] + block = spec_metadata.runtime_draft_len + rank = k3_markov_op._kernel_module().MARKOV_RANK + if ( + not 0 < block <= k3_markov_op.WORKSPACE_MAX_BLOCK + or draft_model.markov_w1.dim() != 2 + or draft_model.markov_w1.shape[1] != rank + or tuple(draft_model.markov_w2.shape[1:]) != (rank,) + ): + return False + if k3_markov_op.pick_grid(shard, block, num_gens) == 0: + logger.warning_once( + f"DSpark Markov head: trtllm::k3_markov cannot split a {shard}-row vocab shard for " + f"{num_gens} requests; the draft logits are all-gathered and the chain runs unfused.", + key=f"dspark_k3_markov_shard_{num_gens}", + ) + return False + return True + + def _draft_block_logits( + self, + draft_model, + gen_hidden_states: torch.Tensor, + attn_metadata, + spec_metadata, + ) -> torch.Tensor: + """This rank's bf16 shard of the block logits when ``k3_markov`` takes them (it converts them to fp32 + exactly, as ``.float()`` would), from the drafter's logits processor's head kernel where it takes the rows, + else the head's own GEMM (``_k3_head_shard``); otherwise the base class' fp32 logits.""" + num_gens = gen_hidden_states.shape[0] // max(spec_metadata.runtime_draft_len, 1) + self._k3_sharded_block_logits = self._keep_draft_logits_sharded( + draft_model, spec_metadata, num_gens + ) + if self._k3_sharded_block_logits: + return self._k3_head_shard( + draft_model.logits_processor, draft_model.lm_head, gen_hidden_states + ) + return super()._draft_block_logits( + draft_model, gen_hidden_states, attn_metadata, spec_metadata + ) + + def sample_and_accept_draft_tokens(self, logits, attn_metadata, spec_metadata): + """The base acceptance; under ``k3_decode`` its outputs are kept for ``k3_markov``'s next_new_tokens.""" + accepted_tokens, num_accepted_tokens = super().sample_and_accept_draft_tokens( + logits, attn_metadata, spec_metadata + ) + if self.k3_decode: + self._on_acceptance(accepted_tokens, num_accepted_tokens, attn_metadata, spec_metadata) + return accepted_tokens, num_accepted_tokens + + def _on_acceptance( + self, accepted_tokens, num_accepted_tokens, attn_metadata, spec_metadata + ) -> None: + """Keeps this step's acceptance for ``k3_markov``'s next_new_tokens (and the metadata whose KV lengths it + rewinds).""" + self._k3_acceptance = ( + accepted_tokens, num_accepted_tokens, attn_metadata.num_contexts, spec_metadata, attn_metadata + ) # fmt: skip + + def _k3_markov_chain( + self, + draft_model, + gen_logits: torch.Tensor, + first_prev_tokens: torch.Tensor, + vocab_slice: slice, + ) -> torch.Tensor: + """The Markov chain as ``trtllm::k3_markov`` on this rank's shard of the block logits. + + Returns the corrected logits (fp32); keeps the kernel's greedy tokens for + ``sample_draft_tokens`` and its next_new_tokens (built from this step's acceptance) for + ``_prepare_next_new_tokens``. + """ + from ..cute_dsl_kernels.k3_markov import op as k3_markov_op + + acceptance = getattr(self, "_k3_acceptance", None) + if acceptance is None: + raise RuntimeError("DSpark Markov head: k3_markov needs this step's acceptance first") + accepted_tokens, num_accepted_tokens, num_contexts, spec_metadata, attn_metadata = ( + acceptance + ) + num_gens = gen_logits.shape[0] + # The draft forward's pending KV-length rewind (see _apply_kv_rewind_after_draft) runs in the kernel: every + # reader of kv_lens_cuda in the draft forward precedes it. In warmup the restore of kv_lens_cuda that + # follows overwrites it, as it would the Python rewind. + rewind = getattr(self, "_kv_rewind_amount", None) + fold = ( + getattr(self, "_kv_rewind_pending", False) + and rewind is not None + and getattr(attn_metadata, "kv_lens_cuda", None) is not None + and rewind.dtype == torch.int32 + and self._kv_rewind_bs - self._kv_rewind_nc == rewind.numel() <= num_gens + ) + corrected, tokens, next_new = k3_markov_op.markov_chain( + self.mapping, + gen_logits.contiguous(), + first_prev_tokens.long(), + draft_model.markov_w1, + draft_model.markov_w2[vocab_slice], + vocab_slice.start, + accepted_tokens, + num_accepted_tokens[num_contexts : num_contexts + num_gens], + spec_metadata.batch_indices_cuda[num_contexts : num_contexts + num_gens], + kv_lens=attn_metadata.kv_lens_cuda if fold else None, + rewind=rewind if fold else None, + rewind_first=self._kv_rewind_nc if fold else 0, + ) + if fold: + self._kv_rewind_amount = None + self._kv_rewind_pending = False + self._k3_markov = (corrected, tokens, next_new, accepted_tokens, num_accepted_tokens) + return corrected + + def sample_draft_tokens( + self, + logits, + spec_metadata, + batch_size, + *, + num_contexts=0, + draft_step=None, + mapping_lm_head_tp=None, + ): + """Greedy block drafts straight from ``k3_markov``. + + When ``logits`` are the corrected logits the kernel just returned, its per-position global + argmax (first maximum, lowest vocabulary index among equal values) is the token the + TP-gathered greedy sampler would pick. Anything else goes to the base sampler. + """ + chain = getattr(self, "_k3_markov", None) + if ( + chain is not None + and chain[0] is logits + and mapping_lm_head_tp is None + and not spec_metadata.wants_advanced_draft_sampling + ): + self._k3_markov_next = (chain[1], chain[2], chain[3], chain[4]) + return chain[1] + self._k3_markov_next = None + return super().sample_draft_tokens( + logits, + spec_metadata, + batch_size, + num_contexts=num_contexts, + draft_step=draft_step, + mapping_lm_head_tp=mapping_lm_head_tp, + ) + + def _prepare_next_new_tokens( + self, + accepted_tokens, + next_draft_tokens, + batch_indices_cuda, + batch_size, + num_accepted_tokens, + ): + """``k3_markov``'s next_new_tokens when the drafts are its tokens for the whole batch (no context + requests); otherwise the base assembly.""" + chain = getattr(self, "_k3_markov_next", None) + self._k3_markov_next = None + if ( + chain is not None + and chain[0] is next_draft_tokens + and chain[2] is accepted_tokens + and chain[3] is num_accepted_tokens + and next_draft_tokens.shape[0] == batch_size + ): + return chain[1] + return super()._prepare_next_new_tokens( + accepted_tokens, next_draft_tokens, batch_indices_cuda, batch_size, num_accepted_tokens + ) diff --git a/tests/unittest/_torch/speculative/hw_agnostic/test_k3_decode_worker.py b/tests/unittest/_torch/speculative/hw_agnostic/test_k3_decode_worker.py new file mode 100644 index 000000000000..ad9a82b6ed89 --- /dev/null +++ b/tests/unittest/_torch/speculative/hw_agnostic/test_k3_decode_worker.py @@ -0,0 +1,754 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""The DFlash / DSpark worker's Kimi K3 decode kernels behind ``k3_decode`` (host-side, fakes only). + +* Off by default: the target logits are the logits processor's, the draft logits the drafter's processor's, and no + kernel predicate reads past the gate. +* Each kernel's predicate takes an eligible step and declines a step that fails one of its conditions: + ``trtllm::k3_spec_accept`` (``_k3_accept_applies``) and its vocabulary-sharded target logits (``target_logits``), + ``trtllm::k3_ctx_kv`` (``_k3_ctx_kv_applies``) and ``trtllm::k3_markov`` (``_keep_draft_logits_sharded``). +* DSpark's chain: sharded block logits go to ``k3_markov``, whose tokens and next_new_tokens are the step's. +""" + +from types import SimpleNamespace + +import pytest +import torch +from torch import nn + +from tensorrt_llm._torch.cute_dsl_kernels.k3_ctx_kv import op as ctx_op +from tensorrt_llm._torch.cute_dsl_kernels.k3_markov import op as markov_op +from tensorrt_llm._torch.cute_dsl_kernels.k3_spec_accept import op as accept_op +from tensorrt_llm._torch.distributed.ops import MNNVLAllReduce +from tensorrt_llm._torch.models import modeling_dflash +from tensorrt_llm._torch.modules.linear import TensorParallelMode +from tensorrt_llm._torch.pyexecutor.kv_cache.mamba_cache_manager import MambaHybridCacheManagerV2 +from tensorrt_llm._torch.speculative.dflash import DFlashWorker +from tensorrt_llm._torch.speculative.dspark import DSparkWorker +from tensorrt_llm._torch.speculative.interface import SpecWorkerBase +from tensorrt_llm.mapping import Mapping + +pytestmark = pytest.mark.cpu_only + +NUM_GENS, K, BLOCK, HIDDEN = 2, 7, 8, 32 +TP, RANK_IN_TP = 4, 2 +VOCAB = 64 +SHARD = VOCAB // TP +MARKOV_RANK = 256 +TP4 = Mapping(world_size=TP, rank=RANK_IN_TP, tp_size=TP) +WORKSPACE = { + "uc": "uc", + "mc": "mc", + "flags": "flags", + "rank": RANK_IN_TP, + "slots": TP, + "push_copies": 1, +} + + +@pytest.fixture(autouse=True) +def eager(monkeypatch): + """No CUDA-graph capture on the host.""" + monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: False) + + +@pytest.fixture +def mnnvl(monkeypatch): + """The model's all-reduces own an MNNVL workspace for TP4.""" + workspaces = {TP4: object()} + monkeypatch.setattr(MNNVLAllReduce, "allreduce_mnnvl_workspaces", workspaces) + return workspaces + + +def _worker(cls=DFlashWorker, **attrs): + """A worker without ``__init__`` (it needs flashinfer and a drafter): the attributes a test sets.""" + worker = cls.__new__(cls) + nn.Module.__init__(worker) + worker.guided_decoder = None + for name, value in attrs.items(): + setattr(worker, name, value) + return worker + + +class _OwnStateUpdate(MambaHybridCacheManagerV2): + """A manager whose recurrent-state update is not the V2 one the kernel reproduces.""" + + def update_mamba_states(self, *args, **kwargs): + pass + + +def _kda_manager(cls=MambaHybridCacheManagerV2): + """The V2 hybrid manager with the KDA replay record, without ``__init__``.""" + mgr = object.__new__(cls) + mgr._use_kda_replay_update = True + mgr.prev_num_accepted_tokens = torch.zeros(4, dtype=torch.int32) + mgr._dummy_request_mask = torch.zeros(4, dtype=torch.bool) + return mgr + + +def _drafter(): + embed = SimpleNamespace(weight=torch.zeros(VOCAB, HIDDEN, dtype=torch.bfloat16)) + return SimpleNamespace( + draft_model_full=SimpleNamespace(model=SimpleNamespace(embed_tokens=embed)) + ) + + +def _accept_case(): + """An eligible decode step for ``trtllm::k3_spec_accept``: two generation requests of K drafts, all greedy.""" + spec = SimpleNamespace( + is_all_greedy_sample=True, + use_rejection_sampling=False, + enable_penalty=False, + runtime_draft_len=K, + draft_tokens=torch.zeros(NUM_GENS * K, dtype=torch.int32), + ) + attn = SimpleNamespace( + num_contexts=0, + num_seqs=NUM_GENS, + kv_cache_manager=_kda_manager(), + mamba_metadata=SimpleNamespace(state_indices=torch.zeros(4, dtype=torch.int32)), + draft_kv_cache_block_offsets=torch.zeros(1, 4, 2, 3, dtype=torch.int32), + kv_lens_cuda=torch.zeros(4, dtype=torch.int32), + ) + worker = _worker( + k3_decode=True, + mapping=TP4, + _ctx_block_tables=torch.zeros(4, 3, dtype=torch.int32), + _compute_block_size=BLOCK, + ) + return SimpleNamespace( + worker=worker, + attn=attn, + spec=spec, + drafter=_drafter(), + logits=torch.zeros(NUM_GENS * (K + 1), VOCAB), + num_contexts=0, + num_gens=NUM_GENS, + supported=True, + ) + + +def _accept_applies(case, monkeypatch): + calls = [] + monkeypatch.setattr(accept_op, "supports", lambda *args: calls.append(args) or case.supported) + applies = case.worker._k3_accept_applies( + case.logits, case.attn, case.spec, case.drafter, case.num_contexts, case.num_gens + ) + return applies, calls + + +def test_k3_decode_is_off_by_default(): + assert DFlashWorker.k3_decode is False + assert DSparkWorker.k3_decode is False + assert _worker().k3_decode is False + assert _worker(DSparkWorker).k3_decode is False + + +def test_target_logits_are_the_logits_processors_when_off(): + """The stock path: the processor's gathered logits, nothing else read.""" + gathered = torch.zeros(3, VOCAB) + calls = [] + processor = SimpleNamespace( + forward=lambda *args: calls.append(args) or gathered, + lm_head_shard=lambda *args: pytest.fail("the shard hook ran with k3_decode off"), + ) + worker = _worker(mapping=TP4) + hidden, lm_head, attn = torch.zeros(3, HIDDEN), object(), object() + + assert worker.target_logits(hidden, lm_head, processor, attn, object(), object()) is gathered + assert worker._k3_step_shard is None + assert len(calls) == 1 + assert calls[0][0] is hidden and calls[0][1] is lm_head and calls[0][2] is attn + assert calls[0][3] is True + + +def test_k3_spec_accept_takes_an_eligible_decode_step(monkeypatch): + case = _accept_case() + applies, calls = _accept_applies(case, monkeypatch) + assert applies + assert calls == [(VOCAB, NUM_GENS, BLOCK, K, HIDDEN)] + + +def _set(path, value): + """``case. = value``; ``value`` may be a function of the case.""" + *owners, name = path.split(".") + + def mutate(case): + owner = case + for attr in owners: + owner = getattr(owner, attr) + setattr(owner, name, value(case) if callable(value) else value) + + return mutate + + +_ACCEPT_DECLINES = { + "k3_decode off": _set("worker.k3_decode", False), + "a context request": _set("num_contexts", 1), + "no generation request": _set("num_gens", 0), + "nine generation requests": _set("num_gens", 9), + "guided decoding": _set("worker.guided_decoder", object()), + "a non-greedy batch": _set("spec.is_all_greedy_sample", False), + "occurrence penalties": _set("spec.enable_penalty", True), + "no drafts": _set("spec.runtime_draft_len", 0), + "no draft tokens": _set("spec.draft_tokens", None), + "int64 draft tokens": _set("spec.draft_tokens", lambda c: c.spec.draft_tokens.long()), + "draft tokens of one request": _set("spec.draft_tokens", lambda c: c.spec.draft_tokens[:K]), + "another KV cache manager": _set( + "attn.kv_cache_manager", SimpleNamespace(use_kda_replay_update=True) + ), + "another recurrent-state update": _set( + "attn.kv_cache_manager", lambda c: _kda_manager(_OwnStateUpdate) + ), + "no KDA replay": _set("attn.kv_cache_manager._use_kda_replay_update", False), + "no replay record": _set("attn.kv_cache_manager.prev_num_accepted_tokens", None), + "no dummy-request mask": _set("attn.kv_cache_manager._dummy_request_mask", None), + "no Mamba metadata": _set("attn.mamba_metadata", None), + "int64 state indices": _set( + "attn.mamba_metadata.state_indices", torch.zeros(4, dtype=torch.int64) + ), + "the private context arena": _set("worker._ctx_block_tables", None), + "no draft block offsets": _set("attn.draft_kv_cache_block_offsets", None), + "no KV lengths": _set("attn.kv_lens_cuda", None), + "a vocab-sharded draft embedding": _set( + "drafter.draft_model_full.model.embed_tokens.tp_size", 2 + ), + "an fp32 draft embedding": _set( + "drafter.draft_model_full.model.embed_tokens.weight", torch.zeros(VOCAB, HIDDEN) + ), + "bf16 target logits": _set("logits", lambda c: c.logits.bfloat16()), + "target logits of one request": _set("logits", lambda c: c.logits[: K + 1]), + "a vocabulary the kernel does not split": _set("supported", False), +} + + +@pytest.mark.parametrize("mutate", list(_ACCEPT_DECLINES.values()), ids=list(_ACCEPT_DECLINES)) +def test_k3_spec_accept_declines(mutate, monkeypatch): + case = _accept_case() + mutate(case) + applies, _ = _accept_applies(case, monkeypatch) + assert not applies + + +def _head(**overrides): + """This rank's shard of a column-parallel bf16 LM head; ``apply_linear`` is its own GEMM, recorded.""" + head = SimpleNamespace( + tp_mode=TensorParallelMode.COLUMN, + gather_output=True, + padding_size=0, + gather_output_sizes=None, + bias=None, + has_any_quant=False, + weight=torch.zeros(SHARD, HIDDEN, dtype=torch.bfloat16), + gemm_calls=[], + ) + head.apply_linear = lambda rows, bias: ( + head.gemm_calls.append((rows, bias)) + or torch.zeros(rows.shape[0], SHARD, dtype=torch.bfloat16) + ) + for name, value in overrides.items(): + setattr(head, name, value) + return head + + +def _processor(shard=None, hook=True): + """A logits processor: ``forward`` gathers (recorded); with ``hook``, ``lm_head_shard`` returns ``shard``.""" + processor = SimpleNamespace( + gathered=torch.zeros(NUM_GENS * (K + 1), VOCAB), forward_calls=[], shard_calls=[] + ) + processor.forward = lambda *args: processor.forward_calls.append(args) or processor.gathered + if hook: + processor.lm_head_shard = ( + lambda rows, head: processor.shard_calls.append((rows, head)) or shard + ) + return processor + + +@pytest.fixture +def shard_kernel(monkeypatch, mnnvl): + """The exchange kernel takes the shard and the batch; ``workspace`` allocates (recorded).""" + allocated = [] + monkeypatch.setattr(accept_op, "supports_columns", lambda columns, slots: True) + monkeypatch.setattr(accept_op, "supports", lambda *args: True) + monkeypatch.setattr( + accept_op, "workspace", lambda mapping: allocated.append(mapping) or WORKSPACE + ) + monkeypatch.setattr(accept_op, "existing_workspace", lambda mapping: None) + return allocated + + +def _target_logits(case, head, processor): + hidden = torch.zeros(case.logits.shape[0], HIDDEN, dtype=torch.bfloat16) + logits = case.worker.target_logits(hidden, head, processor, case.attn, case.spec, case.drafter) + return hidden, logits + + +def test_target_logits_stay_vocabulary_sharded(shard_kernel): + """This rank's shard from the processor's head kernel; the exchange workspace and first column kept.""" + case, head = _accept_case(), _head() + shard = torch.ones(NUM_GENS * (K + 1), SHARD, dtype=torch.bfloat16) + processor = _processor(shard) + + hidden, logits = _target_logits(case, head, processor) + + assert logits is shard + assert case.worker._k3_step_shard[0] is shard + assert case.worker._k3_step_shard[1] == (WORKSPACE, RANK_IN_TP * SHARD) + assert shard_kernel == [TP4] + assert len(processor.shard_calls) == 1 and processor.shard_calls[0][0] is hidden + assert processor.shard_calls[0][1] is head + assert not processor.forward_calls and not head.gemm_calls + + +@pytest.mark.parametrize( + "hook", [True, False], ids=["the head kernel declines the rows", "no head kernel"] +) +def test_sharded_target_logits_fall_back_to_the_heads_gemm(shard_kernel, hook): + case, head = _accept_case(), _head() + processor = _processor(None, hook=hook) + + hidden, logits = _target_logits(case, head, processor) + + assert logits.shape == (NUM_GENS * (K + 1), SHARD) and logits.dtype == torch.bfloat16 + assert len(head.gemm_calls) == 1 + assert head.gemm_calls[0][0] is hidden and head.gemm_calls[0][1] is None + assert case.worker._k3_step_shard[0] is logits + assert not processor.forward_calls + + +_SHARD_DECLINES = { + "k3_decode off": _set("worker.k3_decode", False), + "one rank": _set("worker.mapping", Mapping()), + "attention DP": _set( + "worker.mapping", + Mapping(world_size=TP, rank=RANK_IN_TP, tp_size=TP, enable_attention_dp=True), + ), + "a row-parallel head": _set("head.tp_mode", TensorParallelMode.ROW), + "a head that does not gather": _set("head.gather_output", False), + "a padded vocabulary": _set("head.padding_size", 8), + "a quantized head": _set("head.has_any_quant", True), + "a head with a bias": _set("head.bias", torch.zeros(SHARD, dtype=torch.bfloat16)), + "an fp32 head": _set("head.weight", torch.zeros(SHARD, HIDDEN)), + "a non-greedy batch": _set("spec.is_all_greedy_sample", False), + "a context request": _set("attn.num_contexts", 1), + "rows of another step": _set("logits", lambda c: c.logits[: K + 1]), +} + + +@pytest.mark.parametrize("mutate", list(_SHARD_DECLINES.values()), ids=list(_SHARD_DECLINES)) +def test_target_logits_are_gathered_outside_the_sharded_mode(shard_kernel, mutate): + case = _accept_case() + case.head = _head() + mutate(case) + processor = _processor(torch.ones(1)) + + _, logits = _target_logits(case, case.head, processor) + + assert logits is processor.gathered + assert case.worker._k3_step_shard is None + assert not processor.shard_calls and not case.head.gemm_calls + + +@pytest.mark.parametrize( + "kernel", + ["supports_columns", "supports"], + ids=["a shard the kernel does not split", "a batch the kernel does not split"], +) +def test_target_logits_are_gathered_where_the_kernel_declines(shard_kernel, monkeypatch, kernel): + monkeypatch.setattr(accept_op, kernel, lambda *args: False) + case, processor = _accept_case(), _processor(torch.ones(1)) + + _, logits = _target_logits(case, _head(), processor) + + assert logits is processor.gathered and case.worker._k3_step_shard is None + + +def test_target_logits_wait_for_an_mnnvl_workspace(shard_kernel, mnnvl): + """Without the model's MNNVL workspace the logits are gathered, and the check is repeated on the next step.""" + case, head = _accept_case(), _head() + shard = torch.ones(NUM_GENS * (K + 1), SHARD, dtype=torch.bfloat16) + workspace = mnnvl.pop(TP4) + + _, logits = _target_logits(case, head, _processor(shard)) + assert logits is not shard and case.worker._k3_step_shard is None + + mnnvl[TP4] = workspace + _, logits = _target_logits(case, head, _processor(shard)) + assert logits is shard + + +def test_target_logits_under_capture_need_a_workspace_from_warmup(shard_kernel, monkeypatch): + monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: True) + case, head = _accept_case(), _head() + shard = torch.ones(NUM_GENS * (K + 1), SHARD, dtype=torch.bfloat16) + + _, logits = _target_logits(case, head, _processor(shard)) + assert logits is not shard and not shard_kernel + + monkeypatch.setattr(accept_op, "existing_workspace", lambda mapping: WORKSPACE) + _, logits = _target_logits(case, head, _processor(shard)) + assert logits is shard and not shard_kernel + assert case.worker._k3_step_shard[1] == (WORKSPACE, RANK_IN_TP * SHARD) + + +def test_forward_refuses_logits_other_than_the_shard_it_returned(): + worker = _worker(k3_decode=True) + worker._k3_step_shard = (torch.zeros(1), (WORKSPACE, 0)) + attn = SimpleNamespace(num_seqs=1, num_contexts=0) + with pytest.raises(RuntimeError, match="not the vocabulary shard"): + worker._forward_impl(None, None, None, torch.zeros(1), attn, None, None) + assert worker._k3_step_shard is None + + +# trtllm::k3_ctx_kv: the drafter's context K / V. + +CTX_TOKENS = NUM_GENS * (K + 1) +POOL_VIEW = ("base", "layer offsets", 1, 2, 3) + + +def _ctx_case(monkeypatch, split=8): + """An eligible context write: the manager-bound paged pool (two layers, one K / V head of 64) and a fused K / V + weight with k_norm and NeoX RoPE from flashinfer's fp32 cache.""" + splits = [] + monkeypatch.setattr(ctx_op, "pool_view", lambda layers: POOL_VIEW) + monkeypatch.setattr(ctx_op, "pick_split", lambda *args: splits.append(args) or split) + monkeypatch.setattr(modeling_dflash, "_flashinfer_rope", object()) + worker = _worker( + k3_decode=True, + spec_config=SimpleNamespace(max_draft_len=K), + _ctx_paged=True, + _ctx_block_tables=torch.zeros(4, 3, dtype=torch.int32), + _ctx_kv_buf=[torch.zeros(4, 2, 1, 32, 64, dtype=torch.bfloat16) for _ in range(2)], + ) + drafter = SimpleNamespace( + _fused_kv_weight=torch.zeros(2 * 2 * 64, HIDDEN, dtype=torch.bfloat16), + _fused_kv_bias=None, + _input_ln_eps=None, + _k_norm_stacked=torch.ones(2, 64, dtype=torch.bfloat16), + _is_neox=True, + _num_kv_heads=1, + cos_sin=torch.zeros(128, 64), + ) + drafter._get_cos_sin_cache = lambda: drafter.cos_sin + projected = torch.zeros(CTX_TOKENS, HIDDEN, dtype=torch.bfloat16) + return SimpleNamespace(worker=worker, drafter=drafter, projected=projected, splits=splits) + + +def test_k3_ctx_kv_takes_an_eligible_context(monkeypatch): + case = _ctx_case(monkeypatch) + assert case.worker._k3_ctx_kv_applies(case.drafter, case.projected) + assert case.worker._k3_ctx_pool is POOL_VIEW + assert case.splits == [(2 * 2 * 64, HIDDEN, 1, K + 1, CTX_TOKENS, case.projected.device)] + + +def _ctx_pool(dtype=torch.bfloat16, halves=2): + return [torch.zeros(4, halves, 1, 32, 64, dtype=dtype) for _ in range(2)] + + +_CTX_DECLINES = { + "the private context arena": _set("worker._ctx_paged", False), + "no manager block table": _set("worker._ctx_block_tables", None), + "fp32 projections": _set("projected", lambda c: c.projected.float()), + "an fp32 pool": _set("worker._ctx_kv_buf", lambda c: _ctx_pool(torch.float32)), + "a K / V bias": _set("drafter._fused_kv_bias", torch.zeros(2 * 2 * 64, dtype=torch.bfloat16)), + "a context input norm": _set("drafter._input_ln_eps", 1e-6), + "no k_norm": _set("drafter._k_norm_stacked", None), + "GPT-J RoPE": _set("drafter._is_neox", False), + "a partial-rotary cache": _set("drafter.cos_sin", torch.zeros(128, 32)), + "a single-latent pool": _set("worker._ctx_kv_buf", lambda c: _ctx_pool(halves=1)), +} + + +@pytest.mark.parametrize("mutate", list(_CTX_DECLINES.values()), ids=list(_CTX_DECLINES)) +def test_k3_ctx_kv_declines(monkeypatch, mutate): + case = _ctx_case(monkeypatch) + mutate(case) + assert not case.worker._k3_ctx_kv_applies(case.drafter, case.projected) + assert case.worker._k3_ctx_pool is None + + +@pytest.mark.parametrize( + "decline", + ["rope", "view", "split"], + ids=["no flashinfer RoPE", "per-layer allocations", "a context the kernel does not split"], +) +def test_k3_ctx_kv_declines_where_its_helpers_do(monkeypatch, decline): + case = _ctx_case(monkeypatch, split=0 if decline == "split" else 8) + if decline == "rope": + monkeypatch.setattr(modeling_dflash, "_flashinfer_rope", None) + if decline == "view": + monkeypatch.setattr(ctx_op, "pool_view", lambda layers: None) + assert not case.worker._k3_ctx_kv_applies(case.drafter, case.projected) + assert case.worker._k3_ctx_pool is None + + +def test_k3_ctx_kv_decides_per_token_count_and_bound_pool(monkeypatch): + case = _ctx_case(monkeypatch) + worker = case.worker + assert worker._k3_ctx_kv_applies(case.drafter, case.projected) + + monkeypatch.setattr(ctx_op, "pick_split", lambda *args: 0) + assert worker._k3_ctx_kv_applies(case.drafter, case.projected) # kept for this token count + assert not worker._k3_ctx_kv_applies(case.drafter, case.projected[: K + 1]) + + # KV cache estimation rebinds the pool; the replaced one stays alive so its addresses are not reused. + replaced, worker._ctx_kv_buf = worker._ctx_kv_buf, _ctx_pool() + with pytest.raises(RuntimeError, match="rebound"): + worker._k3_ctx_kv(case.drafter, case.projected, None, None, None, None) + assert not worker._k3_ctx_kv_applies(case.drafter, case.projected) + assert worker._k3_ctx_pool is None + assert len(replaced) == 2 + + +# trtllm::k3_markov: DSpark's draft logits kept vocab-sharded. + + +def _markov_drafter(): + drafter = SimpleNamespace( + lm_head=_head(), + has_markov_head=True, + markov_w1=torch.zeros(VOCAB, MARKOV_RANK, dtype=torch.bfloat16), + markov_w2=torch.zeros(VOCAB, MARKOV_RANK, dtype=torch.bfloat16), + chain_calls=[], + ) + drafter.apply_markov_chain_logits = lambda logits, first, argmax_fn, vocab_slice: ( + drafter.chain_calls.append((logits, first, vocab_slice)) or logits + ) + return drafter + + +@pytest.fixture +def markov_kernel(monkeypatch, mnnvl): + """The Markov kernel's rank; ``pick_grid`` splits every shard (recorded).""" + grids = [] + monkeypatch.setattr( + markov_op, "_kernel_module", lambda: SimpleNamespace(MARKOV_RANK=MARKOV_RANK) + ) + monkeypatch.setattr(markov_op, "pick_grid", lambda *args: grids.append(args) or 128) + return grids + + +def _markov_case(): + return SimpleNamespace( + worker=_worker(DSparkWorker, k3_decode=True, mapping=TP4), + drafter=_markov_drafter(), + spec=SimpleNamespace(wants_advanced_draft_sampling=False, runtime_draft_len=K), + ) + + +def test_k3_markov_keeps_eligible_draft_logits_sharded(markov_kernel): + case = _markov_case() + assert case.worker._keep_draft_logits_sharded(case.drafter, case.spec, NUM_GENS) + assert markov_kernel == [(SHARD, K, NUM_GENS)] + + +_MARKOV_DECLINES = { + "k3_decode off": _set("worker.k3_decode", False), + "no lm_head": _set("drafter.lm_head", None), + "no Markov head": _set("drafter.has_markov_head", False), + "one rank": _set("worker.mapping", Mapping()), + "attention DP": _set( + "worker.mapping", + Mapping(world_size=TP, rank=RANK_IN_TP, tp_size=TP, enable_attention_dp=True), + ), + "advanced draft sampling": _set("spec.wants_advanced_draft_sampling", True), + "a row-parallel head": _set("drafter.lm_head.tp_mode", TensorParallelMode.ROW), + "a head that does not gather": _set("drafter.lm_head.gather_output", False), + "a head with a bias": _set("drafter.lm_head.bias", torch.zeros(SHARD, dtype=torch.bfloat16)), + "an fp32 head": _set("drafter.lm_head.weight", torch.zeros(SHARD, HIDDEN)), + "fp32 markov_w1": _set("drafter.markov_w1", torch.zeros(VOCAB, MARKOV_RANK)), + "fp32 markov_w2": _set("drafter.markov_w2", torch.zeros(VOCAB, MARKOV_RANK)), + "shards that do not tile the Markov vocabulary": _set( + "drafter.markov_w2", torch.zeros(VOCAB - SHARD, MARKOV_RANK, dtype=torch.bfloat16) + ), + "a block longer than the kernel's": _set( + "spec.runtime_draft_len", markov_op.WORKSPACE_MAX_BLOCK + 1 + ), + "another Markov rank": lambda c: ( + _set("drafter.markov_w1", torch.zeros(VOCAB, 128, dtype=torch.bfloat16))(c), + _set("drafter.markov_w2", torch.zeros(VOCAB, 128, dtype=torch.bfloat16))(c), + ), +} + + +@pytest.mark.parametrize("mutate", list(_MARKOV_DECLINES.values()), ids=list(_MARKOV_DECLINES)) +def test_k3_markov_declines(markov_kernel, mutate): + case = _markov_case() + mutate(case) + assert not case.worker._keep_draft_logits_sharded(case.drafter, case.spec, NUM_GENS) + + +def test_k3_markov_declines_without_mnnvl(markov_kernel, mnnvl): + mnnvl.clear() + case = _markov_case() + assert not case.worker._keep_draft_logits_sharded(case.drafter, case.spec, NUM_GENS) + + +def test_k3_markov_declines_a_shard_it_does_not_split(markov_kernel, monkeypatch): + monkeypatch.setattr(markov_op, "pick_grid", lambda *args: 0) + case = _markov_case() + assert not case.worker._keep_draft_logits_sharded(case.drafter, case.spec, NUM_GENS) + + +def test_dspark_draft_logits_are_the_processors_when_off(): + worker = _worker(DSparkWorker, mapping=TP4) + gathered = torch.zeros(NUM_GENS * K, VOCAB) + calls = [] + drafter = SimpleNamespace( + lm_head=object(), logits_processor=lambda *args: calls.append(args) or gathered + ) + rows, attn = torch.zeros(NUM_GENS * K, HIDDEN), object() + + assert ( + worker._draft_block_logits(drafter, rows, attn, SimpleNamespace(runtime_draft_len=K)) + is gathered + ) + assert not worker._k3_sharded_block_logits + assert len(calls) == 1 + assert ( + calls[0][0] is rows + and calls[0][1] is drafter.lm_head + and calls[0][2] is attn + and calls[0][3] is True + ) + + +def test_dspark_draft_logits_stay_sharded(markov_kernel): + case = _markov_case() + shard = torch.ones(NUM_GENS * K, SHARD, dtype=torch.bfloat16) + case.drafter.logits_processor = _processor(shard) + rows = torch.zeros(NUM_GENS * K, HIDDEN, dtype=torch.bfloat16) + + assert case.worker._draft_block_logits(case.drafter, rows, None, case.spec) is shard + assert case.worker._k3_sharded_block_logits + calls = case.drafter.logits_processor.shard_calls + assert len(calls) == 1 and calls[0][0] is rows and calls[0][1] is case.drafter.lm_head + assert not case.drafter.logits_processor.forward_calls + + +def test_sharded_block_logits_run_the_markov_kernel(): + """The flag from ``_draft_block_logits`` sends this step's chain to ``k3_markov`` with this rank's slice, once.""" + case = _markov_case() + worker, drafter = case.worker, case.drafter + corrected = torch.zeros(NUM_GENS, K, SHARD) + chained = [] + worker._k3_markov_chain = lambda draft_model, logits, first, vocab_slice: ( + chained.append((logits, first, vocab_slice)) or corrected + ) + logits = torch.zeros(NUM_GENS, K, SHARD, dtype=torch.bfloat16) + first = torch.zeros(NUM_GENS, dtype=torch.long) + vocab_slice = slice(RANK_IN_TP * SHARD, (RANK_IN_TP + 1) * SHARD) + + worker._k3_sharded_block_logits = True + assert worker._apply_dspark_markov_bias(drafter, logits, first, case.spec) is corrected + assert len(chained) == 1 and chained[0][0] is logits and chained[0][1] is first + assert chained[0][2] == vocab_slice + assert not worker._k3_sharded_block_logits and not drafter.chain_calls + + # The flag is spent: the next chain is the unfused one. + worker._apply_dspark_markov_bias(drafter, logits, first, case.spec) + assert len(chained) == 1 and len(drafter.chain_calls) == 1 + assert drafter.chain_calls[0][2] == vocab_slice + + +def test_dspark_acceptance_is_kept_only_under_k3_decode(monkeypatch): + accepted = torch.zeros(NUM_GENS, K + 1, dtype=torch.int32) + num_accepted = torch.ones(NUM_GENS, dtype=torch.int32) + monkeypatch.setattr( + SpecWorkerBase, + "sample_and_accept_draft_tokens", + lambda self, *args: (accepted, num_accepted), + ) + attn, spec = SimpleNamespace(num_contexts=0), object() + + off = _worker(DSparkWorker) + result = off.sample_and_accept_draft_tokens(None, attn, spec) + assert result[0] is accepted and result[1] is num_accepted + assert getattr(off, "_k3_acceptance", None) is None + + on = _worker(DSparkWorker, k3_decode=True) + on.sample_and_accept_draft_tokens(None, attn, spec) + kept = on._k3_acceptance + assert kept[0] is accepted and kept[1] is num_accepted and kept[2] == 0 + assert kept[3] is spec and kept[4] is attn + + +@pytest.mark.parametrize("pending", [True, False], ids=["a pending KV rewind", "no pending rewind"]) +def test_k3_markov_chain_drafts_and_next_new_tokens(monkeypatch, pending): + """``k3_markov`` gets this rank's slice and the step's acceptance, folds a pending KV-length rewind, and its + tokens and next_new_tokens are the step's drafts and next inputs.""" + worker = _worker(DSparkWorker, k3_decode=True, mapping=TP4) + accepted = torch.zeros(NUM_GENS, K + 1, dtype=torch.int32) + num_accepted = torch.ones(NUM_GENS, dtype=torch.int32) + spec = SimpleNamespace( + batch_indices_cuda=torch.arange(4, dtype=torch.int32), wants_advanced_draft_sampling=False + ) + attn = SimpleNamespace(num_contexts=0, kv_lens_cuda=torch.zeros(4, dtype=torch.int32)) + worker._on_acceptance(accepted, num_accepted, attn, spec) + rewind = torch.zeros(NUM_GENS, dtype=torch.int32) + worker._kv_rewind_amount, worker._kv_rewind_pending = rewind, pending + worker._kv_rewind_nc, worker._kv_rewind_bs = 0, NUM_GENS + outputs = ( + torch.zeros(NUM_GENS, K, SHARD), + torch.zeros(NUM_GENS, K, dtype=torch.int32), + torch.zeros(NUM_GENS, K + 1, dtype=torch.int32), + ) + calls = [] + monkeypatch.setattr( + markov_op, "markov_chain", lambda *args, **kwargs: calls.append((args, kwargs)) or outputs + ) + drafter = _markov_drafter() + logits = torch.zeros(NUM_GENS, K, SHARD, dtype=torch.bfloat16) + first = torch.zeros(NUM_GENS, dtype=torch.long) + vocab_slice = slice(RANK_IN_TP * SHARD, (RANK_IN_TP + 1) * SHARD) + + corrected = worker._k3_markov_chain(drafter, logits, first, vocab_slice) + + assert corrected is outputs[0] + ((args, kwargs),) = calls + assert ( + args[0] is TP4 and args[1] is logits and args[2] is first and args[3] is drafter.markov_w1 + ) + assert torch.equal(args[4], drafter.markov_w2[vocab_slice]) and args[5] == vocab_slice.start + assert args[6] is accepted and torch.equal(args[7], num_accepted) + assert torch.equal(args[8], spec.batch_indices_cuda[:NUM_GENS]) + if pending: + assert kwargs["kv_lens"] is attn.kv_lens_cuda and kwargs["rewind"] is rewind + assert worker._kv_rewind_amount is None and not worker._kv_rewind_pending + else: + assert kwargs["kv_lens"] is None and kwargs["rewind"] is None + assert worker._kv_rewind_amount is rewind + assert kwargs["rewind_first"] == 0 + + drafts = worker.sample_draft_tokens(corrected, spec, NUM_GENS, num_contexts=0) + assert drafts is outputs[1] + next_new = worker._prepare_next_new_tokens( + accepted, drafts, spec.batch_indices_cuda, NUM_GENS, num_accepted + ) + assert next_new is outputs[2] + + +def test_dspark_drafts_from_the_base_sampler_without_the_kernel(monkeypatch): + sampled = torch.zeros(NUM_GENS, K, dtype=torch.int32) + assembled = torch.zeros(NUM_GENS, K + 1, dtype=torch.int32) + monkeypatch.setattr( + SpecWorkerBase, "sample_draft_tokens", lambda self, *args, **kwargs: sampled + ) + monkeypatch.setattr(SpecWorkerBase, "_prepare_next_new_tokens", lambda self, *args: assembled) + worker = _worker(DSparkWorker) + spec = SimpleNamespace(wants_advanced_draft_sampling=False) + + assert worker.sample_draft_tokens(torch.zeros(NUM_GENS, K, VOCAB), spec, NUM_GENS) is sampled + assert worker._prepare_next_new_tokens(None, sampled, None, NUM_GENS, None) is assembled From e85d7a6a5f98c1b388916a672ebd861543263a2e Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 11:03:15 -0700 Subject: [PATCH 129/161] [None][feat] modeling_v2 Kimi K3 tp16_moetp4ep4: the speculative worker's decode kernels - post_load_weights turns the DFlash / DSpark worker's k3_decode on (_gate_spec_worker_kernels) where this target's decode path runs: every attention all-reduce over MNNVL (the TP group's K3DecodeComm is built) and the LM head on gemm/k3_head_gemv. Off otherwise, and without such a worker nothing is set. - K3LogitsProcessor.lm_head_shard: this rank's vocabulary shard of the head kernel's logits without the gather (K3DecodeGemvs.lm_head_logits with gather False), or None where the kernel does not take the rows. The worker's vocabulary-sharded target and draft logits come from it, so they hold the values of the logits the processor gathers. The tp16_moetp16ep1 copy does not carry this change yet; until it does, test_modeling_v2_kimi_k3_drift fails on decode_gemv.py and modeling.py. test_kimi_k3_spec_worker_gate.py (CPU, fakes): the worker's k3_decode is on with the MNNVL decode state and the head kernel, off without either; the processor hands out the kernel's shard; gather False skips the all-gather. Signed-off-by: Vasanth Sabavat --- .../decode_gemv.py | 19 ++- .../modeling.py | 26 +++- .../test_kimi_k3_spec_worker_gate.py | 143 ++++++++++++++++++ 3 files changed, 183 insertions(+), 5 deletions(-) create mode 100644 tests/unittest/_torch/speculative/hw_agnostic/test_kimi_k3_spec_worker_gate.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py index 33fbbac23ca7..0729fe05b91e 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_gemv.py @@ -326,9 +326,12 @@ def _project( self._ran.add(key) return y.view(*x.shape[:-1], spec.n) - def lm_head_logits(self, rows: torch.Tensor, lm_head: nn.Module) -> Optional[torch.Tensor]: + def lm_head_logits( + self, rows: torch.Tensor, lm_head: nn.Module, gather: bool = True + ) -> Optional[torch.Tensor]: """``lm_head(rows)``, the gathered bf16 logits ``[M, vocab]``, with this rank's shard on - ``gemm/k3_head_gemv``; None where it does not take the call (more than `MAX_ROWS` rows, another head, ...).""" + ``gemm/k3_head_gemv`` (``gather`` False: the shard alone, ``[M, vocab / TP]``); None where it does not take + the call (more than `MAX_ROWS` rows, another head, ...).""" workspace = self.head_workspace if workspace is None or not _head_takes_module(lm_head): return None @@ -345,13 +348,13 @@ def lm_head_logits(self, rows: torch.Tensor, lm_head: nn.Module) -> Optional[tor if _capturing() and ("lm_head",) not in self._ran: return None group = lm_head.mapping.tp_group - if len(group) > 1 and mpi_disabled(): + if gather and len(group) > 1 and mpi_disabled(): return None x = _dense_rows(rows) if not _head_op.supports(x, weight): return None local = k3_head_gemv(x, weight, workspace) - if len(group) == 1: + if not gather or len(group) == 1: return local gathered = allgather(local, None, group) return concat(list(split(gathered, rows.shape[0], dim=0)), dim=-1) @@ -454,3 +457,11 @@ def forward( ) -> torch.Tensor: head = lm_head if self.gemvs is None else _K3Head(self.gemvs, lm_head) return self.stock.forward(hidden_states, head, attn_metadata, return_context_logits) + + def lm_head_shard(self, rows: torch.Tensor, lm_head: nn.Module) -> Optional[torch.Tensor]: + """This rank's bf16 vocabulary shard of ``lm_head(rows)`` on ``gemm/k3_head_gemv``, without the gather, or + None where ``gemvs`` does not take the rows. A speculative worker that keeps its logits vocabulary-sharded + computes them here, so they hold the values of the logits this processor gathers.""" + if self.gemvs is None: + return None + return self.gemvs.lm_head_logits(rows, lm_head, gather=False) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index 36cf98f2eac2..d3d7a5f0d739 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -2872,7 +2872,7 @@ def post_load_weights(self) -> None: attention all-reduce runs over MNNVL, the TP group's collective state (``K3DecodeComm``, collective: every rank builds it here) handed to every layer, with whether its sandwich takes the layer's o_proj, the decode path's one-shot ceiling on every stock MNNVL all-reduce (``use_decode_one_shot``), and the MoE decode path - (``_build_decode_moe``).""" + (``_build_decode_moe``); then the speculative worker's decode kernels (``_gate_spec_worker_kernels``).""" super().post_load_weights() kda = [layer.linear_attn for layer in self.model.layers if layer.is_kda] mla = [layer.self_attn.mixer for layer in self.model.layers if not layer.is_kda] @@ -2915,6 +2915,7 @@ def post_load_weights(self) -> None: ) + f"; MoE on k3_moe_front, k3_moe and the row-parallel tail ({moe_layers} layers)" ) + self._gate_spec_worker_kernels(comm) def _build_decode_moe(self, comm: _decode_comm.K3DecodeComm) -> int: """The MoE decode path (``decode_moe.py``) on every MoE layer it takes: the shared state (collective: the @@ -2956,6 +2957,29 @@ def _build_decode_moe(self, comm: _decode_comm.K3DecodeComm) -> int: comm.compile_tail(takes[0].moe_hidden_size, first.shared_cols, first.tail_weight) return len(takes) + def _gate_spec_worker_kernels(self, comm: Optional[_decode_comm.K3DecodeComm]) -> bool: + """Turn the DFlash / DSpark worker's Kimi K3 decode kernels (its ``k3_decode``: ``trtllm::k3_spec_accept``, + ``k3_ctx_kv`` and ``k3_markov``, on target and draft logits kept vocabulary-sharded) on where this target's + decode path runs, off elsewhere. On needs the TP group's collective state over MNNVL (``comm``) and the LM + head on ``gemm/k3_head_gemv``; the worker still checks each step's own conditions. Returns the setting (False + without such a worker).""" + worker = getattr(self, "spec_worker", None) + if not hasattr(worker, "k3_decode"): + return False + gemvs = self.model.decode_gemvs + worker.k3_decode = ( + comm is not None and gemvs is not None and gemvs.head_workspace is not None + ) + logger.info( + f"Kimi K3 decode kernels: {type(worker).__name__} k3_spec_accept, k3_ctx_kv and k3_markov " + + ( + "on" + if worker.k3_decode + else "off (no MNNVL decode state or no k3_head_gemv LM head)" + ) + ) + return worker.k3_decode + def _check_step_contract(self, attn_metadata: AttentionMetadata) -> None: """First-forward checks of the engine surface and the per-engine settings.""" objects = { diff --git a/tests/unittest/_torch/speculative/hw_agnostic/test_kimi_k3_spec_worker_gate.py b/tests/unittest/_torch/speculative/hw_agnostic/test_kimi_k3_spec_worker_gate.py new file mode 100644 index 000000000000..8ee9cc4373e4 --- /dev/null +++ b/tests/unittest/_torch/speculative/hw_agnostic/test_kimi_k3_spec_worker_gate.py @@ -0,0 +1,143 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""The Kimi K3 target ``kimi_k3_mxfp4__sm_100__tp16_moetp4ep4`` and its speculative worker's decode kernels (host-side, +fakes only). + +* ``_gate_spec_worker_kernels`` turns the DFlash / DSpark worker's ``k3_decode`` on only alongside the target's decode + path: the TP group's MNNVL decode state and the LM head on ``gemm/k3_head_gemv``. +* ``K3LogitsProcessor.lm_head_shard`` hands the worker this rank's vocabulary shard of the head kernel's logits, + without the gather. +""" + +from types import SimpleNamespace + +import pytest +import torch +from torch import nn + +from tensorrt_llm._torch._experimental.modeling_v2.models.kimi_k3_vl.kimi_k3_mxfp4__sm_100__tp16_moetp4ep4 import ( # noqa: E501 + decode_gemv, +) +from tensorrt_llm._torch._experimental.modeling_v2.models.kimi_k3_vl.kimi_k3_mxfp4__sm_100__tp16_moetp4ep4.modeling import ( # noqa: E501 + ModelingV2KimiK3Mxfp4Sm100Tp16Moetp4ep4 as Target, +) +from tensorrt_llm._torch.speculative.dflash import DFlashWorker +from tensorrt_llm._torch.speculative.dspark import DSparkWorker + +pytestmark = pytest.mark.cpu_only + + +def _worker(cls): + """A worker without ``__init__`` (it needs flashinfer and a drafter).""" + worker = cls.__new__(cls) + nn.Module.__init__(worker) + return worker + + +def _target(worker, gemvs=True, head_workspace=True): + """The target fields the gate reads: its speculative worker and the decode GEMVs' state.""" + state = SimpleNamespace(head_workspace=object() if head_workspace else None) if gemvs else None + return SimpleNamespace(spec_worker=worker, model=SimpleNamespace(decode_gemvs=state)) + + +@pytest.mark.parametrize("cls", [DFlashWorker, DSparkWorker]) +def test_worker_kernels_on_with_the_targets_decode_path(cls): + worker = _worker(cls) + assert Target._gate_spec_worker_kernels(_target(worker), comm=object()) + assert worker.k3_decode is True + assert cls.k3_decode is False + + +@pytest.mark.parametrize( + "comm,gemvs,head_workspace", + [(None, True, True), (object(), True, False), (object(), False, False)], + ids=[ + "an attention all-reduce not over MNNVL", + "no k3_head_gemv LM head", + "no decode GEMV state", + ], +) +def test_worker_kernels_off_without_the_targets_decode_path(comm, gemvs, head_workspace): + worker = _worker(DSparkWorker) + worker.k3_decode = True + assert not Target._gate_spec_worker_kernels(_target(worker, gemvs, head_workspace), comm) + assert worker.k3_decode is False + + +@pytest.mark.parametrize( + "worker", + [None, SimpleNamespace()], + ids=["no speculative worker", "a worker without the kernels"], +) +def test_no_worker_kernels_to_gate(worker): + assert not Target._gate_spec_worker_kernels(_target(worker), comm=object()) + assert not hasattr(worker, "k3_decode") + + +def test_logits_processor_hands_out_the_head_shard(): + processor = decode_gemv.K3LogitsProcessor(SimpleNamespace()) + rows, head = torch.zeros(2, 8, dtype=torch.bfloat16), object() + assert processor.lm_head_shard(rows, head) is None # before the decode GEMVs' state is built + + shard = torch.ones(2, 4, dtype=torch.bfloat16) + calls = [] + processor.gemvs = SimpleNamespace( + lm_head_logits=lambda r, h, gather=True: calls.append((r, h, gather)) or shard + ) + assert processor.lm_head_shard(rows, head) is shard + assert len(calls) == 1 and calls[0][0] is rows and calls[0][1] is head and calls[0][2] is False + + +def _head(rows_per_rank=16, k=8, group=(0, 1, 2, 3)): + """A plain vocabulary-parallel bf16 head the head kernel reproduces.""" + return SimpleNamespace( + weight=torch.zeros(rows_per_rank, k, dtype=torch.bfloat16), + mapping=SimpleNamespace(enable_attention_dp=False, tp_group=list(group)), + tp_mode=SimpleNamespace(name="COLUMN"), + gather_output=True, + gather_output_sizes=None, + padding_size=0, + bias=None, + has_any_quant=False, + ) + + +def test_head_kernel_shard_skips_the_gather(monkeypatch): + """``gather`` False: this rank's ``k3_head_gemv`` output itself, with no gather (and no need for the MPI one).""" + head = _head() + workspace = SimpleNamespace(n_out=16, k_in=8, partials=torch.empty(1)) + gemvs = decode_gemv.K3DecodeGemvs(workspace) + local = torch.ones(2, 16, dtype=torch.bfloat16) + kernel_calls = [] + monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: False) + monkeypatch.setattr(decode_gemv._head_op, "supports", lambda x, w: True) + monkeypatch.setattr( + decode_gemv, "k3_head_gemv", lambda x, w, ws: kernel_calls.append((x, w, ws)) or local + ) + monkeypatch.setattr( + decode_gemv, "allgather", lambda *args: pytest.fail("the shard was gathered") + ) + monkeypatch.setattr(decode_gemv, "mpi_disabled", lambda: True) + rows = torch.zeros(2, 8, dtype=torch.bfloat16) + + assert gemvs.lm_head_logits(rows, head, gather=False) is local + assert len(kernel_calls) == 1 + assert kernel_calls[0][0] is rows and kernel_calls[0][1] is head.weight + assert kernel_calls[0][2] is workspace + # The gathered logits of a TP group need the MPI all-gather. + assert gemvs.lm_head_logits(rows, head) is None + assert len(kernel_calls) == 1 + # Above the kernel's rows the shard is declined too; the worker then runs the head's own GEMM. + assert gemvs.lm_head_logits(torch.zeros(9, 8, dtype=torch.bfloat16), head, gather=False) is None From 41667882ba86174549f62d325fdad428bb7d270b Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 11:10:47 -0700 Subject: [PATCH 130/161] [None][doc] modeling_v2 Kimi K3 tp16_moetp4ep4: the spec-worker gate's preconditions _gate_spec_worker_kernels turns the DFlash / DSpark worker's decode kernels on only when every input their path needs exists: the TP group's collective state over MNNVL and the LM head's gemm/k3_head_gemv workspace. K3LogitsProcessor.lm_head_shard needs that workspace to produce the worker's vocabulary-sharded logits, so the workspace is a precondition of the path, not a policy choice. The docstring now says so; the code is unchanged. Signed-off-by: Vasanth Sabavat --- .../kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index d3d7a5f0d739..4a6fe35f8ec7 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -2959,10 +2959,11 @@ def _build_decode_moe(self, comm: _decode_comm.K3DecodeComm) -> int: def _gate_spec_worker_kernels(self, comm: Optional[_decode_comm.K3DecodeComm]) -> bool: """Turn the DFlash / DSpark worker's Kimi K3 decode kernels (its ``k3_decode``: ``trtllm::k3_spec_accept``, - ``k3_ctx_kv`` and ``k3_markov``, on target and draft logits kept vocabulary-sharded) on where this target's - decode path runs, off elsewhere. On needs the TP group's collective state over MNNVL (``comm``) and the LM - head on ``gemm/k3_head_gemv``; the worker still checks each step's own conditions. Returns the setting (False - without such a worker).""" + ``k3_ctx_kv`` and ``k3_markov``, on target and draft logits kept vocabulary-sharded) on only when every input + that path needs exists: the TP group's collective state over MNNVL (``comm``) and the LM head's + ``gemm/k3_head_gemv`` workspace. ``K3LogitsProcessor.lm_head_shard`` needs that workspace to produce the + vocabulary shard, so the workspace is a precondition of the path, not a policy choice. The worker still checks + each step's own conditions. Returns the setting (False without such a worker).""" worker = getattr(self, "spec_worker", None) if not hasattr(worker, "k3_decode"): return False From 2b465f63f3c02274f86d915037f719101b864351 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 11:47:16 -0700 Subject: [PATCH 131/161] [None][chore] modeling_v2 Kimi K3 tp16_moetp16ep1: carry the speculative worker gate into the copy The tp16_moetp4ep4 target turns the DFlash / DSpark worker's k3_decode on where its decode path runs (_gate_spec_worker_kernels, at the end of post_load_weights), and its LM head gains a vocabulary-sharded form (K3DecodeGemvs.lm_head_logits with gather False, and K3LogitsProcessor.lm_head_shard). The copy here takes both: decode_gemv.py byte for byte, and the changes to modeling.py outside route B's blocks, the gate's docstring on its preconditions included. This target decodes without speculation, so it has no speculative worker: _gate_spec_worker_kernels returns at once and logs nothing, and nothing calls lm_head_shard. lm_head_logits still gathers by default, so the LM head this target runs is unchanged. Signed-off-by: Vasanth Sabavat --- .../decode_gemv.py | 19 ++++++++++--- .../modeling.py | 27 ++++++++++++++++++- 2 files changed, 41 insertions(+), 5 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_gemv.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_gemv.py index 33fbbac23ca7..0729fe05b91e 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_gemv.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_gemv.py @@ -326,9 +326,12 @@ def _project( self._ran.add(key) return y.view(*x.shape[:-1], spec.n) - def lm_head_logits(self, rows: torch.Tensor, lm_head: nn.Module) -> Optional[torch.Tensor]: + def lm_head_logits( + self, rows: torch.Tensor, lm_head: nn.Module, gather: bool = True + ) -> Optional[torch.Tensor]: """``lm_head(rows)``, the gathered bf16 logits ``[M, vocab]``, with this rank's shard on - ``gemm/k3_head_gemv``; None where it does not take the call (more than `MAX_ROWS` rows, another head, ...).""" + ``gemm/k3_head_gemv`` (``gather`` False: the shard alone, ``[M, vocab / TP]``); None where it does not take + the call (more than `MAX_ROWS` rows, another head, ...).""" workspace = self.head_workspace if workspace is None or not _head_takes_module(lm_head): return None @@ -345,13 +348,13 @@ def lm_head_logits(self, rows: torch.Tensor, lm_head: nn.Module) -> Optional[tor if _capturing() and ("lm_head",) not in self._ran: return None group = lm_head.mapping.tp_group - if len(group) > 1 and mpi_disabled(): + if gather and len(group) > 1 and mpi_disabled(): return None x = _dense_rows(rows) if not _head_op.supports(x, weight): return None local = k3_head_gemv(x, weight, workspace) - if len(group) == 1: + if not gather or len(group) == 1: return local gathered = allgather(local, None, group) return concat(list(split(gathered, rows.shape[0], dim=0)), dim=-1) @@ -454,3 +457,11 @@ def forward( ) -> torch.Tensor: head = lm_head if self.gemvs is None else _K3Head(self.gemvs, lm_head) return self.stock.forward(hidden_states, head, attn_metadata, return_context_logits) + + def lm_head_shard(self, rows: torch.Tensor, lm_head: nn.Module) -> Optional[torch.Tensor]: + """This rank's bf16 vocabulary shard of ``lm_head(rows)`` on ``gemm/k3_head_gemv``, without the gather, or + None where ``gemvs`` does not take the rows. A speculative worker that keeps its logits vocabulary-sharded + computes them here, so they hold the values of the logits this processor gathers.""" + if self.gemvs is None: + return None + return self.gemvs.lm_head_logits(rows, lm_head, gather=False) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py index 75f43fcec212..da7d9bb0d78c 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py @@ -2790,7 +2790,7 @@ def post_load_weights(self) -> None: attention all-reduce runs over MNNVL, the TP group's collective state (``K3DecodeComm``, collective: every rank builds it here) handed to every layer, with whether its sandwich takes the layer's o_proj, the decode path's one-shot ceiling on every stock MNNVL all-reduce (``use_decode_one_shot``), and the MoE decode path - (``_build_decode_moe``).""" + (``_build_decode_moe``); then the speculative worker's decode kernels (``_gate_spec_worker_kernels``).""" super().post_load_weights() kda = [layer.linear_attn for layer in self.model.layers if layer.is_kda] mla = [layer.self_attn.mixer for layer in self.model.layers if not layer.is_kda] @@ -2838,6 +2838,7 @@ def post_load_weights(self) -> None: ) + f"; MoE on k3_moe_front, k3_moe and the row-parallel tail ({moe_layers} layers)" ) + self._gate_spec_worker_kernels(comm) def _build_decode_moe(self, comm: _decode_comm.K3DecodeComm) -> int: """The MoE decode path (``decode_moe.py``) on every MoE layer it takes: the shared state (collective: the @@ -2879,6 +2880,30 @@ def _build_decode_moe(self, comm: _decode_comm.K3DecodeComm) -> int: comm.compile_tail(takes[0].moe_hidden_size, first.shared_cols, first.tail_weight) return len(takes) + def _gate_spec_worker_kernels(self, comm: Optional[_decode_comm.K3DecodeComm]) -> bool: + """Turn the DFlash / DSpark worker's Kimi K3 decode kernels (its ``k3_decode``: ``trtllm::k3_spec_accept``, + ``k3_ctx_kv`` and ``k3_markov``, on target and draft logits kept vocabulary-sharded) on only when every input + that path needs exists: the TP group's collective state over MNNVL (``comm``) and the LM head's + ``gemm/k3_head_gemv`` workspace. ``K3LogitsProcessor.lm_head_shard`` needs that workspace to produce the + vocabulary shard, so the workspace is a precondition of the path, not a policy choice. The worker still checks + each step's own conditions. Returns the setting (False without such a worker).""" + worker = getattr(self, "spec_worker", None) + if not hasattr(worker, "k3_decode"): + return False + gemvs = self.model.decode_gemvs + worker.k3_decode = ( + comm is not None and gemvs is not None and gemvs.head_workspace is not None + ) + logger.info( + f"Kimi K3 decode kernels: {type(worker).__name__} k3_spec_accept, k3_ctx_kv and k3_markov " + + ( + "on" + if worker.k3_decode + else "off (no MNNVL decode state or no k3_head_gemv LM head)" + ) + ) + return worker.k3_decode + def _check_step_contract(self, attn_metadata: AttentionMetadata) -> None: """First-forward checks of the engine surface and the per-engine settings.""" objects = { From 32d419b8d8c7b6ebddd9e1e3e9bb53d632e45f1a Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 11:11:25 -0700 Subject: [PATCH 132/161] [None][feat] DFlash metadata: capture_view, a captured layer's slot of the capture buffer DFlashSpecMetadata.capture_view(layer_id, num_tokens) returns the strided view of the capture buffer that holds layer_id's tap for the first num_tokens rows, or None when the layer is not captured or there is no buffer. A kernel that produces the tap can write it there directly; what it writes is what get_hidden_states hands the drafter. The metadata of every CUDA-graph bucket shares the buffer, so it hands out the same slot. Nothing calls it yet. test_dflash_capture_view.py (CPU): the view aliases the layer's columns of the buffer and a write through it lands in that layer's slot only; a graph bucket's metadata hands out the same slot; no view for a layer that is not captured or without a capture buffer. Signed-off-by: Vasanth Sabavat --- tensorrt_llm/_torch/speculative/dflash.py | 12 +++ .../hw_agnostic/test_dflash_capture_view.py | 92 +++++++++++++++++++ 2 files changed, 104 insertions(+) create mode 100644 tests/unittest/_torch/speculative/hw_agnostic/test_dflash_capture_view.py diff --git a/tensorrt_llm/_torch/speculative/dflash.py b/tensorrt_llm/_torch/speculative/dflash.py index 746c5d36f0f2..70fede4fbcb5 100644 --- a/tensorrt_llm/_torch/speculative/dflash.py +++ b/tensorrt_llm/_torch/speculative/dflash.py @@ -652,6 +652,18 @@ def maybe_capture_hidden_states( (i + 1) * self.hidden_size, ) + def capture_view(self, layer_id: int, num_tokens: int) -> Optional[torch.Tensor]: + """``layer_id``'s slot of the capture buffer for ``num_tokens`` rows (a strided view a kernel can write the + tap into), or None when the layer is not captured.""" + if self.captured_hidden_states is None: + return None + i = self._layer_to_idx.get(layer_id) + if i is None: + return None + return self.captured_hidden_states[ + :num_tokens, i * self.hidden_size : (i + 1) * self.hidden_size + ] + def get_hidden_states(self, num_tokens: int) -> Optional[torch.Tensor]: """Get captured hidden states (all layers concatenated).""" if self.captured_hidden_states is None: diff --git a/tests/unittest/_torch/speculative/hw_agnostic/test_dflash_capture_view.py b/tests/unittest/_torch/speculative/hw_agnostic/test_dflash_capture_view.py new file mode 100644 index 000000000000..40a9f9bad975 --- /dev/null +++ b/tests/unittest/_torch/speculative/hw_agnostic/test_dflash_capture_view.py @@ -0,0 +1,92 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``DFlashSpecMetadata.capture_view``: a captured layer's slot of the capture buffer, which a kernel can write the tap +into (host-side). The view shares the buffer's storage, so what is written through it is what ``get_hidden_states`` +hands the drafter, in that layer's columns only; a layer that is not captured, or metadata without a capture buffer, +has no view.""" + +import pytest +import torch + +from tensorrt_llm._torch.speculative.dflash import DFlashSpecMetadata +from tensorrt_llm._torch.speculative.interface import SpeculativeDecodingMode + +pytestmark = pytest.mark.cpu_only + +HIDDEN, MAX_TOKENS = 16, 32 + + +@pytest.fixture(autouse=True) +def cpu_buffers(monkeypatch): + """The metadata's buffers on the host.""" + empty = torch.empty + + def cpu_empty(*args, **kwargs): + kwargs["device"] = "cpu" + return empty(*args, **kwargs) + + monkeypatch.setattr(torch, "empty", cpu_empty) + + +def _metadata(layers=(3, 1)): + return DFlashSpecMetadata( + max_num_requests=8, + max_draft_len=3, + max_total_draft_tokens=3, + spec_dec_mode=SpeculativeDecodingMode.DFLASH, + layers_to_capture=None if layers is None else list(layers), + hidden_size=HIDDEN, + max_num_tokens=MAX_TOKENS, + dtype=torch.bfloat16, + ) + + +def _zero(t): + return torch.equal(t, torch.zeros_like(t)) + + +@pytest.mark.parametrize("layer,slot", [(1, 0), (3, 1)]) +def test_capture_view_is_the_layers_slot(layer, slot): + metadata = _metadata() # layers 3 and 1: the buffer holds layer 1's slot first + buffer = metadata.captured_hidden_states + buffer.zero_() + num_tokens = 5 + + view = metadata.capture_view(layer, num_tokens) + + assert view.shape == (num_tokens, HIDDEN) and view.dtype == buffer.dtype + assert view.untyped_storage().data_ptr() == buffer.untyped_storage().data_ptr() + assert view.data_ptr() == buffer[0, slot * HIDDEN :].data_ptr() + assert view.stride() == buffer.stride() + tap = torch.randn(num_tokens, HIDDEN).to(torch.bfloat16) + view.copy_(tap) + captured = metadata.get_hidden_states(num_tokens) + other = 1 - slot + assert torch.equal(captured[:, slot * HIDDEN : (slot + 1) * HIDDEN], tap) + assert _zero(captured[:, other * HIDDEN : (other + 1) * HIDDEN]) + assert _zero(buffer[num_tokens:]) + + +def test_cuda_graph_metadata_views_the_shared_buffer(): + metadata = _metadata() + graph = metadata.create_cuda_graph_metadata(4) + assert graph.capture_view(3, 4).data_ptr() == metadata.capture_view(3, 4).data_ptr() + + +def test_no_view_outside_the_captured_layers(): + assert _metadata().capture_view(2, 4) is None + uncaptured = _metadata(layers=None) + assert uncaptured.captured_hidden_states is None + assert uncaptured.capture_view(1, 4) is None From 14473977951c4f38f51d59afe52e3a87763311ff Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 11:13:45 -0700 Subject: [PATCH 133/161] [None][perf] DFlash: the draft block's rows as a view where its slots are one run The draft forward gathered the block-output rows that produce the draft logits on every step, building the slot ids each time. Those ids depend only on the step's shape, so _draft_block_hidden_states builds them once per shape, outside CUDA-graph capture, and keeps the decision. Where the ids are one run of rows (one request, or DSpark's shift_label convention with a block of K slots, e.g. T = K = 7, where each request's slots start where the previous one's end) the rows are a view of the block outputs and no gather runs. Other ids, e.g. plain DFlash's slots 1..K of a block of K + 1, keep the gather. A shape first seen under capture is gathered and not kept. test_dflash_draft_block_rows.py (CPU): for every (num_gens 1..8, K 2 or 7, block K or K + 1, shift_label) shape the rows equal the stock gather of the clamped slot ids, they are a view exactly where the ids are one run, and the decision is reused; DSpark's decode shapes are views and plain DFlash with several requests gathers; the ids are built once per shape; capture gathers an unseen shape without keeping it. Signed-off-by: Vasanth Sabavat --- tensorrt_llm/_torch/speculative/dflash.py | 47 +++++- .../test_dflash_draft_block_rows.py | 153 ++++++++++++++++++ 2 files changed, 193 insertions(+), 7 deletions(-) create mode 100644 tests/unittest/_torch/speculative/hw_agnostic/test_dflash_draft_block_rows.py diff --git a/tensorrt_llm/_torch/speculative/dflash.py b/tensorrt_llm/_torch/speculative/dflash.py index 70fede4fbcb5..032c7aa920aa 100644 --- a/tensorrt_llm/_torch/speculative/dflash.py +++ b/tensorrt_llm/_torch/speculative/dflash.py @@ -2192,13 +2192,9 @@ def _forward_impl( # Which block slots carry them is a drafter-family convention, # resolved through _draft_slot_ids. block_size = self._compute_block_size - gen_gather_ids = self._draft_slot_ids(draft_model, num_gens, block_size, K) - # Shields only the last request: at block_size == K with - # shift_label off, slots run 1..K, so every request reads the - # next one's slot 0 and the last overruns. Degrades, never raises. - gen_gather_ids = gen_gather_ids.clamp(max=hidden_states_out.shape[0] - 1) - - gen_hidden_states = hidden_states_out[gen_gather_ids] + gen_hidden_states = self._draft_block_hidden_states( + draft_model, hidden_states_out, num_gens, block_size, K + ) gen_logits = self._draft_block_logits( draft_model, gen_hidden_states, attn_metadata, spec_metadata ) @@ -2309,6 +2305,43 @@ def _draft_block_width(self, draft_model) -> int: """ return self.max_draft_len + 1 + def _draft_block_hidden_states( + self, + draft_model, + hidden_states_out: torch.Tensor, + num_gens: int, + block_size: int, + num_draft_tokens: int, + ) -> torch.Tensor: + """The block-output rows that produce the K draft logits per gen request (``_draft_slot_ids``). + + The slot ids depend only on the shapes, so they are built once per shape, outside CUDA-graph + capture, and cached. When they are one contiguous run of rows (a single request, or the + shift_label convention with a block of K slots, whose requests' slots follow each other) the + rows are returned as a view, without a gather; otherwise they are gathered. + """ + rows = hidden_states_out.shape[0] + key = (num_gens, block_size, num_draft_tokens, rows) + cache = self.__dict__.setdefault("_draft_block_rows", {}) + entry = cache.get(key) + if entry is None: + ids = self._draft_slot_ids(draft_model, num_gens, block_size, num_draft_tokens) + # Shields only the last request: at block_size == K with + # shift_label off, slots run 1..K, so every request reads the + # next one's slot 0 and the last overruns. Degrades, never raises. + ids = ids.clamp(max=rows - 1) + if torch.cuda.is_current_stream_capturing(): + return hidden_states_out[ids] + host = ids.tolist() + if host and host == list(range(host[0], host[0] + len(host))): + entry = host[0] + else: + entry = ids + cache[key] = entry + if isinstance(entry, int): + return hidden_states_out[entry : entry + num_gens * num_draft_tokens] + return hidden_states_out[entry] + def _draft_slot_ids( self, draft_model, num_gens: int, block_size: int, num_draft_tokens: int ) -> torch.Tensor: diff --git a/tests/unittest/_torch/speculative/hw_agnostic/test_dflash_draft_block_rows.py b/tests/unittest/_torch/speculative/hw_agnostic/test_dflash_draft_block_rows.py new file mode 100644 index 000000000000..f2f36360339f --- /dev/null +++ b/tests/unittest/_torch/speculative/hw_agnostic/test_dflash_draft_block_rows.py @@ -0,0 +1,153 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``DFlashWorker._draft_block_hidden_states``: the block-output rows that produce the draft logits (host-side). + +For every (num_gens, block, K, shift_label) shape the rows equal the stock gather of the clamped slot ids +(``dflash_draft_slot_ids``). Where those ids are one run of rows (one request, or DSpark's shift_label convention +with a block of K slots) the rows are a view of the block outputs; otherwise they are a gather. The ids are built +once per shape, outside CUDA-graph capture. +""" + +from types import SimpleNamespace + +import pytest +import torch +from torch import nn + +from tensorrt_llm._torch.speculative import dflash as dflash_module +from tensorrt_llm._torch.speculative import dspark as dspark_module +from tensorrt_llm._torch.speculative.dflash import DFlashWorker, dflash_draft_slot_ids +from tensorrt_llm._torch.speculative.dspark import DSparkWorker + +pytestmark = pytest.mark.cpu_only + +HIDDEN = 4 + + +@pytest.fixture(autouse=True) +def host_ids(monkeypatch): + """The workers' slot ids on the host, and no CUDA-graph capture.""" + + def on_host(num_gens, block_size, num_draft_tokens, shift_label, device="cuda"): + return dflash_draft_slot_ids(num_gens, block_size, num_draft_tokens, shift_label, "cpu") + + monkeypatch.setattr(dflash_module, "dflash_draft_slot_ids", on_host) + monkeypatch.setattr(dspark_module, "dflash_draft_slot_ids", on_host) + monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: False) + + +def _worker(cls=DSparkWorker): + """A worker without ``__init__`` (it needs flashinfer and a drafter).""" + worker = cls.__new__(cls) + nn.Module.__init__(worker) + return worker + + +def _block(num_gens, block): + """Block outputs whose rows are distinct.""" + rows = num_gens * block + return torch.arange(rows * HIDDEN, dtype=torch.float32).reshape(rows, HIDDEN) + + +def _stock(out, num_gens, block, k, shift_label): + """The gather the draft forward ran: every request's slot ids, the last one clamped to the block outputs.""" + ids = dflash_draft_slot_ids(num_gens, block, k, shift_label, "cpu").clamp(max=out.shape[0] - 1) + return out[ids] + + +def _one_run(num_gens, block, k, shift_label): + """Whether the slots are one run of rows: one request whose slots fit its block, or the shift_label + convention with a block of K slots (each request's slots start where the previous one's end).""" + first = 0 if shift_label else 1 + return first + k <= block and (num_gens == 1 or (shift_label and block == k)) + + +def _is_view(rows, out): + return rows.untyped_storage().data_ptr() == out.untyped_storage().data_ptr() + + +@pytest.mark.parametrize("shift_label", [True, False], ids=["shift_label", "dflash slots"]) +@pytest.mark.parametrize("block_extra", [0, 1], ids=["block K", "block K+1"]) +@pytest.mark.parametrize("k", [2, 7]) +@pytest.mark.parametrize("num_gens", range(1, 9)) +def test_rows_equal_the_stock_gather(num_gens, k, block_extra, shift_label): + block = k + block_extra + worker = _worker() + drafter = SimpleNamespace(_dspark_shift_label=shift_label) + out = _block(num_gens, block) + + rows = worker._draft_block_hidden_states(drafter, out, num_gens, block, k) + + assert torch.equal(rows, _stock(out, num_gens, block, k, shift_label)) + assert _is_view(rows, out) == _one_run(num_gens, block, k, shift_label) + if _is_view(rows, out): + first = 0 if shift_label else 1 + assert rows.data_ptr() == out[first].data_ptr() + # The cached decision gives the same rows on the next step. + again = worker._draft_block_hidden_states(drafter, out, num_gens, block, k) + assert torch.equal(again, rows) and _is_view(again, out) == _is_view(rows, out) + + +def test_dspark_decode_shape_is_a_view(): + """DSpark's shift_label drafter at a block of K = 7: every batch of the decode path reads its rows in place.""" + worker = _worker() + drafter = SimpleNamespace(_dspark_shift_label=True) + for num_gens in range(1, 9): + out = _block(num_gens, 7) + rows = worker._draft_block_hidden_states(drafter, out, num_gens, 7, 7) + assert _is_view(rows, out) and rows.shape == (num_gens * 7, HIDDEN) + assert torch.equal(rows, out) + + +def test_dflash_slots_of_several_requests_are_gathered(): + """Plain DFlash reads slots 1..K of each block of K + 1: a gap per request, so the rows are gathered.""" + worker = _worker(DFlashWorker) + out = _block(3, 8) + rows = worker._draft_block_hidden_states(object(), out, 3, 8, 7) + assert not _is_view(rows, out) + assert torch.equal(rows, _stock(out, 3, 8, 7, False)) + + +def test_slot_ids_are_built_once_per_shape(monkeypatch): + worker = _worker() + drafter = SimpleNamespace(_dspark_shift_label=True) + built = [] + real = DSparkWorker._draft_slot_ids + monkeypatch.setattr( + DSparkWorker, + "_draft_slot_ids", + lambda self, *args: built.append(args[1:]) or real(self, *args), + ) + for _ in range(3): + worker._draft_block_hidden_states(drafter, _block(2, 7), 2, 7, 7) + worker._draft_block_hidden_states(drafter, _block(4, 7), 4, 7, 7) + assert built == [(2, 7, 7), (4, 7, 7)] + + +def test_capture_gathers_a_shape_it_has_not_seen(monkeypatch): + """Under capture an unseen shape is gathered and not kept (building the run needs a host read); a shape seen + before capture keeps its view.""" + worker = _worker() + drafter = SimpleNamespace(_dspark_shift_label=True) + seen, unseen = _block(2, 7), _block(3, 7) + assert _is_view(worker._draft_block_hidden_states(drafter, seen, 2, 7, 7), seen) + + monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: True) + rows = worker._draft_block_hidden_states(drafter, unseen, 3, 7, 7) + assert not _is_view(rows, unseen) and torch.equal(rows, unseen) + assert _is_view(worker._draft_block_hidden_states(drafter, seen, 2, 7, 7), seen) + + monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: False) + assert _is_view(worker._draft_block_hidden_states(drafter, unseen, 3, 7, 7), unseen) From 33d35864f043b08876a3572c1f8490ac2c3e1267 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 10:39:39 -0700 Subject: [PATCH 134/161] [None][test] Kimi K3 DSpark decode kernel tests: list them in l0_b200 and l0_gb200_multi_gpus test_k3_spec_accept.py and test_k3_ctx_kv.py run on one sm_100 GPU and go in l0_b200.yml, next to the DSpark drafter attention tests. test_k3_spec_accept_sharded.py and test_k3_markov.py exchange across 4 ranks and go in l0_gb200_multi_gpus.yml, with the other Kimi K3 collective kernel tests. Signed-off-by: Vasanth Sabavat --- tests/integration/test_lists/test-db/l0_b200.yml | 4 ++++ tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml | 3 +++ 2 files changed, 7 insertions(+) diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index 991644a9f59e..6f437d362e9f 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -286,6 +286,10 @@ l0_b200: - unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn.py - unittest/_torch/modeling_v2/attention/test_modeling_v2_k3_drafter_attn_qknorm.py - unittest/_torch/modeling_v2/attention/test_modeling_v2_fused_qk_norm_rope.py + # Kimi K3 DSpark decode kernels (CuTe DSL, SM 100), one GPU: the acceptance against vocabulary-sharded target + # logits and the drafter's context K / V. Their multi-rank tests are in l0_gb200_multi_gpus.yml. + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_spec_accept.py + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctx_kv.py - unittest/_torch/thop/parallel TIMEOUT (90) - unittest/_torch/visual_gen/kernels/parallel - unittest/_torch/thop/serial diff --git a/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml b/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml index 56b07476585b..7b6e4fa63f17 100644 --- a/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml +++ b/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml @@ -58,6 +58,9 @@ l0_gb200_multi_gpus: - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_latent_reduce.py - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py + # Kimi K3 DSpark decode kernels' cross-rank exchanges (sm_100): the sharded acceptance and the Markov chain, 4 ranks + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_spec_accept_sharded.py + - unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_markov.py - unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_sandwich_oproj_op_matrix.py - unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_sandwich_tail_op_matrix.py - unittest/_torch/modeling_v2/comm/test_modeling_v2_k3_sandwich_plain_op_matrix.py From 5c193acae60fbf02d239ed12ca341f446afdc12c Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 11:41:38 -0700 Subject: [PATCH 135/161] [None][test] Kimi K3 Markov chain test: run its checks under pytest on 4 ranks test_k3_markov.py ran its checks under pytest only when pytest itself was started on two or more MPI ranks (srun -n W python3 -m pytest), so the single pytest process CI starts skipped every test. Its one pytest test now runs the file's report mode on 4 ranks of the node: a fresh interpreter starts mpirun -n 4 (pytest's own process has initialized MPI), under a deadline that kills the job's process group, and the test checks the exit code and the report's ALL PASS. The module skips with fewer than 4 SM100 GPUs. The checks and the script mode are unchanged. Signed-off-by: Vasanth Sabavat --- .../kimi_k3/test_k3_markov.py | 145 +++++++----------- 1 file changed, 52 insertions(+), 93 deletions(-) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_markov.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_markov.py index c276699af735..965d3a115f33 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_markov.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_markov.py @@ -58,14 +58,16 @@ W = 2 or 4 GPUs of one NVLink domain, every GPU of the node visible to every rank (rank r runs on GPU r % count). Without a mode the checks run, then the timing (unless ``--skip-perf``); the exit code is nonzero when a check fails. -The checks also run under pytest on W >= 2 ranks (``srun -n W --mpi=pmix python3 -m pytest -p no:cacheprovider -test_k3_markov.py``); with fewer ranks the module is skipped. +Under pytest (``python3 -m pytest test_k3_markov.py``, 4 GPUs visible) one test runs the checks on 4 ranks: this +file under a local ``mpirun -n 4``, started from a fresh interpreter, under a deadline; with fewer GPUs the module is +skipped. """ import argparse -import contextlib import os +import signal import statistics +import subprocess import sys import traceback import zlib @@ -118,12 +120,14 @@ def _world_size() -> int: return MPI.COMM_WORLD.Get_size() -pytestmark = [ - pytest.mark.skipif(not _sm100(), reason="needs SM100 (MNNVL multicast, TMA bulk copies, PDL)"), - pytest.mark.skipif( - _world_size() < 2, reason="needs >= 2 MPI ranks: srun -n W --mpi=pmix python3 -m pytest" - ), -] +# The ranks of the pytest run (one GB200 tray) and its deadline: a broken collective hangs rather than raising. +WORLD = 4 +DEADLINE_S = 3600 + +pytestmark = pytest.mark.skipif( + not _sm100() or torch.cuda.device_count() < WORLD, + reason=f"needs {WORLD} SM100 GPUs (MNNVL multicast, TMA bulk copies, PDL)", +) def case_seed(*key) -> int: @@ -867,93 +871,45 @@ def fused(i): # ---------------------------------------------------------------------------------------------------------------- -# pytest (srun -n W --mpi=pmix python3 -m pytest -p no:cacheprovider test_k3_markov.py): every rank runs the same -# tests in the same order, and every collective of a test happens before its assert. +# pytest (python3 -m pytest test_k3_markov.py): the checks on WORLD ranks of this node. # ---------------------------------------------------------------------------------------------------------------- -_state = {} - - -def group() -> Group: - """This process' rank of the group, built by the first test (its workspaces are allocated collectively).""" - if "group" not in _state: - _state["group"] = Group() - return _state["group"] - - -@contextlib.contextmanager -def collective(): - """Inference mode; an exception on one rank aborts the job (its peers would wait for it in a collective).""" - with torch.inference_mode(): - try: - yield - except Exception: - traceback.print_exc() - from mpi4py import MPI - - MPI.COMM_WORLD.Abort(1) - raise - - -def test_shard_support(): - with collective(): - g = group() - results = [(shard, *shard_support(g, shard, copies)) for shard, copies in g.shards()] - for shard, rejected, consistent, line in results: - assert consistent, line - assert not (shard == TP16_SHARD and rejected), line - - -@pytest.mark.parametrize("dtype", DTYPES, ids=[DTYPE_NAMES[d] for d in DTYPES]) -@pytest.mark.parametrize("block", BLOCKS) -@pytest.mark.parametrize("batch", BATCHES) -@pytest.mark.parametrize("shard_kind", ["tp16", "full"]) -def test_split(shard_kind, batch, block, dtype): - with collective(): - g = group() - shard, copies = g.shard(shard_kind) - if g.op.pick_grid(shard, block, batch) == 0: - pytest.skip(f"pick_grid rejects S = {shard} (see test_shard_support)") - row = random_case(g, shard, copies, batch, block, dtype) - assert row["ok"], row - - -@pytest.mark.parametrize( - "index", range(len(CRAFTED_SPLITS)), ids=[f"{b}x{k}" for b, k in CRAFTED_SPLITS] -) -@pytest.mark.parametrize("kind", list(CRAFTED)) -@pytest.mark.parametrize("shard_kind", ["tp16", "full"]) -def test_crafted(shard_kind, kind, index): - batch, block = CRAFTED_SPLITS[index] - with collective(): - g = group() - shard, copies = g.shard(shard_kind) - if g.op.pick_grid(shard, block, batch) == 0: - pytest.skip(f"pick_grid rejects S = {shard} (see test_shard_support)") - row = crafted_case(g, kind, shard, copies, batch, block, DTYPES[index % len(DTYPES)]) - assert row["ok"], row - - -@pytest.mark.parametrize("batch", MIXED_GENS) -@pytest.mark.parametrize("contexts", MIXED_CONTEXTS) -def test_mixed_step(contexts, batch): - with collective(): - row = mixed_step_case(group(), contexts, batch) - assert row["ok"], row - -def test_eager_interleave(): - with collective(): - row = eager_interleave(group()) - assert row["ok"], row +@pytest.mark.no_xdist +def test_k3_markov(): + """``report`` on WORLD ranks: ``launch`` in a fresh interpreter, which starts this file under ``mpirun`` (a + process that has initialized MPI, as pytest's has, cannot start mpirun).""" + visible = [d for d in os.environ.get("CUDA_VISIBLE_DEVICES", "").split(",") if d.strip()] + devices = (visible or [str(i) for i in range(torch.cuda.device_count())])[:WORLD] + env = dict(os.environ, CUDA_VISIBLE_DEVICES=",".join(devices)) + done = subprocess.run( + [sys.executable, os.path.abspath(__file__), "launch"], + env=env, + capture_output=True, + text=True, + timeout=DEADLINE_S + 300, + ) + print(done.stdout, flush=True) + assert done.returncode == 0 and "ALL PASS" in done.stdout, ( + done.stdout[-20000:] + done.stderr[-20000:] + ) -@pytest.mark.parametrize("dtype", DTYPES, ids=[DTYPE_NAMES[d] for d in DTYPES]) -@pytest.mark.parametrize("batch,block", GRAPH_SPLITS, ids=[f"{b}x{k}" for b, k in GRAPH_SPLITS]) -def test_graph_replay(batch, block, dtype): - with collective(): - row = graph_replays(group(), batch, block, dtype) - assert row["ok"], row +def launch() -> int: + """This file's ``report`` under ``mpirun -n WORLD`` (one rank per visible device), killed with its process + group at the deadline.""" + command = ["mpirun", "-n", str(WORLD), sys.executable, os.path.abspath(__file__), "report"] + print(f"[launch] {' '.join(command)}", flush=True) + process = subprocess.Popen(command, start_new_session=True) + try: + return process.wait(timeout=DEADLINE_S) + except subprocess.TimeoutExpired: + os.killpg(process.pid, signal.SIGKILL) + process.wait() + print( + f"[launch] the {WORLD}-rank run did not finish in {DEADLINE_S} s (wedged)", flush=True + ) + return 1 # ---------------------------------------------------------------------------------------------------------------- @@ -968,8 +924,9 @@ def main() -> int: parser.add_argument( "mode", nargs="?", - choices=("report", "time"), - help="the checks only, or the timing only (default: the checks, then the timing)", + choices=("report", "time", "launch"), + help="the checks only, the timing only, or the checks on WORLD ranks under a local mpirun (launch; " + "the pytest form) (default: the checks, then the timing)", ) parser.add_argument( "--copies", @@ -982,6 +939,8 @@ def main() -> int: args = parser.parse_args() if args.copies is not None and args.copies < 1: parser.error("--copies must be >= 1") + if args.mode == "launch": + return launch() if not _sm100() or _world_size() < 2: print( "needs SM100 GPUs and >= 2 MPI ranks: srun -n W --mpi=pmix python3 test_k3_markov.py", From 490f289223bdfe15c5773f69c2893b4a58fbe0ad Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 11:18:30 -0700 Subject: [PATCH 136/161] [None][test] Kimi K3 sharded k3_spec_accept test: run under pytest on 4 local MPI ranks The test ran only inside an MPI job (srun -n W python3 -m pytest); a single pytest process, as in CI, skipped it. pytest's one test now runs this file's launch mode in a fresh interpreter. That process runs the checks under `mpirun -n 4`, one rank per visible GPU, with a deadline that kills the process group. The test asserts exit status 0 and ALL PASS. The module skips with fewer than 4 SM100 GPUs. The checks and the srun script mode are unchanged. Signed-off-by: Vasanth Sabavat --- .../kimi_k3/test_k3_spec_accept_sharded.py | 92 ++++++++++++------- 1 file changed, 60 insertions(+), 32 deletions(-) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_spec_accept_sharded.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_spec_accept_sharded.py index 16ea213c4c56..6e5671f3f980 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_spec_accept_sharded.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_spec_accept_sharded.py @@ -34,14 +34,15 @@ Shapes: the rank's shard V / W of V = 163840 (W exchange slots), and TP16's 10240-column shard with every rank filling 16 / W slots (the exchange of 16 ranks; ``--copies`` overrides). - srun -n W --mpi=pmix python3 test_k3_spec_accept_sharded.py [--copies C] [--time [--base]] - srun -n W --mpi=pmix python3 -m pytest test_k3_spec_accept_sharded.py + srun -n W --mpi=pmix python3 test_k3_spec_accept_sharded.py [report] [--copies C] [--time [--base]] ``--time`` also reports us per call at batch 1 and every split of the sharded kernel and of the kernel on the gathered fp32 logits (CUDA graphs of 20 back-to-back calls, the slowest rank of each replay, median over 12 replays in alternating order; the all-gather, cat and cast that the sharded path removes are not in the second number); ``--base`` adds the installed base package's kernel (``$K3_BASE_TRTLLM``) on the same sharded calls and workspace. -Without an MPI job of 2 or more ranks the module skips. +Under pytest (``python3 -m pytest test_k3_spec_accept_sharded.py``, 4 GPUs visible) one test runs the checks on 4 +ranks: this file under a local ``mpirun -n 4``, started from a fresh interpreter, under a deadline; with fewer GPUs the +module is skipped. """ import argparse @@ -49,7 +50,9 @@ import importlib.util import os import random +import signal import statistics +import subprocess import sys import time import traceback @@ -97,19 +100,13 @@ GRAPH_KINDS = ("plain", "ties", "rank_ties", "nan", "negzero") -def launcher_world_size() -> int: - """The world size an MPI launcher (mpirun, MPICH, srun) gave this process, from its environment: whether to skip - is decided without initializing MPI or CUDA.""" - for name in ("OMPI_COMM_WORLD_SIZE", "PMI_SIZE", "SLURM_STEP_NUM_TASKS"): - value = os.environ.get(name, "") - if value.isdigit(): - return int(value) - return 1 - +# The ranks of the pytest run (one GB200 tray) and its deadline: a broken exchange hangs rather than raising. +WORLD = 4 +DEADLINE_S = 1200 pytestmark = pytest.mark.skipif( - launcher_world_size() < 2 or not _sm100(), - reason="needs an MPI job of 2 or more ranks on SM100 GPUs (srun -n W --mpi=pmix python3 -m pytest ...)", + not _sm100() or torch.cuda.device_count() < WORLD, + reason=f"needs {WORLD} SM100 GPUs (MNNVL multicast)", ) _env = {} @@ -456,27 +453,50 @@ def run_config(env: dict, name: str, vocab: int, copies: int, time_it: bool = Fa return None, results, timings -@pytest.mark.parametrize("config", [0, 1], ids=["V-over-W", "TP16-shard"]) -def test_sharded(config): - env = mpi_env() - if env["world"] < 2: - pytest.skip("needs an MPI job of 2 or more ranks") - name, vocab, copies = configs(env["world"])[config] - with torch.inference_mode(): - skip, results, _ = on_every_rank(env, run_config, env, name, vocab, copies) - if skip: - pytest.skip(skip) - bad = [ - {k: r[k] for k in ("config", "split", "case", "checks", "bad")} - for r in results - if not r["ok"] - ] - oks = env["comm"].allgather(not bad) - assert all(oks), (f"ranks ok: {oks}", bad) +# ---------------------------------------------------------------------------------------------------------------- +# pytest (python3 -m pytest test_k3_spec_accept_sharded.py): the checks on WORLD ranks of this node. +# ---------------------------------------------------------------------------------------------------------------- + + +@pytest.mark.no_xdist +def test_k3_spec_accept_sharded(): + """``report`` on WORLD ranks: ``launch`` in a fresh interpreter, which starts this file under ``mpirun`` (a + process that has initialized MPI, as pytest's has, cannot start mpirun).""" + visible = [d for d in os.environ.get("CUDA_VISIBLE_DEVICES", "").split(",") if d.strip()] + devices = (visible or [str(i) for i in range(torch.cuda.device_count())])[:WORLD] + env = dict(os.environ, CUDA_VISIBLE_DEVICES=",".join(devices)) + done = subprocess.run( + [sys.executable, os.path.abspath(__file__), "launch"], + env=env, + capture_output=True, + text=True, + timeout=DEADLINE_S + 300, + ) + print(done.stdout, flush=True) + assert done.returncode == 0 and "ALL PASS" in done.stdout, ( + done.stdout[-20000:] + done.stderr[-20000:] + ) + + +def launch() -> int: + """This file's ``report`` under ``mpirun -n WORLD`` (one rank per visible device), killed with its process + group at the deadline.""" + command = ["mpirun", "-n", str(WORLD), sys.executable, os.path.abspath(__file__), "report"] + print(f"[launch] {' '.join(command)}", flush=True) + process = subprocess.Popen(command, start_new_session=True) + try: + return process.wait(timeout=DEADLINE_S) + except subprocess.TimeoutExpired: + os.killpg(process.pid, signal.SIGKILL) + process.wait() + print( + f"[launch] the {WORLD}-rank run did not finish in {DEADLINE_S} s (wedged)", flush=True + ) + return 1 # ---------------------------------------------------------------------------------------------------------------- -# srun -n W --mpi=pmix python3 test_k3_spec_accept_sharded.py [--copies C] [--time] +# srun -n W --mpi=pmix python3 test_k3_spec_accept_sharded.py [report] [--copies C] [--time] # ---------------------------------------------------------------------------------------------------------------- @@ -521,6 +541,12 @@ def main() -> int: parser = argparse.ArgumentParser( description="trtllm::k3_spec_accept: vocabulary-sharded vs gathered logits" ) + parser.add_argument( + "mode", + nargs="?", + choices=("report", "launch"), + help="the checks (the default), or the checks on WORLD ranks under a local mpirun (launch; the pytest form)", + ) parser.add_argument("--copies", type=int, default=None, help="exchange slots every rank fills in the TP16 shape (default 16 / W)") # fmt: skip parser.add_argument("--time", action="store_true", @@ -533,6 +559,8 @@ def main() -> int: ) parser.add_argument("--time-only", action="store_true", help="the timing without the checks") args = parser.parse_args() + if args.mode == "launch": + return launch() time_splits = None if args.splits: time_splits = [tuple(int(v) for v in sp.split("x")) for sp in args.splits.split(",")] From f85b7090684d5b22bf884ca71c23663f367d550e Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 11:40:19 -0700 Subject: [PATCH 137/161] [None][perf] modeling_v2 Kimi K3 tp16_moetp4ep4: split the DSpark drafter's fc over the TP group The DSpark drafter (K3DSparkDrafter) kept its context projection's fc replicated: every rank read the whole [7168, 35840] bf16 weight, 513.8 MB, on every step, then applied hidden_norm in a separate kernel. Where the target builds the TP group's collective state (every attention all-reduce over MNNVL; TP16 is asserted at construction), post_load_weights now hands it to the drafter (_gate_drafter_comm, use_decode_comm). The drafter then keeps only this rank's contiguous block of fc's input columns (K3FcSlice: 2240 columns, 32 MB at TP16), and project_target_hidden runs: - this rank's partial product of its columns (cuBLAS on the strided column block of the captured features); - up to a decode step's rows (8 requests of 8 tokens), the sum over the group with hidden_norm in one comm/mnnvl_fusion_allreduce call (RESIDUAL_RMS_NORM with a zero residual) on the target's MNNVL workspace, one-shot up to the decode path's 4 MiB ceiling (K3DecodeComm.allreduce_norm); - more rows (a prefill chunk): the drafter's own TP all-reduce, then hidden_norm. Without the collective state, or where fc is not a bias-free bf16 projection whose columns split evenly over the group, fc stays replicated and the projection is the stock one. A later load_weights splits fc again. test_kimi_k3_drafter_comm.py (cpu_only): the column blocks of 1, 2, 4 and 16 ranks tile fc's columns, each is its columns of the full weight, and the partial products sum to the full product; fc splits only on the collective state and where it splits evenly; the projection hands the fused all-reduce this rank's partial, a zero residual and hidden_norm, and above a decode step's rows uses the TP all-reduce; the gate hands the state to the K3 drafter only; the TP16 workspace holds up to 144 rows. test_modeling_v2_kimi_k3_drafter.py (sm_100): on a one-rank collective state whose all-reduce is a counted torch stand-in, the split projection matches the replicated one at 1, 8, 64 and 200 rows and takes the fused all-reduce exactly up to 64 rows. The tp16_moetp16ep1 copy does not carry this change yet; until it does, test_modeling_v2_kimi_k3_drift fails on decode_comm.py and modeling.py. Signed-off-by: Vasanth Sabavat --- .../decode_comm.py | 43 ++- .../modeling.py | 157 +++++++- .../test_modeling_v2_kimi_k3_drafter.py | 92 ++++- .../hw_agnostic/test_kimi_k3_drafter_comm.py | 350 ++++++++++++++++++ 4 files changed, 630 insertions(+), 12 deletions(-) create mode 100644 tests/unittest/_torch/speculative/hw_agnostic/test_kimi_k3_drafter_comm.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py index 2ef7dbe27458..88311a249d60 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py @@ -17,6 +17,11 @@ the same kind of collective: `K3DecodeComm.sandwich_tail`, `comm/k3_sandwich_tail`, runs the tail GEMV, its all-reduce and the next layer's residual update (the final norm's, after the last layer) in one kernel. +The DSpark drafter (`K3DSparkDrafter`) runs its collectives with a plain residual add and RMSNorm on the same state: +`K3DecodeComm.allreduce_norm`, `comm/mnnvl_fusion_allreduce`, all-reduces an unreduced projection output of any token +count the MNNVL workspace holds, with the residual add and the RMSNorm in its epilogue (the context projection's, +with a zero residual). + The state is one `MnnvlWorkspace` and one `K3SandwichWorkspace` of the TP group (`K3DecodeComm.create`): collective over the group and eager, built by the target in `post_load_weights` before any CUDA-graph capture. Every rank must make the same calls on each in the same order. Which call a step takes is decided from its token count and kind and @@ -44,6 +49,10 @@ from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.mnnvl_allreduce_attn_res import ( mnnvl_allreduce_attn_res, ) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.mnnvl_fusion_allreduce import ( + mnnvl_fusion_allreduce, + required_buffer_bytes, +) from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.mnnvl_workspace import ( MnnvlWorkspace, ) @@ -104,8 +113,10 @@ class PendingTail(NamedTuple): def _eps(norm: nn.Module) -> float: - """The epsilon of a KimiK3RMSNorm (``eps``) or a stock RMSNorm (``variance_epsilon``).""" - return float(norm.eps if hasattr(norm, "eps") else norm.variance_epsilon) + """The epsilon of a KimiK3RMSNorm or ``torch.nn.RMSNorm`` (``eps``; for torch's None, its default: the machine + epsilon of the weight's dtype) or a stock RMSNorm (``variance_epsilon``).""" + eps = norm.eps if hasattr(norm, "eps") else norm.variance_epsilon + return float(torch.finfo(norm.weight.dtype).eps if eps is None else eps) def _res_args(res_proj: nn.Module, res_norm: nn.Module, out_norm: nn.Module) -> tuple: @@ -291,3 +302,31 @@ def sandwich_oproj( *_res_args(res_proj, res_norm, out_norm), self.sandwich, ) + + def takes_allreduce_norm(self, rows: int, hidden: int) -> bool: + """Whether ``allreduce_norm`` takes ``rows`` bf16 rows of ``hidden`` columns: the MNNVL workspace holds the + call at the decode path's one-shot ceiling (`DECODE_AR_ONE_SHOT_MAX_BYTES`, two-shot above it).""" + if rows <= 0 or hidden <= 0 or hidden % 8: + return False + world, buffer_bytes = self.mnnvl.world_size, self.mnnvl.buffer_bytes + need = required_buffer_bytes( + rows, hidden, world, torch.bfloat16, DECODE_AR_ONE_SHOT_MAX_BYTES + ) + two_shot = rows * hidden * world * 2 > DECODE_AR_ONE_SHOT_MAX_BYTES + return need <= buffer_bytes and not (two_shot and buffer_bytes % 32) + + def allreduce_norm( + self, partial: torch.Tensor, residual: torch.Tensor, norm: nn.Module + ) -> Tuple[torch.Tensor, torch.Tensor]: + """``(normed, updated)`` in one ``comm/mnnvl_fusion_allreduce`` call: ``updated = residual + + allreduce(partial)``, ``normed = norm(updated)`` for a plain RMSNorm ``norm``, sent one-shot up to + `DECODE_AR_ONE_SHOT_MAX_BYTES`. ``partial`` is this rank's unreduced ``[rows, hidden]`` bf16 output of a + row-parallel projection, of a shape ``takes_allreduce_norm`` holds for.""" + return mnnvl_fusion_allreduce( + partial.contiguous(), + self.mnnvl, + DECODE_AR_ONE_SHOT_MAX_BYTES, + residual.contiguous(), + norm.weight, + _eps(norm), + ) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index 4a6fe35f8ec7..ac7ad94127b5 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -71,8 +71,10 @@ **Speculative decoding** goes through the stock one-engine shell: DSpark or DFlash with an external drafter checkpoint, and SA. The DSpark drafter is this target's `K3DSparkDrafter`, the stock GQA drafter with a decode step's block on the drafter entries (`attention/k3_drafter_attn_qknorm` and the drafter's decode GEMV sites); the shell -builds it through `_build_draft_model`. The worker and its kernels stay upstream code; this target does not own a -worker. +builds it through `_build_draft_model`. Where every attention all-reduce runs over MNNVL, the drafter also runs on the +TP group's collective state: its context projection's `fc` is split by input feature over the group, `hidden_norm` +applied in the all-reduce of the partial products (`comm/mnnvl_fusion_allreduce`). The worker and its kernels stay +upstream code; this target does not own a worker. """ from __future__ import annotations @@ -188,10 +190,12 @@ "k3_moe", "k3_route_quant", "mnnvl_allgather_split", - # The DSpark drafter's decode path (K3DSparkDrafter): its GEMV sites and its block attention. + # The DSpark drafter's decode path (K3DSparkDrafter): its GEMV sites, its block attention and the all-reduce of + # its split context projection with hidden_norm. "k3_ctm_gemv", "k3_ctm_gemv_swiglu", "k3_drafter_attn_qknorm", + "mnnvl_fusion_allreduce", ) #: The engine surface the first forward checks before this target relies on it: per object, the attributes read. @@ -233,8 +237,8 @@ "tensorrt_llm._torch.distributed.AllReduce", "tensorrt_llm._torch.modules.multi_stream_utils.maybe_execute_in_parallel", # The DSpark drafter: the stock GQA drafter K3DSparkDrafter extends (its block forward where the drafter entries - # do not take a block, its context projection and context k / v, its heads and its weight load), and the stock - # builder's checks for which drafter a checkpoint gets. + # do not take a block, its context projection without the TP group's collective state, its context k / v, its + # heads and its weight load), and the stock builder's checks for which drafter a checkpoint gets. "tensorrt_llm._torch.models.modeling_dspark.GQADSparkForCausalLM", "tensorrt_llm._torch.models.modeling_dspark.draft_is_embedded_in_target", "tensorrt_llm._torch.models.modeling_dflash.DFlashForCausalLM", @@ -2401,6 +2405,42 @@ def _k3_decode_view(self, attn_metadata: AttentionMetadata, num_tokens: int) -> _DRAFTER_ROPE_BASE = 10000.0 +def fc_columns(in_features: int, tp_size: int, tp_rank: int) -> Optional[Tuple[int, int]]: + """Rank ``tp_rank``'s input columns ``[start, end)`` of the drafter's context projection ``fc`` split by input + feature over ``tp_size`` ranks: equal contiguous blocks in rank order. None where they do not split evenly.""" + if tp_size < 1 or not 0 <= tp_rank < tp_size or in_features <= 0 or in_features % tp_size: + return None + width = in_features // tp_size + return tp_rank * width, (tp_rank + 1) * width + + +class K3FcSlice(nn.Module): + """This rank's block of the DSpark drafter's context projection ``fc`` (`fc_columns`): ``weight`` is + ``fc.weight[:, start:end]``, contiguous. Its output is this rank's partial product; the sum over the TP group's + ranks is ``fc``'s output.""" + + def __init__(self, weight: torch.Tensor, start: int, end: int) -> None: + super().__init__() + self.weight = nn.Parameter(weight, requires_grad=False) + self.start = start + self.end = end + + @classmethod + def of(cls, fc_weight: torch.Tensor, tp_size: int, tp_rank: int) -> Optional["K3FcSlice"]: + """Rank ``tp_rank``'s block of the full ``fc_weight`` ``[out_features, in_features]``: a copy of its columns + (on one rank, the weight itself); None where the input columns do not split evenly over ``tp_size`` ranks.""" + columns = fc_columns(fc_weight.shape[1], tp_size, tp_rank) + if columns is None: + return None + start, end = columns + return cls(fc_weight.detach()[:, start:end].contiguous(), start, end) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + """``hidden_states[:, start:end] @ weight.T`` for the full features ``hidden_states`` ``[N, in_features]``: the + GEMM reads the column block in place, through the rows' stride.""" + return torch.nn.functional.linear(hidden_states[:, self.start : self.end], self.weight) + + class K3DSparkDrafter(GQADSparkForCausalLM): """Kimi K3's DSpark drafter: the stock GQA drafter, with a decode step's block forward on the drafter entries. @@ -2417,8 +2457,13 @@ class K3DSparkDrafter(GQADSparkForCausalLM): The norms and the residual adds are the stock modules'; a projection whose site does not take its rows runs its module. Every other block runs the stock ``dflash_forward``: a split `DRAFTER_ATTN_SPLITS` does not list, another attention backend, head layout, RoPE or normalization, a cache the kernel does not read, or a compile key of the - attention that has not run eagerly, under CUDA-graph capture. The worker that calls it (the stock - ``DSparkWorker``), the context projection, the context k / v and the Markov head stay upstream code. + attention that has not run eagerly, under CUDA-graph capture. + + Where the target hands the drafter the TP group's collective state (``use_decode_comm``), the context projection's + ``fc`` holds only this rank's block of input columns (`K3FcSlice`), and ``project_target_hidden`` sums the ranks' + partial products with ``hidden_norm`` in one all-reduce; otherwise the context projection is the stock one, with + ``fc`` replicated. The worker that calls the drafter (the stock ``DSparkWorker``), the context k / v and the Markov + head stay upstream code. """ def __init__( @@ -2432,6 +2477,88 @@ def __init__( # The attention's compile keys that ran eagerly ((more than one request, page stride)); a capture takes # only these. self._k3_attn_ran: set = set() + # The TP group's collective state (decode_comm.py), set by the target where it builds one (use_decode_comm). + self.decode_comm: Optional[_decode_comm.K3DecodeComm] = None + # The zero residual of the split context projection's all-reduce, up to a decode step's rows. + self._k3_zero_rows: Optional[torch.Tensor] = None + + def use_decode_comm(self, comm: _decode_comm.K3DecodeComm) -> None: + """Run the drafter's collectives on ``comm``, the TP group's decode state (``decode_comm.K3DecodeComm``) the + target builds where every attention all-reduce runs over MNNVL. The target calls it on every rank once the + weights are loaded, before any CUDA-graph capture. + + The context projection's ``fc`` then keeps only this rank's block of input columns (`K3FcSlice`), where they + split evenly over the group and the drafter's layers have their TP all-reduce: ``project_target_hidden`` sums + the ranks' partial products and applies ``hidden_norm`` in one ``comm/mnnvl_fusion_allreduce`` call.""" + self.decode_comm = comm + self._k3_split_fc() + + def _k3_tp_all_reduce(self) -> Optional[nn.Module]: + """The drafter's own TP all-reduce (its first output projection's row-parallel all-reduce module), or None.""" + o_proj = self.model.layers[0].self_attn.o_proj + if getattr(getattr(o_proj, "tp_mode", None), "name", None) != "ROW": + return None + return getattr(o_proj, "all_reduce", None) + + def _k3_split_fc(self) -> None: + """Replace the replicated ``fc`` with this rank's block of its input columns (`K3FcSlice`) and size the zero + residual of the projection's all-reduce, where the drafter runs on the TP group's collective state, ``fc`` is + a bias-free bf16 projection whose input columns split evenly over the group, and the drafter has its TP + all-reduce (for rows the MNNVL workspace does not hold). Otherwise ``fc`` stays replicated.""" + fc = getattr(self, "fc", None) + if self.decode_comm is None or fc is None or isinstance(fc, K3FcSlice): + return + mapping = self.model_config.mapping + weight = getattr(fc, "weight", None) + sliced = None + if ( + isinstance(weight, torch.Tensor) + and weight.dim() == 2 + and weight.dtype == torch.bfloat16 + and getattr(fc, "bias", None) is None + and self._k3_tp_all_reduce() is not None + ): + sliced = K3FcSlice.of(weight, mapping.tp_size, mapping.tp_rank) + if sliced is None: + logger.info( + "Kimi K3 DSpark drafter: fc stays replicated (it is not a bias-free bf16 projection whose input " + f"columns split evenly over TP{mapping.tp_size}, or the layers have no TP all-reduce)" + ) + return + self.fc = sliced + self._k3_zero_rows = weight.new_zeros( + MAX_REQUESTS * MAX_TOKENS_PER_REQUEST, weight.shape[0] + ) + logger.info( + f"Kimi K3 DSpark drafter: fc split by input feature over TP{mapping.tp_size}: rank {mapping.tp_rank} " + f"holds columns [{sliced.start}, {sliced.end}), hidden_norm in the all-reduce" + ) + + def load_weights(self, weights, weight_mapper=None, **kwargs): + """The stock load; on the TP group's collective state, ``fc`` is split again (`_k3_split_fc`).""" + result = super().load_weights(weights, weight_mapper=weight_mapper, **kwargs) + self._k3_split_fc() + return result + + def project_target_hidden(self, hidden_states: torch.Tensor) -> torch.Tensor: + """``hidden_norm(fc(hidden_states))`` of the captured target features ``[N, in_features]``. + + With ``fc`` split (`K3FcSlice`): this rank's partial product of its columns, then the sum over the TP group + with ``hidden_norm`` applied in one ``comm/mnnvl_fusion_allreduce`` call (a zero residual) up to a decode + step's rows the MNNVL workspace holds; more rows (a prefill) go through the drafter's TP all-reduce, then + ``hidden_norm``. Otherwise the stock projection.""" + fc = self.fc + if not isinstance(fc, K3FcSlice): + return super().project_target_hidden(hidden_states) + partial = fc(hidden_states.to(fc.weight.dtype)) + rows = partial.shape[0] + zeros = self._k3_zero_rows + if rows <= zeros.shape[0] and self.decode_comm.takes_allreduce_norm(rows, partial.shape[1]): + normed, _ = self.decode_comm.allreduce_norm(partial, zeros[:rows], self.hidden_norm) + return normed + if rows == 0: + return self.hidden_norm(partial) + return self.hidden_norm(self._k3_tp_all_reduce()(partial)) def dflash_forward( self, @@ -2872,7 +2999,8 @@ def post_load_weights(self) -> None: attention all-reduce runs over MNNVL, the TP group's collective state (``K3DecodeComm``, collective: every rank builds it here) handed to every layer, with whether its sandwich takes the layer's o_proj, the decode path's one-shot ceiling on every stock MNNVL all-reduce (``use_decode_one_shot``), and the MoE decode path - (``_build_decode_moe``); then the speculative worker's decode kernels (``_gate_spec_worker_kernels``).""" + (``_build_decode_moe``); then the speculative worker's decode kernels (``_gate_spec_worker_kernels``) and the + DSpark drafter's collectives (``_gate_drafter_comm``).""" super().post_load_weights() kda = [layer.linear_attn for layer in self.model.layers if layer.is_kda] mla = [layer.self_attn.mixer for layer in self.model.layers if not layer.is_kda] @@ -2916,6 +3044,7 @@ def post_load_weights(self) -> None: + f"; MoE on k3_moe_front, k3_moe and the row-parallel tail ({moe_layers} layers)" ) self._gate_spec_worker_kernels(comm) + self._gate_drafter_comm(comm) def _build_decode_moe(self, comm: _decode_comm.K3DecodeComm) -> int: """The MoE decode path (``decode_moe.py``) on every MoE layer it takes: the shared state (collective: the @@ -2981,6 +3110,18 @@ def _gate_spec_worker_kernels(self, comm: Optional[_decode_comm.K3DecodeComm]) - ) return worker.k3_decode + def _gate_drafter_comm(self, comm: Optional[_decode_comm.K3DecodeComm]) -> bool: + """Hand the DSpark drafter (`K3DSparkDrafter`) the TP group's collective state ``comm`` where this target built + it (every attention all-reduce over MNNVL; TP16 is a construction assert): the drafter then runs its context + projection split over the group, ``hidden_norm`` in the all-reduce (``K3DSparkDrafter.use_decode_comm``). + Without ``comm`` it keeps the stock replicated ``fc``. Returns whether the drafter took the state (False + without such a drafter).""" + drafter = getattr(self, "draft_model", None) + if comm is None or not isinstance(drafter, K3DSparkDrafter): + return False + drafter.use_decode_comm(comm) + return True + def _check_step_contract(self, attn_metadata: AttentionMetadata) -> None: """First-forward checks of the engine surface and the per-engine settings.""" objects = { diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py index f2dad0161298..7ecc2b53232d 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py @@ -11,15 +11,24 @@ * A split the entry does not certify runs the stock forward, bit for bit, without the entry. * Under CUDA-graph capture, the entries take a block only once its attention compile key has run eagerly. * Negative control: a weight changed between the two forwards fails the comparison. + +On the TP group's collective state (``use_decode_comm``), here a group of one rank whose collectives are torch +stand-ins counted per call (the ops themselves are certified by their multi-GPU op matrices): + +* The context projection runs on the split ``fc`` (on one rank, the whole weight): up to a decode step's rows through + ``comm/mnnvl_fusion_allreduce`` with ``hidden_norm``, more rows through the drafter's TP all-reduce, and matches the + replicated projection. """ import math +from types import SimpleNamespace import pytest import torch from transformers import Qwen3Config from tensorrt_llm._torch._experimental.modeling_v2.models.kimi_k3_vl.kimi_k3_mxfp4__sm_100__tp16_moetp4ep4 import ( # noqa: E501 + decode_comm, decode_gemv, ) from tensorrt_llm._torch._experimental.modeling_v2.models.kimi_k3_vl.kimi_k3_mxfp4__sm_100__tp16_moetp4ep4 import ( # noqa: E501 @@ -100,8 +109,7 @@ def rnd(*shape, scale=0.02): return w -@pytest.fixture(scope="module") -def drafter(): +def _load_drafter(): model_config = ModelConfig(pretrained_config=_config(), attn_backend="TRTLLM") module = target.K3DSparkDrafter(model_config, dflash_attention_backend="TRTLLM").to("cuda") module.load_weights(_weights()) @@ -109,6 +117,68 @@ def drafter(): return module +@pytest.fixture(scope="module") +def drafter(): + return _load_drafter() + + +# The TP group's collective state for a group of one rank: the workspaces' sizes are what the drafter reads; their +# collectives are the stand-ins below. +ONE_RANK_COMM = decode_comm.K3DecodeComm( + mnnvl=SimpleNamespace(world_size=1, buffer_bytes=decode_comm.MNNVL_BUFFER_BYTES), + sandwich=SimpleNamespace(world_size=1), +) +DECODE_ROWS = target.MAX_REQUESTS * target.MAX_TOKENS_PER_REQUEST + + +def _rms_norm(x, weight, eps): + """``RMSNorm(x) * weight`` in fp32, rounded once to bf16.""" + x = x.float() + return (x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + eps) * weight.float()).to( + torch.bfloat16 + ) + + +def _fusion_allreduce_one_rank( + input, workspace, one_shot_max_bytes, residual=None, norm_weight=None, eps=None +): + """``comm/mnnvl_fusion_allreduce`` over one rank: the sum is the input, then the residual add and the RMSNorm.""" + assert workspace is ONE_RANK_COMM.mnnvl + assert one_shot_max_bytes == decode_comm.DECODE_AR_ONE_SHOT_MAX_BYTES + assert input.dtype == residual.dtype == norm_weight.dtype == torch.bfloat16 + assert input.is_contiguous() and residual.is_contiguous() and input.shape == residual.shape + updated = (input.float() + residual.float()).to(torch.bfloat16) + return _rms_norm(updated, norm_weight, eps), updated + + +def _one_rank_collectives(monkeypatch, calls): + """The drafter's collectives replaced by their one-rank stand-ins, each call recorded as (op, rows).""" + + def fusion(input, *args, **kwargs): + calls.append(("mnnvl_fusion_allreduce", input.shape[0])) + return _fusion_allreduce_one_rank(input, *args, **kwargs) + + monkeypatch.setattr(decode_comm, "mnnvl_fusion_allreduce", fusion) + + +@pytest.fixture(scope="module") +def fused_drafter(): + """The drafter on the one-rank collective state: its fc split over one rank.""" + with pytest.MonkeyPatch.context() as mp: + _one_rank_collectives(mp, []) + module = _load_drafter() + module.use_decode_comm(ONE_RANK_COMM) + return module + + +@pytest.fixture +def collectives(monkeypatch): + """The one-rank stand-ins of the drafter's collectives, counted: the list of (op, rows) calls.""" + calls = [] + _one_rank_collectives(monkeypatch, calls) + return calls + + @pytest.fixture(scope="module") def gemvs(): return decode_gemv.K3DecodeGemvs.create(None, sites=SITES) @@ -245,3 +315,21 @@ def test_negative_control_a_changed_weight_fails(drafter): weight.copy_(saved) ref = DFlashForCausalLM.dflash_forward(drafter, **stock_inputs) assert _rel_l2(out, ref) > REL_L2 + + +@pytest.mark.parametrize("rows", [1, 8, DECODE_ROWS, 200]) +def test_split_fc_matches_the_replicated_projection(drafter, fused_drafter, collectives, rows): + """On one rank the block is the whole fc. Up to a decode step's rows the projection's all-reduce applies + hidden_norm; above, the drafter's TP all-reduce (the identity on one rank) runs and then hidden_norm.""" + fc = fused_drafter.fc + assert isinstance(fc, target.K3FcSlice) and (fc.start, fc.end) == (0, 2 * HIDDEN) + assert torch.equal(fc.weight, drafter.fc.weight) + g = torch.Generator(device="cuda").manual_seed(rows) + features = torch.randn(rows, 2 * HIDDEN, generator=g, device="cuda").to(torch.bfloat16) + out = fused_drafter.project_target_hidden(features) + ref = drafter.project_target_hidden(features) + torch.cuda.synchronize() + assert out.shape == ref.shape == (rows, HIDDEN) + err = _rel_l2(out, ref) + assert err <= REL_L2, err + assert collectives == ([("mnnvl_fusion_allreduce", rows)] if rows <= DECODE_ROWS else []) diff --git a/tests/unittest/_torch/speculative/hw_agnostic/test_kimi_k3_drafter_comm.py b/tests/unittest/_torch/speculative/hw_agnostic/test_kimi_k3_drafter_comm.py new file mode 100644 index 000000000000..ca88d0906f2f --- /dev/null +++ b/tests/unittest/_torch/speculative/hw_agnostic/test_kimi_k3_drafter_comm.py @@ -0,0 +1,350 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""The Kimi K3 target ``kimi_k3_mxfp4__sm_100__tp16_moetp4ep4``'s DSpark drafter on the TP group's collective state +(host-side: CPU tensors and fakes, the collective replaced by a recorder). + +* ``fc_columns`` / ``K3FcSlice``: the context projection's ``fc`` split by input feature over 1, 2, 4 and 16 ranks. + Each rank's block is its columns of the full weight, copied; the blocks tile the columns in rank order; the ranks' + partial products sum to the full product. Uneven splits are refused. +* ``K3DSparkDrafter._k3_split_fc`` splits ``fc`` only on the collective state, for a bias-free bf16 weight whose + columns split evenly, with the layers' TP all-reduce. ``project_target_hidden`` hands this rank's partial product, + a zero residual and ``hidden_norm`` to the fused all-reduce up to a decode step's rows (the ranks' partials sum to + the full product), and above them sums the partials with the drafter's TP all-reduce, then applies + ``hidden_norm``; either way the result is the stock projection's. +* ``_gate_drafter_comm`` hands the state to a ``K3DSparkDrafter`` only, and only where the target built it. +* ``K3DecodeComm.takes_allreduce_norm`` / ``allreduce_norm``: which calls the TP16 MNNVL workspace holds, and the + catalog call's arguments. +""" + +from types import SimpleNamespace + +import pytest +import torch +from torch import nn + +from tensorrt_llm._torch._experimental.modeling_v2.models.kimi_k3_vl.kimi_k3_mxfp4__sm_100__tp16_moetp4ep4 import ( # noqa: E501 + decode_comm, +) +from tensorrt_llm._torch._experimental.modeling_v2.models.kimi_k3_vl.kimi_k3_mxfp4__sm_100__tp16_moetp4ep4 import ( # noqa: E501 + modeling as target, +) + +pytestmark = pytest.mark.cpu_only + +# A small context projection: 4 captured layers of 32 features into 32 (K3's is 5 x 7168 into 7168). +IN, OUT = 128, 32 +DECODE_ROWS = target.MAX_REQUESTS * target.MAX_TOKENS_PER_REQUEST + + +def _fc(dtype=torch.bfloat16, seed=3, bias=False): + g = torch.Generator().manual_seed(seed) + fc = nn.Linear(IN, OUT, bias=bias, dtype=dtype) + with torch.no_grad(): + fc.weight.copy_(torch.randn(OUT, IN, generator=g) * 0.1) + return fc + + +def _hidden_norm(eps=1e-6): + norm = nn.RMSNorm(OUT, eps=eps, dtype=torch.bfloat16) + with torch.no_grad(): + norm.weight.copy_(torch.linspace(0.5, 1.5, OUT)) + return norm + + +class _Sum: + """A TP all-reduce over the fake group: records its inputs, returns ``total`` (or the input).""" + + def __init__(self, total=None): + self.inputs = [] + self.total = total + + def __call__(self, x): + self.inputs.append(x) + return x if self.total is None else self.total + + +def _comm(world=1, buffer_bytes=4 << 20): + return decode_comm.K3DecodeComm( + mnnvl=SimpleNamespace(world_size=world, buffer_bytes=buffer_bytes), + sandwich=SimpleNamespace(world_size=world), + ) + + +def _drafter(tp_size=1, tp_rank=0, comm=None, fc=None, all_reduce=True, row_parallel=True): + """A K3DSparkDrafter without ``__init__`` (that needs a GPU drafter checkpoint): the fields the context projection + reads.""" + drafter = target.K3DSparkDrafter.__new__(target.K3DSparkDrafter) + nn.Module.__init__(drafter) + drafter.model_config = SimpleNamespace( + mapping=SimpleNamespace(tp_size=tp_size, tp_rank=tp_rank) + ) + o_proj = SimpleNamespace( + tp_mode=SimpleNamespace(name="ROW" if row_parallel else "COLUMN"), + all_reduce=_Sum() if all_reduce else None, + ) + drafter.model = SimpleNamespace( + layers=[SimpleNamespace(self_attn=SimpleNamespace(o_proj=o_proj))] + ) + drafter.fc = _fc() if fc is None else fc + drafter.hidden_norm = _hidden_norm() + drafter.decode_comm = comm + drafter._k3_zero_rows = None + return drafter + + +@pytest.fixture +def fused_calls(monkeypatch): + """Records ``comm/mnnvl_fusion_allreduce`` calls; each returns the one-rank result (the sum is the input).""" + calls = [] + + def one_rank(input, workspace, one_shot_max_bytes, residual=None, norm_weight=None, eps=None): + calls.append( + dict( + input=input, + workspace=workspace, + one_shot_max_bytes=one_shot_max_bytes, + residual=residual, + norm_weight=norm_weight, + eps=eps, + ) + ) + updated = (input.float() + residual.float()).to(input.dtype) + normed = torch.nn.functional.rms_norm(updated.float(), (updated.shape[-1],), eps=eps) + return (normed * norm_weight.float()).to(input.dtype), updated + + monkeypatch.setattr(decode_comm, "mnnvl_fusion_allreduce", one_rank) + return calls + + +@pytest.mark.parametrize("tp_size", [1, 2, 4, 16]) +def test_fc_columns_tile_the_inputs_in_rank_order(tp_size): + in_features = 5 * 7168 + columns = [target.fc_columns(in_features, tp_size, rank) for rank in range(tp_size)] + assert columns[0][0] == 0 and columns[-1][1] == in_features + assert all(end - start == in_features // tp_size for start, end in columns) + assert all(a[1] == b[0] for a, b in zip(columns, columns[1:])) + assert target.fc_columns(5 * 7168, 16, 3) == (6720, 8960) + + +@pytest.mark.parametrize( + "in_features,tp_size,tp_rank", + [(130, 4, 0), (128, 0, 0), (128, 4, 4), (128, 4, -1), (0, 4, 0)], + ids=["uneven", "no ranks", "rank past the group", "negative rank", "no columns"], +) +def test_fc_columns_refuse_what_does_not_split(in_features, tp_size, tp_rank): + assert target.fc_columns(in_features, tp_size, tp_rank) is None + + +@pytest.mark.parametrize("tp_size", [1, 2, 4, 16]) +def test_each_ranks_slice_is_its_columns_and_the_partials_sum_to_fc(tp_size): + g = torch.Generator().manual_seed(tp_size) + weight = torch.randn(OUT, IN, generator=g, dtype=torch.float64) + x = torch.randn(9, IN, generator=g, dtype=torch.float64) + slices = [target.K3FcSlice.of(weight, tp_size, rank) for rank in range(tp_size)] + for rank, block in enumerate(slices): + start, end = target.fc_columns(IN, tp_size, rank) + assert (block.start, block.end) == (start, end) + assert torch.equal(block.weight, weight[:, start:end]) + assert block.weight.is_contiguous() and not block.weight.requires_grad + # A block of some columns is a copy; one rank's block is the whole weight. + assert (block.weight.data_ptr() == weight.data_ptr()) == (tp_size == 1) + # Reassembled: the blocks side by side are the full weight, and the partial products sum to the full product. + assert torch.equal(torch.cat([block.weight for block in slices], dim=1), weight) + total = sum(block(x) for block in slices) + torch.testing.assert_close(total, x @ weight.T, rtol=1e-12, atol=1e-12) + assert target.K3FcSlice.of(weight[:, :-1], tp_size, 0) is None or tp_size == 1 + + +def test_fc_splits_only_on_the_collective_state(): + drafter = _drafter(tp_size=4, tp_rank=1) + drafter._k3_split_fc() + assert isinstance(drafter.fc, nn.Linear) and drafter._k3_zero_rows is None + + full = drafter.fc.weight.detach().clone() + drafter.decode_comm = _comm(world=4) + drafter._k3_split_fc() + assert isinstance(drafter.fc, target.K3FcSlice) + assert (drafter.fc.start, drafter.fc.end) == (32, 64) + assert torch.equal(drafter.fc.weight, full[:, 32:64]) + assert drafter._k3_zero_rows.shape == (DECODE_ROWS, OUT) + assert drafter._k3_zero_rows.dtype == torch.bfloat16 and not drafter._k3_zero_rows.any() + # Split once: a second call keeps the block. + block = drafter.fc + drafter._k3_split_fc() + assert drafter.fc is block + + +@pytest.mark.parametrize( + "case", + ["bias", "fp32 weight", "uneven columns", "no TP all-reduce", "column-parallel o_proj"], +) +def test_fc_stays_replicated_where_it_does_not_split(case): + kwargs = dict(tp_size=4, tp_rank=0, comm=_comm(world=4)) + if case == "bias": + kwargs["fc"] = _fc(bias=True) + elif case == "fp32 weight": + kwargs["fc"] = _fc(dtype=torch.float32) + elif case == "uneven columns": + kwargs["tp_size"], kwargs["comm"] = 3, _comm(world=3) + elif case == "no TP all-reduce": + kwargs["all_reduce"] = False + else: + kwargs["row_parallel"] = False + drafter = _drafter(**kwargs) + fc = drafter.fc + drafter._k3_split_fc() + assert drafter.fc is fc and drafter._k3_zero_rows is None + + +def test_reload_splits_fc_again(monkeypatch): + drafter = _drafter(tp_size=2, tp_rank=1, comm=_comm(world=2)) + drafter._k3_split_fc() + reloaded = _fc(seed=9) + + def stock_load(self, weights, weight_mapper=None, **kwargs): + self.fc = reloaded + + monkeypatch.setattr(target.GQADSparkForCausalLM, "load_weights", stock_load) + drafter.load_weights({}) + assert isinstance(drafter.fc, target.K3FcSlice) + assert torch.equal(drafter.fc.weight, reloaded.weight[:, IN // 2 :]) + + +@pytest.mark.parametrize("rows", [1, 8, DECODE_ROWS]) +def test_split_projection_takes_the_fused_all_reduce(fused_calls, rows): + stock = _drafter() + drafter = _drafter(comm=_comm()) + drafter.fc = stock.fc + drafter._k3_split_fc() + x = torch.randn(rows, IN, generator=torch.Generator().manual_seed(rows)).to(torch.bfloat16) + + out = drafter.project_target_hidden(x) + ref = stock.project_target_hidden(x) + assert len(fused_calls) == 1 + call = fused_calls[0] + torch.testing.assert_close(call["input"], drafter.fc(x), rtol=0, atol=0) + assert call["workspace"] is drafter.decode_comm.mnnvl + assert call["one_shot_max_bytes"] == decode_comm.DECODE_AR_ONE_SHOT_MAX_BYTES + assert call["residual"].shape == (rows, OUT) and not call["residual"].any() + assert call["norm_weight"] is drafter.hidden_norm.weight and call["eps"] == 1e-6 + assert drafter.model.layers[0].self_attn.o_proj.all_reduce.inputs == [] + torch.testing.assert_close(out, ref, rtol=2e-2, atol=2e-2) + + +def test_split_projection_above_a_decode_step_uses_the_tp_all_reduce(fused_calls): + stock = _drafter() + drafter = _drafter(comm=_comm()) + drafter.fc = stock.fc + drafter._k3_split_fc() + rows = DECODE_ROWS + 1 + x = torch.randn(rows, IN, generator=torch.Generator().manual_seed(5)).to(torch.bfloat16) + + out = drafter.project_target_hidden(x) + assert fused_calls == [] + all_reduce = drafter.model.layers[0].self_attn.o_proj.all_reduce + assert len(all_reduce.inputs) == 1 and all_reduce.inputs[0].shape == (rows, OUT) + torch.testing.assert_close(out, stock.project_target_hidden(x), rtol=2e-2, atol=2e-2) + + +def test_split_projection_where_the_workspace_does_not_hold_the_rows(fused_calls): + drafter = _drafter(comm=_comm(buffer_bytes=16)) + drafter._k3_split_fc() + out = drafter.project_target_hidden(torch.ones(2, IN, dtype=torch.bfloat16)) + assert fused_calls == [] and out.shape == (2, OUT) + assert len(drafter.model.layers[0].self_attn.o_proj.all_reduce.inputs) == 1 + + +def test_split_partials_of_a_tp_group_sum_to_the_projection(fused_calls): + """Four ranks, each with its block: the partial each hands the all-reduce sums, over the ranks, to the full + product; the all-reduce's result is the stock projection.""" + tp = 4 + stock = _drafter() + x = torch.randn(8, IN, generator=torch.Generator().manual_seed(4)).to(torch.bfloat16) + for rank in range(tp): + drafter = _drafter(tp_size=tp, tp_rank=rank, comm=_comm(world=tp)) + drafter.fc = stock.fc + drafter._k3_split_fc() + drafter.project_target_hidden(x) + assert len(fused_calls) == tp + total = sum(call["input"].float() for call in fused_calls) + full = x.float() @ stock.fc.weight.float().T + assert ((total - full).norm() / full.norm()).item() < 1e-2 + + +def test_unsplit_projection_is_the_stock_one(fused_calls): + drafter = _drafter() + x = torch.randn(4, IN, generator=torch.Generator().manual_seed(1)).to(torch.bfloat16) + out = drafter.project_target_hidden(x) + assert fused_calls == [] + torch.testing.assert_close(out, drafter.hidden_norm(drafter.fc(x)), rtol=0, atol=0) + + +def test_gate_hands_the_state_to_the_k3_drafter_only(): + comm = _comm(world=16) + taken = [] + drafter = _drafter() + drafter.use_decode_comm = taken.append + assert target.ModelingV2KimiK3Mxfp4Sm100Tp16Moetp4ep4._gate_drafter_comm( + SimpleNamespace(draft_model=drafter), comm + ) + assert taken == [comm] + assert not target.ModelingV2KimiK3Mxfp4Sm100Tp16Moetp4ep4._gate_drafter_comm( + SimpleNamespace(draft_model=drafter), None + ) + assert taken == [comm] + for other in (None, SimpleNamespace(use_decode_comm=taken.append)): + assert not target.ModelingV2KimiK3Mxfp4Sm100Tp16Moetp4ep4._gate_drafter_comm( + SimpleNamespace(draft_model=other), comm + ) + assert taken == [comm] + + +def test_use_decode_comm_keeps_the_state_and_splits_fc(): + comm = _comm(world=2) + drafter = _drafter(tp_size=2, tp_rank=0) + drafter.use_decode_comm(comm) + assert drafter.decode_comm is comm + assert isinstance(drafter.fc, target.K3FcSlice) + assert (drafter.fc.start, drafter.fc.end) == (0, 64) + + +def test_tp16_workspace_holds_a_decode_steps_rows(): + """The target's 4 MiB MNNVL buffers at TP16: rows of 7168 go one-shot up to 18, two-shot up to 144.""" + comm = _comm(world=16, buffer_bytes=decode_comm.MNNVL_BUFFER_BYTES) + assert all(comm.takes_allreduce_norm(rows, 7168) for rows in range(1, 145)) + assert not comm.takes_allreduce_norm(145, 7168) + assert not comm.takes_allreduce_norm(0, 7168) + assert not comm.takes_allreduce_norm(8, 7164) + assert 18 * 7168 * 16 * 2 <= decode_comm.DECODE_AR_ONE_SHOT_MAX_BYTES < 19 * 7168 * 16 * 2 + # A two-shot call needs a buffer of whole 32-byte units. + assert not _comm(world=16, buffer_bytes=(4 << 20) + 16).takes_allreduce_norm(64, 7168) + assert _comm(world=16, buffer_bytes=(4 << 20) + 16).takes_allreduce_norm(8, 7168) + + +@pytest.mark.parametrize("eps,expected", [(1e-5, 1e-5), (None, torch.finfo(torch.bfloat16).eps)]) +def test_allreduce_norm_passes_the_workspace_ceiling_and_norm(fused_calls, eps, expected): + comm = _comm(world=16) + norm = nn.RMSNorm(OUT, eps=eps, dtype=torch.bfloat16) + partial = torch.ones(3, OUT, dtype=torch.bfloat16)[:, :] + residual = torch.zeros(3, OUT, dtype=torch.bfloat16) + comm.allreduce_norm(partial, residual, norm) + (call,) = fused_calls + assert call["input"] is partial and call["residual"] is residual + assert call["workspace"] is comm.mnnvl and call["norm_weight"] is norm.weight + assert call["one_shot_max_bytes"] == decode_comm.DECODE_AR_ONE_SHOT_MAX_BYTES + assert call["eps"] == pytest.approx(expected) + stock = SimpleNamespace(variance_epsilon=1e-6, weight=norm.weight) + comm.allreduce_norm(partial, residual, stock) + assert fused_calls[1]["eps"] == pytest.approx(1e-6) From 135ab0cdffd5d0629097b6aef5bcb9bed769ef34 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 11:48:18 -0700 Subject: [PATCH 138/161] [None][perf] modeling_v2 Kimi K3 tp16_moetp4ep4: the DSpark drafter's residual adds and RMSNorms in its all-reduces On the drafter entries' decode blocks, each layer ran its output and down projections' plain all-reduces, then the stock fused add + RMSNorm, and the block copied its input to start the residual. With the TP group's collective state (use_decode_comm, from the previous change), a block whose layers take it now runs each residual add and RMSNorm in the all-reduce before it: o_proj's with the post-attention norm, the MLP's with the next layer's input norm, the last layer's with the final norm. The fused all-reduces read the residual and return the updated one, so the block's input is the first residual, uncopied. - Up to 8 rows: comm/k3_sandwich_plain runs o_proj (drafter_o's arithmetic), its all-reduce, the residual add and the norm in one launch, and its SiLU-and-mul form does the same for the down projection after drafter_gate_up (drafter_down's arithmetic). K3DecodeComm.sandwich_plain calls it on the target's sandwich workspace; by the kernel's statement it is bit for bit the GEMV site followed by the MNNVL one-shot RESIDUAL_RMS_NORM all-reduce. - Otherwise (more rows, or a sandwich form that did not compile): the projection on its site, or its module without its all-reduce (decode_comm.skip_all_reduce), then comm/mnnvl_fusion_allreduce with RESIDUAL_RMS_NORM (K3DecodeComm.allreduce_norm). - use_decode_comm compiles each sandwich form whose shape every layer shares with one zero-row call on the group's sandwich workspace (collective, like the target's own sandwich compiles in post_load_weights), and uses only the forms that compiled. - The stock all-reduces and norms stay without the collective state, for a block the MNNVL workspace does not hold, and for layers whose norms are not plain bf16 RMSNorms of the hidden width or whose output and down projections are not row parallel with their all-reduce. test_modeling_v2_kimi_k3_drafter.py (sm_100), on a one-rank collective state whose two collectives are counted torch stand-ins: a block of every certified split, with the torch GEMMs and with the decode GEMV sites, matches the stock block forward (rel L2 <= 1e-2), with the sandwiches up to 8 rows and the fused all-reduce above, two per layer; the block's input and the cache stay untouched; both sandwich forms compiled; the stock norms run where the workspace does not hold the rows or the norms do not take the fused form, and a form that did not compile gives way to the fused all-reduce; a captured block makes the same calls and replays the eager result; the stock drafter makes no collective call. test_kimi_k3_drafter_comm.py (cpu_only): which norms and projections take the fused all-reduces; use_decode_comm compiles each shared form once with a zero row of a zero weight, skips a form the layers do not share or the kernel does not take, and compiles nothing without the fused norms; sandwich_plain / takes_plain / compile_plain hand the catalog entry its arguments; skip_all_reduce turns a module's all-reduce off. The tp16_moetp16ep1 copy does not carry this change yet; until it does, test_modeling_v2_kimi_k3_drift fails on decode_comm.py and modeling.py. Signed-off-by: Vasanth Sabavat --- .../decode_comm.py | 70 +++++- .../modeling.py | 217 +++++++++++++++--- .../test_modeling_v2_kimi_k3_drafter.py | 148 +++++++++++- .../hw_agnostic/test_kimi_k3_drafter_comm.py | 204 +++++++++++++++- 4 files changed, 596 insertions(+), 43 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py index 88311a249d60..dbe64ba7a513 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py @@ -18,9 +18,13 @@ all-reduce and the next layer's residual update (the final norm's, after the last layer) in one kernel. The DSpark drafter (`K3DSparkDrafter`) runs its collectives with a plain residual add and RMSNorm on the same state: -`K3DecodeComm.allreduce_norm`, `comm/mnnvl_fusion_allreduce`, all-reduces an unreduced projection output of any token -count the MNNVL workspace holds, with the residual add and the RMSNorm in its epilogue (the context projection's, -with a zero residual). + +* `K3DecodeComm.sandwich_plain`, `comm/k3_sandwich_plain`: a row-parallel projection of at most `SANDWICH_MAX_TOKENS` + tokens (a layer's attention output projection, or its MLP's SiLU-and-mul and down projection), its all-reduce, the + residual add and the RMSNorm in one kernel. +* `K3DecodeComm.allreduce_norm`, `comm/mnnvl_fusion_allreduce`: the all-reduce of an unreduced projection output of + any token count the MNNVL workspace holds, with the residual add and the RMSNorm in its epilogue (the context + projection's with a zero residual). The state is one `MnnvlWorkspace` and one `K3SandwichWorkspace` of the TP group (`K3DecodeComm.create`): collective over the group and eager, built by the target in `post_load_weights` before any CUDA-graph capture. Every rank must @@ -43,6 +47,9 @@ K3SandwichWorkspace, k3_sandwich_oproj, ) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.k3_sandwich_plain import ( + k3_sandwich_plain, +) from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.k3_sandwich_tail import ( k3_sandwich_tail, ) @@ -87,6 +94,12 @@ def use_decode_one_shot(model: nn.Module) -> None: mnnvl.one_shot_max_bytes = DECODE_AR_ONE_SHOT_MAX_BYTES +def skip_all_reduce() -> AllReduceParams: + """All-reduce parameters under which a row-parallel module returns this rank's unreduced output (its projection + without its all-reduce): a ``Linear``'s ``all_reduce_params``, a ``GatedMLP``'s ``final_all_reduce_params``.""" + return AllReduceParams(enable_allreduce=False) + + def wide_all_reduce(all_reduce: nn.Module, x: torch.Tensor) -> torch.Tensor: """``all_reduce(x)`` of a wide decode step (a stock ``AllReduce`` module, no fusion): its MNNVL all-reduce with the `WIDE_AR_ONE_SHOT_MAX_BYTES` ceiling, else the module itself.""" @@ -303,6 +316,57 @@ def sandwich_oproj( self.sandwich, ) + def compile_plain(self, weight: torch.Tensor, swiglu: bool = False) -> bool: + """Compile the plain sandwich (``sandwich_plain``) for a projection of ``weight``'s shape (with ``swiglu``, a + down projection after the SiLU-and-mul) with one call on a zero row of a zero weight, before any capture; + False, with nothing launched, where the kernel does not take that shape. Collective: every rank of the group + makes the call; it advances the sandwich workspace on every rank alike.""" + zero = torch.zeros_like(weight) + hidden, k_in = zero.shape + x = zero.new_zeros(1, 2 * k_in if swiglu else k_in) + residual = zero.new_zeros(1, hidden) + ones = zero.new_ones(hidden) + if not _sandwich_op.supports_plain(x, zero, residual, ones, swiglu): + return False + k3_sandwich_plain(x, zero, residual, ones, 1e-6, self.sandwich, swiglu=swiglu) + torch.cuda.synchronize(zero.device) + return True + + @staticmethod + def takes_plain( + x: torch.Tensor, + weight: torch.Tensor, + residual: torch.Tensor, + norm: nn.Module, + swiglu: bool = False, + ) -> bool: + """Whether ``sandwich_plain`` takes the call: at most `SANDWICH_MAX_TOKENS` contiguous bf16 rows ``x`` of a + row-parallel slice ``weight`` [7168, K] (K a multiple of 128 up to 896; with ``swiglu`` K 896 and ``x`` the + ``[gate | up]`` rows, 2 K wide), ``residual`` [rows, 7168], and ``norm``'s bf16 [7168] weight.""" + return _sandwich_op.supports_plain(x, weight, residual, norm.weight, swiglu) + + def sandwich_plain( + self, + x: torch.Tensor, + weight: torch.Tensor, + residual: torch.Tensor, + norm: nn.Module, + swiglu: bool = False, + ) -> Tuple[torch.Tensor, torch.Tensor]: + """``(normed, updated)`` in one ``comm/k3_sandwich_plain`` call: ``updated = residual + allreduce(x @ + weight.T)`` (with ``swiglu``, of ``silu_and_mul(x) @ weight.T``), ``normed = norm(updated)`` for a plain + RMSNorm ``norm``. Bit for bit, by the kernel's statement, the projection on ``k3_ctm_gemv`` at split 1 (with + ``swiglu``, ``k3_ctm_gemv_swiglu`` at split 2) followed by ``allreduce_norm``'s one-shot call.""" + return k3_sandwich_plain( + x.contiguous(), + weight, + residual.contiguous(), + norm.weight, + _eps(norm), + self.sandwich, + swiglu=swiglu, + ) + def takes_allreduce_norm(self, rows: int, hidden: int) -> bool: """Whether ``allreduce_norm`` takes ``rows`` bf16 rows of ``hidden`` columns: the MNNVL workspace holds the call at the decode path's one-shot ceiling (`DECODE_AR_ONE_SHOT_MAX_BYTES`, two-shot above it).""" diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index ac7ad94127b5..b9b1fa9549f0 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -73,8 +73,9 @@ block on the drafter entries (`attention/k3_drafter_attn_qknorm` and the drafter's decode GEMV sites); the shell builds it through `_build_draft_model`. Where every attention all-reduce runs over MNNVL, the drafter also runs on the TP group's collective state: its context projection's `fc` is split by input feature over the group, `hidden_norm` -applied in the all-reduce of the partial products (`comm/mnnvl_fusion_allreduce`). The worker and its kernels stay -upstream code; this target does not own a worker. +applied in the all-reduce of the partial products (`comm/mnnvl_fusion_allreduce`), and a block's residual adds and +RMSNorms run in its all-reduces (`comm/k3_sandwich_plain` with the projection up to 8 rows, else +`comm/mnnvl_fusion_allreduce`). The worker and its kernels stay upstream code; this target does not own a worker. """ from __future__ import annotations @@ -190,11 +191,12 @@ "k3_moe", "k3_route_quant", "mnnvl_allgather_split", - # The DSpark drafter's decode path (K3DSparkDrafter): its GEMV sites, its block attention and the all-reduce of - # its split context projection with hidden_norm. + # The DSpark drafter's decode path (K3DSparkDrafter): its GEMV sites, its block attention, and its all-reduces + # with the residual add and RMSNorm (the split context projection's with hidden_norm). "k3_ctm_gemv", "k3_ctm_gemv_swiglu", "k3_drafter_attn_qknorm", + "k3_sandwich_plain", "mnnvl_fusion_allreduce", ) @@ -2454,15 +2456,26 @@ class K3DSparkDrafter(GQADSparkForCausalLM): * gate / up on ``drafter_gate_up`` and the down projection with the SiLU-and-mul on ``drafter_down``, then the module's all-reduce. - The norms and the residual adds are the stock modules'; a projection whose site does not take its rows runs its - module. Every other block runs the stock ``dflash_forward``: a split `DRAFTER_ATTN_SPLITS` does not list, another - attention backend, head layout, RoPE or normalization, a cache the kernel does not read, or a compile key of the - attention that has not run eagerly, under CUDA-graph capture. - - Where the target hands the drafter the TP group's collective state (``use_decode_comm``), the context projection's - ``fc`` holds only this rank's block of input columns (`K3FcSlice`), and ``project_target_hidden`` sums the ranks' - partial products with ``hidden_norm`` in one all-reduce; otherwise the context projection is the stock one, with - ``fc`` replicated. The worker that calls the drafter (the stock ``DSparkWorker``), the context k / v and the Markov + The norms and the residual adds are the stock modules' unless the drafter runs on the TP group's collective state + (below); a projection whose site does not take its rows runs its module. Every other block runs the stock + ``dflash_forward``: a split `DRAFTER_ATTN_SPLITS` does not list, another attention backend, head layout, RoPE or + normalization, a cache the kernel does not read, or a compile key of the attention that has not run eagerly, under + CUDA-graph capture. + + Where the target hands the drafter the TP group's collective state (``use_decode_comm``): + + * the context projection's ``fc`` holds only this rank's block of input columns (`K3FcSlice`), and + ``project_target_hidden`` sums the ranks' partial products with ``hidden_norm`` in one all-reduce; + * a block the entries take runs each residual add and RMSNorm in the all-reduce before it (o_proj's with the + post-attention norm, the MLP's with the next layer's input norm, the last one's with the final norm), the + block's input serving as the first residual, uncopied. Up to 8 rows ``comm/k3_sandwich_plain`` runs the + projection, its all-reduce and the norm in one launch (o_proj, and with the SiLU-and-mul the down projection + after ``drafter_gate_up``); otherwise the projection (its site, or its module without the all-reduce) is + followed by ``comm/mnnvl_fusion_allreduce``. A block whose rows the MNNVL workspace does not hold, or layers + whose norms or projections the fused all-reduces do not reproduce, keep the stock all-reduces and norms. + + Otherwise ``fc`` stays replicated, with the stock context projection, and the norms and residual adds are the + stock modules'. The worker that calls the drafter (the stock ``DSparkWorker``), the context k / v and the Markov head stay upstream code. """ @@ -2481,17 +2494,80 @@ def __init__( self.decode_comm: Optional[_decode_comm.K3DecodeComm] = None # The zero residual of the split context projection's all-reduce, up to a decode step's rows. self._k3_zero_rows: Optional[torch.Tensor] = None + # Whether a block's residual adds and RMSNorms run in its all-reduces, and the comm/k3_sandwich_plain forms + # ("o_proj", "down") compiled for the layers' shapes; set with the collective state. + self._k3_norms_fuse = False + self._k3_sandwich_forms: frozenset = frozenset() def use_decode_comm(self, comm: _decode_comm.K3DecodeComm) -> None: """Run the drafter's collectives on ``comm``, the TP group's decode state (``decode_comm.K3DecodeComm``) the - target builds where every attention all-reduce runs over MNNVL. The target calls it on every rank once the - weights are loaded, before any CUDA-graph capture. - - The context projection's ``fc`` then keeps only this rank's block of input columns (`K3FcSlice`), where they - split evenly over the group and the drafter's layers have their TP all-reduce: ``project_target_hidden`` sums - the ranks' partial products and applies ``hidden_norm`` in one ``comm/mnnvl_fusion_allreduce`` call.""" + target builds where every attention all-reduce runs over MNNVL. The target calls it once the weights are + loaded, before any CUDA-graph capture. Collective: every rank of the group calls it at the same point. + + * The context projection's ``fc`` keeps only this rank's block of input columns (`K3FcSlice`), where they + split evenly over the group and the drafter's layers have their TP all-reduce: ``project_target_hidden`` + sums the ranks' partial products and applies ``hidden_norm`` in one ``comm/mnnvl_fusion_allreduce`` call. + * Where the layers' norms and projections take it (`_k3_norms_take_comm`), a block's residual adds and + RMSNorms run in its all-reduces (``_k3_block_forward``). Each form of ``comm/k3_sandwich_plain`` whose + shape every layer shares compiles here, with one call on a zero row of a zero weight; a form that does not + compile is not used.""" self.decode_comm = comm self._k3_split_fc() + self._k3_norms_fuse = self._k3_norms_take_comm() + forms = set() + if self._k3_norms_fuse: + layers = self.model.layers + for form, swiglu, weights in ( + ("o_proj", False, [layer.self_attn.o_proj.weight for layer in layers]), + ("down", True, [layer.mlp.down_proj.weight for layer in layers]), + ): + if len({tuple(w.shape) for w in weights}) == 1 and comm.compile_plain( + weights[0], swiglu=swiglu + ): + forms.add(form) + self._k3_sandwich_forms = frozenset(forms) + logger.info( + "Kimi K3 DSpark drafter: residual adds and RMSNorms " + + ( + f"in the all-reduces (k3_sandwich_plain: {sorted(forms) or 'none'}, else " + "mnnvl_fusion_allreduce)" + if self._k3_norms_fuse + else "on the stock modules (a layer norm or projection the fused all-reduces do not reproduce)" + ) + ) + + def _k3_norms_take_comm(self) -> bool: + """Whether a block's residual adds and RMSNorms can run in its all-reduces: every layer's input and + post-attention norms and the final norm are plain bf16 RMSNorms of the hidden width (no Gemma offset, no + quantized output), and every output and down projection is row parallel with its all-reduce.""" + hidden = self.config.hidden_size + + def plain(norm: nn.Module) -> bool: + weight = getattr(norm, "weight", None) + return ( + isinstance(weight, torch.Tensor) + and weight.dtype == torch.bfloat16 + and tuple(weight.shape) == (hidden,) + and hasattr(norm, "variance_epsilon") + and not getattr(norm, "use_gemma", True) + and not getattr(norm, "is_nvfp4", True) + and not getattr(norm, "return_hp_output", True) + ) + + def reduces(linear: nn.Module) -> bool: + return ( + getattr(getattr(linear, "tp_mode", None), "name", None) == "ROW" + and getattr(linear, "reduce_output", False) + and getattr(linear, "all_reduce", None) is not None + ) + + return plain(self.model.norm) and all( + plain(layer.input_layernorm) + and plain(layer.post_attention_layernorm) + and reduces(layer.self_attn.o_proj) + and reduces(layer.mlp.down_proj) + for layer in self.model.layers + ) def _k3_tp_all_reduce(self) -> Optional[nn.Module]: """The drafter's own TP all-reduce (its first output projection's row-parallel all-reduce module), or None.""" @@ -2705,14 +2781,22 @@ def _k3_block_forward( page_table = ctx_page_table.index_select(0, ctx_cache_batch_idx.to(torch.long)) positions = query_positions.reshape(-1).contiguous() hidden = noise_embedding.reshape(rows, -1) + layers = self.model.layers + comm = self._k3_fused_comm(hidden) residual = None - for layer_idx, layer in enumerate(self.model.layers): + if comm is not None: + # The fused all-reduces read the residual and return the updated one: the block's input serves as the + # first residual, uncopied. + residual = hidden + normed = layers[0].input_layernorm(hidden) + for layer_idx, layer in enumerate(layers): attn = layer.self_attn - if residual is None: - residual = hidden.clone() - normed = layer.input_layernorm(hidden) - else: - normed, residual = layer.input_layernorm(hidden, residual) + if comm is None: + if residual is None: + residual = hidden.clone() + normed = layer.input_layernorm(hidden) + else: + normed, residual = layer.input_layernorm(hidden, residual) qkv = self._k3_project("drafter_qkv", normed, attn.qkv_proj) out = torch.empty(rows, attn.q_size, dtype=torch.bfloat16, device=qkv.device) k3_drafter_attn_qknorm( @@ -2729,12 +2813,88 @@ def _k3_block_forward( attn.num_key_value_heads, out, ) + if comm is not None: + # o_proj's all-reduce applies the post-attention norm; the MLP's the next layer's input norm, after + # the last layer the final norm. + next_norm = ( + layers[layer_idx + 1].input_layernorm + if layer_idx + 1 < len(layers) + else self.model.norm + ) + normed, residual = self._k3_project_norm( + comm, out, attn.o_proj, residual, layer.post_attention_layernorm + ) + normed, residual = self._k3_mlp_norm(comm, layer.mlp, normed, residual, next_norm) + continue hidden = self._k3_project("drafter_o", out, attn.o_proj) hidden, residual = layer.post_attention_layernorm(hidden, residual) hidden = self._k3_mlp(layer.mlp, hidden) + if comm is not None: + return normed out, _ = self.model.norm(hidden, residual) return out + def _k3_fused_comm(self, hidden: torch.Tensor) -> Optional[_decode_comm.K3DecodeComm]: + """The TP group's collective state where a block of rows ``hidden`` runs its residual adds and RMSNorms in its + all-reduces: the drafter runs on it (``use_decode_comm``) with layers that take it, and its MNNVL workspace + holds an all-reduce of the block's rows; else None (the stock all-reduces and norms).""" + comm = self.decode_comm + if comm is None or not self._k3_norms_fuse or not comm.takes_allreduce_norm(*hidden.shape): + return None + return comm + + def _k3_project_norm( + self, + comm: _decode_comm.K3DecodeComm, + x: torch.Tensor, + linear: nn.Module, + residual: torch.Tensor, + norm: nn.Module, + ) -> Tuple[torch.Tensor, torch.Tensor]: + """``(norm(updated), updated)``, ``updated = residual + linear(x)`` for the row-parallel output projection, + the residual add and ``norm`` in its all-reduce: ``comm/k3_sandwich_plain`` where it takes the rows (the + projection in the ``drafter_o`` site's arithmetic, in the same launch), else the projection on ``drafter_o`` + where that takes the rows, or the module without its all-reduce, then ``comm/mnnvl_fusion_allreduce``.""" + if "o_proj" in self._k3_sandwich_forms and comm.takes_plain( + x, linear.weight, residual, norm + ): + return comm.sandwich_plain(x, linear.weight, residual, norm) + gemvs = self.decode_gemvs + partial = None if gemvs is None else gemvs.project("drafter_o", x, linear.weight) + if partial is None: + partial = linear(x, all_reduce_params=_decode_comm.skip_all_reduce()) + return comm.allreduce_norm(partial, residual, norm) + + def _k3_mlp_norm( + self, + comm: _decode_comm.K3DecodeComm, + mlp: nn.Module, + x: torch.Tensor, + residual: torch.Tensor, + norm: nn.Module, + ) -> Tuple[torch.Tensor, torch.Tensor]: + """``(norm(updated), updated)``, ``updated = residual + mlp(x)``, the residual add and ``norm`` in the down + projection's all-reduce: gate / up on ``drafter_gate_up``, then ``comm/k3_sandwich_plain``'s SiLU-and-mul + form where it takes the rows (the ``drafter_down`` site's arithmetic), else the SiLU-and-mul and down + projection on ``drafter_down``; where those sites do not take the rows, the module without its all-reduce; + then ``comm/mnnvl_fusion_allreduce``.""" + gemvs = self.decode_gemvs + gate_up = ( + None if gemvs is None else gemvs.project("drafter_gate_up", x, mlp.gate_up_proj.weight) + ) + partial = None + if gate_up is not None: + if "down" in self._k3_sandwich_forms and comm.takes_plain( + gate_up, mlp.down_proj.weight, residual, norm, swiglu=True + ): + return comm.sandwich_plain( + gate_up, mlp.down_proj.weight, residual, norm, swiglu=True + ) + partial = gemvs.project("drafter_down", gate_up, mlp.down_proj.weight) + if partial is None: + partial = mlp(x, final_all_reduce_params=_decode_comm.skip_all_reduce()) + return comm.allreduce_norm(partial, residual, norm) + def _k3_project(self, site: str, x: torch.Tensor, linear: nn.Module) -> torch.Tensor: """``linear(x)``, its GEMM on the ``site`` decode GEMV where that takes the rows, then a row-parallel projection's all-reduce.""" @@ -3113,9 +3273,10 @@ def _gate_spec_worker_kernels(self, comm: Optional[_decode_comm.K3DecodeComm]) - def _gate_drafter_comm(self, comm: Optional[_decode_comm.K3DecodeComm]) -> bool: """Hand the DSpark drafter (`K3DSparkDrafter`) the TP group's collective state ``comm`` where this target built it (every attention all-reduce over MNNVL; TP16 is a construction assert): the drafter then runs its context - projection split over the group, ``hidden_norm`` in the all-reduce (``K3DSparkDrafter.use_decode_comm``). - Without ``comm`` it keeps the stock replicated ``fc``. Returns whether the drafter took the state (False - without such a drafter).""" + projection split over the group, ``hidden_norm`` in the all-reduce, and its blocks' residual adds and RMSNorms + in their all-reduces (``K3DSparkDrafter.use_decode_comm``, collective: it compiles the drafter's sandwich on + the group's workspace). Without ``comm`` it keeps the stock replicated ``fc``, all-reduces and norms. Returns + whether the drafter took the state (False without such a drafter).""" drafter = getattr(self, "draft_model", None) if comm is None or not isinstance(drafter, K3DSparkDrafter): return False diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py index 7ecc2b53232d..e0e26a6bdfbc 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py @@ -18,6 +18,14 @@ * The context projection runs on the split ``fc`` (on one rank, the whole weight): up to a decode step's rows through ``comm/mnnvl_fusion_allreduce`` with ``hidden_norm``, more rows through the drafter's TP all-reduce, and matches the replicated projection. +* A decode block of every certified split runs its residual adds and RMSNorms in its all-reduces and matches the + stock block forward: up to 8 rows ``comm/k3_sandwich_plain`` for o_proj (and, on the decode GEMV sites, its + SiLU-and-mul form for the down projection), else ``comm/mnnvl_fusion_allreduce``, two per layer, the last with the + final norm. Both sandwich forms compiled with the state. +* The stock all-reduces and norms run without the collective state, where the workspace does not hold the block's + rows, or where the layers' norms do not take the fused form; a sandwich form that did not compile gives way to + ``comm/mnnvl_fusion_allreduce``. +* A fused block captured in a CUDA graph makes the same collective calls and replays the eager result. """ import math @@ -151,19 +159,44 @@ def _fusion_allreduce_one_rank( return _rms_norm(updated, norm_weight, eps), updated +def _sandwich_plain_one_rank(x, weight, residual, norm_weight, eps, workspace, swiglu=False): + """``comm/k3_sandwich_plain`` over one rank: the projection (of ``silu(gate) * up`` with ``swiglu``), then the + residual add and the RMSNorm.""" + assert workspace is ONE_RANK_COMM.sandwich + k = weight.shape[1] + assert ( + x.shape == (residual.shape[0], 2 * k if swiglu else k) + and residual.shape[1] == weight.shape[0] + ) + assert x.is_contiguous() and residual.is_contiguous() and isinstance(eps, float) + a = x.float() + if swiglu: + a = (torch.nn.functional.silu(a[:, :k]) * a[:, k:]).to(torch.bfloat16).float() + partial = (a @ weight.float().T).to(torch.bfloat16) + updated = (residual.float() + partial.float()).to(torch.bfloat16) + return _rms_norm(updated, norm_weight, eps), updated + + def _one_rank_collectives(monkeypatch, calls): - """The drafter's collectives replaced by their one-rank stand-ins, each call recorded as (op, rows).""" + """The drafter's collectives replaced by their one-rank stand-ins, each call recorded as (op, rows) or, for the + sandwich, (op, rows, swiglu).""" def fusion(input, *args, **kwargs): calls.append(("mnnvl_fusion_allreduce", input.shape[0])) return _fusion_allreduce_one_rank(input, *args, **kwargs) + def sandwich(x, *args, swiglu=False): + calls.append(("k3_sandwich_plain", x.shape[0], swiglu)) + return _sandwich_plain_one_rank(x, *args, swiglu=swiglu) + monkeypatch.setattr(decode_comm, "mnnvl_fusion_allreduce", fusion) + monkeypatch.setattr(decode_comm, "k3_sandwich_plain", sandwich) @pytest.fixture(scope="module") def fused_drafter(): - """The drafter on the one-rank collective state: its fc split over one rank.""" + """The drafter on the one-rank collective state: its fc split over one rank, both sandwich forms compiled (their + compile calls run on the stand-in).""" with pytest.MonkeyPatch.context() as mp: _one_rank_collectives(mp, []) module = _load_drafter() @@ -238,7 +271,9 @@ def counted(*args): @pytest.mark.parametrize("use_gemvs", [False, True], ids=["torch", "gemv"]) @pytest.mark.parametrize("split", sorted(target.DRAFTER_ATTN_SPLITS)) -def test_block_matches_the_stock_forward(drafter, gemvs, attn_calls, monkeypatch, use_gemvs, split): +def test_block_matches_the_stock_forward( + drafter, gemvs, attn_calls, collectives, monkeypatch, use_gemvs, split +): batch, block = split projected = [] project = gemvs.project @@ -261,6 +296,8 @@ def recorded(site, x, weight): err = _rel_l2(out, ref) assert err <= REL_L2, err assert all(torch.equal(c, b) for c, b in zip(inputs["ctx_kv_cache"], before)) + # Without the TP group's collective state the module all-reduces and the stock norms run. + assert collectives == [] if use_gemvs: # At most 8 rows every site takes its projection; above, each declines and its module runs (the down # projection's site is not asked once gate / up's declined). @@ -333,3 +370,108 @@ def test_split_fc_matches_the_replicated_projection(drafter, fused_drafter, coll err = _rel_l2(out, ref) assert err <= REL_L2, err assert collectives == ([("mnnvl_fusion_allreduce", rows)] if rows <= DECODE_ROWS else []) + + +def test_fused_drafter_compiled_both_sandwich_forms(fused_drafter): + assert fused_drafter.decode_comm is ONE_RANK_COMM + assert fused_drafter._k3_norms_fuse + assert fused_drafter._k3_sandwich_forms == {"o_proj", "down"} + + +def _fused_calls(rows, use_gemvs): + """One layer's collective calls on the fused path: up to 8 rows the o_proj sandwich, and the down projection's + sandwich where the gate / up site produced its input; else the projection then the fused all-reduce.""" + if rows > decode_gemv.MAX_ROWS: + return [("mnnvl_fusion_allreduce", rows)] * 2 + down = ("k3_sandwich_plain", rows, True) if use_gemvs else ("mnnvl_fusion_allreduce", rows) + return [("k3_sandwich_plain", rows, False), down] + + +@pytest.mark.parametrize("use_gemvs", [False, True], ids=["torch", "gemv"]) +@pytest.mark.parametrize("split", sorted(target.DRAFTER_ATTN_SPLITS)) +def test_fused_block_matches_the_stock_forward( + fused_drafter, gemvs, attn_calls, collectives, use_gemvs, split +): + batch, block = split + rows = batch * block + fused_drafter.decode_gemvs = gemvs if use_gemvs else None + inputs = _block(batch, block, seed=batch * 10 + block + 1) + stock_inputs = _copy(inputs) + before = [c.clone() for c in inputs["ctx_kv_cache"]] + noise = inputs["noise_embedding"].clone() + out = fused_drafter.dflash_forward(**inputs) + calls = list(collectives) + ref = DFlashForCausalLM.dflash_forward(fused_drafter, **stock_inputs) + torch.cuda.synchronize() + assert attn_calls == [rows] * LAYERS + assert calls == _fused_calls(rows, use_gemvs) * LAYERS + assert out.shape == ref.shape == (rows, HIDDEN) + err = _rel_l2(out, ref) + assert err <= REL_L2, err + # The block's input served as the first residual, read only; the cache is untouched. + assert torch.equal(inputs["noise_embedding"], noise) + assert all(torch.equal(c, b) for c, b in zip(inputs["ctx_kv_cache"], before)) + + +def _stock_norms_case(fused_drafter, monkeypatch, case): + if case == "workspace": + small = decode_comm.K3DecodeComm( + mnnvl=SimpleNamespace(world_size=1, buffer_bytes=1024), + sandwich=ONE_RANK_COMM.sandwich, + ) + monkeypatch.setattr(fused_drafter, "decode_comm", small) + else: + monkeypatch.setattr(fused_drafter, "_k3_norms_fuse", False) + + +@pytest.mark.parametrize("case", ["workspace", "norms"]) +def test_fused_norms_fall_back_to_the_stock_norms( + fused_drafter, gemvs, attn_calls, collectives, monkeypatch, case +): + """A block whose rows the MNNVL workspace does not hold, or layers whose norms the fused all-reduces do not + reproduce, keep the module all-reduces and the stock norms.""" + _stock_norms_case(fused_drafter, monkeypatch, case) + fused_drafter.decode_gemvs = gemvs + inputs = _block(1, 7, seed=17) + stock_inputs = _copy(inputs) + out = fused_drafter.dflash_forward(**inputs) + ref = DFlashForCausalLM.dflash_forward(fused_drafter, **stock_inputs) + torch.cuda.synchronize() + assert collectives == [] + assert attn_calls == [7] * LAYERS + err = _rel_l2(out, ref) + assert err <= REL_L2, err + + +def test_uncompiled_sandwich_gives_way_to_the_fused_all_reduce( + fused_drafter, gemvs, collectives, monkeypatch +): + monkeypatch.setattr(fused_drafter, "_k3_sandwich_forms", frozenset({"down"})) + fused_drafter.decode_gemvs = gemvs + inputs = _block(1, 8, seed=88) + stock_inputs = _copy(inputs) + out = fused_drafter.dflash_forward(**inputs) + ref = DFlashForCausalLM.dflash_forward(fused_drafter, **stock_inputs) + torch.cuda.synchronize() + assert collectives == [("mnnvl_fusion_allreduce", 8), ("k3_sandwich_plain", 8, True)] * LAYERS + err = _rel_l2(out, ref) + assert err <= REL_L2, err + + +@pytest.mark.parametrize("split", [(1, 7), (2, 7)]) +def test_fused_block_replays_under_capture(fused_drafter, gemvs, collectives, split): + """Eager first (the attention compiles for the block's key), then captured and replayed.""" + fused_drafter.decode_gemvs = gemvs + inputs = _block(*split, seed=70 + split[0]) + eager = fused_drafter.dflash_forward(**_copy(inputs)) + torch.cuda.synchronize() + eager_calls = list(collectives) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + captured = fused_drafter.dflash_forward(**inputs) + graph.replay() + torch.cuda.synchronize() + assert eager_calls == _fused_calls(split[0] * split[1], True) * LAYERS + assert collectives[len(eager_calls) :] == eager_calls + err = _rel_l2(captured, eager) + assert err <= 1e-3, err diff --git a/tests/unittest/_torch/speculative/hw_agnostic/test_kimi_k3_drafter_comm.py b/tests/unittest/_torch/speculative/hw_agnostic/test_kimi_k3_drafter_comm.py index ca88d0906f2f..4334a1349157 100644 --- a/tests/unittest/_torch/speculative/hw_agnostic/test_kimi_k3_drafter_comm.py +++ b/tests/unittest/_torch/speculative/hw_agnostic/test_kimi_k3_drafter_comm.py @@ -26,6 +26,11 @@ * ``_gate_drafter_comm`` hands the state to a ``K3DSparkDrafter`` only, and only where the target built it. * ``K3DecodeComm.takes_allreduce_norm`` / ``allreduce_norm``: which calls the TP16 MNNVL workspace holds, and the catalog call's arguments. +* The fused norms: ``_k3_norms_take_comm`` holds for plain bf16 RMSNorms of the hidden width and row-parallel output + and down projections with their all-reduce, and for nothing else; ``use_decode_comm`` compiles each + ``comm/k3_sandwich_plain`` form whose shape every layer shares (one zero-row call each, on the group's sandwich + workspace), and none without the fused norms; ``K3DecodeComm.compile_plain`` / ``takes_plain`` / + ``sandwich_plain`` hand the catalog entry its arguments; ``skip_all_reduce`` turns a module's all-reduce off. """ from types import SimpleNamespace @@ -40,11 +45,14 @@ from tensorrt_llm._torch._experimental.modeling_v2.models.kimi_k3_vl.kimi_k3_mxfp4__sm_100__tp16_moetp4ep4 import ( # noqa: E501 modeling as target, ) +from tensorrt_llm._torch.modules.rms_norm import RMSNorm pytestmark = pytest.mark.cpu_only -# A small context projection: 4 captured layers of 32 features into 32 (K3's is 5 x 7168 into 7168). +# A small context projection: 4 captured layers of 32 features into 32 (K3's is 5 x 7168 into 7168), and the input +# widths of the layers' output and down projections. IN, OUT = 128, 32 +K_O, K_DOWN = 16, 24 DECODE_ROWS = target.MAX_REQUESTS * target.MAX_TOKENS_PER_REQUEST @@ -82,25 +90,49 @@ def _comm(world=1, buffer_bytes=4 << 20): ) +def _norm(**kwargs): + """A stock RMSNorm of the hidden width.""" + kwargs.setdefault("dtype", torch.bfloat16) + return RMSNorm(hidden_size=kwargs.pop("hidden_size", OUT), eps=1e-6, **kwargs) + + +def _projection(k, all_reduce=True, row_parallel=True, reduce_output=True): + """A row-parallel projection's fields: its mode, its all-reduce and its [OUT, k] slice.""" + return SimpleNamespace( + tp_mode=SimpleNamespace(name="ROW" if row_parallel else "COLUMN"), + reduce_output=reduce_output, + all_reduce=_Sum() if all_reduce else None, + weight=torch.full((OUT, k), 0.5, dtype=torch.bfloat16), + ) + + +def _layer(all_reduce=True, row_parallel=True): + return SimpleNamespace( + input_layernorm=_norm(), + post_attention_layernorm=_norm(), + self_attn=SimpleNamespace(o_proj=_projection(K_O, all_reduce, row_parallel)), + mlp=SimpleNamespace(down_proj=_projection(K_DOWN)), + ) + + def _drafter(tp_size=1, tp_rank=0, comm=None, fc=None, all_reduce=True, row_parallel=True): - """A K3DSparkDrafter without ``__init__`` (that needs a GPU drafter checkpoint): the fields the context projection - reads.""" + """A K3DSparkDrafter without ``__init__`` (that needs a GPU drafter checkpoint): the fields its context projection + and fused norms read, two layers.""" drafter = target.K3DSparkDrafter.__new__(target.K3DSparkDrafter) nn.Module.__init__(drafter) drafter.model_config = SimpleNamespace( mapping=SimpleNamespace(tp_size=tp_size, tp_rank=tp_rank) ) - o_proj = SimpleNamespace( - tp_mode=SimpleNamespace(name="ROW" if row_parallel else "COLUMN"), - all_reduce=_Sum() if all_reduce else None, - ) + drafter.config = SimpleNamespace(hidden_size=OUT) drafter.model = SimpleNamespace( - layers=[SimpleNamespace(self_attn=SimpleNamespace(o_proj=o_proj))] + layers=[_layer(all_reduce, row_parallel) for _ in range(2)], norm=_norm() ) drafter.fc = _fc() if fc is None else fc drafter.hidden_norm = _hidden_norm() drafter.decode_comm = comm drafter._k3_zero_rows = None + drafter._k3_norms_fuse = False + drafter._k3_sandwich_forms = frozenset() return drafter @@ -128,6 +160,36 @@ def one_rank(input, workspace, one_shot_max_bytes, residual=None, norm_weight=No return calls +@pytest.fixture +def sandwich_calls(monkeypatch): + """Records ``comm/k3_sandwich_plain`` calls (each returns its residual twice) and lets the kernel's shape check + pass on CPU tensors unless ``unsupported`` holds the call's weight width; device syncs are no-ops.""" + calls = [] + unsupported = set() + + def sandwich(x, weight, residual, norm_weight, eps, workspace, swiglu=False): + calls.append( + dict( + x=x, + weight=weight, + residual=residual, + norm_weight=norm_weight, + eps=eps, + workspace=workspace, + swiglu=swiglu, + ) + ) + return residual, residual + + def supports(x, weight, residual, norm_weight, swiglu=False): + return weight.shape[1] not in unsupported + + monkeypatch.setattr(decode_comm, "k3_sandwich_plain", sandwich) + monkeypatch.setattr(decode_comm._sandwich_op, "supports_plain", supports) + monkeypatch.setattr(torch.cuda, "synchronize", lambda *args, **kwargs: None) + return SimpleNamespace(calls=calls, unsupported=unsupported) + + @pytest.mark.parametrize("tp_size", [1, 2, 4, 16]) def test_fc_columns_tile_the_inputs_in_rank_order(tp_size): in_features = 5 * 7168 @@ -311,7 +373,7 @@ def test_gate_hands_the_state_to_the_k3_drafter_only(): assert taken == [comm] -def test_use_decode_comm_keeps_the_state_and_splits_fc(): +def test_use_decode_comm_keeps_the_state_and_splits_fc(sandwich_calls): comm = _comm(world=2) drafter = _drafter(tp_size=2, tp_rank=0) drafter.use_decode_comm(comm) @@ -348,3 +410,127 @@ def test_allreduce_norm_passes_the_workspace_ceiling_and_norm(fused_calls, eps, stock = SimpleNamespace(variance_epsilon=1e-6, weight=norm.weight) comm.allreduce_norm(partial, residual, stock) assert fused_calls[1]["eps"] == pytest.approx(1e-6) + + +def test_norms_take_the_fused_all_reduces(): + assert _drafter()._k3_norms_take_comm() + + +@pytest.mark.parametrize( + "case", + [ + "gemma input norm", + "fp32 post-attention norm", + "final norm of another width", + "nvfp4 norm output", + "high-precision norm output", + "o_proj without its all-reduce", + "column-parallel o_proj", + "down projection that does not reduce", + ], +) +def test_norms_do_not_take_the_fused_all_reduces(case): + drafter = _drafter() + layer = drafter.model.layers[1] + if case == "gemma input norm": + layer.input_layernorm = _norm(use_gemma=True) + elif case == "fp32 post-attention norm": + layer.post_attention_layernorm = _norm(dtype=torch.float32) + elif case == "final norm of another width": + drafter.model.norm = _norm(hidden_size=OUT + 8) + elif case == "nvfp4 norm output": + layer.input_layernorm.is_nvfp4 = True + elif case == "high-precision norm output": + layer.post_attention_layernorm.return_hp_output = True + elif case == "o_proj without its all-reduce": + layer.self_attn.o_proj = _projection(K_O, all_reduce=False) + elif case == "column-parallel o_proj": + layer.self_attn.o_proj = _projection(K_O, row_parallel=False) + else: + layer.mlp.down_proj = _projection(K_DOWN, reduce_output=False) + assert not drafter._k3_norms_take_comm() + + +def test_use_decode_comm_compiles_each_sandwich_form(sandwich_calls): + comm = _comm(world=16) + drafter = _drafter(tp_size=16, tp_rank=3) + drafter.use_decode_comm(comm) + assert drafter._k3_norms_fuse + assert drafter._k3_sandwich_forms == {"o_proj", "down"} + assert isinstance(drafter.fc, target.K3FcSlice) + assert (drafter.fc.start, drafter.fc.end) == (24, 32) + plain, swiglu = sandwich_calls.calls + assert not plain["swiglu"] and swiglu["swiglu"] + assert plain["x"].shape == (1, K_O) and plain["weight"].shape == (OUT, K_O) + assert swiglu["x"].shape == (1, 2 * K_DOWN) and swiglu["weight"].shape == (OUT, K_DOWN) + for call in (plain, swiglu): + # One zero row of a zero weight on the group's sandwich workspace: it compiles the kernel and adds nothing. + assert call["workspace"] is comm.sandwich + assert call["residual"].shape == (1, OUT) and call["norm_weight"].shape == (OUT,) + assert not call["x"].any() and not call["weight"].any() and not call["residual"].any() + + +def test_use_decode_comm_skips_a_form_the_layers_do_not_share(sandwich_calls): + drafter = _drafter() + drafter.model.layers[1].self_attn.o_proj = _projection(2 * K_O) + drafter.use_decode_comm(_comm()) + assert drafter._k3_sandwich_forms == {"down"} + assert [call["swiglu"] for call in sandwich_calls.calls] == [True] + + +def test_use_decode_comm_skips_a_form_the_kernel_does_not_take(sandwich_calls): + sandwich_calls.unsupported.add(K_DOWN) + drafter = _drafter() + drafter.use_decode_comm(_comm()) + assert drafter._k3_norms_fuse and drafter._k3_sandwich_forms == {"o_proj"} + assert [call["swiglu"] for call in sandwich_calls.calls] == [False] + + +def test_use_decode_comm_without_the_fused_norms_compiles_nothing(sandwich_calls): + drafter = _drafter() + drafter.model.norm = _norm(use_gemma=True) + drafter.use_decode_comm(_comm()) + assert not drafter._k3_norms_fuse and drafter._k3_sandwich_forms == frozenset() + assert sandwich_calls.calls == [] + # The context projection's split does not depend on the norms. + assert isinstance(drafter.fc, target.K3FcSlice) + + +def test_compile_plain_declines_a_shape_the_kernel_does_not_take(sandwich_calls): + sandwich_calls.unsupported.add(K_O) + assert not _comm().compile_plain(torch.ones(OUT, K_O, dtype=torch.bfloat16)) + assert sandwich_calls.calls == [] + + +def test_sandwich_plain_passes_its_arguments(sandwich_calls): + comm = _comm(world=16) + norm = _norm() + x = torch.ones(3, 2 * K_DOWN, dtype=torch.bfloat16) + weight = torch.ones(OUT, K_DOWN, dtype=torch.bfloat16) + residual = torch.zeros(3, OUT, dtype=torch.bfloat16) + comm.sandwich_plain(x, weight, residual, norm, swiglu=True) + (call,) = sandwich_calls.calls + assert call["x"] is x and call["weight"] is weight and call["residual"] is residual + assert call["norm_weight"] is norm.weight and call["eps"] == pytest.approx(1e-6) + assert call["swiglu"] and call["workspace"] is comm.sandwich + + +def test_takes_plain_asks_the_kernel_with_the_norm_weight(monkeypatch): + asked = [] + monkeypatch.setattr( + decode_comm._sandwich_op, + "supports_plain", + lambda *args: asked.append(args) or True, + ) + norm = _norm() + x, weight, residual = torch.ones(2, K_O), torch.ones(OUT, K_O), torch.zeros(2, OUT) + assert decode_comm.K3DecodeComm.takes_plain(x, weight, residual, norm, swiglu=True) + (args,) = asked + assert all(a is b for a, b in zip(args, (x, weight, residual, norm.weight))) + assert args[4] is True + + +def test_skip_all_reduce_turns_the_modules_all_reduce_off(): + params = decode_comm.skip_all_reduce() + assert params.enable_allreduce is False and params.residual is None + assert decode_comm.skip_all_reduce() is not params From 00e43940392d59dc044883f42e59aac4b1e0ef10 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 11:51:09 -0700 Subject: [PATCH 139/161] [None][perf] DFlash / Kimi K3 DSpark drafter: the gen requests' page-table rows as a view On every draft step the Kimi K3 drafter's block forward gathered its requests' rows of the context page table with index_select, though the worker hands it the manager's block table, keyed by batch position, where the gen requests' rows are always the contiguous run [num_contexts, num_contexts + num_gens). - DFlashWorker.prepare_1st_drafter_inputs now also returns that run's start, ctx_rows_start (_ctx_rows_start): num_contexts with the manager's block table and gen requests; None for the private arena, whose table is keyed by slot. - The worker passes it to dflash_forward only for a drafter that takes the keyword (dflash_ctx_rows_kwargs, the signature checked once per class), so the stock drafters are called as before. - K3DSparkDrafter.dflash_forward takes ctx_rows_start, and its block forward reads the rows as a view of the table (the attention entry takes rows at any row stride), one gather kernel fewer per draft step. Without it, it gathers as before. test_dflash_ctx_rows_start.py (cpu_only): the start is num_contexts for the manager's table and None for the private arena or without gen requests; only K3DSparkDrafter receives it, not the stock DFlash, GQA or MLA DSpark drafters. test_modeling_v2_kimi_k3_drafter.py (sm_100): with the batch's rows inside a larger table, the view's block output equals the gather's bit for bit, and the attention entry reads the table's own rows instead of a copy. The tp16_moetp16ep1 copy does not carry this change yet; until it does, test_modeling_v2_kimi_k3_drift fails on modeling.py. Signed-off-by: Vasanth Sabavat --- .../modeling.py | 13 +++- tensorrt_llm/_torch/speculative/dflash.py | 36 +++++++++ .../test_modeling_v2_kimi_k3_drafter.py | 34 +++++++++ .../hw_agnostic/test_dflash_ctx_rows_start.py | 73 +++++++++++++++++++ 4 files changed, 154 insertions(+), 2 deletions(-) create mode 100644 tests/unittest/_torch/speculative/hw_agnostic/test_dflash_ctx_rows_start.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index b9b1fa9549f0..9ceea1273845 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -2646,9 +2646,12 @@ def dflash_forward( ctx_cache_batch_idx: torch.Tensor, ctx_kv_cache: Optional[torch.Tensor] = None, ctx_page_table: Optional[torch.Tensor] = None, + ctx_rows_start: Optional[int] = None, ) -> torch.Tensor: """The block's hidden states ``[B * block, hidden]``: on the drafter entries where they take it (the class - docstring), else the stock forward.""" + docstring), else the stock forward. ``ctx_rows_start``: where ``ctx_cache_batch_idx`` is the contiguous rows + ``[ctx_rows_start, ctx_rows_start + B)`` of ``ctx_page_table``, their start; the entries then read those rows + as a view instead of gathering them.""" keys = self._k3_block_keys(noise_embedding, ctx_kv_cache, ctx_page_table) if keys is None: return super().dflash_forward( @@ -2668,6 +2671,7 @@ def dflash_forward( ctx_cache_batch_idx, ctx_kv_cache, ctx_page_table, + ctx_rows_start, ) if not torch.cuda.is_current_stream_capturing(): self._k3_attn_ran |= keys @@ -2774,11 +2778,16 @@ def _k3_block_forward( ctx_cache_batch_idx: torch.Tensor, ctx_kv_cache: torch.Tensor, ctx_page_table: torch.Tensor, + ctx_rows_start: Optional[int] = None, ) -> torch.Tensor: batch, block = noise_embedding.shape[:2] rows = batch * block ctx_len = num_ctx_per_req[:batch].to(torch.int32) - page_table = ctx_page_table.index_select(0, ctx_cache_batch_idx.to(torch.long)) + if ctx_rows_start is None: + page_table = ctx_page_table.index_select(0, ctx_cache_batch_idx.to(torch.long)) + else: + # The batch's page-table rows are one contiguous run: a view, no gather. + page_table = ctx_page_table[ctx_rows_start : ctx_rows_start + batch] positions = query_positions.reshape(-1).contiguous() hidden = noise_embedding.reshape(rows, -1) layers = self.model.layers diff --git a/tensorrt_llm/_torch/speculative/dflash.py b/tensorrt_llm/_torch/speculative/dflash.py index 032c7aa920aa..0c5709df8205 100644 --- a/tensorrt_llm/_torch/speculative/dflash.py +++ b/tensorrt_llm/_torch/speculative/dflash.py @@ -13,6 +13,8 @@ # See the License for the specific language governing permissions and # limitations under the License. +import functools +import inspect import os from collections import deque from dataclasses import dataclass @@ -496,6 +498,24 @@ def dflash_allocated_ctx_limit( return (block_counts * page_size - block_size).clamp(min=0) +@functools.lru_cache(maxsize=None) +def _takes_ctx_rows_start_cls(cls) -> bool: + return "ctx_rows_start" in inspect.signature(cls.dflash_forward).parameters + + +def _takes_ctx_rows_start(draft_model) -> bool: + """Whether the draft model's ``dflash_forward`` accepts ``ctx_rows_start``.""" + return _takes_ctx_rows_start_cls(type(draft_model)) + + +def dflash_ctx_rows_kwargs(draft_model, ctx_rows_start: Optional[int]) -> dict: + """The ``ctx_rows_start`` keyword (from ``prepare_1st_drafter_inputs``) for a draft model whose ``dflash_forward`` + takes it; no keyword for another draft model or without a start.""" + if ctx_rows_start is None or not _takes_ctx_rows_start(draft_model): + return {} + return {"ctx_rows_start": ctx_rows_start} + + @dataclass class DFlashSpecMetadata(SpecMetadata): """Metadata for DFlash speculative decoding. @@ -2185,6 +2205,7 @@ def _forward_impl( ctx_cache_batch_idx=inputs["ctx_cache_batch_idx"], ctx_kv_cache=inputs["ctx_kv_cache"], ctx_page_table=inputs["ctx_page_table"], + **dflash_ctx_rows_kwargs(draft_model, inputs["ctx_rows_start"]), ) # Gather K logits per gen request from the block outputs. @@ -2468,6 +2489,17 @@ def _dflash2_global_top_k( block_logits = gen_logits.new_full((*gen_logits.shape[:-1], full_vocab), float("-inf")) return candidate_ids, unary_logits, block_logits + def _ctx_rows_start(self, num_contexts: int, num_gens: int) -> Optional[int]: + """The first row of the gen requests' page-table rows where those rows are one contiguous run, else None. + + The manager's block table is keyed by batch position, so the gen requests' rows are [num_contexts, + num_contexts + num_gens): a drafter can view them instead of gathering ``ctx_cache_batch_idx``'s rows. The + private arena's table is keyed by slot, which need not be contiguous. + """ + if num_gens > 0 and self._ctx_block_tables is not None: + return num_contexts + return None + def prepare_1st_drafter_inputs( self, input_ids: torch.LongTensor, @@ -2492,6 +2524,9 @@ def prepare_1st_drafter_inputs( - num_ctx_per_req: per-request context length in the pool - ctx_k_cache / ctx_v_cache / ctx_cache_batch_idx: slot-indexed views of the persistent per-layer K/V pool. + - ctx_rows_start: where ctx_cache_batch_idx is the contiguous rows + [ctx_rows_start, ctx_rows_start + num_gens) of ctx_page_table, their + start (``_ctx_rows_start``), else None. """ num_contexts = attn_metadata.num_contexts batch_size = attn_metadata.num_seqs @@ -2686,4 +2721,5 @@ def prepare_1st_drafter_inputs( if self._ctx_block_tables is not None else self._ctx_page_table ), + "ctx_rows_start": self._ctx_rows_start(num_contexts, num_gens), } diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py index e0e26a6bdfbc..8bc1d83d9569 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py @@ -11,6 +11,8 @@ * A split the entry does not certify runs the stock forward, bit for bit, without the entry. * Under CUDA-graph capture, the entries take a block only once its attention compile key has run eagerly. * Negative control: a weight changed between the two forwards fails the comparison. +* With ``ctx_rows_start`` the entries read the batch's page-table rows as a view of a larger table (the manager's, + rows from the first gen request on), not a gathered copy, and give the gather's result bit for bit. On the TP group's collective state (``use_decode_comm``), here a group of one rank whose collectives are torch stand-ins counted per call (the ops themselves are certified by their multi-GPU op matrices): @@ -354,6 +356,38 @@ def test_negative_control_a_changed_weight_fails(drafter): assert _rel_l2(out, ref) > REL_L2 +@pytest.mark.parametrize("split", [(1, 7), (3, 1), (8, 7)]) +def test_page_table_rows_as_a_view(drafter, gemvs, monkeypatch, split): + batch, block = split + start = 2 + tables = [] + entry = target.k3_drafter_attn_qknorm + + def recorded(*args): + tables.append(args[7]) # the page table + return entry(*args) + + monkeypatch.setattr(target, "k3_drafter_attn_qknorm", recorded) + drafter.decode_gemvs = gemvs + inputs = _block(batch, block, seed=batch * 10 + block + 2) + # The batch's rows sit at [start, start + batch) of a larger table; the other rows hold no page. + table = inputs["ctx_page_table"] + padded = torch.full((start + batch + 1, table.shape[1]), -1, dtype=torch.int32, device="cuda") + padded[start : start + batch] = table + inputs["ctx_page_table"] = padded + inputs["ctx_cache_batch_idx"] = torch.arange( + start, start + batch, dtype=torch.long, device="cuda" + ) + gathered = drafter.dflash_forward(**_copy(inputs)) + viewed = drafter.dflash_forward(**inputs, ctx_rows_start=start) + torch.cuda.synchronize() + assert torch.equal(viewed, gathered) + assert len(tables) == 2 * LAYERS and all(torch.equal(t, table) for t in tables) + # The gather's rows are a copy; the view's are the table's own rows. + assert all(t.data_ptr() != padded[start].data_ptr() for t in tables[:LAYERS]) + assert all(t.data_ptr() == padded[start].data_ptr() for t in tables[LAYERS:]) + + @pytest.mark.parametrize("rows", [1, 8, DECODE_ROWS, 200]) def test_split_fc_matches_the_replicated_projection(drafter, fused_drafter, collectives, rows): """On one rank the block is the whole fc. Up to a decode step's rows the projection's all-reduce applies diff --git a/tests/unittest/_torch/speculative/hw_agnostic/test_dflash_ctx_rows_start.py b/tests/unittest/_torch/speculative/hw_agnostic/test_dflash_ctx_rows_start.py new file mode 100644 index 000000000000..96a86dbfdf2e --- /dev/null +++ b/tests/unittest/_torch/speculative/hw_agnostic/test_dflash_ctx_rows_start.py @@ -0,0 +1,73 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# 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. +"""``ctx_rows_start``: the DFlash worker's start of the gen requests' page-table rows, and who receives it +(host-side). + +* ``DFlashWorker._ctx_rows_start`` is the number of context requests where the drafter reads the manager's block + table (keyed by batch position, so the gen requests' rows are one run), and None for the private arena (keyed by + slot) or a step without gen requests. +* ``dflash_ctx_rows_kwargs`` hands it only to a drafter whose ``dflash_forward`` takes it (the Kimi K3 target's + ``K3DSparkDrafter``), never to the stock drafters, and nothing where there is no start. +""" + +import pytest +import torch +from torch import nn + +from tensorrt_llm._torch._experimental.modeling_v2.models.kimi_k3_vl.kimi_k3_mxfp4__sm_100__tp16_moetp4ep4.modeling import ( # noqa: E501 + K3DSparkDrafter, +) +from tensorrt_llm._torch.models.modeling_dflash import DFlashForCausalLM +from tensorrt_llm._torch.models.modeling_dspark import GQADSparkForCausalLM, MLADSparkForCausalLM +from tensorrt_llm._torch.speculative.dflash import DFlashWorker, dflash_ctx_rows_kwargs +from tensorrt_llm._torch.speculative.dspark import DSparkWorker + +pytestmark = pytest.mark.cpu_only + + +def _bare(cls): + """An instance without ``__init__`` (the workers need flashinfer, the drafters a checkpoint).""" + obj = cls.__new__(cls) + nn.Module.__init__(obj) + return obj + + +@pytest.mark.parametrize("cls", [DFlashWorker, DSparkWorker]) +@pytest.mark.parametrize("num_contexts,num_gens", [(0, 1), (0, 8), (3, 2), (5, 1)]) +def test_rows_start_where_the_manager_table_is_read(cls, num_contexts, num_gens): + worker = _bare(cls) + worker._ctx_block_tables = torch.zeros(num_contexts + num_gens + 1, 4, dtype=torch.int32) + assert worker._ctx_rows_start(num_contexts, num_gens) == num_contexts + + +@pytest.mark.parametrize("cls", [DFlashWorker, DSparkWorker]) +def test_no_rows_start_for_the_private_arena_or_without_gen_requests(cls): + worker = _bare(cls) + worker._ctx_block_tables = None + assert worker._ctx_rows_start(0, 4) is None + worker._ctx_block_tables = torch.zeros(4, 4, dtype=torch.int32) + assert worker._ctx_rows_start(3, 0) is None + + +def test_the_k3_drafter_takes_the_rows_start(): + drafter = _bare(K3DSparkDrafter) + assert dflash_ctx_rows_kwargs(drafter, 3) == {"ctx_rows_start": 3} + assert dflash_ctx_rows_kwargs(drafter, 0) == {"ctx_rows_start": 0} + assert dflash_ctx_rows_kwargs(drafter, None) == {} + + +@pytest.mark.parametrize("cls", [DFlashForCausalLM, GQADSparkForCausalLM, MLADSparkForCausalLM]) +def test_stock_drafters_do_not_get_the_rows_start(cls): + assert dflash_ctx_rows_kwargs(_bare(cls), 3) == {} From ac6e2876f5a83f9487fa6aaf999dd2a5e4422e81 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 11:59:36 -0700 Subject: [PATCH 140/161] [None][doc] modeling_v2 Kimi K3 tp16_moetp4ep4: the drafter gate requires only its comm _gate_drafter_comm's docstring states that the gate requires exactly what the drafter's fused path uses, the TP group's collective state, and not the LM head's k3_head_gemv workspace, which only the speculative worker's path reads. Signed-off-by: Vasanth Sabavat --- .../kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index 9ceea1273845..5fc908147bee 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -3284,8 +3284,10 @@ def _gate_drafter_comm(self, comm: Optional[_decode_comm.K3DecodeComm]) -> bool: it (every attention all-reduce over MNNVL; TP16 is a construction assert): the drafter then runs its context projection split over the group, ``hidden_norm`` in the all-reduce, and its blocks' residual adds and RMSNorms in their all-reduces (``K3DSparkDrafter.use_decode_comm``, collective: it compiles the drafter's sandwich on - the group's workspace). Without ``comm`` it keeps the stock replicated ``fc``, all-reduces and norms. Returns - whether the drafter took the state (False without such a drafter).""" + the group's workspace). Without ``comm`` it keeps the stock replicated ``fc``, all-reduces and norms. The gate + requires exactly what this path uses: ``comm`` and nothing else, so not the LM head's ``k3_head_gemv`` + workspace, which only the speculative worker's path (`_gate_spec_worker_kernels`) reads. Returns whether the + drafter took the state (False without such a drafter).""" drafter = getattr(self, "draft_model", None) if comm is None or not isinstance(drafter, K3DSparkDrafter): return False From 147b87d0816916d847ab79cf1bb89d4b4cf17c30 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 12:00:16 -0700 Subject: [PATCH 141/161] [None][test] k3_ctx_kv test: drop the base-port identity cases, zero the Python path's dummy query test_batch1_identity compared the kernel with an unmodified package named by K3_BASE_TRTLLM, which CI never sets, so its 9 cases always skipped there. They and their helpers (base_kernel, base_call, the timing table's base column) are removed. Evidence of the identity they checked: run against the base port's installed package (921f6e4229), all 9 cases were bit-identical. The Python reference path's dummy query is allocated with new_zeros instead of new_empty: flashinfer's RoPE reads it, and initcheck reported ~12k uninitialized reads there. The reference results are unchanged (72 / 72 with a zero-filled copy). Signed-off-by: Vasanth Sabavat --- .../kimi_k3/test_k3_ctx_kv.py | 118 ++---------------- 1 file changed, 13 insertions(+), 105 deletions(-) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctx_kv.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctx_kv.py index 177009b19ba4..eb344eee6c8f 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctx_kv.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctx_kv.py @@ -26,13 +26,9 @@ reference (the kernel's error must be the Python path's), masked rows zero, the pool's other elements untouched, ctx_len and num_ctx exact, reruns bit-identical; CUDA-graph replays with rewritten inputs. -Batch-1 identity against the unmodified kernel (N <= 8): set ``K3_BASE_TRTLLM`` to an unmodified ``tensorrt_llm`` -package directory. Timing: ``python3 test_k3_ctx_kv.py time [--base]``; error table: ``python3 test_k3_ctx_kv.py -report``. +Timing: ``python3 test_k3_ctx_kv.py time``; error table: ``python3 test_k3_ctx_kv.py report``. """ -import importlib.util -import os import statistics import sys @@ -122,7 +118,7 @@ def python_path( k = F.rms_norm(k, (HEAD,), eps=EPS) k = k * k_norm.view(1, LAYERS, 1, HEAD) pos = cpos.reshape(-1).to(torch.int32).repeat_interleave(LAYERS) - dummy_q = k.new_empty(n * LAYERS, HEAD) + dummy_q = k.new_zeros(n * LAYERS, HEAD) rope(pos, dummy_q, k.view(n * LAYERS, nkv * HEAD), HEAD, cs, True) offs = torch.arange(k1, device="cuda") mask = (offs[None, :] < num_acc.long()[:, None]).reshape(-1).view(-1, 1, 1, 1).to(k.dtype) @@ -463,92 +459,7 @@ def test_index_offset(arg): # ---------------------------------------------------------------------------------------------------------------- -# Batch-1 identity against the unmodified kernel (K3_BASE_TRTLLM: an unmodified tensorrt_llm package directory). -# ---------------------------------------------------------------------------------------------------------------- - -_base = {} - - -def base_kernel(): - root = os.environ.get("K3_BASE_TRTLLM") - if not root: - return None - if "mod" not in _base: - path = os.path.join(root, "_torch", "cute_dsl_kernels", "k3_ctx_kv", "k3_ctx_kv_kernel.py") - spec = importlib.util.spec_from_file_location("k3_ctx_kv_kernel_base", path) - mod = importlib.util.module_from_spec(spec) - spec.loader.exec_module(mod) - _base["mod"] = mod - return _base["mod"] - - -def base_call( - x, w, k_norm, cs, cpos, num_acc, ctx_len, slots, rows, table, counts, layers, block_size, nkv -): - """The unmodified op (N <= 8) on the same arguments.""" - import cuda.bindings.driver as cuda_driver - import cutlass.cute as cute - from cutlass.cute.runtime import from_dlpack - - kern = base_kernel() - - def arg(t): - return from_dlpack(t.detach(), assumed_align=16).mark_layout_dynamic( - leading_dim=t.dim() - 1 - ) - - flat, layer_off, ps, kvs, hs = pool_view(layers) - batch, k1 = cpos.shape - n_tokens, k_in = x.shape - n_rows = w.shape[0] - sms = torch.cuda.get_device_properties(x.device).multi_processor_count - split = next(s for s in (8, 4, 2) if (n_rows // kern.CTA_M) * s <= sms - and kern.supports(n_rows, k_in, s, nkv, k1, n_tokens)) # fmt: skip - ring = kern.pick_ring(k_in, split) - if "counter" not in _base: - _base["counter"] = torch.zeros(1, dtype=torch.int32, device="cuda") - num_ctx = torch.empty(batch, dtype=torch.int32, device="cuda") - args = (arg(w), arg(x), arg(k_norm.reshape(-1)), arg(cs.reshape(-1)), arg(cpos.reshape(-1)), arg(num_acc), - arg(ctx_len), arg(slots), arg(rows), arg(table.reshape(-1)), arg(counts), arg(flat), arg(layer_off), - arg(num_ctx), arg(_base["counter"])) # fmt: skip - scalars = ( - float(EPS), - int(MAX_CTX), - int(PAGE), - int(block_size), - int(table.stride(0)), - int(ps), - int(kvs), - int(hs), - ) - consts = (n_rows, k_in, split, ring, nkv, k1, n_tokens, True) - stream = cuda_driver.CUstream(torch.cuda.current_stream().cuda_stream) - fn = _base.get(consts) - if fn is None: - fn = _base[consts] = cute.compile(kern.k3_ctx_kv, *args, *scalars, *consts, True, stream) - fn(*args, *scalars, stream) - return num_ctx - - -@pytest.mark.skipif( - not os.environ.get("K3_BASE_TRTLLM"), reason="K3_BASE_TRTLLM (unmodified package) not set" -) -@pytest.mark.parametrize("batch,k1", [s for s in SPLITS if s[0] * s[1] <= 8], ids=[f"{b}x{k}" for b, k in SPLITS - if b * k <= 8]) # fmt: skip -def test_batch1_identity(batch, k1): - with torch.inference_mode(): - for seed, (style, clamp) in enumerate((("v1", False), ("arena", False), ("v1", True))): - gen = torch.Generator(device="cuda").manual_seed(55 + 10 * batch + k1 + seed) - st = Step(gen, 1, batch, k1, style, clamp, seed) - buf_k, ctx_k, nc_k, _ = st.run_kernel() - buf_b, ctx_b, nc_b, _ = st.run_kernel(fn=base_call) - torch.cuda.synchronize() - assert torch.equal(buf_k.view(torch.int16), buf_b.view(torch.int16)), (style, clamp) - assert torch.equal(ctx_k, ctx_b) and torch.equal(nc_k, nc_b), (style, clamp) - - -# ---------------------------------------------------------------------------------------------------------------- -# Timing (python3 test_k3_ctx_kv.py time [--base]) and the error table (report) +# Timing (python3 test_k3_ctx_kv.py time) and the error table (report) # ---------------------------------------------------------------------------------------------------------------- @@ -577,7 +488,7 @@ def time_graph(body, calls, replays=15): return statistics.median(per_call), min(per_call), max(per_call) -def timing(with_base: bool) -> None: +def timing() -> None: """Graphs of back-to-back calls with the weight rotating over 160 MB of copies (HBM-cold), TP16.""" _op() gen = torch.Generator(device="cuda").manual_seed(11) @@ -590,8 +501,8 @@ def timing(with_base: bool) -> None: calls = 2 * copies print(f"{torch.cuda.get_device_name()}; graphs of {calls} calls, weights rotating over {copies} copies, " "15 replays: median (min-max) us per call") # fmt: skip - print("| split | N | k3_ctx_kv | base |") - print("| :-- | --: | --: | --: |") + print("| split | N | k3_ctx_kv |") + print("| :-- | --: | --: |") with torch.inference_mode(): for b, k1 in SPLITS: st = Step(gen, 1, b, k1, "v1", seed=5) @@ -599,21 +510,18 @@ def timing(with_base: bool) -> None: ctx = st.ctx0.clone() arms = [lambda i, fn=fn: fn(st.x, ws[i % copies], st.k_norm, st.cs, st.cpos, st.num_acc, ctx, st.slots, st.rows, st.table, st.counts, st.layers, BLOCK, 1) - for fn in ([kernel_call, base_call] if with_base and b * k1 <= 8 else [kernel_call])] # fmt: skip + for fn in (kernel_call,)] # fmt: skip res = [[] for _ in arms] for rep in range(3): for a in range(len(arms)) if rep % 2 == 0 else reversed(range(len(arms))): ctx.copy_(st.ctx0) res[a].append(time_graph(arms[a], calls)) cells = [] - for a in range(2): - if a < len(arms): - meds = sorted(x[0] for x in res[a]) - cells.append( - f"{meds[1]:.2f} ({min(x[1] for x in res[a]):.2f}-{max(x[2] for x in res[a]):.2f})" - ) - else: - cells.append("") + for a in range(len(arms)): + meds = sorted(x[0] for x in res[a]) + cells.append( + f"{meds[1]:.2f} ({min(x[1] for x in res[a]):.2f}-{max(x[2] for x in res[a]):.2f})" + ) print(f"| {b}x{k1} | {b * k1} | " + " | ".join(cells) + " |", flush=True) @@ -643,7 +551,7 @@ def report() -> int: if __name__ == "__main__": if len(sys.argv) > 1 and sys.argv[1] == "time": - timing("--base" in sys.argv) + timing() elif len(sys.argv) > 1 and sys.argv[1] == "report": sys.exit(report()) else: From b7ea3f6c3a4cecac25f6d32ecc3bdeeabdbfbeb0 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 12:01:56 -0700 Subject: [PATCH 142/161] [None][chore] modeling_v2 Kimi K3 tp16_moetp16ep1: carry the DSpark drafter's fc split and fused all-reduces into the copy The tp16_moetp4ep4 target's DSpark drafter now splits its fc over the TP group and runs its residual adds and RMSNorms inside its all-reduces (K3DecodeComm's mnnvl_fusion_allreduce and k3_sandwich_plain forms), and reads the generation requests' page-table rows as a view. The copy here takes those changes: decode_comm.py byte for byte, and the changes to modeling.py outside route B's blocks. Its module docstring keeps route B's own block text: this target decodes without speculation. This target never builds a drafter, so none of the carried drafter code runs here; the drift test requires it to be present. Signed-off-by: Vasanth Sabavat --- .../decode_comm.py | 107 +++++- .../modeling.py | 344 +++++++++++++++++- 2 files changed, 432 insertions(+), 19 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_comm.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_comm.py index 2ef7dbe27458..dbe64ba7a513 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_comm.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_comm.py @@ -17,6 +17,15 @@ the same kind of collective: `K3DecodeComm.sandwich_tail`, `comm/k3_sandwich_tail`, runs the tail GEMV, its all-reduce and the next layer's residual update (the final norm's, after the last layer) in one kernel. +The DSpark drafter (`K3DSparkDrafter`) runs its collectives with a plain residual add and RMSNorm on the same state: + +* `K3DecodeComm.sandwich_plain`, `comm/k3_sandwich_plain`: a row-parallel projection of at most `SANDWICH_MAX_TOKENS` + tokens (a layer's attention output projection, or its MLP's SiLU-and-mul and down projection), its all-reduce, the + residual add and the RMSNorm in one kernel. +* `K3DecodeComm.allreduce_norm`, `comm/mnnvl_fusion_allreduce`: the all-reduce of an unreduced projection output of + any token count the MNNVL workspace holds, with the residual add and the RMSNorm in its epilogue (the context + projection's with a zero residual). + The state is one `MnnvlWorkspace` and one `K3SandwichWorkspace` of the TP group (`K3DecodeComm.create`): collective over the group and eager, built by the target in `post_load_weights` before any CUDA-graph capture. Every rank must make the same calls on each in the same order. Which call a step takes is decided from its token count and kind and @@ -38,12 +47,19 @@ K3SandwichWorkspace, k3_sandwich_oproj, ) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.k3_sandwich_plain import ( + k3_sandwich_plain, +) from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.k3_sandwich_tail import ( k3_sandwich_tail, ) from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.mnnvl_allreduce_attn_res import ( mnnvl_allreduce_attn_res, ) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.mnnvl_fusion_allreduce import ( + mnnvl_fusion_allreduce, + required_buffer_bytes, +) from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.mnnvl_workspace import ( MnnvlWorkspace, ) @@ -78,6 +94,12 @@ def use_decode_one_shot(model: nn.Module) -> None: mnnvl.one_shot_max_bytes = DECODE_AR_ONE_SHOT_MAX_BYTES +def skip_all_reduce() -> AllReduceParams: + """All-reduce parameters under which a row-parallel module returns this rank's unreduced output (its projection + without its all-reduce): a ``Linear``'s ``all_reduce_params``, a ``GatedMLP``'s ``final_all_reduce_params``.""" + return AllReduceParams(enable_allreduce=False) + + def wide_all_reduce(all_reduce: nn.Module, x: torch.Tensor) -> torch.Tensor: """``all_reduce(x)`` of a wide decode step (a stock ``AllReduce`` module, no fusion): its MNNVL all-reduce with the `WIDE_AR_ONE_SHOT_MAX_BYTES` ceiling, else the module itself.""" @@ -104,8 +126,10 @@ class PendingTail(NamedTuple): def _eps(norm: nn.Module) -> float: - """The epsilon of a KimiK3RMSNorm (``eps``) or a stock RMSNorm (``variance_epsilon``).""" - return float(norm.eps if hasattr(norm, "eps") else norm.variance_epsilon) + """The epsilon of a KimiK3RMSNorm or ``torch.nn.RMSNorm`` (``eps``; for torch's None, its default: the machine + epsilon of the weight's dtype) or a stock RMSNorm (``variance_epsilon``).""" + eps = norm.eps if hasattr(norm, "eps") else norm.variance_epsilon + return float(torch.finfo(norm.weight.dtype).eps if eps is None else eps) def _res_args(res_proj: nn.Module, res_norm: nn.Module, out_norm: nn.Module) -> tuple: @@ -291,3 +315,82 @@ def sandwich_oproj( *_res_args(res_proj, res_norm, out_norm), self.sandwich, ) + + def compile_plain(self, weight: torch.Tensor, swiglu: bool = False) -> bool: + """Compile the plain sandwich (``sandwich_plain``) for a projection of ``weight``'s shape (with ``swiglu``, a + down projection after the SiLU-and-mul) with one call on a zero row of a zero weight, before any capture; + False, with nothing launched, where the kernel does not take that shape. Collective: every rank of the group + makes the call; it advances the sandwich workspace on every rank alike.""" + zero = torch.zeros_like(weight) + hidden, k_in = zero.shape + x = zero.new_zeros(1, 2 * k_in if swiglu else k_in) + residual = zero.new_zeros(1, hidden) + ones = zero.new_ones(hidden) + if not _sandwich_op.supports_plain(x, zero, residual, ones, swiglu): + return False + k3_sandwich_plain(x, zero, residual, ones, 1e-6, self.sandwich, swiglu=swiglu) + torch.cuda.synchronize(zero.device) + return True + + @staticmethod + def takes_plain( + x: torch.Tensor, + weight: torch.Tensor, + residual: torch.Tensor, + norm: nn.Module, + swiglu: bool = False, + ) -> bool: + """Whether ``sandwich_plain`` takes the call: at most `SANDWICH_MAX_TOKENS` contiguous bf16 rows ``x`` of a + row-parallel slice ``weight`` [7168, K] (K a multiple of 128 up to 896; with ``swiglu`` K 896 and ``x`` the + ``[gate | up]`` rows, 2 K wide), ``residual`` [rows, 7168], and ``norm``'s bf16 [7168] weight.""" + return _sandwich_op.supports_plain(x, weight, residual, norm.weight, swiglu) + + def sandwich_plain( + self, + x: torch.Tensor, + weight: torch.Tensor, + residual: torch.Tensor, + norm: nn.Module, + swiglu: bool = False, + ) -> Tuple[torch.Tensor, torch.Tensor]: + """``(normed, updated)`` in one ``comm/k3_sandwich_plain`` call: ``updated = residual + allreduce(x @ + weight.T)`` (with ``swiglu``, of ``silu_and_mul(x) @ weight.T``), ``normed = norm(updated)`` for a plain + RMSNorm ``norm``. Bit for bit, by the kernel's statement, the projection on ``k3_ctm_gemv`` at split 1 (with + ``swiglu``, ``k3_ctm_gemv_swiglu`` at split 2) followed by ``allreduce_norm``'s one-shot call.""" + return k3_sandwich_plain( + x.contiguous(), + weight, + residual.contiguous(), + norm.weight, + _eps(norm), + self.sandwich, + swiglu=swiglu, + ) + + def takes_allreduce_norm(self, rows: int, hidden: int) -> bool: + """Whether ``allreduce_norm`` takes ``rows`` bf16 rows of ``hidden`` columns: the MNNVL workspace holds the + call at the decode path's one-shot ceiling (`DECODE_AR_ONE_SHOT_MAX_BYTES`, two-shot above it).""" + if rows <= 0 or hidden <= 0 or hidden % 8: + return False + world, buffer_bytes = self.mnnvl.world_size, self.mnnvl.buffer_bytes + need = required_buffer_bytes( + rows, hidden, world, torch.bfloat16, DECODE_AR_ONE_SHOT_MAX_BYTES + ) + two_shot = rows * hidden * world * 2 > DECODE_AR_ONE_SHOT_MAX_BYTES + return need <= buffer_bytes and not (two_shot and buffer_bytes % 32) + + def allreduce_norm( + self, partial: torch.Tensor, residual: torch.Tensor, norm: nn.Module + ) -> Tuple[torch.Tensor, torch.Tensor]: + """``(normed, updated)`` in one ``comm/mnnvl_fusion_allreduce`` call: ``updated = residual + + allreduce(partial)``, ``normed = norm(updated)`` for a plain RMSNorm ``norm``, sent one-shot up to + `DECODE_AR_ONE_SHOT_MAX_BYTES`. ``partial`` is this rank's unreduced ``[rows, hidden]`` bf16 output of a + row-parallel projection, of a shape ``takes_allreduce_norm`` holds for.""" + return mnnvl_fusion_allreduce( + partial.contiguous(), + self.mnnvl, + DECODE_AR_ONE_SHOT_MAX_BYTES, + residual.contiguous(), + norm.weight, + _eps(norm), + ) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py index da7d9bb0d78c..7dd3ff5b634a 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py @@ -188,10 +188,13 @@ "k3_moe", "k3_route_quant", "mnnvl_allgather_split", - # The DSpark drafter's decode path (K3DSparkDrafter): its GEMV sites and its block attention. + # The DSpark drafter's decode path (K3DSparkDrafter): its GEMV sites, its block attention, and its all-reduces + # with the residual add and RMSNorm (the split context projection's with hidden_norm). "k3_ctm_gemv", "k3_ctm_gemv_swiglu", "k3_drafter_attn_qknorm", + "k3_sandwich_plain", + "mnnvl_fusion_allreduce", ) #: The engine surface the first forward checks before this target relies on it: per object, the attributes read. @@ -234,8 +237,8 @@ "tensorrt_llm._torch.distributed.AllReduce", "tensorrt_llm._torch.modules.multi_stream_utils.maybe_execute_in_parallel", # The DSpark drafter: the stock GQA drafter K3DSparkDrafter extends (its block forward where the drafter entries - # do not take a block, its context projection and context k / v, its heads and its weight load), and the stock - # builder's checks for which drafter a checkpoint gets. + # do not take a block, its context projection without the TP group's collective state, its context k / v, its + # heads and its weight load), and the stock builder's checks for which drafter a checkpoint gets. "tensorrt_llm._torch.models.modeling_dspark.GQADSparkForCausalLM", "tensorrt_llm._torch.models.modeling_dspark.draft_is_embedded_in_target", "tensorrt_llm._torch.models.modeling_dflash.DFlashForCausalLM", @@ -2314,6 +2317,42 @@ def _k3_decode_view(self, attn_metadata: AttentionMetadata, num_tokens: int) -> _DRAFTER_ROPE_BASE = 10000.0 +def fc_columns(in_features: int, tp_size: int, tp_rank: int) -> Optional[Tuple[int, int]]: + """Rank ``tp_rank``'s input columns ``[start, end)`` of the drafter's context projection ``fc`` split by input + feature over ``tp_size`` ranks: equal contiguous blocks in rank order. None where they do not split evenly.""" + if tp_size < 1 or not 0 <= tp_rank < tp_size or in_features <= 0 or in_features % tp_size: + return None + width = in_features // tp_size + return tp_rank * width, (tp_rank + 1) * width + + +class K3FcSlice(nn.Module): + """This rank's block of the DSpark drafter's context projection ``fc`` (`fc_columns`): ``weight`` is + ``fc.weight[:, start:end]``, contiguous. Its output is this rank's partial product; the sum over the TP group's + ranks is ``fc``'s output.""" + + def __init__(self, weight: torch.Tensor, start: int, end: int) -> None: + super().__init__() + self.weight = nn.Parameter(weight, requires_grad=False) + self.start = start + self.end = end + + @classmethod + def of(cls, fc_weight: torch.Tensor, tp_size: int, tp_rank: int) -> Optional["K3FcSlice"]: + """Rank ``tp_rank``'s block of the full ``fc_weight`` ``[out_features, in_features]``: a copy of its columns + (on one rank, the weight itself); None where the input columns do not split evenly over ``tp_size`` ranks.""" + columns = fc_columns(fc_weight.shape[1], tp_size, tp_rank) + if columns is None: + return None + start, end = columns + return cls(fc_weight.detach()[:, start:end].contiguous(), start, end) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + """``hidden_states[:, start:end] @ weight.T`` for the full features ``hidden_states`` ``[N, in_features]``: the + GEMM reads the column block in place, through the rows' stride.""" + return torch.nn.functional.linear(hidden_states[:, self.start : self.end], self.weight) + + class K3DSparkDrafter(GQADSparkForCausalLM): """Kimi K3's DSpark drafter: the stock GQA drafter, with a decode step's block forward on the drafter entries. @@ -2327,11 +2366,27 @@ class K3DSparkDrafter(GQADSparkForCausalLM): * gate / up on ``drafter_gate_up`` and the down projection with the SiLU-and-mul on ``drafter_down``, then the module's all-reduce. - The norms and the residual adds are the stock modules'; a projection whose site does not take its rows runs its - module. Every other block runs the stock ``dflash_forward``: a split `DRAFTER_ATTN_SPLITS` does not list, another - attention backend, head layout, RoPE or normalization, a cache the kernel does not read, or a compile key of the - attention that has not run eagerly, under CUDA-graph capture. The worker that calls it (the stock - ``DSparkWorker``), the context projection, the context k / v and the Markov head stay upstream code. + The norms and the residual adds are the stock modules' unless the drafter runs on the TP group's collective state + (below); a projection whose site does not take its rows runs its module. Every other block runs the stock + ``dflash_forward``: a split `DRAFTER_ATTN_SPLITS` does not list, another attention backend, head layout, RoPE or + normalization, a cache the kernel does not read, or a compile key of the attention that has not run eagerly, under + CUDA-graph capture. + + Where the target hands the drafter the TP group's collective state (``use_decode_comm``): + + * the context projection's ``fc`` holds only this rank's block of input columns (`K3FcSlice`), and + ``project_target_hidden`` sums the ranks' partial products with ``hidden_norm`` in one all-reduce; + * a block the entries take runs each residual add and RMSNorm in the all-reduce before it (o_proj's with the + post-attention norm, the MLP's with the next layer's input norm, the last one's with the final norm), the + block's input serving as the first residual, uncopied. Up to 8 rows ``comm/k3_sandwich_plain`` runs the + projection, its all-reduce and the norm in one launch (o_proj, and with the SiLU-and-mul the down projection + after ``drafter_gate_up``); otherwise the projection (its site, or its module without the all-reduce) is + followed by ``comm/mnnvl_fusion_allreduce``. A block whose rows the MNNVL workspace does not hold, or layers + whose norms or projections the fused all-reduces do not reproduce, keep the stock all-reduces and norms. + + Otherwise ``fc`` stays replicated, with the stock context projection, and the norms and residual adds are the + stock modules'. The worker that calls the drafter (the stock ``DSparkWorker``), the context k / v and the Markov + head stay upstream code. """ def __init__( @@ -2345,6 +2400,151 @@ def __init__( # The attention's compile keys that ran eagerly ((more than one request, page stride)); a capture takes # only these. self._k3_attn_ran: set = set() + # The TP group's collective state (decode_comm.py), set by the target where it builds one (use_decode_comm). + self.decode_comm: Optional[_decode_comm.K3DecodeComm] = None + # The zero residual of the split context projection's all-reduce, up to a decode step's rows. + self._k3_zero_rows: Optional[torch.Tensor] = None + # Whether a block's residual adds and RMSNorms run in its all-reduces, and the comm/k3_sandwich_plain forms + # ("o_proj", "down") compiled for the layers' shapes; set with the collective state. + self._k3_norms_fuse = False + self._k3_sandwich_forms: frozenset = frozenset() + + def use_decode_comm(self, comm: _decode_comm.K3DecodeComm) -> None: + """Run the drafter's collectives on ``comm``, the TP group's decode state (``decode_comm.K3DecodeComm``) the + target builds where every attention all-reduce runs over MNNVL. The target calls it once the weights are + loaded, before any CUDA-graph capture. Collective: every rank of the group calls it at the same point. + + * The context projection's ``fc`` keeps only this rank's block of input columns (`K3FcSlice`), where they + split evenly over the group and the drafter's layers have their TP all-reduce: ``project_target_hidden`` + sums the ranks' partial products and applies ``hidden_norm`` in one ``comm/mnnvl_fusion_allreduce`` call. + * Where the layers' norms and projections take it (`_k3_norms_take_comm`), a block's residual adds and + RMSNorms run in its all-reduces (``_k3_block_forward``). Each form of ``comm/k3_sandwich_plain`` whose + shape every layer shares compiles here, with one call on a zero row of a zero weight; a form that does not + compile is not used.""" + self.decode_comm = comm + self._k3_split_fc() + self._k3_norms_fuse = self._k3_norms_take_comm() + forms = set() + if self._k3_norms_fuse: + layers = self.model.layers + for form, swiglu, weights in ( + ("o_proj", False, [layer.self_attn.o_proj.weight for layer in layers]), + ("down", True, [layer.mlp.down_proj.weight for layer in layers]), + ): + if len({tuple(w.shape) for w in weights}) == 1 and comm.compile_plain( + weights[0], swiglu=swiglu + ): + forms.add(form) + self._k3_sandwich_forms = frozenset(forms) + logger.info( + "Kimi K3 DSpark drafter: residual adds and RMSNorms " + + ( + f"in the all-reduces (k3_sandwich_plain: {sorted(forms) or 'none'}, else " + "mnnvl_fusion_allreduce)" + if self._k3_norms_fuse + else "on the stock modules (a layer norm or projection the fused all-reduces do not reproduce)" + ) + ) + + def _k3_norms_take_comm(self) -> bool: + """Whether a block's residual adds and RMSNorms can run in its all-reduces: every layer's input and + post-attention norms and the final norm are plain bf16 RMSNorms of the hidden width (no Gemma offset, no + quantized output), and every output and down projection is row parallel with its all-reduce.""" + hidden = self.config.hidden_size + + def plain(norm: nn.Module) -> bool: + weight = getattr(norm, "weight", None) + return ( + isinstance(weight, torch.Tensor) + and weight.dtype == torch.bfloat16 + and tuple(weight.shape) == (hidden,) + and hasattr(norm, "variance_epsilon") + and not getattr(norm, "use_gemma", True) + and not getattr(norm, "is_nvfp4", True) + and not getattr(norm, "return_hp_output", True) + ) + + def reduces(linear: nn.Module) -> bool: + return ( + getattr(getattr(linear, "tp_mode", None), "name", None) == "ROW" + and getattr(linear, "reduce_output", False) + and getattr(linear, "all_reduce", None) is not None + ) + + return plain(self.model.norm) and all( + plain(layer.input_layernorm) + and plain(layer.post_attention_layernorm) + and reduces(layer.self_attn.o_proj) + and reduces(layer.mlp.down_proj) + for layer in self.model.layers + ) + + def _k3_tp_all_reduce(self) -> Optional[nn.Module]: + """The drafter's own TP all-reduce (its first output projection's row-parallel all-reduce module), or None.""" + o_proj = self.model.layers[0].self_attn.o_proj + if getattr(getattr(o_proj, "tp_mode", None), "name", None) != "ROW": + return None + return getattr(o_proj, "all_reduce", None) + + def _k3_split_fc(self) -> None: + """Replace the replicated ``fc`` with this rank's block of its input columns (`K3FcSlice`) and size the zero + residual of the projection's all-reduce, where the drafter runs on the TP group's collective state, ``fc`` is + a bias-free bf16 projection whose input columns split evenly over the group, and the drafter has its TP + all-reduce (for rows the MNNVL workspace does not hold). Otherwise ``fc`` stays replicated.""" + fc = getattr(self, "fc", None) + if self.decode_comm is None or fc is None or isinstance(fc, K3FcSlice): + return + mapping = self.model_config.mapping + weight = getattr(fc, "weight", None) + sliced = None + if ( + isinstance(weight, torch.Tensor) + and weight.dim() == 2 + and weight.dtype == torch.bfloat16 + and getattr(fc, "bias", None) is None + and self._k3_tp_all_reduce() is not None + ): + sliced = K3FcSlice.of(weight, mapping.tp_size, mapping.tp_rank) + if sliced is None: + logger.info( + "Kimi K3 DSpark drafter: fc stays replicated (it is not a bias-free bf16 projection whose input " + f"columns split evenly over TP{mapping.tp_size}, or the layers have no TP all-reduce)" + ) + return + self.fc = sliced + self._k3_zero_rows = weight.new_zeros( + MAX_REQUESTS * MAX_TOKENS_PER_REQUEST, weight.shape[0] + ) + logger.info( + f"Kimi K3 DSpark drafter: fc split by input feature over TP{mapping.tp_size}: rank {mapping.tp_rank} " + f"holds columns [{sliced.start}, {sliced.end}), hidden_norm in the all-reduce" + ) + + def load_weights(self, weights, weight_mapper=None, **kwargs): + """The stock load; on the TP group's collective state, ``fc`` is split again (`_k3_split_fc`).""" + result = super().load_weights(weights, weight_mapper=weight_mapper, **kwargs) + self._k3_split_fc() + return result + + def project_target_hidden(self, hidden_states: torch.Tensor) -> torch.Tensor: + """``hidden_norm(fc(hidden_states))`` of the captured target features ``[N, in_features]``. + + With ``fc`` split (`K3FcSlice`): this rank's partial product of its columns, then the sum over the TP group + with ``hidden_norm`` applied in one ``comm/mnnvl_fusion_allreduce`` call (a zero residual) up to a decode + step's rows the MNNVL workspace holds; more rows (a prefill) go through the drafter's TP all-reduce, then + ``hidden_norm``. Otherwise the stock projection.""" + fc = self.fc + if not isinstance(fc, K3FcSlice): + return super().project_target_hidden(hidden_states) + partial = fc(hidden_states.to(fc.weight.dtype)) + rows = partial.shape[0] + zeros = self._k3_zero_rows + if rows <= zeros.shape[0] and self.decode_comm.takes_allreduce_norm(rows, partial.shape[1]): + normed, _ = self.decode_comm.allreduce_norm(partial, zeros[:rows], self.hidden_norm) + return normed + if rows == 0: + return self.hidden_norm(partial) + return self.hidden_norm(self._k3_tp_all_reduce()(partial)) def dflash_forward( self, @@ -2356,9 +2556,12 @@ def dflash_forward( ctx_cache_batch_idx: torch.Tensor, ctx_kv_cache: Optional[torch.Tensor] = None, ctx_page_table: Optional[torch.Tensor] = None, + ctx_rows_start: Optional[int] = None, ) -> torch.Tensor: """The block's hidden states ``[B * block, hidden]``: on the drafter entries where they take it (the class - docstring), else the stock forward.""" + docstring), else the stock forward. ``ctx_rows_start``: where ``ctx_cache_batch_idx`` is the contiguous rows + ``[ctx_rows_start, ctx_rows_start + B)`` of ``ctx_page_table``, their start; the entries then read those rows + as a view instead of gathering them.""" keys = self._k3_block_keys(noise_embedding, ctx_kv_cache, ctx_page_table) if keys is None: return super().dflash_forward( @@ -2378,6 +2581,7 @@ def dflash_forward( ctx_cache_batch_idx, ctx_kv_cache, ctx_page_table, + ctx_rows_start, ) if not torch.cuda.is_current_stream_capturing(): self._k3_attn_ran |= keys @@ -2484,21 +2688,34 @@ def _k3_block_forward( ctx_cache_batch_idx: torch.Tensor, ctx_kv_cache: torch.Tensor, ctx_page_table: torch.Tensor, + ctx_rows_start: Optional[int] = None, ) -> torch.Tensor: batch, block = noise_embedding.shape[:2] rows = batch * block ctx_len = num_ctx_per_req[:batch].to(torch.int32) - page_table = ctx_page_table.index_select(0, ctx_cache_batch_idx.to(torch.long)) + if ctx_rows_start is None: + page_table = ctx_page_table.index_select(0, ctx_cache_batch_idx.to(torch.long)) + else: + # The batch's page-table rows are one contiguous run: a view, no gather. + page_table = ctx_page_table[ctx_rows_start : ctx_rows_start + batch] positions = query_positions.reshape(-1).contiguous() hidden = noise_embedding.reshape(rows, -1) + layers = self.model.layers + comm = self._k3_fused_comm(hidden) residual = None - for layer_idx, layer in enumerate(self.model.layers): + if comm is not None: + # The fused all-reduces read the residual and return the updated one: the block's input serves as the + # first residual, uncopied. + residual = hidden + normed = layers[0].input_layernorm(hidden) + for layer_idx, layer in enumerate(layers): attn = layer.self_attn - if residual is None: - residual = hidden.clone() - normed = layer.input_layernorm(hidden) - else: - normed, residual = layer.input_layernorm(hidden, residual) + if comm is None: + if residual is None: + residual = hidden.clone() + normed = layer.input_layernorm(hidden) + else: + normed, residual = layer.input_layernorm(hidden, residual) qkv = self._k3_project("drafter_qkv", normed, attn.qkv_proj) out = torch.empty(rows, attn.q_size, dtype=torch.bfloat16, device=qkv.device) k3_drafter_attn_qknorm( @@ -2515,12 +2732,88 @@ def _k3_block_forward( attn.num_key_value_heads, out, ) + if comm is not None: + # o_proj's all-reduce applies the post-attention norm; the MLP's the next layer's input norm, after + # the last layer the final norm. + next_norm = ( + layers[layer_idx + 1].input_layernorm + if layer_idx + 1 < len(layers) + else self.model.norm + ) + normed, residual = self._k3_project_norm( + comm, out, attn.o_proj, residual, layer.post_attention_layernorm + ) + normed, residual = self._k3_mlp_norm(comm, layer.mlp, normed, residual, next_norm) + continue hidden = self._k3_project("drafter_o", out, attn.o_proj) hidden, residual = layer.post_attention_layernorm(hidden, residual) hidden = self._k3_mlp(layer.mlp, hidden) + if comm is not None: + return normed out, _ = self.model.norm(hidden, residual) return out + def _k3_fused_comm(self, hidden: torch.Tensor) -> Optional[_decode_comm.K3DecodeComm]: + """The TP group's collective state where a block of rows ``hidden`` runs its residual adds and RMSNorms in its + all-reduces: the drafter runs on it (``use_decode_comm``) with layers that take it, and its MNNVL workspace + holds an all-reduce of the block's rows; else None (the stock all-reduces and norms).""" + comm = self.decode_comm + if comm is None or not self._k3_norms_fuse or not comm.takes_allreduce_norm(*hidden.shape): + return None + return comm + + def _k3_project_norm( + self, + comm: _decode_comm.K3DecodeComm, + x: torch.Tensor, + linear: nn.Module, + residual: torch.Tensor, + norm: nn.Module, + ) -> Tuple[torch.Tensor, torch.Tensor]: + """``(norm(updated), updated)``, ``updated = residual + linear(x)`` for the row-parallel output projection, + the residual add and ``norm`` in its all-reduce: ``comm/k3_sandwich_plain`` where it takes the rows (the + projection in the ``drafter_o`` site's arithmetic, in the same launch), else the projection on ``drafter_o`` + where that takes the rows, or the module without its all-reduce, then ``comm/mnnvl_fusion_allreduce``.""" + if "o_proj" in self._k3_sandwich_forms and comm.takes_plain( + x, linear.weight, residual, norm + ): + return comm.sandwich_plain(x, linear.weight, residual, norm) + gemvs = self.decode_gemvs + partial = None if gemvs is None else gemvs.project("drafter_o", x, linear.weight) + if partial is None: + partial = linear(x, all_reduce_params=_decode_comm.skip_all_reduce()) + return comm.allreduce_norm(partial, residual, norm) + + def _k3_mlp_norm( + self, + comm: _decode_comm.K3DecodeComm, + mlp: nn.Module, + x: torch.Tensor, + residual: torch.Tensor, + norm: nn.Module, + ) -> Tuple[torch.Tensor, torch.Tensor]: + """``(norm(updated), updated)``, ``updated = residual + mlp(x)``, the residual add and ``norm`` in the down + projection's all-reduce: gate / up on ``drafter_gate_up``, then ``comm/k3_sandwich_plain``'s SiLU-and-mul + form where it takes the rows (the ``drafter_down`` site's arithmetic), else the SiLU-and-mul and down + projection on ``drafter_down``; where those sites do not take the rows, the module without its all-reduce; + then ``comm/mnnvl_fusion_allreduce``.""" + gemvs = self.decode_gemvs + gate_up = ( + None if gemvs is None else gemvs.project("drafter_gate_up", x, mlp.gate_up_proj.weight) + ) + partial = None + if gate_up is not None: + if "down" in self._k3_sandwich_forms and comm.takes_plain( + gate_up, mlp.down_proj.weight, residual, norm, swiglu=True + ): + return comm.sandwich_plain( + gate_up, mlp.down_proj.weight, residual, norm, swiglu=True + ) + partial = gemvs.project("drafter_down", gate_up, mlp.down_proj.weight) + if partial is None: + partial = mlp(x, final_all_reduce_params=_decode_comm.skip_all_reduce()) + return comm.allreduce_norm(partial, residual, norm) + def _k3_project(self, site: str, x: torch.Tensor, linear: nn.Module) -> torch.Tensor: """``linear(x)``, its GEMM on the ``site`` decode GEMV where that takes the rows, then a row-parallel projection's all-reduce.""" @@ -2790,7 +3083,8 @@ def post_load_weights(self) -> None: attention all-reduce runs over MNNVL, the TP group's collective state (``K3DecodeComm``, collective: every rank builds it here) handed to every layer, with whether its sandwich takes the layer's o_proj, the decode path's one-shot ceiling on every stock MNNVL all-reduce (``use_decode_one_shot``), and the MoE decode path - (``_build_decode_moe``); then the speculative worker's decode kernels (``_gate_spec_worker_kernels``).""" + (``_build_decode_moe``); then the speculative worker's decode kernels (``_gate_spec_worker_kernels``) and the + DSpark drafter's collectives (``_gate_drafter_comm``).""" super().post_load_weights() kda = [layer.linear_attn for layer in self.model.layers if layer.is_kda] mla = [layer.self_attn.mixer for layer in self.model.layers if not layer.is_kda] @@ -2839,6 +3133,7 @@ def post_load_weights(self) -> None: + f"; MoE on k3_moe_front, k3_moe and the row-parallel tail ({moe_layers} layers)" ) self._gate_spec_worker_kernels(comm) + self._gate_drafter_comm(comm) def _build_decode_moe(self, comm: _decode_comm.K3DecodeComm) -> int: """The MoE decode path (``decode_moe.py``) on every MoE layer it takes: the shared state (collective: the @@ -2904,6 +3199,21 @@ def _gate_spec_worker_kernels(self, comm: Optional[_decode_comm.K3DecodeComm]) - ) return worker.k3_decode + def _gate_drafter_comm(self, comm: Optional[_decode_comm.K3DecodeComm]) -> bool: + """Hand the DSpark drafter (`K3DSparkDrafter`) the TP group's collective state ``comm`` where this target built + it (every attention all-reduce over MNNVL; TP16 is a construction assert): the drafter then runs its context + projection split over the group, ``hidden_norm`` in the all-reduce, and its blocks' residual adds and RMSNorms + in their all-reduces (``K3DSparkDrafter.use_decode_comm``, collective: it compiles the drafter's sandwich on + the group's workspace). Without ``comm`` it keeps the stock replicated ``fc``, all-reduces and norms. The gate + requires exactly what this path uses: ``comm`` and nothing else, so not the LM head's ``k3_head_gemv`` + workspace, which only the speculative worker's path (`_gate_spec_worker_kernels`) reads. Returns whether the + drafter took the state (False without such a drafter).""" + drafter = getattr(self, "draft_model", None) + if comm is None or not isinstance(drafter, K3DSparkDrafter): + return False + drafter.use_decode_comm(comm) + return True + def _check_step_contract(self, attn_metadata: AttentionMetadata) -> None: """First-forward checks of the engine surface and the per-engine settings.""" objects = { From 19ef7792a7774f096aba46487c874ff308fff0ca Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 12:00:40 -0700 Subject: [PATCH 143/161] [None][test] Kimi K3 kernel tests: make their pool workers' functions picklable Under pytest, test_k3_latent_reduce, test_k3_sandwich, test_k3_moe_front, test_k3_mnnvl_comm and test_k3_moe_push send their per-rank checks to a pool of MPI workers, which get the module's functions by value (cloudpickle). Every pool entry failed before reaching a GPU: - torch 2.14 keeps torch.ops in sys.modules, and cloudpickle 3.1 adds it to the state of every function whose own code names torch.ops. torch.ops cannot be pickled: "cannot pickle '_Ops' object". The torch.ops.trtllm calls now go through a helper that names torch.ops only in a nested function, the fix that tests/unittest/_torch/multi_gpu/test_allocate_output_buffer_kinds.py describes. - test_k3_moe_push's expert cache was a functools.lru_cache wrapper, which the workers get by reference, from a module they cannot import. It is a plain dict now. The checks are unchanged. Signed-off-by: Vasanth Sabavat --- .../kimi_k3/test_k3_latent_reduce.py | 16 +++++++-- .../kimi_k3/test_k3_mnnvl_comm.py | 22 +++++++++--- .../kimi_k3/test_k3_moe_front.py | 24 +++++++++---- .../kimi_k3/test_k3_moe_push.py | 34 ++++++++++++++----- .../kimi_k3/test_k3_sandwich.py | 30 +++++++++++----- 5 files changed, 96 insertions(+), 30 deletions(-) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_latent_reduce.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_latent_reduce.py index 6a665ae2554c..b56f110a0e50 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_latent_reduce.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_latent_reduce.py @@ -44,6 +44,18 @@ cloudpickle.register_pickle_by_value(sys.modules[__name__]) MPI.pickle.__init__(cloudpickle.dumps, cloudpickle.loads, pickle.HIGHEST_PROTOCOL) + +def _trtllm(): + """``torch.ops.trtllm``, named only in a nested function: the pool gets this module's functions by value, and + cloudpickle cannot pickle a function whose own code names ``torch.ops`` (it adds ``sys.modules["torch.ops"]`` to + the function's state).""" + + def namespace(): + return torch.ops.trtllm + + return namespace() + + WORLD = 4 LATENT = 3584 M_CASES = list(range(1, 9)) @@ -119,7 +131,7 @@ def _push(ctx, x, half): def _reduce(ctx, m, ctas): - out = torch.ops.trtllm.k3_latent_reduce(ctx.ex.uc, ctx.ex.flags, m, ctas) + out = _trtllm().k3_latent_reduce(ctx.ex.uc, ctx.ex.flags, m, ctas) ctx.count += 1 return out @@ -183,7 +195,7 @@ def check_graph(ctx): for i, m in enumerate(ms): # the halves alternate, so the captured pairs follow the eager calls' parity _push(ctx, inputs[i], (base + i) & 1) - outs[i] = torch.ops.trtllm.k3_latent_reduce(ctx.ex.uc, ctx.ex.flags, m, 0) + outs[i] = _trtllm().k3_latent_reduce(ctx.ex.uc, ctx.ex.flags, m, 0) torch.cuda.synchronize() ctx.comm.Barrier() for rep in range(3): diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mnnvl_comm.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mnnvl_comm.py index e4e7c2746ddd..9e7189cf6a63 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mnnvl_comm.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_mnnvl_comm.py @@ -49,6 +49,18 @@ cloudpickle.register_pickle_by_value(sys.modules[__name__]) MPI.pickle.__init__(cloudpickle.dumps, cloudpickle.loads, pickle.HIGHEST_PROTOCOL) + +def _trtllm(): + """``torch.ops.trtllm``, named only in a nested function: the pool gets this module's functions by value, and + cloudpickle cannot pickle a function whose own code names ``torch.ops`` (it adds ``sys.modules["torch.ops"]`` to + the function's state).""" + + def namespace(): + return torch.ops.trtllm + + return namespace() + + WORLD = 4 H, LATENT, EXPERTS = 7168, 3584, 896 EPS, OUT_EPS = 1e-6, 1e-5 @@ -153,12 +165,12 @@ def _unfused(ctx, partial, prefix, block, res_w, rms_w, out_w): if s == 0: return None, (reduced if prefix is None else (prefix.float() + reduced.float()).bfloat16()) if prefix is None: - out = torch.ops.trtllm.attn_res_rmsnorm_fwd(reduced.reshape(m, 1, H), block.reshape(s, m, 1, H), res_w, rms_w, - out_w, EPS, OUT_EPS) # fmt: skip + out = _trtllm().attn_res_rmsnorm_fwd(reduced.reshape(m, 1, H), block.reshape(s, m, 1, H), res_w, rms_w, + out_w, EPS, OUT_EPS) # fmt: skip return out.reshape(m, H), reduced - updated, out = torch.ops.trtllm.attn_res_add_rmsnorm_fwd(prefix.reshape(m, 1, H), reduced.reshape(m, 1, H), - block.reshape(s, m, 1, H), res_w, rms_w, out_w, EPS, - OUT_EPS) # fmt: skip + updated, out = _trtllm().attn_res_add_rmsnorm_fwd(prefix.reshape(m, 1, H), reduced.reshape(m, 1, H), + block.reshape(s, m, 1, H), res_w, rms_w, out_w, EPS, + OUT_EPS) # fmt: skip return out.reshape(m, H), updated.reshape(m, H) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py index f789e9808487..546d301d7c00 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_front.py @@ -69,6 +69,18 @@ cloudpickle.register_pickle_by_value(sys.modules[__name__]) MPI.pickle.__init__(cloudpickle.dumps, cloudpickle.loads, pickle.HIGHEST_PROTOCOL) + +def _trtllm(): + """``torch.ops.trtllm``, named only in a nested function: the pool gets this module's functions by value, and + cloudpickle cannot pickle a function whose own code names ``torch.ops`` (it adds ``sys.modules["torch.ops"]`` to + the function's state).""" + + def namespace(): + return torch.ops.trtllm + + return namespace() + + WORLD = 4 HIDDEN, LATENT, EXPERTS, TOP_K, SV = 7168, 3584, 896, 16, 32 SHARED_INTER = 6144 # two shared experts of 3072 @@ -197,8 +209,8 @@ def _quiet_check(ctx, fn): def _front(ctx, x): - return torch.ops.trtllm.k3_moe_front(x, ctx.front, ctx.bias, RSF, ctx.inter, GATE_CAP, LINEAR_CAP, *ctx.ag, - ctx.world) # fmt: skip + return _trtllm().k3_moe_front(x, ctx.front, ctx.bias, RSF, ctx.inter, GATE_CAP, LINEAR_CAP, *ctx.ag, + ctx.world) # fmt: skip def _reference(ctx, x): @@ -208,8 +220,8 @@ def _reference(ctx, x): parts = [torch.from_numpy(a).cuda() for a in ctx.comm.allgather(head.cpu().numpy())] latent = torch.cat([p[:, : ctx.wl] for p in parts], dim=1).bfloat16().contiguous() logits = torch.cat([p[:, ctx.wl :] for p in parts], dim=1).contiguous() - ids, w, q, s = torch.ops.trtllm.kimi_k3_noaux_tc_mxfp8_quant(logits, ctx.bias, latent, RSF) - shared = torch.ops.trtllm.situ_and_mul(F.linear(x, ctx.gate_up), GATE_CAP, LINEAR_CAP) + ids, w, q, s = _trtllm().kimi_k3_noaux_tc_mxfp8_quant(logits, ctx.bias, latent, RSF) + shared = _trtllm().situ_and_mul(F.linear(x, ctx.gate_up), GATE_CAP, LINEAR_CAP) key = torch.sigmoid(logits) + ctx.bias top = key.sort(dim=1, descending=True).values return ids, w, q, s, shared, top[:, TOP_K - 1] - top[:, TOP_K], latent @@ -294,7 +306,7 @@ def _fused(ctx, x, layer=None, bias=None): (y, shared). A head_flags layer's front publishes the ready words (``ag_ready``) and its k3_moe acquires them.""" layer = layer or ctx.layer head = layer.state.head_flags - ids, w, q, s, shared = torch.ops.trtllm.k3_moe_front( + ids, w, q, s, shared = _trtllm().k3_moe_front( x, ctx.front, ctx.bias if bias is None else bias, RSF, ctx.inter, GATE_CAP, LINEAR_CAP, *ctx.ag, ctx.world, ag_ready=ctx.ws.ready if head else None, ) # fmt: skip @@ -312,7 +324,7 @@ def _runner(ctx, ids, w, q, s): p = ctx.experts alpha = torch.full((E_LOCAL,), GATE_CAP, dtype=torch.float32, device="cuda") beta = torch.full((E_LOCAL,), LINEAR_CAP, dtype=torch.float32, device="cuda") - return torch.ops.trtllm.mxe4m3_mxe2m1_block_scale_moe_runner( + return _trtllm().mxe4m3_mxe2m1_block_scale_moe_runner( None, None, q, s.view(-1), p["w31"], p["w31s"], None, alpha, beta, None, p["w2"], p["w2s"], None, EXPERTS, TOP_K, 1, 1, I_TP, LATENT, I_TP, ctx.offset, E_LOCAL, 1.0, int(RoutingMethodType.DeepSeekV3), int(ActType_TrtllmGen.SiTu), topk_weights=w, topk_ids=ids) # fmt: skip diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py index c0a8176c33e1..60645af9be11 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py @@ -38,7 +38,6 @@ srun -N1 -n4 --mpi=pmix python3 test_k3_moe_push.py [push k3_moe_push] """ -import functools import hashlib import math import os @@ -60,6 +59,18 @@ cloudpickle.register_pickle_by_value(sys.modules[__name__]) MPI.pickle.__init__(cloudpickle.dumps, cloudpickle.loads, pickle.HIGHEST_PROTOCOL) + +def _trtllm(): + """``torch.ops.trtllm``, named only in a nested function: the pool gets this module's functions by value, and + cloudpickle cannot pickle a function whose own code names ``torch.ops`` (it adds ``sys.modules["torch.ops"]`` to + the function's state).""" + + def namespace(): + return torch.ops.trtllm + + return namespace() + + WORLD = 4 H, NUM_EXPERTS, SV = 3584, 896, 32 # TP16: a rank's 192-wide intermediate slice, zero-padded to whole tiles (256) by the loader. @@ -112,11 +123,17 @@ def _rand_mxfp4(rows, k, k_full, gen): return (codes[:, 0::2] | (codes[:, 1::2] << 4)).contiguous(), exps -@functools.lru_cache(maxsize=None) +# _experts' buffers per seed. A dict, not functools.lru_cache: the pool's workers would get an lru_cache wrapper by +# reference, from a module they cannot import. +_EXPERTS = {} + + def _experts(seed: int): """This rank's TP16 experts through W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod's loader: the buffers the engines read, the 192-wide shard generated as rank 0 of tensors that hold exactly it (the loader slices it, then pads it to 256). - """ + Built once per seed.""" + if seed in _EXPERTS: + return _EXPERTS[seed] from tensorrt_llm._torch.moe.fused_moe.quantization import W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod method = W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod() @@ -137,7 +154,8 @@ def _experts(seed: int): method.load_expert_w3_w1_weight_scale_mxfp4(module, gate_s, up_s, w31s[e]) method.load_expert_w2_weight_scale_mxfp4(module, down_s, w2s[e]) torch.cuda.synchronize() - return w31, w31s, w2, w2s + _EXPERTS[seed] = w31, w31s, w2, w2s + return _EXPERTS[seed] def _tokens(m: int, seed: int): @@ -147,7 +165,7 @@ def _tokens(m: int, seed: int): x = torch.randn(m, H, generator=gen, device="cuda").bfloat16() logits = torch.randn(m, NUM_EXPERTS, generator=gen, device="cuda") bias = (torch.randn(NUM_EXPERTS, generator=gen, device="cuda") * 0.05).float() - ids, weights, x_fp8, x_sf = torch.ops.trtllm.k3_route_quant(logits, bias, x, RSF, True) + ids, weights, x_fp8, x_sf = _trtllm().k3_route_quant(logits, bias, x, RSF, True) return x_fp8, x_sf, ids, weights @@ -178,7 +196,7 @@ def __init__(self, ctx, slots: int, count: int): ctx.comm.Barrier() def reduce(self, m: int) -> torch.Tensor: - out = torch.ops.trtllm.k3_latent_reduce(self.uc, self.flags, m, 0) + out = _trtllm().k3_latent_reduce(self.uc, self.flags, m, 0) self.count = (self.count + 1 + 2**31) % 2**32 - 2**31 # int32 two's complement return out @@ -327,11 +345,11 @@ def _k3_moe_routed(producer, inputs, front): argument order. k3_moe_front is collective over the run's ranks (the head all-gather on ``front.head``).""" if producer == "k3_route_quant": x, logits, bias = inputs - ids, weights, x_fp8, x_sf = torch.ops.trtllm.k3_route_quant(logits, bias, x, RSF, True) + ids, weights, x_fp8, x_sf = _trtllm().k3_route_quant(logits, bias, x, RSF, True) return x_fp8, x_sf, ids, weights, None x, bias = inputs head = front.head - ids, weights, x_fp8, x_sf, shared = torch.ops.trtllm.k3_moe_front( + ids, weights, x_fp8, x_sf, shared = _trtllm().k3_moe_front( x, front.weight, bias, RSF, front.inter, GATE_CAP, LINEAR_CAP, head.uc, head.mc, head.flags, head.rank, head.world_size, ) # fmt: skip diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py index 50da0aeee64d..2794a65d3b14 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_sandwich.py @@ -52,6 +52,18 @@ cloudpickle.register_pickle_by_value(sys.modules[__name__]) MPI.pickle.__init__(cloudpickle.dumps, cloudpickle.loads, pickle.HIGHEST_PROTOCOL) + +def _trtllm(): + """``torch.ops.trtllm``, named only in a nested function: the pool gets this module's functions by value, and + cloudpickle cannot pickle a function whose own code names ``torch.ops`` (it adds ``sys.modules["torch.ops"]`` to + the function's state).""" + + def namespace(): + return torch.ops.trtllm + + return namespace() + + WORLD = 4 H, K_O, LATENT, WIDTH, PAD, ACT = 7168, 768, 3584, 224, 256, 384 PLAIN_K, DOWN_K = 384, 896 @@ -126,20 +138,20 @@ def _all_ranks(ctx, good) -> bool: def _oproj(ctx, core, w, prefix, block, res_w, rms_w, out_w): ws = ctx.ws - return torch.ops.trtllm.k3_sandwich_oproj(core, w, prefix, block, res_w, rms_w, out_w, EPS, EPS, ws.uc, ws.mc, - ws.flags, ws.rank) # fmt: skip + return _trtllm().k3_sandwich_oproj(core, w, prefix, block, res_w, rms_w, out_w, EPS, EPS, ws.uc, ws.mc, + ws.flags, ws.rank) # fmt: skip def _tail(ctx, latent, act, w, lo, prefix, block, res_w, rms_w, out_w, **extra): ws = ctx.ws - return torch.ops.trtllm.k3_sandwich_tail(latent, act, w, lo, LAT_EPS, prefix, block, res_w, rms_w, out_w, EPS, - EPS, ws.uc, ws.mc, ws.flags, ws.rank, **extra) # fmt: skip + return _trtllm().k3_sandwich_tail(latent, act, w, lo, LAT_EPS, prefix, block, res_w, rms_w, out_w, EPS, + EPS, ws.uc, ws.mc, ws.flags, ws.rank, **extra) # fmt: skip def _plain(ctx, x, w, residual, norm_w, swiglu=False): ws = ctx.ws - return torch.ops.trtllm.k3_sandwich_plain(x, w, residual, norm_w, EPS, ws.uc, ws.mc, ws.flags, ws.rank, - swiglu=swiglu) # fmt: skip + return _trtllm().k3_sandwich_plain(x, w, residual, norm_w, EPS, ws.uc, ws.mc, ws.flags, ws.rank, + swiglu=swiglu) # fmt: skip def _attn_res_ar(ctx, partial, prefix, block, res_w, rms_w, out_w): @@ -263,9 +275,9 @@ def check_oproj(ctx): def _tap_mixture(updated, block, res_w, rms_w): """The unfused path's pre-norm attention-residual mixture: trtllm::attn_res_fwd on the updated row and the bank.""" m, s = updated.shape[0], block.shape[0] - out, _, _, _ = torch.ops.trtllm.attn_res_fwd(updated.reshape(m, 1, H).contiguous(), - block.reshape(s, m, 1, H).contiguous(), res_w.reshape(-1).contiguous(), - rms_w.contiguous(), EPS) # fmt: skip + out, _, _, _ = _trtllm().attn_res_fwd(updated.reshape(m, 1, H).contiguous(), + block.reshape(s, m, 1, H).contiguous(), res_w.reshape(-1).contiguous(), + rms_w.contiguous(), EPS) # fmt: skip return out.reshape(m, H) From 428f6f2a77083870bf70a4ff577736455397f1bb Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 11:57:17 -0700 Subject: [PATCH 144/161] [None][perf] modeling_v2 Kimi K3 tp16_moetp4ep4: the DSpark tap from the sandwich tail At a DSpark-tapped layer whose pre-attention step is comm/k3_sandwich_tail (the previous MoE layer deferred its tail), the kernel now stores the tap straight into the layer's capture slot: the pre-norm attention-residual mixture, or updated with the prefix-only aux stream. The layer takes the slot from the speculative metadata's capture_view and skips the split tap, which was a separate attn_res kernel plus the capture copy at each tapped layer of a DSpark decode step. Where the metadata has no capture_view, the layer keeps the split tap. K3DecodeComm.sandwich_tail forwards tap and tap_updated to the catalog entry, which certifies the tapped mixture within 2e-2 of an fp32 reference and every other output bit for bit with and without the tap. The tap is not bit-equal to the split path's attn_res kernel (each rounds an fp32 mixture once, in its own order), so the drafter's input at those layers can move by bf16 rounding. _kimi_k3_decode_comm_op_matrix.py gains check_moe_tail_tap: a tapped consumer of a deferred tail with and without capture_view, the two taps within TOL and every other output bit for bit the same. Signed-off-by: Vasanth Sabavat --- .../decode_comm.py | 7 +- .../modeling.py | 19 ++++- .../comm/_kimi_k3_decode_comm_op_matrix.py | 73 ++++++++++++++++++- 3 files changed, 91 insertions(+), 8 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py index dbe64ba7a513..2884f85e7f8a 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_comm.py @@ -276,11 +276,14 @@ def sandwich_tail( res_norm: nn.Module, out_norm: nn.Module, updated_out: Optional[torch.Tensor] = None, + tap: Optional[torch.Tensor] = None, + tap_updated: bool = False, ) -> Tuple[torch.Tensor, torch.Tensor]: """``(normed, updated)`` as ``allreduce_attn_res`` of a MoE layer's row-parallel tail ``pending`` (``[RMSNorm(latent)[:, lo:lo + 224] | act] @ weight.T``), in one ``comm/k3_sandwich_tail`` call. ``updated_out``: a bf16 ``[T, H]`` tensor the call stores ``updated`` into (the consumer's snapshot bank row), - returned as ``updated``.""" + returned as ``updated``. ``tap``: a bf16 ``[T, H]`` view (a speculative capture slot) the call also stores + the pre-norm attention-residual mixture into, or ``updated`` with ``tap_updated``.""" return k3_sandwich_tail( pending.latent, pending.act, @@ -291,6 +294,8 @@ def sandwich_tail( block_residual, *_res_args(res_proj, res_norm, out_norm), self.sandwich, + tap=tap, + tap_updated=tap_updated, updated_out=updated_out, ) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index 5fc908147bee..8ddb40e548b9 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -1538,7 +1538,10 @@ def forward( ``hidden_states`` is the prefix sum without it, and this layer's pre-attention step reduces and adds it: a ``PendingTail`` in one ``K3DecodeComm.sandwich_tail`` call, a wide decode step's tensor by its - all-reduce and the fused add + attn_res + RMSNorm. + all-reduce and the fused add + attn_res + RMSNorm. On a tapped layer + the ``sandwich_tail`` call also stores the tap into the capture slot + where the speculative metadata exposes it (``capture_view``); other + tapped layers take the tap after the step. ``defer_moe_tail``: return ``(prefix_sum, num_snapshots, partial)`` instead, ``partial`` this layer's MoE output unreduced, for the next @@ -1557,6 +1560,12 @@ def forward( snapshot_row = None if tail is not None and self.layer_idx % self.attn_res_block_size == 0: snapshot_row = block_residual[num_snapshots] + # A tapped layer whose pre-attention step is the sandwich tail has the kernel store the tap straight into the + # layer's capture slot, where the speculative metadata exposes that slot as a view (``capture_view``). + tap_view = None + if tail is not None and capture is not None: + view_of = getattr(capture[0], "capture_view", None) + tap_view = view_of(capture[1], prefix_sum.shape[0]) if view_of is not None else None if prenormed: assert num_snapshots == 0 and self.layer_idx % self.attn_res_block_size == 0 @@ -1570,6 +1579,8 @@ def forward( self.self_attention_res_norm, self.input_layernorm, updated_out=snapshot_row, + tap=tap_view, + tap_updated=not _AUX_ATTN_RES_STREAM_ENABLED, ) elif pending_moe_partial is not None: prefix_sum, hidden_states = _apply_attn_res_add_and_rmsnorm( @@ -1610,9 +1621,9 @@ def forward( else: hidden_states = self.input_layernorm(hidden_states) - if capture is not None and pending_moe_partial is not None: - # The tapped layer handed its MoE output on: the step above reduced it into prefix_sum. Tap that value's - # pre-norm attn_res mixture, what the split path captures. + if capture is not None and pending_moe_partial is not None and tap_view is None: + # The tapped layer handed its MoE output on and no kernel wrote the tap: the step above reduced the output + # into prefix_sum. Tap that value's pre-norm attn_res mixture, what the split path captures. tapped = ( _apply_attn_res( prefix_sum, diff --git a/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py index 786445cadfdc..c869abe7e518 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py @@ -37,6 +37,10 @@ comm/k3_sandwich_tail (layer 24's kernel stores the prefix sum into the bank row it takes), against the same layers with the tail reduced in torch and added before the built-in pre-attention step: the consumer's attention input, its outputs and the bank within TOL; the deferred chain captured and replayed equals eager bit for bit. + * MoE tail tap: layer 23 or 24 as a tapped layer (its speculative metadata a stand-in) consuming layer 22's tail. + With the metadata's ``capture_view`` the sandwich tail stores the tap into the capture slot and the layer captures + nothing more; without it the layer captures the split path's tap (``_apply_attn_res``). The two taps within TOL, + every other output of the chain bit for bit the same. * Every output bitwise equal across the ranks. Every rank draws the replicated tensors (prefix sum, snapshot bank, norms) from one seed and its own o_proj, core and @@ -473,9 +477,12 @@ def _pending(seed, tokens): return T._decode_comm.PendingTail(latent, act, weight.contiguous(), R.rank * LAT_SLICE, 1e-5) -def _tail_chain(producer, consumer, x, bank, step, pending, defer, consumer_core, producer_core): +def _tail_chain( + producer, consumer, x, bank, step, pending, defer, consumer_core, producer_core, capture=None +): """The producer layer (its MoE handing ``pending`` on with ``defer``, else adding its reduced tail), then the - consumer layer. Returns the consumer's attention input, its returned prefix sum, its MoE input and the bank.""" + consumer layer (with ``capture``, a tapped layer). Returns the consumer's attention input, its returned prefix + sum, its MoE input and the bank.""" _attention(producer).core = producer_core _attention(consumer).core = consumer_core producer.block_sparse_moe.pending = pending @@ -491,7 +498,13 @@ def _tail_chain(producer, consumer, x, bank, step, pending, defer, consumer_core else: (prefix, snapshots), partial = out, None prefix_c, snapshots_c = consumer( - prefix, bank, snapshots, SimpleNamespace(), step=step, pending_moe_partial=partial + prefix, + bank, + snapshots, + SimpleNamespace(), + capture=capture, + step=step, + pending_moe_partial=partial, ) return ( _attention(consumer).last_input, @@ -531,6 +544,59 @@ def check_moe_tail_deferral(): assert R.same_on_ranks(*got), (consumer_idx, tokens, "ranks differ") +class _CaptureMetadata: + """Stand-in for the speculative metadata a tapped layer captures into: ``maybe_capture_hidden_states`` records + each value it is handed, and with ``views`` the metadata also exposes ``capture_view``, slot 2 of a NaN-filled + ``[T, 5 x H]`` capture buffer, as DSpark's does.""" + + SLOT = 2 + + def __init__(self, tokens, views): + self.buf = torch.full((tokens, 5 * H), float("nan"), dtype=torch.bfloat16, device="cuda") + self.captured = [] + if views: + self.capture_view = self.view + + def view(self, layer_id, num_tokens): + return self.buf[:num_tokens, self.SLOT * H : (self.SLOT + 1) * H] + + def maybe_capture_hidden_states(self, layer_id, hidden_states, residual=None): + self.captured.append(hidden_states.clone()) + + +def check_moe_tail_tap(): + seed = 3500 + for consumer_idx in (23, 24): + producer, consumer = LAYERS_BUILT[22], LAYERS_BUILT[consumer_idx] + for tokens in (1, 3, 8): + seed += 1 + step = T.DecodeStep(tokens, tokens, 1) + case = Case(seed, 22, tokens) + gr = _gen(seed * 13 + 5 + R.rank) + core_c = ls.exact_bf16(gr, (tokens, K_IN), -4, 5, 1 / 8) + pending = _pending(seed, tokens) + fused, split = _CaptureMetadata(tokens, True), _CaptureMetadata(tokens, False) + got = _tail_chain(producer, consumer, case.x.clone(), case.bank.clone(), step, pending, True, core_c, + case.core, capture=(fused, consumer_idx - 1)) # fmt: skip + got = [t.clone() for t in got] + want = _tail_chain(producer, consumer, case.x.clone(), case.bank.clone(), step, pending, True, core_c, + case.core, capture=(split, consumer_idx - 1)) # fmt: skip + torch.cuda.synchronize() + assert not fused.captured and len(split.captured) == 1, (consumer_idx, tokens) + tap, split_tap = fused.view(consumer_idx - 1, tokens), split.captured[0] + err = _err(tap, split_tap) + differ = int((tap.view(torch.int16) != split_tap.view(torch.int16)).sum()) + if R.rank == 0: + print( + f"[rank 0] tail tap {22}->{consumer_idx} T={tokens}: kernel vs split {err:.2e}, " + f"{differ} of {tap.numel()} elements differ", + flush=True, + ) + assert torch.isfinite(tap.float()).all() and err < TOL, (consumer_idx, tokens, err) + assert all(torch.equal(a, b) for a, b in zip(got, want)), (consumer_idx, tokens) + assert R.same_on_ranks(tap, *got), (consumer_idx, tokens, "ranks differ") + + def check_moe_tail_graph(): producer, consumer = LAYERS_BUILT[22], LAYERS_BUILT[24] tokens = 8 @@ -715,6 +781,7 @@ def check_decode_moe_from_model_parameters(): check_sandwich_vs_mnnvl_bitwise, check_graph_capture_and_replay, check_moe_tail_deferral, + check_moe_tail_tap, check_moe_tail_graph, ] From fe9eaf52030689127d272bcbf9b46069c5837835 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 11:50:35 -0700 Subject: [PATCH 145/161] [None][chore] modeling_v2 Kimi K3 tp16_moetp16ep1: carry the sandwich-tail tap into the copy The tp16_moetp4ep4 target has a DSpark-tapped layer's sandwich tail store the tap straight into its capture slot (capture_view), and K3DecodeComm.sandwich_tail forwards tap and tap_updated. The copy here takes both: decode_comm.py byte for byte, and the change to modeling.py outside route B's blocks. This target decodes without speculation, so no layer has a capture here and the tap is never requested: sandwich_tail runs as before. Signed-off-by: Vasanth Sabavat --- .../decode_comm.py | 7 ++++++- .../modeling.py | 19 +++++++++++++++---- 2 files changed, 21 insertions(+), 5 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_comm.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_comm.py index dbe64ba7a513..2884f85e7f8a 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_comm.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_comm.py @@ -276,11 +276,14 @@ def sandwich_tail( res_norm: nn.Module, out_norm: nn.Module, updated_out: Optional[torch.Tensor] = None, + tap: Optional[torch.Tensor] = None, + tap_updated: bool = False, ) -> Tuple[torch.Tensor, torch.Tensor]: """``(normed, updated)`` as ``allreduce_attn_res`` of a MoE layer's row-parallel tail ``pending`` (``[RMSNorm(latent)[:, lo:lo + 224] | act] @ weight.T``), in one ``comm/k3_sandwich_tail`` call. ``updated_out``: a bf16 ``[T, H]`` tensor the call stores ``updated`` into (the consumer's snapshot bank row), - returned as ``updated``.""" + returned as ``updated``. ``tap``: a bf16 ``[T, H]`` view (a speculative capture slot) the call also stores + the pre-norm attention-residual mixture into, or ``updated`` with ``tap_updated``.""" return k3_sandwich_tail( pending.latent, pending.act, @@ -291,6 +294,8 @@ def sandwich_tail( block_residual, *_res_args(res_proj, res_norm, out_norm), self.sandwich, + tap=tap, + tap_updated=tap_updated, updated_out=updated_out, ) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py index 7dd3ff5b634a..918e9fc2c7ee 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py @@ -1538,7 +1538,10 @@ def forward( ``hidden_states`` is the prefix sum without it, and this layer's pre-attention step reduces and adds it: a ``PendingTail`` in one ``K3DecodeComm.sandwich_tail`` call, a wide decode step's tensor by its - all-reduce and the fused add + attn_res + RMSNorm. + all-reduce and the fused add + attn_res + RMSNorm. On a tapped layer + the ``sandwich_tail`` call also stores the tap into the capture slot + where the speculative metadata exposes it (``capture_view``); other + tapped layers take the tap after the step. ``defer_moe_tail``: return ``(prefix_sum, num_snapshots, partial)`` instead, ``partial`` this layer's MoE output unreduced, for the next @@ -1557,6 +1560,12 @@ def forward( snapshot_row = None if tail is not None and self.layer_idx % self.attn_res_block_size == 0: snapshot_row = block_residual[num_snapshots] + # A tapped layer whose pre-attention step is the sandwich tail has the kernel store the tap straight into the + # layer's capture slot, where the speculative metadata exposes that slot as a view (``capture_view``). + tap_view = None + if tail is not None and capture is not None: + view_of = getattr(capture[0], "capture_view", None) + tap_view = view_of(capture[1], prefix_sum.shape[0]) if view_of is not None else None if prenormed: assert num_snapshots == 0 and self.layer_idx % self.attn_res_block_size == 0 @@ -1570,6 +1579,8 @@ def forward( self.self_attention_res_norm, self.input_layernorm, updated_out=snapshot_row, + tap=tap_view, + tap_updated=not _AUX_ATTN_RES_STREAM_ENABLED, ) elif pending_moe_partial is not None: prefix_sum, hidden_states = _apply_attn_res_add_and_rmsnorm( @@ -1610,9 +1621,9 @@ def forward( else: hidden_states = self.input_layernorm(hidden_states) - if capture is not None and pending_moe_partial is not None: - # The tapped layer handed its MoE output on: the step above reduced it into prefix_sum. Tap that value's - # pre-norm attn_res mixture, what the split path captures. + if capture is not None and pending_moe_partial is not None and tap_view is None: + # The tapped layer handed its MoE output on and no kernel wrote the tap: the step above reduced the output + # into prefix_sum. Tap that value's pre-norm attn_res mixture, what the split path captures. tapped = ( _apply_attn_res( prefix_sum, From 4bf35122304d0d45453d1d44302d839ddcd1c1e8 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 10:55:34 -0700 Subject: [PATCH 146/161] [None][perf] modeling_v2 Kimi K3 target: the decode MoE's latent all-reduce as the routed experts' push form On a pure decode step of at most 8 tokens, captured into a CUDA graph, whose attention layers all run the decode kernels, every MoE layer now pushes its routed partial into the TP group's latent exchange (moe/k3_moe's push form) and comm/k3_latent_reduce sums the partials, in place of k3_moe and the routed experts' MNNVL one-shot all-reduce. The reduce sums in the one-shot's order, so the latent and the step's outputs keep their bits. The exchange is built in post_load_weights beside the head workspace, and the warm-up compiles the push build and the reduce. Every other step (eager, with a context request, on the built-in KDA verify, wide) keeps the one-shot. The decode-comm matrix checks the push against the one-shot bit for bit at 4 ranks with random MXFP4 experts: every M 1..8, a 3-layer sequence with wide and one-shot steps and a late rank, a captured step replayed, and the exchange's ring reused under 16 replays at every M. The decode-step test covers which steps push. Signed-off-by: Vasanth Sabavat --- .../decode_moe.py | 79 ++++++-- .../modeling.py | 65 +++++- .../comm/_kimi_k3_decode_comm_op_matrix.py | 189 ++++++++++++++++++ .../test_modeling_v2_kimi_k3_decode_step.py | 45 +++++ 4 files changed, 358 insertions(+), 20 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_moe.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_moe.py index 951c928e5e0e..a2fa6889283b 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_moe.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_moe.py @@ -8,7 +8,11 @@ GEMV, the slices' all-gather over the TP group's `K3MoeHeadWorkspace`, the top-16 routing, the MXFP8 latent, and the shared experts' gate_up + SiTU; * `moe/k3_moe`: this rank's routed partial over its experts (on a `K3MoeState`); -* the latent all-reduce: the routed experts' all-reduce (one-shot, as `decode_comm.use_decode_one_shot` sets); +* the latent all-reduce. On a pushing step (`DecodeStep.latent_push`: a pure decode step captured into a CUDA graph + whose attention layers all run the decode kernels) the routed experts run as `moe/k3_moe`'s push form, which stores + this rank's partial into every rank's `K3LatentExchange`, and `comm/k3_latent_reduce` sums the partials in the MNNVL + one-shot's order, so with the same bits, while the experts' grid completes. On every other step it is the routed + experts' all-reduce (one-shot, as `decode_comm.use_decode_one_shot` sets); * the tail. The latent norm's weight is folded into the latent up projection at load, so `[RMSNorm(latent) slice | shared activation] @ [latent up columns | shared down]` is this rank's share of the MoE output. Where the next layer's pre-attention step (or the final norm) reduces it, the layer hands it on unreduced @@ -24,10 +28,16 @@ The GEMVs run on the decode GEMV sites of `decode_gemv.py` where they take the call, else on the stock GEMM ops. -`K3DecodeMoe` holds what every MoE layer shares: the head workspace (collective over the TP group), the two -`k3_moe` builds' scratch, and the TP group's MNNVL workspace (`decode_comm.K3DecodeComm`'s). `K3DecodeMoeLayer` -holds one layer's decode weights and its `k3_moe` counters. The target builds both in `post_load_weights`, before -any CUDA-graph capture, and runs every kernel once there so none compiles under a capture. +`K3DecodeMoe` holds what every MoE layer shares: the head workspace and the latent exchange (both collective over +the TP group), the two `k3_moe` builds' scratch, and the TP group's MNNVL workspace (`decode_comm.K3DecodeComm`'s). +`K3DecodeMoeLayer` holds one layer's decode weights and its `k3_moe` counters. The target builds both in +`post_load_weights`, before any CUDA-graph capture, and runs every kernel once there so none compiles under a capture. + +Every MoE layer's push and reduce go to the one exchange, in the stream's order: each push is followed by exactly one +reduce of the same token count before the next push, on every rank in the same order (the exchange's call-order +invariant, `comm/k3_latent_reduce`). Every kernel between a reduce and the next push waits for its predecessor (or +launches without programmatic dependent launch), which the decode kernels a pushing step runs do; pushes run only in +CUDA-graph replays, so no other step's kernels sit between two pushes. """ from __future__ import annotations @@ -38,6 +48,10 @@ import torch from torch import nn +from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.k3_latent_reduce import ( + K3LatentExchange, + k3_latent_reduce, +) from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.mnnvl_allgather_split import ( mnnvl_allgather_split, ) @@ -51,6 +65,7 @@ K3MoeWideState, is_supported, k3_moe, + k3_moe_push, ) from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.k3_moe_front import ( front_weight, @@ -75,25 +90,38 @@ class K3DecodeMoe: """What every MoE layer's decode path shares on one device: the TP group's ``K3MoeHeadWorkspace`` (the front's all-gather), the ``k3_moe`` builds for up to 8 and up to 64 tokens (their scratch; the layers run one at a time on - one stream), and the TP group's ``MnnvlWorkspace`` for a wide step's head all-gather. Built by `create`.""" + one stream), the TP group's ``MnnvlWorkspace`` for a wide step's head all-gather, and the TP group's + ``K3LatentExchange`` for a pushing step's latent all-reduce (None: every step keeps the routed experts' + all-reduce). Built by `create`.""" head: K3MoeHeadWorkspace small: K3MoeState wide: K3MoeWideState mnnvl: MnnvlWorkspace + exchange: Optional[K3LatentExchange] = None @classmethod def create( - cls, mapping, device, i_tp: int, num_local: int, mnnvl: MnnvlWorkspace + cls, mapping, device, i_tp: int, num_local: int, mnnvl: MnnvlWorkspace, push: bool = True ) -> "K3DecodeMoe": """The state for ``mapping``'s TP group on ``device``, for experts of ``i_tp`` intermediate columns per rank, - ``num_local`` of them on this rank. Collective (the head workspace): every rank of the group calls it at the - same point, eagerly, before any CUDA-graph capture.""" + ``num_local`` of them on this rank. Collective (the head workspace and the latent exchange): every rank of the + group calls it at the same point, eagerly, before any CUDA-graph capture. ``push``: build the latent exchange + (where it takes the TP size: 4, 8 or 16 ranks); without it every step keeps the routed experts' all-reduce.""" + head = K3MoeHeadWorkspace.create(mapping) + exchange = None + if push: + try: + exchange = K3LatentExchange.create(mapping) + # Raised before any collective step, on every rank of the group alike: a TP size the exchange does not take. + except ValueError: + exchange = None return cls( - K3MoeHeadWorkspace.create(mapping), + head, K3MoeState(device, i_tp, num_local), K3MoeWideState(device, i_tp, num_local), mnnvl, + exchange, ) @@ -237,12 +265,18 @@ def create( def warm_up(self, moe: nn.Module) -> None: """One call of every kernel of the decode path on zero inputs (M = 1), so none compiles under a capture: the - front (collective: every rank makes the same call), both ``k3_moe`` builds and ``k3_route_quant``.""" + front (collective: every rank makes the same call), both ``k3_moe`` builds, ``k3_route_quant`` and, with the + latent exchange, the push build and the reduce (one push + reduce pair, collective). Once per model: every + layer's calls compile the same builds.""" device = self.front_weight.device x = torch.zeros(1, moe.hidden_size, dtype=torch.bfloat16, device=device) ids, weights, x_fp8, x_sf, _ = self._front(moe, x) offset = moe.routed_experts.backend.slot_start k3_moe(x_fp8, x_sf, ids, weights, offset, self.small) + exchange = self.state.exchange + if exchange is not None: + k3_moe_push(x_fp8, x_sf, ids, weights, offset, self.small, exchange) + k3_latent_reduce(1, exchange) logits = torch.zeros(1, moe.num_experts, dtype=torch.float32, device=device) latent = torch.zeros(1, moe.moe_hidden_size, dtype=torch.bfloat16, device=device) ids, weights, x_fp8, x_sf = k3_route_quant( @@ -271,18 +305,29 @@ def takes(self, hidden_states: torch.Tensor, step, partial_tail: bool) -> bool: return 0 < rows <= MAX_TOKENS or (partial_tail and step.wide and rows <= WIDE_MAX_TOKENS) def forward( - self, moe: nn.Module, hidden_states: torch.Tensor, gemvs, partial_tail: bool + self, + moe: nn.Module, + hidden_states: torch.Tensor, + gemvs, + partial_tail: bool, + push: bool = False, ) -> Union[torch.Tensor, PendingTail]: """The MoE output of ``hidden_states`` (``takes`` holds): with ``partial_tail``, this rank's unreduced share (a ``PendingTail`` at most `MAX_TOKENS` tokens); else the reduced output. ``gemvs``: the decode GEMVs' state, - or None.""" + or None. ``push`` (the step's ``DecodeStep.latent_push``): at most `MAX_TOKENS` tokens, the latent all-reduce + is the routed experts' push form plus ``comm/k3_latent_reduce`` on the state's exchange (the same bits as the + routed experts' all-reduce); a state without the exchange ignores it.""" if hidden_states.shape[0] > MAX_TOKENS: return self._wide(moe, hidden_states, gemvs) ids, weights, x_fp8, x_sf, shared_act = self._front(moe, hidden_states) - routed = k3_moe( - x_fp8, x_sf, ids, weights, moe.routed_experts.backend.slot_start, self.small - ) - latent = moe.routed_experts.all_reduce(routed) + offset = moe.routed_experts.backend.slot_start + exchange = self.state.exchange + if push and exchange is not None: + k3_moe_push(x_fp8, x_sf, ids, weights, offset, self.small, exchange) + latent = k3_latent_reduce(hidden_states.shape[0], exchange) + else: + routed = k3_moe(x_fp8, x_sf, ids, weights, offset, self.small) + latent = moe.routed_experts.all_reduce(routed) if partial_tail: return PendingTail( latent.contiguous(), diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index 8ddb40e548b9..40710ba4666b 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -53,7 +53,8 @@ A MoE layer on a step of at most 8 tokens runs `decode_moe.py`: `moe/k3_moe_front` (this rank's head slice, its all-gather, the routing, the MXFP8 latent and the shared experts' gate_up + SiTU in one kernel), `moe/k3_moe`, the -latent all-reduce, then the row-parallel tail, which the next layer's pre-attention step (the final norm's, after the +latent all-reduce (on a pushing step, `DecodeStep.latent_push`, `moe/k3_moe`'s push form plus `comm/k3_latent_reduce`), +then the row-parallel tail, which the next layer's pre-attention step (the final norm's, after the last layer) runs with its all-reduce and residual update as one `comm/k3_sandwich_tail` kernel. A wide decode step's MoE keeps the sharded head and the row-parallel tail, on `moe/k3_route_quant`, `moe/k3_moe` and M-general ops. @@ -83,7 +84,7 @@ import copy import math import os -from dataclasses import dataclass +from dataclasses import dataclass, replace from typing import TYPE_CHECKING, Any, Callable, Dict, Literal, NamedTuple, Optional, Tuple, Union import torch @@ -189,6 +190,7 @@ "k3_sandwich_tail", "k3_moe_front", "k3_moe", + "k3_latent_reduce", "k3_route_quant", "mnnvl_allgather_split", # The DSpark drafter's decode path (K3DSparkDrafter): its GEMV sites, its block attention, and its all-reduces @@ -1199,7 +1201,9 @@ def forward( most 8 tokens, a tensor on a wide decode step.""" decode = self.decode_moe if decode is not None and decode.takes(hidden_states, step, partial_tail): - return decode.forward(self, hidden_states, self.decode_gemvs, partial_tail) + return decode.forward( + self, hidden_states, self.decode_gemvs, partial_tail, push=step.latent_push + ) if partial_tail: raise RuntimeError("the row-parallel MoE tail needs the decode path to take the step") identity = hidden_states @@ -1823,6 +1827,23 @@ def kda_token_states(self) -> bool: and all(layer.linear_attn.takes_k3_kernels for layer in self.layers if layer.is_kda) ) + def _latent_push(self, attn_metadata: AttentionMetadata, step: DecodeStep) -> bool: + """Whether the MoE layers push on ``step`` (``latent_push``): read only while a CUDA graph is being captured, + from what the graph's key fixes (the step's shape, the attention metadata the decode kernels read) and from + load-time state, so every rank and every replay of the graph decides alike.""" + if not torch.cuda.is_current_stream_capturing(): + return False + kda = [layer.linear_attn for layer in self.layers if layer.is_kda] + mla = [layer.self_attn for layer in self.layers if not layer.is_kda] + return latent_push( + step, + capturing=True, + breakable=is_in_breakable_cuda_graph(), + kda_token_states=self.kda_token_states, + kda_decode_kernels=all(m.takes_k3_kernels and m.k3_buffers is not None for m in kda), + mla_decode_branch=all(m.will_run_decode_branch(attn_metadata, step) for m in mla[:1]), + ) + def _defer_moe_tail( self, i: int, @@ -1860,6 +1881,8 @@ def forward( num_tokens = (input_ids if inputs_embeds is None else inputs_embeds).shape[0] step = decode_step(attn_metadata, num_tokens) + if step is not None and self._latent_push(attn_metadata, step): + step = replace(step, latent_push=True) # A decode step embeds and norms for layer 0 in one launch, the embedding written as layer 0's first snapshot. prenormed = None if ( @@ -2987,6 +3010,9 @@ class DecodeStep: num_tokens: int num_requests: Optional[int] = None tokens_per_request: Optional[int] = None + #: Whether the MoE layers' latent all-reduce is the routed experts' push form plus ``comm/k3_latent_reduce`` + #: (``decode_moe.py``): set by ``KimiLinearModel.forward`` from ``latent_push``. + latent_push: bool = False @property def small(self) -> bool: @@ -3005,6 +3031,30 @@ def wide(self) -> bool: return self.decode and not self.small +def latent_push( + step: Optional[DecodeStep], + *, + capturing: bool, + breakable: bool, + kda_token_states: bool, + kda_decode_kernels: bool, + mla_decode_branch: bool, +) -> bool: + """Whether a step's MoE layers push their routed partials (``DecodeStep.latent_push``): a pure decode step of at + most 8 tokens whose attention layers all run the decode kernels, captured into a CUDA graph and not inside a + breakable one. The KDA layers run them on one token per request with every KDA layer taking the K3 kernels + (``kda_decode_kernels``), and on every verify width with the per-token states (``kda_token_states``); the MLA + layers where their decode branch takes the step (``mla_decode_branch``, the same for every MLA layer of a step). + Every other step keeps the routed experts' all-reduce: the exchange's call-order invariant needs every kernel + between a reduce and the next push to wait for its predecessor, which only those kernels were checked for, and a + graph launch orders every pushing replay behind whatever ran before it.""" + if step is None or not (step.decode and step.small) or not capturing or breakable: + return False + if not (kda_token_states or (step.tokens_per_request == 1 and kda_decode_kernels)): + return False + return mla_decode_branch + + def decode_step(attn_metadata: AttentionMetadata, num_tokens: int) -> Optional[DecodeStep]: """The step's shape if any Kimi K3 decode kernel takes it, else None (the generic path runs). @@ -3263,6 +3313,15 @@ def _build_decode_moe(self, comm: _decode_comm.K3DecodeComm) -> int: moe.decode_gemvs = self.model.decode_gemvs first = takes[0].decode_moe first.warm_up(takes[0]) + logger.info( + "Kimi K3 MoE decode path: the latent all-reduce at <= 8 tokens on " + + ( + "the routed experts' push form and k3_latent_reduce on pushing steps (a pure decode step captured " + "into a CUDA graph on the decode kernels), else the routed experts' all-reduce" + if state.exchange is not None + else f"the routed experts' all-reduce (no latent exchange at TP {mapping.tp_size})" + ) + ) comm.compile_tail(takes[0].moe_hidden_size, first.shared_cols, first.tail_weight) return len(takes) diff --git a/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py index c869abe7e518..56c97eb41bfe 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py @@ -41,12 +41,22 @@ With the metadata's ``capture_view`` the sandwich tail stores the tap into the capture slot and the layer captures nothing more; without it the layer captures the split path's tap (``_apply_attn_res``). The two taps within TOL, every other output of the chain bit for bit the same. + * decode MoE push: the MoE decode path with random MXFP4 routed experts (route A's layout: 224 local experts of + intermediate 768, every rank its own EP shard) on the TP group's latent exchange, its latent all-reduce as the + routed experts' push form plus comm/k3_latent_reduce (``push=True``) against the routed experts' all-reduce + (``push=False``), bit for bit: one call at every M 1..8, both tails; 3 layers x 12 steps with M dipping and growing + back, wide steps and one step without the push in between, a random rank late on every pushing call; a 2-layer step + captured in a CUDA graph and replayed with rewritten inputs; and, at every M 1..8, a 3-layer step captured and + replayed 16 times (48 push + reduce calls, the exchange's two halves rotating at every layer position). After all + of it the exchange is empty and its call count has advanced by one per pushing call. * Every output bitwise equal across the ranks. Every rank draws the replicated tensors (prefix sum, snapshot bank, norms) from one seed and its own o_proj, core and tail from a rank seed. """ +import math +import random import sys from pathlib import Path from types import SimpleNamespace @@ -774,6 +784,184 @@ def check_decode_moe_from_model_parameters(): assert R.same_on_ranks(y), (rows, "ranks differ") +PUSH_LAYERS = 3 +PUSH_STEPS = (8, 3, 1, 8, 16, 5, 2, 8, 1, 16, 7, 8) # tokens per step; 16: a wide step +ONE_SHOT_STEP = 5 # the sequence's step that keeps the routed experts' all-reduce +RING_REPLAYS = 16 +EXCHANGE_EMPTY = -(2**31) # an exchange word no push has written + + +def _rand_mxfp4(rows, k, gen): + """Random checkpoint-format MXFP4 (packed [rows, k / 2], one E8M0 per 32 k), as test_modeling_v2_k3_moe.py draws + it.""" + codes = torch.randint(0, 16, (rows, k), dtype=torch.uint8, device="cuda", generator=gen) + base = 127 + round(0.5 * math.log2(0.01057 / k)) + exps = torch.randint( + base, base + 6, (rows, k // 32), dtype=torch.uint8, device="cuda", generator=gen + ) + return (codes[:, 0::2] | (codes[:, 1::2] << 4)).contiguous(), exps + + +def _load_experts(backend, seed): + """Random MXFP4 routed experts through the TRTLLM-Gen W4A8_MXFP4_MXFP8 loader into ``backend``'s buffers, at + route A's layout: intermediate 768 per rank (moe_tp 4), 224 local experts.""" + from tensorrt_llm._torch.moe.fused_moe.quantization import W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod + + method = W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod() + i_full = I_TP * 4 + module = SimpleNamespace( + tp_size=4, + tp_rank=0, + scaling_vector_size=32, + intermediate_size=i_full, + intermediate_size_per_partition=I_TP, + hidden_size=LATENT, + ) + gen = _gen(seed) + with torch.no_grad(): + for e in range(E_LOCAL): + w1, w1s = _rand_mxfp4(i_full, LATENT, gen) # gate + w3, w3s = _rand_mxfp4(i_full, LATENT, gen) # up + w2, w2s = _rand_mxfp4(LATENT, i_full, gen) # down + method.load_expert_w3_w1_weight(module, w1, w3, backend.w3_w1_weight[e]) + method.load_expert_w2_weight(module, w2, backend.w2_weight[e]) + method.load_expert_w3_w1_weight_scale_mxfp4( + module, w1s, w3s, backend.w3_w1_weight_scale[e] + ) + method.load_expert_w2_weight_scale_mxfp4(module, w2s, backend.w2_weight_scale[e]) + torch.cuda.synchronize() + + +def _outs(out): + """A decode MoE call's tensors: a PendingTail's latent and shared activation, else the output.""" + return (out.latent, out.act) if isinstance(out, T._decode_comm.PendingTail) else (out,) + + +def _same_bits(got, want) -> bool: + return all( + g.shape == w.shape + and torch.equal(g.contiguous().view(torch.int16), w.contiguous().view(torch.int16)) + for g, w in zip(got, want) + ) + + +def check_decode_moe_push(): + """The decode MoE's latent all-reduce as the routed experts' push form plus comm/k3_latent_reduce, against the + routed experts' all-reduce, bit for bit (see the module docstring).""" + if R.world not in (4, 8, 16): + if R.rank == 0: + print(f"[rank 0] decode MoE push skipped at world {R.world}", flush=True) + return + dm = T._decode_moe + moe = _decode_moe_layer() + _load_experts(moe.routed_experts.backend, 6000 + R.rank) + device = torch.device("cuda", torch.cuda.current_device()) + state = dm.K3DecodeMoe.create(R.mapping, device, I_TP, E_LOCAL, COMM.mnnvl) + ex = state.exchange + assert ex is not None, "no latent exchange at a TP size it takes" + dm.fold_latent_norm(moe) + layers = [dm.K3DecodeMoeLayer.create(moe, state, R.rank, R.world) for _ in range(PUSH_LAYERS)] + layers[0].warm_up(moe) + R.barrier() + count0 = int(ex.flags[0].item()) + pushes = 0 + + def call(layer, x, push, tail): + nonlocal pushes + out = _outs(layer.forward(moe, x, None, partial_tail=tail, push=push)) + pushes += int(push and x.shape[0] <= dm.MAX_TOKENS) + return out + + # One call at every M, both tails. + for m in range(1, dm.MAX_TOKENS + 1): + x = _normal(_gen(6100 + m), (m, H), 0.5) + for tail in (False, True): + want = call(layers[0], x, False, tail) + got = call(layers[0], x, True, tail) + torch.cuda.synchronize() + assert _same_bits(got, want), (m, tail, "push != one-shot") + assert got[0].float().abs().sum().item() > 0, (m, tail, "a zero output proves nothing") + # The latent and the replicated output are the same on every rank; the shared activation is the rank's. + assert R.same_on_ranks(got[0]), (m, tail, "ranks differ") + + # A sequence of steps over every layer, both paths; a random rank late on every pushing call. + late = random.Random(11) + runs = {} + for push in (False, True): + outs = [] + for i, m in enumerate(PUSH_STEPS): + x = _normal(_gen(6200 + i), (m, H), 0.5) + for j, layer in enumerate(layers): + pushing = push and i != ONE_SHOT_STEP + R.late(late.randrange(R.world) if pushing else None) + outs.append(call(layer, x, pushing, tail=j % 2 == 0)) + torch.cuda.synchronize() + runs[push] = outs + assert all(_same_bits(g, w) for g, w in zip(runs[True], runs[False])), ( + "sequence: push != one-shot" + ) + + # A 2-layer step at 8 tokens in a CUDA graph, replayed with rewritten inputs, against eager calls without the push. + x_static = _normal(_gen(6300), (dm.MAX_TOKENS, H), 0.5) + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + captured = [ + _outs(layer.forward(moe, x_static, None, partial_tail=j == 0, push=True)) + for j, layer in enumerate(layers[:2]) + ] + for rep in range(3): + x_static.copy_(_normal(_gen(6400 + rep), (dm.MAX_TOKENS, H), 0.5)) + graph.replay() + pushes += 2 + torch.cuda.synchronize() + got = [tuple(t.clone() for t in outs) for outs in captured] + want = [call(layer, x_static, False, j == 0) for j, layer in enumerate(layers[:2])] + torch.cuda.synchronize() + assert all(_same_bits(g, w) for g, w in zip(got, want)), ( + rep, + "graph replay != eager one-shot", + ) + del graph, captured + + # Ring reuse: at every M, a 3-layer step captured once and replayed RING_REPLAYS times. + for m in range(1, dm.MAX_TOKENS + 1): + x_static = _normal(_gen(6500 + m), (m, H), 0.5) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + captured = [ + _outs(layer.forward(moe, x_static, None, partial_tail=j == 1, push=True)) + for j, layer in enumerate(layers) + ] + for rep in range(RING_REPLAYS): + x_static.copy_(_normal(_gen(6600 + 64 * m + rep), (m, H), 0.5)) + R.late(late.randrange(R.world)) + graph.replay() + pushes += PUSH_LAYERS + torch.cuda.synchronize() + got = [tuple(t.clone() for t in outs) for outs in captured] + want = [call(layer, x_static, False, j == 1) for j, layer in enumerate(layers)] + torch.cuda.synchronize() + assert all(_same_bits(g, w) for g, w in zip(got, want)), ( + m, + rep, + "ring replay != eager one-shot", + ) + del graph, captured + + # The exchange is empty, and its count advanced by one per pushing call since the warm-up. + R.barrier() + assert bool((ex.uc == EXCHANGE_EMPTY).all().item()), "exchange words left after the last reduce" + assert int(ex.flags[2].item()) == 0, "arrival word not cleared" + assert int(ex.flags[0].item()) == count0 + pushes, (int(ex.flags[0].item()), count0, pushes) + if R.rank == 0: + print( + f"[rank 0] decode MoE push: {pushes} push + reduce calls, all equal to the one-shot", + flush=True, + ) + + CHECKS = [ check_takes_oproj, check_decode_one_shot, @@ -783,6 +971,7 @@ def check_decode_moe_from_model_parameters(): check_moe_tail_deferral, check_moe_tail_tap, check_moe_tail_graph, + check_decode_moe_push, ] LAYERS_BUILT = None diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_step.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_step.py index 4f221e60a1dd..4c845ab1a39f 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_step.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_decode_step.py @@ -12,6 +12,7 @@ DecodeStep, _attn_res_max_tokens, decode_step, + latent_push, ) pytestmark = pytest.mark.cpu_only @@ -85,3 +86,47 @@ def test_attn_res_epilogue_ceiling(step, ceiling): """The fused attn_res kernels take a classified step up to one token tile and a wide decode step up to 32 tokens; any other step keeps the generic path's ceiling.""" assert _attn_res_max_tokens(step) == ceiling + + +@pytest.mark.parametrize( + "step,given,pushes", + [ + (DecodeStep(1, 1, 1), {}, True), # one token per request on the KDA decode kernels + (DecodeStep(8, 8, 1), {}, True), + ( + DecodeStep(8, 1, 8), + {"kda_token_states": True}, + True, + ), # a DSpark verify with per-token states + (DecodeStep(4, 2, 2), {"kda_token_states": True}, True), + (DecodeStep(8, 1, 8), {}, False), # the built-in KDA verify + (DecodeStep(5), {}, False), # a context request, or rows that are not the step's tokens + ( + DecodeStep(16, 2, 8), + {"kda_token_states": True}, + False, + ), # wide: the routed experts' all-reduce + (DecodeStep(1, 1, 1), {"capturing": False}, False), # an eager step + (DecodeStep(1, 1, 1), {"breakable": True}, False), + ( + DecodeStep(1, 1, 1), + {"kda_decode_kernels": False}, + False, + ), # a KDA layer on the built-in kernels + (DecodeStep(1, 1, 1), {"mla_decode_branch": False}, False), # MLA on the built-in path + (None, {}, False), + ], +) +def test_latent_push(step, given, pushes): + """Which steps the MoE layers push on (``latent_push``): a pure decode step of at most 8 tokens, captured into a + CUDA graph, whose attention layers all run the decode kernels; every other step keeps the routed experts' + all-reduce.""" + flags = dict( + capturing=True, + breakable=False, + kda_token_states=False, + kda_decode_kernels=True, + mla_decode_branch=True, + ) + flags.update(given) + assert latent_push(step, **flags) == pushes From 41b22452df46340fcd930feaaea7e44f39931676 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 12:00:58 -0700 Subject: [PATCH 147/161] [None][test] Kimi K3 k3_moe push test: route A's expert layout check_k3_moe_push also runs on route A's experts (tp16_moetp4ep4): 224 local experts with a 768-wide intermediate slice, at offset (rank % 4) x 224. On one tray the four ranks then hold four different expert sets. Route A's decode MoE pushes with this build (i_tp 768, 224 local experts), which no test ran before. The TP16 rows are unchanged. Signed-off-by: Vasanth Sabavat --- .../kimi_k3/test_k3_moe_push.py | 81 ++++++++++++------- 1 file changed, 52 insertions(+), 29 deletions(-) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py index 60645af9be11..ebe42e9d8f83 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py @@ -23,6 +23,9 @@ trtllm::k3_moe_front: K3MoeLayer.push into the run's exchange, the moe/k3_moe entry's k3_moe_push into the 16-slot one. The front's head and shared experts are sharded over the run's ranks; its shared activation must be the same bits in every call of a set. +k3_moe_push_route_a + the same on route A's experts (tp16_moetp4ep4): a 768-wide intermediate slice of 224 local experts at + offset (rank % 4) x 224, so the ranks of one tray hold 4 different expert sets. Per op, token count and routing set, against the plain call (Layer.__call__, K3MoeLayer.__call__ after the same producer), whose partial must be nonzero: @@ -38,6 +41,7 @@ srun -N1 -n4 --mpi=pmix python3 test_k3_moe_push.py [push k3_moe_push] """ +import functools import hashlib import math import os @@ -75,6 +79,8 @@ def namespace(): H, NUM_EXPERTS, SV = 3584, 896, 32 # TP16: a rank's 192-wide intermediate slice, zero-padded to whole tiles (256) by the loader. I_TP, I_PAD, MOE_TP = 192, 256, 16 +# Route A (tp16_moetp4ep4): a rank's 768-wide slice (whole tiles) of 224 local experts. +I_TP_A, E_LOCAL_A, MOE_TP_A = 768, 224, 4 HIDDEN, SHARED_INTER = 7168, 6144 # the front's input width; two shared experts of 3072 GATE_CAP, LINEAR_CAP = 4.0, 25.0 # the SiTU caps RSF = 2.827 @@ -123,39 +129,44 @@ def _rand_mxfp4(rows, k, k_full, gen): return (codes[:, 0::2] | (codes[:, 1::2] << 4)).contiguous(), exps -# _experts' buffers per seed. A dict, not functools.lru_cache: the pool's workers would get an lru_cache wrapper by -# reference, from a module they cannot import. +# _experts' buffers per seed and layout. A dict, not functools.lru_cache: the pool's workers would get an lru_cache +# wrapper by reference, from a module they cannot import. _EXPERTS = {} -def _experts(seed: int): - """This rank's TP16 experts through W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod's loader: the buffers the engines read, - the 192-wide shard generated as rank 0 of tensors that hold exactly it (the loader slices it, then pads it to 256). - Built once per seed.""" - if seed in _EXPERTS: - return _EXPERTS[seed] +def _experts(seed: int, route_a: bool = False): + """This rank's experts through W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod's loader: the buffers the engines read, the + rank's intermediate slice generated alone. TP16: all 896 experts, the 192-wide slice as rank 0 of tensors that + hold exactly it (the loader slices it, then pads it to 256); route A: 224 experts, the 768-wide slice whole tiles, + which the loader takes as one shard. Built once per seed and layout.""" + if (seed, route_a) in _EXPERTS: + return _EXPERTS[seed, route_a] from tensorrt_llm._torch.moe.fused_moe.quantization import W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod + if route_a: + i_tp, i_pad, shards, moe_tp, experts = I_TP_A, I_TP_A, 1, MOE_TP_A, E_LOCAL_A + else: + i_tp, i_pad, shards, moe_tp, experts = I_TP, I_PAD, MOE_TP, MOE_TP, NUM_EXPERTS method = W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod() - module = SimpleNamespace(tp_size=MOE_TP, tp_rank=0, scaling_vector_size=SV, intermediate_size=I_TP * MOE_TP, - intermediate_size_per_partition=I_TP, hidden_size=H) # fmt: skip + module = SimpleNamespace(tp_size=shards, tp_rank=0, scaling_vector_size=SV, intermediate_size=i_tp * shards, + intermediate_size_per_partition=i_tp, hidden_size=H) # fmt: skip kw = dict(dtype=torch.uint8, device="cuda") - w31 = torch.zeros(NUM_EXPERTS, 2 * I_PAD, H // 2, **kw) - w31s = torch.zeros(NUM_EXPERTS, 2 * I_PAD, H // SV, **kw) - w2 = torch.zeros(NUM_EXPERTS, H, I_PAD // 2, **kw) - w2s = torch.zeros(NUM_EXPERTS, H, I_PAD // SV, **kw) + w31 = torch.zeros(experts, 2 * i_pad, H // 2, **kw) + w31s = torch.zeros(experts, 2 * i_pad, H // SV, **kw) + w2 = torch.zeros(experts, H, i_pad // 2, **kw) + w2s = torch.zeros(experts, H, i_pad // SV, **kw) gen = torch.Generator(device="cuda").manual_seed(seed) - for e in range(NUM_EXPERTS): - gate, gate_s = _rand_mxfp4(I_TP, H, H, gen) - up, up_s = _rand_mxfp4(I_TP, H, H, gen) - down, down_s = _rand_mxfp4(H, I_TP, I_TP * MOE_TP, gen) + for e in range(experts): + gate, gate_s = _rand_mxfp4(i_tp, H, H, gen) + up, up_s = _rand_mxfp4(i_tp, H, H, gen) + down, down_s = _rand_mxfp4(H, i_tp, i_tp * moe_tp, gen) method.load_expert_w3_w1_weight(module, gate, up, w31[e]) method.load_expert_w2_weight(module, down, w2[e]) method.load_expert_w3_w1_weight_scale_mxfp4(module, gate_s, up_s, w31s[e]) method.load_expert_w2_weight_scale_mxfp4(module, down_s, w2s[e]) torch.cuda.synchronize() - _EXPERTS[seed] = w31, w31s, w2, w2s - return _EXPERTS[seed] + _EXPERTS[seed, route_a] = w31, w31s, w2, w2s + return _EXPERTS[seed, route_a] def _tokens(m: int, seed: int): @@ -321,11 +332,13 @@ def _front(ctx): head=K3MoeHeadWorkspace.create(ctx.mapping, fabric_handle=fabric)) # fmt: skip -def _k3_moe_layer(weights): +def _k3_moe_layer(weights, route_a: bool = False): """This rank's experts as a K3MoeLayer of the 8-token build; a call with an exchange takes its push build.""" from tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe.op import K3MoeState device = torch.device("cuda", torch.cuda.current_device()) + if route_a: + return K3MoeState(device, I_TP_A, E_LOCAL_A).layer(*weights) return K3MoeState(device, I_PAD, NUM_EXPERTS).layer(*weights) @@ -356,12 +369,16 @@ def _k3_moe_routed(producer, inputs, front): return x_fp8, x_sf, ids, weights, shared -def check_k3_moe_push(ctx): +def check_k3_moe_push(ctx, route_a: bool = False): from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.k3_moe import k3_moe_push results = [] front = _front(ctx) - layer = _k3_moe_layer(_experts(20260928 + ctx.rank)) + seed = 20260928 + ctx.rank + layer = _k3_moe_layer(_experts(seed, route_a), route_a) + # Route A's experts of this rank: global ids [offset, offset + 224), the rank's EP index being rank % 4. + offset = ctx.rank * E_LOCAL_A % NUM_EXPERTS if route_a else 0 + layout = ", route A" if route_a else "" copies16 = 16 // ctx.world for start in COUNT_STARTS: ex, ex16 = _Exchange(ctx, ctx.world, start), _Exchange(ctx, 16, start) @@ -370,12 +387,12 @@ def check_k3_moe_push(ctx): for si in range(K3_MOE_SETS): inputs = _k3_moe_inputs(producer, m, 2000 + 10 * si + m) *routed, shared = _k3_moe_routed(producer, inputs, front) - y = layer(*routed, 0) + y = layer(*routed, offset) ref = _allreduce(ctx, y) got, pushed_shared = [], [] for _ in range(3): *routed, pushed = _k3_moe_routed(producer, inputs, front) - layer.push(*routed, 0, ex, ctx.rank) + layer.push(*routed, offset, ex, ctx.rank) pushed_shared.append(pushed) got.append(ex.reduce(m)) state = ex.state_ok(ctx) @@ -383,13 +400,14 @@ def check_k3_moe_push(ctx): # TP16's receive side: this rank's partial in slots 4 r .. 4 r + 3, one push each. for c in range(copies16): *routed, pushed = _k3_moe_routed(producer, inputs, front) - k3_moe_push(*routed, 0, layer, ex16, ctx.rank * copies16 + c) + k3_moe_push(*routed, offset, layer, ex16, ctx.rank * copies16 + c) pushed_shared.append(pushed) got16 = ex16.reduce(m) state16 = ex16.state_ok(ctx) shared_eq = shared is None or all(_same(s, shared) for s in pushed_shared) - results.append(_row(ctx, f"K3MoeLayer.push after {producer}", f"set{si}_count{start}", m, y, ref, - got, state, rows, copies16, got16, state16, shared_eq=shared_eq)) # fmt: skip + results.append(_row(ctx, f"K3MoeLayer.push after {producer}{layout}", f"set{si}_count{start}", m, + y, ref, got, state, rows, copies16, got16, state16, + shared_eq=shared_eq)) # fmt: skip return results @@ -475,7 +493,12 @@ def push_reduce(li, tokens): return results -CHECKS = {"push": check_push, "k3_moe_push": check_k3_moe_push, "sequences": check_sequences} +CHECKS = { + "push": check_push, + "k3_moe_push": check_k3_moe_push, + "k3_moe_push_route_a": functools.partial(check_k3_moe_push, route_a=True), + "sequences": check_sequences, +} def _run_checks(names): From aed1c94be29789b49f76a0bf1e6e8140ab7c490e Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 12:22:48 -0700 Subject: [PATCH 148/161] [None][chore] modeling_v2 Kimi K3 tp16_moetp16ep1: carry the latent push into the copy The tp16_moetp4ep4 target runs a pushing step's latent all-reduce as the routed experts' push form plus comm/k3_latent_reduce. In decode_moe.py: the decode MoE path's latent exchange and the layer's push argument. In modeling.py: DecodeStep.latent_push, latent_push, KimiLinearModel._latent_push and its use in the forward, the k3_latent_reduce requirement and _build_decode_moe's log line. The copy here takes them: decode_moe.py byte for byte, and the changes to modeling.py outside route B's blocks. The module docstring keeps this target's text. _latent_push reads the text model's kda_token_states, which this target dropped since it decodes without speculation. Without that property, every CUDA-graph capture of a classified step would raise AttributeError here. So the route B block that drops tp16_moetp4ep4's property now holds a kda_token_states that answers False, the engine's getattr default. test_modeling_v2_kimi_k3_construction.py checks it and the decision it feeds. This target builds no MoE decode path yet, so nothing pushes here. Signed-off-by: Vasanth Sabavat --- .../decode_moe.py | 79 +++++++++++++++---- .../modeling.py | 68 +++++++++++++++- .../test_modeling_v2_kimi_k3_construction.py | 21 +++++ 3 files changed, 149 insertions(+), 19 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py index 951c928e5e0e..a2fa6889283b 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py @@ -8,7 +8,11 @@ GEMV, the slices' all-gather over the TP group's `K3MoeHeadWorkspace`, the top-16 routing, the MXFP8 latent, and the shared experts' gate_up + SiTU; * `moe/k3_moe`: this rank's routed partial over its experts (on a `K3MoeState`); -* the latent all-reduce: the routed experts' all-reduce (one-shot, as `decode_comm.use_decode_one_shot` sets); +* the latent all-reduce. On a pushing step (`DecodeStep.latent_push`: a pure decode step captured into a CUDA graph + whose attention layers all run the decode kernels) the routed experts run as `moe/k3_moe`'s push form, which stores + this rank's partial into every rank's `K3LatentExchange`, and `comm/k3_latent_reduce` sums the partials in the MNNVL + one-shot's order, so with the same bits, while the experts' grid completes. On every other step it is the routed + experts' all-reduce (one-shot, as `decode_comm.use_decode_one_shot` sets); * the tail. The latent norm's weight is folded into the latent up projection at load, so `[RMSNorm(latent) slice | shared activation] @ [latent up columns | shared down]` is this rank's share of the MoE output. Where the next layer's pre-attention step (or the final norm) reduces it, the layer hands it on unreduced @@ -24,10 +28,16 @@ The GEMVs run on the decode GEMV sites of `decode_gemv.py` where they take the call, else on the stock GEMM ops. -`K3DecodeMoe` holds what every MoE layer shares: the head workspace (collective over the TP group), the two -`k3_moe` builds' scratch, and the TP group's MNNVL workspace (`decode_comm.K3DecodeComm`'s). `K3DecodeMoeLayer` -holds one layer's decode weights and its `k3_moe` counters. The target builds both in `post_load_weights`, before -any CUDA-graph capture, and runs every kernel once there so none compiles under a capture. +`K3DecodeMoe` holds what every MoE layer shares: the head workspace and the latent exchange (both collective over +the TP group), the two `k3_moe` builds' scratch, and the TP group's MNNVL workspace (`decode_comm.K3DecodeComm`'s). +`K3DecodeMoeLayer` holds one layer's decode weights and its `k3_moe` counters. The target builds both in +`post_load_weights`, before any CUDA-graph capture, and runs every kernel once there so none compiles under a capture. + +Every MoE layer's push and reduce go to the one exchange, in the stream's order: each push is followed by exactly one +reduce of the same token count before the next push, on every rank in the same order (the exchange's call-order +invariant, `comm/k3_latent_reduce`). Every kernel between a reduce and the next push waits for its predecessor (or +launches without programmatic dependent launch), which the decode kernels a pushing step runs do; pushes run only in +CUDA-graph replays, so no other step's kernels sit between two pushes. """ from __future__ import annotations @@ -38,6 +48,10 @@ import torch from torch import nn +from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.k3_latent_reduce import ( + K3LatentExchange, + k3_latent_reduce, +) from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.mnnvl_allgather_split import ( mnnvl_allgather_split, ) @@ -51,6 +65,7 @@ K3MoeWideState, is_supported, k3_moe, + k3_moe_push, ) from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.k3_moe_front import ( front_weight, @@ -75,25 +90,38 @@ class K3DecodeMoe: """What every MoE layer's decode path shares on one device: the TP group's ``K3MoeHeadWorkspace`` (the front's all-gather), the ``k3_moe`` builds for up to 8 and up to 64 tokens (their scratch; the layers run one at a time on - one stream), and the TP group's ``MnnvlWorkspace`` for a wide step's head all-gather. Built by `create`.""" + one stream), the TP group's ``MnnvlWorkspace`` for a wide step's head all-gather, and the TP group's + ``K3LatentExchange`` for a pushing step's latent all-reduce (None: every step keeps the routed experts' + all-reduce). Built by `create`.""" head: K3MoeHeadWorkspace small: K3MoeState wide: K3MoeWideState mnnvl: MnnvlWorkspace + exchange: Optional[K3LatentExchange] = None @classmethod def create( - cls, mapping, device, i_tp: int, num_local: int, mnnvl: MnnvlWorkspace + cls, mapping, device, i_tp: int, num_local: int, mnnvl: MnnvlWorkspace, push: bool = True ) -> "K3DecodeMoe": """The state for ``mapping``'s TP group on ``device``, for experts of ``i_tp`` intermediate columns per rank, - ``num_local`` of them on this rank. Collective (the head workspace): every rank of the group calls it at the - same point, eagerly, before any CUDA-graph capture.""" + ``num_local`` of them on this rank. Collective (the head workspace and the latent exchange): every rank of the + group calls it at the same point, eagerly, before any CUDA-graph capture. ``push``: build the latent exchange + (where it takes the TP size: 4, 8 or 16 ranks); without it every step keeps the routed experts' all-reduce.""" + head = K3MoeHeadWorkspace.create(mapping) + exchange = None + if push: + try: + exchange = K3LatentExchange.create(mapping) + # Raised before any collective step, on every rank of the group alike: a TP size the exchange does not take. + except ValueError: + exchange = None return cls( - K3MoeHeadWorkspace.create(mapping), + head, K3MoeState(device, i_tp, num_local), K3MoeWideState(device, i_tp, num_local), mnnvl, + exchange, ) @@ -237,12 +265,18 @@ def create( def warm_up(self, moe: nn.Module) -> None: """One call of every kernel of the decode path on zero inputs (M = 1), so none compiles under a capture: the - front (collective: every rank makes the same call), both ``k3_moe`` builds and ``k3_route_quant``.""" + front (collective: every rank makes the same call), both ``k3_moe`` builds, ``k3_route_quant`` and, with the + latent exchange, the push build and the reduce (one push + reduce pair, collective). Once per model: every + layer's calls compile the same builds.""" device = self.front_weight.device x = torch.zeros(1, moe.hidden_size, dtype=torch.bfloat16, device=device) ids, weights, x_fp8, x_sf, _ = self._front(moe, x) offset = moe.routed_experts.backend.slot_start k3_moe(x_fp8, x_sf, ids, weights, offset, self.small) + exchange = self.state.exchange + if exchange is not None: + k3_moe_push(x_fp8, x_sf, ids, weights, offset, self.small, exchange) + k3_latent_reduce(1, exchange) logits = torch.zeros(1, moe.num_experts, dtype=torch.float32, device=device) latent = torch.zeros(1, moe.moe_hidden_size, dtype=torch.bfloat16, device=device) ids, weights, x_fp8, x_sf = k3_route_quant( @@ -271,18 +305,29 @@ def takes(self, hidden_states: torch.Tensor, step, partial_tail: bool) -> bool: return 0 < rows <= MAX_TOKENS or (partial_tail and step.wide and rows <= WIDE_MAX_TOKENS) def forward( - self, moe: nn.Module, hidden_states: torch.Tensor, gemvs, partial_tail: bool + self, + moe: nn.Module, + hidden_states: torch.Tensor, + gemvs, + partial_tail: bool, + push: bool = False, ) -> Union[torch.Tensor, PendingTail]: """The MoE output of ``hidden_states`` (``takes`` holds): with ``partial_tail``, this rank's unreduced share (a ``PendingTail`` at most `MAX_TOKENS` tokens); else the reduced output. ``gemvs``: the decode GEMVs' state, - or None.""" + or None. ``push`` (the step's ``DecodeStep.latent_push``): at most `MAX_TOKENS` tokens, the latent all-reduce + is the routed experts' push form plus ``comm/k3_latent_reduce`` on the state's exchange (the same bits as the + routed experts' all-reduce); a state without the exchange ignores it.""" if hidden_states.shape[0] > MAX_TOKENS: return self._wide(moe, hidden_states, gemvs) ids, weights, x_fp8, x_sf, shared_act = self._front(moe, hidden_states) - routed = k3_moe( - x_fp8, x_sf, ids, weights, moe.routed_experts.backend.slot_start, self.small - ) - latent = moe.routed_experts.all_reduce(routed) + offset = moe.routed_experts.backend.slot_start + exchange = self.state.exchange + if push and exchange is not None: + k3_moe_push(x_fp8, x_sf, ids, weights, offset, self.small, exchange) + latent = k3_latent_reduce(hidden_states.shape[0], exchange) + else: + routed = k3_moe(x_fp8, x_sf, ids, weights, offset, self.small) + latent = moe.routed_experts.all_reduce(routed) if partial_tail: return PendingTail( latent.contiguous(), diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py index 918e9fc2c7ee..1ff1ce88697d 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py @@ -82,7 +82,7 @@ import copy import math import os -from dataclasses import dataclass +from dataclasses import dataclass, replace from typing import TYPE_CHECKING, Any, Callable, Dict, Literal, NamedTuple, Optional, Tuple, Union import torch @@ -186,6 +186,7 @@ "k3_sandwich_tail", "k3_moe_front", "k3_moe", + "k3_latent_reduce", "k3_route_quant", "mnnvl_allgather_split", # The DSpark drafter's decode path (K3DSparkDrafter): its GEMV sites, its block attention, and its all-reduces @@ -1199,7 +1200,9 @@ def forward( most 8 tokens, a tensor on a wide decode step.""" decode = self.decode_moe if decode is not None and decode.takes(hidden_states, step, partial_tail): - return decode.forward(self, hidden_states, self.decode_gemvs, partial_tail) + return decode.forward( + self, hidden_states, self.decode_gemvs, partial_tail, push=step.latent_push + ) if partial_tail: raise RuntimeError("the row-parallel MoE tail needs the decode path to take the step") identity = hidden_states @@ -1810,8 +1813,31 @@ def __init__(self, model_config: ModelConfig): ) # >>> route B: no per-token KDA verify states (kda_token_states): no speculative decoding + @property + def kda_token_states(self) -> bool: + """False: this target decodes without speculation, so the hybrid cache manager keeps no KDA state after + each verify token, and a captured step's latent-push decision (``_latent_push``) reads that.""" + return False + # <<< route B + def _latent_push(self, attn_metadata: AttentionMetadata, step: DecodeStep) -> bool: + """Whether the MoE layers push on ``step`` (``latent_push``): read only while a CUDA graph is being captured, + from what the graph's key fixes (the step's shape, the attention metadata the decode kernels read) and from + load-time state, so every rank and every replay of the graph decides alike.""" + if not torch.cuda.is_current_stream_capturing(): + return False + kda = [layer.linear_attn for layer in self.layers if layer.is_kda] + mla = [layer.self_attn for layer in self.layers if not layer.is_kda] + return latent_push( + step, + capturing=True, + breakable=is_in_breakable_cuda_graph(), + kda_token_states=self.kda_token_states, + kda_decode_kernels=all(m.takes_k3_kernels and m.k3_buffers is not None for m in kda), + mla_decode_branch=all(m.will_run_decode_branch(attn_metadata, step) for m in mla[:1]), + ) + def _defer_moe_tail( self, i: int, @@ -1849,6 +1875,8 @@ def forward( num_tokens = (input_ids if inputs_embeds is None else inputs_embeds).shape[0] step = decode_step(attn_metadata, num_tokens) + if step is not None and self._latent_push(attn_metadata, step): + step = replace(step, latent_push=True) # A decode step embeds and norms for layer 0 in one launch, the embedding written as layer 0's first snapshot. prenormed = None if ( @@ -2897,6 +2925,9 @@ class DecodeStep: num_tokens: int num_requests: Optional[int] = None tokens_per_request: Optional[int] = None + #: Whether the MoE layers' latent all-reduce is the routed experts' push form plus ``comm/k3_latent_reduce`` + #: (``decode_moe.py``): set by ``KimiLinearModel.forward`` from ``latent_push``. + latent_push: bool = False @property def small(self) -> bool: @@ -2915,6 +2946,30 @@ def wide(self) -> bool: return self.decode and not self.small +def latent_push( + step: Optional[DecodeStep], + *, + capturing: bool, + breakable: bool, + kda_token_states: bool, + kda_decode_kernels: bool, + mla_decode_branch: bool, +) -> bool: + """Whether a step's MoE layers push their routed partials (``DecodeStep.latent_push``): a pure decode step of at + most 8 tokens whose attention layers all run the decode kernels, captured into a CUDA graph and not inside a + breakable one. The KDA layers run them on one token per request with every KDA layer taking the K3 kernels + (``kda_decode_kernels``), and on every verify width with the per-token states (``kda_token_states``); the MLA + layers where their decode branch takes the step (``mla_decode_branch``, the same for every MLA layer of a step). + Every other step keeps the routed experts' all-reduce: the exchange's call-order invariant needs every kernel + between a reduce and the next push to wait for its predecessor, which only those kernels were checked for, and a + graph launch orders every pushing replay behind whatever ran before it.""" + if step is None or not (step.decode and step.small) or not capturing or breakable: + return False + if not (kda_token_states or (step.tokens_per_request == 1 and kda_decode_kernels)): + return False + return mla_decode_branch + + def decode_step(attn_metadata: AttentionMetadata, num_tokens: int) -> Optional[DecodeStep]: """The step's shape if any Kimi K3 decode kernel takes it, else None (the generic path runs). @@ -3183,6 +3238,15 @@ def _build_decode_moe(self, comm: _decode_comm.K3DecodeComm) -> int: moe.decode_gemvs = self.model.decode_gemvs first = takes[0].decode_moe first.warm_up(takes[0]) + logger.info( + "Kimi K3 MoE decode path: the latent all-reduce at <= 8 tokens on " + + ( + "the routed experts' push form and k3_latent_reduce on pushing steps (a pure decode step captured " + "into a CUDA graph on the decode kernels), else the routed experts' all-reduce" + if state.exchange is not None + else f"the routed experts' all-reduce (no latent exchange at TP {mapping.tp_size})" + ) + ) comm.compile_tail(takes[0].moe_hidden_size, first.shared_cols, first.tail_weight) return len(takes) diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_construction.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_construction.py index 6c0478735a57..52171384a30c 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_construction.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_construction.py @@ -74,3 +74,24 @@ def test_tp16_moetp16ep1_decodes_without_speculation(): AssertionError, match="moe_tensor_parallel_size 4 and moe_expert_parallel_size 4" ): route_b._check_construction(_config(16, 1, spec_config=object())) + + +def test_tp16_moetp16ep1_keeps_no_per_token_kda_states(monkeypatch): + """The text model answers ``kda_token_states`` (False, without speculation), which a captured step's latent-push + decision reads: a pure decode step on the decode kernels pushes, as on ``tp16_moetp4ep4`` without per-token + states.""" + model = route_b.KimiLinearModel.__new__(route_b.KimiLinearModel) + assert model.kda_token_states is False + monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: True) + monkeypatch.setattr(route_b, "is_in_breakable_cuda_graph", lambda: False) + kda = types.SimpleNamespace( + is_kda=True, + linear_attn=types.SimpleNamespace(takes_k3_kernels=True, k3_buffers=object()), + ) + mla = types.SimpleNamespace( + is_kda=False, + self_attn=types.SimpleNamespace(will_run_decode_branch=lambda metadata, step: True), + ) + object.__setattr__(model, "layers", [kda, mla]) + assert model._latent_push(None, route_b.DecodeStep(1, 1, 1)) + assert not model._latent_push(None, route_b.DecodeStep(8, 1, 8)) From d6613a1d0ab76d80c3cff312d78c4e9fc5471d01 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 12:53:56 -0700 Subject: [PATCH 149/161] [None][feat] modeling_v2 Kimi K3 tp16_moetp16ep1: the MoE decode path on route B's engines A MoE layer of tp16_moetp16ep1 (every expert on each rank) now runs the decode path on steps of at most 8 tokens, as route B blocks in decode_moe.py and modeling.py: moe/k3_moe_front, then the routed experts on moe/k3_moe_m1 at one token, moe/k3_moe_m2 at two and moe/k3_moe over all 896 experts up to 8, the latent all-reduce and the row-parallel tail. On a pushing step (DecodeStep.latent_push, route A's rule as carried into this copy: a pure decode step captured into a non-breakable CUDA graph whose attention layers all run the decode kernels) the engine of the token count runs its push form into the carried K3DecodeMoe.exchange, and comm/k3_latent_reduce sums the partials. Every other step returns the partials to the routed experts' MNNVL all-reduce. Both sum in the one-shot's order, so the bits are the same. K3DecodeMoe.create compiles the engines' push builds along with the exchange. The routing is the front's (the noaux_tc arithmetic of moe/kimi_k3_noaux_tc_mxfp8_quant), not the generic path's TRTLLM-Gen routing, so at a near-tie a layer can select another expert than before. Steps of more than 8 tokens keep the generic path: k3_moe's wide build does not fit 896 local experts. The engines compile the SiTU caps 4 and 25 in, so a checkpoint with other caps keeps the generic path. Tests, listed in l0_b200.yml and l0_gb200_multi_gpus.yml: test_modeling_v2_kimi_k3_route_b_moe.py (the engines' SiTU caps against the kernels', the decline on other caps, the steps the path takes, the engine and latent all-reduce a step runs) and the 4-rank op matrix test_modeling_v2_kimi_k3_route_b_decode_moe_op_matrix.py (each engine against k3_moe and the all-reduce, push and reduce over sequences, graph replays pushing against eager steps, the routing against noaux_tc in torch). Signed-off-by: Vasanth Sabavat --- .../decode_moe.py | 175 ++++++-- .../modeling.py | 31 +- .../test_lists/test-db/l0_b200.yml | 2 + .../test-db/l0_gb200_multi_gpus.yml | 2 + .../_kimi_k3_route_b_decode_moe_op_matrix.py | 398 ++++++++++++++++++ ...v2_kimi_k3_route_b_decode_moe_op_matrix.py | 32 ++ .../test_modeling_v2_kimi_k3_route_b_moe.py | 148 +++++++ 7 files changed, 735 insertions(+), 53 deletions(-) create mode 100644 tests/unittest/_torch/modeling_v2/comm/_kimi_k3_route_b_decode_moe_op_matrix.py create mode 100644 tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_kimi_k3_route_b_decode_moe_op_matrix.py create mode 100644 tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_route_b_moe.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py index a2fa6889283b..bbdb744ca13e 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py @@ -1,18 +1,21 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""The decode path's latent MoE on the catalog's Kimi K3 MoE entries. +# >>> route B: this target's MoE engines (every expert on each rank) +"""The decode path's latent MoE on the catalog's Kimi K3 MoE entries, with every expert on each rank (moe TP16 x EP1). **At most 8 tokens** (`MAX_TOKENS`), a MoE layer runs: * `moe/k3_moe_front`, one kernel: this rank's slice of the MoE head (its latent-down rows and its router rows) as one GEMV, the slices' all-gather over the TP group's `K3MoeHeadWorkspace`, the top-16 routing, the MXFP8 latent, and - the shared experts' gate_up + SiTU; -* `moe/k3_moe`: this rank's routed partial over its experts (on a `K3MoeState`); + the shared experts' gate_up + SiTU. The routing is the noaux_tc arithmetic of `moe/kimi_k3_noaux_tc_mxfp8_quant`; + the generic path's TRTLLM-Gen MoE routes inside its own kernel, so a near-tie can select another expert there; +* the routed experts of all 896 experts over this rank's intermediate slice: `moe/k3_moe_m1` at one token, + `moe/k3_moe_m2` at two, `moe/k3_moe` (on a `K3MoeState`) at three to eight; * the latent all-reduce. On a pushing step (`DecodeStep.latent_push`: a pure decode step captured into a CUDA graph - whose attention layers all run the decode kernels) the routed experts run as `moe/k3_moe`'s push form, which stores - this rank's partial into every rank's `K3LatentExchange`, and `comm/k3_latent_reduce` sums the partials in the MNNVL - one-shot's order, so with the same bits, while the experts' grid completes. On every other step it is the routed - experts' all-reduce (one-shot, as `decode_comm.use_decode_one_shot` sets); + whose attention layers all run the decode kernels) the engine runs its push form, which stores this rank's partial + into every rank's `K3LatentExchange`, and `comm/k3_latent_reduce` sums the partials in the MNNVL one-shot's order, + so with the same bits, while the experts' grid completes. On every other step it is the routed experts' all-reduce + (one-shot, as `decode_comm.use_decode_one_shot` sets); * the tail. The latent norm's weight is folded into the latent up projection at load, so `[RMSNorm(latent) slice | shared activation] @ [latent up columns | shared down]` is this rank's share of the MoE output. Where the next layer's pre-attention step (or the final norm) reduces it, the layer hands it on unreduced @@ -20,18 +23,17 @@ kernel (`decode_comm.py`). Elsewhere the replicated tail runs: the latent RMS applied to the fp32 output of one GEMV with the folded latent up weight, plus the shared experts' down projection and its all-reduce. -**A wide decode step** (9 to 64 tokens, `WIDE_MAX_TOKENS`) keeps the sharded head and the row-parallel tail on -M-general ops: the head GEMV, `comm/mnnvl_allgather_split`, then `moe/k3_route_quant` and `moe/k3_moe` (on a -`K3MoeWideState`) beside the shared gate_up + SiTU, the latent all-reduce, and one GEMV of -`[RMSNorm(latent) slice | padding | shared activation]` with the tail weight. That is this rank's unreduced share, -which the consumer reduces with a plain all-reduce. +**More than 8 tokens** run the generic path: `k3_moe`'s wide build does not fit 896 local experts, so the wide step +below (`_wide`, `tp16_moetp4ep4`'s) never runs here. The engines compile the SiTU caps 4 and 25 in +(`ENGINE_SITU_CAPS`): a checkpoint with other caps keeps the generic path (`layout_gaps`). The GEMVs run on the decode GEMV sites of `decode_gemv.py` where they take the call, else on the stock GEMM ops. -`K3DecodeMoe` holds what every MoE layer shares: the head workspace and the latent exchange (both collective over -the TP group), the two `k3_moe` builds' scratch, and the TP group's MNNVL workspace (`decode_comm.K3DecodeComm`'s). -`K3DecodeMoeLayer` holds one layer's decode weights and its `k3_moe` counters. The target builds both in -`post_load_weights`, before any CUDA-graph capture, and runs every kernel once there so none compiles under a capture. +`K3DecodeMoe` holds what every MoE layer shares: the head workspace and the latent exchange (both collective over the +TP group), the engines' workspaces (`K3MoeM1State`, `K3MoeM2State` and the `k3_moe` build's scratch), and the TP +group's MNNVL workspace (`decode_comm.K3DecodeComm`'s). `K3DecodeMoeLayer` holds one layer's decode weights and its +engine handles. The target builds both in `post_load_weights`, before any CUDA-graph capture, and runs every kernel +once there so none compiles under a capture. Every MoE layer's push and reduce go to the one exchange, in the stream's order: each push is followed by exactly one reduce of the same token count before the next push, on every rank in the same order (the exchange's call-order @@ -39,6 +41,7 @@ launches without programmatic dependent launch), which the decode kernels a pushing step runs do; pushes run only in CUDA-graph replays, so no other step's kernels sit between two pushes. """ +# <<< route B from __future__ import annotations @@ -72,6 +75,22 @@ k3_moe_front, weight_supported, ) + +# >>> route B: the one- and two-token engines +from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.k3_moe_m1 import ( + K3MoeM1Layer, + K3MoeM1State, + k3_moe_m1, + k3_moe_m1_push, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.k3_moe_m2 import ( + K3MoeM2Layer, + K3MoeM2State, + k3_moe_m2, + k3_moe_m2_push, +) + +# <<< route B from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.k3_route_quant import k3_route_quant from tensorrt_llm._torch.modules.multi_stream_utils import maybe_execute_in_parallel @@ -85,29 +104,53 @@ # (comm/k3_sandwich_tail takes the TP16 tail weight [7168, 256 + 384]). TAIL_K_TILE = 256 +# >>> route B: the routed experts' SiTU caps (activation_situ_beta, activation_situ_linear_beta) that k3_moe_m1, +# k3_moe_m2 and k3_moe compile in (their kernels' SITU_GATE_CAP and SITU_LINEAR_CAP) +ENGINE_SITU_CAPS = (4.0, 25.0) +# <<< route B + @dataclass(eq=False) class K3DecodeMoe: + # >>> route B: the engines' workspaces; no wide build """What every MoE layer's decode path shares on one device: the TP group's ``K3MoeHeadWorkspace`` (the front's - all-gather), the ``k3_moe`` builds for up to 8 and up to 64 tokens (their scratch; the layers run one at a time on - one stream), the TP group's ``MnnvlWorkspace`` for a wide step's head all-gather, and the TP group's + all-gather), the engines' workspaces (``k3_moe_m1``, ``k3_moe_m2`` and the ``k3_moe`` build for up to 8 tokens; + the layers run one at a time on one stream), the TP group's ``MnnvlWorkspace``, and the TP group's ``K3LatentExchange`` for a pushing step's latent all-reduce (None: every step keeps the routed experts' - all-reduce). Built by `create`.""" + all-reduce). ``wide`` is None: ``k3_moe``'s wide build does not fit 896 local experts. Built by `create`.""" + + # <<< route B head: K3MoeHeadWorkspace small: K3MoeState wide: K3MoeWideState mnnvl: MnnvlWorkspace exchange: Optional[K3LatentExchange] = None + # >>> route B: the one- and two-token engines' workspaces (``create`` builds both) + m1: Optional[K3MoeM1State] = None + m2: Optional[K3MoeM2State] = None + # <<< route B + # >>> route B: every expert on each rank: the engines with their push builds for the exchange and the k3_moe build + # for up to 8 tokens; no wide build @classmethod def create( - cls, mapping, device, i_tp: int, num_local: int, mnnvl: MnnvlWorkspace, push: bool = True + cls, + mapping, + device, + i_tp: int, + num_local: int, + mnnvl: MnnvlWorkspace, + push: bool = True, + *, + i_logical: int, ) -> "K3DecodeMoe": - """The state for ``mapping``'s TP group on ``device``, for experts of ``i_tp`` intermediate columns per rank, - ``num_local`` of them on this rank. Collective (the head workspace and the latent exchange): every rank of the - group calls it at the same point, eagerly, before any CUDA-graph capture. ``push``: build the latent exchange - (where it takes the TP size: 4, 8 or 16 ranks); without it every step keeps the routed experts' all-reduce.""" + """The state for ``mapping``'s TP group on ``device``, for experts of ``i_tp`` intermediate columns per rank + (the loader's zero-padded width; ``i_logical`` of them hold the checkpoint's slice, which k3_moe_m1 and + k3_moe_m2 stream), ``num_local`` of them on this rank. Collective (the head workspace and the latent + exchange): every rank of the group calls it at the same point, eagerly, before any CUDA-graph capture. + ``push``: build the latent exchange (where it takes the TP size: 4, 8 or 16 ranks) and the engines' push builds + for it; without it every step keeps the routed experts' all-reduce.""" head = K3MoeHeadWorkspace.create(mapping) exchange = None if push: @@ -116,14 +159,19 @@ def create( # Raised before any collective step, on every rank of the group alike: a TP size the exchange does not take. except ValueError: exchange = None + builds = () if exchange is None else ((mapping.tp_size, 1),) return cls( head, K3MoeState(device, i_tp, num_local), - K3MoeWideState(device, i_tp, num_local), + None, mnnvl, exchange, + K3MoeM1State.create(device, i_logical, i_tp, num_local, push=builds), + K3MoeM2State.create(device, i_logical, i_tp, num_local, push=builds), ) + # <<< route B + def _experts(moe: nn.Module) -> tuple: """The routed experts' TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers ``k3_moe`` reads in place.""" @@ -145,6 +193,10 @@ def layout_gaps(moe: nn.Module, tp_size: int, max_snapshots: int) -> list: gate_up, down = shared.gate_up_proj.weight, shared.down_proj.weight if not moe._reduce_routed_output: return ["a routed output the model does not reduce"] + # >>> route B: k3_moe_m1, k3_moe_m2 and k3_moe compile the SiTU caps in + if tuple(moe._situ_betas) != ENGINE_SITU_CAPS: + return [f"SiTU caps {tuple(moe._situ_betas)}; the MoE engines compile {ENGINE_SITU_CAPS}"] + # <<< route B if (moe.num_experts, moe.top_k, moe.moe_hidden_size, moe.hidden_size) != (896, 16, 3584, 7168): return [ f"experts / top-k / latent / hidden {(moe.num_experts, moe.top_k, moe.moe_hidden_size)}" @@ -222,6 +274,10 @@ class K3DecodeMoeLayer: shared_cols: int small: K3MoeLayer wide: K3MoeLayer + # >>> route B: the one- and two-token engines' handles + m1: K3MoeM1Layer + m2: K3MoeM2Layer + # <<< route B @classmethod def create( @@ -259,30 +315,35 @@ def create( lo=lo, width=width, shared_cols=gate_up.shape[0] // 2, + # >>> route B: the engines' handles; no wide build small=state.small.layer(*weights), - wide=state.wide.layer(*weights), + wide=None, + m1=state.m1.layer(*weights), + m2=state.m2.layer(*weights), + # <<< route B ) def warm_up(self, moe: nn.Module) -> None: - """One call of every kernel of the decode path on zero inputs (M = 1), so none compiles under a capture: the - front (collective: every rank makes the same call), both ``k3_moe`` builds, ``k3_route_quant`` and, with the - latent exchange, the push build and the reduce (one push + reduce pair, collective). Once per model: every + # >>> route B: the engines; no wide build + """One call of every kernel of the decode path on zero inputs, so none compiles under a capture: the front + (collective: every rank makes the same call), each engine (``k3_moe_m1``, ``k3_moe_m2``, ``k3_moe``) and, with + the latent exchange, its push build, each push followed by the reduce (collective). Once per model: every layer's calls compile the same builds.""" + # <<< route B device = self.front_weight.device x = torch.zeros(1, moe.hidden_size, dtype=torch.bfloat16, device=device) ids, weights, x_fp8, x_sf, _ = self._front(moe, x) offset = moe.routed_experts.backend.slot_start - k3_moe(x_fp8, x_sf, ids, weights, offset, self.small) + # >>> route B: every engine returning and, with the latent exchange, pushing (k3_moe's builds compile here), + # each push followed by the reduce; no wide build exchange = self.state.exchange - if exchange is not None: - k3_moe_push(x_fp8, x_sf, ids, weights, offset, self.small, exchange) - k3_latent_reduce(1, exchange) - logits = torch.zeros(1, moe.num_experts, dtype=torch.float32, device=device) - latent = torch.zeros(1, moe.moe_hidden_size, dtype=torch.bfloat16, device=device) - ids, weights, x_fp8, x_sf = k3_route_quant( - logits, self.bias, latent, float(moe.gate.routed_scaling_factor), early_trigger=True - ) - k3_moe(x_fp8, x_sf, ids, weights, offset, self.wide) + for rows in (1, 2, 3): + args = [t.expand(rows, -1).contiguous() for t in (x_fp8, x_sf, ids, weights)] + self._routed(*args, offset) + if exchange is not None: + self._routed(*args, offset, push=True) + k3_latent_reduce(rows, exchange) + # <<< route B torch.cuda.synchronize(device) def _front(self, moe: nn.Module, x: torch.Tensor): @@ -297,12 +358,16 @@ def _front(self, moe: nn.Module, x: torch.Tensor): ) def takes(self, hidden_states: torch.Tensor, step, partial_tail: bool) -> bool: + # >>> route B: no wide step """Whether this decode path runs the layer on ``step`` (bf16 rows ``hidden_states``): any step of at most - `MAX_TOKENS` tokens; with ``partial_tail``, also a wide decode step.""" + `MAX_TOKENS` tokens, with or without ``partial_tail``.""" + # <<< route B rows = hidden_states.shape[0] if hidden_states.dtype != torch.bfloat16 or step is None: return False - return 0 < rows <= MAX_TOKENS or (partial_tail and step.wide and rows <= WIDE_MAX_TOKENS) + # >>> route B: no wide step (k3_moe's wide build does not fit 896 local experts) + return 0 < rows <= MAX_TOKENS + # <<< route B def forward( self, @@ -322,12 +387,13 @@ def forward( ids, weights, x_fp8, x_sf, shared_act = self._front(moe, hidden_states) offset = moe.routed_experts.backend.slot_start exchange = self.state.exchange + # >>> route B: the engine of the token count if push and exchange is not None: - k3_moe_push(x_fp8, x_sf, ids, weights, offset, self.small, exchange) + self._routed(x_fp8, x_sf, ids, weights, offset, push=True) latent = k3_latent_reduce(hidden_states.shape[0], exchange) else: - routed = k3_moe(x_fp8, x_sf, ids, weights, offset, self.small) - latent = moe.routed_experts.all_reduce(routed) + latent = moe.routed_experts.all_reduce(self._routed(x_fp8, x_sf, ids, weights, offset)) + # <<< route B if partial_tail: return PendingTail( latent.contiguous(), @@ -346,6 +412,27 @@ def forward( ) return (up * scale + shared_out.float()).bfloat16() + # >>> route B: the engine of a token count + def _routed(self, x_fp8, x_sf, ids, weights, offset: int, push: bool = False): + """This rank's routed partial of the front's outputs on ``k3_moe_m1`` at one token, ``k3_moe_m2`` at two, else + ``k3_moe`` (up to `MAX_TOKENS`); with ``push``, stored into every rank's latent exchange instead (None).""" + rows = x_fp8.shape[0] + if push: + exchange = self.state.exchange + if rows == 1: + k3_moe_m1_push(x_fp8, x_sf, ids, weights, offset, self.m1, exchange) + elif rows == 2: + k3_moe_m2_push(x_fp8, x_sf, ids, weights, offset, self.m2, exchange) + else: + k3_moe_push(x_fp8, x_sf, ids, weights, offset, self.small, exchange) + return None + if rows == 1: + return k3_moe_m1(x_fp8, x_sf, ids, weights, offset, self.m1) + if rows == 2: + return k3_moe_m2(x_fp8, x_sf, ids, weights, offset, self.m2) + return k3_moe(x_fp8, x_sf, ids, weights, offset, self.small) + + # <<< route B def _wide(self, moe: nn.Module, hidden_states: torch.Tensor, gemvs) -> torch.Tensor: """A wide decode step's MoE: this rank's unreduced share of the output, ``[M, hidden]`` bf16.""" x = hidden_states.contiguous() diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py index 1ff1ce88697d..e336b329d484 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py @@ -51,10 +51,14 @@ attention workspace, the decode GEMVs' state, the TP group's MNNVL and sandwich workspaces) lives in typed objects this target creates in `post_load_weights`, before any graph capture. -The MoE layers run the generic path on every step. `decode_moe.py` is `tp16_moetp4ep4`'s MoE decode path, which this -target does not build (its wide `k3_moe` build does not fit 896 local experts); this layout's own MoE engines -(`moe/k3_moe_m1` and `moe/k3_moe_m2` at one and two tokens, `moe/k3_moe` over all 896 experts up to 8) come with their -wiring as route B blocks. +A MoE layer on a step of at most 8 tokens runs `decode_moe.py` with every expert on each rank: `moe/k3_moe_front` +(this rank's head slice, its all-gather, the top-16 routing, the MXFP8 latent and the shared experts' gate_up + SiTU), +then the routed experts as `moe/k3_moe_m1` at one token, `moe/k3_moe_m2` at two and `moe/k3_moe` over all 896 experts +up to 8, the latent all-reduce and the row-parallel tail. On a pushing step (`DecodeStep.latent_push`) the engines +push their partials into the TP group's `K3LatentExchange` and `comm/k3_latent_reduce` sums them; other steps use the +routed experts' all-reduce, which sums in the same order. Its routing is the front's (the noaux_tc arithmetic), not +the generic path's TRTLLM-Gen routing. Wider steps, and every step of a checkpoint whose SiTU caps differ from the ones +the engines compile in, run the generic path: `k3_moe`'s wide build does not fit 896 local experts. **What this target asserts rather than adapts**: SM 10.0; the topology above, with the expert split set explicitly; no speculative decoding; the MXFP4 checkpoint's quantization (W4A16_MXFP4 with no per-layer declarations, so the @@ -189,6 +193,10 @@ "k3_latent_reduce", "k3_route_quant", "mnnvl_allgather_split", + # >>> route B: the MoE engines of one and two tokens + "k3_moe_m1", + "k3_moe_m2", + # <<< route B # The DSpark drafter's decode path (K3DSparkDrafter): its GEMV sites, its block attention, and its all-reduces # with the residual add and RMSNorm (the split context projection's with hidden_norm). "k3_ctm_gemv", @@ -3177,10 +3185,7 @@ def post_load_weights(self) -> None: for layer in layers: layer.decode_comm = comm _decode_comm.use_decode_one_shot(self) - # >>> route B: no MoE decode path yet (its build is tp16_moetp4ep4's; k3_moe's wide build does not fit 896 - # local experts). The MoE layers run the generic path until route B's own engines are wired. - moe_layers = 0 - # <<< route B + moe_layers = 0 if comm is None else self._build_decode_moe(comm) logger.info( # >>> route B: no verify kernels "Kimi K3 decode kernels: KDA on k3_kda_decode_attn " @@ -3196,7 +3201,9 @@ def post_load_weights(self) -> None: if comm is not None else "unfused (an attention all-reduce does not run over MNNVL)" ) - + f"; MoE on k3_moe_front, k3_moe and the row-parallel tail ({moe_layers} layers)" + # >>> route B: this target's MoE engines + + f"; MoE on k3_moe_front, k3_moe_m1 / k3_moe_m2 / k3_moe and the row-parallel tail ({moe_layers} layers)" + # <<< route B ) self._gate_spec_worker_kernels(comm) self._gate_drafter_comm(comm) @@ -3229,6 +3236,12 @@ def _build_decode_moe(self, comm: _decode_comm.K3DecodeComm) -> int: backend.w3_w1_weight.shape[1] // 2, backend.expert_size_per_partition, comm.mnnvl, + # >>> route B: the experts' logical slice (192 of the loader's 256), which k3_moe_m1 / k3_moe_m2 stream + i_logical=( + getattr(backend.quant_method, "intermediate_size_per_partition_lean", None) + or backend.intermediate_size_per_partition + ), + # <<< route B ) for moe in takes: _decode_moe.fold_latent_norm(moe) diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index 6f437d362e9f..678d46093dec 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -170,6 +170,8 @@ l0_b200: # l0_b300's modeling_v2 entry, the only other one that collects them, is waived (nvbugs/6853741). - unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_construction.py - unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drift.py + # tp16_moetp16ep1's MoE decode path: the layers and steps it takes (its kernels: l0_gb200_multi_gpus.yml). + - unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_route_b_moe.py # KDA runtime: host-derived prefill metadata and the bf16 state pool # round-trip. Both are single-device cases that use GPU 0 only. - unittest/_torch/modules/kimi_kda/test_kda_host_metadata.py diff --git a/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml b/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml index 7b6e4fa63f17..86bd73561dda 100644 --- a/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml +++ b/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml @@ -71,6 +71,8 @@ l0_gb200_multi_gpus: - unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_front_op_matrix.py # The Kimi K3 target's decode-path collectives over those entries, 4 ranks - unittest/_torch/modeling_v2/comm/test_modeling_v2_kimi_k3_decode_comm_op_matrix.py + # tp16_moetp16ep1's MoE decode path: the engines pushing into the latent exchange, 4 ranks + - unittest/_torch/modeling_v2/comm/test_modeling_v2_kimi_k3_route_b_decode_moe_op_matrix.py - unittest/_torch/moe/test_moe_comm.py::TestMoEComm::test_mnnvl_checkpoint_preserves_moe_graph_addresses - unittest/_torch/moe/test_moe_comm.py::TestMoEComm::test_mnnvl_engine_checkpoint_coordination - unittest/_torch/moe/test_moe_comm.py::TestMoEComm::test_mnnvl_checkpoint_failure_is_collective_and_bounded diff --git a/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_route_b_decode_moe_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_route_b_decode_moe_op_matrix.py new file mode 100644 index 000000000000..afda4d0112f7 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_route_b_decode_moe_op_matrix.py @@ -0,0 +1,398 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""The MoE decode path of the Kimi K3 target ``kimi_k3_mxfp4__sm_100__tp16_moetp16ep1`` (``decode_moe.py``: every +expert on each rank, the routed experts pushing into the latent exchange) at W ranks of one GB200 tray (default 4), +at the TP16 per-rank shapes. + + CUDA_VISIBLE_DEVICES=0,1,2,3 python _kimi_k3_route_b_decode_moe_op_matrix.py [--world-size 4] + srun -n 4 --mpi=pmix python _kimi_k3_route_b_decode_moe_op_matrix.py --launcher srun --world-size 4 + +Not a pytest module: one fixed sequence of checks inside one W-rank job, sharing the decode path's collective state. +The collected entry point is ``test_modeling_v2_kimi_k3_route_b_decode_moe_op_matrix.py``. + +Every rank builds MoE layers the way the target's ``post_load_weights`` does (``K3DecodeMoe.create``, +``fold_latent_norm``, ``K3DecodeMoeLayer.create``, ``warm_up``) from stand-in layers holding what the path reads: the +target's KimiK3MoEGate, nn.Linear latent projections, the stock RMSNorm and shared GatedMLP at this TP, and the routed +experts of a TP16 rank (all 896, the 192-wide intermediate slice zero-padded to 256 by TRT-LLM's loader, random +MXFP4). The reference of a step is the same front with the returned-partial ``k3_moe`` of its ``K3MoeState`` and the +routed experts' MNNVL all-reduce. + +Checks: + * build: outside inference mode with autograd on, as post_load_weights runs (the gate's parameters require grad). + * engines (eager steps, which return the partials to the all-reduce): for 1..8 tokens, the latent the path hands on + (its PendingTail) against the reference: bit for bit at 3..8 tokens, within ENGINE_TOL at 1 and 2 (k3_moe_m1 / + k3_moe_m2 against k3_moe, which round the FC1 sums apart); the engine each count ran on (the engines' epochs); + the replicated output against the reference tail on the same latent; the routing against the noaux_tc arithmetic + in torch (the same top-16 experts, the weights within ROUTE_TOL). + * push: three layers on one state over steps of 1, 2, 8, 3, 1, 5, 2 and 7 tokens, each engine's push and + k3_latent_reduce against the same engine's returned partial and the all-reduce, bit for bit; the exchange's + halves alternate across every engine and layer. + * graph: three layers of 2, 8 and 1 tokens captured once pushing (``push=True``, as a pushing step's + ``DecodeStep.latent_push`` has the MoE runtime pass it) and replayed + with rewritten inputs: each replay equal to eager steps (which do not push), bit for bit. + * Every output bitwise equal across the ranks. +""" + +import functools +import math +import sys +from pathlib import Path +from types import SimpleNamespace + +import torch +from torch import nn + +sys.path.insert(0, str(Path(__file__).resolve().parent)) +import _lockstep as ls # noqa: E402 + +assert torch.cuda.is_available(), "the Kimi K3 target's MoE decode path requires CUDA devices" + +DEADLINE_S = 1500 +H = 7168 +LATENT = 3584 +NUM_EXPERTS = 896 +# A TP16 rank's 192-wide intermediate slice, zero-padded to whole tiles (256) by the loader. +I_TP, I_PAD, MOE_TP = 192, 256, 16 +SV = 32 +SHARED_PER_RANK = 384 +SITU_CAPS = (4.0, 25.0) +ENGINE_TOL = 8e-3 # max |err| / max |ref| of the latent, k3_moe_m1 / k3_moe_m2 against k3_moe +ROUTE_TOL = 2e-2 # the bf16 routing weights against the fp32 torch routing +SEQUENCE = (1, 2, 8, 3, 1, 5, 2, 7) +GRAPH_ROWS = (2, 8, 1) + +R = None +T = None # the target module +DM = None # its decode_moe module + + +def _gen(seed): + return torch.Generator(device="cuda").manual_seed(seed) + + +def _normal(g, shape, scale=1.0, offset=0.0): + return (offset + scale * torch.randn(shape, generator=g, device="cuda")).bfloat16() + + +def _rand_mxfp4(rows, k, k_full, gen): + """Random checkpoint-format MXFP4: packed [rows, k / 2] (low nibble = even k), E8M0 per 32 k, scaled so a + k_full-long dot product lands near std 3.""" + codes = torch.randint(0, 16, (rows, k), dtype=torch.uint8, device="cuda", generator=gen) + base = 127 + round(0.5 * math.log2(0.01057 / k_full)) + exps = torch.randint( + base, base + 6, (rows, k // SV), dtype=torch.uint8, device="cuda", generator=gen + ) + return (codes[:, 0::2] | (codes[:, 1::2] << 4)).contiguous(), exps + + +@functools.lru_cache(maxsize=None) +def _experts(seed: int): + """A TP16 rank's 896 experts through W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod's loader (the buffers the engines read): + the 192-wide shard generated as rank 0 of tensors that hold exactly it, sliced and padded to 256 by the loader.""" + from tensorrt_llm._torch.moe.fused_moe.quantization import W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod + + method = W4A8MXFP4MXFP8TRTLLMGenFusedMoEMethod() + module = SimpleNamespace(tp_size=MOE_TP, tp_rank=0, scaling_vector_size=SV, intermediate_size=I_TP * MOE_TP, + intermediate_size_per_partition=I_TP, hidden_size=LATENT) # fmt: skip + kw = dict(dtype=torch.uint8, device="cuda") + w31 = torch.zeros(NUM_EXPERTS, 2 * I_PAD, LATENT // 2, **kw) + w31s = torch.zeros(NUM_EXPERTS, 2 * I_PAD, LATENT // SV, **kw) + w2 = torch.zeros(NUM_EXPERTS, LATENT, I_PAD // 2, **kw) + w2s = torch.zeros(NUM_EXPERTS, LATENT, I_PAD // SV, **kw) + gen = _gen(seed) + for e in range(NUM_EXPERTS): + gate, gate_s = _rand_mxfp4(I_TP, LATENT, LATENT, gen) + up, up_s = _rand_mxfp4(I_TP, LATENT, LATENT, gen) + down, down_s = _rand_mxfp4(LATENT, I_TP, I_TP * MOE_TP, gen) + method.load_expert_w3_w1_weight(module, gate, up, w31[e]) + method.load_expert_w2_weight(module, down, w2[e]) + method.load_expert_w3_w1_weight_scale_mxfp4(module, gate_s, up_s, w31s[e]) + method.load_expert_w2_weight_scale_mxfp4(module, down_s, w2s[e]) + torch.cuda.synchronize() + return w31, w31s, w2, w2s + + +def _moe_layer(seed: int, experts): + """A MoE layer with what the decode path reads, from the modules KimiK3MoERuntime builds (the gate's and the latent + projections' parameters require grad, as in the model), and ``experts`` (every expert on this rank).""" + from tensorrt_llm._torch.distributed import AllReduce + from tensorrt_llm._torch.model_config import ModelConfig + from tensorrt_llm._torch.modules.gated_mlp import GatedMLP + from tensorrt_llm._torch.modules.rms_norm import RMSNorm + from tensorrt_llm._torch.modules.situ import SituAndMul + from tensorrt_llm.functional import AllReduceStrategy + + cfg = SimpleNamespace( + num_experts_per_token=16, + num_experts=NUM_EXPERTS, + routed_scaling_factor=2.827, + moe_router_activation_func="sigmoid", + num_expert_group=1, + topk_group=1, + moe_renormalize=True, + hidden_size=H, + ) + with torch.device("cuda"): + gate = T.KimiK3MoEGate(cfg, logits_gemm_dtype=torch.bfloat16) + down = nn.Linear(H, LATENT, bias=False, dtype=torch.bfloat16) + up = nn.Linear(LATENT, H, bias=False, dtype=torch.bfloat16) + norm = RMSNorm(hidden_size=LATENT, eps=1e-5, dtype=torch.bfloat16) + shared = GatedMLP( + hidden_size=H, + intermediate_size=SHARED_PER_RANK * R.world, + bias=False, + activation=SituAndMul( + beta=SITU_CAPS[0], linear_beta=SITU_CAPS[1], use_fused_activation=True + ), + dtype=torch.bfloat16, + config=ModelConfig(mapping=R.mapping, allreduce_strategy=AllReduceStrategy.MNNVL), + reduce_output=True, + layer_idx=seed, + is_shared_expert=True, + ) + g = _gen(5000 + seed) # replicated + gr = _gen(5100 + seed * 16 + R.rank) # this rank's shared expert slices + with torch.no_grad(): + gate.weight.copy_(_normal(g, gate.weight.shape, 0.02)) + gate.e_score_correction_bias.copy_( + 0.05 * torch.randn(NUM_EXPERTS, generator=g, device="cuda") + ) + down.weight.copy_(_normal(g, down.weight.shape, 0.02)) + up.weight.copy_(_normal(g, up.weight.shape, 0.02)) + norm.weight.copy_(_normal(g, norm.weight.shape, 0.1, 1.0)) + shared.gate_up_proj.weight.copy_(_normal(gr, shared.gate_up_proj.weight.shape, 0.02)) + shared.down_proj.weight.copy_(_normal(gr, shared.down_proj.weight.shape, 0.02)) + backend = SimpleNamespace( + w3_w1_weight=experts[0], + w3_w1_weight_scale=experts[1], + w2_weight=experts[2], + w2_weight_scale=experts[3], + expert_size_per_partition=NUM_EXPERTS, + intermediate_size_per_partition=I_TP, + quant_method=SimpleNamespace(intermediate_size_per_partition_lean=I_TP), + slot_start=0, + ) + all_reduce = AllReduce( + mapping=R.mapping, strategy=AllReduceStrategy.MNNVL, dtype=torch.bfloat16 + ) + return SimpleNamespace( + num_experts=NUM_EXPERTS, + top_k=cfg.num_experts_per_token, + moe_hidden_size=LATENT, + hidden_size=H, + gate=gate, + routed_expert_down_proj=down, + routed_expert_up_proj=up, + routed_expert_norm=norm, + shared_experts=shared, + routed_experts=SimpleNamespace(backend=backend, all_reduce=all_reduce), + _situ_betas=SITU_CAPS, + _reduce_routed_output=True, + moe_main_event=torch.cuda.Event(), + moe_shared_event=torch.cuda.Event(), + shared_expert_stream=torch.cuda.Stream(), + ) + + +BUILT = None + + +def check_build(): + """The decode path built as post_load_weights builds it, outside inference mode with autograd on: the shared state + (collective: the head workspace, the latent exchange; the engines' push builds compile), three MoE layers' folded + latent norms and decode weights, and one call of every kernel (warm_up: the front, each engine's push, the + reduce).""" + global BUILT + if R.world not in (4, 8, 16): + raise AssertionError( + f"the front and the latent exchange run 4, 8 or 16 ranks, not {R.world}" + ) + assert torch.is_grad_enabled() and not torch.is_inference_mode_enabled() + experts = _experts(7) + moes = [_moe_layer(seed, experts) for seed in (22, 23, 24)] + assert moes[0].gate.e_score_correction_bias.requires_grad + mnnvl = T._decode_comm.MnnvlWorkspace.create(R.mapping, T._decode_comm.MNNVL_BUFFER_BYTES) + device = torch.device("cuda", torch.cuda.current_device()) + backend = moes[0].routed_experts.backend + state = DM.K3DecodeMoe.create( + R.mapping, + device, + backend.w3_w1_weight.shape[1] // 2, + backend.expert_size_per_partition, + mnnvl, + i_logical=backend.quant_method.intermediate_size_per_partition_lean, + ) + assert state.wide is None and state.exchange is not None + assert state.m1.push_compiled(R.world) and state.m2.push_compiled(R.world) + layers = [] + for moe in moes: + DM.fold_latent_norm(moe) + layers.append(DM.K3DecodeMoeLayer.create(moe, state, R.rank, R.world)) + layers[0].warm_up(moes[0]) + BUILT = SimpleNamespace(state=state, moes=moes, layers=layers) + if R.rank == 0: + print(f"[rank 0] built 3 MoE layers on 896 local experts at world {R.world}", flush=True) + + +def _x(seed, rows): + """The MoE input of a step: the same rows on every rank, as the decode path receives them.""" + return _normal(_gen(9000 + seed), (rows, H), 0.5) + + +def _reference(layer, moe, x): + """The front, the returned-partial k3_moe of the layer's K3MoeState, the routed experts' MNNVL all-reduce: + (latent [rows, 3584], shared activation, routing ids, routing weights).""" + ids, weights, x_fp8, x_sf, shared_act = layer._front(moe, x) + routed = DM.k3_moe(x_fp8, x_sf, ids, weights, 0, layer.small) + return moe.routed_experts.all_reduce(routed), shared_act, ids, weights + + +def _wired(layer, moe, x, push=False): + pending = layer.forward(moe, x, None, partial_tail=True, push=push) + assert isinstance(pending, T._decode_comm.PendingTail), type(pending) + return pending + + +def _engine_epochs(): + st = BUILT.state + return st.m1.epochs.clone(), st.m2.epochs.clone() + + +def _routing_ref(moe, x): + """The noaux_tc arithmetic in fp32 torch: sigmoid scores, the top 16 of scores + bias, the weights renormalized and + scaled.""" + logits = x.float() @ moe.gate.weight.float().t() + scores = torch.sigmoid(logits) + ids = torch.topk(scores + moe.gate.e_score_correction_bias.float(), 16, dim=-1).indices + picked = scores.gather(1, ids) + weights = picked / picked.sum(-1, keepdim=True) * moe.gate.routed_scaling_factor + return ids, weights + + +def _compare(name, got, want, exact, tol=ENGINE_TOL): + err = ls.rel_err(got, want) + same = bool( + torch.equal(got.contiguous().view(torch.int16), want.contiguous().view(torch.int16)) + ) + ok = same if exact else err <= tol + assert torch.isfinite(got.float()).all(), name + assert ok, f"{name}: bitwise {same}, rel err {err:.3e}" + assert R.same_on_ranks(got), f"{name}: ranks differ" + return same, err + + +def check_engines(): + """Every token count 1..8 on layer 0: the latent against the reference (bit for bit from 3 tokens), the engine it + ran on, the replicated output against the reference tail on the same latent, and the routing.""" + layer, moe = BUILT.layers[0], BUILT.moes[0] + for rows in range(1, DM.MAX_TOKENS + 1): + x = _x(rows, rows) + m1_before, m2_before = _engine_epochs() + pending = _wired(layer, moe, x) + m1_after, m2_after = _engine_epochs() + ran = ("k3_moe_m1" if not torch.equal(m1_before, m1_after) else "") + ( + "k3_moe_m2" if not torch.equal(m2_before, m2_after) else "" + ) + want_engine = {1: "k3_moe_m1", 2: "k3_moe_m2"}.get(rows, "") + assert ran == want_engine, (rows, ran, want_engine) + latent, shared_act, ids, weights = _reference(layer, moe, x) + same, err = _compare(f"latent {rows}", pending.latent, latent, exact=rows > 2) + assert torch.equal(pending.act, shared_act), rows + # The replicated tail: the latent RMS on the folded up projection's fp32 output, plus the shared expert. + y = layer.forward(moe, x, None, partial_tail=False) + shared = moe.shared_experts.down_proj(shared_act, layer_idx=moe.shared_experts.layer_idx) + up = DM._gemv(None, "moe_up", latent, moe.routed_expert_up_proj.weight, out_fp32=True) + scale = torch.rsqrt(latent.float().pow(2).mean(-1, keepdim=True) + 1e-5) + y_ref = (up * scale + shared.float()).bfloat16() + _compare(f"output {rows}", y, y_ref, exact=rows > 2) + # The front's routing against the noaux_tc arithmetic in torch. + ref_ids, ref_weights = _routing_ref(moe, x) + order = torch.argsort(ids.long(), dim=-1) + ref_order = torch.argsort(ref_ids, dim=-1) + assert torch.equal(ids.long().gather(1, order), ref_ids.gather(1, ref_order)), ( + rows, + "routing ids", + ) + route_err = ls.rel_err(weights.float().gather(1, order), ref_weights.gather(1, ref_order)) + assert route_err <= ROUTE_TOL, (rows, route_err) + torch.cuda.synchronize() + if R.rank == 0: + print(f"[rank 0] {rows} tokens on {ran or 'k3_moe'}: latent bitwise {same} (rel err {err:.2e}), " + f"routing weights {route_err:.2e}", flush=True) # fmt: skip + + +def check_push(): + """Three layers on one state over steps of SEQUENCE tokens: each layer's engine pushing into the latent exchange, + then k3_latent_reduce, against the same engine's returned partial and the routed experts' all-reduce, bit for + bit.""" + for step, rows in enumerate(SEQUENCE): + for i, (layer, moe) in enumerate(zip(BUILT.layers, BUILT.moes)): + x = _x(100 + 10 * step + i, rows) + ids, weights, x_fp8, x_sf, _ = layer._front(moe, x) + layer._routed(x_fp8, x_sf, ids, weights, 0, push=True) + pushed = DM.k3_latent_reduce(rows, BUILT.state.exchange) + returned = moe.routed_experts.all_reduce(layer._routed(x_fp8, x_sf, ids, weights, 0)) + _compare(f"push step {step} layer {i} ({rows} tokens)", pushed, returned, exact=True) + torch.cuda.synchronize() + if R.rank == 0: + print( + f"[rank 0] push: {len(SEQUENCE)} steps x 3 layers equal to the all-reduce", flush=True + ) + + +def check_graph(): + """Layers 0, 1, 2 at GRAPH_ROWS tokens captured in one CUDA graph pushing (the engines push and k3_latent_reduce + sums, as on a pushing step), replayed three times with rewritten inputs: each replay's latents and shared + activations equal eager steps on the same inputs (the returned partials and the all-reduce), bit for bit.""" + xs = [_x(300 + i, rows) for i, rows in enumerate(GRAPH_ROWS)] + chain = list(zip(BUILT.layers, BUILT.moes, xs)) + for layer, moe, x in chain: # eager: no push + _wired(layer, moe, x) + torch.cuda.synchronize() + R.barrier() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + outs = [_wired(layer, moe, x, push=True) for layer, moe, x in chain] + for replay in range(3): + for i, x in enumerate(xs): + x.copy_(_x(400 + 10 * replay + i, x.shape[0])) + R.barrier() + graph.replay() + torch.cuda.synchronize() + captured = [(p.latent.clone(), p.act.clone()) for p in outs] + R.barrier() + eager = [_wired(layer, moe, x) for layer, moe, x in chain] + torch.cuda.synchronize() + for i, ((latent, act), e) in enumerate(zip(captured, eager)): + _compare(f"graph replay {replay} layer {i}", latent, e.latent, exact=True) + # This rank's shared expert slice: its own activation, not one equal across the ranks. + assert torch.equal(act.view(torch.int16), e.act.view(torch.int16)), (replay, i) + if R.rank == 0: + print("[rank 0] graph: 3 replays (pushing) equal to eager steps", flush=True) + + +def _run_one_rank(args) -> int: + global R, T, DM + R = ls.Rank(args) + import importlib + + T = importlib.import_module( + "tensorrt_llm._torch._experimental.modeling_v2.models.kimi_k3_vl." + "kimi_k3_mxfp4__sm_100__tp16_moetp16ep1.modeling" + ) + DM = T._decode_moe + # First, before anything exists as an inference tensor: the decode MoE as post_load_weights builds it. + code = ls.run_checks(R, [check_build]) + if code: + return code + with torch.inference_mode(): + code = ls.run_checks(R, [check_engines, check_push, check_graph]) + return code + + +if __name__ == "__main__": + ARGS = ls.parse_args(sys.argv[1:]) + if ARGS.rank_worker or ARGS.launcher == "srun": + sys.exit(_run_one_rank(ARGS)) + ls.spawn(__file__, ARGS, DEADLINE_S) + print("OK") diff --git a/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_kimi_k3_route_b_decode_moe_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_kimi_k3_route_b_decode_moe_op_matrix.py new file mode 100644 index 000000000000..a7d1b8df1a09 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_kimi_k3_route_b_decode_moe_op_matrix.py @@ -0,0 +1,32 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Collected entry point for the MoE decode path of the Kimi K3 target ``kimi_k3_mxfp4__sm_100__tp16_moetp16ep1`` +(``decode_moe.py``: every expert on each rank, the routed experts pushing into the latent exchange). + +The checks are ``_kimi_k3_route_b_decode_moe_op_matrix.py`` beside this file, its own W-rank launcher (the layers share +the path's collective state, so one job, not independent cases); see ``_rank_job`` for why that is left intact. +""" + +import _rank_job +import pytest +import torch + +assert torch.cuda.is_available(), "the Kimi K3 target's MoE decode path requires CUDA devices" + +if torch.cuda.get_device_capability() != (10, 0): + # The target is certified on sm_100 (GB200) only, as are the catalog entries it calls. + pytest.skip("the Kimi K3 target runs on sm_100 only", allow_module_level=True) + +if torch.cuda.device_count() < _rank_job.WORLD_SIZE: + # One rank per device: fewer visible devices cannot host the check's world size. + pytest.skip( + f"the check runs {_rank_job.WORLD_SIZE} ranks, one per device; " + f"{torch.cuda.device_count()} visible", + allow_module_level=True, + ) + + +# The case starts its own W-rank mpirun over the visible devices; under xdist several workers would fight for them. +@pytest.mark.no_xdist +def test_kimi_k3_route_b_decode_moe_op_matrix() -> None: + _rank_job.run("kimi_k3_route_b_decode_moe") diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_route_b_moe.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_route_b_moe.py new file mode 100644 index 000000000000..d15c5a3b5c9e --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_route_b_moe.py @@ -0,0 +1,148 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""The MoE decode path of ``tp16_moetp16ep1`` (route B), host-side: which layers and steps it takes, and the engine and +latent all-reduce a step runs. Its kernels' calls are checked at 4 ranks by +``comm/test_modeling_v2_kimi_k3_route_b_decode_moe_op_matrix.py``.""" + +import ast +import types +from pathlib import Path + +import torch + +import tensorrt_llm._torch.cute_dsl_kernels.k3_fused_moe as _k3_fused_moe +from tensorrt_llm._torch._experimental.modeling_v2.models.kimi_k3_vl.kimi_k3_mxfp4__sm_100__tp16_moetp16ep1 import ( # noqa: E501 + decode_moe as route_b_moe, +) + + +def _module_floats(path: Path) -> dict: + """A kernel file's module-level float constants, ``NAME = `` or ``NAME = float(_cfg("key", ))`` + (the default), read without importing it.""" + found = {} + for node in ast.walk(ast.parse(path.read_text())): + if not (isinstance(node, ast.Assign) and len(node.targets) == 1): + continue + name, value = getattr(node.targets[0], "id", None), node.value + if isinstance(value, ast.Constant) and isinstance(value.value, float): + found[name] = value.value + elif ( + isinstance(value, ast.Call) + and getattr(value.func, "id", None) == "float" + and value.args + and isinstance(value.args[0], ast.Call) + and getattr(value.args[0].func, "id", None) == "_cfg" + ): + found[name] = float(value.args[0].args[1].value) + return found + + +def _moe(betas, top_k=16): + """A MoE layer with the fields layout_gaps reads before its shape checks.""" + linear = types.SimpleNamespace(weight=None, bias=None) + return types.SimpleNamespace( + _reduce_routed_output=True, + _situ_betas=betas, + num_experts=896, + top_k=top_k, + moe_hidden_size=3584, + hidden_size=7168, + gate=None, + routed_experts=types.SimpleNamespace(backend=None), + shared_experts=types.SimpleNamespace(gate_up_proj=linear, down_proj=linear), + ) + + +def test_engine_situ_caps_are_the_kernels(): + """ENGINE_SITU_CAPS, against which layout_gaps checks a checkpoint's SiTU caps, are the caps k3_moe_m1, k3_moe_m2 + and k3_moe compile in.""" + kernels = Path(_k3_fused_moe.__file__).resolve().parent + for name in ("k3_moe_m1_kernel.py", "k3_moe_m2_kernel.py", "k3_moe_kernel.py"): + consts = _module_floats(kernels / name) + caps = (consts["SITU_GATE_CAP"], consts["SITU_LINEAR_CAP"]) + assert caps == route_b_moe.ENGINE_SITU_CAPS, (name, caps) + + +def test_other_situ_caps_keep_the_generic_path(): + """A MoE layer whose SiTU caps (the checkpoint's activation_situ_beta / activation_situ_linear_beta) differ from + the engines' gets one gap naming them, so the target leaves it on the generic path. The engines' caps pass that + check: the gap that follows is the next check's (here the top-k).""" + for betas in ((5.0, 25.0), (4.0, 20.0)): + gaps = route_b_moe.layout_gaps(_moe(betas), 16, 8) + assert len(gaps) == 1 and "SiTU caps" in gaps[0] and str(betas) in gaps[0], gaps + gaps = route_b_moe.layout_gaps(_moe(route_b_moe.ENGINE_SITU_CAPS, top_k=8), 16, 8) + assert len(gaps) == 1 and "experts / top-k" in gaps[0], gaps + + +def test_the_path_takes_at_most_8_tokens(): + """Steps of 1..8 tokens, with or without the deferred tail; no wide step (k3_moe's wide build does not fit 896 + local experts), so 9 tokens and more stay on the generic path, as does a step decode_step did not classify.""" + takes = route_b_moe.K3DecodeMoeLayer.takes + step = types.SimpleNamespace(wide=True) + + def rows(n): + return torch.empty(n, 7168, dtype=torch.bfloat16) + + for partial_tail in (False, True): + for n in range(1, route_b_moe.MAX_TOKENS + 1): + assert takes(None, rows(n), step, partial_tail) + for n in (9, 16, 64): + assert not takes(None, rows(n), step, partial_tail) + assert not takes(None, rows(4), None, False) + + +def test_a_pushing_step_runs_the_engines_push_form(monkeypatch): + """The latent all-reduce a MoE layer runs (forward's ``push``, the step's ``DecodeStep.latent_push``): with push + and the state's latent exchange, the engine of the token count pushes (k3_moe_m1 at one token, k3_moe_m2 at two, + k3_moe above) and k3_latent_reduce sums; without push, or without an exchange, the engine returns its partial to + the routed experts' all-reduce.""" + calls = [] + + def returning(name): + return lambda x_fp8, *args: calls.append(name) or torch.zeros(x_fp8.shape[0], 3584) + + def pushing(name): + return lambda *args: calls.append(name) + + for name in ("k3_moe_m1", "k3_moe_m2", "k3_moe"): + monkeypatch.setattr(route_b_moe, name, returning(name)) + monkeypatch.setattr(route_b_moe, f"{name}_push", pushing(f"{name}_push")) + monkeypatch.setattr( + route_b_moe, + "k3_latent_reduce", + lambda rows, exchange: calls.append("k3_latent_reduce") or torch.zeros(rows, 3584), + ) + monkeypatch.setattr( + route_b_moe.K3DecodeMoeLayer, + "_front", + lambda self, moe, x: ( + None, + None, + torch.zeros(x.shape[0], 3584), + None, + torch.zeros(x.shape[0], 384), + ), + ) + moe = types.SimpleNamespace( + routed_experts=types.SimpleNamespace( + backend=types.SimpleNamespace(slot_start=0), + all_reduce=lambda y: calls.append("all_reduce") or y, + ), + routed_expert_norm=types.SimpleNamespace(variance_epsilon=1e-5), + ) + for exchange in (object(), None): + layer = route_b_moe.K3DecodeMoeLayer( + state=types.SimpleNamespace(exchange=exchange), front_weight=None, head_weight=None, + tail_weight=None, tail_pad=None, bias=None, lo=0, width=224, shared_cols=384, + small=None, wide=None, m1=None, m2=None, + ) # fmt: skip + for rows, engine in ((1, "k3_moe_m1"), (2, "k3_moe_m2"), (3, "k3_moe"), (8, "k3_moe")): + x = torch.zeros(rows, 7168, dtype=torch.bfloat16) + for push in (False, True): + calls.clear() + pending = layer.forward(moe, x, None, partial_tail=True, push=push) + assert pending.latent.shape == (rows, 3584) + if push and exchange is not None: + assert calls == [f"{engine}_push", "k3_latent_reduce"], (rows, push, calls) + else: + assert calls == [engine, "all_reduce"], (rows, push, exchange, calls) From 3656bcd5d2880c5ffa38ba59a561452798eaf147 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 11:38:03 -0700 Subject: [PATCH 150/161] [None][perf] Kimi K3 generic MoE path: the fused route + MXFP8 quant on k3_route_quant The TRTLLM-Gen W4A8 MXFP4 MXFP8 backend's fused Kimi K3 route + MXFP8 quant (at most 64 tokens: the generic path's prefill, mixed and wide steps) now runs as trtllm::k3_route_quant, the CuTe DSL form of trtllm::kimi_k3_noaux_tc_mxfp8_quant: the same outputs bit for bit, top-16 order included, in about a third of the time. test_backend_fused_route_quant_matches_kimi_k3_noaux_tc_mxfp8_quant checks the backend's outputs against kimi_k3_noaux_tc_mxfp8_quant at 1, 8, 9 and 64 tokens, bit for bit; it is listed in l0_b200.yml. Signed-off-by: Vasanth Sabavat --- .../trtllm_gen/trtllm_w4a8_mxfp4_mxfp8.py | 9 ++++- .../test_lists/test-db/l0_b200.yml | 1 + .../_torch/moe/test_kimi_k3_moe_gate.py | 33 +++++++++++++++++++ 3 files changed, 42 insertions(+), 1 deletion(-) diff --git a/tensorrt_llm/_torch/moe/fused_moe/trtllm_gen/trtllm_w4a8_mxfp4_mxfp8.py b/tensorrt_llm/_torch/moe/fused_moe/trtllm_gen/trtllm_w4a8_mxfp4_mxfp8.py index 5870093b3ba6..dec8431e5a93 100644 --- a/tensorrt_llm/_torch/moe/fused_moe/trtllm_gen/trtllm_w4a8_mxfp4_mxfp8.py +++ b/tensorrt_llm/_torch/moe/fused_moe/trtllm_gen/trtllm_w4a8_mxfp4_mxfp8.py @@ -61,6 +61,10 @@ def try_fused_route_quant( 896 experts, top-16, hidden 3584 and at most 64 tokens. The checks below mirror its ``TORCH_CHECK``s so a miss declines quietly instead of raising, which keeps every other model and shape on the unfused path. + + It runs as ``trtllm::k3_route_quant``, the CuTe DSL form of + ``trtllm::kimi_k3_noaux_tc_mxfp8_quant``: the same outputs bit for bit + (top-16 order included), in about a third of the time. """ if os.environ.get("TLLM_K3_DISABLE_FUSED_ROUTE_QUANT", "0") == "1" or isinstance( x, MxFp8QuantizedTensor @@ -93,6 +97,9 @@ def try_fused_route_quant( ): return None - return torch.ops.trtllm.kimi_k3_noaux_tc_mxfp8_quant( + # Registers trtllm::k3_route_quant. + from ....cute_dsl_kernels.k3_route_quant import op as _k3_route_quant_op # noqa: F401 + + return torch.ops.trtllm.k3_route_quant( router_logits, bias, x, routing.routed_scaling_factor ) diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index 678d46093dec..eb1e8576c7f2 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -212,6 +212,7 @@ l0_b200: - unittest/_torch/moe/fused_moe/test_deepgemm_fused_expand_quant.py # ------------- MoE: test_moe_backend (by backend) --------------- - unittest/_torch/moe/test_kimi_k3_moe_gate.py::test_fused_route_quant_matches_unfused_chain + - unittest/_torch/moe/test_kimi_k3_moe_gate.py::test_backend_fused_route_quant_matches_kimi_k3_noaux_tc_mxfp8_quant - unittest/_torch/moe/test_moe_backend.py::test_kimi_fused_route_quant_skips_prequantized_input - unittest/_torch/moe/test_moe_backend.py::test_kimi_mxfp8_quantized_tensor_handoff - unittest/_torch/moe/test_megamoe_streaming_load.py diff --git a/tests/unittest/_torch/moe/test_kimi_k3_moe_gate.py b/tests/unittest/_torch/moe/test_kimi_k3_moe_gate.py index 5b8b52bc1d3b..782116fc3c77 100644 --- a/tests/unittest/_torch/moe/test_kimi_k3_moe_gate.py +++ b/tests/unittest/_torch/moe/test_kimi_k3_moe_gate.py @@ -3,6 +3,7 @@ """Parity tests for the production Kimi K3 MoE routing method.""" import dataclasses +import types import pytest import torch @@ -103,3 +104,35 @@ def test_fused_route_quant_matches_unfused_chain(num_tokens): assert torch.equal(scales.view(torch.int16), ref_scales.to(torch.bfloat16).view(torch.int16)) assert torch.equal(quantized.view(torch.uint8), ref_quantized.view(torch.uint8)) assert torch.equal(quant_scales, ref_quant_scales.view(num_tokens, -1)) + + +@pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.get_device_capability() != (10, 0), + reason="the TRTLLM-Gen backend's fused route+quant runs trtllm::k3_route_quant (CuTe DSL), sm_100", +) +@pytest.mark.parametrize("num_tokens", [1, 8, 9, 64]) +def test_backend_fused_route_quant_matches_kimi_k3_noaux_tc_mxfp8_quant(num_tokens): + """The TRTLLM-Gen W4A8 MXFP4 backend's fused Kimi K3 route + MXFP8 quant returns what + kimi_k3_noaux_tc_mxfp8_quant returns for the same inputs, bit for bit (the top-16 order included).""" + from tensorrt_llm._torch.moe.fused_moe.routing import DeepSeekV3MoeRoutingMethod + from tensorrt_llm._torch.moe.fused_moe.trtllm_gen.trtllm_w4a8_mxfp4_mxfp8 import ( + TrtllmTrtllmGenW4a8Mxfp4Mxfp8Impl, + ) + + torch.manual_seed(0x5EED + num_tokens) + scores = torch.randn(num_tokens, 896, dtype=torch.float32, device="cuda") + bias = torch.randn(896, dtype=torch.float32, device="cuda") + hidden_states = torch.randn(num_tokens, 3584, dtype=torch.bfloat16, device="cuda") + routed_scaling_factor = 2.446 + routing = DeepSeekV3MoeRoutingMethod(16, 1, 1, routed_scaling_factor, lambda: bias) + backend = types.SimpleNamespace(routing_method=routing) + + got = TrtllmTrtllmGenW4a8Mxfp4Mxfp8Impl.try_fused_route_quant(backend, hidden_states, scores) + want = torch.ops.trtllm.kimi_k3_noaux_tc_mxfp8_quant( + scores, bias, hidden_states, routed_scaling_factor + ) + + assert got is not None + for name, g, w in zip(("experts", "scales", "quantized", "quant_scales"), got, want): + assert g.dtype == w.dtype and g.shape == w.shape, name + assert torch.equal(g.contiguous().view(torch.uint8), w.contiguous().view(torch.uint8)), name From 74839d10e148485f179c95eaaab34e0e8885f197 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 12:55:09 -0700 Subject: [PATCH 151/161] [None][perf] Kimi K3 targets: route the generic MoE outside the TRTLLM-Gen kernel Both Kimi K3 targets' gates now return KimiK3MoeRoutingMethod, a DeepSeekV3MoeRoutingMethod whose requires_separated_routing is True. A generic TRTLLM-Gen MoE call (prefill, mixed steps, and decode steps the decode path does not take) then routes outside the kernel: the backend's fused route + MXFP8 quantize (trtllm::k3_route_quant) up to 64 tokens, noaux_tc_op above. The kernel takes the top-16 ids and weights. The measured stack (5596533407) routed every generic call this way. Numerics at op level (work/moe8/item9/sep_routing_proof.py): - Setup: one rank of each target, random MXFP4 experts, SiTU 4 / 25. Router logits are random or come from the checkpoint's gates of layers 1, 46 and 92, plus exact-tie and zero-logit rows. Token counts 1-8192, plus 200 batches each of 8 and 64 tokens. - The MoE output is bitwise equal to the in-kernel routing's on all but 27 of 411,480 tokens. Each of those 27 has a top-16 weight within 2^-20 of a bf16 rounding boundary, where the two fp32 normalizations round one bf16 ulp apart. - No difference points at expert choice or order. Every differing token has such a weight, and rows of exact ties (five logit levels, no bias) and of zero logits match bitwise. - The MXFP8 input is bitwise the same, and the separated weights are within half a bf16 ulp of an fp64 reference. Time per generic MoE call (CUDA graph replays, one tp16_moetp16ep1 / tp16_moetp4ep4 rank): - Faster at small counts: -10 to -20 us at 5-128 tokens, -3 us at 512. - Slower from 1024 tokens: +5 / +7 us at 1024, +14 / +13 at 2048, +27 / +30 at 4096 and +64 / +62 at 8192. There noaux_tc_op (90 us at 8192 tokens) costs more than the kernel's own routing. A prefill chunk of 8192 tokens then takes about 5.8 ms more over the 92 MoE layers. - Steps the decode path takes (at most 8 tokens on both targets, and tp16_moetp4ep4's wide decode steps) never make this call. Route B's decode_moe.py docstring now says its front routes with the generic path's arithmetic. test_modeling_v2_kimi_k3_moe_routing.py (l0_b200) checks the gate's method, and one ConfigurableMoE call routed outside the kernel against the same call routed inside it: bitwise on the tokens whose routing cannot round two ways, within a bf16 ulp elsewhere. Signed-off-by: Vasanth Sabavat --- .../decode_moe.py | 4 +- .../modeling.py | 35 ++- .../modeling.py | 24 +- .../test_lists/test-db/l0_b200.yml | 2 + .../test_modeling_v2_kimi_k3_moe_routing.py | 245 ++++++++++++++++++ 5 files changed, 295 insertions(+), 15 deletions(-) create mode 100644 tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_moe_routing.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py index bbdb744ca13e..b113b6ad6f75 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py @@ -7,8 +7,8 @@ * `moe/k3_moe_front`, one kernel: this rank's slice of the MoE head (its latent-down rows and its router rows) as one GEMV, the slices' all-gather over the TP group's `K3MoeHeadWorkspace`, the top-16 routing, the MXFP8 latent, and - the shared experts' gate_up + SiTU. The routing is the noaux_tc arithmetic of `moe/kimi_k3_noaux_tc_mxfp8_quant`; - the generic path's TRTLLM-Gen MoE routes inside its own kernel, so a near-tie can select another expert there; + the shared experts' gate_up + SiTU. The routing is the noaux_tc arithmetic of `moe/kimi_k3_noaux_tc_mxfp8_quant`, + as the generic path's (`KimiK3MoeRoutingMethod`), on the router logits of the front's head GEMV; * the routed experts of all 896 experts over this rank's intermediate slice: `moe/k3_moe_m1` at one token, `moe/k3_moe_m2` at two, `moe/k3_moe` (on a `K3MoeState`) at three to eight; * the latent all-reduce. On a pushing step (`DecodeStep.latent_push`: a pure decode step captured into a CUDA graph diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py index e336b329d484..b5fd68473e90 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/modeling.py @@ -29,9 +29,10 @@ request has one token. The request-aware kernels take it: MLA attention and its KV store. Every other step (prefill, mixed steps, decode steps above those bounds) runs the **generic path**: this target's -text model (`KimiLinearModel` below: decoder layers, attention residuals, the MLA / KDA / MoE runtimes), computed -exactly as the built-in Kimi K3 text model computes it, on stock modules and ops that have no catalog entries yet. -`UNCERTIFIED_GENERIC_CALLS` names them. +text model (`KimiLinearModel` below: decoder layers, attention residuals, the MLA / KDA / MoE runtimes), computed as +the built-in Kimi K3 text model computes it, on stock modules and ops that have no catalog entries yet +(`UNCERTIFIED_GENERIC_CALLS` names them), except that the routed experts' top-16 is computed outside the TRTLLM-Gen +kernel (`KimiK3MoeRoutingMethod`), with the kernel's arithmetic. The text model hands each step's classification to its attention modules, which run a **decode step** on the K3 decode kernels' catalog entries: @@ -56,9 +57,10 @@ then the routed experts as `moe/k3_moe_m1` at one token, `moe/k3_moe_m2` at two and `moe/k3_moe` over all 896 experts up to 8, the latent all-reduce and the row-parallel tail. On a pushing step (`DecodeStep.latent_push`) the engines push their partials into the TP group's `K3LatentExchange` and `comm/k3_latent_reduce` sums them; other steps use the -routed experts' all-reduce, which sums in the same order. Its routing is the front's (the noaux_tc arithmetic), not -the generic path's TRTLLM-Gen routing. Wider steps, and every step of a checkpoint whose SiTU caps differ from the ones -the engines compile in, run the generic path: `k3_moe`'s wide build does not fit 896 local experts. +routed experts' all-reduce, which sums in the same order. Its routing is the front's: the noaux_tc arithmetic, as on +the generic path, on the router logits of the front's head GEMV. Wider steps, and every step of a checkpoint whose +SiTU caps differ from the ones the engines compile in, run the generic path: `k3_moe`'s wide build does not fit 896 +local experts. **What this target asserts rather than adapts**: SM 10.0; the topology above, with the expert split set explicitly; no speculative decoding; the MXFP4 checkpoint's quantization (W4A16_MXFP4 with no per-layer declarations, so the @@ -289,6 +291,21 @@ _KIMI_K3_MLA_MAX_POSITIONS_ENV = "KIMI_K3_MLA_MAX_POSITIONS" +class KimiK3MoeRoutingMethod(DeepSeekV3MoeRoutingMethod): + """DeepSeek-V3 routing, computed outside the TRTLLM-Gen MoE kernel. + + The kernel's own top-16 of 896 experts is the slow part of a generic MoE call of a few tokens. Outside it, the + scheduler routes with the backend's fused route + MXFP8 quantize (``trtllm::k3_route_quant``) up to 64 tokens and + with ``noaux_tc_op`` above, and the kernel takes the expert ids and weights. Both compute the kernel's arithmetic + and break ties the same way; a bf16 weight can differ by one ulp where it lies within fp32 rounding of a bf16 + rounding boundary. From about a thousand tokens on, ``noaux_tc_op`` takes longer than the kernel's routing. + """ + + @property + def requires_separated_routing(self) -> bool: + return True + + class KimiK3MoEGate(nn.Module): """Kimi K3 gate weights and routing method for ``ConfigurableMoE``.""" @@ -341,15 +358,15 @@ def compute_logits(self, hidden_states: torch.Tensor) -> torch.Tensor: ) @property - def routing_method(self) -> DeepSeekV3MoeRoutingMethod: - """Return the shared DeepSeek-V3 router used by ``ConfigurableMoE``.""" + def routing_method(self) -> KimiK3MoeRoutingMethod: + """Return the DeepSeek-V3 router used by ``ConfigurableMoE``, computed outside the kernel.""" if self.moe_router_activation_func != "sigmoid": raise ValueError("Kimi K3 ConfigurableMoE routing requires sigmoid scores.") if not self.moe_renormalize: raise ValueError( "Kimi K3 ConfigurableMoE routing requires top-k weight renormalization." ) - return DeepSeekV3MoeRoutingMethod( + return KimiK3MoeRoutingMethod( top_k=self.top_k, n_group=self.num_expert_group, topk_group=self.topk_group, diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py index 40710ba4666b..6edf941ebdba 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/modeling.py @@ -28,7 +28,8 @@ text model (`KimiLinearModel` below: decoder layers, attention residuals, the MLA / KDA / MoE runtimes), computed as the built-in Kimi K3 text model computes it, on stock modules and ops that have no catalog entries yet. One weight differs: where the MoE decode path takes a layer, its latent norm's weight is folded into the latent up projection -(the same function, rounded differently). `UNCERTIFIED_GENERIC_CALLS` names the stock code. +(the same function, rounded differently). And the routed experts' top-16 is computed outside the TRTLLM-Gen kernel +(`KimiK3MoeRoutingMethod`), with the kernel's arithmetic. `UNCERTIFIED_GENERIC_CALLS` names the stock code. The text model hands each step's classification to its attention modules, which run a **decode step** on the K3 decode kernels' catalog entries: @@ -284,6 +285,21 @@ _KIMI_K3_MLA_MAX_POSITIONS_ENV = "KIMI_K3_MLA_MAX_POSITIONS" +class KimiK3MoeRoutingMethod(DeepSeekV3MoeRoutingMethod): + """DeepSeek-V3 routing, computed outside the TRTLLM-Gen MoE kernel. + + The kernel's own top-16 of 896 experts is the slow part of a generic MoE call of a few tokens. Outside it, the + scheduler routes with the backend's fused route + MXFP8 quantize (``trtllm::k3_route_quant``) up to 64 tokens and + with ``noaux_tc_op`` above, and the kernel takes the expert ids and weights. Both compute the kernel's arithmetic + and break ties the same way; a bf16 weight can differ by one ulp where it lies within fp32 rounding of a bf16 + rounding boundary. From about a thousand tokens on, ``noaux_tc_op`` takes longer than the kernel's routing. + """ + + @property + def requires_separated_routing(self) -> bool: + return True + + class KimiK3MoEGate(nn.Module): """Kimi K3 gate weights and routing method for ``ConfigurableMoE``.""" @@ -336,15 +352,15 @@ def compute_logits(self, hidden_states: torch.Tensor) -> torch.Tensor: ) @property - def routing_method(self) -> DeepSeekV3MoeRoutingMethod: - """Return the shared DeepSeek-V3 router used by ``ConfigurableMoE``.""" + def routing_method(self) -> KimiK3MoeRoutingMethod: + """Return the DeepSeek-V3 router used by ``ConfigurableMoE``, computed outside the kernel.""" if self.moe_router_activation_func != "sigmoid": raise ValueError("Kimi K3 ConfigurableMoE routing requires sigmoid scores.") if not self.moe_renormalize: raise ValueError( "Kimi K3 ConfigurableMoE routing requires top-k weight renormalization." ) - return DeepSeekV3MoeRoutingMethod( + return KimiK3MoeRoutingMethod( top_k=self.top_k, n_group=self.num_expert_group, topk_group=self.topk_group, diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index eb1e8576c7f2..1e58e12c3895 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -172,6 +172,8 @@ l0_b200: - unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drift.py # tp16_moetp16ep1's MoE decode path: the layers and steps it takes (its kernels: l0_gb200_multi_gpus.yml). - unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_route_b_moe.py + # The Kimi K3 targets' generic-path MoE routing, outside the TRTLLM-Gen kernel, against the kernel's own. + - unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_moe_routing.py # KDA runtime: host-derived prefill metadata and the bf16 state pool # round-trip. Both are single-device cases that use GPU 0 only. - unittest/_torch/modules/kimi_kda/test_kda_host_metadata.py diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_moe_routing.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_moe_routing.py new file mode 100644 index 000000000000..310536526654 --- /dev/null +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_moe_routing.py @@ -0,0 +1,245 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""The Kimi K3 targets' routed-expert routing on the generic path: the gate's ``KimiK3MoeRoutingMethod`` computes the +top-16 outside the TRTLLM-Gen kernel. + +Host-side, the method each target's gate hands ``ConfigurableMoE``. On SM 10.0, one generic MoE call as +``KimiK3MoERuntime`` builds it (W4A8_MXFP4_MXFP8 on TRTLLM-Gen, SiTU), routed outside the kernel, against the same +call on the same experts routed inside it.""" + +import types + +import pytest +import torch + +from tensorrt_llm._torch._experimental.modeling_v2.models.kimi_k3_vl.kimi_k3_mxfp4__sm_100__tp16_moetp4ep4 import ( # noqa: E501 + modeling as route_a, +) +from tensorrt_llm._torch._experimental.modeling_v2.models.kimi_k3_vl.kimi_k3_mxfp4__sm_100__tp16_moetp16ep1 import ( # noqa: E501 + modeling as route_b, +) +from tensorrt_llm._torch.moe.fused_moe.routing import DeepSeekV3MoeRoutingMethod, RoutingMethodType + +TARGETS = pytest.mark.parametrize( + "target", [route_a, route_b], ids=["tp16_moetp4ep4", "tp16_moetp16ep1"] +) + +NUM_EXPERTS, TOP_K, HIDDEN, LATENT = 896, 16, 7168, 3584 +# One tp16_moetp16ep1 rank's expert width: 192 of 3072, zero-padded to 256 by the loader. +INTERMEDIATE = 192 +SITU_CAPS = (4.0, 25.0) + + +def _gate_config(routed_scaling_factor=2.5): + """The checkpoint config fields KimiK3MoEGate reads.""" + return types.SimpleNamespace( + num_experts_per_token=TOP_K, + num_experts=NUM_EXPERTS, + routed_scaling_factor=routed_scaling_factor, + moe_router_activation_func="sigmoid", + num_expert_group=1, + topk_group=1, + moe_renormalize=True, + hidden_size=HIDDEN, + ) + + +@TARGETS +def test_gate_routes_outside_the_kernel(target): + """The gate's method routes outside the kernel and is a DeepSeek-V3 method to everything that dispatches on one: + the backend's fused route + quant, the routing arguments the kernel receives and the routing type it is told.""" + gate = target.KimiK3MoEGate(_gate_config()) + method = gate.routing_method + assert type(method) is target.KimiK3MoeRoutingMethod + assert isinstance(method, DeepSeekV3MoeRoutingMethod) + assert method.requires_separated_routing + assert method.routing_method_type == RoutingMethodType.DeepSeekV3 + impl = method.routing_impl + assert (impl.top_k, impl.n_group, impl.topk_group, impl.routed_scaling_factor) == ( + TOP_K, + 1, + 1, + 2.5, + ) + assert impl.is_fused + assert method.e_score_correction_bias is gate.e_score_correction_bias + + +def _sm_100(): + return torch.cuda.is_available() and torch.cuda.get_device_capability() == (10, 0) + + +def _gate(target): + """The target's gate, random bf16 weights and fp32 correction bias.""" + g = torch.Generator(device="cuda").manual_seed(71) + gate = target.KimiK3MoEGate(_gate_config(), logits_gemm_dtype=torch.bfloat16, device="cuda") + with torch.no_grad(): + gate.weight.copy_(0.02 * torch.randn(gate.weight.shape, generator=g, device="cuda")) + gate.e_score_correction_bias.copy_( + 0.05 * torch.randn(NUM_EXPERTS, generator=g, device="cuda") + ) + return gate + + +def _routed_experts(routing_method): + """The routed experts of one tp16_moetp16ep1 rank as KimiK3MoERuntime builds them (create_moe: ConfigurableMoE, + W4A8_MXFP4_MXFP8 on TRTLLM-Gen, SiTU), every expert loaded with the same random packed MXFP4 checkpoint slice + through the per-expert loader the target's weight load uses.""" + from transformers.configuration_utils import PretrainedConfig + + from tensorrt_llm._torch.model_config import ModelConfig + from tensorrt_llm._torch.moe.fused_moe import ConfigurableMoE, SiTuActivation, create_moe + from tensorrt_llm.mapping import Mapping + from tensorrt_llm.models.modeling_utils import QuantAlgo, QuantConfig + + pretrained = PretrainedConfig() + pretrained.num_experts = NUM_EXPERTS + pretrained.hidden_size = LATENT + pretrained.intermediate_size = INTERMEDIATE + pretrained.torch_dtype = torch.bfloat16 + pretrained.activation_situ_beta, pretrained.activation_situ_linear_beta = SITU_CAPS + moe = create_moe( + routing_method=routing_method, + num_experts=NUM_EXPERTS, + hidden_size=LATENT, + intermediate_size=INTERMEDIATE, + dtype=torch.bfloat16, + reduce_results=False, + model_config=ModelConfig( + pretrained_config=pretrained, mapping=Mapping(), moe_backend="TRTLLM" + ), + override_quant_config=QuantConfig(quant_algo=QuantAlgo.W4A8_MXFP4_MXFP8), + layer_idx=0, + communication_method=None, + activation=SiTuActivation(gate_softcap=SITU_CAPS[0], linear_softcap=SITU_CAPS[1]), + ).cuda() + assert isinstance(moe, ConfigurableMoE) + backend = moe.backend + assert type(backend).__name__ == "TrtllmTrtllmGenW4a8Mxfp4Mxfp8Impl", type(backend).__name__ + + # The checkpoint's TP16 rank-0 slice (192 rows of 3072), as the loader slices and pads it. + loader = types.SimpleNamespace( + expert_size_per_partition=backend.expert_size_per_partition, + initial_local_expert_ids=backend.initial_local_expert_ids, + scaling_vector_size=backend.scaling_vector_size, + intermediate_size=INTERMEDIATE * 16, + intermediate_size_per_partition=INTERMEDIATE, + tp_size=16, + tp_rank=0, + w3_w1_weight=backend.w3_w1_weight, + w2_weight=backend.w2_weight, + w3_w1_weight_scale=backend.w3_w1_weight_scale, + w2_weight_scale=backend.w2_weight_scale, + ) + g = torch.Generator(device="cuda").manual_seed(101) + + def packed(*shape): + return torch.randint(0, 256, shape, generator=g, dtype=torch.uint8, device="cuda") + + def scales(*shape): + return torch.randint(118, 124, shape, generator=g, dtype=torch.uint8, device="cuda") + + for e in range(NUM_EXPERTS): + backend.quant_method.load_packed_mxfp4_expert( + loader, + global_expert_id=e, + local_slot_id=e, + w1_weight=packed(INTERMEDIATE, LATENT // 2), + w1_weight_scale=scales(INTERMEDIATE, LATENT // 32), + w2_weight=packed(LATENT, INTERMEDIATE // 2), + w2_weight_scale=scales(LATENT, INTERMEDIATE // 32), + w3_weight=packed(INTERMEDIATE, LATENT // 2), + w3_weight_scale=scales(INTERMEDIATE, LATENT // 32), + ) + backend._weights_transformed = False + moe.post_load_weights() + return moe + + +def _clear_rows(logits, bias, routed_scaling_factor): + """Tokens whose routing fp32 arithmetic cannot round two ways: no two of the top 17 biased scores within 1e-6 of + each other (the top-16 and its order are the same however fp32 rounds), and no top-16 weight within 2^-20 of a + bf16 rounding boundary (its bf16 value is the same however the fp32 normalization rounds). Computed in fp64.""" + scores = torch.sigmoid(logits.double()) + top = torch.topk(scores + bias.double(), TOP_K + 1, dim=1) + no_tie = (top.values[:, :-1] - top.values[:, 1:]).min(dim=1).values > 1e-6 + w = scores.gather(1, top.indices[:, :TOP_K]) + w = w / w.sum(dim=1, keepdim=True) * routed_scaling_factor + ulp = torch.pow(2.0, torch.floor(torch.log2(w)) - 7) + frac = w / ulp - torch.floor(w / ulp) + no_boundary = ((frac - 0.5).abs() * ulp / w >= 2.0**-20).all(dim=1) + return no_tie & no_boundary + + +@pytest.mark.skipif( + not _sm_100(), reason="the TRTLLM-Gen MXFP4 cubins and k3_route_quant run on sm_100" +) +@TARGETS +def test_routing_outside_the_kernel_matches_inside(target, monkeypatch): + """One generic MoE call routed outside the kernel (the gate's method: the backend's fused route + MXFP8 quant up to + 64 tokens, noaux_tc_op above) equals the same call on the same experts routed inside it + (DeepSeekV3MoeRoutingMethod), bit for bit, on every token whose routing cannot round two ways (``_clear_rows``); + the others within a bf16 ulp of the row's largest value.""" + from tensorrt_llm._torch.moe.fused_moe.trtllm_gen.trtllm_w4a8_mxfp4_mxfp8 import ( + TrtllmTrtllmGenW4a8Mxfp4Mxfp8Impl, + ) + + gate = _gate(target) + outside = gate.routing_method + inside = DeepSeekV3MoeRoutingMethod( + TOP_K, + 1, + 1, + outside.routing_impl.routed_scaling_factor, + lambda: gate.e_score_correction_bias, + ) + moe_outside = _routed_experts(outside) + moe_inside = _routed_experts(inside) + + fused, applied = [], [] + fused_route_quant = TrtllmTrtllmGenW4a8Mxfp4Mxfp8Impl.try_fused_route_quant + apply = target.KimiK3MoeRoutingMethod.apply + + def spy_fused(self, x, router_logits): + out = fused_route_quant(self, x, router_logits) + fused.append((self.routing_method is outside, x.shape[0], out is not None)) + return out + + def spy_apply(self, router_logits, input_ids=None): + applied.append(router_logits.shape[0]) + return apply(self, router_logits, input_ids) + + monkeypatch.setattr(TrtllmTrtllmGenW4a8Mxfp4Mxfp8Impl, "try_fused_route_quant", spy_fused) + monkeypatch.setattr(target.KimiK3MoeRoutingMethod, "apply", spy_apply) + + g = torch.Generator(device="cuda").manual_seed(5) + with torch.inference_mode(): + for num_tokens in (1, 9, 64, 65, 300): + hidden = torch.randn(num_tokens, HIDDEN, generator=g, device="cuda").to(torch.bfloat16) + x = torch.randn(num_tokens, LATENT, generator=g, device="cuda").to(torch.bfloat16) + logits = gate.compute_logits(hidden) + fused.clear() + applied.clear() + out = moe_outside(x, logits) + # The scheduler routed outside the kernel: the fused route + quant up to 64 tokens, the method above. + assert fused == [(True, num_tokens, num_tokens <= 64)], fused + assert applied == ([] if num_tokens <= 64 else [num_tokens]), applied + fused.clear() + ref = moe_inside(x, logits) + assert fused == [], fused + assert torch.isfinite(ref.float()).all() + + clear = _clear_rows( + logits, gate.e_score_correction_bias, outside.routing_impl.routed_scaling_factor + ) + assert clear.float().mean() > 0.9, clear.float().mean() + assert torch.equal(out[clear].view(torch.int16), ref[clear].view(torch.int16)), ( + num_tokens + ) + rest = ~clear + if rest.any(): + err = (out[rest].float() - ref[rest].float()).abs().amax(dim=1) + assert (err <= ref[rest].float().abs().amax(dim=1) * 2.0**-8).all(), ( + num_tokens, + err, + ) From 4657c202af5822e137f808bdfe29be377f87c31f Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 14:06:29 -0700 Subject: [PATCH 152/161] [None][test] Kimi K3 decode-comm matrix: route A's shared-expert slice at every world The decode MoE checks built the shared expert with 384 x 4 intermediate columns, so each rank held a TP16 rank's 384-column slice only at world 4. At world 16 each rank held 96 columns. k3_moe_front's weight layout refuses that (it needs a multiple of 64 shared columns), so check_decode_moe_from_model_parameters failed and the job stopped before any decode MoE check ran. Route A's model holds 384 columns per rank at TP16 (two shared experts of 3072 over 16 ranks). The shared expert now has 384 columns per rank at every world. Nothing changes at world 4. Signed-off-by: Vasanth Sabavat --- .../modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py index 56c97eb41bfe..5d5817e80772 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py @@ -653,9 +653,8 @@ def load(): # The MoE decode path at this check's TP: a moetp4ep4 rank's routed experts (intermediate columns, experts), and a -# shared expert whose rank slice at TP 4 is a TP16 rank's (384 activation columns). +# shared expert whose rank slice is a TP16 rank's (TAIL_ACT = 384 activation columns) at every world. I_TP, E_LOCAL, NUM_EXPERTS = 768, 224, 896 -SHARED = 384 * 4 SITU_CAPS = (4.0, 25.0) MOE_TOL = 3e-2 # the front's fused shared gate_up + SiTU against the shared expert's own GEMMs @@ -689,7 +688,7 @@ def _decode_moe_layer(): norm = RMSNorm(hidden_size=LATENT, eps=1e-5, dtype=torch.bfloat16) shared = GatedMLP( hidden_size=H, - intermediate_size=SHARED, + intermediate_size=TAIL_ACT * R.world, bias=False, activation=SituAndMul( beta=SITU_CAPS[0], linear_beta=SITU_CAPS[1], use_fused_activation=True From 0da8a0267c16ec4d089c6faf21bd3ea1572dcb2c Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 15:47:04 -0700 Subject: [PATCH 153/161] [None][fix] Kimi K3 decode MoE warm-up: no back-to-back k3_moe calls on one layer The decode MoE warm-ups ran k3_moe and its push build on the same layer back to back. k3_moe claims its first tile from the layer's counters before its grid-dependency wait. With programmatic dependent launch, the push build could therefore take tiles from the plain call's queue. The call that lost its tiles, or the next call on that layer, then hung, and the latent reduce waited for the missing push on every rank. The first warm-up in a process compiles the push build between the two calls, which hides the race. A second warm-up, with every build already compiled, hung in 3 of 4 runs at 4 ranks. Repeating the warm-up hung within 10 to 20 calls. Both targets' warm-ups now synchronize the device before the push build. With the sync, 100 warm-ups in a row and the whole decode-comm push check pass. Pushing on a second handle of the layer, which has its own counters, also passes, so the shared state is the counters. The k3_moe catalog entry now states the rule. Signed-off-by: Vasanth Sabavat --- .../_experimental/modeling_v2/catalog/moe/k3_moe.py | 9 +++++++-- .../kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py | 3 +++ .../kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_moe.py | 3 +++ 3 files changed, 13 insertions(+), 2 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.py index ddd44bed1a4d..909a30f8fd10 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/k3_moe.py @@ -52,7 +52,11 @@ def k3_moe( acquires the front's outputs through the workspace's ready words and advances its epoch, so the front call before it must have published them (``moe/k3_moe_front`` with ``publish=True``). ``out``: bf16, contiguous, at least ``[M, 3584]``; the call writes its first M rows and returns an empty ``[0, 3584]`` tensor. Writes the state's - slab (left armed) and partial rows, and the layer's counters (left zero).""" + slab (left armed) and partial rows, and the layer's counters (left zero). + + Two calls on one layer, in either form, must not run back to back: a call claims its first tile from the layer's + counters before its grid-dependency wait. Between the two, some kernel must wait for its predecessor before it + triggers its dependents (``griddepcontrol.wait``, then ``launch_dependents``), or the stream is synchronized.""" state = layer.state if (head is not None) != state.head_flags: raise ValueError( @@ -80,7 +84,8 @@ def k3_moe_push( (default ``exchange.rank``) of every rank's ``exchange`` (a TP group's ``K3LatentExchange``) instead of returned. One ``comm/k3_latent_reduce`` of the M tokens on that exchange must follow before the next push, on every rank in the same order. ``head`` as in :func:`k3_moe`. Writes the state's slab (left armed) and partial rows, the layer's - counters (left zero), and every rank's exchange.""" + counters (left zero), and every rank's exchange. As for :func:`k3_moe`, two calls on one layer must not run back to + back.""" state = layer.state if (head is not None) != state.head_flags: raise ValueError( diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py index b113b6ad6f75..d2014be7c85d 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1/decode_moe.py @@ -341,6 +341,9 @@ def warm_up(self, moe: nn.Module) -> None: args = [t.expand(rows, -1).contiguous() for t in (x_fp8, x_sf, ids, weights)] self._routed(*args, offset) if exchange is not None: + # Two calls on one layer must not run back to back (moe/k3_moe): the push build starts once the + # call above has ended. + torch.cuda.synchronize(device) self._routed(*args, offset, push=True) k3_latent_reduce(rows, exchange) # <<< route B diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_moe.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_moe.py index a2fa6889283b..89e2c571e935 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_moe.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp4ep4/decode_moe.py @@ -275,6 +275,9 @@ def warm_up(self, moe: nn.Module) -> None: k3_moe(x_fp8, x_sf, ids, weights, offset, self.small) exchange = self.state.exchange if exchange is not None: + # Two calls on one layer must not run back to back (moe/k3_moe): the push build starts once the call + # above has ended. + torch.cuda.synchronize(device) k3_moe_push(x_fp8, x_sf, ids, weights, offset, self.small, exchange) k3_latent_reduce(1, exchange) logits = torch.zeros(1, moe.num_experts, dtype=torch.float32, device=device) From 459e4462b583060f6e37a04c25ca8ef2ac95a3aa Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 17:23:04 -0700 Subject: [PATCH 154/161] [None][doc] Kimi K3 sandwich: polled input and fold store before the final grid wait k3_sandwich's polled input (src_slab, x_src 1) and folded latent all-reduce (lat_uc, x_src 2) take their inputs without a grid-dependency wait and store the op's outputs before the kernel's final wait. Under programmatic dependent launch a predecessor may still be running then, and an output block the caching allocator recycled after that predecessor's launch can still be in its reads. No model call and no catalog wrapper reaches either form; only test_k3_sandwich.py's check_fold_wrap does, and it synchronizes before every call. State the caller's obligation in the op's module docstring and in the two catalog entries that list the options as inert: a stream synchronization before the call, or TRTLLM_ENABLE_PDL=0. Documentation only; no code path changes. Signed-off-by: Vasanth Sabavat --- .../modeling_v2/catalog/comm/k3_sandwich_oproj.md | 4 +++- .../modeling_v2/catalog/comm/k3_sandwich_tail.md | 4 +++- tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py | 6 ++++++ 3 files changed, 12 insertions(+), 2 deletions(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md index 8c186a7916b0..f525b4177c21 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_oproj.md @@ -68,7 +68,9 @@ take `S` 0-8. The inputs are read only. Inert (not exposed by the wrapper): `x_slab=None, slab_buf=0, src_slab=None, src_buf=0` — the op can also publish `normed` into a Lamport slab the next kernel polls, and poll `core` from its producer's slab. Each slab is cross-call state of its own (three sentinel-armed buffers the caller rotates by the call's ordinal), not part of the workspace; -see *Notes*. +see *Notes*. Polling stores the op's outputs before its final grid-dependency wait: a caller passing `src_slab` must +make sure no predecessor still in flight reads memory those outputs may occupy (a stream synchronization before the +call, or `TRTLLM_ENABLE_PDL=0`). ## State diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md index e1a703e7d932..dce43aee8774 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/k3_sandwich_tail.md @@ -101,7 +101,9 @@ def k3_sandwich_tail( Inert (not exposed by the wrapper): `x_slab`, `slab_buf`, `src_slab`, `src_buf` — publishing `normed` into a Lamport slab and polling the reduced latent from its producer's slab, cross-call state of their own as for `comm/k3_sandwich_oproj`; and `lat_uc`, `lat_flags` — the latent all-reduce folded into this op over a second state -object, a `K3SandwichLatentExchange` (see *Notes*). +object, a `K3SandwichLatentExchange` (see *Notes*). Polling and the fold store the op's outputs before its final +grid-dependency wait: a caller passing `src_slab` or `lat_uc` must make sure no predecessor still in flight reads +memory those outputs may occupy (a stream synchronization before the call, or `TRTLLM_ENABLE_PDL=0`). ## State diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py index a467d0db8c16..684d328f4044 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_sandwich/op.py @@ -37,6 +37,12 @@ ``k3_sandwich_oproj``, the reduced latent for ``k3_sandwich_tail``) publishes it as such a slab (int32 [3][8][cols / 2], sentinel 0xFFFFFFFF) and launches its dependents only after its own grid wait, the kernel polls buffer ``src_buf`` of it instead of waiting for the producer's grid; ``core`` / ``latent`` then give only the shape. + +The polled input and the folded latent all-reduce store their outputs before the kernel's final grid-dependency +wait. A caller of either form must make sure that no predecessor still in flight reads memory those outputs may +occupy, such as a block the caching allocator recycled after that predecessor's launch: synchronize the stream before +the call (as ``test_k3_sandwich.py``'s ``check_fold_wrap`` does), or run without programmatic dependent launch +(``TRTLLM_ENABLE_PDL=0``). """ from __future__ import annotations From c48e85e738507821bd10c8d7da8b82c32fa754fc Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 18:43:34 -0700 Subject: [PATCH 155/161] [None][test] Kimi K3 decode-comm matrix: free a graph's outputs before the next capture check_graph_capture_and_replay captures one graph per case (decode8, then unclassified12) and binds each graph's outputs to outs inside its capture. The second capture's assignment therefore freed the first graph's outputs while capturing. Under cudaMallocAsync that free is recorded in the second graph, a free of memory the graph does not own, and its replay fails with "CUDA error: invalid argument": decode8 replays bit for bit, unclassified12 fails. Delete the graph and its outputs at the end of each case, outside any capture, as the file's other capture loop does. The stand-ins' records rebound inside a capture free eager tensors the graph does not use. PyTorch defers those frees to the capture's end ("freeAsync() was called on an uncaptured allocation during graph capture" is a warning only), so they are left as they are. Test-only: with the caching allocator, as in CI, the check passed before and after. Signed-off-by: Vasanth Sabavat --- .../_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py index 5d5817e80772..c5e8df7fc299 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_kimi_k3_decode_comm_op_matrix.py @@ -472,7 +472,9 @@ def check_graph_capture_and_replay(): ) if R.rank == 0: print(f"[rank 0] graph {name}: 3 replays == eager bit for bit", flush=True) - del graph + # Freed here, outside any capture: rebinding outs inside the next case's capture would free this graph's + # outputs there, a captured free of memory the new graph does not own. + del graph, outs def _pending(seed, tokens): From 2f25c1d6ca8013b6d6aa38235bb156bd30ea04ff Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 18:34:14 -0700 Subject: [PATCH 156/161] [None][fix] DSpark K3 Markov path: release the step's acceptance and chain outputs DSparkWorker kept _k3_acceptance (the step's accepted tokens, spec metadata and attention metadata) and _k3_markov (k3_markov's corrected logits, tokens and next_new_tokens) after the step; only _k3_markov_next was cleared. When the last Python-run step of an executor is a CUDA-graph capture, the two kept that capture's graph-pool blocks and graph attention metadata alive through teardown and empty_cache(). Clear both in _prepare_next_new_tokens, where the step ends, beside _k3_markov_next. test_k3_markov_step_state_is_released_after_next_new_tokens (kernel drafts and base-sampler drafts) fails before the change and passes after it. Signed-off-by: Vasanth Sabavat --- tensorrt_llm/_torch/speculative/dspark.py | 4 ++ .../hw_agnostic/test_k3_decode_worker.py | 65 ++++++++++++++++++- 2 files changed, 68 insertions(+), 1 deletion(-) diff --git a/tensorrt_llm/_torch/speculative/dspark.py b/tensorrt_llm/_torch/speculative/dspark.py index 7422b03471b0..ae81f4223019 100644 --- a/tensorrt_llm/_torch/speculative/dspark.py +++ b/tensorrt_llm/_torch/speculative/dspark.py @@ -1186,7 +1186,11 @@ def _prepare_next_new_tokens( """``k3_markov``'s next_new_tokens when the drafts are its tokens for the whole batch (no context requests); otherwise the base assembly.""" chain = getattr(self, "_k3_markov_next", None) + # The step ends here: release its acceptance (which holds the step's attention and spec metadata) and the + # kernel's outputs, so they do not keep a captured graph's pool or metadata alive after the step. self._k3_markov_next = None + self._k3_acceptance = None + self._k3_markov = None if ( chain is not None and chain[0] is next_draft_tokens diff --git a/tests/unittest/_torch/speculative/hw_agnostic/test_k3_decode_worker.py b/tests/unittest/_torch/speculative/hw_agnostic/test_k3_decode_worker.py index ad9a82b6ed89..05db5fd0af8d 100644 --- a/tests/unittest/_torch/speculative/hw_agnostic/test_k3_decode_worker.py +++ b/tests/unittest/_torch/speculative/hw_agnostic/test_k3_decode_worker.py @@ -19,9 +19,12 @@ * Each kernel's predicate takes an eligible step and declines a step that fails one of its conditions: ``trtllm::k3_spec_accept`` (``_k3_accept_applies``) and its vocabulary-sharded target logits (``target_logits``), ``trtllm::k3_ctx_kv`` (``_k3_ctx_kv_applies``) and ``trtllm::k3_markov`` (``_keep_draft_logits_sharded``). -* DSpark's chain: sharded block logits go to ``k3_markov``, whose tokens and next_new_tokens are the step's. +* DSpark's chain: sharded block logits go to ``k3_markov``, whose tokens and next_new_tokens are the step's; after + next_new_tokens the worker holds nothing of the step. """ +import gc +import weakref from types import SimpleNamespace import pytest @@ -740,6 +743,66 @@ def test_k3_markov_chain_drafts_and_next_new_tokens(monkeypatch, pending): assert next_new is outputs[2] +class _StepMetadata: + """A step's attention metadata stand-in: a plain object, so a weak reference tells when it is released.""" + + +@pytest.mark.parametrize( + "kernel_drafts", [True, False], ids=["kernel drafts", "base sampler drafts"] +) +def test_k3_markov_step_state_is_released_after_next_new_tokens(monkeypatch, kernel_drafts): + """Once ``_prepare_next_new_tokens`` has run, the worker holds nothing of the step: neither its acceptance, + which carries the step's attention and spec metadata, nor ``k3_markov``'s outputs. Held, they would keep a + captured step's graph pool and metadata alive after it.""" + worker = _worker(DSparkWorker, k3_decode=True, mapping=TP4) + accepted = torch.zeros(NUM_GENS, K + 1, dtype=torch.int32) + num_accepted = torch.ones(NUM_GENS, dtype=torch.int32) + spec = SimpleNamespace( + batch_indices_cuda=torch.arange(4, dtype=torch.int32), + wants_advanced_draft_sampling=not kernel_drafts, + ) + attn = _StepMetadata() + attn.num_contexts, attn.kv_lens_cuda = 0, None + outputs = ( + torch.zeros(NUM_GENS, K, SHARD), + torch.zeros(NUM_GENS, K, dtype=torch.int32), + torch.zeros(NUM_GENS, K + 1, dtype=torch.int32), + ) + sampled = torch.zeros(NUM_GENS, K, dtype=torch.int32) + assembled = torch.zeros(NUM_GENS, K + 1, dtype=torch.int32) + monkeypatch.setattr(markov_op, "markov_chain", lambda *args, **kwargs: outputs) + monkeypatch.setattr( + SpecWorkerBase, "sample_draft_tokens", lambda self, *args, **kwargs: sampled + ) + monkeypatch.setattr(SpecWorkerBase, "_prepare_next_new_tokens", lambda self, *args: assembled) + vocab_slice = slice(RANK_IN_TP * SHARD, (RANK_IN_TP + 1) * SHARD) + + worker._on_acceptance(accepted, num_accepted, attn, spec) + corrected = worker._k3_markov_chain( + _markov_drafter(), + torch.zeros(NUM_GENS, K, SHARD, dtype=torch.bfloat16), + torch.zeros(NUM_GENS, dtype=torch.long), + vocab_slice, + ) + drafts = worker.sample_draft_tokens(corrected, spec, NUM_GENS, num_contexts=0) + next_new = worker._prepare_next_new_tokens( + accepted, drafts, spec.batch_indices_cuda, NUM_GENS, num_accepted + ) + assert drafts is (outputs[1] if kernel_drafts else sampled) + assert next_new is (outputs[2] if kernel_drafts else assembled) + + held = [ + name + for name in ("_k3_acceptance", "_k3_markov", "_k3_markov_next") + if getattr(worker, name, None) is not None + ] + assert not held, f"the worker still holds {held} after the step" + step_metadata = weakref.ref(attn) + del attn + gc.collect() + assert step_metadata() is None, "the step's attention metadata outlives the step" + + def test_dspark_drafts_from_the_base_sampler_without_the_kernel(monkeypatch): sampled = torch.zeros(NUM_GENS, K, dtype=torch.int32) assembled = torch.zeros(NUM_GENS, K + 1, dtype=torch.int32) From 3411d482b2037f53fece9ca29eeaaef25476fc72 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 19:23:45 -0700 Subject: [PATCH 157/161] [None][test] Kimi K3 DSpark drafter test: give the one-rank drafter's o_proj its TP all-reduce test_modeling_v2_kimi_k3_drafter builds the drafter on a one-rank Mapping, where Qwen3DecoderLayer builds o_proj without an all-reduce (it reduces only at tp_size > 1). K3DSparkDrafter keys the fc split and the fused norms on that all-reduce (_k3_tp_all_reduce, _k3_norms_take_comm), so on one rank the drafter declined both and every fused_drafter test failed (38 eager, 2 under CUDA-graph capture). Before use_decode_comm, give each layer's o_proj the one-rank group's TP all-reduce, an identity module: the fused path engages (fc split, both sandwich forms) and the stock forward through o_proj keeps its bits. At TP16 the layers build that all-reduce themselves. Signed-off-by: Vasanth Sabavat --- .../test_modeling_v2_kimi_k3_drafter.py | 24 ++++++++++++++++++- 1 file changed, 23 insertions(+), 1 deletion(-) diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py index 8bc1d83d9569..dbaf003daf05 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py @@ -15,7 +15,8 @@ rows from the first gen request on), not a gathered copy, and give the gather's result bit for bit. On the TP group's collective state (``use_decode_comm``), here a group of one rank whose collectives are torch -stand-ins counted per call (the ops themselves are certified by their multi-GPU op matrices): +stand-ins counted per call (the ops themselves are certified by their multi-GPU op matrices). The layers' o_proj gets +the group's TP all-reduce, the identity on one rank, which a one-rank Qwen3 layer does not build: * The context projection runs on the split ``fc`` (on one rank, the whole weight): up to a decode step's rows through ``comm/mnnvl_fusion_allreduce`` with ``hidden_norm``, more rows through the drafter's TP all-reduce, and matches the @@ -195,6 +196,26 @@ def sandwich(x, *args, swiglu=False): monkeypatch.setattr(decode_comm, "k3_sandwich_plain", sandwich) +class _OneRankTPAllReduce(torch.nn.Module): + """The TP all-reduce of a group of one rank: the identity.""" + + def forward(self, input, all_reduce_params=None): + return input + + def uses_nccl_symmetric_memory_window(self) -> bool: + return False + + +def _one_rank_tp_all_reduce(module): + """Every layer's o_proj reduces over the one-rank group, as it reduces over the TP group on more ranks + (``Qwen3DecoderLayer`` builds o_proj without an all-reduce when tp_size is 1). The drafter's fused path keys on that + all-reduce; the stock forward through o_proj keeps its bits.""" + for layer in module.model.layers: + o_proj = layer.self_attn.o_proj + o_proj.reduce_output = True + o_proj.all_reduce = _OneRankTPAllReduce() + + @pytest.fixture(scope="module") def fused_drafter(): """The drafter on the one-rank collective state: its fc split over one rank, both sandwich forms compiled (their @@ -202,6 +223,7 @@ def fused_drafter(): with pytest.MonkeyPatch.context() as mp: _one_rank_collectives(mp, []) module = _load_drafter() + _one_rank_tp_all_reduce(module) module.use_decode_comm(ONE_RANK_COMM) return module From b3b4e2724f78a92672f8462035ca59a6fa23f5e7 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 20:13:11 -0700 Subject: [PATCH 158/161] [None][test] Kimi K3 k3_ctx_kv: back-to-back calls share the arrival counter Launch seven splits (different grids) twice each on one stream, with no host synchronization between the calls, as consecutive decode steps run. Every call's pool, ctx_len and num_ctx must equal the same call made alone, and the device's arrival counter must be zero once the device is idle. Until now a test reused the counter back to back only once, and no test checked that it returns to zero. Signed-off-by: Vasanth Sabavat --- .../kimi_k3/test_k3_ctx_kv.py | 50 +++++++++++++++++++ 1 file changed, 50 insertions(+) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctx_kv.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctx_kv.py index eb344eee6c8f..7d7e4eaf2006 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctx_kv.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_ctx_kv.py @@ -323,6 +323,56 @@ def test_split_tp4(batch, k1): assert res["ok"], res +BACK_TO_BACK_SPLITS = [(1, 8), (8, 8), (2, 4), (8, 1), (4, 8), (3, 4), (8, 7)] + + +def test_back_to_back_calls_share_the_arrival_counter(): + """Calls of different splits (different grids) launched back to back on one stream, two of each, with no host + synchronization between them, as consecutive decode steps run: they share the device's arrival counter, which + each call's last arrival resets. Every call's pool, ctx_len and num_ctx equal the same call made alone, and the + counter is zero once the device is idle.""" + op = _op() + with torch.inference_mode(): + steps = [] + for i, (batch, k1) in enumerate(BACK_TO_BACK_SPLITS): + gen = torch.Generator(device="cuda").manual_seed(20261003 + 10 * i) + steps.append(Step(gen, 1, batch, k1, seed=i)) + alone = [] + for st in steps: + buf, ctx, num_ctx, _ = st.run_kernel() + torch.cuda.synchronize() + alone.append((buf, ctx, num_ctx)) + # Inputs, pools and pool views first, so that nothing between the calls waits on the device. + calls = [] + for st in steps: + for _ in range(2): + buf, ctx = st.buf.clone(), st.ctx0.clone() + calls.append((st, buf, ctx, pool_view(views_of(buf, st.layers, st.style)))) + torch.cuda.synchronize() + outs = [] + for st, buf, ctx, (flat, layer_off, ps, kvs, hs) in calls: + num_ctx = torch.ops.trtllm.k3_ctx_kv( + st.x, st.w, st.k_norm, st.cs, st.cpos, st.num_acc, ctx, st.slots, st.rows, st.table, st.counts, + flat, layer_off, ps, kvs, hs, EPS, MAX_CTX, PAGE, BLOCK, st.nkv, + ) # fmt: skip + outs.append((buf, ctx, num_ctx.long())) + torch.cuda.synchronize() + for i, (buf, ctx, num_ctx) in enumerate(outs): + want = alone[i // 2] + split = BACK_TO_BACK_SPLITS[i // 2] + assert torch.equal(buf.view(torch.int16), want[0].view(torch.int16)), ( + f"pool, {split} call {i % 2}" + ) + assert torch.equal(ctx, want[1]) and torch.equal(num_ctx, want[2]), ( + f"ctx_len / num_ctx, {split}" + ) + device = torch.cuda.current_device() + counters = [c for d, c in op._counters.items() if d.index == device] + assert counters and all(int(c.item()) == 0 for c in counters), ( + "the arrival counter is not at rest" + ) + + def test_supported_shapes(): """TP16 takes every step up to 8 x 8; TP4 (shared memory for the resident tokens) up to 32 tokens.""" op = _op() From ae6e9a95624b942b299fc2e43aa7cf75dabb1afd Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 20:26:30 -0700 Subject: [PATCH 159/161] [None][test] Kimi K3 k3_spec_accept and k3_markov: arrival counters at rest after back-to-back calls k3_spec_accept: launch five splits (block widths K and K + 1), three steps each, back to back on one stream with the state carried and no host synchronization, and require every grid's arrival counter to be zero once the device is idle. k3_markov: the multi-rank report also requires, on every rank, that the arrival word of each k3_markov workspace the run used is zero after the run (and that the run used at least one). A leftover count would otherwise surface only later, as an unrelated failure. Signed-off-by: Vasanth Sabavat --- .../kimi_k3/test_k3_markov.py | 10 +++++++- .../kimi_k3/test_k3_spec_accept.py | 25 +++++++++++++++++++ 2 files changed, 34 insertions(+), 1 deletion(-) diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_markov.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_markov.py index 965d3a115f33..d0a95dbf3648 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_markov.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_markov.py @@ -796,7 +796,15 @@ def report(g: Group) -> int: not_run = [f"S = {shard}: {len(rejected)} splits" for shard, _, rejected in runs if rejected] if not_run: g.say("Not run (rejected by pick_grid, see above): " + ", ".join(not_run)) - ok = g.all_ranks(ok and all(row["ok"] for row in rows)) + # The last CTA of every call zeroes its workspace's arrival word (flags[1]): at rest once the device is idle. + torch.cuda.synchronize() + workspaces = list(g.op._workspaces.values()) + at_rest = g.all_ranks(bool(workspaces) and all(int(ws["flags"][1]) == 0 for ws in workspaces)) + g.say( + f"Arrival words of the {len(workspaces)} k3_markov workspaces at rest after the run: " + f"{'yes' if at_rest else 'NO'}" + ) + ok = g.all_ranks(ok and at_rest and all(row["ok"] for row in rows)) g.say("ALL PASS" if ok else "FAIL") return 0 if ok else 1 diff --git a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_spec_accept.py b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_spec_accept.py index 8b231a8d2f77..4510ac9571b2 100644 --- a/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_spec_accept.py +++ b/tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_spec_accept.py @@ -548,6 +548,31 @@ def test_force_mode(): assert op.force_mode(force_value(drafts, "over"), drafts) == want, drafts +def test_arrival_counter_at_rest_after_back_to_back_calls(): + """Calls of several splits and block widths launched back to back on one stream, the state carried and no host + synchronization between them, as consecutive decode steps run: each grid's arrival counter (``CTAs arrived``, + which the last CTA of every call zeroes) is zero once the device is idle.""" + op = _op() + gen = torch.Generator(device="cuda").manual_seed(20261003) + with torch.inference_mode(): + calls = [] + for batch, tokens, block_delta in [(1, 8, 0), (8, 8, 0), (4, 2, 1), (8, 4, 1), (2, 8, 1)]: + drafts = tokens - 1 + st = make_state(gen, batch) + for step in range(3): + logits, draft = step_inputs(gen, batch, drafts, "plain", step) + calls.append((st, logits, draft, drafts + block_delta)) + torch.cuda.synchronize() + for st, logits, draft, block in calls: + fused(st, logits, draft, 0.0, block) + torch.cuda.synchronize() + device = torch.cuda.current_device() + counters = [scratch[1] for (dev, _), scratch in op._scratch.items() if dev.index == device] + assert counters and all(int(c.item()) == 0 for c in counters), ( + "an arrival counter is not at rest" + ) + + def test_supports(): """Every split the engine runs is supported (block K .. 8, the whole vocabulary or a TP shard); the limits are rejected, and unsupported calls raise before launching.""" From b52143183b0fdc9271bd7a5af2deccb959a5b501 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 19:10:50 -0700 Subject: [PATCH 160/161] [None][test] MNNVL catalog matrices: armed checks wait for every rank In the allgather_split and fusion_allreduce matrices, check_workspaces_are_armed_and_sized and check_create_refuses_on_every_rank assert locally that every word of a new workspace is -0.0, and each rank then makes its first call on that workspace. Nothing collective comes between, while create() lets peers push into a rank as soon as it returns. A rank that falls behind can therefore read its buffer after its peers' first call has written into it, and fails "every word -0.0". Seen under compute-sanitizer initcheck, with another process's kernel spinning on rank 0's GPU: rank 0 found 16,128 words changed, exactly the three peers' rows of the next call, AG(8500, 8, 896, 224). Gather the result over the ranks instead. The allgather is also the barrier that keeps every rank from pushing until every rank has looked, as the latent-reduce matrix's assert_clean does. Test-only. Signed-off-by: Vasanth Sabavat --- .../comm/_mnnvl_allgather_split_op_matrix.py | 12 +++++++----- .../comm/_mnnvl_fusion_allreduce_op_matrix.py | 12 +++++++----- 2 files changed, 14 insertions(+), 10 deletions(-) diff --git a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py index 0814b7cb9554..e96a2b09e3bf 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_allgather_split_op_matrix.py @@ -409,9 +409,11 @@ def check_workspaces_are_armed_and_sized() -> None: assert ws.buffer_bytes == BUFFER_BYTES and BUFFER_BYTES % 32 == 0 assert ws.comm_buffer(torch.bfloat16).shape == (3, BUFFER_BYTES // 2) armed = ws.lamport.view(torch.int32) - assert bool((armed == torch.tensor(INT32_MIN, dtype=torch.int32, device="cuda")).all()), ( - "every word -0.0" + every_word = bool( + (armed == torch.tensor(INT32_MIN, dtype=torch.int32, device="cuda")).all() ) + # The allgather is also the barrier that keeps every rank from pushing until every rank has looked. + assert R.all_true(every_word), "every word -0.0" rot(ws).unchanged("armed") for b, f in SPLITS: assert required_buffer_bytes(max(TOKENS), b, f, R.world) <= BUFFER_BYTES @@ -649,9 +651,9 @@ def check_create_refuses_on_every_rank() -> None: WS_C = MnnvlWorkspace.create(R.mapping, BUFFER_BYTES, fabric_handle=R.fabric) ROT[id(WS_C)] = Rotation(WS_C) armed = WS_C.lamport.view(torch.int32) - assert bool((armed == torch.tensor(INT32_MIN, dtype=torch.int32, device="cuda")).all()), ( - "the workspace created after the refusals: every word -0.0" - ) + every_word = bool((armed == torch.tensor(INT32_MIN, dtype=torch.int32, device="cuda")).all()) + # The allgather is also the barrier that keeps every rank from pushing until every rank has looked. + assert R.all_true(every_word), "the workspace created after the refusals: every word -0.0" rot(WS_C).unchanged("created after the refusals") call_and_check(AG(8500, 8, b, f), WS_C, "first call on the new workspace") call_and_check(AG(8501, 8, b, f), WS_A, "after the refusals") diff --git a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_fusion_allreduce_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_fusion_allreduce_op_matrix.py index 58ac651c2d65..474932eb2f23 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_mnnvl_fusion_allreduce_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_mnnvl_fusion_allreduce_op_matrix.py @@ -428,9 +428,11 @@ def check_workspaces_are_armed_and_sized() -> None: assert ws.buffer_bytes == BUFFER_BYTES and BUFFER_BYTES % 32 == 0 assert ws.comm_buffer(torch.bfloat16).shape == (3, BUFFER_BYTES // 2) armed = ws.lamport.view(torch.int32) - assert bool((armed == torch.tensor(INT32_MIN, dtype=torch.int32, device="cuda")).all()), ( - "every word -0.0" + every_word = bool( + (armed == torch.tensor(INT32_MIN, dtype=torch.int32, device="cuda")).all() ) + # The allgather is also the barrier that keeps every rank from pushing until every rank has looked. + assert R.all_true(every_word), "every word -0.0" rot(ws).unchanged("armed") one_shot = required_buffer_bytes(max(TOKENS), H_MODEL, R.world, torch.bfloat16, 1 << 62) assert one_shot == BUFFER_BYTES, f"the largest one-shot call needs {one_shot} bytes" @@ -700,9 +702,9 @@ def check_create_refuses_on_every_rank() -> None: WS_C = MnnvlWorkspace.create(R.mapping, BUFFER_BYTES, fabric_handle=R.fabric) ROT[id(WS_C)] = Rotation(WS_C) armed = WS_C.lamport.view(torch.int32) - assert bool((armed == torch.tensor(INT32_MIN, dtype=torch.int32, device="cuda")).all()), ( - "the workspace created after the refusals: every word -0.0" - ) + every_word = bool((armed == torch.tensor(INT32_MIN, dtype=torch.int32, device="cuda")).all()) + # The allgather is also the barrier that keeps every rank from pushing until every rank has looked. + assert R.all_true(every_word), "the workspace created after the refusals: every word -0.0" rot(WS_C).unchanged("created after the refusals") call_and_check(AR(8500, 8, H_MODEL, residual=True), WS_C, "first call on the new workspace") call_and_check(AR(8501, 8, H_MODEL, residual=True, path="two"), WS_A, "after the refusals") From 13b26f42b44caa04cb0dae2691fee0ecfb3cefa8 Mon Sep 17 00:00:00 2001 From: Vasanth Sabavat Date: Sat, 3 Oct 2026 20:35:28 -0700 Subject: [PATCH 161/161] [None][doc] Kimi K3 k3_spec_accept and k3_ctx_kv: calls on a device are ordered on one stream Each op keeps a module-level arrival counter per device (k3_spec_accept: per device and logits width) that every call shares; the last CTA to arrive does the step's bookkeeping and re-arms the counter to zero. Two calls not ordered on one stream would mix their arrivals. State that contract in each module docstring, as k3_head_gemv does for its workspace. Docstrings only; no code change. Signed-off-by: Vasanth Sabavat --- tensorrt_llm/_torch/cute_dsl_kernels/k3_ctx_kv/op.py | 3 ++- tensorrt_llm/_torch/cute_dsl_kernels/k3_spec_accept/op.py | 2 ++ 2 files changed, 4 insertions(+), 1 deletion(-) diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_ctx_kv/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_ctx_kv/op.py index 26c330ba8152..f0cf37f1a05e 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_ctx_kv/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_ctx_kv/op.py @@ -18,7 +18,8 @@ of every drafter layer, k_norm, NeoX RoPE, the write mask and the paged store into the drafter's context pool, then ``ctx_len += num_accepted`` (clamped) and the context length each request may advertise, in one launch (see ``k3_ctx_kv_kernel``). Compiled on the first call for its shape (B, K + 1), which must happen outside CUDA-graph -capture. +capture. Every call on a device shares one arrival counter, so the calls must be ordered on one stream: two calls not +ordered on one stream would mix their arrivals. """ from __future__ import annotations diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/k3_spec_accept/op.py b/tensorrt_llm/_torch/cute_dsl_kernels/k3_spec_accept/op.py index 82bea70f884f..369536f40f02 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/k3_spec_accept/op.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/k3_spec_accept/op.py @@ -21,6 +21,8 @@ With a :func:`workspace`, the target logits may stay vocabulary-sharded (this rank's bf16 columns of a TP column-parallel head): the ranks exchange their row maxima over a multicast Lamport buffer instead of all-gathering the logits, with the same argmax. Compiled on the first call per shape, which must happen outside CUDA-graph capture. +Every call on a device with the same number of logits columns shares one arrival counter, so the calls must be ordered +on one stream: two calls not ordered on one stream would mix their arrivals. """ from __future__ import annotations