Skip to content

[None][feat] Kimi K3 on modeling_v2: stateful catalog entries, MNNVL and collective decode kernels, drafter attention, MoE route B - #19841

Draft
vsabavat wants to merge 178 commits into
NVIDIA:mainfrom
vsabavat:k3-up/k3-rest
Draft

vsabavat wants to merge 178 commits into
NVIDIA:mainfrom
vsabavat:k3-up/k3-rest

Conversation

@vsabavat

@vsabavat vsabavat commented Oct 3, 2026 •

Copy link
Copy Markdown

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_res

ModelingV2's catalog had no convention for an op whose result depends on state that outlives the call.

Conventions (catalog/index.yaml, README.md):

  • The state is a typed object the caller owns and passes in. An explicit constructor builds it: collective
    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.
  • The contract gains a ## State section: the object's contents and size; who creates it, and when; which ops
    may share one object; the call order every rank keeps across layers and steps; what a later launch reads; how it
    is re-armed.
  • The test drives call sequences (chained calls across steps, capture and replay, two objects interleaved) and a
    negative control that breaks the call order, not only single calls.
  • Receipts: an entry first called by a GB200 (sm_100) target is certified on sm_100, and a collective's receipt
    records the world size its matrix ran at.
  • Launchers: the collective matrices' launcher takes --world-size and --launcher srun, so the same rank body
    records a multi-node run. _rank_job.run takes the world size.

trtllm::mnnvl_allreduce_attn_res (mnnvlAllreduceKernels.cu/.h, thop/allreduceOp.cpp)

  • A one-shot MNNVL all-reduce whose epilogue is Kimi K3's residual update: updated = prefix_sum + the sum over the ranks, and normed = RMSNorm of the attention-residual selection over the snapshot bank and updated.
  • It rounds like the unfused all-reduce followed by attn_res_add_rmsnorm_fwd. Its reduction order is fixed, so
    every rank gets the same bits.
  • One cluster of hidden / 1024 CTAs owns a token. The per-token statistics are summed in a fixed order, the
    cluster part through distributed shared memory. Each thread arrives on the cluster barrier early and waits before
    its first write into a peer CTA's shared memory.
  • reduceOneshotLamport is the one-shot kernel's Lamport reduction as a function, for the new kernel. Part 1 leaves
    the 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 3
    then moves the one-shot kernel onto the same function and adds its early trigger.
  • The schema declares both comm_buffer and buffer_flags mutable.
  • MNNVLAllReduce.allreduce_attn_res_rmsnorm calls the op on the module's own workspace.

comm/mnnvl_allreduce_attn_res

  • The catalog entry for the op, over a caller-owned MnnvlWorkspace (comm/mnnvl_workspace.py): three Lamport
    buffers, the flag words and the multicast handle.
  • MnnvlWorkspace.create is collective over the TP group only: its communicator comes from the new
    _get_mnnvl_tp_group_comm (distributed/ops.py, MPI_Comm_create_group over mapping.tp_group). Before
    allocating, the ranks agree that each of them can; if one cannot, every rank raises and frees that communicator.
  • The contract's ## State section states the call-order rule. A swapped pair of same-shaped calls is silently
    wrong 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_attn and trtllm::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:

  • Up to 8 requests each bring a draft block of up to 8 tokens.
  • Each block attends densely (non-causal) to its request's cached context in the paged KV cache (HND pages of 64)
    and to the block's own K / V. Those are read from the projection output and not written to the cache.
  • GQA with 6 query heads per KV head, head dim 64.
  • k3_drafter_attn_qknorm takes the raw projection output and applies the per-head q / k RMSNorm and the NeoX RoPE
    inside the kernel, with fused_qk_norm_rope's arithmetic.
  • The kernel compiles on its first call for its head counts, page stride and whether more than one request is
    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.
  • For this drafter's shapes it does the work of DFlash's TRTLLM path (append the block's K / V, then the trtllm-gen
    context kernel) in one kernel, without writing the cache. The op's test compares the two.

Catalog entries (modeling_v2/catalog/attention/):

  • k3_drafter_attn and k3_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_103
    receipt 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, new mnnvlAllGatherKernels.{h,cu},
