- 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
24 lines
765 B
Python
24 lines
765 B
Python
"""Stateless wrapper for the directly callable BF16 GEMV primitive."""
|
|
|
|
from typing import Optional
|
|
|
|
import torch
|
|
|
|
from astrai.extension.loader import get_module
|
|
|
|
|
|
def bf16_gemv(
|
|
x: torch.Tensor,
|
|
weight: torch.Tensor,
|
|
bias: Optional[torch.Tensor] = None,
|
|
) -> 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, 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
|
|
fallback or model-level dispatch.
|
|
"""
|
|
return get_module("bf16_gemv").bf16_gemv(x, weight, bias)
|