- delete csrc/kernels/gemm.cu and swiglu.cu and drop their CMake and setup.py registration - remove the ops wrappers plus backend/linear.py and backend/swiglu.py so Linear and MLP call F.linear directly - drop the four gemm and swiglu kernel test files and prune the stale cuda_kernels.md sections - add csrc/bench benchmarks for the remaining kernels: attention decode prefill paged decode paged prefill versus single-launch SDPA references, rotary versus the torch fallback, fp8 quantize and mm_fp8 versus torch baselines - attention, rotary_emb, and fp8_ops kernels are unchanged
28 lines
594 B
Python
28 lines
594 B
Python
"""Backend selection, fallbacks, and execution policies."""
|
|
|
|
from astrai.extension.backend.attention import (
|
|
ATTN_BACKEND,
|
|
AttentionBackend,
|
|
AttentionBackendFactory,
|
|
CudaBackend,
|
|
FlashAttnBackend,
|
|
TorchNativeBackend,
|
|
attention,
|
|
attn_backend,
|
|
get_backend,
|
|
)
|
|
from astrai.extension.backend.rotary import apply_rotary_emb
|
|
|
|
__all__ = [
|
|
"ATTN_BACKEND",
|
|
"AttentionBackend",
|
|
"AttentionBackendFactory",
|
|
"CudaBackend",
|
|
"FlashAttnBackend",
|
|
"TorchNativeBackend",
|
|
"apply_rotary_emb",
|
|
"attention",
|
|
"attn_backend",
|
|
"get_backend",
|
|
]
|