attn_res kernel latency improvements (#50185)

This commit is contained in:
gnovack
2026-08-06 08:53:20 -07:00
committed by GitHub
parent 41e7746b82
commit 7b4ed49628
+82 -37
View File
@@ -91,7 +91,8 @@ __device__ __forceinline__ void mbarrier_wait(uint64_t& barrier, int phase) {
"bra WAIT;\n"
"DONE:\n"
"}\n" ::"r"(barrier_addr),
"r"(phase));
"r"(phase)
: "memory");
}
__device__ __forceinline__ void mbarrier_arrive(uint64_t& barrier) {
@@ -101,7 +102,8 @@ __device__ __forceinline__ void mbarrier_arrive(uint64_t& barrier) {
"{\n"
".reg .b64 state;\n"
"mbarrier.arrive.shared::cta.b64 state, [%0];\n"
"}\n" ::"r"(barrier_addr));
"}\n" ::"r"(barrier_addr)
: "memory");
}
__device__ __forceinline__ void fence_mbarrier_init() {
@@ -227,14 +229,14 @@ __device__ __forceinline__ void cp_async_bulk(void* smem_dst,
: "memory");
}
template <int H, int NC = N_CHUNK_DEFAULT, bool RELEASE_TMEM = false,
bool HAS_DELTA = false, bool HAS_OUTPUT_NORM = false,
bool OUTPUT_NORM_IN_SMEM = false>
template <int H, int N, int NC = N_CHUNK_DEFAULT, int B = 1,
bool RELEASE_TMEM = false, bool HAS_DELTA = false,
bool HAS_OUTPUT_NORM = false, bool OUTPUT_NORM_IN_SMEM = false>
__global__ void __launch_bounds__(BLK, 1) attn_res_fwd_online_v2_kernel(
const bf16_t* __restrict__ block_res, bf16_t* __restrict__ layer_res,
const bf16_t* __restrict__ delta, const bf16_t* __restrict__ res_w,
const bf16_t* __restrict__ rms_w, bf16_t* __restrict__ output, int N, int T,
int B, int block_stride_m, int block_stride_r, float rms_eps,
const bf16_t* __restrict__ rms_w, bf16_t* __restrict__ output, int T,
int block_stride_m, int block_stride_r, float rms_eps,
const bf16_t* __restrict__ output_norm_weight, float output_norm_eps) {
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 1000 && __CUDA_ARCH__ < 1100
constexpr float LOG2_E = 1.4426950408889634f;
@@ -258,7 +260,7 @@ __global__ void __launch_bounds__(BLK, 1) attn_res_fwd_online_v2_kernel(
const int lane = tid & 31;
const int TB = T * B;
const int num_ctas = gridDim.x;
const int num_chunks = (N + N_CHUNK - 1) / N_CHUNK;
constexpr int num_chunks = (N + N_CHUNK - 1) / N_CHUNK;
const int comp_wid = wid - 1;
const int comp_tid = tid - 32;
@@ -329,11 +331,16 @@ __global__ void __launch_bounds__(BLK, 1) attn_res_fwd_online_v2_kernel(
if constexpr (H == 7168) {
if (si == SLICES_PER_GROUP - 1) {
int h_base = 6 * K_TILE + group * (K_TILE / 2) + ct_in_group * 4;
int2 rms_v = *reinterpret_cast<const int2*>(rms_w + h_base);
int2 res_v = *reinterpret_cast<const int2*>(res_w + h_base);
auto* rms2 = reinterpret_cast<__nv_bfloat162*>(&rms_v);
auto* res2 = reinterpret_cast<__nv_bfloat162*>(&res_v);
#pragma unroll
for (int j = 0; j < 4; j++) {
int h = h_base + j;
q_cache[si * VEC + j] =
__bfloat162float(rms_w[h]) * __bfloat162float(res_w[h]);
for (int k = 0; k < 2; k++) {
float2 rf = __bfloat1622float2(rms2[k]);
float2 sf = __bfloat1622float2(res2[k]);
q_cache[si * VEC + 2 * k] = rf.x * sf.x;
q_cache[si * VEC + 2 * k + 1] = rf.y * sf.y;
}
continue;
}
@@ -341,11 +348,16 @@ __global__ void __launch_bounds__(BLK, 1) attn_res_fwd_online_v2_kernel(
int dt = si * CONSUMER_GROUPS + group;
if (dt >= NHT) continue;
int h_base = dt * K_TILE + k_local;
int4 rms_v = *reinterpret_cast<const int4*>(rms_w + h_base);
int4 res_v = *reinterpret_cast<const int4*>(res_w + h_base);
auto* rms2 = reinterpret_cast<__nv_bfloat162*>(&rms_v);
auto* res2 = reinterpret_cast<__nv_bfloat162*>(&res_v);
#pragma unroll
for (int j = 0; j < VEC; j++) {
int h = h_base + j;
q_cache[si * VEC + j] =
__bfloat162float(rms_w[h]) * __bfloat162float(res_w[h]);
for (int k = 0; k < 4; k++) {
float2 rf = __bfloat1622float2(rms2[k]);
float2 sf = __bfloat1622float2(res2[k]);
q_cache[si * VEC + 2 * k] = rf.x * sf.x;
q_cache[si * VEC + 2 * k + 1] = rf.y * sf.y;
}
}
}
@@ -404,6 +416,7 @@ __global__ void __launch_bounds__(BLK, 1) attn_res_fwd_online_v2_kernel(
acc32[i] = 0.f;
}
#pragma unroll
for (int ci = 0; ci < num_chunks; ci++, gci++) {
int ns = ci * N_CHUNK;
int an = min(N_CHUNK, N - ns);
@@ -854,13 +867,13 @@ __global__ void __launch_bounds__(BLK, 1) attn_res_fwd_online_v2_kernel(
#endif
}
template <int H, int NC = N_CHUNK_DEFAULT, bool RELEASE_TMEM = false,
template <int H, int N, int NC = N_CHUNK_DEFAULT, bool RELEASE_TMEM = false,
bool HAS_DELTA = false, bool HAS_OUTPUT_NORM = false,
bool OUTPUT_NORM_IN_SMEM = false>
static void launch_fwd(const bf16_t* block_residual, bf16_t* layer_residual,
const bf16_t* delta, const bf16_t* res_weight,
const bf16_t* rms_weight, bf16_t* output, int N, int T,
int B, float rms_eps, int num_sm, cudaStream_t stream,
const bf16_t* rms_weight, bf16_t* output, int T,
float rms_eps, int num_sm, cudaStream_t stream,
const bf16_t* output_norm_weight = nullptr,
float output_norm_eps = 0.f, int block_stride_m = 0,
int block_stride_r = 0) {
@@ -870,7 +883,7 @@ static void launch_fwd(const bf16_t* block_residual, bf16_t* layer_residual,
sizeof(FwdSmemPlan<NC>) + 15) &
~size_t(15);
auto kernel =
&attn_res_fwd_online_v2_kernel<H, NC, RELEASE_TMEM, HAS_DELTA,
&attn_res_fwd_online_v2_kernel<H, N, NC, 1, RELEASE_TMEM, HAS_DELTA,
HAS_OUTPUT_NORM, OUTPUT_NORM_IN_SMEM>;
static bool attrs_set = false;
if (!attrs_set) {
@@ -892,7 +905,7 @@ static void launch_fwd(const bf16_t* block_residual, bf16_t* layer_residual,
config.attrs = attrs;
config.numAttrs = 1;
cudaLaunchKernelEx(&config, kernel, block_residual, layer_residual, delta,
res_weight, rms_weight, output, N, T, B, block_stride_m,
res_weight, rms_weight, output, T, block_stride_m,
block_stride_r, rms_eps, output_norm_weight,
output_norm_eps);
}
@@ -926,33 +939,65 @@ void kimi_k3_attn_res(torch::stable::Tensor& prefix,
// Two-source chunks and two resident CTAs are beneficial once setup is
// amortized by the long, full eight-block prefill workload.
if (num_blocks == 8 && num_tokens >= 4096) {
launch_fwd<7168, 2, true, true, true, true>(
launch_fwd<7168, 9, 2, true, true, true, true>(
static_cast<bf16_t const*>(blocks.data_ptr()),
static_cast<bf16_t*>(prefix.data_ptr()),
static_cast<bf16_t const*>(delta.data_ptr()),
static_cast<bf16_t const*>(qk_weight.data_ptr()),
static_cast<bf16_t const*>(norm_weight.data_ptr()),
static_cast<bf16_t*>(output.data_ptr()),
static_cast<int>(num_blocks) + 1, num_tokens, 1,
static_cast<bf16_t*>(output.data_ptr()), num_tokens,
static_cast<float>(eps), properties->multiProcessorCount,
get_current_cuda_stream(device),
static_cast<bf16_t const*>(output_norm_weight.data_ptr()),
static_cast<float>(output_norm_eps), static_cast<int>(blocks.stride(0)),
static_cast<int>(blocks.stride(1)));
} else {
launch_fwd<7168, 4, false, true, true, true>(
static_cast<bf16_t const*>(blocks.data_ptr()),
static_cast<bf16_t*>(prefix.data_ptr()),
static_cast<bf16_t const*>(delta.data_ptr()),
static_cast<bf16_t const*>(qk_weight.data_ptr()),
static_cast<bf16_t const*>(norm_weight.data_ptr()),
static_cast<bf16_t*>(output.data_ptr()),
static_cast<int>(num_blocks) + 1, num_tokens, 1,
static_cast<float>(eps), properties->multiProcessorCount,
get_current_cuda_stream(device),
static_cast<bf16_t const*>(output_norm_weight.data_ptr()),
static_cast<float>(output_norm_eps), static_cast<int>(blocks.stride(0)),
static_cast<int>(blocks.stride(1)));
auto dispatch = [&](auto nsrc_tok) {
constexpr int NSRC = decltype(nsrc_tok)::value;
launch_fwd<7168, NSRC, 3, false, true, true, true>(
static_cast<bf16_t const*>(blocks.data_ptr()),
static_cast<bf16_t*>(prefix.data_ptr()),
static_cast<bf16_t const*>(delta.data_ptr()),
static_cast<bf16_t const*>(qk_weight.data_ptr()),
static_cast<bf16_t const*>(norm_weight.data_ptr()),
static_cast<bf16_t*>(output.data_ptr()), num_tokens,
static_cast<float>(eps), properties->multiProcessorCount,
get_current_cuda_stream(device),
static_cast<bf16_t const*>(output_norm_weight.data_ptr()),
static_cast<float>(output_norm_eps),
static_cast<int>(blocks.stride(0)),
static_cast<int>(blocks.stride(1)));
};
// The source count is num_blocks + 1; specialising on it makes the chunk
// count, per-chunk source count and prefix-chunk test compile-time.
switch (num_blocks) {
case 1:
dispatch(std::integral_constant<int, 2>{});
break;
case 2:
dispatch(std::integral_constant<int, 3>{});
break;
case 3:
dispatch(std::integral_constant<int, 4>{});
break;
case 4:
dispatch(std::integral_constant<int, 5>{});
break;
case 5:
dispatch(std::integral_constant<int, 6>{});
break;
case 6:
dispatch(std::integral_constant<int, 7>{});
break;
case 7:
dispatch(std::integral_constant<int, 8>{});
break;
case 8:
dispatch(std::integral_constant<int, 9>{});
break;
default:
STD_TORCH_CHECK(false, "Kimi K3 AttnRes: num_blocks must be 1..8");
}
}
cudaError_t const error = cudaGetLastError();
STD_TORCH_CHECK(