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:
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user