Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
20 commits
Select commit Hold shift + click to select a range
6c54a34
[None][perf] Fuse DeepSeek-V4 O-projection
mingyangHao Jul 13, 2026
7362239
[None][perf] Tune DSV4 O_b scale pipeline
mingyangHao Jul 14, 2026
cb845b8
[None][perf] Use split-K 4 for tiny DSV4 O_b
mingyangHao Jul 14, 2026
9381f32
[None][fix] Harden DSV4 fused O-projection dispatch
mingyangHao Jul 14, 2026
a58236a
[None][fix] Address DSV4 O-projection review feedback
mingyangHao Jul 14, 2026
a284018
fix: avoid runtime shape JIT in DSV4 output projection
mingyangHao Jul 16, 2026
7b56c73
perf: autotune DSV4 O_b runtime tactics
mingyangHao Jul 17, 2026
7214cb4
[None][perf] DSV4 O_b: fix (1,1)-cluster grid sizing and expand SK1 t…
mingyangHao Jul 19, 2026
c58c6b9
[None][fix] DSV4 O_b: pair-local scale release mask for size-4 clusters
mingyangHao Jul 19, 2026
839ef9e
fix: gate DSV4 split O_b on fused HC
mingyangHao Jul 20, 2026
80aefc4
fix: default DSV4 O_b to DeepGEMM
mingyangHao Jul 20, 2026
78b0633
perf: write DSV4 O-projection outputs directly
mingyangHao Jul 20, 2026
17c7600
fix: prewarm DSV4 DeepGEMM O-projection
mingyangHao Jul 22, 2026
7075baf
test: verify DSV4 O_b warmup cache hits
mingyangHao Jul 29, 2026
02586bf
fix: reuse FP8 swap-ab warmup for DSV4 O_b
mingyangHao Jul 29, 2026
564dfc9
perf: overlap DSV4 O_b scale transfers
mingyangHao Aug 3, 2026
c2ee904
fix: tighten DSV4 fused O-proj contracts
mingyangHao Aug 3, 2026
fbb2f38
fix: keep DSV4 O-proj fallback on standard path
mingyangHao Aug 3, 2026
c139757
fix: precompile DSV4 O_a epilogue variants
mingyangHao Aug 3, 2026
e565933
perf: optimize DSV4 O_b and mHC split path
mingyangHao Aug 5, 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
500 changes: 342 additions & 158 deletions cpp/tensorrt_llm/kernels/mhcKernels/fused_tf32_pmap_gemm.cuh

Large diffs are not rendered by default.

114 changes: 97 additions & 17 deletions cpp/tensorrt_llm/kernels/mhcKernels/mhcFusedHcKernel.cu
Original file line number Diff line number Diff line change
Expand Up @@ -297,14 +297,75 @@ static constexpr uint32_t fhcSmemSize()
using FusedRoutFn = void (*)(
uint32_t, CUtensorMap, CUtensorMap, CUtensorMap, CUtensorMap, float*, float const*, float const*, float*);

template <uint32_t Hidden, uint32_t KS>
template <uint32_t Hidden, uint32_t KS, uint32_t XS = 1>
static FusedRoutFn fhcInstance()
{
static_assert(isSupportedFhcHidden<Hidden>(), "Unsupported fused-HC hidden size");
static_assert(isSupportedFhcMmaKS<Hidden, KS>(), "Unsupported fused-HC MMA kNumSplits for hidden size");
return &fused_mhc::fused_tf32_pmap_gemm_rout_atomic_impl<FHC_SHAPE_N, Hidden, FHC_HC_MULT, FHC_BLOCK_M, FHC_BLOCK_N,
FHC_BLOCK_K, FHC_SWIZZLE_CD, FHC_N_B_STAGES, FHC_N_INPUT_STG, FHC_NUM_MMA_TH, FHC_NUM_PMAP_TH, KS,
/*kEarlyRelease=*/false>;
/*kEarlyRelease=*/false, XS>;
}

