Conversation
added 30 commits
October 2, 2026 09:36
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 <[email protected]>
…ache 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 <[email protected]>
…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 <[email protected]>
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 <[email protected]>
…t 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 <[email protected]>
…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 <[email protected]>
…ive 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 <[email protected]>
…lti-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 <[email protected]>
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 <[email protected]>
…re 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 <[email protected]>
… 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 <[email protected]>
…luster 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 <[email protected]>
…ads 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 <[email protected]>
…eases 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 <[email protected]>
…h 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 <[email protected]>
…ce 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 <[email protected]>
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 <[email protected]>
…e the lane-0 block-reduce stores) The ssm/kda_decode entry's receipts run the fixed kernel. NVIDIA#19817 (this branch's base) and NVIDIA#19830 touch disjoint files. Signed-off-by: Vasanth Sabavat <[email protected]>
…rier fix Signed-off-by: Vasanth Sabavat <[email protected]>
…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 <[email protected]>
…ract (target logits) Signed-off-by: Vasanth Sabavat <[email protected]>
Signed-off-by: Vasanth Sabavat <[email protected]>
… + 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 <[email protected]>
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 <[email protected]>
…de 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 <[email protected]>
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 <[email protected]>
…uction 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 <[email protected]>
…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 <[email protected]>
…tries 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 <[email protected]>
…V, 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 <[email protected]>
added 30 commits
October 3, 2026 12:32
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 <[email protected]>
…_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 <[email protected]>
…er'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 <[email protected]>
…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 <[email protected]>
…ive 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 <[email protected]>
…f 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 <[email protected]>
… 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 <[email protected]>
… 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 <[email protected]>
…n 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 <[email protected]>
… 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 <[email protected]>
…fter'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 <[email protected]>
… 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 <[email protected]>
…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 <[email protected]>
…ires 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 <[email protected]>
…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 <[email protected]>
…rafter'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 <[email protected]>
… 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 <[email protected]>
…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 <[email protected]>
…-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 <[email protected]>
…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 <[email protected]>
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 <[email protected]>
…ush 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 <[email protected]>
… 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 <[email protected]>
…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 <[email protected]>
…M-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 <[email protected]>
…e 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 <[email protected]>
…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 <[email protected]>
…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 <[email protected]>
…e 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 <[email protected]>
…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 <[email protected]>
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Description
Depends on #19839 at ebfd380 (its follow-up commits: F59 routing fix, trim, l0_cpu decode_step entry, sm_100
count, declared ops; through it on #19833 and #19832), #19828, #19836, #19816, #19814 and #19827: this branch
carries their commits through merges.
Carries P11's commit (#19842, DSpark trained mask-token row) so DSpark acceptance is right; it drops out of k3rest
when #19842 merges.
The rest of Kimi K3 on ModelingV2, in parts that cannot merge independently: the collective kernels and the drafter
attention build on the catalog's conventions for stateful entries and on its MNNVL workspace (part 1), the MoE
engines for route B's expert split use the collectives' API (part 4), and route B's target (part 5) is the layout
those engines serve. They land together. Route B's target does not call the part 4 engines yet: its routed experts
run the generic MoE path, and the wiring follows on this branch (see Follow-ups). Section 6 moves the rest of
#19839's decode step onto parts 1 and 3, and section 7 moves its DSpark drafter onto part 2's entries.
1. Catalog conventions for stateful entries, and
comm/mnnvl_allreduce_attn_resModelingV2's catalog had no convention for an op whose result depends on state that outlives the call.
Conventions (
catalog/index.yaml,README.md):where the op is, eager, before any CUDA-graph capture. There is no module-level dict and no environment switch.
The type sits beside its entries and is not an entry itself; it launches nothing per call.
## Statesection: the object's contents and size; who creates it, and when; which opsmay share one object; the call order every rank keeps across layers and steps; what a later launch reads; how it
is re-armed.
negative control that breaks the call order, not only single calls.
records the world size its matrix ran at.
--world-sizeand--launcher srun, so the same rank bodyrecords a multi-node run.
_rank_job.runtakes the world size.trtllm::mnnvl_allreduce_attn_res(mnnvlAllreduceKernels.cu/.h,thop/allreduceOp.cpp)updated = prefix_sum + the sum over the ranks, andnormed= RMSNorm of the attention-residual selection over the snapshot bank andupdated.attn_res_add_rmsnorm_fwd. Its reduction order is fixed, soevery rank gets the same bits.
hidden / 1024CTAs owns a token. The per-token statistics are summed in a fixed order, thecluster part through distributed shared memory. Each thread arrives on the cluster barrier early and waits before
its first write into a peer CTA's shared memory.
reduceOneshotLamportis the one-shot kernel's Lamport reduction as a function, for the new kernel. Part 1 leavesthe existing kernels as they are: on a build of part 1 alone their SASS (
oneshotAllreduceFusionKernel,twoshotAllreduceKernel,rmsNormLamport, every instantiation) is identical to that of main + [None][fix] MNNVL all-reduce: cluster barrier use in the RMSNorm-fused kernels #19836. Part 3then moves the one-shot kernel onto the same function and adds its early trigger.
comm_bufferandbuffer_flagsmutable.MNNVLAllReduce.allreduce_attn_res_rmsnormcalls the op on the module's own workspace.comm/mnnvl_allreduce_attn_resMnnvlWorkspace(comm/mnnvl_workspace.py): three Lamportbuffers, the flag words and the multicast handle.
MnnvlWorkspace.createis collective over the TP group only: its communicator comes from the new_get_mnnvl_tp_group_comm(distributed/ops.py,MPI_Comm_create_groupovermapping.tp_group). Beforeallocating, the ranks agree that each of them can; if one cannot, every rank raises and frees that communicator.
## Statesection states the call-order rule. A swapped pair of same-shaped calls is silentlywrong on every rank, and two ranks ordering calls on two workspaces differently on one stream deadlock.
2. Kimi K3 DSpark drafter attention
trtllm::k3_drafter_attnandtrtllm::k3_drafter_attn_qknorm(tensorrt_llm/_torch/cute_dsl_kernels/k3_drafter/,CuTe DSL, SM 100). The block attention of Kimi K3's DSpark drafter:
and to the block's own K / V. Those are read from the projection output and not written to the cache.
k3_drafter_attn_qknormtakes the raw projection output and applies the per-head q / k RMSNorm and the NeoX RoPEinside the kernel, with
fused_qk_norm_rope's arithmetic.present; that first call must be outside CUDA-graph capture. Lengths, page tables and qkv are read on the device,
so a captured call replays with them rewritten in place.
context kernel) in one kernel, without writing the cache. The op's test compares the two.
Catalog entries (
modeling_v2/catalog/attention/):k3_drafter_attnandk3_drafter_attn_qknorm: contracts and wrappers, certified on sm_100.fused_qk_norm_rope, which the drafter also calls, gains the drafter's cell and an sm_100 receipt. Its sm_103receipt stays, noted as pre-dating that cell, and CI's l0_b300 run keeps re-certifying the other tests: the K3
cell skips off SM 10.0.
3. Kimi K3 collective decode kernels and their catalog entries
Kimi K3's decode kernels that exchange data across the TP group, the MNNVL additions they need, and their
modeling_v2 catalog entries, at the TP16 per-rank shapes on sm_100.
MNNVL (
thop/allreduceOp.cpp,communicationKernels/mnnvlAllreduceKernels.cu, newmnnvlAllGatherKernels.{h,cu},distributed/ops.py,custom_ops/cpp_custom_ops.py):trtllm::mnnvl_fusion_allreducetakesone_shot_max_bytes: messages up to that size go one-shot, larger onestwo-shot. The default stays main's 1 MiB.
MNNVLAllReducekeeps it 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).MNNVLAllReduceincluded. The one-shot fusion kernel nowreleases its programmatic dependents right after its own grid-dependency wait; main's released them after the
reduction. A dependent kernel (a GEMV) can then launch and stream its weights while the all-reduce waits for the
other ranks. A dependent still reads the output and the Lamport flags only after its own grid-dependency wait,
which waits for this whole grid, so results do not change. The kernel's Lamport reduction also moves into
reduceOneshotLamport, the function the attention-residual one-shot kernel of part 1 uses: the code moves, theorder of the reduction does not. Evidence, on GB200 (4 ranks): main's
multi_gpu/test_mnnvl_allreduce.pybodies pass on this build (110 cases at 4 ranks, 109 at 2, the graph-capture cases); main's
MNNVLAllReduceat[M, 7168]bf16, M 1 to 2048, plain and RMSNorm-fused, one-shot and two-shot, returns the previous kernels' outputsbit for bit, and its per-call time in a captured chain of dependent calls is within -0.56 / +0.12 us of theirs
(table below). The SASS of every MNNVL kernel changes, the two-shot, RMSNorm and attention-residual kernels'
included, because the kernels' parameter struct gains the
earlyTriggerfield; only the one-shot kernel's sourcechanges.
trtllm::mnnvl_fusion_allreducenow declaresbuffer_flagsmutable (Tensor(b!)): every call writes it.trtllm::mnnvl_allgather_split(input, bf16_columns, world_size, comm_buffer, buffer_flags)(
MNNVLAllReduce.allgather_split): a one-shot all-gather of fp32 rows whose leading columns travel and arrive asbf16, the rest as fp32, on the all-reduce's workspace and its Lamport rotation. Kimi K3 gathers its row-sharded MoE
head with it. Its schema declares
comm_bufferandbuffer_flagsmutable, it checksworld_sizeagainst theworkspace, and it has a fake implementation.
Kimi K3 collective kernels (
tensorrt_llm/_torch/cute_dsl_kernels/, CuTe DSL, sm_100). The kernel files are theones the Kimi K3 deployment runs, unchanged (one reformatted by the repo's hooks, AST-identical). The routed-experts
kernel's fused all-reduce push paths compile only for a push build, which part 4's wrappers request.
k3_sandwich/:trtllm::k3_sandwich_oproj,_tailand_plain, a row-parallel projection, its TP all-reduceand the residual update in one kernel for up to 8 tokens (the post-attention step, the pre-attention step after
the MoE tail, and the drafter's plain residual add + RMSNorm). Their arithmetic is that of the projection followed
by the MNNVL one-shot all-reduce they replace, bit for bit. The tail's fold mode (
K3SandwichLatentExchange) shipswith the kernel but has no caller and no catalog cell here: out of scope.
k3_fused_moe/:trtllm::k3_moe_front(the MoE front: head GEMV, head all-gather, top-16 routing, MXFP8 latent,shared gate_up + SiTU in one kernel);
trtllm::k3_moe, the persistent routed-experts kernel (FC1 + SiTU + FC2 withthe routing-weighted combine over the TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers) for up to 8 tokens and, in its wide
build, 64; and
trtllm::k3_latent_reduce, the latent all-reduce at decode size as the consumer of pushed partials.k3_route_quant/:trtllm::k3_route_quant, the top-16 routing and MXFP8 input quantization, bit-identical totrtllm::kimi_k3_noaux_tc_mxfp8_quant.Every stateful op runs on state the caller creates and passes. No module-level dict holds buffers, counters or
Lamport words:
K3SandwichWorkspace(the sandwiches' all-reduce buffer, shared by the three sandwich ops of a TP group),K3MoeHeadWorkspace(the front's head all-gather buffers and ready words) andK3LatentExchange(the latentexchange) are built by an eager
create(mapping, fabric_handle=None)that is collective over the TP group's ranksonly (its communicator is made from them alone, as part 1's
MnnvlWorkspace.createdoes), with the same failuremodel: 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 each frees thecommunicator made for the call. A failure returned by the allocation is agreed the same way. A rank that fails
inside the allocation's handle exchange can still leave its peers waiting there.
K3MoeState/K3MoeWideStatehold a device'sk3_moescratch (up to 8 / 64 tokens) andK3MoeLayera layer'sexperts and counters; their constructors refuse capture, and the state is per rank.
mutates_argsnames the buffers it writes,trtllm::k3_moeincluded (the state's scratch, the layer'scounters, and with a head_flags state the head workspace's ready words and epoch).
re-exports the type it takes.
Catalog entries (
_experimental/modeling_v2/catalog/), each a contract with a## Statesection where the ophas state, a one-call wrapper, and a GPU test that drives call sequences on real state:
comm/k3_sandwich_oproj,comm/k3_sandwich_tail,comm/k3_sandwich_plain,comm/mnnvl_fusion_allreduce,comm/mnnvl_allgather_split,comm/k3_latent_reduce,moe/k3_moe_front(publish=Truereleases the ready words a head_flagsk3_moeacquires),
moe/k3_moe, and the statelessmoe/k3_route_quant.comm/allgather, which the Kimi K3 LM head calls,gains an
sm_100receipt; itssm_103receipt stays.4. MoE route B kernels
Kimi K3's decode MoE runs every routed expert of a token on every rank when the experts are split 16 ways by
intermediate (moe TP16 x EP1). At one or two tokens each routed expert is a GEMV, and
trtllm::k3_moe's 128-rowtiles leave most of the GPU idle. Part 4 adds the two engines for those steps, the push builds, and the engines' catalog
entries. No target calls the engines yet: route B's target (part 5) runs its routed experts on the generic MoE path,
and wiring it to these engines is a follow-up.
trtllm::k3_moe_m1(k3_moe_m1_kernel.py): the routed experts of one or two tokens as one weight-stream CuTeDSL kernel. The experts' weight rows are spread over every CTA. FC1 (MXFP4 x MXFP8), SiTU, the MXFP8 intermediate,
FC2 and the routing-weighted combine run in one launch. The combine keeps
trtllm::k3_moe's order, so the outputhas its bits except where FC1's two partial sums round an intermediate value differently (at most one bf16 ulp of
the row's max in the tests). The weights are the TRTLLM-Gen W4A8_MXFP4_MXFP8 buffers read in place, including the
loader's padding of a 192-wide shard to 256; only the 192 real values are streamed.
trtllm::k3_moe_m2(k3_moe_m2_kernel.py): the two-token form for intermediates up to 256. Each distinctexpert's intermediate is computed for both tokens, and FC2 starts per group of 4 experts as soon as that group's
FC1 rows are complete.
trtllm::k3_latent_reduce(part 3), each engine andtrtllm::k3_moecan store this rank's partial into its slot of every rank'sK3LatentExchangeinstead ofreturning it:
K3MoeM1Layer.push,K3MoeM2Layer.push, andK3MoeLayer.pushaftertrtllm::k3_route_quantortrtllm::k3_moe_front. A push followed by the reduce equalsMNNVLAllReduce's one-shot of the plain partials bitfor bit.
K3MoeM1State/K3MoeM2Stateown the small workspace every layer on one device shares: theintermediate rows, two hand-off count sets by epoch parity, one epoch per CTA. Every call zeroes the set the next
call uses and advances the epochs, so the workspace never needs a reset.
create()allocates it and compiles theplain and listed push builds before any CUDA-graph capture.
moe/k3_moe_m1andmoe/k3_moe_m2, with## Statesections, wrappers for the plain and push forms,and their tests; and the push form of
moe/k3_moe(k3_moe_push).5. Route B
A second ModelingV2 target for Kimi K3,
models/kimi_k3_vl/kimi_k3_mxfp4__sm_100__tp16_moetp16ep1: the samecheckpoint, GPU and attention layout as #19839's
kimi_k3_mxfp4__sm_100__tp16_moetp4ep4, with the routed expertssplit 16 ways by tensor. Every rank holds all 896 experts at a sixteenth of their intermediate width (192 of 3072
values, zero-padded to 256 when they load). It is Kimi K3's layout for decoding without speculation; the
tp16_moetp4ep4target serves speculative decoding.Routing (
models/kimi_k3_vl/routing.py): the K3 routing tree gains the second key. The expert split decides thetarget:
moe_tensor_parallel_size4,moe_expert_parallel_size4tp16_moetp4ep4moe_tensor_parallel_size16,moe_expert_parallel_size1, both settp16_moetp16ep1moe_tensor_parallel_size1,moe_expert_parallel_size16"None" means:
autoruns the built-in model, andrequireraises with the routing trace. The trace'sparallelline prints
split_set, soexplainshows why an unset split misses.The target is a copy of
tp16_moetp4ep4's modules, changed only in marked blocks. Targets share no files(
test_targets_do_not_share_files), so this target ships its ownmodeling.py,weights.pyanddecode_gemv.py,copied from
tp16_moetp4ep4's. Each place the copy differs is a block between# >>> route B: <why>and# <<< route B:split set explicitly, and no speculative decoding config (the message names the 4 x 4 split for DSpark, DFlash and
SA);
ModelingV2KimiK3Mxfp4Sm100Tp16Moetp16ep1;(
kda_token_states) and no KDA verify kernels (k3_kda_attn,k3_kda_verify,trtllm::kda_mtp_decode), withREQUIRED_TRTLLM_OPSandUNCERTIFIED_GENERIC_CALLSto match.decode_gemv.pyis byte-identical, and so are the text model, the step classification, the decode kernels' wiring,the weight load and the engine checks. The causal LM is still 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.
Section 7's drafter is in the copy too (
K3DSparkDrafter, the_build_draft_modeloverride and the drafter's fourdecode GEMV sites) and is never built here;
K3DecodeGemvs.createwarms the four sites once on zero weights atstartup, as it does on
tp16_moetp4ep4.test_modeling_v2_kimi_k3_drift.pyreads both targets' files, imports neither, and fails on any difference outsidethe blocks, so a change to
tp16_moetp4ep4reaches this copy, or becomes one of its blocks, in the same PR.Route B's routed experts run the generic MoE path (the TRTLLM-Gen MoE on the padded experts) on every step. This
target does not call part 4's engines yet; wiring them in, as route B blocks, follows on this branch (see
Follow-ups).
6. The fused decode path of the
tp16_moetp4ep4target#19839's target runs a decode step's attention, GEMVs and dense MLP on the Kimi K3 decode kernels. These commits move
the rest of a decode step onto the entries of parts 1 and 3: the collectives around the attention, the MoE layers,
and the residual updates between them.
decode_comm.py,K3DecodeComm. The TP group'sMnnvlWorkspace(4 MiB buffers) andK3SandwichWorkspace.The target creates both collectively in
post_load_weights, before any CUDA-graph capture, wherever everyattention all-reduce runs over MNNVL.
its unreduced o_proj output.
comm/mnnvl_allreduce_attn_resthen runs the all-reduce with the residual update asits epilogue.
comm/k3_sandwich_oprojkernel, bit for bit o_proj followed by the MNNVL entry.comm/k3_sandwich_tailruns the tail GEMV, its all-reduce and the next layer's residual update in one kernel. After the last layer it is
the final norm's update.
decode_moe.py. A MoE layer of at most 8 tokens runs:moe/k3_moe_front: this rank's slice of the MoE head (latent-down and router rows) as one GEMV, the slices'all-gather, the top-16 routing, the MXFP8 latent, and the shared experts' gate_up + SiTU;
moe/k3_moeon aK3MoeState, then the routed experts' all-reduce of the latent;[RMSNorm(latent) slice | shared activation] @ [latent up columns | shared down].PendingTail.A wide decode step (9-64 tokens) keeps the sharded head (
comm/mnnvl_allgather_split) and the row-parallel tail onM-general ops:
moe/k3_route_quantandmoe/k3_moeon aK3MoeWideState. The consumer reduces its unreducedoutput with a plain all-reduce and then runs the fused add + attn_res + RMSNorm.
The decode step's numerics.
the same function, rounded differently.
The attention's interface.
will_run_decode_branch(attn_metadata, step)tells the layer, before the attentionruns, whether the attention will take its decode branch.
reduce_output=Falsereturns the attention's outputunreduced, on every path.
project_output=Falsereturns it before o_proj: on the decode branch for MLA, on everypath for KDA.
Logging. The startup line gives the number of layers that
comm/k3_sandwich_oprojand the MoE decode path take.Each reason a MoE layer stays on the generic path is logged once.
7. The route A target's DSpark drafter on part 2's entries
The target builds DSpark's standalone GQA drafter itself:
K3DSparkDrafter, the stockGQADSparkForCausalLMwitha decode block on part 2's entries and the decode GEMV sites. The worker (
DSparkWorker, throughget_spec_worker),its context projection and context K / V, the Markov head and the acceptance stay upstream code. The weights load
through the stock draft loader, and the embedding and LM head are the target's.
Engine change (outside modeling_v2):
SpecDecOneEngineForCausalLM._build_draft_model(), which the one-engineshell calls once to build its drafter. Its default is the mode registry's
get_draft_modelwith the samearguments, so every existing model is unchanged. A model that owns its drafter overrides it instead of copying the
shell's speculative setup (draft config, separate draft KV cache, epilogue, logits processor, worker). A
hardware-agnostic cpu_only test pins both behaviours.
The decode block, per layer:
drafter_qkvsite, [512, 7168];attention/k3_drafter_attn_qknorm: the q / k RMSNorm, the NeoX RoPE and the block attending to its paged contextand its own k / v, without storing the block's k / v;
drafter_o([7168, 384],k3_ctm_gemv), then the module's all-reduce;drafter_gate_up([1792, 7168]), and down with the SiLU-and-mul ondrafter_down([7168, 896],k3_ctm_gemv_swiglu), then the all-reduce.The four sites are cells the GEMV contracts certify at 1..8 rows, at this drafter's TP16 per-rank shapes. Above 8
rows each projection runs its module. The norms and the residual adds are the stock modules'.
Every other block runs the stock block forward: a split
attention/k3_drafter_attn_qknormdoes not certify(
DRAFTER_ATTN_SPLITSmirrors its contract), another attention backend, head layout, RoPE base, epsilon or MLP, acache layout the kernel does not read, or, under CUDA-graph capture, an attention compile key that has not run
eagerly. The stock drafter classes and the builder's checks are declared in
UNCERTIFIED_GENERIC_CALLS; the newops are in
REQUIRED_TRTLLM_OPS.R x 7: DSpark's draft block under
shift_labelismax_draft_lentokens, 7 here (DSparkWorker._draft_block_width).Part 2's tests and contracts gain 1x7..8x7 (the drafter reviewed them), and the target takes those splits.
Follow-ups
Landed as new commits after the first review round, without rewrites:
trtllm::k3_moecompiles on routing that requires grad. A load-time warm-up outside inference mode routes withe_score_correction_bias, annn.Parameter, and the first (compiling) call raisedBufferError: Can't export tensors that require gradientfrom DLPack._viewnow exports a detached view of the same storage, with no copy.test_routing_from_a_parameter_that_requires_gradfails before the fix and passes after.moe/k3_moe,moe/k3_moe_m1,moe/k3_moe_m2); part 3'seight contracts record theirs.
k3_spec_accept(acceptance with vocabulary-sharded target logits),k3_ctx_kvandk3_markov, one launch each where the worker ran 47, 52 and 119 small torch kernels. The worker takes them behindits Kimi K3 gate (
k3_decode). Tests are listed inl0_b200.ymlandl0_gb200_multi_gpus.yml; the sharded andMarkov tests run under pytest on 4 local MPI ranks.
capture_view, a captured layer's slot of the capture buffer. The draft block's rows and thegen requests' page-table rows are now views where the slots are one run.
all-reduces (
comm/k3_sandwich_plain, or the fused MNNVL all-reduce above the plain path's rows). The DSpark tapis written by the sandwich tail.
moe/k3_moe'spush form plus
comm/k3_latent_reduce. It sums in the one-shot's order, so the bits are the same: the 8 fixed-promptoutputs are identical with and without it.
k3_moe_m1at one token,k3_moe_m2at two,k3_moeover all 896experts up to 8, with the push form and the latent reduce. Route B no-spec median TPOT drops from 8.83 / 9.44 /
10.37 / 13.54 ms to 4.32 / 4.71 / 6.34 / 7.83 ms at BS 1 / 2 / 4 / 8.
trtllm::k3_route_quant, bit for bitthe C++ op's outputs. Both targets' gates return
KimiK3MoeRoutingMethod, so the generic TRTLLM-Gen MoE routesoutside the kernel.
shared-expert slice at every world.
k3_moeand its push buildon the same layer back to back.
k3_moeclaims its first tile from the layer's counters before its grid-dependencywait, so with PDL the push build could take tiles from the plain call's queue. That call, or the next call on the
layer, then hung, and the latent reduce waited on every rank.
warm-up with every build compiled hung in 3 of 4 runs at 16 ranks.
moe/k3_moestates the rule.all-reduce,
k3_latent_reduce, or a graph boundary.With every published PR merged locally, the follow-ups bring DSpark at forced acceptance 6 to the speed of the
measured stack. Median TPOT is 1.010 / 1.697 / 1.929 / 2.541 ms against 1.008 / 1.701 / 1.940 / 2.536 ms at BS 1 / 2 /
4 / 8, each gap inside the run-to-run spread. GSM8K (1319 problems) is 96.21 with DSpark, 96.44 on route A without
speculation and 96.74 on route B, each within one point of its baseline.
Design choices
create()that the caller (the target'spost_load_weights) runs before capture, with the failure model of part 1 (the ranks agree before allocating; every rank raises together). No module-level dicts, no global dispatch switches. The contract gets a## Statesection (contents and size; creator and when; which ops may share one object; call order across ranks and layers; what a later launch reads; re-arm). Tests drive real-state call sequences (layers x steps, calls queued across ranks without host synchronization, capture + replay, two objects interleaved) plus a negative control.mutates_argsnames every written buffer. Here:MnnvlWorkspace(part 1),K3SandwichWorkspace,K3LatentExchange,K3MoeHeadWorkspace(collectivecreate(), over the TP group's ranks only);K3MoeState/K3MoeWideStateandK3MoeLayer(per rank, constructors)._rank_jobtakes the world size, and the stateful matrices take--world-sizeand--launcher mpirun|srun(_lockstep). CI runs every multi-GPU test at 2 or 4 ranks on one GB200 tray (l0_gb200_multi_gpus.yml); there is no 16-GPU CI. The 16-rank results below are manual runs; the contracts record both world sizes.sm_100receipts and skip on other architectures. A reused entry gains its sm_100 cell in the PR of its first Kimi K3 caller and keeps itssm_103receipt (it pre-dates the Kimi K3 cells and covers the existing functions, which l0_b300 re-certifies):fused_qk_norm_rope(part 2),comm/allgather(part 3). Single-GPU tests go inl0_b200.yml, 4-rank tests inl0_gb200_multi_gpus.yml.comm/mnnvl_allreduce_attn_res, part 3'scomm/mnnvl_fusion_allreduceandcomm/mnnvl_allgather_split). One caller-ownedMnnvlWorkspace(comm_buffer,buffer_flags, multicast handle) is shared by all three;one_shot_max_bytesis per call.get_spec_workerand thetarget_logitshook (#19814). The worker's own kernels stay upstream code. The drafter model is target code, so its ops are catalog entries.attention/; there is nospec/category.modeling_v2/README.mdgoverns layout and records.trtllm::k3_moeis, so their writes are declared:mutates_argsnames the workspace, the exchange's words andout.K3MoeM1State/K3MoeM2Stateare built by an explicit, eagercreate()that the target runs inpost_load_weights, before capture, and that compiles the plain and listed push builds. No module-level state, no environment switch; compiled builds are cached per device and configuration (code, not state).tp16_moetp4ep4as this target's base (TARGET_BASES), the isolation rule would allow importing exactly that target, a claims test would keep the exception to layouts of one checkpoint (same checkpoint and SM, another parallel segment, no chains), and the target would subclass route A's class with its own layout: about 30 lines instead of ~2.4k copied. It changes the rule, so it is the staircase owners' call; if they choose it, it lands as a follow-up commit that replaces the copy.Mapping.moe_tp_ep_user_specified), so routing reads the same flag and the target asserts it: the engine and the target always agree on the layout.explaincannot name one), so the target asserts it is absent at construction. DSpark's wide steps need the wide k3_moe build, whose per-group tables do not fit 896 local experts.K3DecodeComm,K3DecodeMoeandK3DecodeMoeLayerare the target's caller-owned state, built once inpost_load_weights(part 1's conventions). Their collective parts (MnnvlWorkspace,K3SandwichWorkspace,K3MoeHeadWorkspace) are created on every rank at the same point. Every kernel that compiles is run once there, on every rank, so no CUDA-graph capture compiles one.will_run_decode_branchbefore calling the attention. MLA's KV-cache writes allow one attention run per step, so there is no fallback after the fact.moe/k3_moe's push form (part 4) into aK3LatentExchangepluscomm/k3_latent_reducesums in the same order; moving to it is a separate performance change.SpecDecOneEngineForCausalLM._build_draft_model(), whose default isget_draft_modelwith the same arguments, so every existing model is unchanged. The Kimi K3 target overrides it to buildK3DSparkDrafterinstead of copying the shell's speculative setup. A hardware-agnostic cpu_only test pins both behaviours.Test Coverage
Part 1 (GB200, one tray, 4 GPUs):
comm/mnnvl_allreduce_attn_resmatrix (_mnnvl_allreduce_attn_res_op_matrix.py, 4 ranks): 9 / 9 checkspass, started as CI starts it (pytest on
test_modeling_v2_mnnvl_allreduce_attn_res_op_matrix.py, mpirun through_rank_job) and with--launcher srun:cover every snapshot count 0-11);
MnnvlWorkspace.createwith one rank short of memory: every rank raises and frees its communicator, and thenext call is correct;
MnnvlWorkspace.createunder TP 2 x PP 2: one TP group creates a workspace while the other group's ranks makeno MNNVL call.
updatedequals the exact sum bit for bit;normedis within 7.8e-3 (relative to the largest value) of thefp32 reference; every output is bitwise equal on all ranks.
test_k3_mnnvl_comm.py(attn_res), 4 ranks:MNNVLAllReduce.allreduce_attn_res_rmsnormagainst the unfusedpath (the default MNNVL all-reduce, then
attn_res_add_rmsnorm_fwd): PASS.of
oneshotAllreduceFusionKernel,twoshotAllreduceKernelandrmsNormLamport(276 functions) is identical tothat of a build of main + [None][fix] MNNVL all-reduce: cluster barrier use in the RMSNorm-fused kernels #19836, and
test_mnnvl_allreduce.py's 42 RESIDUAL_RMS_NORM cases pass at 4 and 2ranks. Part 3 changes those kernels; its evidence is under part 3.
Part 2 (GB200, single GPU per file, listed in
l0_b200.yml):tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_drafter_attn.py(the op's test) covers the splits 1x1 to5x1, 1x8, 2x4, 4x2, 8x1, R x 8 and R x 7 for R <= 8 (R x 7 is DSpark's block under
shift_label), with mixedcontext lengths: an fp32 reference and the DFlash
TRTLLM path it does the work of (
append_paged_kv_cache+ the trtllm-gen context kernel, non-causal); no NaN (thepool's unused rows are NaN) and the pool left untouched; reruns bit-identical and requests isolated; CUDA-graph
replays with rewritten inputs; the norm / RoPE variant against
fused_qk_norm_rope+k3_drafter_attn. 105 passed.test_modeling_v2_k3_drafter_attn.pyand..._qknorm.py: the entries' cell-by-cell tests against fp64references written in the tests, the R x 7 splits included. 36 and 16 passed.
test_modeling_v2_fused_qk_norm_rope.pygainstest_bf16_neox_kimi_k3_drafter: the whole file, 10 passed.Part 3 (G4):
CI lists.
l0_b200.yml(single GPU):kimi_k3/test_k3_fused_moe.py,test_k3_moe_wide.py,test_k3_route_quant.py,test_k3_moe_front_geometry.py(no GPU work: the front's head geometry per TP size),modeling_v2/moe/test_modeling_v2_k3_moe.py,moe/test_modeling_v2_k3_route_quant.py.l0_gb200_multi_gpus.yml(4 ranks):kimi_k3/test_k3_sandwich.py,test_k3_moe_front.py,test_k3_latent_reduce.py,test_k3_mnnvl_comm.py(its split all-gather half), and thecollected matrices of the seven stateful entries plus
comm/allgather's.Results on GB200 (sm_100), 4 GPUs of one tray, this branch's head on a build of its C++:
test_k3_fused_moe.pytest_k3_moe_wide.pytest_k3_route_quant.py,test_k3_moe_front_geometry.py, the claims / routing tests, the two kernel lintstest_modeling_v2_k3_route_quant.pymoe/test_modeling_v2_k3_moe.pymoe/test_modeling_v2_k3_route_quant.pycomm/test_modeling_v2_{k3_sandwich_oproj,k3_sandwich_tail,k3_sandwich_plain,mnnvl_fusion_allreduce,mnnvl_allgather_split,k3_latent_reduce,allgather}_op_matrix.py,moe/test_modeling_v2_k3_moe_front_op_matrix.py), pytest startingmpirunwith 4 rankssruntest_k3_sandwich.py,test_k3_moe_front.py,test_k3_latent_reduce.py,test_k3_mnnvl_comm.py(4 ranks, theirmain()) and their one-process testsmulti_gpu/test_mnnvl_allreduce.pybodiestest_k3_moe_wide.py's former 256 router-dump cases needed a recorded dump and skipped without one; they are nowgenerated in the test (seeded), so nothing in G4's tests skips on sm_100.
MNNVL one-shot change, stock-caller timing (main's
MNNVLAllReduce,[M, 7168]bf16, 4 ranks, a CUDA graph of20 back-to-back calls each the programmatic dependent of the previous, 40 replays; the slowest rank's median per
call; the previous kernels = G1's build, alternated with this build A B B A, nothing else on the GPUs):
Outputs equal the previous kernels' bit for bit in every case.
Part 4 (RB):
tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_m1.py,test_k3_moe_m2.py(one GPU,l0_b200):element, 4 ulp relative RMS) and against
trtllm::k3_moeon the same routing (1 ulp), at the TP16 shape and,for one token, the TP4 x EP4 shape;
k3_moe_m2's epochs across the int32 wrap;The tests run only the cells each engine serves, and skip none.
tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_k3_moe_m1.py,..._m2.py(one GPU,l0_b200), throughthe catalog wrappers:
create()refusing capture and compiling every listed build.Every call is checked bit for bit against the same call on a reference state, and the next call's count words
are checked after each call.
tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_moe_push.py(4 ranks,l0_gb200_multi_gpus): each rankholds its own TP16 experts and runs the same tokens.
k3_moe_m1at 1 and 2 tokens,k3_moe_m2at 2.K3MoeLayer.push(into the run's exchange) and themoe/k3_moeentry'sk3_moe_push(into the16-slot one) at 1, 3 and 8 tokens, after
trtllm::k3_route_quantand aftertrtllm::k3_moe_front(whose sharedactivation must be the same bits in every call).
leave the exchange empty with its count advanced.
with the exchange's int32 count wrapping.
call, plus a captured step. The negative control has one rank swap two push + reduce pairs: every rank's sums
are wrong, nothing raises or hangs, and the next pair is right.
tests/unittest/_torch/cute_dsl_kernels/kimi_k3/test_k3_tcgen05_fences.py: the two kernels' sources join thetcgen05 fence lint. On the kernels before their fences, both rows fail.
Results on GB200 (sm_100), RB's head
f0fd94cfbdon part 3's build (tray 7633130):test_k3_moe_m1.pyandtest_k3_moe_m2.py7 + 7 passed; the catalog tests 7 + 7, and 7 formoe/k3_moe's push form;test_k3_moe_push.pyat 4 ranks: push 24 / 24, k3_moe_push 24 / 24, sequences 3 / 3; the fence lint 10 / 10.Part 3's tests are unchanged with part 4 in (
moe/test_modeling_v2_k3_moe.py62,test_k3_fused_moe.py44,test_k3_moe_wide.py+test_k3_route_quant.py663; none skip).Part 5 (GB200, main's build plus this branch's Python; CPU-level checks, one tray):
test_modeling_v2_routing.py: route B matches with the split set, inautoandrequire, from atext_configobject and from a dict, with the checkpoint's quantization as the engine reads it (W4A16_MXFP4); both K3 targets
register and count as external; near misses on
parallel(the unset split, an explicit EP16 split, attention DP onroute B's split, pipeline parallelism).
test_modeling_v2_kimi_k3_drift.py(import-free): every module oftp16_moetp4ep4has a copy here, and each copyequals its original outside the route B blocks. A control checks that the comparison flags a change outside the
blocks and only there.
test_modeling_v2_kimi_k3_construction.py(host-side): each target's construction checks pass on its own layoutand fail on the other split and on attention DP, naming the target; route B also fails on an unset split and on a
speculative config.
test_modeling_v2_claims.py,test_modeling_v2_no_stale_claims.pyandtest_modeling_v2_target_contract.pycover the new target through the routing tables; [None][feat] modeling_v2: Kimi K3 MXFP4 target, sm_100, tp16 attention + moe tp4 x ep4 #19839's decode step and decode GEMV tests cover route B's
identical copies of that code.
fa24aca986: claims 23, routing 43, no_stale_claims 1, kimi_k3_drift 5,target_contract 8, kimi_k3_construction 4, kimi_k3_decode_step 22, kimi_k3_decode_gemv 23; all pass, none skip.
Part 6 (GB200, one tray, 4 GPUs; a build of this branch's C++ plus its Python):
comm/test_modeling_v2_kimi_k3_decode_comm_op_matrix.py(l0_gb200_multi_gpus.yml) starts_kimi_k3_decode_comm_op_matrix.pyat 4 ranks.keep the interface the layer calls, real workspaces and the real MNNVL all-reduce.
post_load_weightsruns: the MoE decode pathbuilt from a MoE layer's own modules (the target's gate, whose parameters require grad, the nn.Linear latent
projections, the stock RMSNorm and a shared GatedMLP at TP 4; zeroed routed experts in the TRTLLM-Gen layout),
every kernel warmed up, then a 3-token step and a wide 16-token step against the shared expert's forward
(within 2.8e-3 and bit for bit), ranks bitwise equal;
unclassified tokens, a wide 2 x 8 step).
updatedis bit for bit in 24 / 24 cases;normedis within 7.1e-3(relative to the largest value), and so is the path each call took;
comm/k3_sandwich_tail, with and without a snapshot) vs the tail reduced in torch:within 8.3e-3, and its capture and replay bit for bit;
mpirun): passed in 29 s;srun: 7 / 7 checks on every rank at 4 ranks and at 2.l0_gb300_multi_gpus.ymlcollectsmodeling_v2/comm);with one GPU visible it reports 1 skipped.
test_modeling_v2_kimi_k3_decode_step.py: the attn_res epilogue ceilings per step kind. 36 passed together withtest_modeling_v2_target_contract.py, whoseREQUIRED_TRTLLM_OPSnow names the C3 ops (both targets).test_modeling_v2_kimi_k3_decode_gemv.py: the MoE sites (fp32 outputs, wide row ceilings). 28 passed.test_modeling_v2_kimi_k3_drift.pytests: 74 passed (drift 7 / 7: routeB's copy carries these changes).
Part 7 (GB200; this branch's C++ build plus its Python):
tests/unittest/_torch/speculative/hw_agnostic/test_build_draft_model_hook.py(cpu_only, l0_cpu): the default_build_draft_modelcallsget_draft_model(model_config, draft_config, lm_head, model);__init__builds thedrafter through a subclass override, once, with the resolved draft config.
tests/unittest/_torch/modeling_v2/test_modeling_v2_kimi_k3_drafter.py(1 GPU, sm_100, l0_b200), at one TP16rank's drafter shapes, two layers:
DFlashForCausalLM.dflash_forwardon the samemodule and context (rel. L2 ≤ 1e-2), with the torch GEMMs and with the GEMV sites (taken at ≤ 8 rows, declined
above), and leaves the cache untouched;
test_modeling_v2_kimi_k3_decode_gemv.py: the four drafter sites, the SiLU-and-mul site against the float64product of torch's bf16
silu(gate) * up.The merged tree (this branch's head on a build of its own C++, GB200, one tray of 4 GPUs):
k3_route_quantentry and part 4'sk3_moe_m1kernel test: 98;k3_moeentry and part 4'sk3_moe_m1entry: 114;k3_moe_m2kernel and entry: 59.srun, every rank passing: part 1's matrix (9 checks);test_k3_mnnvl_comm.py(attn_res,split all-gather); part 3's fusion all-reduce (10), plain sandwich (9) and MoE front (10) matrices; part 4's
test_k3_moe_push.py(push, k3_moe_push, sequences).mpirun, 4 ranks): 10 passed(allgather, reducescatter, part 1's attn_res matrix, part 3's seven stateful matrices and the MoE front).
P11's commit (#19842; on the same build):
test_gqa_dspark_keeps_its_trained_mask_row_over_the_target_embeddingpasses, and fails on this branch without the commit.
test_kimi_k3_dspark_semantics.pyas a whole: 47 of 48pass. The one failure,
test_mla_dspark_auto_backend_resolves_to_cutedsl, expects the MLA drafter's AUTObackend to be CUTEDSL while the drafter defaults to TRTLLM; that test and that default are main's, unchanged here.
test_dflash_dummy_slot.py,test_dspark_worker.pyandtest_dspark_draft.pypass.Test and CI-list follow-ups, and route B's carry of section 7 (GB200, this branch's tree unless noted):
new rows).
test_kimi_k3_situ_moe.py: 85 passed and 1 xfailed (an NVFP4 case, xfailed before too); itsTP-sharded check now also runs at TP16.
test_modeling_v2_kimi_k3_drift.py7 / 7 (5 / 7 before: section 7 had changedtp16_moetp4ep4'sdecode_gemv.pyandmodeling.pyonly). The top-level ModelingV2 tests (claims,no_stale_claims, routing, target-contract, drift, construction, decode step): 114 passed. Route B's
decode_gemv.pyis byte-identical totp16_moetp4ep4's again;test_modeling_v2_kimi_k3_decode_gemv.pypasses32 / 32 against either copy.
l0_b200.ymlselects them (-m "not cpu_only"): the Kimi K3 routing tests (-k "kimi_k3") 21 passed and 22deselected, the claims stock-import check 1, construction 4, drift 7.
test_k3_mla_decode_view.py: 29 passed.mpirun, 4 ranks), on main with every Kimi K3 PR applied andthese test files byte for byte: 11 passed. The two generic MNNVL op matrices, now sm_100 only, run there and are
not skipped.
tests, and no listed path is missing.
test_k3_mla_decode_view.pyis cpu_only throughout, so itsl0_b200.ymlentry selected nothing and pytestexited 5; it is now in
l0_cpu.ymlinstead.entry,
l0_b300.yml's modeling_v2 one, which is waived; they now havel0_b200.ymlentries.Manual 16-rank reference
On 16 GB200 GPUs (4 trays in one NVLink domain, fabric handles); no CI stage has 16 ranks.
srun -N4 -n16 --ntasks-per-node=4 --mpi=pmix python3 _mnnvl_allreduce_attn_res_op_matrix.py --launcher srun --world-size 169926793012runs/session-7631980/pre/g1/matrix16.logsrun -N4 -n16 --ntasks-per-node=4 --mpi=pmix python3 -u _torch/modeling_v2/comm/_<entry>_op_matrix.py --launcher srun --world-size 16(intests/unittest)db8ca58bd0runs/session-7634394/pre/g4/<entry>.logmain(), world 16srun -N4 -n16 ... python3 _torch/cute_dsl_kernels/kimi_k3/test_k3_{sandwich,moe_front,latent_reduce,mnnvl_comm}.pydb8ca58bd0runs/session-7634394/pre/g4.logk3_moe_m1M 1 / 2,k3_moe_m2,K3MoeLayer.push/k3_moe_push, sequences)srun -N4 -n16 --ntasks-per-node=4 --mpi=pmix python3 -u test_k3_moe_push.py <check>, = push, k3_moe_push, sequencesf0fd94cfbdlogs/up-session-rb16-7634700.log,runs/rb16-7634700/rb/{push,k3_moe_push,sequences}.logTRTLLM_MODELING_V2=require, no speculation, max_batch_size 8, 8 fixed prompts; on main's C++ build (without #19830's__syncwarpin the legacy KDA decode kernel, which mixed steps run), with main's built-in model at the same layout (off) as the controlnospec16@8(moe_tensor_parallel_size: 16,moe_expert_parallel_size: 1)fa24aca986runs/session-7631980/09-nospec16_8_u5b-main-fa24aca986-require_fixed/serve.log(control:06-nospec16_8_u5b-main-fa24aca986-off_fixed)moe_tp=16, moe_ep=1; 8 / 8 fixed prompts coherent (128 tokens each). 1 / 8 texts equal the control's; the other 7 diverge after 12-150 characters, as the decode kernels round differently from main's built-in path. Accuracy is gated on GSM8K (session 2).requirevsoff(each >= 96.5 - tol); main's C++ build as aboveeval_nospec16@8(1319 samples, max_batch_size 8)fa24aca986runs/session-7631980/13-eval_nospec16_8_u5b-main-fa24aca986-require_gsm8kfull/eval-gsm8kfull.log,runs/session-7634394/01-eval_nospec16_8_u5b-main-fa24aca986-off_gsm8kfull/eval-gsm8kfull.logTRTLLM_MODELING_V2=require, no speculation, max_batch_size 8, 8 fixed greedy prompts one at a time; this branch's C++ buildnospec@8ba53c84359runs/session-7636347/01-nospec_8_c4-k3rest-ba53c84359-require_fixedrequire: >= 96.5 - toleval_nospec@8ba53c84359runs/session-7636347/02-eval_nospec_8_c4-k3rest-ba53c84359-require_gsm8kfullbenchmark_servingat concurrency 8, ISL / OSL 1024 (wide steps of 9-64 tokens)nat@8ba53c84359runs/session-7636347/03-nat_8_c4-k3rest-ba53c84359-require_fixed_bench8require: >= 96.5 - toleval_dspark@8ba53c84359runs/session-7636347/06-eval_dspark_8_c4-k3rest-ba53c84359-require_gsm8kfullrequire, no speculation, max_batch_size 8, the fixed promptsnospec16@8ba53c84359runs/session-7636347/05-nospec16_8_c4-k3rest-ba53c84359-require_fixedrequire: >= 96.5 - toleval_nospec16@8ba53c84359runs/session-7636347/08-eval_nospec16_8_c4-k3rest-ba53c84359-require_gsm8kfullnat@8%c4-k3rest-0d16d09861-require:fixed,bench8(arm 4) vs C3'snat@8%c4-k3rest-ba53c84359-require:fixed,bench8(arm 3)runs/session-7636347/04-nat_8_c4-k3rest-0d16d09861-require_fixed_bench8/serve.logrequire); acceptance lengtheval_dspark@8%c4-k3rest-0d16d09861-require:gsm8kfull(arm 7) vs C3's (arm 6)runs/session-7636347/07-eval_dspark_8_c4-k3rest-0d16d09861-require_gsm8kfull/eval-gsm8kfull.lognat@8%c4-k3rest-0d16d09861-require:bench1,bench8(arm 10) vs C3's (arm 11), read with arms 4 / 3runs/session-7636347/10-nat_8_c4-k3rest-0d16d09861-require_bench1_bench8/serve.logeval_dspark@1%c4-k3rest-0d16d09861-require:gsm8k200(arm 12) vs C3's (arm 13)runs/session-7636347/12-eval_dspark_1_c4-k3rest-0d16d09861-require_gsm8k200/eval-gsm8k200.logsite_check.py: every TP16 rank's slice x 5 layers of the drafter checkpoint, 1..8 rows, vs torch bf16 and float64runs/tokens/c4-sitecheck-7638408/site_check.logdrafter_obit-identical to torch; the others differ in 0.03-0.12 % of outputs by 1 ulp of accumulation orderrequire: the same text as before the carryeval_nospec16@88806ae3a6b's route B tree (431f05db21), on main with every Kimi K3 PR applied (C++ and Python)runs/session-7640089/10-eval_nospec16_8_k3main-57d7238926-require_gsm8kfull/eval-gsm8kfull.logmoe_tp=16, moe_ep=1Part 7's rows pair C4's arms on this head (
c4-k3rest-0d16d09861-require; the tree also carries a count hook for thedrafter's entries, outside this PR) with C3's arms on the previous head (
c4-k3rest-ba53c84359-require, where thetarget builds the stock drafter), in the same session. The reference is that pair, not
off:offalso swaps thetarget for the built-in model, so a gap would mix the target and the drafter. A part 7 row passes when:
the drafter GEMV sites" line is on every rank of C4's serve.log and absent from C3's, and the count hook sees blocks
on the entries,
k3_drafter_attn_qknormand the fourdrafter_*sites on every rank, eager and at capture. Ifevery block fell back to the stock forward, texts, acceptance and bench would all equal C3's and the check would
false-pass;
Texts are evidence only: a same-session pair can match 8 / 8, but up to 3 / 8 differ across launches.
PR Checklist
Please review the following before submitting your PR:
PR description clearly explains what and why. If using CodeRabbit's summary, please make sure it makes sense.
PR Follows TRT-LLM CODING GUIDELINES to the best of your knowledge.
Test cases are provided for new code paths (see test instructions)
If PR introduces API changes, an appropriate PR label is added - either
api-compatibleorapi-breaking. Forapi-breaking, includeBREAKINGin the PR title. (api-compatible:trtllm::mnnvl_fusion_allreducegains an optional trailing argument with main's default and declaresbuffer_flagsmutable;MNNVLAllReduce.forwardgains an optional keyword. New ops, catalog entries and a new ModelingV2 target otherwise.)Any new dependencies have been scanned for license and vulnerabilities (No new dependencies.)
CODEOWNERS updated if ownership changes (No ownership change.)
Documentation updated as needed (The catalog's
index.yaml,README.mdand the contracts state the conventions and each entry.)Update tava architecture diagram if there is a significant design change in PR. (No design change outside ModelingV2.)
The reviewers assigned automatically/manually are appropriate for the PR. (To check once the PR is open.)
Please check this after reviewing the above items as appropriate for this PR.
GitHub Bot Help
To see a list of available CI bot commands, please comment
/bot help.