perf: rebuild decode gemm dispatch around shape-driven tile configs

- split-K removed entirely: tiled kernel walks K in one pass, no partials/semas workspace, no memset, single launch per call
- skinny GEMM (M<=8) dispatch table replaces the hand-written switch
- shape-driven four-family table replaces plan_gemm: wide-N (n>=4096) default {16,64,64,3,128} with BM=32 at M>16; narrow-N deep-K rings {16,32,256,2,64} while the grid fits one wave, {16,32,128,2,64} past it
- narrow-N is K-serial: widening the grid measurably does nothing (BN 64->32 ties, doubled m_tiles tie, kv at 4 blocks ties q/o at 24); deeper K chunks win until 72KB smem forces one CTA per SM and past one wave the 2-wave quantization loses to BK=128
- launch-check macros in common/launch.cuh; smem opt-in for the 72KB/60KB rings
- rename kernels/bf16_*.cu to gemm.cu/swiglu.cu; module names unchanged
- Python gate: lm_head (N>32768) falls back to cuBLAS, band narrows to M<=32
- drop the stale per-op benchmark narratives; fold the live numbers into cuda_kernels.md

Benchmark: NVIDIA L20 (sm_89, 92 SMs), CUDA 12.8, bf16, L2-thrash weight rotation, per-call medians at M=16: q/o 9.5us, kv 8.6us, gate/up 33.3us, down 33.7us (down -29% vs prior default). End-to-end 1B decode (gen 128, 3 trials, tokens/s vs cuBLAS): B=1 260 vs 252, B=8 1660 vs 1446, B=16 2464 vs 2437, B=32 3620 vs 3690. Prior split-K dispatch measured B=16 2243 / B=32 3393.
This commit is contained in:
2026-09-04 22:41:39 +08:00
parent 8e39d9d8c9
commit 1798474316
21 changed files with 1098 additions and 697 deletions
@@ -2,48 +2,65 @@ import pytest
import torch
import torch.nn.functional as F
from astrai.extension import bf16_gemv, is_available
from astrai.extension import bf16_gemm, is_available
GEMV_AVAILABLE = (
GEMM_AVAILABLE = (
torch.cuda.is_available()
and is_available("bf16_gemv")
and is_available("bf16_gemm")
and torch.cuda.get_device_capability() >= (8, 0)
)
skip_no_gemv = pytest.mark.skipif(
not GEMV_AVAILABLE,
reason="BF16 GEMV requires a built kernel and compute capability 8.0+",
skip_no_gemm = pytest.mark.skipif(
not GEMM_AVAILABLE,
reason="BF16 GEMM requires a built kernel and compute capability 8.0+",
)
@skip_no_gemv
def _assert_close_fp64(actual, x, weight, bias=None):
"""Compare bf16 kernel output vs fp64-exact with ulp-scaled tolerance.
Avoids false failures from cuBLAS default bf16 split-K partial reduction
(which can introduce ~2 ulp diffs on near-tie rounding at long K). Used
for tiled-path tests (M > 8, K >= 4096) where the bf16 accumulation tie
pattern may differ from cuBLAS's."""
exact = x.double() @ weight.double().T
if bias is not None:
exact = exact + bias.double()
ulp = (exact.abs() * 2**-9).clamp(min=2**-9)
max_ulp = ((actual.double() - exact).abs() / ulp).max().item()
assert max_ulp < 8, (
f"max_ulp={max_ulp:.1f} exceeds 8 (exact fp32 accumulation should stay within ~2 ulps)"
)
@skip_no_gemm
@pytest.mark.parametrize(
"n,k",
[(256, 1536), (1536, 1536), (6912, 1536), (1536, 6912), (100000, 1536)],
)
def test_bf16_gemv_matches_linear_shape_families(n, k):
def test_bf16_gemm_matches_linear_shape_families(n, k):
torch.manual_seed(17)
x = torch.randn(k, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(n, k, device="cuda", dtype=torch.bfloat16)
actual = bf16_gemv(x, weight)
actual = bf16_gemm(x, weight)
expected = F.linear(x, weight)
assert actual.shape == (n,)
torch.testing.assert_close(actual, expected, rtol=0.02, atol=0.25)
@skip_no_gemv
@skip_no_gemm
@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):
def test_bf16_gemm_matches_small_decode_batches(m, n, k):
torch.manual_seed(19 + m)
x = torch.randn(m, k, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(n, k, device="cuda", dtype=torch.bfloat16)
actual = bf16_gemv(x, weight)
actual = bf16_gemm(x, weight)
expected = F.linear(x, weight)
assert actual.shape == (m, n)
torch.testing.assert_close(actual, expected, rtol=0.02, atol=0.5)
@skip_no_gemv
@skip_no_gemm
@pytest.mark.parametrize("m", [2, 4])
@pytest.mark.parametrize(
"n,k",
@@ -72,17 +89,17 @@ def test_bf16_gemv_matches_small_decode_batches(m, n, k):
(2048, 8192),
],
)
def test_bf16_gemv_matches_common_transformer_shapes(m, n, k):
def test_bf16_gemm_matches_common_transformer_shapes(m, n, k):
torch.manual_seed(2026 + m + n + k)
x = torch.randn(m, k, device="cuda", dtype=torch.bfloat16)
weight = torch.empty(n, k, device="cuda", dtype=torch.bfloat16)
weight.normal_(mean=0.0, std=0.02)
actual = bf16_gemv(x, weight)
actual = bf16_gemm(x, weight)
expected = F.linear(x, weight)
torch.testing.assert_close(actual, expected, rtol=0.02, atol=0.25)
@skip_no_gemv
@skip_no_gemm
@pytest.mark.parametrize(
"m,n,k",
[
@@ -93,63 +110,63 @@ def test_bf16_gemv_matches_common_transformer_shapes(m, n, k):
(8, 2048, 8192),
],
)
def test_bf16_gemv_matches_m8_edge_bands(m, n, k):
def test_bf16_gemm_matches_m8_edge_bands(m, n, k):
torch.manual_seed(2026 + m + n + k)
x = torch.randn(m, k, device="cuda", dtype=torch.bfloat16)
weight = torch.empty(n, k, device="cuda", dtype=torch.bfloat16)
weight.normal_(mean=0.0, std=0.02)
actual = bf16_gemv(x, weight)
actual = bf16_gemm(x, weight)
expected = F.linear(x, weight)
torch.testing.assert_close(actual, expected, rtol=0.02, atol=0.25)
@skip_no_gemv
def test_bf16_gemv_preserves_singleton_batch_and_fuses_bias():
@skip_no_gemm
def test_bf16_gemm_preserves_singleton_batch_and_fuses_bias():
torch.manual_seed(23)
x = torch.randn(1, 1536, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(1536, 1536, device="cuda", dtype=torch.bfloat16)
bias = torch.randn(1536, device="cuda", dtype=torch.bfloat16)
actual = bf16_gemv(x, weight, bias)
actual = bf16_gemm(x, weight, bias)
expected = F.linear(x, weight, bias)
assert actual.shape == (1, 1536)
torch.testing.assert_close(actual, expected, rtol=0.02, atol=0.25)
@skip_no_gemv
def test_bf16_gemv_small_batch_fuses_bias():
@skip_no_gemm
def test_bf16_gemm_small_batch_fuses_bias():
torch.manual_seed(25)
x = torch.randn(4, 1536, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(256, 1536, device="cuda", dtype=torch.bfloat16)
bias = torch.randn(256, device="cuda", dtype=torch.bfloat16)
actual = bf16_gemv(x, weight, bias)
actual = bf16_gemm(x, weight, bias)
expected = F.linear(x, weight, bias)
torch.testing.assert_close(actual, expected, rtol=0.02, atol=0.25)
@skip_no_gemv
def test_bf16_gemv_uses_current_stream():
@skip_no_gemm
def test_bf16_gemm_uses_current_stream():
x = torch.randn(1536, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(256, 1536, device="cuda", dtype=torch.bfloat16)
torch.cuda.synchronize()
stream = torch.cuda.Stream()
with torch.cuda.stream(stream):
actual = bf16_gemv(x, weight)
actual = bf16_gemm(x, weight)
expected = F.linear(x, weight)
stream.synchronize()
torch.testing.assert_close(actual, expected, rtol=0.02, atol=0.25)
@skip_no_gemv
def test_bf16_gemv_cuda_graph_replay():
@skip_no_gemm
def test_bf16_gemm_cuda_graph_replay():
torch.manual_seed(29)
x = torch.randn(1536, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(1536, 1536, device="cuda", dtype=torch.bfloat16)
for _ in range(3):
bf16_gemv(x, weight)
bf16_gemm(x, weight)
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
actual = bf16_gemv(x, weight)
actual = bf16_gemm(x, weight)
x.copy_(torch.randn_like(x))
graph.replay()
@@ -157,24 +174,24 @@ def test_bf16_gemv_cuda_graph_replay():
torch.testing.assert_close(actual, expected, rtol=0.02, atol=0.25)
@skip_no_gemv
@skip_no_gemm
@pytest.mark.parametrize("n,k", [(64, 7), (64, 12), (33, 100), (256, 1534)])
def test_bf16_gemv_handles_unaligned_k(n, k):
def test_bf16_gemm_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)
actual = bf16_gemm(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)
actual3 = bf16_gemm(x3, weight)
torch.testing.assert_close(actual3, F.linear(x3, weight), rtol=0.02, atol=0.5)
@skip_no_gemv
@skip_no_gemm
@pytest.mark.parametrize("m", [1, 2, 3, 4])
def test_bf16_gemv_accepts_complementary_misalignment(m):
def test_bf16_gemm_accepts_complementary_misalignment(m):
"""Misaligned weight rows plus an x base chosen so the vectorized branch
is entered with a non-16B-aligned ``x`` pointer (regression: the branch
guard checked ``x + whead`` alignment but the uint4 view was rooted at
@@ -189,13 +206,13 @@ def test_bf16_gemv_accepts_complementary_misalignment(m):
x = big_x[5 : 5 + m * k].view(m, k) if m > 1 else big_x[5 : 5 + k]
assert (x.data_ptr() & 15) == 10 and (weight.data_ptr() & 15) == 10
actual = bf16_gemv(x, weight)
actual = bf16_gemm(x, weight)
expected = F.linear(x, weight)
torch.testing.assert_close(actual, expected, rtol=0.02, atol=0.5 if m > 1 else 0.25)
@skip_no_gemv
def test_bf16_gemv_scalar_path_handles_misaligned_weight_only():
@skip_no_gemm
def test_bf16_gemm_scalar_path_handles_misaligned_weight_only():
"""Weight rows misaligned while x stays 16B-aligned take the scalar-x
middle and must stay exact."""
torch.manual_seed(41)
@@ -205,21 +222,21 @@ def test_bf16_gemv_scalar_path_handles_misaligned_weight_only():
x = torch.randn(2, k, device="cuda", dtype=torch.bfloat16)
assert (weight.data_ptr() & 15) == 10 and (x.data_ptr() & 15) == 0
actual = bf16_gemv(x, weight)
actual = bf16_gemm(x, weight)
torch.testing.assert_close(actual, F.linear(x, weight), rtol=0.02, atol=0.5)
@skip_no_gemv
def test_bf16_gemv_small_batch_cuda_graph_replay():
@skip_no_gemm
def test_bf16_gemm_small_batch_cuda_graph_replay():
torch.manual_seed(31)
x = torch.randn(8, 1536, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(1536, 1536, device="cuda", dtype=torch.bfloat16)
for _ in range(3):
bf16_gemv(x, weight)
bf16_gemm(x, weight)
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
actual = bf16_gemv(x, weight)
actual = bf16_gemm(x, weight)
x.copy_(torch.randn_like(x))
graph.replay()
@@ -227,13 +244,13 @@ def test_bf16_gemv_small_batch_cuda_graph_replay():
torch.testing.assert_close(actual, expected, rtol=0.02, atol=0.25)
@skip_no_gemv
@skip_no_gemm
@pytest.mark.parametrize(
"make_args,error",
[
(
lambda: (
torch.randn(9, 16, device="cuda", dtype=torch.bfloat16),
torch.randn(65, 16, device="cuda", dtype=torch.bfloat16),
torch.randn(8, 16, device="cuda", dtype=torch.bfloat16),
),
"M must",
@@ -256,6 +273,124 @@ def test_bf16_gemv_small_batch_cuda_graph_replay():
),
],
)
def test_bf16_gemv_rejects_unsupported_inputs(make_args, error):
def test_bf16_gemm_rejects_unsupported_inputs(make_args, error):
with pytest.raises(RuntimeError, match=error):
bf16_gemv(*make_args())
bf16_gemm(*make_args())
# ---------------------------------------------------------------------------
# Tiled path: M in (8, 64]
# ---------------------------------------------------------------------------
@skip_no_gemm
@pytest.mark.parametrize("m", [9, 12, 16, 17, 24, 32, 33, 48, 64])
@pytest.mark.parametrize("n,k", [(256, 1536), (1536, 1536), (6912, 1536), (1536, 6912)])
def test_bf16_gemm_tiled_matches_decode_batches(m, n, k):
torch.manual_seed(31 + m)
x = torch.randn(m, k, device="cuda", dtype=torch.bfloat16)
weight = torch.empty(n, k, device="cuda", dtype=torch.bfloat16)
weight.normal_(mean=0.0, std=0.02)
actual = bf16_gemm(x, weight)
assert actual.shape == (m, n)
_assert_close_fp64(actual, x, weight)
@skip_no_gemm
@pytest.mark.parametrize("m", [12, 64])
def test_bf16_gemm_tiled_matches_lm_head(m):
# N=100000 fills the SMs with N tiles alone: the splits=1 epilogue.
torch.manual_seed(37 + m)
x = torch.randn(m, 1536, device="cuda", dtype=torch.bfloat16)
weight = torch.empty(100000, 1536, device="cuda", dtype=torch.bfloat16)
weight.normal_(mean=0.0, std=0.02)
actual = bf16_gemm(x, weight)
_assert_close_fp64(actual, x, weight)
@skip_no_gemm
@pytest.mark.parametrize(
"m,n,k",
[
(12, 100, 72),
(12, 96, 1536),
(12, 1632, 1536),
(64, 100, 72),
(17, 200, 8),
(33, 160, 152),
],
)
def test_bf16_gemm_tiled_handles_remainder_tiles(m, n, k):
# N not a multiple of 64 (predicated epilogue columns) and K not a
# multiple of 64 (zero-filled staging chunks).
torch.manual_seed(41 + m + n + k)
x = torch.randn(m, k, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(n, k, device="cuda", dtype=torch.bfloat16)
actual = bf16_gemm(x, weight)
_assert_close_fp64(actual, x, weight)
@skip_no_gemm
@pytest.mark.parametrize("m,n,k", [(12, 1536, 1536), (64, 6912, 1536)])
def test_bf16_gemm_tiled_fuses_bias(m, n, k):
# (12, 1536) exercises the narrow-N deep-K config; (64, 6912) the
# wide-N default.
torch.manual_seed(43 + m)
x = torch.randn(m, k, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(n, k, device="cuda", dtype=torch.bfloat16)
bias = torch.randn(n, device="cuda", dtype=torch.bfloat16)
actual = bf16_gemm(x, weight, bias)
_assert_close_fp64(actual, x, weight, bias)
@skip_no_gemm
def test_bf16_gemm_tiled_deterministic_across_runs():
# Single-pass K accumulation with no atomics: reruns are bitwise
# identical — CUDA Graph replay relies on this.
torch.manual_seed(47)
x = torch.randn(16, 1536, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(1536, 1536, device="cuda", dtype=torch.bfloat16)
first = bf16_gemm(x, weight)
second = bf16_gemm(x, weight)
assert torch.equal(first, second)
@skip_no_gemm
def test_bf16_gemm_tiled_cuda_graph_replay():
torch.manual_seed(53)
x = torch.randn(16, 1536, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(1536, 1536, device="cuda", dtype=torch.bfloat16)
for _ in range(3):
bf16_gemm(x, weight)
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
actual = bf16_gemm(x, weight)
x.copy_(torch.randn_like(x))
graph.replay()
_assert_close_fp64(actual, x, weight)
@skip_no_gemm
def test_bf16_gemm_tiled_rejects_k_not_multiple_of_8():
x = torch.randn(12, 12, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(64, 12, device="cuda", dtype=torch.bfloat16)
with pytest.raises(RuntimeError, match="multiple of 8"):
bf16_gemm(x, weight)
@skip_no_gemm
def test_bf16_gemm_tiled_rejects_misaligned_x():
# A 2-byte storage offset breaks the 16B alignment the tiled path
# stages chunks on; the M <= 8 GEMV path still accepts it.
k = 1536
storage = torch.randn(12 * k + 1, device="cuda", dtype=torch.bfloat16)
x8 = storage[1 : 1 + 8 * k].view(8, k)
x12 = storage[1 : 1 + 12 * k].view(12, k)
weight = torch.randn(1536, k, device="cuda", dtype=torch.bfloat16)
torch.testing.assert_close(
bf16_gemm(x8, weight), F.linear(x8, weight), rtol=0.02, atol=0.5
)
with pytest.raises(RuntimeError, match="16-byte"):
bf16_gemm(x12, weight)
+61 -49
View File
@@ -12,26 +12,26 @@ from astrai.extension.dispatch import explain, op_backend, resolve
# module object explicitly for monkeypatching its private helpers.
linear_module = importlib.import_module("astrai.extension.backend.linear")
GEMV_AVAILABLE = (
GEMM_AVAILABLE = (
torch.cuda.is_available()
and is_available("bf16_gemv")
and is_available("bf16_gemm")
and torch.cuda.get_device_capability() >= (8, 0)
)
skip_no_gemv = pytest.mark.skipif(
not GEMV_AVAILABLE,
reason="BF16 GEMV requires a built kernel and compute capability 8.0+",
skip_no_gemm = pytest.mark.skipif(
not GEMM_AVAILABLE,
reason="BF16 GEMM requires a built kernel and compute capability 8.0+",
)
def _routes_to_gemv(monkeypatch, x, weight, bias=None) -> bool:
"""Patch the GEMV entry point to a sentinel and report whether
def _routes_to_gemm(monkeypatch, x, weight, bias=None) -> bool:
"""Patch the GEMM entry point to a sentinel and report whether
``linear`` selected it (torch fallback would compute a real tensor)."""
sentinel = object()
def fake_gemv(x, weight, bias):
def fake_gemm(x, weight, bias):
return sentinel
monkeypatch.setattr(linear_module, "_inference_bf16_gemv", fake_gemv)
monkeypatch.setattr(linear_module, "_inference_bf16_gemm", fake_gemm)
return linear(x, weight, bias) is sentinel
@@ -52,7 +52,7 @@ def test_model_linear_routes_through_backend(monkeypatch):
def test_invalid_mode_warns_and_uses_auto(monkeypatch, caplog):
monkeypatch.setenv("ASTRAI_GEMV", "invalid-test-mode")
monkeypatch.setenv("ASTRAI_GEMM", "invalid-test-mode")
x = torch.randn(2, 8)
weight = torch.randn(4, 8)
with caplog.at_level(logging.WARNING):
@@ -62,7 +62,7 @@ def test_invalid_mode_warns_and_uses_auto(monkeypatch, caplog):
def test_cpu_and_training_calls_fall_back_to_torch(monkeypatch):
monkeypatch.setenv("ASTRAI_GEMV", "1")
monkeypatch.setenv("ASTRAI_GEMM", "1")
x = torch.randn(2, 8, requires_grad=True)
weight = torch.randn(4, 8, requires_grad=True)
actual = linear(x, weight)
@@ -73,71 +73,83 @@ def test_cpu_and_training_calls_fall_back_to_torch(monkeypatch):
assert weight.grad is not None
@skip_no_gemv
@pytest.mark.parametrize("m", [2, 3, 4])
@skip_no_gemm
@pytest.mark.parametrize("m", [1, 2, 3, 4, 5, 6, 7, 8])
def test_auto_selects_small_decode_batches(monkeypatch, m):
monkeypatch.setenv("ASTRAI_GEMV", "auto")
monkeypatch.setenv("ASTRAI_GEMM", "auto")
x = torch.randn(m, 1536, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(1536, 1536, device="cuda", dtype=torch.bfloat16)
with torch.no_grad():
assert _routes_to_gemv(monkeypatch, x, weight)
assert _routes_to_gemm(monkeypatch, x, weight)
@skip_no_gemv
@pytest.mark.parametrize("m", [1, 5, 8, 9])
@skip_no_gemm
@pytest.mark.parametrize("m", [12, 16, 24, 32])
def test_auto_selects_larger_decode_batches(monkeypatch, m):
"""Auto covers M up to 32; 48+ loses to cuBLAS on long-K shapes."""
monkeypatch.setenv("ASTRAI_GEMM", "auto")
x = torch.randn(m, 1536, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(1536, 1536, device="cuda", dtype=torch.bfloat16)
with torch.no_grad():
assert _routes_to_gemm(monkeypatch, x, weight)
@skip_no_gemm
@pytest.mark.parametrize("m", [48, 64, 65])
def test_auto_falls_back_outside_band(monkeypatch, m):
monkeypatch.setenv("ASTRAI_GEMV", "auto")
"""M beyond 32 falls back to cuBLAS (measured regression at M=48+)."""
monkeypatch.setenv("ASTRAI_GEMM", "auto")
x = torch.randn(m, 1536, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(1536, 1536, device="cuda", dtype=torch.bfloat16)
with torch.no_grad():
assert not _routes_to_gemv(monkeypatch, x, weight)
assert not _routes_to_gemm(monkeypatch, x, weight)
torch.testing.assert_close(
linear(x, weight), F.linear(x, weight), rtol=0.02, atol=0.25
)
@skip_no_gemv
def test_mode_zero_disables_gemv(monkeypatch):
monkeypatch.setenv("ASTRAI_GEMV", "0")
@skip_no_gemm
def test_mode_zero_disables_gemm(monkeypatch):
monkeypatch.setenv("ASTRAI_GEMM", "0")
x = torch.randn(2, 1536, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(1536, 1536, device="cuda", dtype=torch.bfloat16)
with torch.no_grad():
assert not _routes_to_gemv(monkeypatch, x, weight)
assert not _routes_to_gemm(monkeypatch, x, weight)
torch.testing.assert_close(linear(x, weight), F.linear(x, weight))
@skip_no_gemv
@pytest.mark.parametrize("m", [1, 2, 8])
@skip_no_gemm
@pytest.mark.parametrize("m", [1, 2, 8, 16, 32])
def test_mode_one_forces_every_capable_batch(monkeypatch, m):
monkeypatch.setenv("ASTRAI_GEMV", "1")
monkeypatch.setenv("ASTRAI_GEMM", "1")
x = torch.randn(m, 1536, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(256, 1536, device="cuda", dtype=torch.bfloat16)
with torch.no_grad():
assert _routes_to_gemv(monkeypatch, x, weight)
assert _routes_to_gemm(monkeypatch, x, weight)
@skip_no_gemv
@skip_no_gemm
def test_mode_one_rejects_oversized_batch_and_grad(monkeypatch):
monkeypatch.setenv("ASTRAI_GEMV", "1")
monkeypatch.setenv("ASTRAI_GEMM", "1")
weight = torch.randn(
256, 1536, device="cuda", dtype=torch.bfloat16, requires_grad=True
)
with torch.no_grad():
oversized = torch.randn(9, 1536, device="cuda", dtype=torch.bfloat16)
assert not _routes_to_gemv(monkeypatch, oversized, weight)
assert not _routes_to_gemv(
oversized = torch.randn(65, 1536, device="cuda", dtype=torch.bfloat16)
assert not _routes_to_gemm(monkeypatch, oversized, weight)
assert not _routes_to_gemm(
monkeypatch, torch.randn(2, 1536, device="cuda", dtype=torch.bfloat16), weight
)
@skip_no_gemv
@skip_no_gemm
def test_mode_one_supports_bias_and_vector_input(monkeypatch):
monkeypatch.setenv("ASTRAI_GEMV", "1")
monkeypatch.setenv("ASTRAI_GEMM", "1")
x = torch.randn(1, 1536, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(256, 1536, device="cuda", dtype=torch.bfloat16)
bias = torch.randn(256, device="cuda", dtype=torch.bfloat16)
with torch.no_grad():
assert _routes_to_gemv(monkeypatch, x, weight, bias)
assert _routes_to_gemm(monkeypatch, x, weight, bias)
monkeypatch.undo()
torch.testing.assert_close(
linear(x, weight, bias),
@@ -147,9 +159,9 @@ def test_mode_one_supports_bias_and_vector_input(monkeypatch):
)
@skip_no_gemv
@skip_no_gemm
def test_dispatched_linear_cuda_graph_replay(monkeypatch):
monkeypatch.setenv("ASTRAI_GEMV", "1")
monkeypatch.setenv("ASTRAI_GEMM", "1")
x = torch.randn(1, 1536, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(256, 1536, device="cuda", dtype=torch.bfloat16)
with torch.no_grad():
@@ -176,44 +188,44 @@ def test_linear_family_is_registered_with_shared_dispatcher():
assert "linear" in explain("linear", x, weight)
@skip_no_gemv
@skip_no_gemm
def test_ops_env_override_forces_torch_for_capable_call(monkeypatch):
"""ASTR_OPS=linear=torch must keep working after the M-band rewrite
(regression: the family was silently dropped from the dispatcher, so
the override warned, fell through, and the gemv kernel still ran)."""
the override warned, fell through, and the gemm kernel still ran)."""
monkeypatch.setenv("ASTR_OPS", "linear=torch")
x = torch.randn(2, 1536, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(1536, 1536, device="cuda", dtype=torch.bfloat16)
with torch.no_grad():
assert not _routes_to_gemv(monkeypatch, x, weight)
assert not _routes_to_gemm(monkeypatch, x, weight)
torch.testing.assert_close(
linear(x, weight), F.linear(x, weight), rtol=0.02, atol=0.25
)
@skip_no_gemv
def test_ops_env_override_forces_gemv(monkeypatch):
monkeypatch.setenv("ASTR_OPS", "linear=gemv")
@skip_no_gemm
def test_ops_env_override_forces_gemm(monkeypatch):
monkeypatch.setenv("ASTR_OPS", "linear=gemm")
x = torch.randn(1, 1536, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(256, 1536, device="cuda", dtype=torch.bfloat16)
with torch.no_grad():
# M=1 is outside the auto band but inside the forced gemv record.
assert _routes_to_gemv(monkeypatch, x, weight)
# M=1 is outside the auto band but inside the forced gemm record.
assert _routes_to_gemm(monkeypatch, x, weight)
@skip_no_gemv
@skip_no_gemm
def test_op_backend_context_selects_torch(monkeypatch):
monkeypatch.setenv("ASTRAI_GEMV", "1")
monkeypatch.setenv("ASTRAI_GEMM", "1")
x = torch.randn(2, 1536, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(256, 1536, device="cuda", dtype=torch.bfloat16)
with torch.no_grad(), op_backend(linear="torch"):
assert not _routes_to_gemv(monkeypatch, x, weight)
assert not _routes_to_gemm(monkeypatch, x, weight)
torch.testing.assert_close(
linear(x, weight), F.linear(x, weight), rtol=0.02, atol=0.25
)
# The override is scoped: the forced mode applies again afterwards.
with torch.no_grad():
assert _routes_to_gemv(monkeypatch, x, weight)
assert _routes_to_gemm(monkeypatch, x, weight)
def test_op_backend_rejects_unknown_linear_handle():
+3 -4
View File
@@ -86,10 +86,9 @@ def test_mode_one_forces_supported_shape(monkeypatch):
@skip_no_swiglu
def test_auto_uses_unfused_chain_until_shape_is_qualified(monkeypatch):
# The fusion table is empty, so auto keeps the unfused linear-backend
# chain. The linear backend may still dispatch its own GEMV for M=4,
# hence the relaxed tolerance versus the pure-torch reference.
def test_auto_uses_fused_chain_for_decode_batches(monkeypatch):
# Auto adopts the fused primitive for the decode band; numerics match
# the unfused linear-backend chain within BF16 accumulation-order noise.
monkeypatch.setenv("ASTRAI_SWIGLU", "auto")
x = torch.randn(4, 1536, device="cuda", dtype=torch.bfloat16)
up_weight = torch.randn(6912, 1536, device="cuda", dtype=torch.bfloat16)