docs: align CUDA kernel and RoPE docs with code

- Fix rotary docs to describe cos/sin freqs_cis table, not complex buffer
- Replace attn_prefill with attn_paged_prefill for the CudaBackend path
- Register attn_paged_prefill in kernel overview, layout, and module list
- Add qo_indptr and InferenceWorkspace to architecture class diagram
- Add FrequencyPenaltyStrategy to sampling design patterns
This commit is contained in:
2026-08-03 20:54:40 +08:00
parent 85f0461b3b
commit cdf9145ecf
5 changed files with 57 additions and 25 deletions
+9 -7
View File
@@ -41,12 +41,14 @@ RoPE embeds position into Q/K vectors via complex rotation:
$$ q_i = R_i W_q x_i, \quad k_j = R_j W_k x_j, \quad q_i^T k_j = x_i^T W_q^T R_{i-j} W_k x_j $$
`RotaryEmbedding` pre-computes a complex `freqs_cis` buffer. `forward()` returns
a tensor indexed by `position_ids`. `apply_rotary_emb` applies the rotation:
during training it uses torch complex multiply (autograd-compatible); during
inference it auto-dispatches to a fused CUDA kernel when available. The key
property is that the dot product $q_i^T k_j$ depends only on the relative
position $i - j$, not the absolute positions.
`RotaryEmbedding` pre-computes a cos/sin table `freqs_cis` of shape
`[max_len, dim/2, 2]` (f32 — `[cos, sin]` pairs). `forward()` returns
a `[batch, seq_len, dim/2, 2]` slice indexed by `position_ids`.
`apply_rotary_emb` applies the rotation: during training it uses torch
complex multiply (autograd-compatible); during inference it auto-dispatches
to a fused CUDA kernel when available. The key property is that the dot
product $q_i^T k_j$ depends only on the relative position $i - j$, not the
absolute positions.
**Critical for inference**: RoPE is applied **before** KV cache write, not after. If applied after caching, position encoding drift occurs because cached K/V would have stale rotation factors.
@@ -175,7 +177,7 @@ Three-layer separation (SGLang-inspired):
Attention computation is decoupled from the model via `AttentionBackend` ABC (`astrai/extension/attention_backend.py`):
- **`TorchNativeBackend`** (default): writes K/V to cache, gathers via `req_to_token` indirect indexing, calls `F.scaled_dot_product_attention`.
- **`CudaBackend`**: decode path uses `attn_paged_decode` with `page_size=1` (the `req_to_token` table serves as the page table, each token slot is a single-token "page"); prefill path gathers K/V then calls `attn_prefill`. Falls back to `TorchNativeBackend` when kernel unavailable.
- **`CudaBackend`**: decode path uses `attn_paged_decode` with `page_size=1` (the `req_to_token` table serves as the page table, each token slot is a single-token "page"); prefill path uses the ragged-batch `attn_paged_prefill` (addresses each request via `qo_indptr` + `kv_indptr` directly against the flat pool). Falls back to `TorchNativeBackend` when kernel unavailable.
Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/rotary_backend.py`, which auto-dispatches to the fused CUDA kernel (`rotary_emb.cu`) during inference or torch complex multiply during training (for autograd compatibility). Both attention backends share the same rotary dispatch.