distributed/ops.py, custom_ops/cpp_custom_ops.py):

  • trtllm::mnnvl_fusion_allreduce takes one_shot_max_bytes: messages up to that size go one-shot, larger ones
    two-shot. The default stays main's 1 MiB. MNNVLAllReduce keeps 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).
  • A behaviour change for every one-shot caller, main's MNNVLAllReduce included. The one-shot fusion kernel now
    releases 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, the
    order of the reduction does not. Evidence, on GB200 (4 ranks): main's multi_gpu/test_mnnvl_allreduce.py
    bodies pass on this build (110 cases at 4 ranks, 109 at 2, the graph-capture cases); main's MNNVLAllReduce at
    [M, 7168] bf16, M 1 to 2048, plain and RMSNorm-fused, one-shot and two-shot, returns the previous kernels' outputs
    bit 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 earlyTrigger field; only the one-shot kernel's source
    changes.
  • trtllm::mnnvl_fusion_allreduce now declares buffer_flags mutable (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 as
    bf16, 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_buffer and buffer_flags mutable, it checks world_size against the
    workspace, and it has a fake implementation.

Kimi K3 collective kernels (tensorrt_llm/_torch/cute_dsl_kernels/, CuTe DSL, sm_100). The kernel files are the
ones 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, _tail and _plain, a row-parallel projection, its TP all-reduce
    and 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) ships
    with 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 with
    the 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 to
    trtllm::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) and K3LatentExchange (the latent
    exchange) are built by an eager create(mapping, fabric_handle=None) that is collective over the TP group's ranks
    only (its communicator is made from them alone, as part 1's MnnvlWorkspace.create does), with the same failure
    model: 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 the
    communicator 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 / K3MoeWideState hold a device's k3_moe scratch (up to 8 / 64 tokens) and K3MoeLayer a layer's
    experts and counters; their constructors refuse capture, and the state is per rank.
  • Every op's mutates_args names the buffers it writes, trtllm::k3_moe included (the state's scratch, the layer's
    counters, and with a head_flags state the head workspace's ready words and epoch).
  • The state types sit beside the buffer layouts they allocate, in the kernel packages; each catalog wrapper
    re-exports the type it takes.

Catalog entries (_experimental/modeling_v2/catalog/), each a contract with a ## State section where the op
has 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=True releases the ready words a head_flags k3_moe
acquires), moe/k3_moe, and the stateless moe/k3_route_quant. comm/allgather, which the Kimi K3 LM head calls,
gains an sm_100 receipt; its sm_103 receipt 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-row
tiles 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 CuTe
    DSL 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 output
    has 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 distinct
    expert'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.
  • Push builds. With the latent all-reduce done by trtllm::k3_latent_reduce (part 3), each engine and
    trtllm::k3_moe can store this rank's partial into its slot of every rank's K3LatentExchange instead of
    returning it: K3MoeM1Layer.push, K3MoeM2Layer.push, and K3MoeLayer.push after trtllm::k3_route_quant or
    trtllm::k3_moe_front. A push followed by the reduce equals MNNVLAllReduce's one-shot of the plain partials bit
    for bit.
  • State. K3MoeM1State / K3MoeM2State own the small workspace every layer on one device shares: the
    intermediate 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 the
    plain and listed push builds before any CUDA-graph capture.
  • Catalog: moe/k3_moe_m1 and moe/k3_moe_m2, with ## State sections, 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 same
checkpoint, GPU and attention layout as #19839's kimi_k3_mxfp4__sm_100__tp16_moetp4ep4, with the routed experts
split 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_moetp4ep4 target serves speculative decoding.

Routing (models/kimi_k3_vl/routing.py): the K3 routing tree gains the second key. The expert split decides the
target:

configuration (SM 10.0, the K3 MXFP4 checkpoint, world 16, tp 16, pp 1, no attention DP) target
moe_tensor_parallel_size 4, moe_expert_parallel_size 4 tp16_moetp4ep4
moe_tensor_parallel_size 16, moe_expert_parallel_size 1, both set tp16_moetp16ep1
the split unset (the Mapping resolves it to 16 x 1, and the built-in model runs it as expert parallelism over 16 ranks) none
moe_tensor_parallel_size 1, moe_expert_parallel_size 16 none
attention data parallelism, pipeline parallelism, or another world size none
the NVFP4 requant (reads MIXED_PRECISION), an unquantized checkpoint, or SM 10.3 none

"None" means: auto runs the built-in model, and require raises with the routing trace. The trace's parallel
line prints split_set, so explain shows 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 own modeling.py, weights.py and decode_gemv.py,
copied from tp16_moetp4ep4's. Each place the copy differs is a block between # >>> route B: <why> and
# <<< route B:

  • the docstrings that name the layout: the module's, the loader's, the routed-expert split's;
  • the construction checks: the topology (world 16, tp 16, pp 1, moe_tp 16, moe_ep 1, no attention DP), the expert
    split set explicitly, and no speculative decoding config (the message names the 4 x 4 split for DSpark, DFlash and
    SA);
  • the registered name and class, ModelingV2KimiK3Mxfp4Sm100Tp16Moetp16ep1;
  • the speculative path: no SA / DFlash / DSpark admission, no drafter LM head hand-off, no per-token KDA verify states
    (kda_token_states) and no KDA verify kernels (k3_kda_attn, k3_kda_verify, trtllm::kda_mtp_decode), with
    REQUIRED_TRTLLM_OPS and UNCERTIFIED_GENERIC_CALLS to match.

decode_gemv.py is byte-identical, 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_model override and the drafter's four
decode GEMV sites) and is never built here; K3DecodeGemvs.create warms the four sites once on zero weights at
startup, as it does on tp16_moetp4ep4.

test_modeling_v2_kimi_k3_drift.py reads both targets' files, imports neither, and fails on any difference outside
the blocks, so a change to tp16_moetp4ep4 reaches 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_moetp4ep4 target

