Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
178 commits
Select commit Hold shift + click to select a range
1c1d49d
[None][perf] Cut host work between speculative decoding steps
Oct 1, 2026
f6cb23b
[None][feat] KDA verify: optional per-token states in the V2 hybrid c…
Oct 1, 2026
70ae504
[None][feat] One-model speculative decoding: the worker produces the …
Oct 1, 2026
b44c537
[None][fix] spec_step_copies: store zeros past a row's accepted tokens
Oct 2, 2026
92d67f3
[None][test] spec_step_copies tests: side stream waits for the curren…
Oct 2, 2026
f3cf23f
[None][fix] KV cache manager V2: keep the SSM pool at its live floor …
Oct 2, 2026
5e16d82
[None][fix] MNNVL all-reduce: keep a grown workspace's predecessor al…
Oct 2, 2026
f0ff9a8
MNNVL all-reduce: run the workspace-growth graph test in the GB200 mu…
Oct 2, 2026
551ca16
[None][chore] MNNVL all-reduce test: yapf the workspace-growth test
Oct 2, 2026
df9d9bf
[None][test] MNNVL all-reduce: one-rank test of the workspace's captu…
Oct 2, 2026
fd1bddc
MNNVL all-reduce: count Lamport arrivals per CTA in the RMSNorm-fused…
Oct 2, 2026
7a7ccf3
MNNVL all-reduce: start every CTA of the cluster before the RMSNorm c…
Oct 2, 2026
f467211
[None][fix] Kimi K3 attn_res_fwd online kernel: order the ws_stats re…
Oct 2, 2026
26cb3c3
[None][fix] Kimi K3 attn_res: cross-proxy fence before a consumer rel…
Oct 2, 2026
05e9e4f
[None][test] Kimi K3 attn_res: check the cross-proxy fence before eac…
Oct 2, 2026
af3e384
[None][fix] kdaDecode legacy: __syncwarp before the lane-0 block-redu…
Oct 2, 2026
1e0dceb
DFlash: start the dummy slot at context length 0 every step
Oct 2, 2026
ee0dc11
Merge #19830's head af3e384175 (kdaDecode legacy: __syncwarp before t…
Oct 3, 2026
6c2f854
[None][chore] Stack on the MNNVL RMSNorm-fused all-reduce cluster-bar…
Oct 3, 2026
5db6a14
[None][feat] Kimi K3 decode: KDA and MLA CuTe DSL kernels with their …
Oct 3, 2026
071adfa
[None][chore] Stack on the one-model speculative-decoding worker cont…
Oct 3, 2026
8e48156
[None][chore] Stack on the DFlash dummy-slot fix
Oct 3, 2026
b8edc4a
[None][feat] modeling_v2: Kimi K3 MXFP4 target skeleton, sm_100, tp16…
Oct 3, 2026
b45a946
[None][feat] Kimi K3 MLA attention: a caller-owned workspace
Oct 3, 2026
f51e780
[None][feat] modeling_v2 catalog: Kimi K3's stateful KDA and MLA deco…
Oct 3, 2026
373b10b
[None][chore] Kimi K3 KDA / MLA decode: format with main's hooks
Oct 3, 2026
ce2b3c3
[None][fix] modeling_v2 ssm/kda_decode: the op supports only the prod…
Oct 3, 2026
95fbc35
[None][test] modeling_v2 K3 KDA / MLA entries: sm_100 only, measured …
Oct 3, 2026
8ce0792
[None][test] Kimi K3 KDA / MLA decode: sm_100 receipts and l0_b200 en…
Oct 3, 2026
0ca3a9f
[None][feat] Kimi K3 decode kernels (CuTe DSL): CTM GEMVs, decode GEM…
Oct 3, 2026
a6828d2
[None][perf] Kimi K3 attn_res decode RMSNorm: release the next kernel…
Oct 3, 2026
b6a8fb5
[None][feat] Kimi K3 head GEMV: a caller-owned workspace
Oct 3, 2026
11cefbd
[None][feat] modeling_v2 catalog: the Kimi K3 single-GPU decode entries
Oct 3, 2026
c69ed9c
[None][chore] Kimi K3 decode kernels: format with the repo's pre-comm…
Oct 3, 2026
2700b9d
[None][test] Kimi K3 MLA op tests: drop the batch-1 identity cases
Oct 3, 2026
fe73486
[None][feat] modeling_v2 Kimi K3 target: classify each step for the d…
Oct 3, 2026
abb2564
[None][feat] modeling_v2 Kimi K3 target: its own text model, on the g…
Oct 3, 2026
9bc4dae
[None][test] modeling_v2 attn_res entries: replay once before poisoni…
Oct 3, 2026
8bfb169
[None][test] modeling_v2 claims: a target's generic path declares eve…
Oct 3, 2026
54e84c7
Merge U4 G3's published head 2700b9d4c2 (Kimi K3 KDA / MLA decode ker…
Oct 3, 2026
b280442
[None][test] modeling_v2 Kimi K3 decode entries: sm_100 receipts
Oct 3, 2026
e9b3595
[None][feat] modeling_v2: Kimi K3 MXFP4 target skeleton, sm_100, tp16…
Oct 3, 2026
f2a60d5
Merge U4 G2's head b28044256d (Kimi K3 single-GPU decode kernels and …
Oct 3, 2026
75d8a59
[None][feat] modeling_v2 Kimi K3 target: KDA and MLA attention on the…
Oct 3, 2026
0bf72e4
[None][feat] modeling_v2 Kimi K3 target: decode GEMVs, LM head and em…
Oct 3, 2026
fabb6e6
[None][feat] MNNVL all-reduce: Kimi K3 attention-residual epilogue (t…
Oct 3, 2026
494b1d9
[None][feat] modeling_v2 catalog: stateful entries; comm/mnnvl_allred…
Oct 3, 2026
b97811a
[None][feat] MNNVLAllReduce.allreduce_attn_res_rmsnorm, and the op's …
Oct 3, 2026
c922c7a
[None][doc] comm/mnnvl_allreduce_attn_res: every snapshot count 0-11 …
Oct 3, 2026
bf1f60a
[None][test] mnnvl_allreduce_attn_res matrix: run on sm_100 only
Oct 3, 2026
67f82d4
[None][feat] modeling_v2 Kimi K3 target: decode GEMV sites for the fu…
Oct 3, 2026
376aea8
Merge tokens' k3-up/u5-route-a-c2 75d8a597ec (KDA and MLA attention o…
Oct 3, 2026
754d706
[None][fix] modeling_v2 Kimi K3 target: build the decode GEMVs' state…
Oct 3, 2026
8c66821
[None][fix] MnnvlWorkspace.create: the ranks agree before allocating;…
Oct 3, 2026
a964500
[None][feat] modeling_v2 Kimi K3 target: layer 0's dense MLP on the d…
Oct 3, 2026
f224355
[None][feat] modeling_v2 Kimi K3 target: the attention's projections …
Oct 3, 2026
65cfd47
[None][fix] MnnvlWorkspace.create frees its communicator split when i…
Oct 3, 2026
2f444c9
[None][fix] modeling_v2 Kimi K3 target: no measurement date in a copi…
Oct 3, 2026
549ffd4
[None][fix] modeling_v2 Kimi K3 target: register trtllm::kda_mtp_deco…
Oct 3, 2026
f01a46b
[None][chore] modeling_v2 Kimi K3 target: trim the text model to the …
Oct 3, 2026
9926793
[None][fix] MnnvlWorkspace.create involves only the TP group's ranks
Oct 3, 2026
1af43c6
Merge #19839's head 549ffd4198 (Kimi K3 route A target: the target-ow…
Oct 3, 2026
f39b747
[None][chore] Stack on the Kimi K3 stateless decode entries (G2)
Oct 3, 2026
3691863
[None][feat] MNNVL: one-shot size per call, early PDL trigger; split …
Oct 3, 2026
65def52
[None][feat] Kimi K3 decode: collective CuTe DSL kernels with their o…
Oct 3, 2026
5181297
[None][test] Kimi K3 MoE front and tail sandwich tests: torch references
Oct 3, 2026
6fbd618
[None][refactor] Kimi K3 collectives: caller-owned state objects
Oct 3, 2026
1db2a80
[None][fix] Kimi K3 MoE wide state: refuse graph capture; flag word docs
Oct 3, 2026
bdad277
[None][feat] modeling_v2 catalog: Kimi K3 collective and MoE entries
Oct 3, 2026
20d2953
[None][test] Kimi K3 catalog entries: index rows and CI lists
Oct 3, 2026
84fc18f
[None][test] Kimi K3 MNNVL kernel test: the split all-gather, one-sho…
Oct 3, 2026
7a16d03
[None][test] Kimi K3 kernel lints: the MoE front, k3_moe and sandwich…
Oct 3, 2026
a35362e
[None][test] Kimi K3 collective tests: names codespell accepts, lines…
Oct 3, 2026
22629db
[None][chore] Kimi K3 collectives: format with the repo's pre-commit …
Oct 3, 2026
1de6f83
[None][fix] MNNVL: buffer_flags declared mutable; the split all-gathe…
Oct 3, 2026
823538a
[None][fix] Kimi K3 collective state: the ranks agree before allocating
Oct 3, 2026
bd4a407
[None][refactor] Kimi K3 MoE: trtllm::k3_moe, a torch op over the cal…
Oct 3, 2026
26a3dbc
[None][test] MNNVL catalog matrices: calls queued without host sync; …
Oct 3, 2026
e22949b
[None][test] Kimi K3 sandwich and latent exchange matrices: create() …
Oct 3, 2026
344d340
[None][feat] modeling_v2 catalog: moe/k3_moe over trtllm::k3_moe; moe…
Oct 3, 2026
dcb4d88
[None][test] Kimi K3 MoE front: the head geometry each TP size runs, …
Oct 3, 2026
d0c5e75
[None][chore] MNNVL split all-gather fake: format with the repo's hooks
Oct 3, 2026
f8a29e3
[None][test] Kimi K3 MoE front test: a latent code may differ only at…
Oct 3, 2026
8c0431e
[None][fix] Kimi K3 collectives: review nits
Oct 3, 2026
9fbc51f
[None][fix] Kimi K3 collective state: create() frees its communicator…
Oct 3, 2026
b30c349
[None][chore] Stack on the modeling_v2 catalog conventions and comm/m…
Oct 3, 2026
4ca4bfe
[None][feat] Kimi K3 DSpark drafter attention (trtllm::k3_drafter_att…
Oct 3, 2026
a6a8cb7
[None][feat] modeling_v2 catalog: attention/k3_drafter_attn, k3_draft…
Oct 3, 2026
90638fd
[None][fix] Kimi K3 collective state: create() involves its TP group'…
Oct 3, 2026
90158d9
[None][doc] MNNVL catalog contracts: MnnvlWorkspace.create involves t…
Oct 3, 2026
068241f
[None][feat] modeling_v2 Kimi K3 target: attention outputs for fused …
Oct 3, 2026
10f4a6e
[None][feat] modeling_v2 Kimi K3 tp16_moetp16ep1: tp16_moetp4ep4's mo…
Oct 3, 2026
ae282b3
[None][chore] Merge G1 9926793012: modeling_v2 catalog conventions fo…
Oct 3, 2026
120913c
[None][chore] Merge G6 a6a8cb73a3: Kimi K3 DSpark drafter attention a…
Oct 3, 2026
03df8ee
[None][fix] modeling_v2 Kimi K3 target: route the MXFP4 checkpoint as…
Oct 3, 2026
35eccac
Merge #19839's follow-ups f01a46b67f and 03df8ee849 (the trim; the MX…
Oct 3, 2026
fa24aca
[None][chore] modeling_v2 Kimi K3 tp16_moetp16ep1: carry the trim and…
Oct 3, 2026
de63263
[None][feat] modeling_v2 Kimi K3 target: KDA attention outputs on eve…
Oct 3, 2026
cce3f46
[None][test] Kimi K3 latent exchange: the one-process test checks the…
Oct 3, 2026
0dcbdaf
[None][test] Kimi K3 wide MoE test: a seeded skewed routing instead o…
Oct 3, 2026
a2d912b
[None][chore] Merge u5b fa24aca986: Kimi K3 route B (tp16_moetp16ep1)…
Oct 3, 2026
ba296fc
[None][test] Kimi K3 modeling_v2 target: run the decode-step test in …
Oct 3, 2026
df6da9d
[None][doc] modeling_v2 catalog: 18 entries carry sm_100 receipts
Oct 3, 2026
10eaa93
[None][chore] Merge u5-route-a-next df6da9d9b0: #19839's follow-ups (…
Oct 3, 2026
db8ca58
[None][test] modeling_v2 Kimi K3 collective and MoE entries: sm_100 r…
Oct 3, 2026
5c37c10
[None][chore] Merge G4 db8ca58bd0: Kimi K3 collective decode kernels …
Oct 3, 2026
37a8c92
[None][feat] Kimi K3 MoE: trtllm::k3_moe_m1 / k3_moe_m2, the routed e…
Oct 3, 2026
45cbb24
[None][feat] Kimi K3 MoE: k3_moe's push build into the latent exchange
Oct 3, 2026
81a0235
[None][test] Kimi K3 MoE: the push builds at 4 ranks against MNNVLAll…
Oct 3, 2026
82fc8ad
[None][feat] modeling_v2 catalog: moe/k3_moe_m1 and moe/k3_moe_m2 ove…
Oct 3, 2026
a04ec2e
[None][test] Kimi K3 kernel lints: the k3_moe_m1 and k3_moe_m2 kernels
Oct 3, 2026
f0fd94c
[None][feat] modeling_v2 catalog: moe/k3_moe's push form
Oct 3, 2026
9f7ad02
[None][chore] Merge RB f0fd94cfbd: Kimi K3 route B MoE engines (k3_mo…
Oct 3, 2026
ebfd380
[None][fix] modeling_v2 Kimi K3 target: declare the long GEMV and SiT…
Oct 3, 2026
85c2399
[None][chore] Merge u5-route-a ebfd38001c: #19839's declared-ops fix …
Oct 3, 2026
f32084b
[None][fix] modeling_v2 Kimi K3 tp16_moetp16ep1: declare the long GEM…
Oct 3, 2026
4bf2378
[None][fix] DSpark GQA drafter: keep the checkpoint's trained mask-to…
Oct 3, 2026
05ab11d
[None][doc] Kimi K3 collective contracts: the 16-rank runs passed
Oct 3, 2026
9dd1f27
[None][chore] Merge G4 05ab11d9e8: the collective contracts record th…
Oct 3, 2026
a55d13f
[None][doc] Kimi K3 MoE push form: the 16-rank runs passed
Oct 3, 2026
83924f8
[None][fix] Kimi K3 k3_moe: compile on routing that requires grad
Oct 3, 2026
d017fc9
[None][chore] Merge G4 83924f8281: k3_moe compiles on routing that re…
Oct 3, 2026
05e14ea
[None][chore] Merge RB a55d13f0fc: the push form's 16-rank runs passed
Oct 3, 2026
8ae58f7
[None][chore] Merge u5-route-a-c2 de6326362e: the Kimi K3 target's at…
Oct 3, 2026
eefd225
[None][feat] modeling_v2 Kimi K3 target: a decode step's attention-re…
Oct 3, 2026
93f0ad3
[None][feat] modeling_v2 Kimi K3 target: the post-attention all-reduc…
Oct 3, 2026
d1f297a
[None][feat] modeling_v2 Kimi K3 target: the post-attention step of a…
Oct 3, 2026
3ad8c5a
[None][feat] modeling_v2 Kimi K3 target: the decode path's one-shot c…
Oct 3, 2026
d99cbd6
[None][feat] modeling_v2 Kimi K3 target: the MoE layers of a decode s…
Oct 3, 2026
1c9c287
[None][test] modeling_v2 Kimi K3 target: the decode path's collective…
Oct 3, 2026
ba1db7e
[None][chore] modeling_v2 Kimi K3 tp16_moetp16ep1: carry C3 into the …
Oct 3, 2026
ba53c84
[None][fix] modeling_v2 Kimi K3 target: the MoE decode path hands its…
Oct 3, 2026
7dc60cc
[None][feat] SpecDecOneEngineForCausalLM: build the drafter through a…
Oct 3, 2026
31d76e0
[None][feat] modeling_v2 Kimi K3 target: decode GEMV sites for the DS…
Oct 3, 2026
a867199
[None][feat] modeling_v2 Kimi K3 target: its DSpark drafter on the dr…
Oct 3, 2026
b79745b
[None][test] modeling_v2 Kimi K3 target: the DSpark drafter against t…
Oct 3, 2026
592ea1a
[None][test] Kimi K3 drafter attention: the R x 7 splits
Oct 3, 2026
31b6275
[None][test] Kimi K3 drafter attention tests: name the splits they run
Oct 3, 2026
18f1edd
[None][doc] modeling_v2 catalog: the Kimi K3 drafter attention certif…
Oct 3, 2026
0d16d09
[None][feat] modeling_v2 Kimi K3 target: the DSpark drafter takes R x…
Oct 3, 2026
706cf2d
[None][test] Kimi K3 kernel lints: the KDA, MLA and drafter kernels' …
Oct 3, 2026
f0e7d4c
[None][test] Kimi K3 SiTU MoE: the TP16 shard (192 values, padded to …
Oct 3, 2026
87c7e70
[None][test] modeling_v2 MNNVL op matrices: run on sm_100 only
Oct 3, 2026
52ca9fa
[None][test] Kimi K3 MLA decode view test: list it in l0_cpu, not l0_…
Oct 3, 2026
543ebb3
[None][test] modeling_v2 Kimi K3 construction and drift tests: list t…
Oct 3, 2026
871ef64
[None][test] modeling_v2 Kimi K3 routing and claims tests: list them …
Oct 3, 2026
8806ae3
[None][chore] modeling_v2 Kimi K3 tp16_moetp16ep1: carry C4 into the …
Oct 3, 2026
32a581a
[None][feat] Kimi K3 DSpark decode: k3_spec_accept, k3_ctx_kv and k3_…
Oct 3, 2026
aaef4b7
[None][doc] k3_markov: state the co-residency assumption
Oct 3, 2026
a1d336d
[None][feat] DFlash / DSpark worker: Kimi K3 decode kernels behind k3…
Oct 3, 2026
e85d7a6
[None][feat] modeling_v2 Kimi K3 tp16_moetp4ep4: the speculative work…
Oct 3, 2026
4166788
[None][doc] modeling_v2 Kimi K3 tp16_moetp4ep4: the spec-worker gate'…
Oct 3, 2026
2b465f6
[None][chore] modeling_v2 Kimi K3 tp16_moetp16ep1: carry the speculat…
Oct 3, 2026
32d419b
[None][feat] DFlash metadata: capture_view, a captured layer's slot o…
Oct 3, 2026
1447397
[None][perf] DFlash: the draft block's rows as a view where its slots…
Oct 3, 2026
33d3586
[None][test] Kimi K3 DSpark decode kernel tests: list them in l0_b200…
Oct 3, 2026
5c193ac
[None][test] Kimi K3 Markov chain test: run its checks under pytest o…
Oct 3, 2026
490f289
[None][test] Kimi K3 sharded k3_spec_accept test: run under pytest on…
Oct 3, 2026
f85b709
[None][perf] modeling_v2 Kimi K3 tp16_moetp4ep4: split the DSpark dra…
Oct 3, 2026
135ab0c
[None][perf] modeling_v2 Kimi K3 tp16_moetp4ep4: the DSpark drafter's…
Oct 3, 2026
00e4394
[None][perf] DFlash / Kimi K3 DSpark drafter: the gen requests' page-…
Oct 3, 2026
ac6e287
[None][doc] modeling_v2 Kimi K3 tp16_moetp4ep4: the drafter gate requ…
Oct 3, 2026
147b87d
[None][test] k3_ctx_kv test: drop the base-port identity cases, zero …
Oct 3, 2026
b7ea3f6
[None][chore] modeling_v2 Kimi K3 tp16_moetp16ep1: carry the DSpark d…
Oct 3, 2026
19ef779
[None][test] Kimi K3 kernel tests: make their pool workers' functions…
Oct 3, 2026
428f6f2
[None][perf] modeling_v2 Kimi K3 tp16_moetp4ep4: the DSpark tap from …
Oct 3, 2026
fe9eaf5
[None][chore] modeling_v2 Kimi K3 tp16_moetp16ep1: carry the sandwich…
Oct 3, 2026
4bf3512
[None][perf] modeling_v2 Kimi K3 target: the decode MoE's latent all-…
Oct 3, 2026
41b2245
[None][test] Kimi K3 k3_moe push test: route A's expert layout
Oct 3, 2026
aed1c94
[None][chore] modeling_v2 Kimi K3 tp16_moetp16ep1: carry the latent p…
Oct 3, 2026
d6613a1
[None][feat] modeling_v2 Kimi K3 tp16_moetp16ep1: the MoE decode path…
Oct 3, 2026
3656bcd
[None][perf] Kimi K3 generic MoE path: the fused route + MXFP8 quant …
Oct 3, 2026
74839d1
[None][perf] Kimi K3 targets: route the generic MoE outside the TRTLL…
Oct 3, 2026
4657c20
[None][test] Kimi K3 decode-comm matrix: route A's shared-expert slic…
Oct 3, 2026
0da8a02
[None][fix] Kimi K3 decode MoE warm-up: no back-to-back k3_moe calls …
Oct 3, 2026
459e446
[None][doc] Kimi K3 sandwich: polled input and fold store before the …
Oct 4, 2026
c48e85e
[None][test] Kimi K3 decode-comm matrix: free a graph's outputs befor…
Oct 4, 2026
2f25c1d
[None][fix] DSpark K3 Markov path: release the step's acceptance and …
Oct 4, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions cpp/tensorrt_llm/common/lamportUtils.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -261,6 +261,8 @@ public:
if constexpr (UseCGA)
{
#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900))
// Every thread arrives on the cluster barrier but only rank 0's first warp waits, so a kernel using this
// must not arrive on the cluster barrier again (each thread arrives and waits once per phase).
cg::cluster_group cluster = cg::this_cluster();
__cluster_barrier_arrive();
if (cluster.block_rank() == 0 && threadIdx.x < kWARP_SIZE)
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,165 @@
/*
* Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/

#include "tensorrt_llm/kernels/communicationKernels/mnnvlAllGatherKernels.h"

#include "tensorrt_llm/common/cudaUtils.h"
#include "tensorrt_llm/common/envUtils.h"
#include "tensorrt_llm/common/lamportUtils.cuh"

TRTLLM_NAMESPACE_BEGIN

namespace kernels::mnnvl
{

using tensorrt_llm::common::isLamportDirty;
using tensorrt_llm::common::LamportFlags;
using tensorrt_llm::common::loadPackedVolatile;
using tensorrt_llm::common::VolatilePackedLoad;

namespace
{

constexpr int kThreads = 256;

// Payload vectors (16 bytes) of one rank's slot for one token: the bf16 part, then the fp32 part.
__host__ __device__ inline int slotVectors(int bf16Columns, int fp32Columns)
{
return bf16Columns / 8 + fp32Columns / 4;
}

// No payload word may equal the Lamport sentinel (the fp32 -0.0 word): -0.0 halves of a bf16
// pair and -0.0 fp32 words become +0.0.
__device__ inline uint32_t sanitizeBf16Pair(uint32_t word)
{
if ((word & 0xffffu) == 0x8000u)
{
word &= 0xffff0000u;
}
if ((word >> 16) == 0x8000u)
{
word &= 0x0000ffffu;
}
return word;
}

__device__ inline uint32_t sanitizeFp32(uint32_t word)
{
return word == 0x80000000u ? 0u : word;
}

__device__ inline uint32_t packBf16Pair(float lo, float hi)
{
__nv_bfloat162 const pair = __floats2bfloat162_rn(lo, hi);
return sanitizeBf16Pair(*reinterpret_cast<uint32_t const*>(&pair));
}

// One CTA per token: broadcast this rank's slot of the token, then poll every rank's slot and
// scatter it to the outputs.
__global__ void __launch_bounds__(kThreads) mnnvlAllGatherSplitKernel(AllGatherSplitParams const p)
{
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)
cudaGridDependencySynchronize();
// Dependents read the outputs (and the Lamport flags) only after their own wait for this grid.
cudaTriggerProgrammaticLaunchCompletion();
#endif
int const token = blockIdx.x;
int const bf16Vectors = p.bf16Columns / 8;
int const vectors = slotVectors(p.bf16Columns, p.fp32Columns);
int const inputColumns = p.bf16Columns + p.fp32Columns;

LamportFlags<float4> flag(p.bufferFlags, 1);
auto* lamportMcast = reinterpret_cast<uint4*>(flag.getCurLamportBuf(p.multicastPtr, 0));
auto* lamportLocal = reinterpret_cast<uint4*>(flag.getCurLamportBuf(p.bufferPtrsDev[p.rank], 0));

float const* row = p.input + static_cast<int64_t>(token) * inputColumns;
for (int v = threadIdx.x; v < vectors; v += blockDim.x)
{
uint4 packed;
if (v < bf16Vectors)
{
float4 const a = reinterpret_cast<float4 const*>(row)[2 * v];
float4 const b = reinterpret_cast<float4 const*>(row)[2 * v + 1];
packed = make_uint4(
packBf16Pair(a.x, a.y), packBf16Pair(a.z, a.w), packBf16Pair(b.x, b.y), packBf16Pair(b.z, b.w));
}
else
{
uint4 const words = reinterpret_cast<uint4 const*>(row + p.bf16Columns)[v - bf16Vectors];
packed = make_uint4(
sanitizeFp32(words.x), sanitizeFp32(words.y), sanitizeFp32(words.z), sanitizeFp32(words.w));
}
lamportMcast[(static_cast<int64_t>(token) * p.nRanks + p.rank) * vectors + v] = packed;
}

flag.ctaArrive();
flag.clearDirtyLamportBuf(p.bufferPtrsDev[p.rank], -1);

for (int i = threadIdx.x; i < p.nRanks * vectors; i += blockDim.x)
{
int const r = i / vectors;
int const v = i % vectors;
VolatilePackedLoad<float4> value;
do
{
value
= loadPackedVolatile<float4>(&lamportLocal[(static_cast<int64_t>(token) * p.nRanks + r) * vectors + v]);
} while (isLamportDirty(value));
uint4 const words = make_uint4(value.words[0], value.words[1], value.words[2], value.words[3]);
if (v < bf16Vectors)
{
reinterpret_cast<uint4*>(
p.bf16Output + static_cast<int64_t>(token) * p.nRanks * p.bf16Columns + r * p.bf16Columns)[v]
= words;
}
else
{
reinterpret_cast<uint4*>(p.fp32Output + static_cast<int64_t>(token) * p.nRanks * p.fp32Columns
+ r * p.fp32Columns)[v - bf16Vectors]
= words;
}
}

flag.waitAndUpdate({static_cast<uint32_t>(p.numTokens * p.nRanks * vectors * sizeof(uint4)), 0, 0, 0});
}

} // namespace

int64_t mnnvlAllGatherSplitFootprint(int numTokens, int bf16Columns, int fp32Columns, int nRanks)
{
return static_cast<int64_t>(numTokens) * nRanks * slotVectors(bf16Columns, fp32Columns) * sizeof(uint4);
}

void mnnvlAllGatherSplitOp(AllGatherSplitParams const& params)
{
TLLM_CHECK_WITH_INFO(params.bf16Columns % 8 == 0 && params.fp32Columns % 4 == 0,
"[mnnvlAllGatherSplit] needs bf16 columns in multiples of 8 and fp32 columns in multiples of 4");
TLLM_CHECK_WITH_INFO(params.numTokens > 0, "[mnnvlAllGatherSplit] needs at least one token");
cudaLaunchAttribute attrs[1];
attrs[0].id = cudaLaunchAttributeProgrammaticStreamSerialization;
attrs[0].val.programmaticStreamSerializationAllowed = tensorrt_llm::common::getEnvEnablePDL() ? 1 : 0;
cudaLaunchConfig_t config{};
config.gridDim = dim3(params.numTokens);
config.blockDim = dim3(kThreads);
config.stream = params.stream;
config.attrs = attrs;
config.numAttrs = 1;
TLLM_CUDA_CHECK(cudaLaunchKernelEx(&config, mnnvlAllGatherSplitKernel, params));
}

} // namespace kernels::mnnvl

TRTLLM_NAMESPACE_END
Original file line number Diff line number Diff line change
@@ -0,0 +1,66 @@
/*
* Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/

#pragma once

#include "tensorrt_llm/common/config.h"

#include <cstdint>
#include <cuda_bf16.h>
#include <cuda_runtime_api.h>

TRTLLM_NAMESPACE_BEGIN

namespace kernels::mnnvl
{

/**
* \brief Parameters of mnnvlAllGatherSplitOp: a one-shot all-gather over the MNNVL workspace of
* fp32 rows whose first bf16Columns columns travel, and are gathered, as bf16.
*
* Every rank contributes input [numTokens, bf16Columns + fp32Columns] (fp32). On every rank,
* bf16Output[t, r * bf16Columns + j] = bf16(input_r[t, j]) and
* fp32Output[t, r * fp32Columns + j] = input_r[t, bf16Columns + j], for every rank r. The bf16
* rounding is round-to-nearest, as a GEMV storing bf16 would round. -0.0 arrives as +0.0.
*
* The kernel follows the one-shot all-reduce's Lamport protocol on the shared workspace (it takes
* one turn of the buffer rotation) and triggers its dependents as soon as it starts.
*/
struct AllGatherSplitParams
{
float const* input; //!< [numTokens, bf16Columns + fp32Columns], this rank's slice
__nv_bfloat16* bf16Output; //!< [numTokens, nRanks * bf16Columns]
float* fp32Output; //!< [numTokens, nRanks * fp32Columns]; unused if fp32Columns == 0
int numTokens;
int bf16Columns; //!< Multiple of 8
int fp32Columns; //!< Multiple of 4

int nRanks;
int rank;
void** bufferPtrsDev;
void* multicastPtr;
uint32_t* bufferFlags;
cudaStream_t stream;
};

//! Bytes of one Lamport buffer the all-gather occupies.
int64_t mnnvlAllGatherSplitFootprint(int numTokens, int bf16Columns, int fp32Columns, int nRanks);

void mnnvlAllGatherSplitOp(AllGatherSplitParams const& params);

} // namespace kernels::mnnvl

TRTLLM_NAMESPACE_END
Loading