- add decode-shape benchmark harness - add bf16 GEMV CUDA primitive with head-dim generic kernel - dispatch decode-time linear layers to gemv for M=1 - extend gemv coverage to small decode batches
24 lines
771 B
Python
24 lines
771 B
Python
"""Stateless wrapper for the directly callable BF16 GEMV primitive."""
|
|
|
|
from typing import Optional
|
|
|
|
import torch
|
|
|
|
from astrai.extension.loader import get_module
|
|
|
|
|
|
def bf16_gemv(
|
|
x: torch.Tensor,
|
|
weight: torch.Tensor,
|
|
bias: Optional[torch.Tensor] = None,
|
|
) -> torch.Tensor:
|
|
"""Compute ``F.linear(x, weight, bias)`` for up to eight BF16 rows.
|
|
|
|
``x`` must have shape ``[K]`` or ``[M, K]`` with M in ``{1, 2, 4, 8}``,
|
|
and ``weight`` must be a contiguous row-major ``[N, K]`` tensor. The CUDA
|
|
kernel reuses each weight row across M, accumulates in FP32, and returns
|
|
BF16. This primitive is inference-only and intentionally performs no
|
|
fallback or model-level dispatch.
|
|
"""
|
|
return get_module("bf16_gemv").bf16_gemv(x, weight, bias)
|