#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's MnnvlWorkspace (4 MiB buffers) and K3SandwichWorkspace.
    The target creates both collectively in post_load_weights, before any CUDA-graph capture, wherever every
    attention all-reduce runs over MNNVL.

    • The post-attention step. On a step of at most 16 tokens (wide decode steps aside), the attention hands over
      its unreduced o_proj output. comm/mnnvl_allreduce_attn_res then runs the all-reduce with the residual update as
      its epilogue.
      • At most 8 tokens on the attention's decode branch, o_proj, the all-reduce and the update run as one
        comm/k3_sandwich_oproj kernel, bit for bit o_proj followed by the MNNVL entry.
    • The pre-attention step after a MoE layer. It runs the MoE layer's handed-on tail: comm/k3_sandwich_tail
      runs 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.
      • On a snapshot layer, the kernel also stores the prefix sum into the bank row that layer takes.
  • 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_moe on a K3MoeState, then the routed experts' all-reduce of the latent;
    • the row-parallel tail, [RMSNorm(latent) slice | shared activation] @ [latent up columns | shared down].
      • Where the next layer's pre-attention step, or the final norm, reduces it, the tail goes on unreduced as a
        PendingTail.
      • Elsewhere the replicated tail runs.

    A wide decode step (9-64 tokens) keeps the sharded head (comm/mnnvl_allgather_split) and the row-parallel tail on
    M-general ops: moe/k3_route_quant and moe/k3_moe on a K3MoeWideState. The consumer reduces its unreduced
    output with a plain all-reduce and then runs the fused add + attn_res + RMSNorm.

  • The decode step's numerics.

    • A classified step's attn_res epilogues take the fused kernels up to 8 tokens, or 32 on a wide step.
    • The stock MNNVL all-reduces of the decode path send one-shot up to 4 MiB; a wide step's all-reduces use 1 MiB.
    • Where the MoE decode path takes a layer, the latent norm's weight is folded into the latent up projection. It is
      the same function, rounded differently.
  • The attention's interface. will_run_decode_branch(attn_metadata, step) tells the layer, before the attention
    runs, whether the attention will take its decode branch. reduce_output=False returns the attention's output
    unreduced, on every path. project_output=False returns it before o_proj: on the decode branch for MLA, on every
    path for KDA.

  • Logging. The startup line gives the number of layers that comm/k3_sandwich_oproj and 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 stock GQADSparkForCausalLM with
