perf: tune bf16 gemv and add opt-in fused swiglu

- deepen common-shape BF16 GEMV tuning with warp-row tiling for LLaMA/Qwen2/GPT-NeoX/OPT decode projections
- add fused BF16 up/gate SwiGLU CUDA primitive with ASTRAI_SWIGLU=0/1/auto dispatch
- keep the unfused linear backend as the default path; auto enables no shape until per-architecture checkpoint gates pass
- fall back to the linear/torch chain when kernels are absent, on CPU, in training, or outside supported M/K/dtype shapes
- add gemv/swiglu benchmark scripts, dispatch and parity tests, and kernel documentation

Benchmark: NVIDIA L20 (sm_89), CUDA 12.8, PyTorch 2.11.0+cu128, idle GPU. AstrAI 1B config (24 layers, hidden 1536, vocab 100000), BF16, prompt 128, 32 greedy decode tokens, CUDA graphs enabled, A/B in separate interleaved processes (3 rounds, 8 trials each, medians). Default vs ASTRAI_SWIGLU=1 per generate call: batch 1 134.8->129.1 ms (+4.44%), batch 2 136.2->130.9 ms (+4.06%), batch 4 145.5->140.3 ms (+3.66%). Greedy output identical at batch 1, differs at batch 2/4, so auto stays unfused by default; kernelless fallback verified bit-identical greedy.
This commit is contained in:
0z5a
2026-09-03 04:26:53 +08:00
parent 88c06db096
commit d4a292b36b
20 changed files with 2008 additions and 37 deletions
+141 -5
View File
@@ -12,7 +12,9 @@
namespace {
constexpr int kThreads = 256;
constexpr int kHalfCtaThreads = 128;
constexpr int kWarpSize = 32;
constexpr int kWarpTiledThreads = 128;
__device__ __forceinline__ float warp_sum(float value) {
#pragma unroll
@@ -22,7 +24,7 @@ __device__ __forceinline__ float warp_sum(float value) {
return value;
}
template <int Rows>
template <int Rows, int Threads>
__global__ void bf16_gemv_kernel(
const __nv_bfloat16* __restrict__ x,
const __nv_bfloat16* __restrict__ weight,
@@ -36,7 +38,7 @@ __global__ void bf16_gemv_kernel(
const int warp = threadIdx.x / kWarpSize;
float sums[Rows] = {};
__shared__ float warp_sums[Rows][kThreads / kWarpSize];
__shared__ float warp_sums[Rows][Threads / kWarpSize];
// Weight row: scalar head/tail around a 16-byte-aligned uint4 middle so
// any K is accepted while keeping 128-bit weight loads, which dominate
// bandwidth on decode shapes. x pairs with scalar loads: it is a tiny
@@ -129,7 +131,6 @@ __global__ void bf16_gemv_kernel(
}
}
#pragma unroll
for (int row = 0; row < Rows; ++row) {
sums[row] = warp_sum(sums[row]);
@@ -146,7 +147,7 @@ __global__ void bf16_gemv_kernel(
#pragma unroll
for (int row = 0; row < Rows; ++row) {
float sum =
lane < (kThreads / kWarpSize) ? warp_sums[row][lane] : 0.0f;
lane < (Threads / kWarpSize) ? warp_sums[row][lane] : 0.0f;
sum = warp_sum(sum);
if (lane == 0) {
if (bias != nullptr) {
@@ -158,6 +159,120 @@ __global__ void bf16_gemv_kernel(
}
}
template <int Rows>
__global__ void bf16_gemv_aligned_warp_tiled_kernel(
const __nv_bfloat16* __restrict__ x,
const __nv_bfloat16* __restrict__ weight,
const __nv_bfloat16* __restrict__ bias,
__nv_bfloat16* __restrict__ output,
int n,
int k
) {
constexpr int kWarpsPerBlock = kWarpTiledThreads / kWarpSize;
const int lane = threadIdx.x & (kWarpSize - 1);
const int warp = threadIdx.x / kWarpSize;
const int output_index = blockIdx.x * kWarpsPerBlock + warp;
if (output_index >= n) {
return;
}
// The launcher selects this path only when each row is 16-byte aligned.
// Four independent output rows per CTA remove the block-wide reduction
// barrier and improve occupancy for the medium LLaMA projection bands.
const int vectors = k / 8;
const auto* x4 = reinterpret_cast<const uint4*>(x);
const auto* w4 = reinterpret_cast<const uint4*>(weight) +
static_cast<int64_t>(output_index) * vectors;
float sums[Rows] = {};
for (int vector = lane; vector < vectors; vector += kWarpSize) {
const uint4 wv_raw = w4[vector];
const auto* wv = reinterpret_cast<const __nv_bfloat162*>(&wv_raw);
#pragma unroll
for (int row = 0; row < Rows; ++row) {
const uint4 xv_raw =
x4[static_cast<int64_t>(row) * vectors + vector];
const auto* xv = reinterpret_cast<const __nv_bfloat162*>(&xv_raw);
#pragma unroll
for (int pair = 0; pair < 4; ++pair) {
sums[row] = fmaf(
__bfloat162float(__low2bfloat16(xv[pair])),
__bfloat162float(__low2bfloat16(wv[pair])),
sums[row]
);
sums[row] = fmaf(
__bfloat162float(__high2bfloat16(xv[pair])),
__bfloat162float(__high2bfloat16(wv[pair])),
sums[row]
);
}
}
}
#pragma unroll
for (int row = 0; row < Rows; ++row) {
sums[row] = warp_sum(sums[row]);
if (lane == 0) {
if (bias != nullptr) {
sums[row] += __bfloat162float(bias[output_index]);
}
output[row * n + output_index] = __float2bfloat16_rn(sums[row]);
}
}
}
template <int Rows>
constexpr bool use_warp_tiled_kernel(int n, int k) {
// These bands are intentionally narrow and are validated by the common
// transformer benchmark. The 256-thread cooperative kernel remains the
// fallback for arbitrary K, larger projections, and M=2 (where the
// single-warp reduction regresses the current vectorized kernel).
if constexpr (Rows == 4) {
return (n == 1024 && k == 4096) ||
(n == 4096 && k == 4096) ||
(n == 11008 && k == 4096) ||
(n == 4096 && k == 11008);
}
return false;
}
template <int Rows>
constexpr bool use_half_cta_kernel(int n, int k) {
// A 128-thread CTA reduces synchronization and scheduling overhead for
// selected medium decode projections. Keep the selector exact: long-K
// and bandwidth-saturated shapes regress, and the winning bands differ
// materially with the number of reused input rows.
if constexpr (Rows == 1) {
return n == 8192 && k == 2048;
}
if constexpr (Rows == 2) {
return (n == 4096 && k == 4096) ||
(n == 11008 && k == 4096) ||
(n == 3584 && k == 3584) ||
(n == 2048 && k == 2048) ||
(n == 8192 && k == 2048);
}
if constexpr (Rows == 4) {
return (n == 5120 && k == 5120) ||
(n == 3584 && k == 3584) ||
(n == 2048 && k == 2048) ||
(n == 8192 && k == 2048);
}
if constexpr (Rows == 8) {
return (n == 4096 && k == 4096) ||
(n == 11008 && k == 4096) ||
(n == 4096 && k == 11008) ||
(n == 1024 && k == 4096) ||
(n == 5120 && k == 5120) ||
(n == 512 && k == 3584) ||
(n == 3584 && k == 3584) ||
(n == 1024 && k == 8192) ||
(n == 2048 && k == 2048) ||
(n == 8192 && k == 2048) ||
(n == 2048 && k == 8192);
}
return false;
}
template <int Rows>
void launch_bf16_gemv(
const __nv_bfloat16* x,
@@ -168,7 +283,28 @@ void launch_bf16_gemv(
int k,
cudaStream_t stream
) {
bf16_gemv_kernel<Rows><<<n, kThreads, 0, stream>>>(
const bool aligned_rows = k % 8 == 0 &&
(reinterpret_cast<uintptr_t>(x) & 15u) == 0u &&
(reinterpret_cast<uintptr_t>(weight) & 15u) == 0u;
if constexpr (Rows == 4) {
if (aligned_rows && use_warp_tiled_kernel<Rows>(n, k)) {
constexpr int kWarpsPerBlock = kWarpTiledThreads / kWarpSize;
const int blocks = (n + kWarpsPerBlock - 1) / kWarpsPerBlock;
bf16_gemv_aligned_warp_tiled_kernel<Rows>
<<<blocks, kWarpTiledThreads, 0, stream>>>(
x, weight, bias, output, n, k
);
return;
}
}
if (aligned_rows && use_half_cta_kernel<Rows>(n, k)) {
bf16_gemv_kernel<Rows, kHalfCtaThreads>
<<<n, kHalfCtaThreads, 0, stream>>>(
x, weight, bias, output, n, k
);
return;
}
bf16_gemv_kernel<Rows, kThreads><<<n, kThreads, 0, stream>>>(
x, weight, bias, output, n, k
);
}