From d4a292b36b32614030c4fb9bc6937684320b1a87 Mon Sep 17 00:00:00 2001 From: 0z5a Date: Thu, 3 Sep 2026 04:26:53 +0800 Subject: [PATCH] perf: tune bf16 gemv and add opt-in fused swiglu - deepen common-shape BF16 GEMV tuning with warp-row tiling for LLaMA/Qwen2/GPT-NeoX/OPT decode projections - add fused BF16 up/gate SwiGLU CUDA primitive with ASTRAI_SWIGLU=0/1/auto dispatch - keep the unfused linear backend as the default path; auto enables no shape until per-architecture checkpoint gates pass - fall back to the linear/torch chain when kernels are absent, on CPU, in training, or outside supported M/K/dtype shapes - add gemv/swiglu benchmark scripts, dispatch and parity tests, and kernel documentation Benchmark: NVIDIA L20 (sm_89), CUDA 12.8, PyTorch 2.11.0+cu128, idle GPU. AstrAI 1B config (24 layers, hidden 1536, vocab 100000), BF16, prompt 128, 32 greedy decode tokens, CUDA graphs enabled, A/B in separate interleaved processes (3 rounds, 8 trials each, medians). Default vs ASTRAI_SWIGLU=1 per generate call: batch 1 134.8->129.1 ms (+4.44%), batch 2 136.2->130.9 ms (+4.06%), batch 4 145.5->140.3 ms (+3.66%). Greedy output identical at batch 1, differs at batch 2/4, so auto stays unfused by default; kernelless fallback verified bit-identical greedy. --- astrai/extension/__init__.py | 4 + astrai/extension/backend/__init__.py | 2 + astrai/extension/backend/linear.py | 70 +++- astrai/extension/backend/swiglu.py | 109 ++++++ astrai/extension/ops/__init__.py | 2 + astrai/extension/ops/swiglu.py | 22 ++ astrai/model/components/mlp.py | 3 +- csrc/CMakeLists.txt | 2 + csrc/kernels/gemv/bf16_gemv.cu | 146 ++++++- csrc/kernels/gemv/bf16_swiglu.cu | 368 ++++++++++++++++++ docs/developer/cuda_kernels.md | 123 +++++- docs/developer/decode_linear_benchmark.md | 19 + docs/developer/swiglu_benchmark.md | 58 +++ scripts/tools/benchmark_gemv_common.py | 450 ++++++++++++++++++++++ scripts/tools/benchmark_swiglu.py | 325 ++++++++++++++++ setup.py | 1 + tests/extension/test_gemv.py | 60 +++ tests/extension/test_linear_dispatch.py | 88 ++++- tests/extension/test_swiglu.py | 99 +++++ tests/extension/test_swiglu_dispatch.py | 94 +++++ 20 files changed, 2008 insertions(+), 37 deletions(-) create mode 100644 astrai/extension/backend/swiglu.py create mode 100644 astrai/extension/ops/swiglu.py create mode 100644 csrc/kernels/gemv/bf16_swiglu.cu create mode 100644 docs/developer/swiglu_benchmark.md create mode 100644 scripts/tools/benchmark_gemv_common.py create mode 100644 scripts/tools/benchmark_swiglu.py create mode 100644 tests/extension/test_swiglu.py create mode 100644 tests/extension/test_swiglu_dispatch.py diff --git a/astrai/extension/__init__.py b/astrai/extension/__init__.py index 6e0ba08..bbfcc1f 100644 --- a/astrai/extension/__init__.py +++ b/astrai/extension/__init__.py @@ -27,6 +27,7 @@ from astrai.extension.backend import ( attn_backend, get_backend, linear, + swiglu, ) from astrai.extension.dispatch import ( Axes, @@ -51,6 +52,7 @@ from astrai.extension.ops import ( attn_paged_decode, attn_prefill, bf16_gemv, + bf16_swiglu, ) __all__ = [ @@ -65,10 +67,12 @@ __all__ = [ "attn_backend", "get_backend", "linear", + "swiglu", "attn_decode", "attn_paged_decode", "attn_prefill", "bf16_gemv", + "bf16_swiglu", "is_available", "KERNEL_NAMES", "apply_rotary_emb", diff --git a/astrai/extension/backend/__init__.py b/astrai/extension/backend/__init__.py index 7b43a21..8c4127a 100644 --- a/astrai/extension/backend/__init__.py +++ b/astrai/extension/backend/__init__.py @@ -13,6 +13,7 @@ from astrai.extension.backend.attention import ( ) from astrai.extension.backend.linear import linear from astrai.extension.backend.rotary import apply_rotary_emb +from astrai.extension.backend.swiglu import swiglu __all__ = [ "ATTN_BACKEND", @@ -26,4 +27,5 @@ __all__ = [ "attn_backend", "get_backend", "linear", + "swiglu", ] diff --git a/astrai/extension/backend/linear.py b/astrai/extension/backend/linear.py index 4c85c82..38fd275 100644 --- a/astrai/extension/backend/linear.py +++ b/astrai/extension/backend/linear.py @@ -1,9 +1,9 @@ """Inference-only dispatch for AstrAI linear layers. The CUDA GEMV path is deliberately narrow: automatic selection is enabled -only for single-row BF16 shapes measured to beat ``F.linear`` on a supported -architecture. Every training, prefill, unsupported-layout, and unmeasured -call falls back to PyTorch. +only for small decode batches and BF16 shapes measured to beat ``F.linear`` +on a supported architecture. Every training, prefill, unsupported-layout, +and unmeasured call falls back to PyTorch. """ import logging @@ -30,20 +30,72 @@ from astrai.extension.ops.gemv import bf16_gemv logger = logging.getLogger(__name__) # Shape keys are (N, K) for Y[M, N] = X[M, K] @ W[N, K].T. A band is -# automatic only after both the per-shape >=5% and end-to-end decode >=3% -# gates pass and checkpoint greedy output remains stable. M=1 and M=8 remain -# empty on SM89; the safe M=2/4 bands improve real-engine throughput by -# 11.8-14.0%. +# automatic only after both the per-shape >=5% and projection-chain/engine +# >=3% gates pass and output argmax remains stable. M=1 is limited to OPT 1.3B; +# M=8 remains empty because at least one projection in each measured family +# misses the per-shape gate even when its aggregate chain result is positive. +_COMMON_TRANSFORMER_SM89_SHAPES = frozenset( + { + (1024, 4096), # LLaMA 3 8B K/V + (4096, 4096), # LLaMA 2/3 7B/8B Q/O + (11008, 4096), # LLaMA 2 7B gate/up + (4096, 11008), # LLaMA 2 7B down + (14336, 4096), # LLaMA 3 8B gate/up + (4096, 14336), # LLaMA 3 8B down + (5120, 5120), # LLaMA 2 13B Q/K/V/O + (13824, 5120), # LLaMA 2 13B gate/up + (5120, 13824), # LLaMA 2 13B down + (16384, 4096), # GPT-NeoX MLP up + (4096, 16384), # GPT-NeoX MLP down + } +) +_COMMON_TRANSFORMER_SM89_M4_SHAPES = _COMMON_TRANSFORMER_SM89_SHAPES - { + (4096, 4096), + (11008, 4096), + (4096, 11008), +} +_QWEN2_7B_SM89_SHAPES = frozenset( + { + (512, 3584), # K/V + (3584, 3584), # Q/O + (18944, 3584), # gate/up + (3584, 18944), # down + } +) +_LLAMA3_70B_SM89_SHAPES = frozenset( + { + (1024, 8192), # K/V + (8192, 8192), # Q/O + (28672, 8192), # gate/up + (8192, 28672), # down + } +) +_OPT_1_3B_SM89_SHAPES = frozenset( + { + (2048, 2048), # Q/K/V/O + (8192, 2048), # MLP up + (2048, 8192), # MLP down + } +) + _AUTO_GEMV_SHAPES: dict[tuple[int, int], dict[int, frozenset[tuple[int, int]]]] = { (8, 9): { - 2: frozenset( + 1: _OPT_1_3B_SM89_SHAPES, + 2: _COMMON_TRANSFORMER_SM89_SHAPES + | _QWEN2_7B_SM89_SHAPES + | _LLAMA3_70B_SM89_SHAPES + | _OPT_1_3B_SM89_SHAPES + | frozenset( { (256, 1536), (1536, 1536), (100000, 1536), } ), - 4: frozenset({(256, 1536), (1536, 1536)}), + 4: _COMMON_TRANSFORMER_SM89_M4_SHAPES + | _QWEN2_7B_SM89_SHAPES + | _LLAMA3_70B_SM89_SHAPES + | frozenset({(256, 1536), (1536, 1536)}), } } _AUTO_GEMV_M = frozenset( diff --git a/astrai/extension/backend/swiglu.py b/astrai/extension/backend/swiglu.py new file mode 100644 index 0000000..7f4fe7c --- /dev/null +++ b/astrai/extension/backend/swiglu.py @@ -0,0 +1,109 @@ +"""Inference-only fused SwiGLU selection for dense MLP layers.""" + +import logging +import os +from functools import cache + +import torch +import torch.nn.functional as F +from torch import Tensor + +from astrai.extension.backend.linear import linear +from astrai.extension.loader import is_available +from astrai.extension.ops.swiglu import bf16_swiglu + +logger = logging.getLogger(__name__) + +# Shape keys are (N, K) for the paired up/gate projections. Automatic entries +# are populated only after the primitive, MLP chain, and greedy checkpoint +# gates pass on that architecture. +_AUTO_SWIGLU_SHAPES: dict[tuple[int, int], dict[int, frozenset[tuple[int, int]]]] = {} +_AUTO_SWIGLU_M = frozenset( + m for architecture in _AUTO_SWIGLU_SHAPES.values() for m in architecture +) +_VALID_MODES = {"0", "1", "auto"} +_WARNED_MODES: set[str] = set() + + +def _swiglu_mode() -> str: + mode = os.environ.get("ASTRAI_SWIGLU", "auto").strip().lower() + if mode in _VALID_MODES: + return mode + if mode not in _WARNED_MODES: + _WARNED_MODES.add(mode) + logger.warning( + "ASTRAI_SWIGLU=%r is invalid; expected 0, 1, or auto; using auto", + mode, + ) + return "auto" + + +def _unfused_swiglu(x: Tensor, up_weight: Tensor, gate_weight: Tensor) -> Tensor: + # Keep the existing linear backend in the fallback chain. This preserves + # any independently qualified GEMV shapes instead of making the fusion + # decision suppress linear-level optimizations. + return linear(x, up_weight) * F.silu(linear(x, gate_weight)) + + +def _fused_swiglu(x: Tensor, up_weight: Tensor, gate_weight: Tensor) -> Tensor: + return bf16_swiglu(x.detach(), up_weight.detach(), gate_weight.detach()) + + +@cache +def _device_capability(device_index: int) -> tuple[int, int]: + return torch.cuda.get_device_capability(device_index) + + +def _swiglu_capable(x: Tensor, up_weight: Tensor, gate_weight: Tensor) -> bool: + return not ( + torch.is_grad_enabled() + or not x.is_cuda + or x.dtype != torch.bfloat16 + or up_weight.dtype != torch.bfloat16 + or gate_weight.dtype != torch.bfloat16 + or x.ndim not in (1, 2) + or up_weight.ndim != 2 + or gate_weight.ndim != 2 + or (x.ndim == 2 and not 1 <= x.shape[0] <= 8) + or up_weight.shape != gate_weight.shape + or x.shape[-1] != up_weight.shape[1] + or x.shape[-1] % 8 != 0 + or x.device != up_weight.device + or x.device != gate_weight.device + or not x.is_contiguous() + or not up_weight.is_contiguous() + or not gate_weight.is_contiguous() + or not is_available("bf16_swiglu") + ) + + +def _auto_swiglu_shape(x: Tensor, up_weight: Tensor) -> bool: + capability = _device_capability(x.get_device()) + m = 1 if x.ndim == 1 else x.shape[0] + return (up_weight.shape[0], up_weight.shape[1]) in _AUTO_SWIGLU_SHAPES.get( + capability, {} + ).get(m, ()) + + +def swiglu(x: Tensor, up_weight: Tensor, gate_weight: Tensor) -> Tensor: + """Apply the dense-MLP SwiGLU projection with a safe torch fallback. + + ``ASTRAI_SWIGLU=0`` keeps the unfused linear-backend chain, ``1`` forces + the fused primitive for supported inputs, and ``auto`` uses only + architecture/shape bands backed by benchmark and checkpoint evidence. + """ + mode = _swiglu_mode() + if mode == "0" or (mode == "auto" and not _AUTO_SWIGLU_SHAPES): + return _unfused_swiglu(x, up_weight, gate_weight) + if mode == "auto": + m = 1 if x.ndim == 1 else (x.shape[0] if x.ndim == 2 else None) + if m not in _AUTO_SWIGLU_M: + return _unfused_swiglu(x, up_weight, gate_weight) + if _swiglu_capable(x, up_weight, gate_weight) and ( + mode == "1" or _auto_swiglu_shape(x, up_weight) + ): + return _fused_swiglu(x, up_weight, gate_weight) + return _unfused_swiglu(x, up_weight, gate_weight) + + +__all__ = ["swiglu"] diff --git a/astrai/extension/ops/__init__.py b/astrai/extension/ops/__init__.py index ba06b2f..233ba8f 100644 --- a/astrai/extension/ops/__init__.py +++ b/astrai/extension/ops/__init__.py @@ -9,6 +9,7 @@ from astrai.extension.ops.attention import ( ) from astrai.extension.ops.gemv import bf16_gemv from astrai.extension.ops.rotary import rotary_emb +from astrai.extension.ops.swiglu import bf16_swiglu __all__ = [ "TensorLayout", @@ -17,5 +18,6 @@ __all__ = [ "attn_paged_prefill", "attn_prefill", "bf16_gemv", + "bf16_swiglu", "rotary_emb", ] diff --git a/astrai/extension/ops/swiglu.py b/astrai/extension/ops/swiglu.py new file mode 100644 index 0000000..e1bfc24 --- /dev/null +++ b/astrai/extension/ops/swiglu.py @@ -0,0 +1,22 @@ +"""Stateless wrapper for the directly callable fused BF16 SwiGLU primitive.""" + +import torch + +from astrai.extension.loader import get_module + + +def bf16_swiglu( + x: torch.Tensor, + up_weight: torch.Tensor, + gate_weight: torch.Tensor, +) -> torch.Tensor: + """Compute ``linear(x, up) * silu(linear(x, gate))`` for M in [1, 8]. + + Inputs must be contiguous BF16 CUDA tensors. Both weights use row-major + ``[N, K]`` storage with identical shapes, and K must be divisible by 8. + The primitive is inference-only and performs no fallback. + """ + return get_module("bf16_swiglu").bf16_swiglu(x, up_weight, gate_weight) + + +__all__ = ["bf16_swiglu"] diff --git a/astrai/model/components/mlp.py b/astrai/model/components/mlp.py index 083dad3..defd8d4 100644 --- a/astrai/model/components/mlp.py +++ b/astrai/model/components/mlp.py @@ -5,6 +5,7 @@ import torch.nn as nn import torch.nn.functional as F from torch import Tensor +from astrai.extension.backend.swiglu import swiglu from astrai.factory import BaseFactory from astrai.model.components.linear import Linear @@ -38,7 +39,7 @@ class MLP(nn.Module): self.down = Linear(dim_ffn, dim, init_std=down_init_std) def forward(self, x: Tensor) -> FFNOutput: - gated = self.up(x) * F.silu(self.gate(x)) + gated = swiglu(x, self.up.weight, self.gate.weight) out = self.down(gated) return {"hidden_states": out, "aux_loss": None, "router_stats": None} diff --git a/csrc/CMakeLists.txt b/csrc/CMakeLists.txt index bdd6abb..7f32c23 100644 --- a/csrc/CMakeLists.txt +++ b/csrc/CMakeLists.txt @@ -62,6 +62,7 @@ set(KERNEL_NAMES attn_paged_decode attn_paged_prefill bf16_gemv + bf16_swiglu rotary_emb ) set(KERNEL_SRCS @@ -70,6 +71,7 @@ set(KERNEL_SRCS attention/paged_decode.cu attention/paged_prefill.cu gemv/bf16_gemv.cu + gemv/bf16_swiglu.cu rotary_emb.cu ) diff --git a/csrc/kernels/gemv/bf16_gemv.cu b/csrc/kernels/gemv/bf16_gemv.cu index 70cce50..58278a8 100644 --- a/csrc/kernels/gemv/bf16_gemv.cu +++ b/csrc/kernels/gemv/bf16_gemv.cu @@ -12,7 +12,9 @@ namespace { constexpr int kThreads = 256; +constexpr int kHalfCtaThreads = 128; constexpr int kWarpSize = 32; +constexpr int kWarpTiledThreads = 128; __device__ __forceinline__ float warp_sum(float value) { #pragma unroll @@ -22,7 +24,7 @@ __device__ __forceinline__ float warp_sum(float value) { return value; } -template +template __global__ void bf16_gemv_kernel( const __nv_bfloat16* __restrict__ x, const __nv_bfloat16* __restrict__ weight, @@ -36,7 +38,7 @@ __global__ void bf16_gemv_kernel( const int warp = threadIdx.x / kWarpSize; float sums[Rows] = {}; - __shared__ float warp_sums[Rows][kThreads / kWarpSize]; + __shared__ float warp_sums[Rows][Threads / kWarpSize]; // Weight row: scalar head/tail around a 16-byte-aligned uint4 middle so // any K is accepted while keeping 128-bit weight loads, which dominate // bandwidth on decode shapes. x pairs with scalar loads: it is a tiny @@ -129,7 +131,6 @@ __global__ void bf16_gemv_kernel( } } - #pragma unroll for (int row = 0; row < Rows; ++row) { sums[row] = warp_sum(sums[row]); @@ -146,7 +147,7 @@ __global__ void bf16_gemv_kernel( #pragma unroll for (int row = 0; row < Rows; ++row) { float sum = - lane < (kThreads / kWarpSize) ? warp_sums[row][lane] : 0.0f; + lane < (Threads / kWarpSize) ? warp_sums[row][lane] : 0.0f; sum = warp_sum(sum); if (lane == 0) { if (bias != nullptr) { @@ -158,6 +159,120 @@ __global__ void bf16_gemv_kernel( } } +template +__global__ void bf16_gemv_aligned_warp_tiled_kernel( + const __nv_bfloat16* __restrict__ x, + const __nv_bfloat16* __restrict__ weight, + const __nv_bfloat16* __restrict__ bias, + __nv_bfloat16* __restrict__ output, + int n, + int k +) { + constexpr int kWarpsPerBlock = kWarpTiledThreads / kWarpSize; + const int lane = threadIdx.x & (kWarpSize - 1); + const int warp = threadIdx.x / kWarpSize; + const int output_index = blockIdx.x * kWarpsPerBlock + warp; + if (output_index >= n) { + return; + } + + // The launcher selects this path only when each row is 16-byte aligned. + // Four independent output rows per CTA remove the block-wide reduction + // barrier and improve occupancy for the medium LLaMA projection bands. + const int vectors = k / 8; + const auto* x4 = reinterpret_cast(x); + const auto* w4 = reinterpret_cast(weight) + + static_cast(output_index) * vectors; + float sums[Rows] = {}; + for (int vector = lane; vector < vectors; vector += kWarpSize) { + const uint4 wv_raw = w4[vector]; + const auto* wv = reinterpret_cast(&wv_raw); +#pragma unroll + for (int row = 0; row < Rows; ++row) { + const uint4 xv_raw = + x4[static_cast(row) * vectors + vector]; + const auto* xv = reinterpret_cast(&xv_raw); +#pragma unroll + for (int pair = 0; pair < 4; ++pair) { + sums[row] = fmaf( + __bfloat162float(__low2bfloat16(xv[pair])), + __bfloat162float(__low2bfloat16(wv[pair])), + sums[row] + ); + sums[row] = fmaf( + __bfloat162float(__high2bfloat16(xv[pair])), + __bfloat162float(__high2bfloat16(wv[pair])), + sums[row] + ); + } + } + } + +#pragma unroll + for (int row = 0; row < Rows; ++row) { + sums[row] = warp_sum(sums[row]); + if (lane == 0) { + if (bias != nullptr) { + sums[row] += __bfloat162float(bias[output_index]); + } + output[row * n + output_index] = __float2bfloat16_rn(sums[row]); + } + } +} + +template +constexpr bool use_warp_tiled_kernel(int n, int k) { + // These bands are intentionally narrow and are validated by the common + // transformer benchmark. The 256-thread cooperative kernel remains the + // fallback for arbitrary K, larger projections, and M=2 (where the + // single-warp reduction regresses the current vectorized kernel). + if constexpr (Rows == 4) { + return (n == 1024 && k == 4096) || + (n == 4096 && k == 4096) || + (n == 11008 && k == 4096) || + (n == 4096 && k == 11008); + } + return false; +} + +template +constexpr bool use_half_cta_kernel(int n, int k) { + // A 128-thread CTA reduces synchronization and scheduling overhead for + // selected medium decode projections. Keep the selector exact: long-K + // and bandwidth-saturated shapes regress, and the winning bands differ + // materially with the number of reused input rows. + if constexpr (Rows == 1) { + return n == 8192 && k == 2048; + } + if constexpr (Rows == 2) { + return (n == 4096 && k == 4096) || + (n == 11008 && k == 4096) || + (n == 3584 && k == 3584) || + (n == 2048 && k == 2048) || + (n == 8192 && k == 2048); + } + if constexpr (Rows == 4) { + return (n == 5120 && k == 5120) || + (n == 3584 && k == 3584) || + (n == 2048 && k == 2048) || + (n == 8192 && k == 2048); + } + if constexpr (Rows == 8) { + return (n == 4096 && k == 4096) || + (n == 11008 && k == 4096) || + (n == 4096 && k == 11008) || + (n == 1024 && k == 4096) || + (n == 5120 && k == 5120) || + (n == 512 && k == 3584) || + (n == 3584 && k == 3584) || + (n == 1024 && k == 8192) || + (n == 2048 && k == 2048) || + (n == 8192 && k == 2048) || + (n == 2048 && k == 8192); + } + return false; +} + template void launch_bf16_gemv( const __nv_bfloat16* x, @@ -168,7 +283,28 @@ void launch_bf16_gemv( int k, cudaStream_t stream ) { - bf16_gemv_kernel<<>>( + const bool aligned_rows = k % 8 == 0 && + (reinterpret_cast(x) & 15u) == 0u && + (reinterpret_cast(weight) & 15u) == 0u; + if constexpr (Rows == 4) { + if (aligned_rows && use_warp_tiled_kernel(n, k)) { + constexpr int kWarpsPerBlock = kWarpTiledThreads / kWarpSize; + const int blocks = (n + kWarpsPerBlock - 1) / kWarpsPerBlock; + bf16_gemv_aligned_warp_tiled_kernel + <<>>( + x, weight, bias, output, n, k + ); + return; + } + } + if (aligned_rows && use_half_cta_kernel(n, k)) { + bf16_gemv_kernel + <<>>( + x, weight, bias, output, n, k + ); + return; + } + bf16_gemv_kernel<<>>( x, weight, bias, output, n, k ); } diff --git a/csrc/kernels/gemv/bf16_swiglu.cu b/csrc/kernels/gemv/bf16_swiglu.cu new file mode 100644 index 0000000..99fe0b1 --- /dev/null +++ b/csrc/kernels/gemv/bf16_swiglu.cu @@ -0,0 +1,368 @@ +// Fused small-M BF16 SwiGLU primitive for decode-time dense MLP layers. + +#include +#include +#include +#include +#include + +#include +#include + +namespace { + +constexpr int kThreads = 256; +constexpr int kWarpSize = 32; +constexpr int kWarps = kThreads / kWarpSize; + +__device__ __forceinline__ float warp_sum(float value) { +#pragma unroll + for (int offset = kWarpSize / 2; offset > 0; offset >>= 1) { + value += __shfl_down_sync(0xffffffff, value, offset); + } + return value; +} + +__device__ __forceinline__ float round_bf16(float value) { + return __bfloat162float(__float2bfloat16_rn(value)); +} + +template +__global__ void bf16_swiglu_kernel( + const __nv_bfloat16* __restrict__ x, + const __nv_bfloat16* __restrict__ up_weight, + const __nv_bfloat16* __restrict__ gate_weight, + __nv_bfloat16* __restrict__ output, + int n, + int k +) { + const int output_index = blockIdx.x; + const int lane = threadIdx.x & (kWarpSize - 1); + const int warp = threadIdx.x / kWarpSize; + const int vector_count = k / 8; + + float up_sums[Rows] = {}; + float gate_sums[Rows] = {}; + __shared__ float up_warp_sums[Rows][kWarps]; + __shared__ float gate_warp_sums[Rows][kWarps]; + + const auto* x4 = reinterpret_cast(x); + const auto* up4 = reinterpret_cast( + up_weight + static_cast(output_index) * k + ); + const auto* gate4 = reinterpret_cast( + gate_weight + static_cast(output_index) * k + ); + + // Read each pair of up/gate weight chunks once per CTA, then reuse it for + // every active decode row. The fused epilogue removes two [M, N] + // intermediates and the standalone SiLU and multiply launches. + for (int vector_index = threadIdx.x; + vector_index < vector_count; + vector_index += blockDim.x) { + const uint4 up_raw = up4[vector_index]; + const uint4 gate_raw = gate4[vector_index]; + const auto* up_values = + reinterpret_cast(&up_raw); + const auto* gate_values = + reinterpret_cast(&gate_raw); + +#pragma unroll + for (int row = 0; row < Rows; ++row) { + const uint4 x_raw = + x4[static_cast(row) * vector_count + vector_index]; + const auto* x_values = + reinterpret_cast(&x_raw); +#pragma unroll + for (int pair = 0; pair < 4; ++pair) { + const float2 xv = __bfloat1622float2(x_values[pair]); + const float2 uv = __bfloat1622float2(up_values[pair]); + const float2 gv = __bfloat1622float2(gate_values[pair]); + up_sums[row] = fmaf(xv.x, uv.x, up_sums[row]); + up_sums[row] = fmaf(xv.y, uv.y, up_sums[row]); + gate_sums[row] = fmaf(xv.x, gv.x, gate_sums[row]); + gate_sums[row] = fmaf(xv.y, gv.y, gate_sums[row]); + } + } + } + +#pragma unroll + for (int row = 0; row < Rows; ++row) { + up_sums[row] = warp_sum(up_sums[row]); + gate_sums[row] = warp_sum(gate_sums[row]); + } + if (lane == 0) { +#pragma unroll + for (int row = 0; row < Rows; ++row) { + up_warp_sums[row][warp] = up_sums[row]; + gate_warp_sums[row][warp] = gate_sums[row]; + } + } + __syncthreads(); + + if (warp == 0) { +#pragma unroll + for (int row = 0; row < Rows; ++row) { + float up = lane < kWarps ? up_warp_sums[row][lane] : 0.0f; + float gate = lane < kWarps ? gate_warp_sums[row][lane] : 0.0f; + up = warp_sum(up); + gate = warp_sum(gate); + if (lane == 0) { + // Match the public composition's BF16 rounding boundaries: + // BF16 linear outputs, BF16 SiLU output, then BF16 multiply. + up = round_bf16(up); + gate = round_bf16(gate); + const float silu = round_bf16(gate / (1.0f + expf(-gate))); + output[static_cast(row) * n + output_index] = + __float2bfloat16_rn(up * silu); + } + } + } +} + +template +void launch_bf16_swiglu( + const __nv_bfloat16* x, + const __nv_bfloat16* up_weight, + const __nv_bfloat16* gate_weight, + __nv_bfloat16* output, + int n, + int k, + cudaStream_t stream +) { + bf16_swiglu_kernel<<>>( + x, up_weight, gate_weight, output, n, k + ); +} + +template +__global__ void bf16_swiglu_warp_rows_kernel( + const __nv_bfloat16* __restrict__ x, + const __nv_bfloat16* __restrict__ up_weight, + const __nv_bfloat16* __restrict__ gate_weight, + __nv_bfloat16* __restrict__ output, + int n, + int k +) { + const int output_index = blockIdx.x; + const int row = threadIdx.x / kWarpSize; + const int lane = threadIdx.x & (kWarpSize - 1); + const int vector_count = k / 8; + + float up_sum = 0.0f; + float gate_sum = 0.0f; + const auto* x4 = reinterpret_cast( + x + static_cast(row) * k + ); + const auto* up4 = reinterpret_cast( + up_weight + static_cast(output_index) * k + ); + const auto* gate4 = reinterpret_cast( + gate_weight + static_cast(output_index) * k + ); + + // A warp owns one decode row. Same-address weight reads from sibling + // warps are served through the read-only/L1 path, while each row avoids + // CTA-wide shared-memory reductions and synchronization. + for (int vector_index = lane; + vector_index < vector_count; + vector_index += kWarpSize) { + const uint4 x_raw = x4[vector_index]; + const uint4 up_raw = up4[vector_index]; + const uint4 gate_raw = gate4[vector_index]; + const auto* x_values = reinterpret_cast(&x_raw); + const auto* up_values = + reinterpret_cast(&up_raw); + const auto* gate_values = + reinterpret_cast(&gate_raw); +#pragma unroll + for (int pair = 0; pair < 4; ++pair) { + const float2 xv = __bfloat1622float2(x_values[pair]); + const float2 uv = __bfloat1622float2(up_values[pair]); + const float2 gv = __bfloat1622float2(gate_values[pair]); + up_sum = fmaf(xv.x, uv.x, up_sum); + up_sum = fmaf(xv.y, uv.y, up_sum); + gate_sum = fmaf(xv.x, gv.x, gate_sum); + gate_sum = fmaf(xv.y, gv.y, gate_sum); + } + } + up_sum = warp_sum(up_sum); + gate_sum = warp_sum(gate_sum); + if (lane == 0) { + up_sum = round_bf16(up_sum); + gate_sum = round_bf16(gate_sum); + const float silu = + round_bf16(gate_sum / (1.0f + expf(-gate_sum))); + output[static_cast(row) * n + output_index] = + __float2bfloat16_rn(up_sum * silu); + } +} + +template +void launch_bf16_swiglu_warp_rows( + const __nv_bfloat16* x, + const __nv_bfloat16* up_weight, + const __nv_bfloat16* gate_weight, + __nv_bfloat16* output, + int n, + int k, + cudaStream_t stream +) { + bf16_swiglu_warp_rows_kernel<<>>( + x, up_weight, gate_weight, output, n, k + ); +} + +torch::Tensor bf16_swiglu( + torch::Tensor x, + torch::Tensor up_weight, + torch::Tensor gate_weight +) { + TORCH_CHECK( + x.is_cuda() && up_weight.is_cuda() && gate_weight.is_cuda(), + "x, up_weight, and gate_weight must be CUDA tensors" + ); + TORCH_CHECK( + x.device() == up_weight.device() && x.device() == gate_weight.device(), + "x and weights must share a device" + ); + TORCH_CHECK( + x.scalar_type() == torch::kBFloat16 && + up_weight.scalar_type() == torch::kBFloat16 && + gate_weight.scalar_type() == torch::kBFloat16, + "x and weights must be bf16" + ); + TORCH_CHECK( + x.dim() == 1 || x.dim() == 2, + "x must have shape [K] or [M, K]" + ); + TORCH_CHECK( + up_weight.dim() == 2 && gate_weight.dim() == 2, + "weights must have shape [N, K]" + ); + TORCH_CHECK( + x.is_contiguous() && up_weight.is_contiguous() && + gate_weight.is_contiguous(), + "x and weights must be contiguous" + ); + TORCH_CHECK( + !x.requires_grad() && !up_weight.requires_grad() && + !gate_weight.requires_grad(), + "bf16_swiglu is inference-only and does not support autograd" + ); + + const int64_t m = x.dim() == 1 ? 1 : x.size(0); + const int64_t k = x.size(-1); + const int64_t n = up_weight.size(0); + TORCH_CHECK(m >= 1 && m <= 8, "M must be in [1, 8]"); + TORCH_CHECK( + gate_weight.sizes() == up_weight.sizes(), + "up_weight and gate_weight must have identical shapes" + ); + TORCH_CHECK(up_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 % 8 == 0, "K must be divisible by 8"); + TORCH_CHECK( + k <= std::numeric_limits::max() && + n <= std::numeric_limits::max(), + "N or K exceeds the CUDA launcher limit" + ); + + const at::cuda::OptionalCUDAGuard guard(x.device()); + const auto* properties = at::cuda::getDeviceProperties(x.device().index()); + TORCH_CHECK( + properties->major >= 8, + "bf16_swiglu requires compute capability 8.0+" + ); + auto stream = at::cuda::getCurrentCUDAStream(); + auto output = x.dim() == 1 ? torch::empty({n}, x.options()) + : torch::empty({m, n}, x.options()); + + const auto* x_ptr = + reinterpret_cast(x.data_ptr()); + const auto* up_ptr = + reinterpret_cast(up_weight.data_ptr()); + const auto* gate_ptr = + reinterpret_cast(gate_weight.data_ptr()); + auto* output_ptr = + reinterpret_cast<__nv_bfloat16*>(output.data_ptr()); + const int n_int = static_cast(n); + const int k_int = static_cast(k); + const bool use_warp_rows = n_int == 6912 && k_int == 1536; + + switch (m) { + case 1: + launch_bf16_swiglu<1>( + x_ptr, up_ptr, gate_ptr, output_ptr, n_int, k_int, stream.stream() + ); + break; + case 2: + if (use_warp_rows) { + launch_bf16_swiglu_warp_rows<2>( + x_ptr, up_ptr, gate_ptr, output_ptr, n_int, k_int, stream.stream() + ); + } else { + launch_bf16_swiglu<2>( + x_ptr, up_ptr, gate_ptr, output_ptr, n_int, k_int, stream.stream() + ); + } + break; + case 3: + launch_bf16_swiglu<3>( + x_ptr, up_ptr, gate_ptr, output_ptr, n_int, k_int, stream.stream() + ); + break; + case 4: + if (use_warp_rows) { + launch_bf16_swiglu_warp_rows<4>( + x_ptr, up_ptr, gate_ptr, output_ptr, n_int, k_int, stream.stream() + ); + } else { + launch_bf16_swiglu<4>( + x_ptr, up_ptr, gate_ptr, output_ptr, n_int, k_int, stream.stream() + ); + } + break; + case 5: + launch_bf16_swiglu<5>( + x_ptr, up_ptr, gate_ptr, output_ptr, n_int, k_int, stream.stream() + ); + break; + case 6: + launch_bf16_swiglu<6>( + x_ptr, up_ptr, gate_ptr, output_ptr, n_int, k_int, stream.stream() + ); + break; + case 7: + launch_bf16_swiglu<7>( + x_ptr, up_ptr, gate_ptr, output_ptr, n_int, k_int, stream.stream() + ); + break; + case 8: + if (use_warp_rows) { + launch_bf16_swiglu_warp_rows<8>( + x_ptr, up_ptr, gate_ptr, output_ptr, n_int, k_int, stream.stream() + ); + } else { + launch_bf16_swiglu<8>( + x_ptr, up_ptr, gate_ptr, output_ptr, n_int, k_int, stream.stream() + ); + } + break; + } + C10_CUDA_CHECK(cudaGetLastError()); + return output; +} + +} // namespace + +PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) { + module.def( + "bf16_swiglu", + &bf16_swiglu, + py::arg("x"), + py::arg("up_weight"), + py::arg("gate_weight"), + "M in [1, 8] fused BF16 up/gate projection and SwiGLU" + ); +} diff --git a/docs/developer/cuda_kernels.md b/docs/developer/cuda_kernels.md index b2dbd20..fa13848 100644 --- a/docs/developer/cuda_kernels.md +++ b/docs/developer/cuda_kernels.md @@ -1,9 +1,9 @@ # CUDA Kernels AstrAI includes optional custom CUDA kernels for attention, rotary embedding, -BF16 GEMV, and FP8 GEMM. These are built when `nvcc` is available and CUDA is -detected. BF16 GEMV is directly callable and can be selected by the guarded -model linear dispatcher described below. +BF16 GEMV/SwiGLU, and FP8 GEMM. These are built when `nvcc` is available and +CUDA is detected. BF16 GEMV and SwiGLU are directly callable and can be +selected by guarded model dispatchers described below. ## Overview @@ -15,35 +15,96 @@ model linear dispatcher described below. | `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) | | `bf16_gemv` | `gemv/bf16_gemv.cu` | M=1..8 BF16 linear with FP32 accumulation (sm_80+) | +| `bf16_swiglu` | `gemv/bf16_swiglu.cu` | Fused M=1..8 BF16 up/gate projections and SwiGLU epilogue (sm_80+) | | `fp8_ops` | `fp8/ops.cu` | FP8 quantization + tensor-core GEMM (sm_89+) | ### BF16 GEMV primitive `astrai.extension.bf16_gemv(x, weight, bias=None)` accepts a contiguous BF16 input shaped `[K]` or `[M, K]`, with `M` in `[1, 8]` and any positive `K`, and -row-major weights `[N, K]`. One CTA reduces each output row and computes all M -results together, reusing the weight row across tokens. The weight stream uses -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 -Graph capture-safe, and requires sm_80 or newer. +row-major weights `[N, K]`. The general path assigns one 256-thread CTA to an +output row and computes all M results together, reusing the weight row across +tokens. For measured aligned M=4 medium projections, a 128-thread CTA instead +assigns one output to each of four warps. That removes the CTA-wide reduction +barrier and exposes four neighboring outputs without changing accumulation. + +The weight stream uses 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. The warp-tiled path is used +only when both tensors and every row are 16-byte aligned; all other calls keep +the general arbitrary-K path. Accumulation is FP32; the optional BF16 bias is +fused before the BF16 store. The launcher uses the current CUDA stream, is +CUDA Graph capture-safe, and requires sm_80 or newer. Model `Linear` calls route through the lightweight linear backend. Set `ASTRAI_GEMV=0` for an unconditional `F.linear` fallback, `1` to force the 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. -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: +Measured SM89 small-M bands are enabled as follows: -| M | Automatic `(N, K)` bands | Engine throughput | +| M | Automatic `(N, K)` bands | Validated gain | |---:|---|---:| -| 2 | `(256,1536)`, `(1536,1536)`, `(100000,1536)` | +14.0% | -| 4 | `(256,1536)`, `(1536,1536)` | +11.8% | +| 1 | OPT-1.3B Q/K/V/O and MLP | +4.54% OPT projection chain | +| 2 | AstrAI `(256,1536)`, `(1536,1536)`, `(100000,1536)` plus all common shapes below | +14.0% on AstrAI 1B; +5.66% to +25.20% common chains | +| 4 | AstrAI `(256,1536)`, `(1536,1536)` plus gated common shapes below | +11.8% on AstrAI 1B; +5.67% to +7.71% common chains | +| 8 | none | at least one projection in every measured family missed the per-shape gate | -These A→B→B→A results use the real `InferenceEngine`, including scheduler, +The common set covers LLaMA 2 7B Q/O, gate/up, and down; LLaMA 3 8B K/V, +gate/up, and down; LLaMA 2 13B Q/K/V/O, gate/up, and down; and GPT-NeoX MLP +up/down. In `(N,K)` form it is `(1024,4096)`, `(4096,4096)`, +`(11008,4096)`, `(4096,11008)`, `(14336,4096)`, `(4096,14336)`, +`(5120,5120)`, `(13824,5120)`, `(5120,13824)`, `(16384,4096)`, and +`(4096,16384)`. M=2 enables all eleven. M=4 excludes the three LLaMA 2 7B +bands `(4096,4096)`, `(11008,4096)`, and `(4096,11008)` because their combined +projection chain reached only +1.89%, below the 3% automatic-dispatch gate. + +The extended common set adds Qwen2-7B `(512,3584)`, `(3584,3584)`, +`(18944,3584)`, and `(3584,18944)`; LLaMA 3 70B `(1024,8192)`, +`(8192,8192)`, `(28672,8192)`, and `(8192,28672)`; and OPT-1.3B +`(2048,2048)`, `(8192,2048)`, and `(2048,8192)`. Qwen2 and LLaMA 3 70B are +enabled at M=2/4. OPT-1.3B is enabled at M=1/2. Other rows retain their +previous policy or fall back to PyTorch. + +Inside the primitive, a templated cooperative kernel uses either 256 threads +or a shape-gated 128-thread CTA. The smaller CTA is enabled only where an +interleaved direct-module comparison against the original 256-thread kernel +cleared 5%: OPT up at M=1; selected LLaMA 2 7B, Qwen2, and OPT projections at +M=2; LLaMA 2 13B Q/O, Qwen2 Q/O, and selected OPT projections at M=4; and +selected LLaMA 2, Qwen2, LLaMA 3 KV, and OPT projections at M=8. Confirmed +direct-kernel gains range from +5.37% to +48.54%. Long-K and saturated shapes +keep the 256-thread fallback. This internal selector is separate from model +automatic dispatch, whose Python/wrapper overhead is included in the gates +above. + +On NVIDIA L20 (SM89), the common-shape microbenchmark reports +5.37% to ++114.39% for M=2 and +5.38% to +115.26% for M=4 versus `F.linear`. The paired +main-versus-warp-tiling run used identical interleaved settings; for the four +M=4 selected shapes, candidate latency changed from 0.016292 to 0.016108 ms +for `(4096,4096)`, 0.037939 to 0.028539 ms for `(11008,4096)`, 0.043407 to +0.039803 ms for `(4096,11008)`, and 0.006697 to 0.006390 ms for +`(1024,4096)`. + +The dependent projection-chain gate, which includes Python dispatch and +rotates through distinct weights instead of repeatedly warming one matrix, +measured: + +| Synthetic chain | M=2 | M=4 | Row argmax parity | +|---|---:|---:|---| +| LLaMA 2 7B | +8.49% | fallback (M=4 bands excluded) | exact | +| LLaMA 3 8B | +8.50% | +6.44% | exact | +| LLaMA 2 13B | +5.66% | +5.67% | exact | +| GPT-NeoX 20B | +6.95% | +5.93% | exact | +| Qwen2 7B | +7.48% | +7.48% | exact | +| LLaMA 3 70B | +7.77% | +7.69% | exact | +| OPT 1.3B | +25.20% | fallback (M=4 up projection regresses) | exact | + +OPT 1.3B M=1 is +4.54%. Qwen2 and LLaMA 3 70B M=1, and all three new +families at M=8, remain exact PyTorch fallbacks. + +These are synthetic projection-chain measurements, not whole-model throughput +claims. Reproduce them with `scripts/tools/benchmark_gemv_common.py`. + +The AstrAI 1B A→B→B→A results use the real `InferenceEngine`, including scheduler, sampling, and CUDA Graph. M=8 stays on PyTorch because its remaining greedy-stable winners missed the 3% end-to-end gate. Long-K MLP-down bands are also excluded because their valid BF16 error changed a checkpoint greedy @@ -55,6 +116,30 @@ remain on PyTorch. Use mode `1` only for explicit A/B runs outside this table. The primitive remains directly callable and deliberately has no internal `F.linear` fallback. The model-level backend owns fallback and dispatch policy. +### BF16 SwiGLU primitive + +`astrai.extension.bf16_swiglu(x, up_weight, gate_weight)` fuses the two dense +MLP projections with `up * silu(gate)` into one CUDA launch for contiguous +BF16 inputs with `M` in `[1, 8]` and K divisible by 8. It preserves the BF16 +rounding boundaries of the two projection outputs, SiLU output, and final +product while accumulating dot products in FP32. + +The kernel contains two output-row tilings. A CTA-reuse path reads each up/gate +weight chunk once and applies it to all M rows. The native AstrAI 1B shape +`(N,K)=(6912,1536)` uses one warp per decode row for M=2/4/8; on L20 this +removes the shared reductions and barrier and reduces M=4 CUDA-Graph latency +from 0.0324 ms to 0.0181 ms. Wider LLaMA/GPT-NeoX matrices keep CTA reuse, +because duplicating their weight reads across row warps regressed 1.3-4.2%. + +Dense `MLP` modules route through the SwiGLU backend. `ASTRAI_SWIGLU=0` keeps +the unfused linear backend, and `1` explicitly forces the fused primitive. +`auto` is the default but currently has no enabled bands: although direct +errors are small (maximum absolute error at most 2.4e-4 in the L20 matrix), +the different FP32 reduction order changed greedy checkpoint output for +M=1/2/4. Automatic dispatch therefore remains numerically identical to the +existing path. See [the benchmark protocol](./swiglu_benchmark.md) for raw +operator, engine, and checkpoint evidence. + Additionally, optimized `.cuh` variants with tensor-core MMA (Matrix Multiply-Accumulate) exist: | Variant | File | Optimization | @@ -245,10 +330,12 @@ astrai/extension/ │ ├── attention.py # Stateless attention kernel wrappers │ ├── rotary.py # Stateless rotary kernel wrapper │ ├── gemv.py # Stateless BF16 GEMV primitive +│ ├── swiglu.py # Stateless fused BF16 SwiGLU primitive │ └── fp8.py # Stateless FP8 primitives (custom_op) ├── fp8.py # FP8 strategy layer (fp8_autocast, recipes) └── backend/ ├── attention.py # Backend selection, KV cache I/O, and fallback + ├── swiglu.py # Inference-only fused/unfused SwiGLU policy └── rotary.py # Per-call CUDA/torch rotary dispatch ``` diff --git a/docs/developer/decode_linear_benchmark.md b/docs/developer/decode_linear_benchmark.md index 02bf113..877efc4 100644 --- a/docs/developer/decode_linear_benchmark.md +++ b/docs/developer/decode_linear_benchmark.md @@ -20,3 +20,22 @@ change the decode batch sizes. Compare each GPU architecture only with its own baseline; do not use absolute A100-versus-L20 numbers as a dispatch criterion. Keep the raw JSON as the source of truth and generate tables with `--markdown-output` rather than transcribing measurements by hand. + +For direct A/B coverage of the custom kernel and guarded dispatcher across +traditional LLaMA and GPT-NeoX decode shapes, use: + +```bash +CUDA_VISIBLE_DEVICES=0 PYTHONPATH=. python scripts/tools/benchmark_gemv_common.py \ + --suite all --family traditional --m 2 4 \ + --output results/gemv_common.json +``` + +The kernel suite compares the directly callable primitive with `F.linear`. +Use repeatable `--shape-label` and `--chain-label` filters for a focused run. +The synthetic-chain suite alternates `ASTRAI_GEMV=0` and `auto`, includes +dependent MLP work and Python dispatch, and rotates through distinct weights. +Pass `--candidate-mode 1` to characterize a family before adding it to the +automatic shape table; the checked-in final evidence always uses `auto`. +It is deliberately not labeled a whole-model throughput benchmark. Both +suites report median/p90 CUDA-event latency plus maximum absolute error, +relative L2 error, and row-wise argmax parity. diff --git a/docs/developer/swiglu_benchmark.md b/docs/developer/swiglu_benchmark.md new file mode 100644 index 0000000..54b37df --- /dev/null +++ b/docs/developer/swiglu_benchmark.md @@ -0,0 +1,58 @@ +# Fused SwiGLU benchmark + +`scripts/tools/benchmark_swiglu.py` compares the directly callable fused BF16 +SwiGLU primitive with both `F.linear` and the existing two-GEMV chain. It covers +the native AstrAI 1B MLP plus LLaMA 2 7B/13B, LLaMA 3 8B, and GPT-NeoX 20B +up/gate shapes at M=1/2/4/8 in eager and CUDA Graph modes. + +```bash +CUDA_VISIBLE_DEVICES=0 python scripts/tools/benchmark_swiglu.py \ + --output results/swiglu.json \ + --markdown-output results/swiglu.md \ + --m-values 1,2,4,8 --mode both \ + --warmup 20 --iterations 100 --trials 10 +``` + +Each trial uses A-B-C-C-B-A ordering to balance clock, cache, and temperature +drift. The generated JSON records every timing sample, p50/p90/p99, CUDA launch +count, maximum/mean absolute error, and cosine similarity. + +## L20 findings + +Hardware was one NVIDIA L20 (sm_89), PyTorch 2.11.0+cu128, CUDA 12.8. The +existing GPU5 inference service remained resident (15.4 GiB) but idle at the +sampling boundaries; no process or container was stopped. + +For AstrAI 1B `(N,K)=(6912,1536)`, CUDA Graph medians were: + +| M | torch (ms) | GEMV chain (ms) | fused (ms) | vs best unfused | +|---:|---:|---:|---:|---:| +| 1 | 0.02564 | 0.02298 | 0.01375 | +67.13% | +| 2 | 0.02484 | 0.02628 | 0.01416 | +75.40% | +| 4 | 0.02507 | 0.03839 | 0.01806 | +38.82% | +| 8 | 0.02563 | 0.07007 | 0.03339 | -23.24% | + +The wide traditional shapes are weight-bandwidth dominated. CTA reuse keeps +the fused primitive within roughly -1.2% to +0.9% of the best unfused chain, +so none is eligible for automatic selection. This negative crossover is kept +in the raw evidence rather than hidden by a favorable subset. + +The real 24-layer AstrAI checkpoint was then run through `InferenceEngine`, +including scheduler, sampling, and CUDA Graph. A-B-B-A medians were: + +| Batch | unfused (ms/step) | forced fused (ms/step) | throughput gain | +|---:|---:|---:|---:| +| 1 | 4.125 | 3.925 | +5.10% | +| 2 | 4.245 | 4.055 | +4.69% | +| 4 | 4.475 | 4.305 | +3.95% | + +## Dispatch decision + +Direct correctness stayed close (`max_abs <= 2.4e-4`, cosine approximately +1.0), but deterministic greedy generations changed at M=1, M=2, and M=4. + +For that reason no SM89 shape is enabled in `auto`. The default path stays on +the existing unfused linear backend, including any independently qualified +GEMV dispatch. `ASTRAI_SWIGLU=1` remains an explicit benchmark/experimentation +switch for callers that accept normal BF16 reduction-order variation. A future +automatic band must repeat both the performance and checkpoint-output gates. diff --git a/scripts/tools/benchmark_gemv_common.py b/scripts/tools/benchmark_gemv_common.py new file mode 100644 index 0000000..33b8f0c --- /dev/null +++ b/scripts/tools/benchmark_gemv_common.py @@ -0,0 +1,450 @@ +"""Benchmark the BF16 GEMV primitive and guarded linear dispatcher. + +The kernel suite covers AstrAI's native projections plus common LLaMA and +GPT-NeoX matrix shapes. The chain suite is a synthetic projection/MLP chain; +it measures dispatcher overhead and dependent MLP work, but is deliberately +not presented as a whole-model throughput benchmark. +""" + +import argparse +import gc +import json +import math +import os +import statistics +from collections.abc import Callable +from dataclasses import dataclass +from pathlib import Path + +import torch +import torch.nn.functional as F + +from astrai.extension import bf16_gemv, is_available, linear + + +@dataclass(frozen=True) +class Shape: + label: str + n: int + k: int + + +@dataclass(frozen=True) +class Chain: + label: str + hidden: int + kv: int + intermediate: int + fused_qkv: bool = False + gated_mlp: bool = True + + +@dataclass(frozen=True) +class Timing: + median_ms: float + p90_ms: float + + +ASTRAI_SHAPES = ( + Shape("astrai_qkv", 256, 1536), + Shape("astrai_square", 1536, 1536), + Shape("astrai_up_gate", 6912, 1536), + Shape("astrai_down", 1536, 6912), + Shape("astrai_lm_head", 100000, 1536), +) + +TRADITIONAL_SHAPES = ( + Shape("llama2_7b_qo", 4096, 4096), + Shape("llama2_7b_up_gate", 11008, 4096), + Shape("llama2_7b_down", 4096, 11008), + Shape("llama3_8b_kv", 1024, 4096), + Shape("llama3_8b_up_gate", 14336, 4096), + Shape("llama3_8b_down", 4096, 14336), + Shape("llama2_13b_qo", 5120, 5120), + Shape("llama2_13b_up_gate", 13824, 5120), + Shape("llama2_13b_down", 5120, 13824), + Shape("gpt_neox_up", 16384, 4096), + Shape("gpt_neox_down", 4096, 16384), + Shape("qwen2_7b_kv", 512, 3584), + Shape("qwen2_7b_qo", 3584, 3584), + Shape("qwen2_7b_up_gate", 18944, 3584), + Shape("qwen2_7b_down", 3584, 18944), + Shape("llama3_70b_kv", 1024, 8192), + Shape("llama3_70b_qo", 8192, 8192), + Shape("llama3_70b_up_gate", 28672, 8192), + Shape("llama3_70b_down", 8192, 28672), + Shape("opt_1_3b_qkvo", 2048, 2048), + Shape("opt_1_3b_up", 8192, 2048), + Shape("opt_1_3b_down", 2048, 8192), +) + +CHAINS = ( + Chain("llama2_7b", 4096, 4096, 11008), + Chain("llama3_8b", 4096, 1024, 14336), + Chain("llama2_13b", 5120, 5120, 13824), + Chain("gpt_neox_20b", 4096, 4096, 16384, fused_qkv=True), + Chain("qwen2_7b", 3584, 512, 18944), + Chain("llama3_70b", 8192, 1024, 28672), + Chain("opt_1_3b", 2048, 2048, 8192, gated_mlp=False), +) + + +def _elapsed_ms(fn: Callable[[], torch.Tensor], inner: int) -> float: + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + for _ in range(inner): + fn() + end.record() + end.synchronize() + return start.elapsed_time(end) / inner + + +def _timing(values: list[float]) -> Timing: + ordered = sorted(values) + p90_index = max(0, math.ceil(0.9 * len(ordered)) - 1) + return Timing(statistics.median(ordered), ordered[p90_index]) + + +def _measure_pair( + baseline: Callable[[], torch.Tensor], + candidate: Callable[[], torch.Tensor], + *, + warmup: int, + samples: int, + inner: int, + prepare_baseline: Callable[[], None] = lambda: None, + prepare_candidate: Callable[[], None] = lambda: None, +) -> tuple[Timing, Timing]: + cases = ( + ("baseline", prepare_baseline, baseline), + ("candidate", prepare_candidate, candidate), + ) + for iteration in range(warmup): + _, prepare, fn = cases[iteration % 2] + prepare() + fn() + torch.cuda.synchronize() + + values: dict[str, list[float]] = {"baseline": [], "candidate": []} + for sample in range(samples): + order = cases if sample % 2 == 0 else tuple(reversed(cases)) + for label, prepare, fn in order: + prepare() + values[label].append(_elapsed_ms(fn, inner)) + return _timing(values["baseline"]), _timing(values["candidate"]) + + +def _print_header() -> None: + print( + "suite,label,m,n,k,torch_median_ms,torch_p90_ms," + "candidate_median_ms,candidate_p90_ms,speedup_pct," + "max_abs,relative_l2,argmax_equal" + ) + + +def _print_result( + suite: str, + label: str, + m: int, + n: int, + k: int, + baseline: Timing, + candidate: Timing, + reference: torch.Tensor, + actual: torch.Tensor, +) -> dict[str, object]: + difference = actual.float() - reference.float() + max_abs = difference.abs().max().item() + relative_l2 = difference.norm().item() / max(reference.float().norm().item(), 1e-12) + argmax_equal = torch.equal(actual.argmax(dim=-1), reference.argmax(dim=-1)) + speedup = (baseline.median_ms / candidate.median_ms - 1.0) * 100.0 + result: dict[str, object] = { + "suite": suite, + "label": label, + "m": m, + "n": n, + "k": k, + "torch_median_ms": baseline.median_ms, + "torch_p90_ms": baseline.p90_ms, + "candidate_median_ms": candidate.median_ms, + "candidate_p90_ms": candidate.p90_ms, + "speedup_pct": speedup, + "max_abs": max_abs, + "relative_l2": relative_l2, + "argmax_equal": argmax_equal, + } + print( + f"{suite},{label},{m},{n},{k}," + f"{baseline.median_ms:.6f},{baseline.p90_ms:.6f}," + f"{candidate.median_ms:.6f},{candidate.p90_ms:.6f}," + f"{speedup:+.2f},{max_abs:.6f},{relative_l2:.8f}," + f"{str(argmax_equal).lower()}", + flush=True, + ) + return result + + +def _weight(n: int, k: int, device: torch.device, std: float) -> torch.Tensor: + weight = torch.empty((n, k), device=device, dtype=torch.bfloat16) + weight.normal_(mean=0.0, std=std) + return weight.requires_grad_(True) + + +def _kernel_functions( + x: torch.Tensor, weight: torch.Tensor +) -> tuple[Callable[[], torch.Tensor], Callable[[], torch.Tensor]]: + def baseline() -> torch.Tensor: + return F.linear(x, weight) + + def candidate() -> torch.Tensor: + return bf16_gemv(x, weight.detach()) + + return baseline, candidate + + +def benchmark_kernels( + args: argparse.Namespace, device: torch.device +) -> list[dict[str, object]]: + if args.family == "astrai": + shapes = ASTRAI_SHAPES + elif args.family == "traditional": + shapes = TRADITIONAL_SHAPES + else: + shapes = ASTRAI_SHAPES + TRADITIONAL_SHAPES + if args.shape_label: + requested = set(args.shape_label) + shapes = tuple(shape for shape in shapes if shape.label in requested) + missing = requested - {shape.label for shape in shapes} + if missing: + raise ValueError(f"unknown shape labels: {', '.join(sorted(missing))}") + + results: list[dict[str, object]] = [] + for shape in shapes: + weight = _weight(shape.n, shape.k, device, args.weight_std) + for m in args.m: + x = torch.randn((m, shape.k), device=device, dtype=torch.bfloat16) + baseline_fn, candidate_fn = _kernel_functions(x, weight) + with torch.inference_mode(): + reference = baseline_fn() + actual = candidate_fn() + baseline, candidate = _measure_pair( + baseline_fn, + candidate_fn, + warmup=args.warmup, + samples=args.samples, + inner=args.inner, + ) + results.append( + _print_result( + "kernel", + shape.label, + m, + shape.n, + shape.k, + baseline, + candidate, + reference, + actual, + ) + ) + del baseline_fn, candidate_fn, x, reference, actual + del weight + gc.collect() + torch.cuda.empty_cache() + return results + + +def _set_mode(mode: str) -> None: + os.environ["ASTRAI_GEMV"] = mode + + +def _chain_weights( + spec: Chain, device: torch.device, std: float +) -> dict[str, torch.Tensor]: + weights = { + "o": _weight(spec.hidden, spec.hidden, device, std), + "up": _weight(spec.intermediate, spec.hidden, device, std), + "down": _weight(spec.hidden, spec.intermediate, device, std), + } + if spec.fused_qkv: + weights["qkv"] = _weight(3 * spec.hidden, spec.hidden, device, std) + else: + weights.update( + { + "q": _weight(spec.hidden, spec.hidden, device, std), + "k": _weight(spec.kv, spec.hidden, device, std), + "v": _weight(spec.kv, spec.hidden, device, std), + } + ) + if spec.gated_mlp: + weights["gate"] = _weight(spec.intermediate, spec.hidden, device, std) + return weights + + +def _chain_fn( + x: torch.Tensor, weights: dict[str, torch.Tensor], spec: Chain +) -> Callable[[], torch.Tensor]: + def run() -> torch.Tensor: + output_projection = linear(x, weights["o"]) + up = linear(x, weights["up"]) + if spec.fused_qkv: + attention_projection = linear(x, weights["qkv"])[..., : x.shape[-1]] + hidden = F.gelu(up) + else: + attention_projection = linear(x, weights["q"]) + linear(x, weights["k"]) + linear(x, weights["v"]) + if spec.gated_mlp: + gate = linear(x, weights["gate"]) + hidden = F.silu(gate) * up + else: + hidden = F.gelu(up) + down = linear(hidden, weights["down"]) + return attention_projection + output_projection + down + + return run + + +def benchmark_chains( + args: argparse.Namespace, device: torch.device +) -> list[dict[str, object]]: + results: list[dict[str, object]] = [] + chains = CHAINS + if args.chain_label: + requested = set(args.chain_label) + chains = tuple(chain for chain in chains if chain.label in requested) + missing = requested - {chain.label for chain in chains} + if missing: + raise ValueError(f"unknown chain labels: {', '.join(sorted(missing))}") + for spec in chains: + weights = _chain_weights(spec, device, args.weight_std) + for m in args.m: + x = torch.randn((m, spec.hidden), device=device, dtype=torch.bfloat16) + run = _chain_fn(x, weights, spec) + with torch.inference_mode(): + _set_mode("0") + reference = run() + _set_mode(args.candidate_mode) + actual = run() + baseline, candidate = _measure_pair( + run, + run, + warmup=args.warmup, + samples=args.samples, + inner=args.chain_inner, + prepare_baseline=lambda: _set_mode("0"), + prepare_candidate=lambda: _set_mode(args.candidate_mode), + ) + results.append( + _print_result( + "synthetic_chain", + spec.label, + m, + spec.hidden, + spec.intermediate, + baseline, + candidate, + reference, + actual, + ) + ) + del x, reference, actual + del weights + gc.collect() + torch.cuda.empty_cache() + return results + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--suite", choices=("kernel", "chain", "all"), default="all") + parser.add_argument( + "--family", choices=("astrai", "traditional", "all"), default="all" + ) + parser.add_argument( + "--m", type=int, nargs="+", choices=(1, 2, 4, 8), default=(1, 2, 4, 8) + ) + parser.add_argument( + "--shape-label", + action="append", + help="limit the kernel suite to one or more named shape labels", + ) + parser.add_argument( + "--chain-label", + action="append", + help="limit the chain suite to one or more named model families", + ) + parser.add_argument("--device", type=int, default=0) + parser.add_argument("--warmup", type=int, default=20) + parser.add_argument("--samples", type=int, default=9) + parser.add_argument("--inner", type=int, default=100) + parser.add_argument("--chain-inner", type=int, default=20) + parser.add_argument( + "--candidate-mode", + choices=("auto", "1"), + default="auto", + help="dispatcher mode for the candidate side of the chain suite", + ) + parser.add_argument("--weight-std", type=float, default=0.02) + parser.add_argument("--seed", type=int, default=20260902) + parser.add_argument( + "--output", + type=Path, + help="optional JSON output; stdout always retains the compact CSV table", + ) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + if not torch.cuda.is_available() or not is_available("bf16_gemv"): + raise RuntimeError("benchmark requires CUDA and the built bf16_gemv extension") + if args.warmup < 0 or args.samples < 1 or args.inner < 1 or args.chain_inner < 1: + raise ValueError("warmup must be non-negative and sample/inner counts positive") + + torch.cuda.set_device(args.device) + device = torch.device("cuda", args.device) + torch.manual_seed(args.seed) + torch.cuda.manual_seed_all(args.seed) + properties = torch.cuda.get_device_properties(device) + print( + f"# device={properties.name}, capability={properties.major}.{properties.minor}, " + f"seed={args.seed}, weight_std={args.weight_std}" + ) + _print_header() + results: list[dict[str, object]] = [] + if args.suite in ("kernel", "all"): + results.extend(benchmark_kernels(args, device)) + if args.suite in ("chain", "all"): + results.extend(benchmark_chains(args, device)) + if args.output is not None: + payload = { + "environment": { + "device": properties.name, + "capability": f"{properties.major}.{properties.minor}", + "torch": torch.__version__, + "cuda": torch.version.cuda, + }, + "parameters": { + "suite": args.suite, + "family": args.family, + "m": args.m, + "shape_labels": args.shape_label, + "chain_labels": args.chain_label, + "candidate_mode": args.candidate_mode, + "seed": args.seed, + "weight_std": args.weight_std, + "warmup": args.warmup, + "samples": args.samples, + "inner": args.inner, + "chain_inner": args.chain_inner, + }, + "results": results, + } + args.output.parent.mkdir(parents=True, exist_ok=True) + args.output.write_text(json.dumps(payload, indent=2) + "\n") + + +if __name__ == "__main__": + main() diff --git a/scripts/tools/benchmark_swiglu.py b/scripts/tools/benchmark_swiglu.py new file mode 100644 index 0000000..812b9d9 --- /dev/null +++ b/scripts/tools/benchmark_swiglu.py @@ -0,0 +1,325 @@ +"""Benchmark fused BF16 SwiGLU against torch and unfused GEMV chains.""" + +from __future__ import annotations + +import json +import math +import statistics +from dataclasses import dataclass +from datetime import datetime, timezone +from pathlib import Path +from typing import Callable, Iterable + +import click +import torch +import torch.nn.functional as F + +from astrai.extension import bf16_gemv, bf16_swiglu, is_available + + +@dataclass(frozen=True) +class SwiGLUShape: + name: str + n: int + k: int + + +DEFAULT_SHAPES = ( + SwiGLUShape("astrai_1b", 6912, 1536), + SwiGLUShape("llama2_7b", 11008, 4096), + SwiGLUShape("llama3_8b", 14336, 4096), + SwiGLUShape("llama2_13b", 13824, 5120), + SwiGLUShape("gpt_neox_20b", 16384, 6144), +) + + +def parse_positive_ints(value: str) -> tuple[int, ...]: + try: + values = tuple(dict.fromkeys(int(item.strip()) for item in value.split(","))) + except ValueError as exc: + raise click.BadParameter("expected comma-separated integers") from exc + if not values or any(item <= 0 for item in values): + raise click.BadParameter("values must be positive integers") + return values + + +def parse_shape(value: str) -> SwiGLUShape: + parts = value.split(":") + if len(parts) != 3 or not parts[0]: + raise click.BadParameter("shape must use NAME:N:K") + try: + n, k = (int(item) for item in parts[1:]) + except ValueError as exc: + raise click.BadParameter("N and K must be integers") from exc + if n <= 0 or k <= 0 or k % 8: + raise click.BadParameter("N must be positive and K positive/divisible by 8") + return SwiGLUShape(parts[0], n, k) + + +def percentile(values: Iterable[float], quantile: float) -> float: + ordered = sorted(values) + rank = (len(ordered) - 1) * quantile + lower = math.floor(rank) + upper = math.ceil(rank) + if lower == upper: + return ordered[lower] + fraction = rank - lower + return ordered[lower] * (1 - fraction) + ordered[upper] * fraction + + +def summarize(values: list[float]) -> dict[str, float]: + return { + "median_ms": statistics.median(values), + "p90_ms": percentile(values, 0.90), + "p99_ms": percentile(values, 0.99), + "min_ms": min(values), + "max_ms": max(values), + } + + +def time_operation(operation: Callable[[], torch.Tensor], iterations: int) -> float: + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + for _ in range(iterations): + operation() + end.record() + end.synchronize() + return start.elapsed_time(end) / iterations + + +def count_cuda_kernels( + operation: Callable[[], torch.Tensor], repeats: int = 5 +) -> float: + with torch.profiler.profile( + activities=[ + torch.profiler.ProfilerActivity.CPU, + torch.profiler.ProfilerActivity.CUDA, + ], + acc_events=True, + ) as profile: + for _ in range(repeats): + operation() + torch.cuda.synchronize() + device_type = torch.autograd.DeviceType.CUDA + events = [event for event in profile.events() if event.device_type == device_type] + return len(events) / repeats + + +def capture(operation: Callable[[], torch.Tensor]): + for _ in range(3): + operation() + torch.cuda.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + output = operation() + + def replay() -> torch.Tensor: + graph.replay() + return output + + return replay + + +def make_operations(x, up_weight, gate_weight, mode: str): + operations: dict[str, Callable[[], torch.Tensor]] = { + "torch": lambda: F.linear(x, up_weight) * F.silu(F.linear(x, gate_weight)), + "gemv_chain": lambda: ( + bf16_gemv(x, up_weight) * F.silu(bf16_gemv(x, gate_weight)) + ), + "fused": lambda: bf16_swiglu(x, up_weight, gate_weight), + } + if mode == "graph": + operations = {name: capture(op) for name, op in operations.items()} + return operations + + +def benchmark_case( + shape: SwiGLUShape, + m: int, + mode: str, + *, + warmup: int, + iterations: int, + trials: int, +) -> list[dict[str, object]]: + x = torch.randn((m, shape.k), device="cuda", dtype=torch.bfloat16) * 0.1 + scale = shape.k**-0.5 + up_weight = ( + torch.randn((shape.n, shape.k), device="cuda", dtype=torch.bfloat16) * scale + ) + gate_weight = ( + torch.randn((shape.n, shape.k), device="cuda", dtype=torch.bfloat16) * scale + ) + operations = make_operations(x, up_weight, gate_weight, mode) + for operation in operations.values(): + for _ in range(warmup): + operation() + torch.cuda.synchronize() + + samples = {name: [] for name in operations} + forward_order = tuple(operations) + # A-B-C-C-B-A order balances cache, clock, and temperature drift. + for _ in range(trials): + for name in (*forward_order, *reversed(forward_order)): + samples[name].append(time_operation(operations[name], iterations)) + + with torch.no_grad(): + expected = operations["torch"]().clone() + actual = operations["fused"]().clone() + difference = (actual.float() - expected.float()).abs() + max_abs_error = float(difference.max()) + mean_abs_error = float(difference.mean()) + cosine_similarity = float( + F.cosine_similarity(actual.float().flatten(), expected.float().flatten(), dim=0) + ) + + results = [] + for name, operation in operations.items(): + result: dict[str, object] = { + "shape": shape.name, + "m": m, + "n": shape.n, + "k": shape.k, + "mode": mode, + "implementation": name, + "cuda_kernel_launches_per_call": count_cuda_kernels(operation), + **summarize(samples[name]), + } + if name == "fused": + result.update( + max_abs_error=max_abs_error, + mean_abs_error=mean_abs_error, + cosine_similarity=cosine_similarity, + ) + results.append(result) + return results + + +def device_metadata() -> dict[str, object]: + props = torch.cuda.get_device_properties(0) + return { + "timestamp_utc": datetime.now(timezone.utc).isoformat(), + "gpu_name": props.name, + "compute_capability": f"{props.major}.{props.minor}", + "total_memory_bytes": props.total_memory, + "torch_version": torch.__version__, + "cuda_version": torch.version.cuda, + "dtype": "bfloat16", + } + + +def render_markdown(payload: dict[str, object]) -> str: + metadata = payload["metadata"] + results = payload["results"] + assert isinstance(metadata, dict) + assert isinstance(results, list) + by_case = { + (item["shape"], item["m"], item["mode"], item["implementation"]): item + for item in results + } + cases = sorted({(item["shape"], item["m"], item["mode"]) for item in results}) + lines = [ + "# Fused SwiGLU benchmark", + "", + f"- GPU: {metadata['gpu_name']}", + f"- Compute capability: {metadata['compute_capability']}", + f"- PyTorch / CUDA: {metadata['torch_version']} / {metadata['cuda_version']}", + "", + "| Shape | M | Mode | torch ms | GEMV chain ms | fused ms | " + "vs best unfused | fused kernels | max abs | cosine |", + "|---|---:|---|---:|---:|---:|---:|---:|---:|---:|", + ] + for shape, m, mode in cases: + torch_item = by_case[(shape, m, mode, "torch")] + gemv_item = by_case[(shape, m, mode, "gemv_chain")] + fused_item = by_case[(shape, m, mode, "fused")] + best = min(torch_item["median_ms"], gemv_item["median_ms"]) + improvement = (best / fused_item["median_ms"] - 1) * 100 + lines.append( + f"| {shape} | {m} | {mode} | {torch_item['median_ms']:.5f} | " + f"{gemv_item['median_ms']:.5f} | {fused_item['median_ms']:.5f} | " + f"{improvement:+.2f}% | " + f"{fused_item['cuda_kernel_launches_per_call']:.1f} | " + f"{fused_item['max_abs_error']:.5f} | " + f"{fused_item['cosine_similarity']:.8f} |" + ) + lines.append("") + return "\n".join(lines) + + +@click.command(help=__doc__) +@click.option("--output", type=click.Path(path_type=Path), required=True) +@click.option("--markdown-output", type=click.Path(path_type=Path)) +@click.option("--m-values", default="1,2,4,8", show_default=True) +@click.option("--shape", "shape_values", multiple=True, help="Repeat NAME:N:K.") +@click.option("--mode", type=click.Choice(("eager", "graph", "both")), default="both") +@click.option("--warmup", type=click.IntRange(min=1), default=10, show_default=True) +@click.option( + "--iterations", type=click.IntRange(min=1), default=100, show_default=True +) +@click.option("--trials", type=click.IntRange(min=1), default=10, show_default=True) +@click.option("--seed", type=int, default=0, show_default=True) +def benchmark_command( + output: Path, + markdown_output: Path | None, + m_values: str, + shape_values: tuple[str, ...], + mode: str, + warmup: int, + iterations: int, + trials: int, + seed: int, +) -> None: + if not torch.cuda.is_available(): + raise click.ClickException("CUDA is required") + if not is_available("bf16_gemv") or not is_available("bf16_swiglu"): + raise click.ClickException("built bf16_gemv and bf16_swiglu are required") + shapes = tuple(parse_shape(value) for value in shape_values) or DEFAULT_SHAPES + m_values_parsed = parse_positive_ints(m_values) + if any(m > 8 for m in m_values_parsed): + raise click.BadParameter("fused primitive supports M up to 8") + modes = ("eager", "graph") if mode == "both" else (mode,) + torch.manual_seed(seed) + torch.cuda.manual_seed_all(seed) + + results = [] + with torch.inference_mode(): + for shape in shapes: + for m in m_values_parsed: + for current_mode in modes: + click.echo( + f"{shape.name}: M={m} N={shape.n} K={shape.k} {current_mode}" + ) + results.extend( + benchmark_case( + shape, + m, + current_mode, + warmup=warmup, + iterations=iterations, + trials=trials, + ) + ) + torch.cuda.empty_cache() + + payload: dict[str, object] = { + "metadata": device_metadata(), + "settings": { + "warmup": warmup, + "iterations": iterations, + "trials": trials, + "seed": seed, + "order": "A-B-C-C-B-A", + }, + "results": results, + } + output.parent.mkdir(parents=True, exist_ok=True) + output.write_text(json.dumps(payload, indent=2) + "\n") + if markdown_output is not None: + markdown_output.parent.mkdir(parents=True, exist_ok=True) + markdown_output.write_text(render_markdown(payload)) + + +if __name__ == "__main__": + benchmark_command() diff --git a/setup.py b/setup.py index 0b4313f..baa38a0 100644 --- a/setup.py +++ b/setup.py @@ -122,6 +122,7 @@ class _CMakeBuildExt(_build_ext): "attn_paged_decode", "attn_paged_prefill", "bf16_gemv", + "bf16_swiglu", "rotary_emb", ) missing = [name for name in required if not any(lib_dir.glob(f"{name}.*.so"))] diff --git a/tests/extension/test_gemv.py b/tests/extension/test_gemv.py index c49bf67..1ebecc4 100644 --- a/tests/extension/test_gemv.py +++ b/tests/extension/test_gemv.py @@ -43,6 +43,66 @@ def test_bf16_gemv_matches_small_decode_batches(m, n, k): torch.testing.assert_close(actual, expected, rtol=0.02, atol=0.5) +@skip_no_gemv +@pytest.mark.parametrize("m", [2, 4]) +@pytest.mark.parametrize( + "n,k", + [ + (1024, 4096), + (4096, 4096), + (11008, 4096), + (4096, 11008), + (14336, 4096), + (4096, 14336), + (5120, 5120), + (13824, 5120), + (5120, 13824), + (16384, 4096), + (4096, 16384), + (512, 3584), + (3584, 3584), + (18944, 3584), + (3584, 18944), + (1024, 8192), + (8192, 8192), + (28672, 8192), + (8192, 28672), + (2048, 2048), + (8192, 2048), + (2048, 8192), + ], +) +def test_bf16_gemv_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) + expected = F.linear(x, weight) + torch.testing.assert_close(actual, expected, rtol=0.02, atol=0.25) + + +@skip_no_gemv +@pytest.mark.parametrize( + "m,n,k", + [ + (1, 8192, 2048), + (8, 4096, 11008), + (8, 512, 3584), + (8, 1024, 8192), + (8, 2048, 8192), + ], +) +def test_bf16_gemv_matches_half_cta_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) + 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(): torch.manual_seed(23) diff --git a/tests/extension/test_linear_dispatch.py b/tests/extension/test_linear_dispatch.py index ae9cf01..d23fa92 100644 --- a/tests/extension/test_linear_dispatch.py +++ b/tests/extension/test_linear_dispatch.py @@ -6,6 +6,7 @@ import torch.nn.functional as F from astrai.extension import explain, is_available, linear, op_backend from astrai.extension.backend import linear as public_linear +from astrai.extension.backend.linear import _AUTO_GEMV_SHAPES from astrai.model.components.linear import Linear GEMV_AVAILABLE = ( @@ -23,6 +24,45 @@ def test_linear_backend_is_public(): assert linear is public_linear +def test_sm89_common_shape_policy_keeps_only_validated_families_enabled(): + common = { + (1024, 4096), + (4096, 4096), + (11008, 4096), + (4096, 11008), + (14336, 4096), + (4096, 14336), + (5120, 5120), + (13824, 5120), + (5120, 13824), + (16384, 4096), + (4096, 16384), + } + subthreshold_m4 = {(4096, 4096), (11008, 4096), (4096, 11008)} + qwen2_7b = { + (512, 3584), + (3584, 3584), + (18944, 3584), + (3584, 18944), + } + llama3_70b = { + (1024, 8192), + (8192, 8192), + (28672, 8192), + (8192, 28672), + } + opt_1_3b = {(2048, 2048), (8192, 2048), (2048, 8192)} + policy = _AUTO_GEMV_SHAPES[(8, 9)] + assert policy[1] == opt_1_3b + assert common <= policy[2] + assert qwen2_7b | llama3_70b | opt_1_3b <= policy[2] + assert common - subthreshold_m4 <= policy[4] + assert qwen2_7b | llama3_70b <= policy[4] + assert subthreshold_m4.isdisjoint(policy[4]) + assert opt_1_3b.isdisjoint(policy[4]) + assert 8 not in policy + + def test_model_linear_routes_through_backend(monkeypatch): sentinel = torch.randn(2, 4) @@ -93,7 +133,7 @@ def test_mode_one_dispatches_supported_small_batches(monkeypatch, m): @skip_no_gemv -def test_auto_m1_falls_back_until_end_to_end_gate_passes(monkeypatch): +def test_auto_unmeasured_m1_falls_back(monkeypatch): monkeypatch.setenv("ASTRAI_GEMV", "auto") x = torch.randn(1, 1536, device="cuda", dtype=torch.bfloat16) winning = torch.randn( @@ -109,10 +149,38 @@ def test_auto_m1_falls_back_until_end_to_end_gate_passes(monkeypatch): @skip_no_gemv -def test_auto_selects_measured_sm89_small_batch_winner(monkeypatch): +@pytest.mark.parametrize( + "m,n,k", + [ + (4, 256, 1536), + (2, 1024, 4096), + (2, 11008, 4096), + (2, 4096, 11008), + (2, 14336, 4096), + (4, 4096, 14336), + (2, 5120, 5120), + (4, 13824, 5120), + (2, 5120, 13824), + (4, 16384, 4096), + (2, 4096, 16384), + (2, 512, 3584), + (4, 3584, 3584), + (2, 18944, 3584), + (4, 3584, 18944), + (2, 1024, 8192), + (4, 8192, 8192), + (2, 28672, 8192), + (4, 8192, 28672), + (1, 2048, 2048), + (2, 8192, 2048), + (1, 2048, 8192), + ], +) +def test_auto_selects_measured_sm89_small_batch_winner(monkeypatch, m, n, k): monkeypatch.setenv("ASTRAI_GEMV", "auto") - x = torch.randn(4, 1536, device="cuda", dtype=torch.bfloat16) - weight = torch.randn(256, 1536, device="cuda", dtype=torch.bfloat16) + 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) with torch.no_grad(): trace = explain("linear", x, weight) if torch.cuda.get_device_capability() == (8, 9): @@ -133,6 +201,18 @@ def test_auto_selects_measured_sm89_small_batch_winner(monkeypatch): (4, 100000, 1536), # LM head misses the 5% M=4 gate (4, 1536, 6912), # long-K accumulation changed checkpoint greedy output (8, 256, 1536), # remaining M=8 winners miss the 3% end-to-end gate + (1, 4096, 4096), # isolated M=1 winner misses the projection-chain gate + (8, 1024, 4096), # isolated M=8 winner misses the projection-chain gate + (4, 12288, 4096), # GPT-NeoX fused QKV was not measured as a winner + (4, 4096, 4096), # LLaMA 2 7B M=4 chain misses the 3% gate + (4, 11008, 4096), + (4, 4096, 11008), + (1, 3584, 3584), # Qwen2 M=1 chain misses the 3% gate + (8, 3584, 3584), # Qwen2 Q/O misses the M=8 per-shape gate + (1, 8192, 8192), # LLaMA 3 70B M=1 projections miss the per-shape gate + (8, 1024, 8192), # LLaMA 3 70B K/V loses at wrapper level for M=8 + (4, 8192, 2048), # OPT up loses at wrapper level for M=4 + (8, 2048, 2048), # OPT M=8 chain and Q/K/V/O both regress ], ) def test_auto_rejects_measured_small_batch_losers(monkeypatch, m, n, k): diff --git a/tests/extension/test_swiglu.py b/tests/extension/test_swiglu.py new file mode 100644 index 0000000..04165b7 --- /dev/null +++ b/tests/extension/test_swiglu.py @@ -0,0 +1,99 @@ +import pytest +import torch +import torch.nn.functional as F + +from astrai.extension import bf16_swiglu, is_available + +SWIGLU_AVAILABLE = ( + torch.cuda.is_available() + and is_available("bf16_swiglu") + and torch.cuda.get_device_capability() >= (8, 0) +) +skip_no_swiglu = pytest.mark.skipif( + not SWIGLU_AVAILABLE, + reason="BF16 SwiGLU requires a built kernel and compute capability 8.0+", +) + + +def reference_swiglu(x, up_weight, gate_weight): + return F.linear(x, up_weight) * F.silu(F.linear(x, gate_weight)) + + +@skip_no_swiglu +@pytest.mark.parametrize("m", [1, 2, 4, 8]) +@pytest.mark.parametrize("n,k", [(6912, 1536), (4096, 4096), (11008, 4096)]) +def test_bf16_swiglu_matches_common_dense_mlp_shapes(m, n, k): + torch.manual_seed(37 + m) + x = torch.randn(m, k, device="cuda", dtype=torch.bfloat16) * 0.1 + up_weight = torch.randn(n, k, device="cuda", dtype=torch.bfloat16) * (k**-0.5) + gate_weight = torch.randn(n, k, device="cuda", dtype=torch.bfloat16) * (k**-0.5) + actual = bf16_swiglu(x, up_weight, gate_weight) + expected = reference_swiglu(x, up_weight, gate_weight) + assert actual.shape == (m, n) + torch.testing.assert_close(actual, expected, rtol=0.03, atol=0.01) + + +@skip_no_swiglu +def test_bf16_swiglu_preserves_vector_shape(): + x = torch.randn(1536, device="cuda", dtype=torch.bfloat16) * 0.1 + up_weight = torch.randn(256, 1536, device="cuda", dtype=torch.bfloat16) * 0.02 + gate_weight = torch.randn_like(up_weight) * 0.02 + actual = bf16_swiglu(x, up_weight, gate_weight) + expected = reference_swiglu(x, up_weight, gate_weight) + assert actual.shape == (256,) + torch.testing.assert_close(actual, expected, rtol=0.03, atol=0.01) + + +@skip_no_swiglu +def test_bf16_swiglu_uses_current_stream_and_cuda_graph(): + torch.manual_seed(43) + x = torch.randn(4, 1536, device="cuda", dtype=torch.bfloat16) * 0.1 + up_weight = torch.randn(6912, 1536, device="cuda", dtype=torch.bfloat16) * 0.02 + gate_weight = torch.randn_like(up_weight) * 0.02 + stream = torch.cuda.Stream() + with torch.cuda.stream(stream): + for _ in range(3): + bf16_swiglu(x, up_weight, gate_weight) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + actual = bf16_swiglu(x, up_weight, gate_weight) + x.copy_(torch.randn_like(x) * 0.1) + graph.replay() + stream.synchronize() + expected = reference_swiglu(x, up_weight, gate_weight) + torch.testing.assert_close(actual, expected, rtol=0.03, atol=0.01) + + +@skip_no_swiglu +@pytest.mark.parametrize( + "make_args,error", + [ + ( + lambda: ( + torch.randn(9, 16, device="cuda", dtype=torch.bfloat16), + torch.randn(8, 16, device="cuda", dtype=torch.bfloat16), + torch.randn(8, 16, device="cuda", dtype=torch.bfloat16), + ), + "M must", + ), + ( + lambda: ( + torch.randn(2, 15, device="cuda", dtype=torch.bfloat16), + torch.randn(8, 15, device="cuda", dtype=torch.bfloat16), + torch.randn(8, 15, device="cuda", dtype=torch.bfloat16), + ), + "divisible by 8", + ), + ( + lambda: ( + torch.randn(2, 16, device="cuda", dtype=torch.bfloat16), + torch.randn(8, 16, device="cuda", dtype=torch.bfloat16), + torch.randn(7, 16, device="cuda", dtype=torch.bfloat16), + ), + "identical shapes", + ), + ], +) +def test_bf16_swiglu_rejects_unsupported_inputs(make_args, error): + with pytest.raises(RuntimeError, match=error): + bf16_swiglu(*make_args()) diff --git a/tests/extension/test_swiglu_dispatch.py b/tests/extension/test_swiglu_dispatch.py new file mode 100644 index 0000000..19b8508 --- /dev/null +++ b/tests/extension/test_swiglu_dispatch.py @@ -0,0 +1,94 @@ +import logging + +import pytest +import torch +import torch.nn.functional as F + +from astrai.extension import is_available, swiglu +from astrai.model.components.mlp import MLP + +SWIGLU_AVAILABLE = ( + torch.cuda.is_available() + and is_available("bf16_swiglu") + and torch.cuda.get_device_capability() >= (8, 0) +) +skip_no_swiglu = pytest.mark.skipif( + not SWIGLU_AVAILABLE, + reason="BF16 SwiGLU requires a built kernel and compute capability 8.0+", +) + + +def reference_swiglu(x, up_weight, gate_weight): + return F.linear(x, up_weight) * F.silu(F.linear(x, gate_weight)) + + +def test_cpu_and_training_calls_fall_back_with_gradients(monkeypatch): + monkeypatch.setenv("ASTRAI_SWIGLU", "1") + x = torch.randn(2, 8, requires_grad=True) + up_weight = torch.randn(4, 8, requires_grad=True) + gate_weight = torch.randn(4, 8, requires_grad=True) + actual = swiglu(x, up_weight, gate_weight) + expected = reference_swiglu(x, up_weight, gate_weight) + torch.testing.assert_close(actual, expected) + actual.sum().backward() + assert x.grad is not None + assert up_weight.grad is not None + assert gate_weight.grad is not None + + +def test_invalid_mode_warns_and_uses_auto(monkeypatch, caplog): + monkeypatch.setenv("ASTRAI_SWIGLU", "invalid-test-mode") + with caplog.at_level(logging.WARNING): + actual = swiglu(torch.randn(2, 8), torch.randn(4, 8), torch.randn(4, 8)) + assert actual.shape == (2, 4) + assert "using auto" in caplog.text + + +def test_mlp_routes_through_swiglu_backend(monkeypatch): + sentinel = torch.randn(2, 4) + + def fake_swiglu(x, up_weight, gate_weight): + assert x.shape == (2, 3) + assert up_weight.shape == gate_weight.shape == (4, 3) + return sentinel + + monkeypatch.setattr("astrai.model.components.mlp.swiglu", fake_swiglu) + layer = MLP(3, 4) + output = layer(torch.randn(2, 3)) + assert output["hidden_states"].shape == (2, 3) + + +@skip_no_swiglu +def test_mode_zero_disables_fused_kernel(monkeypatch): + monkeypatch.setenv("ASTRAI_SWIGLU", "0") + x = torch.randn(4, 1536, device="cuda", dtype=torch.bfloat16) * 0.1 + up_weight = torch.randn(6912, 1536, device="cuda", dtype=torch.bfloat16) * 0.02 + gate_weight = torch.randn_like(up_weight) * 0.02 + with torch.no_grad(): + actual = swiglu(x, up_weight, gate_weight) + expected = reference_swiglu(x, up_weight, gate_weight) + torch.testing.assert_close(actual, expected) + + +@skip_no_swiglu +def test_mode_one_forces_supported_shape(monkeypatch): + monkeypatch.setenv("ASTRAI_SWIGLU", "1") + x = torch.randn(4, 1536, device="cuda", dtype=torch.bfloat16) * 0.1 + up_weight = torch.randn(6912, 1536, device="cuda", dtype=torch.bfloat16) * 0.02 + gate_weight = torch.randn_like(up_weight) * 0.02 + with torch.no_grad(): + actual = swiglu(x, up_weight, gate_weight) + expected = reference_swiglu(x, up_weight, gate_weight) + torch.testing.assert_close(actual, expected, rtol=0.03, atol=0.01) + + +@skip_no_swiglu +def test_auto_falls_back_until_shape_is_qualified(monkeypatch): + 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) + gate_weight = torch.randn_like(up_weight) + with torch.no_grad(): + actual = swiglu(x, up_weight, gate_weight) + expected = reference_swiglu(x, up_weight, gate_weight) + torch.testing.assert_close(actual, expected)