template <uint32_t Hidden, uint32_t KS, uint32_t XS>
static FusedRoutFn fhcXSplitInstanceIfSupported()
{
if constexpr (isSupportedFhcMmaKS<Hidden, KS>())
{
return fhcInstance<Hidden, KS, XS>();
}
else
{
TLLM_CHECK_WITH_INFO(false, "mhcFusedHcLaunch: unsupported (kNumSplits=%u, hidden=%u)", KS, Hidden);
return nullptr;
}
}

template <uint32_t Hidden, uint32_t XS>
static FusedRoutFn pickFhcXSplitKs(uint32_t ks)
{
switch (ks)
{
case 1: return fhcInstance<Hidden, 1, XS>();
case 2: return fhcXSplitInstanceIfSupported<Hidden, 2, XS>();
case 4: return fhcXSplitInstanceIfSupported<Hidden, 4, XS>();
case 8: return fhcXSplitInstanceIfSupported<Hidden, 8, XS>();
case 16: return fhcXSplitInstanceIfSupported<Hidden, 16, XS>();
default: break;
}

// XSplit=2 is the production path for M=32..128. Keep its split-K
// coverage aligned with the unsplit dispatcher so those buckets retain
// the high-KS wave occupancy. XSplit=4 is used only for M<=16, where the
// FMA path wins, so avoid instantiating unused high-KS variants for it.
if constexpr (XS == 2)
{
switch (ks)
{
case 7: return fhcXSplitInstanceIfSupported<Hidden, 7, XS>();
case 14: return fhcXSplitInstanceIfSupported<Hidden, 14, XS>();
case 28: return fhcXSplitInstanceIfSupported<Hidden, 28, XS>();
case 32: return fhcXSplitInstanceIfSupported<Hidden, 32, XS>();
case 56: return fhcXSplitInstanceIfSupported<Hidden, 56, XS>();
case 64: return fhcXSplitInstanceIfSupported<Hidden, 64, XS>();
case 112: return fhcXSplitInstanceIfSupported<Hidden, 112, XS>();
default: break;
}
}

TLLM_CHECK_WITH_INFO(false, "mhcFusedHcLaunch: unsupported kNumSplits=%u for x split=%u", ks, XS);
return nullptr;
}

template <uint32_t Hidden>
static FusedRoutFn pickFhcXSplit(uint32_t ks, uint32_t xs)
{
switch (xs)
{
case 2: return pickFhcXSplitKs<Hidden, 2>(ks);
case 4: return pickFhcXSplitKs<Hidden, 4>(ks);
default: TLLM_CHECK_WITH_INFO(false, "mhcFusedHcLaunch: unsupported x split=%u", xs); return nullptr;
}
}

