Files
AstrAI/astrai/extension/__init__.py
T
ViperEkura 7540acb43e perf: dispatch linear gemv by decode batch size and unify extension style
- replace the per-shape auto tables in the linear backend with an M-banded rule (M in [2,4] on compute capability 8.0+) that measured at the HBM bandwidth floor across every family, and fold the capability check into the capable guard
- drop the unreachable swiglu auto shape-table machinery so both backends share one env-mode ladder via the new dispatch.env_mode helper
- add __all__ across extension modules, name the rotary registration records, and unify typing to the typing-module style
- rewrite test_linear_dispatch.py around behavioral routing assertions and document the M-banded policy in the developer docs
- Benchmark: L20 SM89, Python dispatch overhead 2.9us to 1.5us, auto now covers every projection shape at M in [2,4].
2026-09-03 07:23:13 +08:00

97 lines
2.3 KiB
Python

"""CUDA kernel wrappers, operator dispatch, and backend selection.
Public API:
- ``attention``, ``linear``, ``swiglu``, ``apply_rotary_emb`` — op
families with safe torch fallbacks (see ``astrai.extension.backend``)
- ``attn_decode`` / ``attn_prefill`` / ``attn_paged_decode`` /
``attn_paged_prefill`` — direct attention kernel wrappers
- ``bf16_gemv`` / ``bf16_swiglu`` — directly callable linear/MLP kernels
- ``AttentionBackend`` / ``TorchNativeBackend`` / ``CudaBackend`` /
``FlashAttnBackend`` — attention backend strategies
- ``resolve`` / ``explain`` / ``op_backend`` / ``env_mode`` — the shared
operator dispatcher (see ``astrai.extension.dispatch``)
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
(blhd). Scale is always ``1/sqrt(head_dim)``. Wrapper functions call their
compiled CUDA kernels directly; fallback is the backend's responsibility.
"""
from astrai.extension.backend import (
ATTN_BACKEND,
AttentionBackend,
AttentionBackendFactory,
CudaBackend,
FlashAttnBackend,
TorchNativeBackend,
apply_rotary_emb,
attention,
attn_backend,
get_backend,
linear,
swiglu,
)
from astrai.extension.dispatch import (
Axes,
ExplicitSelectionError,
ImplRecord,
Resolution,
Spec,
axis,
env_mode,
explain,
explain_plan,
op_backend,
register_env_alias,
register_family,
resolve,
resolve_plan,
tensor_axes,
)
from astrai.extension.loader import KERNEL_NAMES, is_available
from astrai.extension.ops import (
TensorLayout,
attn_decode,
attn_paged_decode,
attn_prefill,
bf16_gemv,
bf16_swiglu,
)
__all__ = [
"ATTN_BACKEND",
"AttentionBackend",
"AttentionBackendFactory",
"CudaBackend",
"TorchNativeBackend",
"FlashAttnBackend",
"TensorLayout",
"attention",
"attn_backend",
"get_backend",
"linear",
"swiglu",
"attn_decode",
"attn_paged_decode",
"attn_prefill",
"bf16_gemv",
"bf16_swiglu",
"is_available",
"KERNEL_NAMES",
"apply_rotary_emb",
"Axes",
"ExplicitSelectionError",
"ImplRecord",
"Resolution",
"Spec",
"axis",
"env_mode",
"explain",
"explain_plan",
"op_backend",
"register_env_alias",
"register_family",
"resolve",
"resolve_plan",
"tensor_axes",
]