a decode block on part 2's entries and the decode GEMV sites. The worker (DSparkWorker, through get_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-engine
    shell calls once to build its drafter. Its default is the mode registry's get_draft_model with the same
    arguments, 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:

    1. q / k / v on the drafter_qkv site, [512, 7168];
    2. attention/k3_drafter_attn_qknorm: the q / k RMSNorm, the NeoX RoPE and the block attending to its paged context
      and its own k / v, without storing the block's k / v;
    3. the output projection on drafter_o ([7168, 384], k3_ctm_gemv), then the module's all-reduce;
    4. gate / up on drafter_gate_up ([1792, 7168]), and down with the SiLU-and-mul on drafter_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_qknorm does not certify
    (DRAFTER_ATTN_SPLITS mirrors its contract), another attention backend, head layout, RoPE base, epsilon or MLP, a
    cache 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 new
    ops are in REQUIRED_TRTLLM_OPS.

  • R x 7: DSpark's draft block under shift_label is max_draft_len tokens, 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_moe compiles on routing that requires grad. A load-time warm-up outside inference mode routes with
    e_score_correction_bias, an nn.Parameter, and the first (compiling) call raised BufferError: Can't export tensors that require gradient from DLPack. _view now exports a detached view of the same storage, with no copy.
    test_routing_from_a_parameter_that_requires_grad fails before the fix and passes after.
  • Part 4's contracts record the push form's 16-rank runs (moe/k3_moe, moe/k3_moe_m1, moe/k3_moe_m2); part 3's
    eight contracts record theirs.
  • DSpark decode kernels: k3_spec_accept (acceptance with vocabulary-sharded target logits), k3_ctx_kv and
    k3_markov, one launch each where the worker ran 47, 52 and 119 small torch kernels. The worker takes them behind
    its Kimi K3 gate (k3_decode). Tests are listed in l0_b200.yml and l0_gb200_multi_gpus.yml; the sharded and
    Markov tests run under pytest on 4 local MPI ranks.
  • DFlash metadata: capture_view, a captured layer's slot of the capture buffer. The draft block's rows and the
    gen requests' page-table rows are now views where the slots are one run.
  • The DSpark drafter (section 7): its fc is split over the TP group. Its residual adds and RMSNorms run in its
    all-reduces (comm/k3_sandwich_plain, or the fused MNNVL all-reduce above the plain path's rows). The DSpark tap
    is written by the sandwich tail.
  • The MoE decode path's latent all-reduce on a pure decode step captured into a CUDA graph is now moe/k3_moe's
    push form plus comm/k3_latent_reduce. It sums in the one-shot's order, so the bits are the same: the 8 fixed-prompt
    outputs are identical with and without it.
  • Route B's decode MoE on part 4's engines: k3_moe_m1 at one token, k3_moe_m2 at two, k3_moe over all 896
    experts 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.
  • The generic MoE path: the backend's fused route + MXFP8 quant runs as trtllm::k3_route_quant, bit for bit
    the C++ op's outputs. Both targets' gates return KimiK3MoeRoutingMethod, so the generic TRTLLM-Gen MoE routes
    outside the kernel.
  • Route B's copy carries every route A change outside its blocks; the drift test checks it.
  • Tests: the kernel tests' pool workers are picklable under torch 2.14. The decode-comm matrix builds route A's
    shared-expert slice at every world.
  • Fix, decode MoE warm-up (found by the 16-rank decode-comm matrix). The warm-up 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, 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.
    • 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 compiled hung in 3 of 4 runs at 16 ranks.
    • Both targets' warm-ups now synchronize the device before the push build, and moe/k3_moe states the rule.
    • Decode steps always have a kernel that waits before it triggers between two calls on one layer: the MNNVL
      all-reduce, k3_latent_reduce, or a graph boundary.
    • After the fix: the matrix passes 3 / 3 at 16 ranks, and 100 back-to-back warm-ups in one process pass.

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

part ask choice
1, 3, 4 Stateful ops A typed state object per stateful op (workspace, Lamport buffers, counters), built by an explicit, collective, eager create() that the caller (the target's post_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 ## State section (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_args names every written buffer. Here: MnnvlWorkspace (part 1), K3SandwichWorkspace, K3LatentExchange, K3MoeHeadWorkspace (collective create(), over the TP group's ranks only); K3MoeState / K3MoeWideState and K3MoeLayer (per rank, constructors).
1, 3, 4 Multi-rank harness _rank_job takes the world size, and the stateful matrices take --world-size and --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.
1-4 sm_100 The Kimi K3 entries carry sm_100 receipts and skip on other architectures. A reused entry gains its sm_100 cell in the PR of its first Kimi K3 caller and keeps its sm_103 receipt (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 in l0_b200.yml, 4-rank tests in l0_gb200_multi_gpus.yml.
1, 3 MNNVL entries One entry per op (part 1's comm/mnnvl_allreduce_attn_res, part 3's comm/mnnvl_fusion_allreduce and comm/mnnvl_allgather_split). One caller-owned MnnvlWorkspace (comm_buffer, buffer_flags, multicast handle) is shared by all three; one_shot_max_bytes is per call.
2 DSpark shell The target builds the upstream DFlash / DSpark worker through get_spec_worker and the target_logits hook (#19814). The worker's own kernels stay upstream code. The drafter model is target code, so its ops are catalog entries.
2 Naming Categories: the drafter entries are attention/; there is no spec/ category.
all Layout The merged modeling_v2/README.md governs layout and records.
4 Torch ops over caller-owned state Both engines are torch ops, as trtllm::k3_moe is, so their writes are declared: mutates_args names the workspace, the exchange's words and out. K3MoeM1State / K3MoeM2State are built by an explicit, eager create() that the target runs in post_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).
5 Two layouts of one checkpoint Two targets: each rank holds a different shard of the routed experts (both would need ~211 GB per rank). Route B's modules are a copy of route A's, changed only in marked blocks and kept equal everywhere else by a CPU drift test, because targets share no files.
5 The alternative: a declared base Not taken here: the family's routing module would name tp16_moetp4ep4 as 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.
5 The unset expert split It resolves to the same sizes as 16 x 1, but the built-in Kimi K3 model reads that default as expert parallelism (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.
5 Speculative decoding on route B None. Routing may not read the speculative config (explain cannot 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.
5 Padded experts Each rank's 192-wide expert slice is padded to 256 when it loads, a multiple of the TRTLLM-Gen MoE kernels' 128-value alignment.
6 State K3DecodeComm, K3DecodeMoe and K3DecodeMoeLayer are the target's caller-owned state, built once in post_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.
6 Call order across ranks Which path a step takes is decided from its token count and kind and from the load-time layout only, so every rank makes the same calls on each workspace in the same order.
6 When the sandwich runs The layer asks will_run_decode_branch before calling the attention. MLA's KV-cache writes allow one attention run per step, so there is no fallback after the fact.
6 The latent all-reduce The routed experts' stock one-shot all-reduce. moe/k3_moe's push form (part 4) into a K3LatentExchange plus comm/k3_latent_reduce sums in the same order; moving to it is a separate performance change.
6 The MoE tail Row-parallel and handed to its consumer, so the tail GEMV, its all-reduce and the next residual update are one kernel. Where nothing consumes it that way, the replicated tail keeps the layer's output reduced.
7 A target that owns its drafter One engine change, outside ModelingV2: the one-engine shell builds its drafter through SpecDecOneEngineForCausalLM._build_draft_model(), whose default is get_draft_model with the same arguments, so every existing model is unchanged. The Kimi K3 target overrides it to build K3DSparkDrafter instead 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_res matrix (_mnnvl_allreduce_attn_res_op_matrix.py, 4 ranks): 9 / 9 checks
    pass, 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:
    • single calls: 1-8 and 16 tokens x 0, 1, 5, 8 and 11 snapshots, with and without a prefix (the sequences below
      cover every snapshot count 0-11);
    • a call over one buffer raises on every rank, and the next call is correct;
    • layers x steps with the token count dipping and growing back and a random rank late;
    • two workspaces interleaved, and CUDA-graph capture and replay mixed with eager calls;
    • the negative control: one rank swaps two calls, and every rank gets a wrong answer without an error;
    • MnnvlWorkspace.create with one rank short of memory: every rank raises and frees its communicator, and the
      next call is correct;
    • MnnvlWorkspace.create under TP 2 x PP 2: one TP group creates a workspace while the other group's ranks make
      no MNNVL call.
    • updated equals the exact sum bit for bit; normed is within 7.8e-3 (relative to the largest value) of the
      fp32 reference; every output is bitwise equal on all ranks.
  • test_k3_mnnvl_comm.py (attn_res), 4 ranks: MNNVLAllReduce.allreduce_attn_res_rmsnorm against the unfused
    path (the default MNNVL all-reduce, then attn_res_add_rmsnorm_fwd): PASS.
  • Part 1 alone leaves the existing kernels unchanged: on a build of part 1's head, the SASS of every instantiation
    of oneshotAllreduceFusionKernel, twoshotAllreduceKernel and rmsNormLamport (276 functions) is identical to
    that 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 2
    ranks. 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 to
    5x1, 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 mixed
    context 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 (the
    pool'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.py and ..._qknorm.py: the entries' cell-by-cell tests against fp64
    references written in the tests, the R x 7 splits included. 36 and 16 passed.
  • test_modeling_v2_fused_qk_norm_rope.py gains test_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 the
collected 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 result
test_k3_fused_moe.py 44 passed
test_k3_moe_wide.py 643 passed, including 256 cases of a seeded skewed routing with this rank as its busiest EP group
test_k3_route_quant.py, test_k3_moe_front_geometry.py, the claims / routing tests, the two kernel lints 95 passed together with test_modeling_v2_k3_route_quant.py
moe/test_modeling_v2_k3_moe.py 62 passed (the receipt)
moe/test_modeling_v2_k3_route_quant.py 18 passed (the receipt)
the 8 collected op matrices (comm/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 starting mpirun with 4 ranks 8 / 8 passed
the same rank bodies, 4 ranks started by srun checks per rank: oproj 11, tail 10, plain 9, fusion all-reduce 10, split all-gather 10, latent reduce 13, MoE front 10
test_k3_sandwich.py, test_k3_moe_front.py, test_k3_latent_reduce.py, test_k3_mnnvl_comm.py (4 ranks, their main()) and their one-process tests pass
main's multi_gpu/test_mnnvl_allreduce.py bodies 110 / 110 at 4 ranks, 109 / 109 at 2, the checkpoint and workspace-growth graph cases pass

test_k3_moe_wide.py's former 256 router-dump cases needed a recorded dump and skipped without one; they are now
generated 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 of
20 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):

fusion M path previous (us) this PR (us) difference run-to-run noise
none 1 one-shot 3.74 3.67 -0.06 0.03
none 2 one-shot 3.79 3.72 -0.07 0.04
none 4 one-shot 3.84 3.81 -0.02 0.03
none 8 one-shot 4.00 4.03 +0.02 0.01
none 16 one-shot 4.63 4.55 -0.08 0.01
none 32 two-shot 5.75 5.75 +0.00 0.06
none 128 two-shot 8.03 8.14 +0.11 0.00
none 512 two-shot 21.48 21.53 +0.05 0.01
none 2048 two-shot 75.93 75.97 +0.03 0.03
RMSNorm 1 one-shot 5.46 5.50 +0.04 0.01
RMSNorm 2 one-shot 5.46 5.57 +0.12 0.02
RMSNorm 4 one-shot 5.58 5.68 +0.10 0.03
RMSNorm 8 one-shot 5.78 5.90 +0.12 0.03
RMSNorm 16 one-shot 6.71 6.46 -0.25 0.03
RMSNorm 32 two-shot 7.30 7.30 +0.00 0.01
RMSNorm 128 two-shot 9.67 9.61 -0.06 0.05
RMSNorm 512 two-shot 24.94 24.80 -0.14 0.02
RMSNorm 2048 two-shot 96.25 95.69 -0.56 0.10

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):

    • both engines against an fp32 reference over the dequantized checkpoint slices (8 ulp of the row's max per
      element, 4 ulp relative RMS) and against trtllm::k3_moe on the same routing (1 ulp), at the TP16 shape and,
      for one token, the TP4 x EP4 shape;
    • routing cases with all, 4 and no local experts; two tokens on the same and on disjoint experts;
    • two layers interleaved on one state; k3_moe_m2's epochs across the int32 wrap;
    • the two-token builds refuse the TP4 x EP4 intermediate (768).

    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), through
    the catalog wrappers:

    • single calls against the fp32 reference;
    • 12 steps of 3 layers on one state;
    • a captured step replayed 4 times with rewritten inputs and eager calls between replays;
    • two states in an irregular order;
    • the epochs across the int32 wrap;
    • 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 rank
    holds its own TP16 experts and runs the same tokens.

    • push: k3_moe_m1 at 1 and 2 tokens, k3_moe_m2 at 2.
    • k3_moe_push: K3MoeLayer.push (into the run's exchange) and the moe/k3_moe entry's k3_moe_push (into the
      16-slot one) at 1, 3 and 8 tokens, after trtllm::k3_route_quant and after trtllm::k3_moe_front (whose shared
      activation must be the same bits in every call).
    • Each push + reduce must equal the one-shot of the plain partials bit for bit, on every rank and run to run, and
      leave the exchange empty with its count advanced.
    • The same holds for a 16-slot exchange filled 4 slots per rank, against the one-shot's 16-slot order, and again
      with the exchange's int32 count wrapping.
    • sequences: the push form through the catalog wrapper, over 12 steps of 3 layers with a random rank late at each
      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 the
    tcgen05 fence lint. On the kernels before their fences, both rows fail.

Results on GB200 (sm_100), RB's head f0fd94cfbd on part 3's build (tray 7633130): test_k3_moe_m1.py and
test_k3_moe_m2.py 7 + 7 passed; the catalog tests 7 + 7, and 7 for moe/k3_moe's push form;
test_k3_moe_push.py at 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.py 62, test_k3_fused_moe.py 44,
test_k3_moe_wide.py + test_k3_route_quant.py 663; 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, in auto and require, from a text_config
    object 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 on
    route B's split, pipeline parallelism).
  • test_modeling_v2_kimi_k3_drift.py (import-free): every module of tp16_moetp4ep4 has a copy here, and each copy
    equals 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 layout
    and 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.py and test_modeling_v2_target_contract.py
    cover 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.
  • Results on the part's head 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.py at 4 ranks.
    • Setup: the target's own decoder layers at the TP16 per-rank shapes, with stand-in attention and MoE modules that
      keep the interface the layer calls, real workspaces and the real MNNVL all-reduce.
    • Checks:
      • first, outside inference mode with autograd on, as the target's post_load_weights runs: the MoE decode path
        built from a MoE layer's own modules (the target's gate, whose parameters require grad, the 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;
      • the o_proj shapes the sandwich takes;
      • the decode one-shot ceiling on every stock MNNVL all-reduce;
      • fused vs built-in on 8 step kinds (decode of 1 / 3 / 8 tokens, DSpark 1 x 6, a 5-token prefill, 12 and 16
        unclassified tokens, a wide 2 x 8 step). updated is bit for bit in 24 / 24 cases; normed is within 7.1e-3
        (relative to the largest value), and so is the path each call took;
      • the sandwich vs the MNNVL entry on exact payloads: bit for bit;
      • CUDA-graph capture and replay with rewritten inputs vs eager: bit for bit;
      • the deferred MoE tail (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;
      • every output bitwise equal across the ranks.
    • Results:
      • in its CI form (pytest starting mpirun): passed in 29 s;
      • started by srun: 7 / 7 checks on every rank at 4 ranks and at 2.
    • It skips off sm_100 and with fewer than 4 visible GPUs (l0_gb300_multi_gpus.yml collects modeling_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 with
    test_modeling_v2_target_contract.py, whose REQUIRED_TRTLLM_OPS now names the C3 ops (both targets).
  • test_modeling_v2_kimi_k3_decode_gemv.py: the MoE sites (fp32 outputs, wide row ceilings). 28 passed.
  • The claims, no_stale_claims, routing and test_modeling_v2_kimi_k3_drift.py tests: 74 passed (drift 7 / 7: route
    B's copy carries these changes).
  • The attention modules' C3 flags, on real cache managers (one GPU): passed.

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_model calls get_draft_model(model_config, draft_config, lm_head, model); __init__ builds the
    drafter 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 TP16
    rank's drafter shapes, two layers:
    • every certified split runs the entry once per layer and matches DFlashForCausalLM.dflash_forward on the same
      module 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;
    • an uncertified split runs the stock forward bit for bit;
    • under capture only compiled keys are taken;
    • a changed weight fails the comparison.
  • test_modeling_v2_kimi_k3_decode_gemv.py: the four drafter sites, the SiLU-and-mul site against the float64
    product 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):

  • One GPU per group of files, every test passing and none skipped:
    • the claims, routing, target-contract, route B drift and construction tests and the tcgen05 fence lint: 94;
    • part 2's op test, part 3's k3_route_quant entry and part 4's k3_moe_m1 kernel test: 98;
    • part 2's catalog tests, part 3's k3_moe entry and part 4's k3_moe_m1 entry: 114;
    • the Kimi K3 decode step and decode GEMV tests and part 4's k3_moe_m2 kernel and entry: 59.
  • 4 ranks started by 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).
  • The comm directory and the MoE front matrix as CI collects them (pytest starting 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_embedding
passes, and fails on this branch without the commit. test_kimi_k3_dspark_semantics.py as a whole: 47 of 48
pass. The one failure, test_mla_dspark_auto_backend_resolves_to_cutedsl, expects the MLA drafter's AUTO
backend 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.py and test_dspark_draft.py pass.

Test and CI-list follow-ups, and route B's carry of section 7 (GB200, this branch's tree unless noted):

  • The cluster-wait and tcgen05-fence lints with the KDA, MLA and drafter kernels' rows: 10 and 15 passed (6 and 5
    new rows). test_kimi_k3_situ_moe.py: 85 passed and 1 xfailed (an NVFP4 case, xfailed before too); its
    TP-sharded check now also runs at TP16.
  • Route B after it carries section 7: test_modeling_v2_kimi_k3_drift.py 7 / 7 (5 / 7 before: section 7 had changed
    tp16_moetp4ep4's decode_gemv.py and modeling.py only). The top-level ModelingV2 tests (claims,
    no_stale_claims, routing, target-contract, drift, construction, decode step): 114 passed. Route B's
    decode_gemv.py is byte-identical to tp16_moetp4ep4's again; test_modeling_v2_kimi_k3_decode_gemv.py passes
    32 / 32 against either copy.
  • As l0_b200.yml selects them (-m "not cpu_only"): the Kimi K3 routing tests (-k "kimi_k3") 21 passed and 22
    deselected, the claims stock-import check 1, construction 4, drift 7. test_k3_mla_decode_view.py: 29 passed.
  • The comm directory in its CI form (pytest starting mpirun, 4 ranks), on main with every Kimi K3 PR applied and
    these test files byte for byte: 11 passed. The two generic MNNVL op matrices, now sm_100 only, run there and are
    not skipped.
  • The CI lists, read against this branch: every new test is listed for its stage marker, no entry selects zero
    tests, and no listed path is missing.
    • test_k3_mla_decode_view.py is cpu_only throughout, so its l0_b200.yml entry selected nothing and pytest
      exited 5; it is now in l0_cpu.yml instead.
    • The construction and drift tests, the Kimi K3 routing tests and the claims stock-import check had a single
      entry, l0_b300.yml's modeling_v2 one, which is waived; they now have l0_b200.yml entries.

Manual 16-rank reference

On 16 GB200 GPUs (4 trays in one NVLink domain, fabric handles); no CI stage has 16 ranks.

part check command commit job log result
1 attn_res matrix, world 16 srun -N4 -n16 --ntasks-per-node=4 --mpi=pmix python3 _mnnvl_allreduce_attn_res_op_matrix.py --launcher srun --world-size 16 9926793012 7631980 runs/session-7631980/pre/g1/matrix16.log 9 / 9 checks passed on 16 / 16 ranks (the TP-group check at TP 8 x PP 2); max normed rel err 7.692e-3
3 the seven stateful matrices, world 16 srun -N4 -n16 --ntasks-per-node=4 --mpi=pmix python3 -u _torch/modeling_v2/comm/_<entry>_op_matrix.py --launcher srun --world-size 16 (in tests/unittest) db8ca58bd0 7634394 runs/session-7634394/pre/g4/<entry>.log every rank: k3_sandwich_oproj 11 / 11, k3_sandwich_tail 10 / 10, k3_sandwich_plain 9 / 9, mnnvl_fusion_allreduce 10 / 10 (max normed rel err 7.6e-3; one-shot vs two-shot 1.9e-3), mnnvl_allgather_split 10 / 10, k3_latent_reduce 13 / 13, k3_moe_front 10 / 10 (the TP16 half-tile head)
3 the four kernel tests' main(), world 16 srun -N4 -n16 ... python3 _torch/cute_dsl_kernels/kimi_k3/test_k3_{sandwich,moe_front,latent_reduce,mnnvl_comm}.py db8ca58bd0 7634394 runs/session-7634394/pre/g4.log PASS, 0 failing rows each
4 push builds at 16 ranks (k3_moe_m1 M 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, sequences f0fd94cfbd 7634700 logs/up-session-rb16-7634700.log, runs/rb16-7634700/rb/{push,k3_moe_push,sequences}.log pass, every rank: push 24 / 24, k3_moe_push 24 / 24, sequences 3 / 3
5 route B boot, TRTLLM_MODELING_V2=require, no speculation, max_batch_size 8, 8 fixed prompts; on main's C++ build (without #19830's __syncwarp in the legacy KDA decode kernel, which mixed steps run), with main's built-in model at the same layout (off) as the control nospec16@8 (moe_tensor_parallel_size: 16, moe_expert_parallel_size: 1) fa24aca986 7631980 runs/session-7631980/09-nospec16_8_u5b-main-fa24aca986-require_fixed/serve.log (control: 06-nospec16_8_u5b-main-fa24aca986-off_fixed) pass: healthy after 345 s (control 361 s); every rank logs route B's "KDA on k3_kda_decode_attn (69 / 69 layers take them), MLA on k3_mla_qkv and k3_mla_attn_vb_out (24 layers)" and 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).
5 route B GSM8K identity, require vs off (each >= 96.5 - tol); main's C++ build as above eval_nospec16@8 (1319 samples, max_batch_size 8) fa24aca986 require 7631980 (arm 13); off 7634394 (arm 1) runs/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.log pass: require 96.66 ± 0.49, off (main's built-in model, same layout) 96.59 ± 0.50; flexible = strict; R = 8 on 99.2 % of the require leg's 16,228 decode steps
6 route A boot with this section's commits, TRTLLM_MODELING_V2=require, no speculation, max_batch_size 8, 8 fixed greedy prompts one at a time; this branch's C++ build nospec@8 ba53c84359 7636347 runs/session-7636347/01-nospec_8_c4-k3rest-ba53c84359-require_fixed server up in 856 s (the session's first, cold boot); every rank: k3_sandwich_oproj on 93 / 93 layers, the MoE decode path on 92 / 92; texts identical to an exact relaunch (arm 9, 8 / 8); 3 / 8 identical to #19839's target (session 7631980), the others coherent. One text differs from its first token: a near-tie tipped by folding the latent norm into the up projection (the same function, rounded differently); route B, which does not fold it, keeps #19839's text
6 GSM8K, no speculation, batch 8 (decode steps of up to 8 tokens), require: >= 96.5 - tol eval_nospec@8 ba53c84359 7636347 runs/session-7636347/02-eval_nospec_8_c4-k3rest-ba53c84359-require_gsm8kfull 96.21 flexible / 96.13 strict over 1319 samples (#19839's target at the same settings: 96.59); 59.8 % of the responses identical to that run's (#19839's target vs main's built-in model: 59.4 %); paired vs C2: 11 vs 16 discordant of 1319, McNemar p 0.44
6 DSpark, the fixed prompts one at a time (steps of up to 8 tokens), then benchmark_serving at concurrency 8, ISL / OSL 1024 (wide steps of 9-64 tokens) nat@8 ba53c84359 7636347 runs/session-7636347/03-nat_8_c4-k3rest-ba53c84359-require_fixed_bench8 server up in 330 s; 1007 output tok/s, median TPOT 6.84 ms (P90 10.63 ms). Session 7634394, same benchmark: main's built-in model with the same drafter fix (4bf2378) 759 tok/s, 9.39 ms; #19839's target without it 665 tok/s, 11.08 ms
6 GSM8K with DSpark, batch 8 (wide steps, mostly 8 requests x 8 tokens), require: >= 96.5 - tol eval_dspark@8 ba53c84359 7636347 runs/session-7636347/06-eval_dspark_8_c4-k3rest-ba53c84359-require_gsm8kfull 96.97 over 1319 samples, acceptance length 4.292
5, 6 route B boot with this section's collectives (its MoE stays generic), require, no speculation, max_batch_size 8, the fixed prompts nospec16@8 ba53c84359 7636347 runs/session-7636347/05-nospec16_8_c4-k3rest-ba53c84359-require_fixed server up in 346 s; every rank: k3_sandwich_oproj on 93 / 93 layers, the MoE decode path on 0 (route B's block); texts coherent, 2 / 8 identical to route B without these commits (session 7631980) and 3 / 8 to the same code relaunched (session 7634394; route B's generic MoE tunes per launch)
5, 6 route B GSM8K, no speculation, batch 8, require: >= 96.5 - tol eval_nospec16@8 ba53c84359 7636347 runs/session-7636347/08-eval_nospec16_8_c4-k3rest-ba53c84359-require_gsm8kfull 97.12 over 1319 samples (route B without these commits: 96.66, session 7631980)
7 boot, require, DSpark with the target-owned drafter; fixed prompts; bench8 nat@8%c4-k3rest-0d16d09861-require:fixed,bench8 (arm 4) vs C3's nat@8%c4-k3rest-ba53c84359-require:fixed,bench8 (arm 3) 0d16d09 vs ba53c84 7636347 runs/session-7636347/04-nat_8_c4-k3rest-0d16d09861-require_fixed_bench8/serve.log PASS: the drafter line on 16 / 16 ranks (none in C3), entries taken at capture on every rank; 7 / 8 texts identical; bench8 in the ABBA row
7 GSM8K DSpark at 8 requests, C3-head vs C4-head (both require); acceptance length eval_dspark@8%c4-k3rest-0d16d09861-require:gsm8kfull (arm 7) vs C3's (arm 6) 0d16d09 vs ba53c84 7636347 runs/session-7636347/07-eval_dspark_8_c4-k3rest-0d16d09861-require_gsm8kfull/eval-gsm8kfull.log PASS: 96.66 / 96.59 vs 96.97 / 96.97 (-0.30 / -0.38); AL C4 / C3 0.9986 (verify request-steps 32,061 vs 32,017)
7 bench8, ABBA (arms 3 C3, 4 C4, 10 C4, 11 C3); bench1 nat@8%c4-k3rest-0d16d09861-require:bench1,bench8 (arm 10) vs C3's (arm 11), read with arms 4 / 3 0d16d09 vs ba53c84 7636347 runs/session-7636347/10-nat_8_c4-k3rest-0d16d09861-require_bench1_bench8/serve.log bench8 PASS, neutral as an ABBA mean of 2 launches each: median TPOT -0.45 % (C4 7.050 / 6.391 vs C3 6.839 / 6.662 ms), throughput -0.44 %, mean AL -0.15 %. bench1, one pair of 8 random-prompt requests: AL -11 % (mean 4.266 vs 4.798), so TPOT 2.013 vs 1.531 ms; TPOT x AL 7.27 vs 7.37 ms. Real-prompt AL at R = 1 is equal (next row)
7 GSM8K DSpark at 1 request (R = 1, where the drafter GEMV sites run); acceptance length eval_dspark@1%c4-k3rest-0d16d09861-require:gsm8k200 (arm 12) vs C3's (arm 13) 0d16d09 vs ba53c84 7636347 runs/session-7636347/12-eval_dspark_1_c4-k3rest-0d16d09861-require_gsm8k200/eval-gsm8k200.log PASS. AL 4.246 vs 4.241 (C4 +0.12 %; S 4,902 vs 4,908, T 20,813 from session 2 arm 11), R = 1 on 100 % of steps, the 4 sites taken at capture on 16 ranks. GSM8K 95.50 vs 96.00: one question out of 200 with 0 / 1 discordant (doc 57; McNemar p 1.0), so it passes (both right 191, C4 only 0, C3 only 1, neither 8; 179 / 200 responses identical). Device step 7.058 vs 7.207 ms median (-2.1 %)
7 the drafter GEMV sites with real weights, one GPU site_check.py: every TP16 rank's slice x 5 layers of the drafter checkpoint, 1..8 rows, vs torch bf16 and float64 0d16d09 7638408 runs/tokens/c4-sitecheck-7638408/site_check.log PASS. Each site's error vs float64 equals torch's in every (site, rows) cell: max 3.1-3.9e-3 of max |ref|, rel L2 1.66e-3, bias below 1e-5 and equal to torch's to 2 digits. drafter_o bit-identical to torch; the others differ in 0.03-0.12 % of outputs by 1 ulp of accumulation order
5, 7 route B GSM8K after it carries section 7, no speculation, batch 8, require: the same text as before the carry eval_nospec16@8 8806ae3a6b's route B tree (431f05db21), on main with every Kimi K3 PR applied (C++ and Python) 7640089 (arm 10; before the carry: arm 8) runs/session-7640089/10-eval_nospec16_8_k3main-57d7238926-require_gsm8kfull/eval-gsm8kfull.log pass: 97.12 ± 0.46, flexible = strict; generated text identical to before the carry on 1319 / 1319 documents; 16,195 decode steps either way, R = 8 on 99.2 %; every rank logs route B's KDA / MLA line and moe_tp=16, moe_ep=1

Part 7's rows pair C4's arms on this head (c4-k3rest-0d16d09861-require; the tree also carries a count hook for the
drafter's entries, outside this PR) with C3's arms on the previous head (c4-k3rest-ba53c84359-require, where the
target builds the stock drafter), in the same session. The reference is that pair, not off: off also swaps the
target for the built-in model, so a gap would mix the target and the drafter. A part 7 row passes when:

  • the entries ran, not just built: the "Kimi K3 DSpark drafter: 5 layers, decode blocks on k3_drafter_attn_qknorm and
    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_qknorm and the four drafter_* sites on every rank, eager and at capture. If
    every block fell back to the stock forward, texts, acceptance and bench would all equal C3's and the check would
    false-pass;
  • GSM8K is within 0.5 point of C3's;
  • the acceptance length from the verify request-steps is within 1 % of C3's, or higher;
  • bench8 TPOT is within 1 % of C3's, or faster.

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-compatible or api-breaking. For api-breaking, include BREAKING in the PR title. (api-compatible: trtllm::mnnvl_fusion_allreduce gains an optional trailing argument with main's default and declares buffer_flags mutable; MNNVLAllReduce.forward gains 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.md and 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.

Vasanth Sabavat 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]>
…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]>
… + 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]>
Vasanth Sabavat 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

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

Community want to contribute PRs initiated from Community

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants