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
-3
View File
@@ -107,7 +107,6 @@ def _axes(
x_contiguous=x.is_contiguous(), x_contiguous=x.is_contiguous(),
weight_contiguous=weight.is_contiguous(), weight_contiguous=weight.is_contiguous(),
bias_supported=bias_supported, bias_supported=bias_supported,
k_even=k is not None and k % 2 == 0,
) )
@@ -122,7 +121,6 @@ _SPEC_CAPABLE = (
& axis("x_contiguous").truthy() & axis("x_contiguous").truthy()
& axis("weight_contiguous").truthy() & axis("weight_contiguous").truthy()
& axis("bias_supported").truthy() & axis("bias_supported").truthy()
& axis("k_even").truthy()
) )
_SPEC_AUTO = _SPEC_CAPABLE & Spec.of( _SPEC_AUTO = _SPEC_CAPABLE & Spec.of(
@@ -166,7 +164,6 @@ def _gemv_capable(x: Tensor, weight: Tensor, bias: Optional[Tensor]) -> bool:
or x.ndim not in (1, 2) or x.ndim not in (1, 2)
or (x.ndim == 2 and not 1 <= x.shape[0] <= 8) or (x.ndim == 2 and not 1 <= x.shape[0] <= 8)
or x.shape[-1] != weight.shape[1] or x.shape[-1] != weight.shape[1]
or weight.shape[1] % 2 != 0
or x.device != weight.device or x.device != weight.device
or not x.is_contiguous() or not x.is_contiguous()
or not weight.is_contiguous() or not weight.is_contiguous()
+65 -32
View File
@@ -36,26 +36,40 @@ __global__ void bf16_gemv_kernel(
const int warp = threadIdx.x / kWarpSize; const int warp = threadIdx.x / kWarpSize;
float sums[Rows] = {}; float sums[Rows] = {};
if (k % 8 == 0) { __shared__ float warp_sums[Rows][kThreads / kWarpSize];
// 128-bit vectorized loads: eight bf16 elements per access halve the // Weight row: scalar head/tail around a 16-byte-aligned uint4 middle so
// per-thread iteration count on bandwidth-bound decode shapes. // any K is accepted while keeping 128-bit weight loads, which dominate
const int vecs = k / 8; // 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* x4 = reinterpret_cast<const uint4*>(x);
const auto* w4 = reinterpret_cast<const uint4*>(weight) + for (int v = threadIdx.x; v < wvecs; v += blockDim.x) {
static_cast<int64_t>(output_index) * vecs;
for (int v = threadIdx.x; v < vecs; v += blockDim.x) {
const uint4 wv_raw = w4[v]; const uint4 wv_raw = w4[v];
const auto* wv = const auto* wv =
reinterpret_cast<const __nv_bfloat162*>(&wv_raw); 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 #pragma unroll
for (int row = 0; row < Rows; ++row) { for (int row = 0; row < Rows; ++row) {
const uint4 xv_raw =
x4[(static_cast<int64_t>(row) * wvecs) + v];
const auto* xv = const auto* xv =
reinterpret_cast<const __nv_bfloat162*>(&xv_raw[row]); reinterpret_cast<const __nv_bfloat162*>(&xv_raw);
#pragma unroll #pragma unroll
for (int p = 0; p < 4; ++p) { for (int p = 0; p < 4; ++p) {
sums[row] = fmaf( sums[row] = fmaf(
@@ -72,30 +86,50 @@ __global__ void bf16_gemv_kernel(
} }
} }
} else { } else {
const int pairs = k / 2; for (int v = threadIdx.x; v < wvecs; v += blockDim.x) {
const auto* x2 = reinterpret_cast<const __nv_bfloat162*>(x); const uint4 wv_raw = w4[v];
const auto* w2 = const __nv_bfloat16* wv_s =
reinterpret_cast<const __nv_bfloat162*>(weight) + output_index * pairs; reinterpret_cast<const __nv_bfloat16*>(&wv_raw);
for (int pair = threadIdx.x; pair < pairs; pair += blockDim.x) {
const __nv_bfloat162 wv = w2[pair];
#pragma unroll #pragma unroll
for (int row = 0; row < Rows; ++row) { for (int row = 0; row < Rows; ++row) {
const __nv_bfloat162 xv = x2[row * pairs + pair]; const __nv_bfloat16* xv =
sums[row] = fmaf( x + static_cast<int64_t>(row) * k + whead + 8 * v;
__bfloat162float(__low2bfloat16(xv)), #pragma unroll
__bfloat162float(__low2bfloat16(wv)), for (int s = 0; s < 8; ++s) {
sums[row] sums[row] = fmaf(
); __bfloat162float(xv[s]),
sums[row] = fmaf( __bfloat162float(wv_s[s]),
__bfloat162float(__high2bfloat16(xv)), sums[row]
__bfloat162float(__high2bfloat16(wv)), );
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 #pragma unroll
for (int row = 0; row < Rows; ++row) { for (int row = 0; row < Rows; ++row) {
sums[row] = warp_sum(sums[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(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 > 0 && n > 0, "N and K must be positive");
TORCH_CHECK(k % 2 == 0, "K must be even for vectorized bf16 loads");
TORCH_CHECK( TORCH_CHECK(
k <= std::numeric_limits<int>::max() && k <= std::numeric_limits<int>::max() &&
n <= std::numeric_limits<int>::max(), n <= std::numeric_limits<int>::max(),
+10 -6
View File
@@ -14,22 +14,26 @@ model linear dispatcher described below.
| `attn_paged_decode` | `attention/paged_decode.cu` | Paged KV cache decode attention | | `attn_paged_decode` | `attention/paged_decode.cu` | Paged KV cache decode attention |
| `attn_paged_prefill` | `attention/paged_prefill.cu` | Paged KV cache prefill attention (ragged batch) | | `attn_paged_prefill` | `attention/paged_prefill.cu` | Paged KV cache prefill attention (ragged batch) |
| `rotary_emb` | `rotary_emb.cu` | Fused rotary embedding (cos/sin lookup + rotation) | | `rotary_emb` | `rotary_emb.cu` | Fused rotary embedding (cos/sin lookup + rotation) |
| `bf16_gemv` | `gemv/bf16_gemv.cu` | M=1/2/4/8 BF16 linear with FP32 accumulation (sm_80+) | | `bf16_gemv` | `gemv/bf16_gemv.cu` | M=1..8 BF16 linear with FP32 accumulation (sm_80+) |
| `fp8_ops` | `fp8/ops.cu` | FP8 quantization + tensor-core GEMM (sm_89+) | | `fp8_ops` | `fp8/ops.cu` | FP8 quantization + tensor-core GEMM (sm_89+) |
### BF16 GEMV primitive ### BF16 GEMV primitive
`astrai.extension.bf16_gemv(x, weight, bias=None)` accepts a contiguous BF16 `astrai.extension.bf16_gemv(x, weight, bias=None)` accepts a contiguous BF16
input shaped `[K]` or `[M, K]`, with `M` in `{1, 2, 4, 8}`, and row-major input shaped `[K]` or `[M, K]`, with `M` in `[1, 8]` and any positive `K`, and
weights `[N, K]`. One CTA reduces each output row and computes all M results row-major weights `[N, K]`. One CTA reduces each output row and computes all M
together, reusing the weight row across tokens. It uses vectorized results together, reusing the weight row across tokens. The weight stream uses
`__nv_bfloat162` loads and FP32 accumulation; the optional BF16 bias is fused 128-bit vectorized loads anchored at each row's first 16-byte-aligned address
with scalar head/tail sweeps for unaligned remainders, so arbitrary `K` and
storage offsets stay correct; x loads are vectorized when every row base is
16-byte aligned (always true for K % 8 == 0 with allocator-aligned tensors)
and scalar otherwise. Accumulation is FP32; the optional BF16 bias is fused
before the BF16 store. The launcher uses the current CUDA stream, is CUDA before the BF16 store. The launcher uses the current CUDA stream, is CUDA
Graph capture-safe, and requires sm_80 or newer. Graph capture-safe, and requires sm_80 or newer.
Model `Linear` calls route through the lightweight linear backend. Set Model `Linear` calls route through the lightweight linear backend. Set
`ASTRAI_GEMV=0` for an unconditional `F.linear` fallback, `1` to force the `ASTRAI_GEMV=0` for an unconditional `F.linear` fallback, `1` to force the
kernel for any supported M=1/2/4/8 call, or `auto` (the default) to select only kernel for any supported M in [1, 8], or `auto` (the default) to select only
architecture/shape bands that pass both the per-shape and end-to-end gates. architecture/shape bands that pass both the per-shape and end-to-end gates.
M=1 has no automatic SM89 band because isolated winners did not reach the 3% M=1 has no automatic SM89 band because isolated winners did not reach the 3%
whole-graph gate. Measured SM89 small-M bands are enabled as follows: whole-graph gate. Measured SM89 small-M bands are enabled as follows:
+15 -7
View File
@@ -96,6 +96,21 @@ def test_bf16_gemv_cuda_graph_replay():
torch.testing.assert_close(actual, expected, rtol=0.02, atol=0.25) torch.testing.assert_close(actual, expected, rtol=0.02, atol=0.25)
@skip_no_gemv
@pytest.mark.parametrize("n,k", [(64, 7), (64, 12), (33, 100), (256, 1534)])
def test_bf16_gemv_handles_unaligned_k(n, k):
torch.manual_seed(29)
x = torch.randn(k, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(n, k, device="cuda", dtype=torch.bfloat16)
actual = bf16_gemv(x, weight)
expected = F.linear(x, weight)
torch.testing.assert_close(actual, expected, rtol=0.02, atol=0.25)
x3 = torch.randn(3, k, device="cuda", dtype=torch.bfloat16)
actual3 = bf16_gemv(x3, weight)
torch.testing.assert_close(actual3, F.linear(x3, weight), rtol=0.02, atol=0.5)
@skip_no_gemv @skip_no_gemv
def test_bf16_gemv_small_batch_cuda_graph_replay(): def test_bf16_gemv_small_batch_cuda_graph_replay():
torch.manual_seed(31) torch.manual_seed(31)
@@ -125,13 +140,6 @@ def test_bf16_gemv_small_batch_cuda_graph_replay():
), ),
"M must", "M must",
), ),
(
lambda: (
torch.randn(15, device="cuda", dtype=torch.bfloat16),
torch.randn(8, 15, device="cuda", dtype=torch.bfloat16),
),
"even",
),
( (
lambda: ( lambda: (
torch.randn(16, device="cuda", dtype=torch.float16), torch.randn(16, device="cuda", dtype=torch.float16),