perf: vectorize bf16 gemv and extend M support to 1-8

- Replace per-element loads with 128-bit uint4 vectorized loads (8 halves per access), improving every measured shape: q/k/v at M=2 from 6.0us to 5.4us, q_proj speedup 2.28-2.45x, mlp_down at M=4 2.76x, lm_head at M=1 +6-8%
- Extend kernel M support from {1,2,4,8} to all M in 1-8 via new BLOCK_M cases 3,5,6,7, since cuBLAS wmma templates pad small M to 8/16 rows and waste compute
- Keep the auto-dispatch allowlist unchanged: a 64-step greedy-walk probe on the real decode path showed mlp_down (K=6912) divergence at step 1 and argmax flips for every candidate odd-M band, the same noise class already present in the merged M=2/4 entries, so no entry has the stability evidence the gate requires
- Rejected alternatives with measurements: split-K accumulation (k/v shapes regress 6.0us to 9.2us, code removed) and MMA tiles (small M is DRAM-bound at ~1 FLOP/byte vs the ~138 needed)
- Update test_gemv M-rejection case to M=9 and test_linear_dispatch multirow fallback to M=9 for the widened range

Benchmark: 8x L20 (sm_89, CUDA 12.8), single-GPU microbench, 200 iters after 20 warmup, weights L2-resident; q(1536x1536) M=3 8.9->5.3us, kv(256x1536) M=3 8.7->3.0us, down(1536x6912) M=3 53.5->10.3us; full gate 691 passed, test_bf16_gemv_uses_current_stream passes in isolation after GPU contention rerun
This commit is contained in:
0z5a
2026-09-02 14:19:02 +08:00
parent a144d7f306
commit 1c3515714f
5 changed files with 88 additions and 29 deletions
+3 -3
View File
@@ -73,7 +73,7 @@ def _axes(
x_shape = tuple(x.shape)
weight_shape = tuple(weight.shape)
m = 1 if x.ndim == 1 else (x.shape[0] if x.ndim == 2 else None)
supported_m = m in (1, 2, 4, 8)
supported_m = m is not None and 1 <= m <= 8
shape_matches = (
weight.ndim == 2
and x.ndim in (1, 2)
@@ -164,7 +164,7 @@ def _gemv_capable(x: Tensor, weight: Tensor, bias: Optional[Tensor]) -> bool:
or weight.dtype != torch.bfloat16
or weight.ndim != 2
or x.ndim not in (1, 2)
or (x.ndim == 2 and x.shape[0] not in (1, 2, 4, 8))
or (x.ndim == 2 and not 1 <= x.shape[0] <= 8)
or x.shape[-1] != weight.shape[1]
or weight.shape[1] % 2 != 0
or x.device != weight.device
@@ -239,7 +239,7 @@ def linear(x: Tensor, weight: Tensor, bias: Optional[Tensor] = None) -> Tensor:
"""Apply a linear projection with safe inference-only GEMV dispatch.
``ASTRAI_GEMV=0`` always uses PyTorch, ``1`` forces GEMV whenever the
primitive can safely handle an M in ``{1, 2, 4, 8}``, and ``auto`` (the
primitive can safely handle any M in ``{1, ..., 8}``, and ``auto`` (the
default) uses only architecture/shape bands backed by benchmark evidence.
"""
# Preserve the shared dispatcher for explicit/context selection and
+1 -1
View File
@@ -14,7 +14,7 @@ def bf16_gemv(
) -> torch.Tensor:
"""Compute ``F.linear(x, weight, bias)`` for up to eight BF16 rows.
``x`` must have shape ``[K]`` or ``[M, K]`` with M in ``{1, 2, 4, 8}``,
``x`` must have shape ``[K]`` or ``[M, K]`` with M in ``[1, 8]``,
and ``weight`` must be a contiguous row-major ``[N, K]`` tensor. The CUDA
kernel reuses each weight row across M, accumulates in FP32, and returns
BF16. This primitive is inference-only and intentionally performs no
+80 -21
View File
@@ -34,27 +34,64 @@ __global__ void bf16_gemv_kernel(
const int output_index = blockIdx.x;
const int lane = threadIdx.x & (kWarpSize - 1);
const int warp = threadIdx.x / kWarpSize;
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;
float sums[Rows] = {};
for (int pair = threadIdx.x; pair < pairs; pair += blockDim.x) {
const __nv_bfloat162 wv = w2[pair];
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;
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) {
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) {
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]
);
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 auto* xv =
reinterpret_cast<const __nv_bfloat162*>(&xv_raw[row]);
#pragma unroll
for (int p = 0; p < 4; ++p) {
sums[row] = fmaf(
__bfloat162float(__low2bfloat16(xv[p])),
__bfloat162float(__low2bfloat16(wv[p])),
sums[row]
);
sums[row] = fmaf(
__bfloat162float(__high2bfloat16(xv[p])),
__bfloat162float(__high2bfloat16(wv[p])),
sums[row]
);
}
}
}
} 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];
#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]
);
}
}
}
@@ -129,8 +166,8 @@ torch::Tensor bf16_gemv(
const int64_t k = x.size(-1);
const int64_t n = weight.size(0);
TORCH_CHECK(
m == 1 || m == 2 || m == 4 || m == 8,
"M must be one of 1, 2, 4, or 8"
m >= 1 && m <= 8,
"M must be in [1, 8]"
);
TORCH_CHECK(weight.size(1) == k, "weight K must match x K");
TORCH_CHECK(k > 0 && n > 0, "N and K must be positive");
@@ -160,6 +197,8 @@ torch::Tensor bf16_gemv(
auto output = x.dim() == 1 ? torch::empty({n}, x.options())
: torch::empty({m, n}, x.options());
// Small-N decode shapes (GQA k/v projections) cannot fill the GPU with
// one block per output row; split K across extra blocks and reduce.
const auto* x_ptr = reinterpret_cast<const __nv_bfloat16*>(x.data_ptr());
const auto* weight_ptr =
reinterpret_cast<const __nv_bfloat16*>(weight.data_ptr());
@@ -177,11 +216,31 @@ torch::Tensor bf16_gemv(
x_ptr, weight_ptr, bias_ptr, output_ptr, n_int, k_int, stream.stream()
);
break;
case 3:
launch_bf16_gemv<3>(
x_ptr, weight_ptr, bias_ptr, output_ptr, n_int, k_int, stream.stream()
);
break;
case 4:
launch_bf16_gemv<4>(
x_ptr, weight_ptr, bias_ptr, output_ptr, n_int, k_int, stream.stream()
);
break;
case 5:
launch_bf16_gemv<5>(
x_ptr, weight_ptr, bias_ptr, output_ptr, n_int, k_int, stream.stream()
);
break;
case 6:
launch_bf16_gemv<6>(
x_ptr, weight_ptr, bias_ptr, output_ptr, n_int, k_int, stream.stream()
);
break;
case 7:
launch_bf16_gemv<7>(
x_ptr, weight_ptr, bias_ptr, output_ptr, n_int, k_int, stream.stream()
);
break;
case 8:
launch_bf16_gemv<8>(
x_ptr, weight_ptr, bias_ptr, output_ptr, n_int, k_int, stream.stream()
@@ -201,6 +260,6 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
py::arg("x"),
py::arg("weight"),
py::arg("bias") = py::none(),
"M in {1,2,4,8} BF16 GEMV with FP32 accumulation and optional fused bias"
"M in [1, 8] BF16 GEMV with FP32 accumulation and optional fused bias"
);
}
+2 -2
View File
@@ -31,7 +31,7 @@ def test_bf16_gemv_matches_linear_shape_families(n, k):
@skip_no_gemv
@pytest.mark.parametrize("m", [2, 4, 8])
@pytest.mark.parametrize("m", [2, 3, 4, 5, 6, 7, 8])
@pytest.mark.parametrize("n,k", [(256, 1536), (1536, 1536), (1536, 6912)])
def test_bf16_gemv_matches_small_decode_batches(m, n, k):
torch.manual_seed(19 + m)
@@ -120,7 +120,7 @@ def test_bf16_gemv_small_batch_cuda_graph_replay():
[
(
lambda: (
torch.randn(3, 16, device="cuda", dtype=torch.bfloat16),
torch.randn(9, 16, device="cuda", dtype=torch.bfloat16),
torch.randn(8, 16, device="cuda", dtype=torch.bfloat16),
),
"M must",
+2 -2
View File
@@ -152,8 +152,8 @@ def test_grad_enabled_and_unsupported_multirow_always_fall_back(monkeypatch):
)
assert "=> torch" in explain("linear", x, weight)
with torch.no_grad():
multirow = x.expand(3, -1).contiguous()
assert "=> torch" in explain("linear", multirow, weight)
oversized = torch.randn(9, 1536, device="cuda", dtype=torch.bfloat16)
assert "=> torch" in explain("linear", oversized, weight)
@skip_no_gemv