submission 371640
shigao · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 856 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-modal-nvfp4-dual-gemm-371640?include=source"interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp8_e4m3, nvfp4
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:5fe2da00ce1fadb882a48dc25314c3f94e75999fad18f362bd3758d9aaa0fb52
license declaredunknown
license concludedunknown
authorsshigao
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fused-epilogue
template <int K, int BM, int BN, int BK, int STAGES, int POLICY, int CTA_N_MAJOR, int FAST_SILU, int EPILOGUE_SEG, int FUSE_CP_MMA>mbarrier
asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(addr), "r"(count));shared-memory
extern __shared__ __align__(1024) char smem_raw[];tcgen05
asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(sdesc));tma
"cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint "vector-width = half2
half2* out_row0 = reinterpret_cast<half2*>(OUT + row * N);Kernel source
submission.py856 lines
import torch
from torch.utils.cpp_extension import load_inline
_CUDA_SRC = r"""
#include <cuda.h>
#include <cudaTypedefs.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <ATen/core/Tensor.h>
#include <torch/library.h>
// 说明(中文):仅面向 B200(sm_100a) 与评测固定形状做极致特化;不做任何回退。
// 分发器(中文):计分形状启用 FUSE_CP_MMA(scale cp 与 mma 交织);其余形状保持稳定路径。
// 备注(中文):EPILOGUE_SEG 保留为实验开关,当前计分区使用 SEG=32。
constexpr int WARP_SZ = 32;
constexpr int MMA_K64 = 64;
constexpr uint64_t L2_EVICT_FIRST = 0x12F0000000000000ULL;
constexpr uint64_t L2_EVICT_LAST = 0x14F0000000000000ULL;
__device__ __forceinline__ constexpr uint64_t desc_pack(uint64_t x) { return (x & 0x3FFFFULL) >> 4ULL; }
__device__ __forceinline__ uint32_t elect_one() {
uint32_t pred = 0;
asm volatile(
"{\n\t"
".reg .pred %%p;\n\t"
"elect.sync _|%%p, %1;\n\t"
"@%%p mov.s32 %0, 1;\n\t"
"}\n\t"
: "+r"(pred)
: "r"(0xFFFFFFFF)
);
return pred;
}
__device__ __forceinline__ uint64_t l2_policy_first() { return L2_EVICT_FIRST; }
__device__ __forceinline__ uint64_t l2_policy_last() { return L2_EVICT_LAST; }
__device__ __forceinline__ void mbar_init_shared(int addr, int count) {
asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(addr), "r"(count));
}
__device__ __forceinline__ void mbar_wait_parity(int addr, int phase) {
uint32_t ticks = 0x989680;
asm volatile(
"{\n\t"
".reg .pred P;\n\t"
"L_WAIT:\n\t"
"mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P, [%0], %1, %2;\n\t"
"@P bra.uni L_DONE;\n\t"
"bra.uni L_WAIT;\n\t"
"L_DONE:\n\t"
"}\n\t"
:: "r"(addr), "r"(phase), "r"(ticks)
);
}
__device__ __forceinline__ void tma_g2s_bytes(int dst, const void* src, int bytes, int mbar, uint64_t cache) {
asm volatile(
"cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint "
"[%0], [%1], %2, [%3], %4;"
:: "r"(dst), "l"(src), "r"(bytes), "r"(mbar), "l"(cache)
);
}
__device__ __forceinline__ void tma_g2s_3d(int dst, const void* tmap, int x, int y, int z, int mbar, uint64_t cache) {
asm volatile(
"cp.async.bulk.tensor.3d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::1.L2::cache_hint "
"[%0], [%1, {%2, %3, %4}], [%5], %6;"
:: "r"(dst), "l"(tmap), "r"(x), "r"(y), "r"(z), "r"(mbar), "l"(cache)
: "memory"
);
}
__device__ __forceinline__ void tc_scale_cp(uint32_t taddr, uint64_t sdesc) {
asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(sdesc));
}
__device__ __forceinline__ void tc_mma_a_fill(
uint32_t daddr,
uint64_t adesc,
uint64_t bdesc,
uint32_t idesc,
uint32_t scale_a,
uint32_t scale_b,
int enable_d
) {
asm volatile(
"{\n\t"
".reg .pred p;\n\t"
"setp.ne.b32 p, %6, 0;\n\t"
"tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::fill [%0], %1, %2, %3, [%4], [%5], p;\n\t"
"}\n\t"
:: "r"(daddr), "l"(adesc), "l"(bdesc), "r"(idesc), "r"(scale_a), "r"(scale_b), "r"(enable_d)
);
}
__device__ __forceinline__ void tc_mma_a_last(
uint32_t daddr,
uint64_t adesc,
uint64_t bdesc,
uint32_t idesc,
uint32_t scale_a,
uint32_t scale_b,
int enable_d
) {
asm volatile(
"{\n\t"
".reg .pred p;\n\t"
"setp.ne.b32 p, %6, 0;\n\t"
"tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::lastuse [%0], %1, %2, %3, [%4], [%5], p;\n\t"
"}\n\t"
:: "r"(daddr), "l"(adesc), "l"(bdesc), "r"(idesc), "r"(scale_a), "r"(scale_b), "r"(enable_d)
);
}
struct _TC_SH { static constexpr char _16x256b[] = ".16x256b"; };
struct _TC_NM { static constexpr char x4[] = ".x4"; static constexpr char x8[] = ".x8"; };
template <const char* SH, const char* NM>
__device__ __forceinline__ void tc_ld16(float* out, uint32_t addr) {
asm volatile(
"tcgen05.ld.sync.aligned%17%18.b32 "
"{ %0, %1, %2, %3, %4, %5, %6, %7, "
" %8, %9, %10, %11, %12, %13, %14, %15}, [%16];"
: "=f"(out[ 0]), "=f"(out[ 1]), "=f"(out[ 2]), "=f"(out[ 3]), "=f"(out[ 4]), "=f"(out[ 5]), "=f"(out[ 6]), "=f"(out[ 7]),
"=f"(out[ 8]), "=f"(out[ 9]), "=f"(out[10]), "=f"(out[11]), "=f"(out[12]), "=f"(out[13]), "=f"(out[14]), "=f"(out[15])
: "r"(addr), "C"(SH), "C"(NM)
);
}
template <const char* SH, const char* NM>
__device__ __forceinline__ void tc_ld32(float* out, uint32_t addr) {
asm volatile(
"tcgen05.ld.sync.aligned%33%34.b32 "
"{ %0, %1, %2, %3, %4, %5, %6, %7, "
" %8, %9, %10, %11, %12, %13, %14, %15, "
" %16, %17, %18, %19, %20, %21, %22, %23, "
" %24, %25, %26, %27, %28, %29, %30, %31}, [%32];"
: "=f"(out[ 0]), "=f"(out[ 1]), "=f"(out[ 2]), "=f"(out[ 3]), "=f"(out[ 4]), "=f"(out[ 5]), "=f"(out[ 6]), "=f"(out[ 7]),
"=f"(out[ 8]), "=f"(out[ 9]), "=f"(out[10]), "=f"(out[11]), "=f"(out[12]), "=f"(out[13]), "=f"(out[14]), "=f"(out[15]),
"=f"(out[16]), "=f"(out[17]), "=f"(out[18]), "=f"(out[19]), "=f"(out[20]), "=f"(out[21]), "=f"(out[22]), "=f"(out[23]),
"=f"(out[24]), "=f"(out[25]), "=f"(out[26]), "=f"(out[27]), "=f"(out[28]), "=f"(out[29]), "=f"(out[30]), "=f"(out[31])
: "r"(addr), "C"(SH), "C"(NM)
);
}
__device__ __forceinline__ void tc_ld_16x256bx8(float* out, uint32_t addr) { tc_ld32<_TC_SH::_16x256b, _TC_NM::x8>(out, addr); }
__device__ __forceinline__ void tc_ld_16x256bx4(float* out, uint32_t addr) { tc_ld16<_TC_SH::_16x256b, _TC_NM::x4>(out, addr); }
static inline void ck_cu(CUresult err) {
if (err == CUDA_SUCCESS) return;
const char* msg = nullptr;
if (cuGetErrorString(err, &msg) != CUDA_SUCCESS) msg = "cu err";
TORCH_CHECK(false, msg);
}
static inline void ck_cuda(cudaError_t err) {
if (err == cudaSuccess) return;
const char* msg = cudaGetErrorString(err);
TORCH_CHECK(false, msg ? msg : "cuda err");
}
static inline int sm_count_cached() {
static int sm = -1;
if (sm > 0) return sm;
int dev = 0;
ck_cuda(cudaGetDevice(&dev));
cudaDeviceProp prop;
ck_cuda(cudaGetDeviceProperties(&prop, dev));
sm = prop.multiProcessorCount;
return sm;
}
static inline void encode_tmap(
CUtensorMap* tmap,
const char* ptr,
uint64_t h,
uint64_t w,
uint32_t sh,
uint32_t sw,
CUtensorMapL2promotion promo
) {
constexpr uint32_t rank = 3;
uint64_t gdim[rank] = {256, h, w / 256};
uint64_t gstride[rank-1] = {w / 2, 128};
uint32_t bdim[rank] = {256, sh, sw / 256};
uint32_t estride[rank] = {1, 1, 1};
auto err = cuTensorMapEncodeTiled(
tmap,
CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,
rank,
(void*)ptr,
gdim,
gstride,
bdim,
estride,
CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE,
CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B,
promo,
CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE
);
ck_cu(err);
}
struct CacheSel {
uint64_t a_data;
uint64_t a_sf;
uint64_t b_data;
uint64_t b_sf;
};
template <int K, int POLICY>
__device__ __forceinline__ CacheSel cache_sel() {
const uint64_t p_first = l2_policy_first();
const uint64_t p_last = l2_policy_last();
CacheSel r;
if constexpr (K == 4096 || K == 7168) {
if constexpr (POLICY <= 2) {
if constexpr (POLICY == 1) {
r.a_data = p_last; r.a_sf = p_last;
r.b_data = p_first; r.b_sf = p_first;
} else if constexpr (POLICY == 2) {
r.a_data = p_first; r.a_sf = p_first;
r.b_data = p_last; r.b_sf = p_last;
} else {
r.a_data = p_first; r.a_sf = p_last;
r.b_data = p_first; r.b_sf = p_first;
}
} else {
static_assert(POLICY <= (3 + 15));
constexpr int mask = POLICY - 3;
r.a_data = (mask & 0x1) ? p_last : p_first;
r.a_sf = (mask & 0x2) ? p_last : p_first;
r.b_data = (mask & 0x4) ? p_last : p_first;
r.b_sf = (mask & 0x8) ? p_last : p_first;
}
} else {
r.a_data = p_first; r.a_sf = p_first;
r.b_data = p_last; r.b_sf = p_last;
}
return r;
}
template <int K, int BM, int BN, int BK, int STAGES, int POLICY, int CTA_N_MAJOR, int FAST_SILU, int EPILOGUE_SEG, int FUSE_CP_MMA>
__global__ __launch_bounds__(BM + 2 * WARP_SZ, 1)
void kernel_dual_fused(
const __grid_constant__ CUtensorMap A_t,
const __grid_constant__ CUtensorMap B1_t,
const __grid_constant__ CUtensorMap B2_t,
const char* __restrict__ SFA,
const char* __restrict__ SFB1,
const char* __restrict__ SFB2,
half* __restrict__ OUT,
int M,
int N
) {
const int tid = (int)threadIdx.x;
const int lane = tid & 31;
const int warp = tid >> 5;
int bid_m;
int bid_n;
if constexpr (K == 7168) {
if constexpr (CTA_N_MAJOR) { bid_n = (int)blockIdx.x; bid_m = (int)blockIdx.y; }
else { bid_m = (int)blockIdx.x; bid_n = (int)blockIdx.y; }
} else {
bid_n = (int)blockIdx.x;
bid_m = (int)blockIdx.y;
}
const int off_m = bid_m * BM;
const int off_n = bid_n * BN;
constexpr int WARP_CNT = BM / WARP_SZ + 2;
extern __shared__ __align__(1024) char smem_raw[];
const int smem = static_cast<int>(__cvta_generic_to_shared(smem_raw));
constexpr int A_BYTES = BM * BK / 2;
constexpr int B_BYTES = BN * BK / 2;
constexpr int SFA_BYTES = 128 * BK / 16;
constexpr int SFB_BYTES = 128 * BK / 16;
constexpr int STAGE_BYTES = A_BYTES + 2 * B_BYTES + SFA_BYTES + 2 * SFB_BYTES;
constexpr int TMEM_NEED = 2 * BN + 12 * (BK / MMA_K64);
constexpr int TMEM_COLS = (TMEM_NEED <= 256) ? 256 : 512;
static_assert(TMEM_NEED <= 512);
#pragma nv_diag_suppress static_var_with_dynamic_init
__shared__ int64_t mbars[STAGES * 2 + 1];
const int tma_mbar = static_cast<int>(__cvta_generic_to_shared(mbars));
const int mma_mbar = tma_mbar + STAGES * 8;
const int main_mbar = mma_mbar + STAGES * 8;
if (warp == 0 && elect_one()) {
#pragma unroll
for (int i = 0; i < STAGES * 2 + 1; ++i) mbar_init_shared(tma_mbar + i * 8, 1);
asm volatile("fence.mbarrier_init.release.cluster;");
}
if (warp == 1) {
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(smem), "r"(TMEM_COLS));
}
__syncthreads();
constexpr uint32_t tmem_base = 0;
constexpr uint32_t out1_col = 0;
constexpr uint32_t out2_col = (uint32_t)BN;
constexpr uint32_t sfa_col = (uint32_t)(2 * BN);
constexpr uint32_t sfb1_col = (uint32_t)(2 * BN + 4 * (BK / MMA_K64));
constexpr uint32_t sfb2_col = (uint32_t)(2 * BN + 8 * (BK / MMA_K64));
constexpr int iters = K / BK;
if (warp == WARP_CNT - 1 && elect_one()) {
const CacheSel pol = cache_sel<K, POLICY>();
constexpr int REST_K = K / 16 / 4;
constexpr int SF_STEP = (BK / (16 * 4)) * 512;
constexpr int Z_STEP = (BK / 256);
const int off_m128 = off_m >> 7;
const int off_n128 = off_n >> 7;
const char* sfa_src = SFA + off_m128 * REST_K * 512;
const char* sfb1_src = SFB1 + off_n128 * REST_K * 512;
const char* sfb2_src = SFB2 + off_n128 * REST_K * 512;
int stage = 0;
int wraps = 0;
int stage_base = smem;
int z = 0;
for (int iter = 0; iter < iters; ++iter) {
if (iter >= STAGES) {
mbar_wait_parity(mma_mbar + stage * 8, (wraps - 1) & 1);
}
const int mbar = tma_mbar + stage * 8;
const int a_s = stage_base;
const int b1_s = a_s + A_BYTES;
const int b2_s = b1_s + B_BYTES;
const int sfa_s = b2_s + B_BYTES;
const int sfb1_s = sfa_s + SFA_BYTES;
const int sfb2_s = sfb1_s + SFB_BYTES;
tma_g2s_3d(a_s, &A_t, 0, off_m, z, mbar, pol.a_data);
tma_g2s_3d(b1_s, &B1_t, 0, off_n, z, mbar, pol.b_data);
tma_g2s_3d(b2_s, &B2_t, 0, off_n, z, mbar, pol.b_data);
tma_g2s_bytes(sfa_s, sfa_src, SFA_BYTES, mbar, pol.a_sf);
tma_g2s_bytes(sfb1_s, sfb1_src, SFB_BYTES, mbar, pol.b_sf);
tma_g2s_bytes(sfb2_s, sfb2_src, SFB_BYTES, mbar, pol.b_sf);
asm volatile(
"mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
:: "r"(mbar), "r"(STAGE_BYTES)
: "memory"
);
z += Z_STEP;
sfa_src += SF_STEP;
sfb1_src += SF_STEP;
sfb2_src += SF_STEP;
++stage;
stage_base += STAGE_BYTES;
if (stage == STAGES) {
stage = 0;
stage_base = smem;
++wraps;
}
}
} else if (warp == WARP_CNT - 2 && elect_one()) {
constexpr int MMA_N = BN;
constexpr int MMA_M = 128;
constexpr uint32_t idesc =
(1U << 7U) | (1U << 10U) | ((uint32_t)MMA_N >> 3U << 17U) | ((uint32_t)MMA_M >> 7U << 27U);
auto desc_ab = [](int addr) -> uint64_t {
const int sbo = 8 * 128;
return desc_pack(addr) | (desc_pack(sbo) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
};
auto desc_sf = [](int addr) -> uint64_t {
const int sbo = 8 * 16;
return desc_pack(addr) | (desc_pack(sbo) << 32ULL) | (1ULL << 46ULL);
};
const uint64_t ab0 = desc_ab(0);
const uint64_t sf0 = desc_sf(0);
constexpr uint32_t SB_PARTS = (uint32_t)(128 / BN);
static_assert((128 % BN) == 0);
const uint32_t sb_off =
(SB_PARTS == 1) ? 0U : ((uint32_t)bid_n & (SB_PARTS - 1U)) * (uint32_t)(BN / 32);
int stage = 0;
int phase = 0;
int stage_base = smem;
for (int iter = 0; iter < iters; ++iter) {
mbar_wait_parity(tma_mbar + stage * 8, phase);
const int a_s = stage_base;
const int b1_s = a_s + A_BYTES;
const int b2_s = b1_s + B_BYTES;
const int sfa_s = b2_s + B_BYTES;
const int sfb1_s = sfa_s + SFA_BYTES;
const int sfb2_s = sfb1_s + SFB_BYTES;
const uint64_t sfa_desc0 = sf0 + ((uint64_t)sfa_s >> 4ULL);
const uint64_t sfb1_desc0 = sf0 + ((uint64_t)sfb1_s >> 4ULL);
const uint64_t sfb2_desc0 = sf0 + ((uint64_t)sfb2_s >> 4ULL);
const uint64_t a_base = ab0 + ((uint64_t)a_s >> 4ULL);
const uint64_t b1_base = ab0 + ((uint64_t)b1_s >> 4ULL);
const uint64_t b2_base = ab0 + ((uint64_t)b2_s >> 4ULL);
if constexpr (FUSE_CP_MMA) {
static_assert(BK == 256);
uint32_t td_sfa = tmem_base + sfa_col;
uint32_t td_sfb1 = tmem_base + sfb1_col;
uint32_t td_sfb2 = tmem_base + sfb2_col;
uint32_t sa = tmem_base + sfa_col;
uint32_t sb1 = tmem_base + sfb1_col + sb_off;
uint32_t sb2 = tmem_base + sfb2_col + sb_off;
uint64_t sfa_desc = sfa_desc0;
uint64_t sfb1_desc = sfb1_desc0;
uint64_t sfb2_desc = sfb2_desc0;
uint64_t a_desc = a_base;
uint64_t b1_desc = b1_base;
uint64_t b2_desc = b2_base;
#pragma unroll
for (int k2 = 0; k2 < BK / MMA_K64; ++k2) {
tc_scale_cp(td_sfa, sfa_desc);
tc_scale_cp(td_sfb1, sfb1_desc);
tc_scale_cp(td_sfb2, sfb2_desc);
const int en = (k2 == 0) ? iter : 1;
tc_mma_a_fill(tmem_base + out1_col, a_desc, b1_desc, idesc, sa, sb1, en);
tc_mma_a_last(tmem_base + out2_col, a_desc, b2_desc, idesc, sa, sb2, en);
a_desc += 2ULL;
b1_desc += 2ULL;
b2_desc += 2ULL;
sfa_desc += (512ULL >> 4ULL);
sfb1_desc += (512ULL >> 4ULL);
sfb2_desc += (512ULL >> 4ULL);
td_sfa += 4U;
td_sfb1 += 4U;
td_sfb2 += 4U;
sa += 4U;
sb1 += 4U;
sb2 += 4U;
}
} else {
uint32_t td_sfa = tmem_base + sfa_col;
uint32_t td_sfb1 = tmem_base + sfb1_col;
uint32_t td_sfb2 = tmem_base + sfb2_col;
uint64_t sfa_desc = sfa_desc0;
uint64_t sfb1_desc = sfb1_desc0;
uint64_t sfb2_desc = sfb2_desc0;
#pragma unroll
for (int kk = 0; kk < BK / MMA_K64; ++kk) {
tc_scale_cp(td_sfa, sfa_desc);
tc_scale_cp(td_sfb1, sfb1_desc);
tc_scale_cp(td_sfb2, sfb2_desc);
sfa_desc += (512ULL >> 4ULL);
sfb1_desc += (512ULL >> 4ULL);
sfb2_desc += (512ULL >> 4ULL);
td_sfa += 4U;
td_sfb1 += 4U;
td_sfb2 += 4U;
}
uint32_t sa = tmem_base + sfa_col;
uint32_t sb1 = tmem_base + sfb1_col + sb_off;
uint32_t sb2 = tmem_base + sfb2_col + sb_off;
if constexpr (BK == 256) {
uint64_t a_desc = a_base;
uint64_t b1_desc = b1_base;
uint64_t b2_desc = b2_base;
#pragma unroll
for (int k2 = 0; k2 < BK / MMA_K64; ++k2) {
const int en = (k2 == 0) ? iter : 1;
tc_mma_a_fill(tmem_base + out1_col, a_desc, b1_desc, idesc, sa, sb1, en);
tc_mma_a_last(tmem_base + out2_col, a_desc, b2_desc, idesc, sa, sb2, en);
a_desc += 2ULL;
b1_desc += 2ULL;
b2_desc += 2ULL;
sa += 4U;
sb1 += 4U;
sb2 += 4U;
}
} else {
constexpr uint64_t A_STEP = ((uint64_t)BM * 128ULL) >> 4ULL;
constexpr uint64_t B_STEP = ((uint64_t)BN * 128ULL) >> 4ULL;
#pragma unroll
for (int k2 = 0; k2 < BK / MMA_K64; ++k2) {
const int en = (k2 == 0) ? iter : 1;
const int k1 = k2 >> 2;
const int kk = k2 & 3;
const uint64_t a_desc = a_base + (uint64_t)k1 * A_STEP + (uint64_t)kk * 2ULL;
const uint64_t b1_desc = b1_base + (uint64_t)k1 * B_STEP + (uint64_t)kk * 2ULL;
const uint64_t b2_desc = b2_base + (uint64_t)k1 * B_STEP + (uint64_t)kk * 2ULL;
tc_mma_a_fill(tmem_base + out1_col, a_desc, b1_desc, idesc, sa, sb1, en);
tc_mma_a_last(tmem_base + out2_col, a_desc, b2_desc, idesc, sa, sb2, en);
sa += 4U;
sb1 += 4U;
sb2 += 4U;
}
}
}
asm volatile(
"tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
:: "r"(mma_mbar + stage * 8)
: "memory"
);
++stage;
stage_base += STAGE_BYTES;
if (stage == STAGES) {
stage = 0;
stage_base = smem;
phase ^= 1;
}
}
asm volatile(
"tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
:: "r"(main_mbar)
: "memory"
);
} else if (tid < BM) {
mbar_wait_parity(main_mbar, 0);
asm volatile("tcgen05.fence::after_thread_sync;");
constexpr float LOG2E = 1.4426950408889634f;
const int lane_row = lane >> 2;
const int lane_h2 = lane & 3;
constexpr int SEG = EPILOGUE_SEG;
static_assert((BN % SEG) == 0);
constexpr int SEGS = BN / SEG;
#pragma unroll
for (int mm = 0; mm < 2; ++mm) {
const int row0 = warp * 32 + mm * 16;
const uint32_t addr_x0 = tmem_base + (uint32_t)((row0 << 16) | (int)out1_col);
const uint32_t addr_y0 = tmem_base + (uint32_t)((row0 << 16) | (int)out2_col);
const int row = off_m + row0 + lane_row;
half2* out_row0 = reinterpret_cast<half2*>(OUT + row * N);
half2* out_row8 = reinterpret_cast<half2*>(OUT + (row + 8) * N);
#pragma unroll
for (int seg = 0; seg < SEGS; ++seg) {
float x[SEG / 2];
float y[SEG / 2];
const uint32_t col_off = (uint32_t)(seg * SEG);
if constexpr (SEG == 32) {
tc_ld_16x256bx4(x, addr_x0 + col_off);
tc_ld_16x256bx4(y, addr_y0 + col_off);
} else if constexpr (SEG == 64) {
tc_ld_16x256bx8(x, addr_x0 + col_off);
tc_ld_16x256bx8(y, addr_y0 + col_off);
} else {
static_assert(SEG == 32 || SEG == 64);
}
const int col_base_h2 = (off_n >> 1) + seg * (SEG >> 1);
asm volatile("tcgen05.wait::ld.sync.aligned;");
#pragma unroll
for (int i = 0; i < SEG / 8; ++i) {
const int out_col_base = col_base_h2 + i * 4;
const float x00 = x[i * 4 + 0];
const float x01 = x[i * 4 + 1];
const float x80 = x[i * 4 + 2];
const float x81 = x[i * 4 + 3];
const float y00 = y[i * 4 + 0];
const float y01 = y[i * 4 + 1];
const float y80 = y[i * 4 + 2];
const float y81 = y[i * 4 + 3];
float s00, s01, s80, s81;
if constexpr (FAST_SILU) {
float t00, t01, t80, t81;
asm("ex2.approx.f32 %0, %1;" : "=f"(t00) : "f"((-x00) * LOG2E));
asm("ex2.approx.f32 %0, %1;" : "=f"(t01) : "f"((-x01) * LOG2E));
asm("ex2.approx.f32 %0, %1;" : "=f"(t80) : "f"((-x80) * LOG2E));
asm("ex2.approx.f32 %0, %1;" : "=f"(t81) : "f"((-x81) * LOG2E));
const float d00 = 1.0f + t00;
const float d01 = 1.0f + t01;
const float d80 = 1.0f + t80;
const float d81 = 1.0f + t81;
asm("rcp.approx.f32 %0, %1;" : "=f"(s00) : "f"(d00));
asm("rcp.approx.f32 %0, %1;" : "=f"(s01) : "f"(d01));
asm("rcp.approx.f32 %0, %1;" : "=f"(s80) : "f"(d80));
asm("rcp.approx.f32 %0, %1;" : "=f"(s81) : "f"(d81));
} else {
s00 = __fdividef(1.0f, 1.0f + exp2f((-x00) * LOG2E));
s01 = __fdividef(1.0f, 1.0f + exp2f((-x01) * LOG2E));
s80 = __fdividef(1.0f, 1.0f + exp2f((-x80) * LOG2E));
s81 = __fdividef(1.0f, 1.0f + exp2f((-x81) * LOG2E));
}
float2 o0;
float2 o8;
o0.x = (x00 * s00) * y00;
o0.y = (x01 * s01) * y01;
o8.x = (x80 * s80) * y80;
o8.y = (x81 * s81) * y81;
const int out_col_h2 = out_col_base + lane_h2;
out_row0[out_col_h2] = __float22half2_rn(o0);
out_row8[out_col_h2] = __float22half2_rn(o8);
}
}
}
asm volatile("bar.sync 1, %0;" :: "r"(BM) : "memory");
if (warp == 0) asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(TMEM_COLS));
}
}
template <int K, int BM, int BN, int BK, int STAGES, int POLICY, int CTA_N_MAJOR, int FAST_SILU, int EPILOGUE_SEG, int FUSE_CP_MMA>
static inline void launch_cfg(
const at::Tensor& A,
const at::Tensor& B1,
const at::Tensor& B2,
const at::Tensor& SFA,
const at::Tensor& SFB1,
const at::Tensor& SFB2,
at::Tensor& out
) {
auto call = [](const at::Tensor& A,
const at::Tensor& B1,
const at::Tensor& B2,
const at::Tensor& SFA,
const at::Tensor& SFB1,
const at::Tensor& SFB2,
at::Tensor& out,
CUtensorMapL2promotion promo_a,
CUtensorMapL2promotion promo_b) {
const int M = (int)A.size(0);
const int N = (int)B1.size(0);
const char* A_ptr = reinterpret_cast<const char*>(A.data_ptr());
const char* B1_ptr = reinterpret_cast<const char*>(B1.data_ptr());
const char* B2_ptr = reinterpret_cast<const char*>(B2.data_ptr());
const char* SFA_ptr = reinterpret_cast<const char*>(SFA.data_ptr());
const char* SFB1_ptr = reinterpret_cast<const char*>(SFB1.data_ptr());
const char* SFB2_ptr = reinterpret_cast<const char*>(SFB2.data_ptr());
half* Out_ptr = reinterpret_cast<half*>(out.data_ptr());
struct TmapCache {
const char* a_ptr;
const char* b1_ptr;
const char* b2_ptr;
int m;
int n;
CUtensorMap a_t;
CUtensorMap b1_t;
CUtensorMap b2_t;
bool valid;
};
#pragma nv_diag_suppress static_var_with_dynamic_init
static TmapCache cache = {nullptr, nullptr, nullptr, 0, 0, {}, {}, {}, false};
CUtensorMap A_t, B1_t, B2_t;
if (cache.valid && cache.a_ptr == A_ptr && cache.b1_ptr == B1_ptr && cache.b2_ptr == B2_ptr && cache.m == M && cache.n == N) {
A_t = cache.a_t;
B1_t = cache.b1_t;
B2_t = cache.b2_t;
} else {
encode_tmap(&A_t, A_ptr, (uint64_t)M, (uint64_t)K, (uint32_t)BM, (uint32_t)BK, promo_a);
encode_tmap(&B1_t, B1_ptr, (uint64_t)N, (uint64_t)K, (uint32_t)BN, (uint32_t)BK, promo_b);
encode_tmap(&B2_t, B2_ptr, (uint64_t)N, (uint64_t)K, (uint32_t)BN, (uint32_t)BK, promo_b);
cache.a_ptr = A_ptr;
cache.b1_ptr = B1_ptr;
cache.b2_ptr = B2_ptr;
cache.m = M;
cache.n = N;
cache.a_t = A_t;
cache.b1_t = B1_t;
cache.b2_t = B2_t;
cache.valid = true;
}
dim3 grid;
if constexpr (K == 7168) {
if constexpr (CTA_N_MAJOR) grid = dim3((unsigned)(N / BN), (unsigned)(M / BM));
else grid = dim3((unsigned)(M / BM), (unsigned)(N / BN));
} else {
grid = dim3((unsigned)(N / BN), (unsigned)(M / BM));
}
const int tb = BM + 2 * WARP_SZ;
constexpr int A_BYTES = BM * BK / 2;
constexpr int B_BYTES = BN * BK / 2;
constexpr int SFA_BYTES = 128 * BK / 16;
constexpr int SFB_BYTES = 128 * BK / 16;
constexpr int smem_bytes = (A_BYTES + 2 * B_BYTES + SFA_BYTES + 2 * SFB_BYTES) * STAGES;
constexpr int kMaxSmemBytes = 227 * 1024;
TORCH_CHECK(smem_bytes <= kMaxSmemBytes, "smem ", smem_bytes);
auto kptr = kernel_dual_fused<K, BM, BN, BK, STAGES, POLICY, CTA_N_MAJOR, FAST_SILU, EPILOGUE_SEG, FUSE_CP_MMA>;
if constexpr (smem_bytes > 48000) {
static bool attr_set = false;
if (!attr_set) {
ck_cuda(cudaFuncSetAttribute(kptr, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_bytes));
attr_set = true;
}
}
kptr<<<grid, tb, smem_bytes>>>(A_t, B1_t, B2_t, SFA_ptr, SFB1_ptr, SFB2_ptr, Out_ptr, M, N);
};
call(A, B1, B2, SFA, SFB1, SFB2, out,
CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE);
}
template <int K, int BM, int BN, int BK, int STAGES, int FAST_SILU, int EPILOGUE_SEG, int FUSE_CP_MMA>
static inline void launch_policy_auto(
const at::Tensor& A,
const at::Tensor& B1,
const at::Tensor& B2,
const at::Tensor& SFA,
const at::Tensor& SFB1,
const at::Tensor& SFB2,
at::Tensor& out,
int cta_m,
int cta_n
) {
if constexpr (K == 4096 || K == 7168) {
if (cta_n >= (cta_m << 3)) launch_cfg<K, BM, BN, BK, STAGES, 1, 0, FAST_SILU, EPILOGUE_SEG, FUSE_CP_MMA>(A, B1, B2, SFA, SFB1, SFB2, out);
else if (cta_m >= (cta_n << 3)) launch_cfg<K, BM, BN, BK, STAGES, 2, 0, FAST_SILU, EPILOGUE_SEG, FUSE_CP_MMA>(A, B1, B2, SFA, SFB1, SFB2, out);
else launch_cfg<K, BM, BN, BK, STAGES, 0, 0, FAST_SILU, EPILOGUE_SEG, FUSE_CP_MMA>(A, B1, B2, SFA, SFB1, SFB2, out);
} else {
launch_cfg<K, BM, BN, BK, STAGES, 0, 0, FAST_SILU, EPILOGUE_SEG, FUSE_CP_MMA>(A, B1, B2, SFA, SFB1, SFB2, out);
}
}
static inline uint64_t shape_key_u64(int m, int n, int k) {
return ((uint64_t)(uint32_t)k << 32) | ((uint64_t)(uint32_t)m << 16) | (uint64_t)(uint32_t)n;
}
#define KKEY(M, N, K) ((((uint64_t)(K)) << 32) | (((uint64_t)(M)) << 16) | ((uint64_t)(N)))
at::Tensor fused(
const at::Tensor& A,
const at::Tensor& B1,
const at::Tensor& B2,
const at::Tensor& SFA,
const at::Tensor& SFB1,
const at::Tensor& SFB2,
at::Tensor& out
) {
const int M = (int)A.size(0);
const int Kp = (int)A.size(1);
const int N = (int)B1.size(0);
const int K = Kp * 2;
const uint64_t key = shape_key_u64(M, N, K);
switch (key) {
// 计分区:固定 shape 强特化(启用 FUSE_CP_MMA;EPILOGUE_SEG=32)
case KKEY(256, 4096, 7168):
launch_cfg<7168, 128, 64, 256, 5, 1, 1, 0, 32, 1>(A, B1, B2, SFA, SFB1, SFB2, out);
return out;
case KKEY(512, 4096, 7168):
launch_cfg<7168, 128, 128, 256, 4, 3, 1, 0, 32, 1>(A, B1, B2, SFA, SFB1, SFB2, out);
return out;
case KKEY(256, 3072, 4096):
launch_cfg<4096, 128, 64, 256, 5, 3, 0, 0, 32, 1>(A, B1, B2, SFA, SFB1, SFB2, out);
return out;
case KKEY(512, 3072, 7168):
launch_cfg<7168, 128, 128, 256, 4, 3, 1, 0, 32, 1>(A, B1, B2, SFA, SFB1, SFB2, out);
return out;
default:
break;
}
const int sm = sm_count_cached();
const int cta_m = M / 128;
const int cta_n128 = N / 128;
const int cta_128 = cta_m * cta_n128;
const int cta128_threshold = (sm > 96) ? 96 : sm;
const bool use_bn128 = ((N & 127) == 0) && (cta_128 >= cta128_threshold);
const int cta_n64 = N / 64;
// 正确性区:其余形状仅需通过测试(保守 epilogue SEG=32 + x4;仍然使用自定义 CUDA kernel,无回退)
if (K == 7168) {
if (use_bn128) launch_policy_auto<7168, 128, 128, 256, 4, 0, 32, 0>(A, B1, B2, SFA, SFB1, SFB2, out, cta_m, cta_n128);
else launch_policy_auto<7168, 128, 64, 256, 5, 0, 32, 0>(A, B1, B2, SFA, SFB1, SFB2, out, cta_m, cta_n64);
} else if (K == 4096) {
if (use_bn128) launch_policy_auto<4096, 128, 128, 256, 4, 0, 32, 0>(A, B1, B2, SFA, SFB1, SFB2, out, cta_m, cta_n128);
else launch_policy_auto<4096, 128, 64, 256, 5, 0, 32, 0>(A, B1, B2, SFA, SFB1, SFB2, out, cta_m, cta_n64);
} else if (K == 2304) {
launch_policy_auto<2304, 128, 64, 256, 4, 0, 32, 0>(A, B1, B2, SFA, SFB1, SFB2, out, cta_m, cta_n64);
} else if (K == 2048) {
launch_policy_auto<2048, 128, 64, 256, 4, 0, 32, 0>(A, B1, B2, SFA, SFB1, SFB2, out, cta_m, cta_n64);
} else if (K == 1536) {
launch_policy_auto<1536, 128, 64, 256, 4, 0, 32, 0>(A, B1, B2, SFA, SFB1, SFB2, out, cta_m, cta_n64);
} else if (K == 512) {
launch_policy_auto<512, 128, 64, 256, 4, 0, 32, 0>(A, B1, B2, SFA, SFB1, SFB2, out, cta_m, cta_n64);
} else if (K == 256) {
launch_policy_auto<256, 128, 64, 256, 4, 0, 32, 0>(A, B1, B2, SFA, SFB1, SFB2, out, cta_m, cta_n64);
} else {
TORCH_CHECK(false, "k ", K);
}
return out;
}
TORCH_LIBRARY(nvfp4_dual_lib_r215_c77, m) {
m.def("fused(Tensor A, Tensor B1, Tensor B2, Tensor SFA, Tensor SFB1, Tensor SFB2, Tensor(a!) out) -> Tensor");
m.impl("fused", &fused);
}
"""
_READY = False
def _init():
global _READY
if _READY:
return
load_inline(
name="nvfp4_dual_ext_r215_c77",
cpp_sources="",
cuda_sources=_CUDA_SRC,
functions=None,
with_cuda=True,
extra_cuda_cflags=[
"-O3",
"-gencode=arch=compute_100a,code=sm_100a",
"--use_fast_math",
"--expt-relaxed-constexpr",
"--relocatable-device-code=false",
],
extra_ldflags=["-lcuda"],
verbose=False,
is_python_module=False,
no_implicit_headers=True,
)
_READY = True
def custom_kernel(data):
_init()
a, b1, b2, _sfa, _sfb1, _sfb2, sfa_p, sfb1_p, sfb2_p, c = data
return torch.ops.nvfp4_dual_lib_r215_c77.fused(a, b1, b2, sfa_p, sfb1_p, sfb2_p, c)
__all__ = ["custom_kernel"]
scrolls · 856 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 371636.
---import torchfrom torch.utils.cpp_extension import load_inline⋯ 4 unchanged lines#include <cuda_fp16.h>#include <cuda_runtime.h>- #include <torch/library.h>#include <ATen/core/Tensor.h>+ #include <torch/library.h>- constexpr int WARP_SIZE = 32;- constexpr int MMA_K = 64;+ // 说明(中文):仅面向 B200(sm_100a) 与评测固定形状做极致特化;不做任何回退。+ // 分发器(中文):计分形状启用 FUSE_CP_MMA(scale cp 与 mma 交织);其余形状保持稳定路径。+ // 备注(中文):EPILOGUE_SEG 保留为实验开关,当前计分区使用 SEG=32。- constexpr uint64_t EVICT_FIRST = 0x12F0000000000000ULL;- constexpr uint64_t EVICT_LAST = 0x14F0000000000000ULL;+ constexpr int WARP_SZ = 32;+ constexpr int MMA_K64 = 64;+ constexpr uint64_t L2_EVICT_FIRST = 0x12F0000000000000ULL;+ constexpr uint64_t L2_EVICT_LAST = 0x14F0000000000000ULL;- __device__ __forceinline__ constexpr uint64_t desc_encode(uint64_t x) { return (x & 0x3'FFFFULL) >> 4ULL; }+ __device__ __forceinline__ constexpr uint64_t desc_pack(uint64_t x) { return (x & 0x3FFFFULL) >> 4ULL; }- __device__ __forceinline__ uint32_t elect_sync() {+ __device__ __forceinline__ uint32_t elect_one() {uint32_t pred = 0;asm volatile("{\n\t"- ".reg .pred %%px;\n\t"- "elect.sync _|%%px, %1;\n\t"- "@%%px mov.s32 %0, 1;\n\t"+ ".reg .pred %%p;\n\t"+ "elect.sync _|%%p, %1;\n\t"+ "@%%p mov.s32 %0, 1;\n\t""}\n\t": "+r"(pred): "r"(0xFFFFFFFF)⋯ 1 unchanged linesreturn pred;}- __device__ __forceinline__ void mbarrier_init(int mbar_addr, int count) {- asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(mbar_addr), "r"(count));+ __device__ __forceinline__ uint64_t l2_policy_first() { return L2_EVICT_FIRST; }+ __device__ __forceinline__ uint64_t l2_policy_last() { return L2_EVICT_LAST; }++ __device__ __forceinline__ void mbar_init_shared(int addr, int count) {+ asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(addr), "r"(count));}- __device__ __forceinline__ void mbarrier_wait(int mbar_addr, int phase) {+ __device__ __forceinline__ void mbar_wait_parity(int addr, int phase) {uint32_t ticks = 0x989680;asm volatile("{\n\t"- ".reg .pred P1;\n\t"- "LAB_WAIT:\n\t"- "mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P1, [%0], %1, %2;\n\t"- "@P1 bra.uni DONE;\n\t"- "bra.uni LAB_WAIT;\n\t"- "DONE:\n\t"+ ".reg .pred P;\n\t"+ "L_WAIT:\n\t"+ "mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P, [%0], %1, %2;\n\t"+ "@P bra.uni L_DONE;\n\t"+ "bra.uni L_WAIT;\n\t"+ "L_DONE:\n\t""}\n\t"- :: "r"(mbar_addr), "r"(phase), "r"(ticks)+ :: "r"(addr), "r"(phase), "r"(ticks));}- __device__ __forceinline__ void tma_gmem2smem(int dst, const void *src, int size, int mbar_addr, uint64_t cache_policy) {+ __device__ __forceinline__ void tma_g2s_bytes(int dst, const void* src, int bytes, int mbar, uint64_t cache) {asm volatile("cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint ""[%0], [%1], %2, [%3], %4;"- :: "r"(dst), "l"(src), "r"(size), "r"(mbar_addr), "l"(cache_policy)+ :: "r"(dst), "l"(src), "r"(bytes), "r"(mbar), "l"(cache));}- __device__ __forceinline__ void tma_3d_gmem2smem(int dst, const void *tmap_ptr, int x, int y, int z, int mbar_addr, uint64_t cache_policy) {+ __device__ __forceinline__ void tma_g2s_3d(int dst, const void* tmap, int x, int y, int z, int mbar, uint64_t cache) {asm volatile("cp.async.bulk.tensor.3d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::1.L2::cache_hint ""[%0], [%1, {%2, %3, %4}], [%5], %6;"- :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "l"(cache_policy)+ :: "r"(dst), "l"(tmap), "r"(x), "r"(y), "r"(z), "r"(mbar), "l"(cache): "memory");}- __device__ __forceinline__ void tcgen05_cp_nvfp4(int taddr, uint64_t s_desc) {- asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));+ __device__ __forceinline__ void tc_scale_cp(uint32_t taddr, uint64_t sdesc) {+ asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(sdesc));}- __device__ __forceinline__ void tcgen05_mma_nvfp4(- uint64_t a_desc,- uint64_t b_desc,- uint32_t i_desc,- int scale_A_tmem,- int scale_B_tmem,- int enable_input_d+ __device__ __forceinline__ void tc_mma_a_fill(+ uint32_t daddr,+ uint64_t adesc,+ uint64_t bdesc,+ uint32_t idesc,+ uint32_t scale_a,+ uint32_t scale_b,+ int enable_d) {- const int d_tmem = 0;asm volatile("{\n\t"".reg .pred p;\n\t""setp.ne.b32 p, %6, 0;\n\t"- "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [%0], %1, %2, %3, [%4], [%5], p;\n\t"+ "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::fill [%0], %1, %2, %3, [%4], [%5], p;\n\t""}\n\t"- :: "r"(d_tmem), "l"(a_desc), "l"(b_desc), "r"(i_desc),- "r"(scale_A_tmem), "r"(scale_B_tmem), "r"(enable_input_d)+ :: "r"(daddr), "l"(adesc), "l"(bdesc), "r"(idesc), "r"(scale_a), "r"(scale_b), "r"(enable_d));}- struct SHAPE { static constexpr char _16x256b[] = ".16x256b"; };- struct NUM { static constexpr char x8[] = ".x8"; };+ __device__ __forceinline__ void tc_mma_a_last(+ uint32_t daddr,+ uint64_t adesc,+ uint64_t bdesc,+ uint32_t idesc,+ uint32_t scale_a,+ uint32_t scale_b,+ int enable_d+ ) {+ asm volatile(+ "{\n\t"+ ".reg .pred p;\n\t"+ "setp.ne.b32 p, %6, 0;\n\t"+ "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::lastuse [%0], %1, %2, %3, [%4], [%5], p;\n\t"+ "}\n\t"+ :: "r"(daddr), "l"(adesc), "l"(bdesc), "r"(idesc), "r"(scale_a), "r"(scale_b), "r"(enable_d)+ );+ }- template <const char *SHAPE_, const char *NUM_>- __device__ __forceinline__ void tcgen05_ld_32regs(float *tmp, int row, int col) {+ struct _TC_SH { static constexpr char _16x256b[] = ".16x256b"; };+ struct _TC_NM { static constexpr char x4[] = ".x4"; static constexpr char x8[] = ".x8"; };++ template <const char* SH, const char* NM>+ __device__ __forceinline__ void tc_ld16(float* out, uint32_t addr) {asm volatile(+ "tcgen05.ld.sync.aligned%17%18.b32 "+ "{ %0, %1, %2, %3, %4, %5, %6, %7, "+ " %8, %9, %10, %11, %12, %13, %14, %15}, [%16];"+ : "=f"(out[ 0]), "=f"(out[ 1]), "=f"(out[ 2]), "=f"(out[ 3]), "=f"(out[ 4]), "=f"(out[ 5]), "=f"(out[ 6]), "=f"(out[ 7]),+ "=f"(out[ 8]), "=f"(out[ 9]), "=f"(out[10]), "=f"(out[11]), "=f"(out[12]), "=f"(out[13]), "=f"(out[14]), "=f"(out[15])+ : "r"(addr), "C"(SH), "C"(NM)+ );+ }++ template <const char* SH, const char* NM>+ __device__ __forceinline__ void tc_ld32(float* out, uint32_t addr) {+ asm volatile("tcgen05.ld.sync.aligned%33%34.b32 ""{ %0, %1, %2, %3, %4, %5, %6, %7, "" %8, %9, %10, %11, %12, %13, %14, %15, "" %16, %17, %18, %19, %20, %21, %22, %23, "" %24, %25, %26, %27, %28, %29, %30, %31}, [%32];"- : "=f"(tmp[ 0]), "=f"(tmp[ 1]), "=f"(tmp[ 2]), "=f"(tmp[ 3]), "=f"(tmp[ 4]), "=f"(tmp[ 5]), "=f"(tmp[ 6]), "=f"(tmp[ 7]),- "=f"(tmp[ 8]), "=f"(tmp[ 9]), "=f"(tmp[10]), "=f"(tmp[11]), "=f"(tmp[12]), "=f"(tmp[13]), "=f"(tmp[14]), "=f"(tmp[15]),- "=f"(tmp[16]), "=f"(tmp[17]), "=f"(tmp[18]), "=f"(tmp[19]), "=f"(tmp[20]), "=f"(tmp[21]), "=f"(tmp[22]), "=f"(tmp[23]),- "=f"(tmp[24]), "=f"(tmp[25]), "=f"(tmp[26]), "=f"(tmp[27]), "=f"(tmp[28]), "=f"(tmp[29]), "=f"(tmp[30]), "=f"(tmp[31])- : "r"((row << 16) | col), "C"(SHAPE_), "C"(NUM_)+ : "=f"(out[ 0]), "=f"(out[ 1]), "=f"(out[ 2]), "=f"(out[ 3]), "=f"(out[ 4]), "=f"(out[ 5]), "=f"(out[ 6]), "=f"(out[ 7]),+ "=f"(out[ 8]), "=f"(out[ 9]), "=f"(out[10]), "=f"(out[11]), "=f"(out[12]), "=f"(out[13]), "=f"(out[14]), "=f"(out[15]),+ "=f"(out[16]), "=f"(out[17]), "=f"(out[18]), "=f"(out[19]), "=f"(out[20]), "=f"(out[21]), "=f"(out[22]), "=f"(out[23]),+ "=f"(out[24]), "=f"(out[25]), "=f"(out[26]), "=f"(out[27]), "=f"(out[28]), "=f"(out[29]), "=f"(out[30]), "=f"(out[31])+ : "r"(addr), "C"(SH), "C"(NM));}- __device__ __forceinline__ void tcgen05_ld_16x256bx8(float *tmp, int row, int col) {- tcgen05_ld_32regs<SHAPE::_16x256b, NUM::x8>(tmp, row, col);- }+ __device__ __forceinline__ void tc_ld_16x256bx8(float* out, uint32_t addr) { tc_ld32<_TC_SH::_16x256b, _TC_NM::x8>(out, addr); }+ __device__ __forceinline__ void tc_ld_16x256bx4(float* out, uint32_t addr) { tc_ld16<_TC_SH::_16x256b, _TC_NM::x4>(out, addr); }static inline void ck_cu(CUresult err) {if (err == CUDA_SUCCESS) return;- const char *msg = nullptr;+ const char* msg = nullptr;if (cuGetErrorString(err, &msg) != CUDA_SUCCESS) msg = "cu err";TORCH_CHECK(false, msg);}- static inline void init_AB_tmap(- CUtensorMap *tmap,- const char *ptr,- uint64_t global_h, uint64_t global_w,- uint32_t shared_h, uint32_t shared_w+ static inline void ck_cuda(cudaError_t err) {+ if (err == cudaSuccess) return;+ const char* msg = cudaGetErrorString(err);+ TORCH_CHECK(false, msg ? msg : "cuda err");+ }++ static inline int sm_count_cached() {+ static int sm = -1;+ if (sm > 0) return sm;+ int dev = 0;+ ck_cuda(cudaGetDevice(&dev));+ cudaDeviceProp prop;+ ck_cuda(cudaGetDeviceProperties(&prop, dev));+ sm = prop.multiProcessorCount;+ return sm;+ }++ static inline void encode_tmap(+ CUtensorMap* tmap,+ const char* ptr,+ uint64_t h,+ uint64_t w,+ uint32_t sh,+ uint32_t sw,+ CUtensorMapL2promotion promo) {constexpr uint32_t rank = 3;- uint64_t globalDim[rank] = {256, global_h, global_w / 256};- uint64_t globalStrides[rank-1] = {global_w / 2, 128};- uint32_t boxDim[rank] = {256, shared_h, shared_w / 256};- uint32_t elementStrides[rank] = {1, 1, 1};-+ uint64_t gdim[rank] = {256, h, w / 256};+ uint64_t gstride[rank-1] = {w / 2, 128};+ uint32_t bdim[rank] = {256, sh, sw / 256};+ uint32_t estride[rank] = {1, 1, 1};auto err = cuTensorMapEncodeTiled(tmap,CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,rank,- (void *)ptr,- globalDim,- globalStrides,- boxDim,- elementStrides,+ (void*)ptr,+ gdim,+ gstride,+ bdim,+ estride,CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE,CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B,- CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE,+ promo,CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);ck_cu(err);}- template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES>- __global__ __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)- void gemm_f32_kernel(- const __grid_constant__ CUtensorMap A_tmap,- const __grid_constant__ CUtensorMap B_tmap,- const char *SFA_ptr,- const char *SFB_ptr,- float *C_ptr,- int M, int N- ) {- const int tid = threadIdx.x;- const int bid = blockIdx.y;+ struct CacheSel {+ uint64_t a_data;+ uint64_t a_sf;+ uint64_t b_data;+ uint64_t b_sf;+ };- const int lane_id = tid & 31;- const int warp_id = tid >> 5;+ template <int K, int POLICY>+ __device__ __forceinline__ CacheSel cache_sel() {+ const uint64_t p_first = l2_policy_first();+ const uint64_t p_last = l2_policy_last();- const int grid_m = M / BLOCK_M;- const int grid_n = N / BLOCK_N;- const int bid_m = bid / grid_n;- const int bid_n = bid - bid_m * grid_n;-- const int off_m = bid_m * BLOCK_M;- const int off_n = bid_n * BLOCK_N;-- constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;-- extern __shared__ __align__(1024) char smem_ptr[];- const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));- constexpr int A_size = BLOCK_M * BLOCK_K / 2;- constexpr int B_size = BLOCK_N * BLOCK_K / 2;- constexpr int SFA_size = 128 * BLOCK_K / 16;- constexpr int SFB_size = 128 * BLOCK_K / 16;- constexpr int STAGE_SIZE = A_size + B_size + SFA_size + SFB_size;-- #pragma nv_diag_suppress static_var_with_dynamic_init- __shared__ int64_t mbars[NUM_STAGES * 2 + 1];- const int tma_mbar_addr = static_cast<int>(__cvta_generic_to_shared(mbars));- const int mma_mbar_addr = tma_mbar_addr + NUM_STAGES * 8;- const int mainloop_mbar_addr = mma_mbar_addr + NUM_STAGES * 8;-- constexpr int SFA_tmem = BLOCK_N;- constexpr int SFB_tmem = SFA_tmem + 4 * (BLOCK_K / MMA_K);-- if (warp_id == 0 && elect_sync()) {- for (int i = 0; i < NUM_STAGES * 2 + 1; i++) mbarrier_init(tma_mbar_addr + i * 8, 1);- asm volatile("fence.mbarrier_init.release.cluster;");- } else if (warp_id == 1) {- asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(smem), "r"(BLOCK_N * 2));- }- __syncthreads();-- constexpr int num_iters = K / BLOCK_K;-- if (warp_id == NUM_WARPS - 2 && elect_sync()) {- const uint64_t cache_A = EVICT_LAST;- const uint64_t cache_B = EVICT_FIRST;-- auto issue_tma = [&](int iter_k, int stage_id) {- const int mbar_addr = tma_mbar_addr + stage_id * 8;- const int A_smem = smem + stage_id * STAGE_SIZE;- const int B_smem = A_smem + A_size;- const int SFA_smem = B_smem + B_size;- const int SFB_smem = SFA_smem + SFA_size;-- const int off_k = iter_k * BLOCK_K;- tma_3d_gmem2smem(A_smem, &A_tmap, 0, off_m, off_k / 256, mbar_addr, cache_A);- tma_3d_gmem2smem(B_smem, &B_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);-- const int rest_k = K / 16 / 4;- const char *SFA_src = SFA_ptr + ((off_m / 128) * rest_k + off_k / (16 * 4)) * 512;- const char *SFB_src = SFB_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;- tma_gmem2smem(SFA_smem, SFA_src, SFA_size, mbar_addr, cache_A);- tma_gmem2smem(SFB_smem, SFB_src, SFB_size, mbar_addr, cache_B);-- asm volatile(- "mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"- :: "r"(mbar_addr), "r"(STAGE_SIZE)- : "memory"- );- };-- constexpr int PRELOAD = (num_iters < NUM_STAGES) ? num_iters : NUM_STAGES;- for (int iter_k = 0; iter_k < PRELOAD; iter_k++) issue_tma(iter_k, iter_k);- for (int iter_k = NUM_STAGES; iter_k < num_iters; iter_k++) {- const int stage_id = iter_k % NUM_STAGES;- const int mma_phase = (iter_k / NUM_STAGES - 1) & 1;- mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);- issue_tma(iter_k, stage_id);- }- } else if (warp_id == NUM_WARPS - 1 && elect_sync()) {- constexpr int MMA_N = BLOCK_N;- constexpr int MMA_M = 128;- constexpr uint32_t i_desc = (1U << 7U) | (1U << 10U) | ((uint32_t)MMA_N >> 3U << 17U) | ((uint32_t)MMA_M >> 7U << 27U);-- for (int iter_k = 0; iter_k < num_iters; iter_k++) {- const int stage_id = iter_k % NUM_STAGES;- const int tma_phase = (iter_k / NUM_STAGES) & 1;- mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase);-- const int A_smem = smem + stage_id * STAGE_SIZE;- const int B_smem = A_smem + A_size;- const int SFA_smem = B_smem + B_size;- const int SFB_smem = SFA_smem + SFA_size;-- auto make_desc_AB = [](int addr) -> uint64_t {- const int SBO = 8 * 128;- return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);- };- auto make_desc_SF = [](int addr) -> uint64_t {- const int SBO = 8 * 16;- return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);- };-- constexpr uint64_t SF_desc = make_desc_SF(0);- const uint64_t SFA_desc = SF_desc + ((uint64_t)SFA_smem >> 4ULL);- const uint64_t SFB_desc = SF_desc + ((uint64_t)SFB_smem >> 4ULL);-- for (int k = 0; k < BLOCK_K / MMA_K; k++) {- uint64_t sfa_desc = SFA_desc + (uint64_t)k * (512ULL >> 4ULL);- uint64_t sfb_desc = SFB_desc + (uint64_t)k * (512ULL >> 4ULL);- tcgen05_cp_nvfp4(SFA_tmem + k * 4, sfa_desc);- tcgen05_cp_nvfp4(SFB_tmem + k * 4, sfb_desc);+ CacheSel r;+ if constexpr (K == 4096 || K == 7168) {+ if constexpr (POLICY <= 2) {+ if constexpr (POLICY == 1) {+ r.a_data = p_last; r.a_sf = p_last;+ r.b_data = p_first; r.b_sf = p_first;+ } else if constexpr (POLICY == 2) {+ r.a_data = p_first; r.a_sf = p_first;+ r.b_data = p_last; r.b_sf = p_last;+ } else {+ r.a_data = p_first; r.a_sf = p_last;+ r.b_data = p_first; r.b_sf = p_first;}-- for (int k1 = 0; k1 < BLOCK_K / 256; k1++)- for (int k2 = 0; k2 < 256 / MMA_K; k2++) {- uint64_t a_desc = make_desc_AB(A_smem + k1 * BLOCK_M * 128 + k2 * 32);- uint64_t b_desc = make_desc_AB(B_smem + k1 * BLOCK_N * 128 + k2 * 32);-- const int k_sf = k1 * 4 + k2;- const int scale_A_tmem = SFA_tmem + k_sf * 4;- const int scale_B_tmem = SFB_tmem + k_sf * 4 + (bid_n % (128 / BLOCK_N)) * (BLOCK_N / 32);-- const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;- tcgen05_mma_nvfp4(a_desc, b_desc, i_desc, scale_A_tmem, scale_B_tmem, enable_input_d);- }-- asm volatile(- "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"- :: "r"(mma_mbar_addr + stage_id * 8)- : "memory"- );+ } else {+ static_assert(POLICY <= (3 + 15));+ constexpr int mask = POLICY - 3;+ r.a_data = (mask & 0x1) ? p_last : p_first;+ r.a_sf = (mask & 0x2) ? p_last : p_first;+ r.b_data = (mask & 0x4) ? p_last : p_first;+ r.b_sf = (mask & 0x8) ? p_last : p_first;}-- asm volatile(- "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"- :: "r"(mainloop_mbar_addr)- : "memory"- );- } else if (tid < BLOCK_M) {- mbarrier_wait(mainloop_mbar_addr, 0);- asm volatile("tcgen05.fence::after_thread_sync;");-- for (int mm = 0; mm < 2; mm++) {- float tmp[BLOCK_N / 2];- tcgen05_ld_16x256bx8(tmp, warp_id * 32 + mm * 16, 0);- asm volatile("tcgen05.wait::ld.sync.aligned;");-- #pragma unroll- for (int i = 0; i < BLOCK_N / 8; i++) {- const int row = off_m + warp_id * 32 + mm * 16 + lane_id / 4;- const int col = off_n + i * 8 + (lane_id & 3) * 2;- reinterpret_cast<float2 *>(C_ptr + (row + 0) * N + col)[0] = float2{tmp[i * 4 + 0], tmp[i * 4 + 1]};- reinterpret_cast<float2 *>(C_ptr + (row + 8) * N + col)[0] = float2{tmp[i * 4 + 2], tmp[i * 4 + 3]};- }- }-- asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");- if (warp_id == 0) asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(BLOCK_N * 2));+ } else {+ r.a_data = p_first; r.a_sf = p_first;+ r.b_data = p_last; r.b_sf = p_last;}+ return r;}- template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES>- __global__ __launch_bounds__(BLOCK_M + 2 * WARP_SIZE)- void gemm_silu_mul_kernel(- const __grid_constant__ CUtensorMap A_tmap,- const __grid_constant__ CUtensorMap B_tmap,- const char *SFA_ptr,- const char *SFB_ptr,- const float *G1_ptr,- half *Out_ptr,- int M, int N+ template <int K, int BM, int BN, int BK, int STAGES, int POLICY, int CTA_N_MAJOR, int FAST_SILU, int EPILOGUE_SEG, int FUSE_CP_MMA>+ __global__ __launch_bounds__(BM + 2 * WARP_SZ, 1)+ void kernel_dual_fused(+ const __grid_constant__ CUtensorMap A_t,+ const __grid_constant__ CUtensorMap B1_t,+ const __grid_constant__ CUtensorMap B2_t,+ const char* __restrict__ SFA,+ const char* __restrict__ SFB1,+ const char* __restrict__ SFB2,+ half* __restrict__ OUT,+ int M,+ int N) {- const int tid = threadIdx.x;- const int bid = blockIdx.y;+ const int tid = (int)threadIdx.x;+ const int lane = tid & 31;+ const int warp = tid >> 5;- const int lane_id = tid & 31;- const int warp_id = tid >> 5;+ int bid_m;+ int bid_n;+ if constexpr (K == 7168) {+ if constexpr (CTA_N_MAJOR) { bid_n = (int)blockIdx.x; bid_m = (int)blockIdx.y; }+ else { bid_m = (int)blockIdx.x; bid_n = (int)blockIdx.y; }+ } else {+ bid_n = (int)blockIdx.x;+ bid_m = (int)blockIdx.y;+ }- const int grid_m = M / BLOCK_M;- const int grid_n = N / BLOCK_N;- const int bid_m = bid / grid_n;- const int bid_n = bid - bid_m * grid_n;+ const int off_m = bid_m * BM;+ const int off_n = bid_n * BN;- const int off_m = bid_m * BLOCK_M;- const int off_n = bid_n * BLOCK_N;+ constexpr int WARP_CNT = BM / WARP_SZ + 2;- constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;+ extern __shared__ __align__(1024) char smem_raw[];+ const int smem = static_cast<int>(__cvta_generic_to_shared(smem_raw));- extern __shared__ __align__(1024) char smem_ptr[];- const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));- constexpr int A_size = BLOCK_M * BLOCK_K / 2;- constexpr int B_size = BLOCK_N * BLOCK_K / 2;- constexpr int SFA_size = 128 * BLOCK_K / 16;- constexpr int SFB_size = 128 * BLOCK_K / 16;- constexpr int STAGE_SIZE = A_size + B_size + SFA_size + SFB_size;+ constexpr int A_BYTES = BM * BK / 2;+ constexpr int B_BYTES = BN * BK / 2;+ constexpr int SFA_BYTES = 128 * BK / 16;+ constexpr int SFB_BYTES = 128 * BK / 16;+ constexpr int STAGE_BYTES = A_BYTES + 2 * B_BYTES + SFA_BYTES + 2 * SFB_BYTES;+ constexpr int TMEM_NEED = 2 * BN + 12 * (BK / MMA_K64);+ constexpr int TMEM_COLS = (TMEM_NEED <= 256) ? 256 : 512;+ static_assert(TMEM_NEED <= 512);+#pragma nv_diag_suppress static_var_with_dynamic_init- __shared__ int64_t mbars[NUM_STAGES * 2 + 1];- const int tma_mbar_addr = static_cast<int>(__cvta_generic_to_shared(mbars));- const int mma_mbar_addr = tma_mbar_addr + NUM_STAGES * 8;- const int mainloop_mbar_addr = mma_mbar_addr + NUM_STAGES * 8;+ __shared__ int64_t mbars[STAGES * 2 + 1];- constexpr int SFA_tmem = BLOCK_N;- constexpr int SFB_tmem = SFA_tmem + 4 * (BLOCK_K / MMA_K);+ const int tma_mbar = static_cast<int>(__cvta_generic_to_shared(mbars));+ const int mma_mbar = tma_mbar + STAGES * 8;+ const int main_mbar = mma_mbar + STAGES * 8;- if (warp_id == 0 && elect_sync()) {- for (int i = 0; i < NUM_STAGES * 2 + 1; i++) mbarrier_init(tma_mbar_addr + i * 8, 1);+ if (warp == 0 && elect_one()) {+ #pragma unroll+ for (int i = 0; i < STAGES * 2 + 1; ++i) mbar_init_shared(tma_mbar + i * 8, 1);asm volatile("fence.mbarrier_init.release.cluster;");- } else if (warp_id == 1) {- asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(smem), "r"(BLOCK_N * 2));}++ if (warp == 1) {+ asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(smem), "r"(TMEM_COLS));+ }__syncthreads();- constexpr int num_iters = K / BLOCK_K;+ constexpr uint32_t tmem_base = 0;+ constexpr uint32_t out1_col = 0;+ constexpr uint32_t out2_col = (uint32_t)BN;+ constexpr uint32_t sfa_col = (uint32_t)(2 * BN);+ constexpr uint32_t sfb1_col = (uint32_t)(2 * BN + 4 * (BK / MMA_K64));+ constexpr uint32_t sfb2_col = (uint32_t)(2 * BN + 8 * (BK / MMA_K64));- if (warp_id == NUM_WARPS - 2 && elect_sync()) {- const uint64_t cache_A = EVICT_LAST;- const uint64_t cache_B = EVICT_FIRST;+ constexpr int iters = K / BK;- auto issue_tma = [&](int iter_k, int stage_id) {- const int mbar_addr = tma_mbar_addr + stage_id * 8;- const int A_smem = smem + stage_id * STAGE_SIZE;- const int B_smem = A_smem + A_size;- const int SFA_smem = B_smem + B_size;- const int SFB_smem = SFA_smem + SFA_size;+ if (warp == WARP_CNT - 1 && elect_one()) {+ const CacheSel pol = cache_sel<K, POLICY>();+ constexpr int REST_K = K / 16 / 4;+ constexpr int SF_STEP = (BK / (16 * 4)) * 512;+ constexpr int Z_STEP = (BK / 256);- const int off_k = iter_k * BLOCK_K;- tma_3d_gmem2smem(A_smem, &A_tmap, 0, off_m, off_k / 256, mbar_addr, cache_A);- tma_3d_gmem2smem(B_smem, &B_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);+ const int off_m128 = off_m >> 7;+ const int off_n128 = off_n >> 7;+ const char* sfa_src = SFA + off_m128 * REST_K * 512;+ const char* sfb1_src = SFB1 + off_n128 * REST_K * 512;+ const char* sfb2_src = SFB2 + off_n128 * REST_K * 512;- const int rest_k = K / 16 / 4;- const char *SFA_src = SFA_ptr + ((off_m / 128) * rest_k + off_k / (16 * 4)) * 512;- const char *SFB_src = SFB_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512;- tma_gmem2smem(SFA_smem, SFA_src, SFA_size, mbar_addr, cache_A);- tma_gmem2smem(SFB_smem, SFB_src, SFB_size, mbar_addr, cache_B);+ int stage = 0;+ int wraps = 0;+ int stage_base = smem;+ int z = 0;+ for (int iter = 0; iter < iters; ++iter) {+ if (iter >= STAGES) {+ mbar_wait_parity(mma_mbar + stage * 8, (wraps - 1) & 1);+ }+ const int mbar = tma_mbar + stage * 8;+ const int a_s = stage_base;+ const int b1_s = a_s + A_BYTES;+ const int b2_s = b1_s + B_BYTES;+ const int sfa_s = b2_s + B_BYTES;+ const int sfb1_s = sfa_s + SFA_BYTES;+ const int sfb2_s = sfb1_s + SFB_BYTES;++ tma_g2s_3d(a_s, &A_t, 0, off_m, z, mbar, pol.a_data);+ tma_g2s_3d(b1_s, &B1_t, 0, off_n, z, mbar, pol.b_data);+ tma_g2s_3d(b2_s, &B2_t, 0, off_n, z, mbar, pol.b_data);++ tma_g2s_bytes(sfa_s, sfa_src, SFA_BYTES, mbar, pol.a_sf);+ tma_g2s_bytes(sfb1_s, sfb1_src, SFB_BYTES, mbar, pol.b_sf);+ tma_g2s_bytes(sfb2_s, sfb2_src, SFB_BYTES, mbar, pol.b_sf);+asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"- :: "r"(mbar_addr), "r"(STAGE_SIZE)+ :: "r"(mbar), "r"(STAGE_BYTES): "memory");- };- constexpr int PRELOAD = (num_iters < NUM_STAGES) ? num_iters : NUM_STAGES;- for (int iter_k = 0; iter_k < PRELOAD; iter_k++) issue_tma(iter_k, iter_k);- for (int iter_k = NUM_STAGES; iter_k < num_iters; iter_k++) {- const int stage_id = iter_k % NUM_STAGES;- const int mma_phase = (iter_k / NUM_STAGES - 1) & 1;- mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);- issue_tma(iter_k, stage_id);+ z += Z_STEP;+ sfa_src += SF_STEP;+ sfb1_src += SF_STEP;+ sfb2_src += SF_STEP;++ ++stage;+ stage_base += STAGE_BYTES;+ if (stage == STAGES) {+ stage = 0;+ stage_base = smem;+ ++wraps;+ }}- } else if (warp_id == NUM_WARPS - 1 && elect_sync()) {- constexpr int MMA_N = BLOCK_N;+ } else if (warp == WARP_CNT - 2 && elect_one()) {+ constexpr int MMA_N = BN;constexpr int MMA_M = 128;- constexpr uint32_t i_desc = (1U << 7U) | (1U << 10U) | ((uint32_t)MMA_N >> 3U << 17U) | ((uint32_t)MMA_M >> 7U << 27U);+ constexpr uint32_t idesc =+ (1U << 7U) | (1U << 10U) | ((uint32_t)MMA_N >> 3U << 17U) | ((uint32_t)MMA_M >> 7U << 27U);- for (int iter_k = 0; iter_k < num_iters; iter_k++) {- const int stage_id = iter_k % NUM_STAGES;- const int tma_phase = (iter_k / NUM_STAGES) & 1;- mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase);+ auto desc_ab = [](int addr) -> uint64_t {+ const int sbo = 8 * 128;+ return desc_pack(addr) | (desc_pack(sbo) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);+ };+ auto desc_sf = [](int addr) -> uint64_t {+ const int sbo = 8 * 16;+ return desc_pack(addr) | (desc_pack(sbo) << 32ULL) | (1ULL << 46ULL);+ };+ const uint64_t ab0 = desc_ab(0);+ const uint64_t sf0 = desc_sf(0);+ constexpr uint32_t SB_PARTS = (uint32_t)(128 / BN);+ static_assert((128 % BN) == 0);+ const uint32_t sb_off =+ (SB_PARTS == 1) ? 0U : ((uint32_t)bid_n & (SB_PARTS - 1U)) * (uint32_t)(BN / 32);- const int A_smem = smem + stage_id * STAGE_SIZE;- const int B_smem = A_smem + A_size;- const int SFA_smem = B_smem + B_size;- const int SFB_smem = SFA_smem + SFA_size;+ int stage = 0;+ int phase = 0;+ int stage_base = smem;+ for (int iter = 0; iter < iters; ++iter) {+ mbar_wait_parity(tma_mbar + stage * 8, phase);- auto make_desc_AB = [](int addr) -> uint64_t {- const int SBO = 8 * 128;- return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);- };- auto make_desc_SF = [](int addr) -> uint64_t {- const int SBO = 8 * 16;- return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);- };+ const int a_s = stage_base;+ const int b1_s = a_s + A_BYTES;+ const int b2_s = b1_s + B_BYTES;+ const int sfa_s = b2_s + B_BYTES;+ const int sfb1_s = sfa_s + SFA_BYTES;+ const int sfb2_s = sfb1_s + SFB_BYTES;- constexpr uint64_t SF_desc = make_desc_SF(0);- const uint64_t SFA_desc = SF_desc + ((uint64_t)SFA_smem >> 4ULL);- const uint64_t SFB_desc = SF_desc + ((uint64_t)SFB_smem >> 4ULL);+ const uint64_t sfa_desc0 = sf0 + ((uint64_t)sfa_s >> 4ULL);+ const uint64_t sfb1_desc0 = sf0 + ((uint64_t)sfb1_s >> 4ULL);+ const uint64_t sfb2_desc0 = sf0 + ((uint64_t)sfb2_s >> 4ULL);- for (int k = 0; k < BLOCK_K / MMA_K; k++) {- uint64_t sfa_desc = SFA_desc + (uint64_t)k * (512ULL >> 4ULL);- uint64_t sfb_desc = SFB_desc + (uint64_t)k * (512ULL >> 4ULL);- tcgen05_cp_nvfp4(SFA_tmem + k * 4, sfa_desc);- tcgen05_cp_nvfp4(SFB_tmem + k * 4, sfb_desc);- }+ const uint64_t a_base = ab0 + ((uint64_t)a_s >> 4ULL);+ const uint64_t b1_base = ab0 + ((uint64_t)b1_s >> 4ULL);+ const uint64_t b2_base = ab0 + ((uint64_t)b2_s >> 4ULL);+ if constexpr (FUSE_CP_MMA) {+ static_assert(BK == 256);+ uint32_t td_sfa = tmem_base + sfa_col;+ uint32_t td_sfb1 = tmem_base + sfb1_col;+ uint32_t td_sfb2 = tmem_base + sfb2_col;+ uint32_t sa = tmem_base + sfa_col;+ uint32_t sb1 = tmem_base + sfb1_col + sb_off;+ uint32_t sb2 = tmem_base + sfb2_col + sb_off;+ uint64_t sfa_desc = sfa_desc0;+ uint64_t sfb1_desc = sfb1_desc0;+ uint64_t sfb2_desc = sfb2_desc0;+ uint64_t a_desc = a_base;+ uint64_t b1_desc = b1_base;+ uint64_t b2_desc = b2_base;+ #pragma unroll+ for (int k2 = 0; k2 < BK / MMA_K64; ++k2) {+ tc_scale_cp(td_sfa, sfa_desc);+ tc_scale_cp(td_sfb1, sfb1_desc);+ tc_scale_cp(td_sfb2, sfb2_desc);+ const int en = (k2 == 0) ? iter : 1;+ tc_mma_a_fill(tmem_base + out1_col, a_desc, b1_desc, idesc, sa, sb1, en);+ tc_mma_a_last(tmem_base + out2_col, a_desc, b2_desc, idesc, sa, sb2, en);+ a_desc += 2ULL;+ b1_desc += 2ULL;+ b2_desc += 2ULL;+ sfa_desc += (512ULL >> 4ULL);+ sfb1_desc += (512ULL >> 4ULL);+ sfb2_desc += (512ULL >> 4ULL);+ td_sfa += 4U;+ td_sfb1 += 4U;+ td_sfb2 += 4U;+ sa += 4U;+ sb1 += 4U;+ sb2 += 4U;+ }+ } else {+ uint32_t td_sfa = tmem_base + sfa_col;+ uint32_t td_sfb1 = tmem_base + sfb1_col;+ uint32_t td_sfb2 = tmem_base + sfb2_col;+ uint64_t sfa_desc = sfa_desc0;+ uint64_t sfb1_desc = sfb1_desc0;+ uint64_t sfb2_desc = sfb2_desc0;+ #pragma unroll+ for (int kk = 0; kk < BK / MMA_K64; ++kk) {+ tc_scale_cp(td_sfa, sfa_desc);+ tc_scale_cp(td_sfb1, sfb1_desc);+ tc_scale_cp(td_sfb2, sfb2_desc);+ sfa_desc += (512ULL >> 4ULL);+ sfb1_desc += (512ULL >> 4ULL);+ sfb2_desc += (512ULL >> 4ULL);+ td_sfa += 4U;+ td_sfb1 += 4U;+ td_sfb2 += 4U;+ }- for (int k1 = 0; k1 < BLOCK_K / 256; k1++)- for (int k2 = 0; k2 < 256 / MMA_K; k2++) {- uint64_t a_desc = make_desc_AB(A_smem + k1 * BLOCK_M * 128 + k2 * 32);- uint64_t b_desc = make_desc_AB(B_smem + k1 * BLOCK_N * 128 + k2 * 32);+ uint32_t sa = tmem_base + sfa_col;+ uint32_t sb1 = tmem_base + sfb1_col + sb_off;+ uint32_t sb2 = tmem_base + sfb2_col + sb_off;- const int k_sf = k1 * 4 + k2;- const int scale_A_tmem = SFA_tmem + k_sf * 4;- const int scale_B_tmem = SFB_tmem + k_sf * 4 + (bid_n % (128 / BLOCK_N)) * (BLOCK_N / 32);-- const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;- tcgen05_mma_nvfp4(a_desc, b_desc, i_desc, scale_A_tmem, scale_B_tmem, enable_input_d);+ if constexpr (BK == 256) {+ uint64_t a_desc = a_base;+ uint64_t b1_desc = b1_base;+ uint64_t b2_desc = b2_base;+ #pragma unroll+ for (int k2 = 0; k2 < BK / MMA_K64; ++k2) {+ const int en = (k2 == 0) ? iter : 1;+ tc_mma_a_fill(tmem_base + out1_col, a_desc, b1_desc, idesc, sa, sb1, en);+ tc_mma_a_last(tmem_base + out2_col, a_desc, b2_desc, idesc, sa, sb2, en);+ a_desc += 2ULL;+ b1_desc += 2ULL;+ b2_desc += 2ULL;+ sa += 4U;+ sb1 += 4U;+ sb2 += 4U;+ }+ } else {+ constexpr uint64_t A_STEP = ((uint64_t)BM * 128ULL) >> 4ULL;+ constexpr uint64_t B_STEP = ((uint64_t)BN * 128ULL) >> 4ULL;+ #pragma unroll+ for (int k2 = 0; k2 < BK / MMA_K64; ++k2) {+ const int en = (k2 == 0) ? iter : 1;+ const int k1 = k2 >> 2;+ const int kk = k2 & 3;+ const uint64_t a_desc = a_base + (uint64_t)k1 * A_STEP + (uint64_t)kk * 2ULL;+ const uint64_t b1_desc = b1_base + (uint64_t)k1 * B_STEP + (uint64_t)kk * 2ULL;+ const uint64_t b2_desc = b2_base + (uint64_t)k1 * B_STEP + (uint64_t)kk * 2ULL;+ tc_mma_a_fill(tmem_base + out1_col, a_desc, b1_desc, idesc, sa, sb1, en);+ tc_mma_a_last(tmem_base + out2_col, a_desc, b2_desc, idesc, sa, sb2, en);+ sa += 4U;+ sb1 += 4U;+ sb2 += 4U;+ }}+ }asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"- :: "r"(mma_mbar_addr + stage_id * 8)+ :: "r"(mma_mbar + stage * 8): "memory");++ ++stage;+ stage_base += STAGE_BYTES;+ if (stage == STAGES) {+ stage = 0;+ stage_base = smem;+ phase ^= 1;+ }}asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"- :: "r"(mainloop_mbar_addr)+ :: "r"(main_mbar): "memory");- } else if (tid < BLOCK_M) {- mbarrier_wait(mainloop_mbar_addr, 0);+ } else if (tid < BM) {+ mbar_wait_parity(main_mbar, 0);asm volatile("tcgen05.fence::after_thread_sync;");- for (int mm = 0; mm < 2; mm++) {- float tmp[BLOCK_N / 2];- tcgen05_ld_16x256bx8(tmp, warp_id * 32 + mm * 16, 0);- asm volatile("tcgen05.wait::ld.sync.aligned;");+ constexpr float LOG2E = 1.4426950408889634f;+ const int lane_row = lane >> 2;+ const int lane_h2 = lane & 3;+ constexpr int SEG = EPILOGUE_SEG;+ static_assert((BN % SEG) == 0);+ constexpr int SEGS = BN / SEG;++ #pragma unroll+ for (int mm = 0; mm < 2; ++mm) {+ const int row0 = warp * 32 + mm * 16;+ const uint32_t addr_x0 = tmem_base + (uint32_t)((row0 << 16) | (int)out1_col);+ const uint32_t addr_y0 = tmem_base + (uint32_t)((row0 << 16) | (int)out2_col);+ const int row = off_m + row0 + lane_row;+ half2* out_row0 = reinterpret_cast<half2*>(OUT + row * N);+ half2* out_row8 = reinterpret_cast<half2*>(OUT + (row + 8) * N);+#pragma unroll- for (int i = 0; i < BLOCK_N / 8; i++) {- const int row = off_m + warp_id * 32 + mm * 16 + lane_id / 4;- const int col = off_n + i * 8 + (lane_id & 3) * 2;- const float2 x0 = reinterpret_cast<const float2 *>(G1_ptr + (row + 0) * N + col)[0];- const float2 x8 = reinterpret_cast<const float2 *>(G1_ptr + (row + 8) * N + col)[0];- const float2 y0 = float2{tmp[i * 4 + 0], tmp[i * 4 + 1]};- const float2 y8 = float2{tmp[i * 4 + 2], tmp[i * 4 + 3]};+ for (int seg = 0; seg < SEGS; ++seg) {+ float x[SEG / 2];+ float y[SEG / 2];+ const uint32_t col_off = (uint32_t)(seg * SEG);+ if constexpr (SEG == 32) {+ tc_ld_16x256bx4(x, addr_x0 + col_off);+ tc_ld_16x256bx4(y, addr_y0 + col_off);+ } else if constexpr (SEG == 64) {+ tc_ld_16x256bx8(x, addr_x0 + col_off);+ tc_ld_16x256bx8(y, addr_y0 + col_off);+ } else {+ static_assert(SEG == 32 || SEG == 64);+ }+ const int col_base_h2 = (off_n >> 1) + seg * (SEG >> 1);+ asm volatile("tcgen05.wait::ld.sync.aligned;");- float2 o0;- float2 o8;- const float s00 = 1.0f / (1.0f + __expf(-x0.x));- const float s01 = 1.0f / (1.0f + __expf(-x0.y));- const float s80 = 1.0f / (1.0f + __expf(-x8.x));- const float s81 = 1.0f / (1.0f + __expf(-x8.y));- o0.x = (x0.x * s00) * y0.x;- o0.y = (x0.y * s01) * y0.y;- o8.x = (x8.x * s80) * y8.x;- o8.y = (x8.y * s81) * y8.y;+ #pragma unroll+ for (int i = 0; i < SEG / 8; ++i) {+ const int out_col_base = col_base_h2 + i * 4;- reinterpret_cast<half2 *>(Out_ptr + (row + 0) * N + col)[0] = __float22half2_rn(o0);- reinterpret_cast<half2 *>(Out_ptr + (row + 8) * N + col)[0] = __float22half2_rn(o8);+ const float x00 = x[i * 4 + 0];+ const float x01 = x[i * 4 + 1];+ const float x80 = x[i * 4 + 2];+ const float x81 = x[i * 4 + 3];+ const float y00 = y[i * 4 + 0];+ const float y01 = y[i * 4 + 1];+ const float y80 = y[i * 4 + 2];+ const float y81 = y[i * 4 + 3];++ float s00, s01, s80, s81;+ if constexpr (FAST_SILU) {+ float t00, t01, t80, t81;+ asm("ex2.approx.f32 %0, %1;" : "=f"(t00) : "f"((-x00) * LOG2E));+ asm("ex2.approx.f32 %0, %1;" : "=f"(t01) : "f"((-x01) * LOG2E));+ asm("ex2.approx.f32 %0, %1;" : "=f"(t80) : "f"((-x80) * LOG2E));+ asm("ex2.approx.f32 %0, %1;" : "=f"(t81) : "f"((-x81) * LOG2E));+ const float d00 = 1.0f + t00;+ const float d01 = 1.0f + t01;+ const float d80 = 1.0f + t80;+ const float d81 = 1.0f + t81;+ asm("rcp.approx.f32 %0, %1;" : "=f"(s00) : "f"(d00));+ asm("rcp.approx.f32 %0, %1;" : "=f"(s01) : "f"(d01));+ asm("rcp.approx.f32 %0, %1;" : "=f"(s80) : "f"(d80));+ asm("rcp.approx.f32 %0, %1;" : "=f"(s81) : "f"(d81));+ } else {+ s00 = __fdividef(1.0f, 1.0f + exp2f((-x00) * LOG2E));+ s01 = __fdividef(1.0f, 1.0f + exp2f((-x01) * LOG2E));+ s80 = __fdividef(1.0f, 1.0f + exp2f((-x80) * LOG2E));+ s81 = __fdividef(1.0f, 1.0f + exp2f((-x81) * LOG2E));+ }++ float2 o0;+ float2 o8;+ o0.x = (x00 * s00) * y00;+ o0.y = (x01 * s01) * y01;+ o8.x = (x80 * s80) * y80;+ o8.y = (x81 * s81) * y81;++ const int out_col_h2 = out_col_base + lane_h2;+ out_row0[out_col_h2] = __float22half2_rn(o0);+ out_row8[out_col_h2] = __float22half2_rn(o8);+ }}}- asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");- if (warp_id == 0) asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(BLOCK_N * 2));+ asm volatile("bar.sync 1, %0;" :: "r"(BM) : "memory");+ if (warp == 0) asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(TMEM_COLS));}}- template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES>- static inline void launch_gemm_f32(+ template <int K, int BM, int BN, int BK, int STAGES, int POLICY, int CTA_N_MAJOR, int FAST_SILU, int EPILOGUE_SEG, int FUSE_CP_MMA>+ static inline void launch_cfg(const at::Tensor& A,- const at::Tensor& B,+ const at::Tensor& B1,+ const at::Tensor& B2,const at::Tensor& SFA,- const at::Tensor& SFB,- at::Tensor& C+ const at::Tensor& SFB1,+ const at::Tensor& SFB2,+ at::Tensor& out) {- const int M = (int)A.size(0);- const int N = (int)B.size(0);+ auto call = [](const at::Tensor& A,+ const at::Tensor& B1,+ const at::Tensor& B2,+ const at::Tensor& SFA,+ const at::Tensor& SFB1,+ const at::Tensor& SFB2,+ at::Tensor& out,+ CUtensorMapL2promotion promo_a,+ CUtensorMapL2promotion promo_b) {+ const int M = (int)A.size(0);+ const int N = (int)B1.size(0);- auto A_ptr = reinterpret_cast<const char *>(A.data_ptr());- auto B_ptr = reinterpret_cast<const char *>(B.data_ptr());- auto SFA_ptr = reinterpret_cast<const char *>(SFA.data_ptr());- auto SFB_ptr = reinterpret_cast<const char *>(SFB.data_ptr());- auto C_ptr = reinterpret_cast<float *>(C.data_ptr());+ const char* A_ptr = reinterpret_cast<const char*>(A.data_ptr());+ const char* B1_ptr = reinterpret_cast<const char*>(B1.data_ptr());+ const char* B2_ptr = reinterpret_cast<const char*>(B2.data_ptr());+ const char* SFA_ptr = reinterpret_cast<const char*>(SFA.data_ptr());+ const char* SFB1_ptr = reinterpret_cast<const char*>(SFB1.data_ptr());+ const char* SFB2_ptr = reinterpret_cast<const char*>(SFB2.data_ptr());+ half* Out_ptr = reinterpret_cast<half*>(out.data_ptr());- CUtensorMap A_tmap, B_tmap;- init_AB_tmap(&A_tmap, A_ptr, (uint64_t)M, (uint64_t)K, (uint32_t)BLOCK_M, (uint32_t)BLOCK_K);- init_AB_tmap(&B_tmap, B_ptr, (uint64_t)N, (uint64_t)K, (uint32_t)BLOCK_N, (uint32_t)BLOCK_K);+ struct TmapCache {+ const char* a_ptr;+ const char* b1_ptr;+ const char* b2_ptr;+ int m;+ int n;+ CUtensorMap a_t;+ CUtensorMap b1_t;+ CUtensorMap b2_t;+ bool valid;+ };- dim3 grid(1, (unsigned)((M / BLOCK_M) * (N / BLOCK_N)));- const int tb_size = BLOCK_M + 2 * WARP_SIZE;- const int AB_size = (BLOCK_M + BLOCK_N) * (BLOCK_K / 2);- const int SFAB_size = 128 * (BLOCK_K / 16) * 2;- const int smem_size = (AB_size + SFAB_size) * NUM_STAGES;+ #pragma nv_diag_suppress static_var_with_dynamic_init+ static TmapCache cache = {nullptr, nullptr, nullptr, 0, 0, {}, {}, {}, false};- auto kptr = gemm_f32_kernel<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES>;- if (smem_size > 48'000) cudaFuncSetAttribute(kptr, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);- kptr<<<grid, tb_size, smem_size>>>(A_tmap, B_tmap, SFA_ptr, SFB_ptr, C_ptr, M, N);- }+ CUtensorMap A_t, B1_t, B2_t;+ if (cache.valid && cache.a_ptr == A_ptr && cache.b1_ptr == B1_ptr && cache.b2_ptr == B2_ptr && cache.m == M && cache.n == N) {+ A_t = cache.a_t;+ B1_t = cache.b1_t;+ B2_t = cache.b2_t;+ } else {+ encode_tmap(&A_t, A_ptr, (uint64_t)M, (uint64_t)K, (uint32_t)BM, (uint32_t)BK, promo_a);+ encode_tmap(&B1_t, B1_ptr, (uint64_t)N, (uint64_t)K, (uint32_t)BN, (uint32_t)BK, promo_b);+ encode_tmap(&B2_t, B2_ptr, (uint64_t)N, (uint64_t)K, (uint32_t)BN, (uint32_t)BK, promo_b);+ cache.a_ptr = A_ptr;+ cache.b1_ptr = B1_ptr;+ cache.b2_ptr = B2_ptr;+ cache.m = M;+ cache.n = N;+ cache.a_t = A_t;+ cache.b1_t = B1_t;+ cache.b2_t = B2_t;+ cache.valid = true;+ }- template <int K, int BLOCK_M, int BLOCK_N, int BLOCK_K, int NUM_STAGES>- static inline void launch_gemm_silu_mul(- const at::Tensor& A,- const at::Tensor& B,- const at::Tensor& SFA,- const at::Tensor& SFB,- const at::Tensor& g1,- at::Tensor& out- ) {- const int M = (int)A.size(0);- const int N = (int)B.size(0);+ dim3 grid;+ if constexpr (K == 7168) {+ if constexpr (CTA_N_MAJOR) grid = dim3((unsigned)(N / BN), (unsigned)(M / BM));+ else grid = dim3((unsigned)(M / BM), (unsigned)(N / BN));+ } else {+ grid = dim3((unsigned)(N / BN), (unsigned)(M / BM));+ }- auto A_ptr = reinterpret_cast<const char *>(A.data_ptr());- auto B_ptr = reinterpret_cast<const char *>(B.data_ptr());- auto SFA_ptr = reinterpret_cast<const char *>(SFA.data_ptr());- auto SFB_ptr = reinterpret_cast<const char *>(SFB.data_ptr());- auto G1_ptr = reinterpret_cast<const float *>(g1.data_ptr());- auto Out_ptr = reinterpret_cast<half *>(out.data_ptr());+ const int tb = BM + 2 * WARP_SZ;+ constexpr int A_BYTES = BM * BK / 2;+ constexpr int B_BYTES = BN * BK / 2;+ constexpr int SFA_BYTES = 128 * BK / 16;+ constexpr int SFB_BYTES = 128 * BK / 16;+ constexpr int smem_bytes = (A_BYTES + 2 * B_BYTES + SFA_BYTES + 2 * SFB_BYTES) * STAGES;+ constexpr int kMaxSmemBytes = 227 * 1024;+ TORCH_CHECK(smem_bytes <= kMaxSmemBytes, "smem ", smem_bytes);- CUtensorMap A_tmap, B_tmap;- init_AB_tmap(&A_tmap, A_ptr, (uint64_t)M, (uint64_t)K, (uint32_t)BLOCK_M, (uint32_t)BLOCK_K);- init_AB_tmap(&B_tmap, B_ptr, (uint64_t)N, (uint64_t)K, (uint32_t)BLOCK_N, (uint32_t)BLOCK_K);+ auto kptr = kernel_dual_fused<K, BM, BN, BK, STAGES, POLICY, CTA_N_MAJOR, FAST_SILU, EPILOGUE_SEG, FUSE_CP_MMA>;+ if constexpr (smem_bytes > 48000) {+ static bool attr_set = false;+ if (!attr_set) {+ ck_cuda(cudaFuncSetAttribute(kptr, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_bytes));+ attr_set = true;+ }+ }+ kptr<<<grid, tb, smem_bytes>>>(A_t, B1_t, B2_t, SFA_ptr, SFB1_ptr, SFB2_ptr, Out_ptr, M, N);+ };- dim3 grid(1, (unsigned)((M / BLOCK_M) * (N / BLOCK_N)));- const int tb_size = BLOCK_M + 2 * WARP_SIZE;- const int AB_size = (BLOCK_M + BLOCK_N) * (BLOCK_K / 2);- const int SFAB_size = 128 * (BLOCK_K / 16) * 2;- const int smem_size = (AB_size + SFAB_size) * NUM_STAGES;-- auto kptr = gemm_silu_mul_kernel<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES>;- if (smem_size > 48'000) cudaFuncSetAttribute(kptr, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);- kptr<<<grid, tb_size, smem_size>>>(A_tmap, B_tmap, SFA_ptr, SFB_ptr, G1_ptr, Out_ptr, M, N);+ call(A, B1, B2, SFA, SFB1, SFB2, out,+ CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_L2_128B,+ CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE);}- __global__ void silu_mul_f32_vec2(const float* __restrict__ x, const float* __restrict__ y, half* __restrict__ out, int64_t n2) {- const int64_t idx = int64_t(blockIdx.x) * blockDim.x + threadIdx.x;- if (idx >= n2) return;- const float2 fx = reinterpret_cast<const float2*>(x)[idx];- const float2 fy = reinterpret_cast<const float2*>(y)[idx];- float2 o;- const float sx0 = 1.0f / (1.0f + __expf(-fx.x));- const float sx1 = 1.0f / (1.0f + __expf(-fx.y));- o.x = (fx.x * sx0) * fy.x;- o.y = (fx.y * sx1) * fy.y;- reinterpret_cast<half2*>(out)[idx] = __float22half2_rn(o);+ template <int K, int BM, int BN, int BK, int STAGES, int FAST_SILU, int EPILOGUE_SEG, int FUSE_CP_MMA>+ static inline void launch_policy_auto(+ const at::Tensor& A,+ const at::Tensor& B1,+ const at::Tensor& B2,+ const at::Tensor& SFA,+ const at::Tensor& SFB1,+ const at::Tensor& SFB2,+ at::Tensor& out,+ int cta_m,+ int cta_n+ ) {+ if constexpr (K == 4096 || K == 7168) {+ if (cta_n >= (cta_m << 3)) launch_cfg<K, BM, BN, BK, STAGES, 1, 0, FAST_SILU, EPILOGUE_SEG, FUSE_CP_MMA>(A, B1, B2, SFA, SFB1, SFB2, out);+ else if (cta_m >= (cta_n << 3)) launch_cfg<K, BM, BN, BK, STAGES, 2, 0, FAST_SILU, EPILOGUE_SEG, FUSE_CP_MMA>(A, B1, B2, SFA, SFB1, SFB2, out);+ else launch_cfg<K, BM, BN, BK, STAGES, 0, 0, FAST_SILU, EPILOGUE_SEG, FUSE_CP_MMA>(A, B1, B2, SFA, SFB1, SFB2, out);+ } else {+ launch_cfg<K, BM, BN, BK, STAGES, 0, 0, FAST_SILU, EPILOGUE_SEG, FUSE_CP_MMA>(A, B1, B2, SFA, SFB1, SFB2, out);+ }}- static inline void launch_silu_mul_f32(const at::Tensor& g1, const at::Tensor& g2, at::Tensor& out) {- const int64_t n = out.numel();- TORCH_CHECK((n & 1) == 0, "n");- const int64_t n2 = n >> 1;- const int threads = 256;- const int blocks = (int)((n2 + threads - 1) / threads);- silu_mul_f32_vec2<<<blocks, threads>>>(- reinterpret_cast<const float*>(g1.data_ptr()),- reinterpret_cast<const float*>(g2.data_ptr()),- reinterpret_cast<half*>(out.data_ptr()),- n2- );+ static inline uint64_t shape_key_u64(int m, int n, int k) {+ return ((uint64_t)(uint32_t)k << 32) | ((uint64_t)(uint32_t)m << 16) | (uint64_t)(uint32_t)n;}+ #define KKEY(M, N, K) ((((uint64_t)(K)) << 32) | (((uint64_t)(M)) << 16) | ((uint64_t)(N)))+at::Tensor fused(const at::Tensor& A,const at::Tensor& B1,⋯ 1 unchanged linesconst at::Tensor& SFA,const at::Tensor& SFB1,const at::Tensor& SFB2,- at::Tensor& out,- at::Tensor& g1,- at::Tensor& g2+ at::Tensor& out) {⋯ diff truncated
scrolls · 1201 diff lines total
Best evidence level for this revision: reported
JSON