feat: wire up paged decode CUDA kernel to Python extension
- Add attn_paged_decode wrapper in ops.py with gather fallback - Register kernel in loader.py and export from __init__.py - Extract test_utils.cuh shared by all attention unit tests - Rename attn_paged_vs_contiguous.cu to attn_paged_decode_test.cu - Refactor decode/prefill tests to use common bf16 helpers and cpu ref - Fix k_cache dim check in attn_paged_decode.cu
This commit is contained in:
@@ -3,16 +3,18 @@
|
||||
Public API:
|
||||
- ``attn_decode`` — single-query decode attention
|
||||
- ``attn_prefill`` — multi-query prefill attention
|
||||
- ``attn_paged_decode`` — paged decode attention (direct page-table access)
|
||||
|
||||
Each wrapper dispatches to its compiled CUDA kernel (``astrai.extension.attn_*``)
|
||||
when available, otherwise falls back to ``torch.nn.functional.scaled_dot_product_attention``.
|
||||
"""
|
||||
|
||||
from astrai.extension.loader import KERNEL_NAMES, is_available
|
||||
from astrai.extension.ops import attn_decode, attn_prefill
|
||||
from astrai.extension.ops import attn_decode, attn_paged_decode, attn_prefill
|
||||
|
||||
__all__ = [
|
||||
"attn_decode",
|
||||
"attn_paged_decode",
|
||||
"attn_prefill",
|
||||
"is_available",
|
||||
"KERNEL_NAMES",
|
||||
|
||||
Reference in New Issue
Block a user