docs: sync all documentation with current codebase
- remove GenerationRequest and generate_with_request references (class deleted) - document cuda>flash>torch default priority and FlashAttnBackend - add ASTR_BACKEND env var to backend docs, TorchNativeBackend (default) → (fallback) - fix JsonlStore transform routing → DatasetFactory ownership - fix CudaBackend fallback chain description (FlashAttn → TorchNative) - add FlashAttnBackend to architecture strategy table - add router_stats to DecoderOutput/FFNOutput TypedDict diagrams - add decode_o_part/ml_part/decode_out to KVCache diagram - add --append_eos/--no-append_eos to IFD evaluation parameter table - update get-started CUDA kernel note (no longer requires explicit attn_backend activation) - fix python -m scripts.tools.server (no __init__.py) → direct script call
This commit is contained in:
@@ -86,8 +86,12 @@ Each kernel in `astrai/extension/lib` is compiled as an independent pybind11 mod
|
||||
`astrai/extension/attention_backend.py` provides the backend abstraction:
|
||||
|
||||
- **`AttentionBackend`** (ABC): `fwd_decode` / `fwd_prefill` abstract methods, `forward` dispatches by q_len
|
||||
- **`TorchNativeBackend`**: SDPA with indirect KV cache gather (default)
|
||||
- **`CudaBackend`**: CUDA kernel dispatch — decode via `attn_paged_decode` (page_size=1), prefill via `attn_paged_prefill` (ragged batch, `qo_indptr` + `kv_indptr`)
|
||||
- **`CudaBackend`**: CUDA kernel dispatch — decode via `attn_paged_decode` (page_size=1), prefill via `attn_paged_prefill` (ragged batch, `qo_indptr` + `kv_indptr`). Default on GPU.
|
||||
- **`FlashAttnBackend`**: Optional flash-attn dispatch with `flash_attn_with_kvcache` fast path.
|
||||
- **`TorchNativeBackend`**: SDPA with indirect KV cache gather (always-available fallback)
|
||||
|
||||
Default priority: cuda > flash > torch. Set ``ASTR_BACKEND=cuda|torch_native|flash``
|
||||
to override the default.
|
||||
|
||||
Select a backend via context manager (mirrors `torch.nn.attention.sdpa_kernel`):
|
||||
|
||||
@@ -98,7 +102,7 @@ with attn_backend(ATTN_BACKEND.CUDA):
|
||||
engine.generate("hello")
|
||||
```
|
||||
|
||||
`CudaBackend` falls back to `TorchNativeBackend` when a kernel is not available.
|
||||
`CudaBackend` falls back to `FlashAttnBackend` (when flash-attn is installed and supports the input dtype) or `TorchNativeBackend` otherwise.
|
||||
|
||||
### Rotary Backend
|
||||
|
||||
|
||||
Reference in New Issue
Block a user