template <uint32_t Hidden, uint32_t KS>
Expand Down Expand Up @@ -379,7 +440,7 @@ static void mhcFusedHcLaunchImpl(__nv_bfloat16 const* x_prev, __nv_bfloat16 cons
__nv_bfloat16* layer_input_cur, float* y_acc_workspace, float* r_acc_workspace, int M, int hidden_size, int hc_mult,
int num_k_splits, int bigfuse_block_size, float rms_eps, float hc_pre_eps, float hc_sinkhorn_eps,
float hc_post_mult_value, int sinkhorn_repeat, __nv_bfloat16 const* norm_weight, float norm_eps,
cudaStream_t stream)
cudaStream_t stream, int x_num_splits = 1)
{
if (M <= 0)
return;
Expand Down Expand Up @@ -411,7 +472,8 @@ static void mhcFusedHcLaunchImpl(__nv_bfloat16 const* x_prev, __nv_bfloat16 cons
/*swizzleBytes=*/128, sizeof(__nv_bfloat16));

CUtensorMap desc_x = getCachedTma2D(const_cast<__nv_bfloat16*>(x_prev), CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, Hidden,
m_u, FHC_BLOCK_K, FHC_BLOCK_M, static_cast<uint64_t>(Hidden) * sizeof(__nv_bfloat16),
static_cast<uint32_t>(x_num_splits) * m_u, FHC_BLOCK_K, FHC_BLOCK_M,
static_cast<uint64_t>(Hidden) * sizeof(__nv_bfloat16),
/*swizzleBytes=*/128, sizeof(__nv_bfloat16));

CUtensorMap desc_b = getCachedTma2D(const_cast<float*>(w_t), CU_TENSOR_MAP_DATA_TYPE_TFLOAT32, SHAPE_K, FHC_SHAPE_N,
Expand All @@ -423,8 +485,11 @@ static void mhcFusedHcLaunchImpl(__nv_bfloat16 const* x_prev, __nv_bfloat16 cons
/*swizzleBytes=*/128, sizeof(__nv_bfloat16));

// ---- Step 1: fused post-mapping + TF32 GEMM + sqrsum + residual_out ----
constexpr uint32_t fused_smem = fhcSmemSize();
FusedRoutFn fa = pickFhc<Hidden>(ks);
// Split-x adds barriers but reuses the x data buffers.
uint32_t const extra_x_smem = (x_num_splits > 1) ? (2u * FHC_N_INPUT_STG * sizeof(uint64_t)) : 0u;
uint32_t const fused_smem = fhcSmemSize() + extra_x_smem;
FusedRoutFn fa
= (x_num_splits > 1) ? pickFhcXSplit<Hidden>(ks, static_cast<uint32_t>(x_num_splits)) : pickFhc<Hidden>(ks);
TLLM_CUDA_CHECK(cudaFuncSetAttribute(
reinterpret_cast<void const*>(fa), cudaFuncAttributeMaxDynamicSharedMemorySize, fused_smem));

Expand All @@ -447,7 +512,7 @@ void mhcFusedHcLaunch(__nv_bfloat16 const* x_prev, __nv_bfloat16 const* residual
__nv_bfloat16* residual_cur, float* post_mix_cur, float* comb_mix_cur, __nv_bfloat16* layer_input_cur,
float* y_acc_workspace, float* r_acc_workspace, int M, int hidden_size, int hc_mult, int num_k_splits,
int bigfuse_block_size, float rms_eps, float hc_pre_eps, float hc_sinkhorn_eps, float hc_post_mult_value,
int sinkhorn_repeat, __nv_bfloat16 const* norm_weight, float norm_eps, cudaStream_t stream)
int sinkhorn_repeat, __nv_bfloat16 const* norm_weight, float norm_eps, cudaStream_t stream, int x_num_splits)
Comment thread
coderabbitai[bot] marked this conversation as resolved.
{
if (M <= 0)
return;
Expand All @@ -461,13 +526,13 @@ void mhcFusedHcLaunch(__nv_bfloat16 const* x_prev, __nv_bfloat16 const* residual
mhcFusedHcLaunchImpl<FHC_HIDDEN_FLASH>(x_prev, residual_prev, post_mix_prev, comb_mix_prev, w_t, hc_scale,
hc_base, residual_cur, post_mix_cur, comb_mix_cur, layer_input_cur, y_acc_workspace, r_acc_workspace, M,
hidden_size, hc_mult, num_k_splits, bigfuse_block_size, rms_eps, hc_pre_eps, hc_sinkhorn_eps,
hc_post_mult_value, sinkhorn_repeat, norm_weight, norm_eps, stream);
hc_post_mult_value, sinkhorn_repeat, norm_weight, norm_eps, stream, x_num_splits);
return;
case static_cast<int>(FHC_HIDDEN_PRO):
mhcFusedHcLaunchImpl<FHC_HIDDEN_PRO>(x_prev, residual_prev, post_mix_prev, comb_mix_prev, w_t, hc_scale,
hc_base, residual_cur, post_mix_cur, comb_mix_cur, layer_input_cur, y_acc_workspace, r_acc_workspace, M,
hidden_size, hc_mult, num_k_splits, bigfuse_block_size, rms_eps, hc_pre_eps, hc_sinkhorn_eps,
hc_post_mult_value, sinkhorn_repeat, norm_weight, norm_eps, stream);
hc_post_mult_value, sinkhorn_repeat, norm_weight, norm_eps, stream, x_num_splits);
return;
default: return;
}
Expand All @@ -487,19 +552,21 @@ void mhcFusedHcLaunch(__nv_bfloat16 const* x_prev, __nv_bfloat16 const* residual
using FmaKsplitFn = void (*)(__nv_bfloat16 const*, __nv_bfloat16 const*, float const*, float const*, float const*,
float*, float*, int, int, int, __nv_bfloat16*);

template <int TN, int KS>
template <int TN, int KS, int XS = 1>
static FmaKsplitFn fhcFmaInstance()
{
return &fused_fma_kernels::fused_pmap_gemm_fma_ksplit<TN, KS, /*BF16_VEC_OVERRIDE=*/0, /*WRITE_RESIDUAL=*/true>;
return &fused_fma_kernels::fused_pmap_gemm_fma_ksplit<TN, KS, /*BF16_VEC_OVERRIDE=*/0,
/*WRITE_RESIDUAL=*/true, XS>;
}

// Valid (tile_n, num_k_splits) combinations the fused_hc FMA path supports.
// Keep this limited to the small/mid-M sweet spots from profile_fair_report v4.
static FmaKsplitFn pickFhcFma(int tile_n, int ks)
template <int XS>
static FmaKsplitFn pickFhcFmaXSplit(int tile_n, int ks)
{
#define FHCFMA_CASE(TN, KS) \
if (tile_n == (TN) && ks == (KS)) \
return fhcFmaInstance<TN, KS>()
return fhcFmaInstance<TN, KS, XS>()

FHCFMA_CASE(1, 1);
FHCFMA_CASE(1, 2);
Expand All @@ -519,16 +586,29 @@ static FmaKsplitFn pickFhcFma(int tile_n, int ks)
FHCFMA_CASE(12, 1);
FHCFMA_CASE(24, 1);
#undef FHCFMA_CASE
TLLM_CHECK_WITH_INFO(false, "mhcFusedHcFmaLaunch: unsupported (tile_n=%d, ks=%d)", tile_n, ks);
TLLM_CHECK_WITH_INFO(false, "mhcFusedHcFmaLaunch: unsupported (tile_n=%d, ks=%d, x_splits=%d)", tile_n, ks, XS);
return nullptr;
}

static FmaKsplitFn pickFhcFma(int tile_n, int ks, int x_num_splits)
{
switch (x_num_splits)
{
case 1: return pickFhcFmaXSplit<1>(tile_n, ks);
case 2: return pickFhcFmaXSplit<2>(tile_n, ks);
case 4: return pickFhcFmaXSplit<4>(tile_n, ks);
default:
TLLM_CHECK_WITH_INFO(false, "mhcFusedHcFmaLaunch: unsupported x_num_splits=%d", x_num_splits);
return nullptr;
}
}

void mhcFusedHcFmaLaunch(__nv_bfloat16 const* x_prev, __nv_bfloat16 const* residual_prev, float const* post_mix_prev,
float const* comb_mix_prev, float const* w_t, float const* hc_scale, float const* hc_base,
__nv_bfloat16* residual_cur, float* post_mix_cur, float* comb_mix_cur, __nv_bfloat16* layer_input_cur,
float* y_acc_workspace, float* r_acc_workspace, int M, int hidden_size, int hc_mult, int tile_n, int num_k_splits,
int bigfuse_block_size, float rms_eps, float hc_pre_eps, float hc_sinkhorn_eps, float hc_post_mult_value,
int sinkhorn_repeat, __nv_bfloat16 const* norm_weight, float norm_eps, cudaStream_t stream)
int sinkhorn_repeat, __nv_bfloat16 const* norm_weight, float norm_eps, cudaStream_t stream, int x_num_splits)
{
if (M <= 0)
return;
Expand All @@ -544,7 +624,7 @@ void mhcFusedHcFmaLaunch(__nv_bfloat16 const* x_prev, __nv_bfloat16 const* resid
int const N = static_cast<int>(FHC_SHAPE_N);

// ---- Step 1: fused pmap + GEMM + sqrsum + residual_cur (FMA ksplit) ----
FmaKsplitFn fn = pickFhcFma(tile_n, num_k_splits);
FmaKsplitFn fn = pickFhcFma(tile_n, num_k_splits, x_num_splits);
dim3 const grid(static_cast<unsigned>(M), static_cast<unsigned>(N / tile_n), static_cast<unsigned>(num_k_splits));
dim3 const block(256);
tensorrt_llm::common::launchWithPdlWhenEnabled("fused_pmap_gemm_fma_ksplit", fn, grid, block, 0, stream,
Expand Down Expand Up @@ -734,7 +814,7 @@ void mhcFusedHcLaunch(__nv_bfloat16 const* x_prev, __nv_bfloat16 const* residual
__nv_bfloat16* residual_cur, float* post_mix_cur, float* comb_mix_cur, __nv_bfloat16* layer_input_cur,
float* y_acc_workspace, float* r_acc_workspace, int M, int hidden_size, int hc_mult, int num_k_splits,
int bigfuse_block_size, float rms_eps, float hc_pre_eps, float hc_sinkhorn_eps, float hc_post_mult_value,
int sinkhorn_repeat, __nv_bfloat16 const* norm_weight, float norm_eps, cudaStream_t stream)
int sinkhorn_repeat, __nv_bfloat16 const* norm_weight, float norm_eps, cudaStream_t stream, int x_num_splits)
{
TLLM_CHECK_WITH_INFO(
false, "mhcFusedHcLaunch requires BUILD_DEEP_GEMM=ON to compile the TF32 MMA fused-HC backend");
Expand Down
4 changes: 2 additions & 2 deletions cpp/tensorrt_llm/kernels/mhcKernels/mhcKernels.h
Original file line number Diff line number Diff line change
Expand Up @@ -68,7 +68,7 @@ void mhcFusedHcLaunch(__nv_bfloat16 const* x_prev, __nv_bfloat16 const* residual
__nv_bfloat16* residual_cur, float* post_mix_cur, float* comb_mix_cur, __nv_bfloat16* layer_input_cur,
float* y_acc_workspace, float* r_acc_workspace, int M, int hidden_size, int hc_mult, int num_k_splits,
int bigfuse_block_size, float rms_eps, float hc_pre_eps, float hc_sinkhorn_eps, float hc_post_mult_value,
int sinkhorn_repeat, __nv_bfloat16 const* norm_weight, float norm_eps, cudaStream_t stream);
int sinkhorn_repeat, __nv_bfloat16 const* norm_weight, float norm_eps, cudaStream_t stream, int x_num_splits = 1);

// FMA-path fused hyper-connection boundary launcher.
//
Expand Down Expand Up @@ -96,7 +96,7 @@ void mhcFusedHcFmaLaunch(__nv_bfloat16 const* x_prev, __nv_bfloat16 const* resid
__nv_bfloat16* residual_cur, float* post_mix_cur, float* comb_mix_cur, __nv_bfloat16* layer_input_cur,
float* y_acc_workspace, float* r_acc_workspace, int M, int hidden_size, int hc_mult, int tile_n, int num_k_splits,
int bigfuse_block_size, float rms_eps, float hc_pre_eps, float hc_sinkhorn_eps, float hc_post_mult_value,
int sinkhorn_repeat, __nv_bfloat16 const* norm_weight, float norm_eps, cudaStream_t stream);
int sinkhorn_repeat, __nv_bfloat16 const* norm_weight, float norm_eps, cudaStream_t stream, int x_num_splits = 1);

// Single-kernel all-in-one fused hyper-connection boundary launcher (TF32 tcgen05 path).
//
Expand Down
Loading
Loading