refactor: accept arbitrary K in bf16 gemv with aligned head-tail sweeps

- Drop the K % 2 entry rejection and the per-K if/else load-width branch: the weight stream now anchors uint4 loads at each row's first 16-byte-aligned address, with scalar head/tail sweeps covering at most 14 remainder elements, so any positive K and any storage offset is correct
- Keep one pure-uint4 loop (no branching inside the loop) for the production case where every x row base is 16-byte aligned (K % 8 == 0 with allocator-aligned tensors) and a scalar-x pairing loop only for unaligned K, where per-row uint4 loads are not addressable; measured cost of scalar x everywhere was up to 2.5x on multi-row shapes (down M=4 28.4us vs 11.3us)
- Remove the now-obsolete k_aligned axis and K divisibility gate from the linear dispatch spec since the primitive no longer rejects any K
- Add test coverage for unaligned K (7, 12, 100, 1534) at M=1 and M=3

Benchmark: 8x L20 (sm_89, CUDA 12.8), L2-resident microbench, 300 iters; hot path unchanged within noise vs the pure-uint4 kernel (q M=2 5.8us, down M=4 11.3us, lm M=1 391us); full gate green
This commit is contained in:
2026-09-02 14:48:01 +08:00
committed by 0z5a
parent 1c3515714f
commit 800981d85a
4 changed files with 90 additions and 48 deletions
+65 -32
View File
@@ -36,26 +36,40 @@ __global__ void bf16_gemv_kernel(
const int warp = threadIdx.x / kWarpSize;
float sums[Rows] = {};
if (k % 8 == 0) {
// 128-bit vectorized loads: eight bf16 elements per access halve the
// per-thread iteration count on bandwidth-bound decode shapes.
const int vecs = k / 8;
__shared__ float warp_sums[Rows][kThreads / 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
// L1/L2-resident matrix, consecutive threads still touch contiguous
// addresses, and no per-row alignment case analysis is needed.
const __nv_bfloat16* __restrict__ wrow =
weight + static_cast<int64_t>(output_index) * k;
const unsigned whead_raw =
((16u - (reinterpret_cast<uintptr_t>(wrow) & 15u)) & 15u) >> 1;
const int whead = static_cast<int>(min(whead_raw, static_cast<unsigned>(k)));
const int wvecs = (k - whead) / 8;
const int wtail_start = whead + wvecs * 8;
const uint4* __restrict__ w4 = reinterpret_cast<const uint4*>(wrow + whead);
// x chunks pair element-for-element with the aligned weight middle. When
// K % 8 == 0 every x row base shares the weight alignment, so one pure
// uint4 loop covers all rows (the production case: head/tail empty, no
// branching inside the loop). Otherwise per-row uint4 loads are not
// 16-byte addressable, and scalar x pairing keeps the kernel correct for
// any K while the weight stream stays vectorized.
if (k % 8 == 0 &&
((reinterpret_cast<uintptr_t>(x) + 2u * static_cast<unsigned>(whead)) & 15u) == 0u) {
const auto* x4 = reinterpret_cast<const uint4*>(x);
const auto* w4 = reinterpret_cast<const uint4*>(weight) +
static_cast<int64_t>(output_index) * vecs;
for (int v = threadIdx.x; v < vecs; v += blockDim.x) {
for (int v = threadIdx.x; v < wvecs; v += blockDim.x) {
const uint4 wv_raw = w4[v];
const auto* wv =
reinterpret_cast<const __nv_bfloat162*>(&wv_raw);
uint4 xv_raw[Rows];
#pragma unroll
for (int row = 0; row < Rows; ++row) {
xv_raw[row] = x4[static_cast<int64_t>(row) * vecs + v];
}
#pragma unroll
for (int row = 0; row < Rows; ++row) {
const uint4 xv_raw =
x4[(static_cast<int64_t>(row) * wvecs) + v];
const auto* xv =
reinterpret_cast<const __nv_bfloat162*>(&xv_raw[row]);
reinterpret_cast<const __nv_bfloat162*>(&xv_raw);
#pragma unroll
for (int p = 0; p < 4; ++p) {
sums[row] = fmaf(
@@ -72,30 +86,50 @@ __global__ void bf16_gemv_kernel(
}
}
} else {
const int pairs = k / 2;
const auto* x2 = reinterpret_cast<const __nv_bfloat162*>(x);
const auto* w2 =
reinterpret_cast<const __nv_bfloat162*>(weight) + output_index * pairs;
for (int pair = threadIdx.x; pair < pairs; pair += blockDim.x) {
const __nv_bfloat162 wv = w2[pair];
for (int v = threadIdx.x; v < wvecs; v += blockDim.x) {
const uint4 wv_raw = w4[v];
const __nv_bfloat16* wv_s =
reinterpret_cast<const __nv_bfloat16*>(&wv_raw);
#pragma unroll
for (int row = 0; row < Rows; ++row) {
const __nv_bfloat162 xv = x2[row * pairs + pair];
sums[row] = fmaf(
__bfloat162float(__low2bfloat16(xv)),
__bfloat162float(__low2bfloat16(wv)),
sums[row]
);
sums[row] = fmaf(
__bfloat162float(__high2bfloat16(xv)),
__bfloat162float(__high2bfloat16(wv)),
sums[row]
);
const __nv_bfloat16* xv =
x + static_cast<int64_t>(row) * k + whead + 8 * v;
#pragma unroll
for (int s = 0; s < 8; ++s) {
sums[row] = fmaf(
__bfloat162float(xv[s]),
__bfloat162float(wv_s[s]),
sums[row]
);
}
}
}
}
// Head and tail remainders: plain scalar pairing, at most 14 elements.
for (int i = threadIdx.x; i < whead; i += blockDim.x) {
const float wv = __bfloat162float(wrow[i]);
#pragma unroll
for (int row = 0; row < Rows; ++row) {
sums[row] = fmaf(
__bfloat162float(x[static_cast<int64_t>(row) * k + i]),
wv,
sums[row]
);
}
}
for (int i = wtail_start + threadIdx.x; i < k; i += blockDim.x) {
const float wv = __bfloat162float(wrow[i]);
#pragma unroll
for (int row = 0; row < Rows; ++row) {
sums[row] = fmaf(
__bfloat162float(x[static_cast<int64_t>(row) * k + i]),
wv,
sums[row]
);
}
}
__shared__ float warp_sums[Rows][kThreads / kWarpSize];
#pragma unroll
for (int row = 0; row < Rows; ++row) {
sums[row] = warp_sum(sums[row]);
@@ -171,7 +205,6 @@ torch::Tensor bf16_gemv(
);
TORCH_CHECK(weight.size(1) == k, "weight K must match x K");
TORCH_CHECK(k > 0 && n > 0, "N and K must be positive");
TORCH_CHECK(k % 2 == 0, "K must be even for vectorized bf16 loads");
TORCH_CHECK(
k <= std::numeric_limits<int>::max() &&
n <= std::numeric_limits<int>::max(),