Compare commits
17
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8447f88f61 | ||
|
|
b1b65a657e | ||
|
|
3439e3104e | ||
|
|
288ba20db1 | ||
|
|
020e2eff4e | ||
|
|
1c7369f293 | ||
|
|
0fc1b1bd46 | ||
|
|
d7db37a70f | ||
|
|
6d98bb4f9f | ||
|
|
925cbedc93 | ||
|
|
fda82ee232 | ||
|
|
4b25664c79 | ||
|
|
a27c8a819d | ||
|
|
91acaf4b0b | ||
|
|
41dcf0feb9 | ||
|
|
9960f79920 | ||
|
|
7feeb0b93e |
+9
-7
@@ -20,9 +20,6 @@ Run the following checks **in order** — CI will reject if any fail.
|
|||||||
ruff format .
|
ruff format .
|
||||||
```
|
```
|
||||||
|
|
||||||
> **Note**: `ruff format` may rename parameters (e.g. `mask` → `attn_mask`).
|
|
||||||
> Always review the diff after formatting.
|
|
||||||
|
|
||||||
### 2. Import sorting
|
### 2. Import sorting
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
@@ -44,7 +41,7 @@ python -u -m pytest tests/ -v
|
|||||||
|
|
||||||
> Failed tests may leave orphan tempdirs under `%TEMP%`. Clean them manually if needed.
|
> Failed tests may leave orphan tempdirs under `%TEMP%`. Clean them manually if needed.
|
||||||
|
|
||||||
### 4. (Optional) Full pre-commit check
|
### 4. (Optional) Full pre-commit check script
|
||||||
|
|
||||||
If you have Git Bash available:
|
If you have Git Bash available:
|
||||||
|
|
||||||
@@ -52,12 +49,17 @@ If you have Git Bash available:
|
|||||||
bash scripts/pre_commit.sh
|
bash scripts/pre_commit.sh
|
||||||
```
|
```
|
||||||
|
|
||||||
This runs format check, import sort check, and tests in one go.
|
The script installs development dependencies by default, then runs the format
|
||||||
|
check, import sort check, and tests. If dependencies are already installed, use:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
bash scripts/pre_commit.sh --skip-deps
|
||||||
|
```
|
||||||
|
|
||||||
## Commit Style
|
## Commit Style
|
||||||
|
|
||||||
```
|
```
|
||||||
fix/feat/chore/docs/refactor/perf/test/style/ci/build/revert : short description (~50 chars)
|
type: short description (~50 chars)
|
||||||
|
|
||||||
- bullet point body (each ~60 chars)
|
- bullet point body (each ~60 chars)
|
||||||
```
|
```
|
||||||
@@ -73,7 +75,7 @@ fix/feat/chore/docs/refactor/perf/test/style/ci/build/revert : short description
|
|||||||
|---------|-------|-----|
|
|---------|-------|-----|
|
||||||
| `ruff check --select I` fails | Wrong import order | `ruff check . --select I --fix .` then `ruff format .` |
|
| `ruff check --select I` fails | Wrong import order | `ruff check . --select I --fix .` then `ruff format .` |
|
||||||
| `ruff format` changed many files | Not formatted before commit | Review diff carefully before staging |
|
| `ruff format` changed many files | Not formatted before commit | Review diff carefully before staging |
|
||||||
| Pre-commit hook rejects | Tests or lint failed | Fix individually, do not `--no-verify` |
|
| Pre-commit check script fails | Dependency install, tests, or lint failed | Fix the failing step; use `--skip-deps` only when dependencies are already installed |
|
||||||
| Tests fail with tempdir left | Test crash | Clean `%TEMP%` manually |
|
| Tests fail with tempdir left | Test crash | Clean `%TEMP%` manually |
|
||||||
|
|
||||||
## Submitting Changes
|
## Submitting Changes
|
||||||
|
|||||||
@@ -56,6 +56,8 @@ End-to-end walkthrough in 5 steps:
|
|||||||
|
|
||||||
**1. Install**
|
**1. Install**
|
||||||
|
|
||||||
|
AstrAI requires Python 3.12+ and pins PyTorch exactly to `2.11.0`. Training, `scripts/tools/generate.py`, generation evaluations, and the generation demos require CUDA; CPU support is limited to components with an explicit CPU device path, such as the HTTP server and direct-scoring evaluations.
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
git clone https://github.com/ViperEkura/AstrAI.git
|
git clone https://github.com/ViperEkura/AstrAI.git
|
||||||
cd AstrAI
|
cd AstrAI
|
||||||
@@ -132,7 +134,7 @@ Check out the demos in the `scripts/demo/` folder:
|
|||||||
# Download model weights (required before running demos)
|
# Download model weights (required before running demos)
|
||||||
python scripts/demo/download.py # model → params/
|
python scripts/demo/download.py # model → params/
|
||||||
|
|
||||||
# Interactive streaming chat (multi-turn, maintains history)
|
# Single-turn interactive streaming prompt loop (no conversation history)
|
||||||
python scripts/demo/stream_chat.py
|
python scripts/demo/stream_chat.py
|
||||||
# Type your message after >>, type !exit to quit
|
# Type your message after >>, type !exit to quit
|
||||||
|
|
||||||
@@ -183,7 +185,7 @@ docker run --gpus all -v /path/to/data:/data -it astrai:latest
|
|||||||
# Docker Compose (GPU, default)
|
# Docker Compose (GPU, default)
|
||||||
docker compose up -d
|
docker compose up -d
|
||||||
|
|
||||||
# Docker Compose (CPU only)
|
# Docker Compose CPU server profile (CUDA-only generation scripts/demos are unavailable)
|
||||||
docker compose --profile cpu up -d
|
docker compose --profile cpu up -d
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|||||||
@@ -63,6 +63,11 @@ class AutoRegressiveLMConfig(BaseModelConfig):
|
|||||||
n_shared_experts (Optional[int]): Number of shared experts, MoE only. Defaults to None.
|
n_shared_experts (Optional[int]): Number of shared experts, MoE only. Defaults to None.
|
||||||
n_activated_experts (Optional[int]): Number of activated experts per token, MoE only. Defaults to None.
|
n_activated_experts (Optional[int]): Number of activated experts per token, MoE only. Defaults to None.
|
||||||
topk_method (Optional[str]): Top-k routing method, MoE only. Defaults to None.
|
topk_method (Optional[str]): Top-k routing method, MoE only. Defaults to None.
|
||||||
|
moe_intermediate_size (Optional[int]): Expert hidden dim, defaults to intermediate_size if None. MoE only.
|
||||||
|
shared_expert_intermediate_size (Optional[int]): Shared expert hidden dim, defaults to intermediate_size if None. MoE only.
|
||||||
|
norm_topk_prob (bool): Normalize top-k routing probabilities. Defaults to True.
|
||||||
|
decoder_sparse_step (int): Frequency of MoE layers, 1=every layer. Defaults to 1.
|
||||||
|
mlp_only_layers (Optional[list[int]]): Layer indices using dense MLP instead of MoE. Defaults to None.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
vocab_size: Optional[int] = None
|
vocab_size: Optional[int] = None
|
||||||
@@ -87,6 +92,11 @@ class AutoRegressiveLMConfig(BaseModelConfig):
|
|||||||
n_shared_experts: Optional[int] = None
|
n_shared_experts: Optional[int] = None
|
||||||
n_activated_experts: Optional[int] = None
|
n_activated_experts: Optional[int] = None
|
||||||
topk_method: Optional[str] = None
|
topk_method: Optional[str] = None
|
||||||
|
moe_intermediate_size: Optional[int] = None
|
||||||
|
shared_expert_intermediate_size: Optional[int] = None
|
||||||
|
norm_topk_prob: bool = True
|
||||||
|
decoder_sparse_step: int = 1
|
||||||
|
mlp_only_layers: Optional[list[int]] = None
|
||||||
|
|
||||||
@field_validator("attn_type")
|
@field_validator("attn_type")
|
||||||
def _validate_attn_type(cls, v: str) -> str:
|
def _validate_attn_type(cls, v: str) -> str:
|
||||||
@@ -102,6 +112,12 @@ class AutoRegressiveLMConfig(BaseModelConfig):
|
|||||||
raise ValueError(f"ffn_type must be one of {sorted(_FFN_TYPES)}, got {v!r}")
|
raise ValueError(f"ffn_type must be one of {sorted(_FFN_TYPES)}, got {v!r}")
|
||||||
return v
|
return v
|
||||||
|
|
||||||
|
@field_validator("decoder_sparse_step")
|
||||||
|
def _validate_decoder_sparse_step(cls, v: int) -> int:
|
||||||
|
if v < 1:
|
||||||
|
raise ValueError(f"decoder_sparse_step must be at least 1, got {v}")
|
||||||
|
return v
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@ConfigFactory.register("embedding")
|
@ConfigFactory.register("embedding")
|
||||||
|
|||||||
@@ -61,6 +61,7 @@ class TrainConfig(BaseConfig):
|
|||||||
val_split (Optional[float]): Ratio to split from training dataset for validation, e.g. 0.05. Defaults to None.
|
val_split (Optional[float]): Ratio to split from training dataset for validation, e.g. 0.05. Defaults to None.
|
||||||
val_step (int): Number of optimizer steps between validation runs. Defaults to 1000.
|
val_step (int): Number of optimizer steps between validation runs. Defaults to 1000.
|
||||||
neftune_alpha (float): NEFTune noise alpha, 0=disabled, typical: 5.0. Defaults to 0.0.
|
neftune_alpha (float): NEFTune noise alpha, 0=disabled, typical: 5.0. Defaults to 0.0.
|
||||||
|
moe_aux_loss_coef (float): Weight applied to the MoE load-balancing loss. Defaults to 0.01.
|
||||||
rollout_interval (int): Number of optimizer steps between online rollouts. Defaults to 512.
|
rollout_interval (int): Number of optimizer steps between online rollouts. Defaults to 512.
|
||||||
rollout_temperature (float): Sampling temperature for online rollout. Defaults to 0.7.
|
rollout_temperature (float): Sampling temperature for online rollout. Defaults to 0.7.
|
||||||
rollout_top_k (int): Top-k filtering for online rollout, 0=disable. Defaults to 0.
|
rollout_top_k (int): Top-k filtering for online rollout, 0=disable. Defaults to 0.
|
||||||
@@ -112,6 +113,7 @@ class TrainConfig(BaseConfig):
|
|||||||
val_split: Optional[float] = None
|
val_split: Optional[float] = None
|
||||||
val_step: int = 1000
|
val_step: int = 1000
|
||||||
neftune_alpha: float = 0.0
|
neftune_alpha: float = 0.0
|
||||||
|
moe_aux_loss_coef: float = 0.01
|
||||||
|
|
||||||
rollout_interval: int = 512
|
rollout_interval: int = 512
|
||||||
rollout_temperature: float = 0.7
|
rollout_temperature: float = 0.7
|
||||||
@@ -187,7 +189,9 @@ class TrainConfig(BaseConfig):
|
|||||||
raise ValueError(f"rollout_top_p must be in (0, 1], got {v}")
|
raise ValueError(f"rollout_top_p must be in (0, 1], got {v}")
|
||||||
return v
|
return v
|
||||||
|
|
||||||
@field_validator("rollout_top_k", "num_workers", "neftune_alpha")
|
@field_validator(
|
||||||
|
"rollout_top_k", "num_workers", "neftune_alpha", "moe_aux_loss_coef"
|
||||||
|
)
|
||||||
def _validate_non_negative(cls, v):
|
def _validate_non_negative(cls, v):
|
||||||
if v < 0:
|
if v < 0:
|
||||||
raise ValueError(f"must be non-negative, got {v}")
|
raise ValueError(f"must be non-negative, got {v}")
|
||||||
|
|||||||
@@ -25,6 +25,7 @@ from astrai.extension.attention_backend import (
|
|||||||
get_backend,
|
get_backend,
|
||||||
)
|
)
|
||||||
from astrai.extension.attention_ops import (
|
from astrai.extension.attention_ops import (
|
||||||
|
TensorLayout,
|
||||||
attn_decode,
|
attn_decode,
|
||||||
attn_paged_decode,
|
attn_paged_decode,
|
||||||
attn_prefill,
|
attn_prefill,
|
||||||
@@ -37,6 +38,7 @@ __all__ = [
|
|||||||
"AttentionBackend",
|
"AttentionBackend",
|
||||||
"CudaBackend",
|
"CudaBackend",
|
||||||
"TorchNativeBackend",
|
"TorchNativeBackend",
|
||||||
|
"TensorLayout",
|
||||||
"attention",
|
"attention",
|
||||||
"attn_backend",
|
"attn_backend",
|
||||||
"get_backend",
|
"get_backend",
|
||||||
|
|||||||
@@ -38,8 +38,10 @@ import torch
|
|||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from astrai.extension.attention_ops import attn_paged_decode, attn_prefill
|
from astrai.extension.attention_ops import (
|
||||||
from astrai.extension.loader import is_available
|
attn_paged_decode,
|
||||||
|
attn_paged_prefill,
|
||||||
|
)
|
||||||
from astrai.inference.core.cache import KVCache
|
from astrai.inference.core.cache import KVCache
|
||||||
|
|
||||||
_current_backend: contextvars.ContextVar["AttentionBackend"] = contextvars.ContextVar(
|
_current_backend: contextvars.ContextVar["AttentionBackend"] = contextvars.ContextVar(
|
||||||
@@ -272,12 +274,12 @@ class TorchNativeBackend(AttentionBackend):
|
|||||||
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
||||||
|
|
||||||
max_len = kv_cache.max_len
|
max_len = kv_cache.max_len
|
||||||
if kv_cache.page_table is not None:
|
|
||||||
indices = kv_cache.page_table
|
|
||||||
else:
|
|
||||||
indices = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len]
|
indices = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len]
|
||||||
if kv_cache.decode_mask is not None:
|
# Zero out padding positions so gather never touches invalid slots.
|
||||||
pos_mask = kv_cache.decode_mask
|
# Decode: attn_mask[:,0,0] is exactly the per-position validity
|
||||||
|
# mask ([B, max_len], True=keep). Prefill: fall back to seq_lens.
|
||||||
|
if q.size(1) == 1 and attn_mask is not None and attn_mask.dim() == 4:
|
||||||
|
pos_mask = attn_mask[:, 0, 0]
|
||||||
else:
|
else:
|
||||||
pos_mask = (
|
pos_mask = (
|
||||||
torch.arange(max_len, device=q.device)[None, :]
|
torch.arange(max_len, device=q.device)[None, :]
|
||||||
@@ -307,24 +309,19 @@ _default_backend = TorchNativeBackend()
|
|||||||
class CudaBackend(AttentionBackend):
|
class CudaBackend(AttentionBackend):
|
||||||
"""CUDA kernel backend with direct KV cache access.
|
"""CUDA kernel backend with direct KV cache access.
|
||||||
|
|
||||||
Decode path: writes K/V to cache, then calls ``attn_paged_decode``
|
Decode path: writes K/V to the flat pool, then calls
|
||||||
with ``page_size=1`` (each token slot is a single-token "page").
|
``attn_paged_decode`` with req_to_token + kv_indptr.
|
||||||
The ``req_to_token`` table serves directly as the page table.
|
|
||||||
|
|
||||||
Prefill path: writes K/V to cache, gathers full-sequence K/V via
|
Prefill path: writes K/V to the flat pool, then calls
|
||||||
indirect indexing (same as TorchNativeBackend), then calls
|
``attn_paged_prefill`` with ragged-batch support via qo_indptr +
|
||||||
``attn_prefill``.
|
kv_indptr.
|
||||||
|
|
||||||
Training path (``kv_cache is None``): calls ``attn_prefill`` directly
|
``kv_cache is None`` (training) is not handled — use
|
||||||
on the projected q/k/v.
|
``TorchNativeBackend`` for training.
|
||||||
|
|
||||||
Falls back to ``TorchNativeBackend`` for any path where the
|
Raises ``RuntimeError`` if the required kernel is not available.
|
||||||
corresponding CUDA kernel is not available.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self):
|
|
||||||
self._fallback = TorchNativeBackend()
|
|
||||||
|
|
||||||
def fwd_decode(
|
def fwd_decode(
|
||||||
self,
|
self,
|
||||||
q: Tensor,
|
q: Tensor,
|
||||||
@@ -335,47 +332,29 @@ class CudaBackend(AttentionBackend):
|
|||||||
attn_mask: Optional[Tensor] = None,
|
attn_mask: Optional[Tensor] = None,
|
||||||
is_causal: bool = False,
|
is_causal: bool = False,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
if kv_cache is None or not is_available("attn_paged_decode"):
|
if kv_cache is None:
|
||||||
return self._fallback.fwd_decode(
|
raise RuntimeError("CudaBackend does not support training (kv_cache=None)")
|
||||||
q, k, v, kv_cache, layer_id, attn_mask, is_causal
|
|
||||||
)
|
|
||||||
|
|
||||||
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
||||||
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
||||||
|
|
||||||
max_len = kv_cache.max_len
|
b = q.size(0)
|
||||||
|
q_3d = q.squeeze(1)
|
||||||
|
|
||||||
if kv_cache.page_table is not None:
|
kv_indptr = kv_cache.kv_indptr
|
||||||
page_table = kv_cache.page_table
|
|
||||||
else:
|
|
||||||
page_table = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len]
|
|
||||||
|
|
||||||
k_cache = kv_cache.k_buffer[layer_id].unsqueeze(1)
|
|
||||||
v_cache = kv_cache.v_buffer[layer_id].unsqueeze(1)
|
|
||||||
|
|
||||||
if q.size(0) == 1:
|
|
||||||
mask = None
|
|
||||||
elif kv_cache.decode_mask is not None:
|
|
||||||
mask = kv_cache.decode_mask
|
|
||||||
else:
|
|
||||||
mask = (
|
|
||||||
torch.arange(max_len, device=q.device)[None, :]
|
|
||||||
< kv_cache.seq_lens[:, None]
|
|
||||||
)
|
|
||||||
|
|
||||||
out = attn_paged_decode(
|
out = attn_paged_decode(
|
||||||
q,
|
q_3d,
|
||||||
page_table,
|
kv_cache.k_buffer[layer_id],
|
||||||
k_cache,
|
kv_cache.v_buffer[layer_id],
|
||||||
v_cache,
|
kv_cache.req_to_token,
|
||||||
page_size=1,
|
kv_cache.req_pool_indices,
|
||||||
kv_len=max_len,
|
kv_indptr,
|
||||||
mask=mask,
|
kv_cache.max_len,
|
||||||
|
mask=attn_mask,
|
||||||
is_causal=is_causal,
|
is_causal=is_causal,
|
||||||
)
|
)
|
||||||
|
return out.unsqueeze(1).flatten(2)
|
||||||
out = out.flatten(2)
|
|
||||||
return out
|
|
||||||
|
|
||||||
def fwd_prefill(
|
def fwd_prefill(
|
||||||
self,
|
self,
|
||||||
@@ -388,32 +367,32 @@ class CudaBackend(AttentionBackend):
|
|||||||
is_causal: bool = False,
|
is_causal: bool = False,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
if kv_cache is None:
|
if kv_cache is None:
|
||||||
if is_available("attn_prefill"):
|
raise RuntimeError("CudaBackend does not support training (kv_cache=None)")
|
||||||
out = attn_prefill(q, k, v, mask=attn_mask, is_causal=is_causal)
|
|
||||||
return out.flatten(2)
|
|
||||||
return self._fallback.fwd_prefill(
|
|
||||||
q, k, v, kv_cache, layer_id, attn_mask, is_causal
|
|
||||||
)
|
|
||||||
|
|
||||||
if not is_available("attn_prefill"):
|
|
||||||
return self._fallback.fwd_prefill(
|
|
||||||
q, k, v, kv_cache, layer_id, attn_mask, is_causal
|
|
||||||
)
|
|
||||||
|
|
||||||
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
||||||
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
||||||
|
|
||||||
max_len = kv_cache.max_len
|
b = q.size(0)
|
||||||
indices = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len]
|
q_len = q.size(1)
|
||||||
pos_mask = (
|
|
||||||
torch.arange(max_len, device=q.device)[None, :] < kv_cache.seq_lens[:, None]
|
|
||||||
)
|
|
||||||
indices = torch.where(pos_mask, indices, torch.zeros_like(indices))
|
|
||||||
k_full = kv_cache.k_buffer[layer_id, indices]
|
|
||||||
v_full = kv_cache.v_buffer[layer_id, indices]
|
|
||||||
|
|
||||||
out = attn_prefill(q, k_full, v_full, mask=attn_mask, is_causal=is_causal)
|
kv_indptr = kv_cache.kv_indptr
|
||||||
return out.flatten(2)
|
qo_indptr = torch.arange(b + 1, dtype=torch.int32, device=q.device) * q_len
|
||||||
|
|
||||||
|
q_flat = q.reshape(b * q_len, q.size(2), q.size(3))
|
||||||
|
|
||||||
|
out = attn_paged_prefill(
|
||||||
|
q_flat,
|
||||||
|
kv_cache.k_buffer[layer_id],
|
||||||
|
kv_cache.v_buffer[layer_id],
|
||||||
|
kv_cache.req_to_token,
|
||||||
|
kv_cache.req_pool_indices,
|
||||||
|
kv_indptr,
|
||||||
|
qo_indptr,
|
||||||
|
attn_mask,
|
||||||
|
q_len,
|
||||||
|
is_causal=is_causal,
|
||||||
|
)
|
||||||
|
return out.reshape(b, q_len, q.size(2), q.size(3)).flatten(2)
|
||||||
|
|
||||||
|
|
||||||
_BACKEND_REGISTRY: dict[ATTN_BACKEND, type[AttentionBackend]] = {
|
_BACKEND_REGISTRY: dict[ATTN_BACKEND, type[AttentionBackend]] = {
|
||||||
|
|||||||
@@ -12,11 +12,24 @@ Interface (all functions):
|
|||||||
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True=keep)
|
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True=keep)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import enum
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from astrai.extension.loader import _available, _modules
|
from astrai.extension.loader import _available, _modules
|
||||||
|
|
||||||
|
|
||||||
|
class TensorLayout(enum.IntEnum):
|
||||||
|
"""Q/K/V tensor layout, mirrors the C++ ``TensorLayout`` enum in ``attn_common.h``.
|
||||||
|
|
||||||
|
Kernels internally operate on BHLD; BLHD inputs are transposed at entry.
|
||||||
|
"""
|
||||||
|
|
||||||
|
BHLD = 0 # [batch, n_heads, seq_len, head_dim]
|
||||||
|
BLHD = 1 # [batch, seq_len, n_heads, head_dim]
|
||||||
|
|
||||||
|
|
||||||
def _check_available(name: str):
|
def _check_available(name: str):
|
||||||
if not _available.get(name):
|
if not _available.get(name):
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
@@ -29,7 +42,7 @@ def attn_decode(
|
|||||||
q: torch.Tensor,
|
q: torch.Tensor,
|
||||||
k: torch.Tensor,
|
k: torch.Tensor,
|
||||||
v: torch.Tensor,
|
v: torch.Tensor,
|
||||||
mask: torch.Tensor | None = None,
|
mask: Optional[torch.Tensor] = None,
|
||||||
is_causal: bool = False,
|
is_causal: bool = False,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""GQA decode attention (q_len == 1).
|
"""GQA decode attention (q_len == 1).
|
||||||
@@ -47,7 +60,7 @@ def attn_decode(
|
|||||||
_check_available("attn_decode")
|
_check_available("attn_decode")
|
||||||
causal_offset = (k.size(1) - 1) if is_causal else -1
|
causal_offset = (k.size(1) - 1) if is_causal else -1
|
||||||
return _modules["attn_decode"].attn_decode(
|
return _modules["attn_decode"].attn_decode(
|
||||||
q, k, v, mask=mask, causal_offset=causal_offset, layout=1
|
q, k, v, mask=mask, causal_offset=causal_offset, layout=TensorLayout.BLHD
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -55,7 +68,7 @@ def attn_prefill(
|
|||||||
q: torch.Tensor,
|
q: torch.Tensor,
|
||||||
k: torch.Tensor,
|
k: torch.Tensor,
|
||||||
v: torch.Tensor,
|
v: torch.Tensor,
|
||||||
mask: torch.Tensor | None = None,
|
mask: Optional[torch.Tensor] = None,
|
||||||
is_causal: bool = False,
|
is_causal: bool = False,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""GQA prefill attention (q_len > 1).
|
"""GQA prefill attention (q_len > 1).
|
||||||
@@ -73,45 +86,100 @@ def attn_prefill(
|
|||||||
_check_available("attn_prefill")
|
_check_available("attn_prefill")
|
||||||
causal_offset = (k.size(1) - q.size(1)) if is_causal else -1
|
causal_offset = (k.size(1) - q.size(1)) if is_causal else -1
|
||||||
return _modules["attn_prefill"].attn_prefill(
|
return _modules["attn_prefill"].attn_prefill(
|
||||||
q, k, v, mask=mask, causal_offset=causal_offset, layout=1
|
q, k, v, mask=mask, causal_offset=causal_offset, layout=TensorLayout.BLHD
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def attn_paged_decode(
|
def attn_paged_decode(
|
||||||
q: torch.Tensor,
|
q: torch.Tensor,
|
||||||
page_table: torch.Tensor,
|
|
||||||
k_cache: torch.Tensor,
|
k_cache: torch.Tensor,
|
||||||
v_cache: torch.Tensor,
|
v_cache: torch.Tensor,
|
||||||
page_size: int,
|
req_to_token: torch.Tensor,
|
||||||
kv_len: int,
|
req_pool_indices: torch.Tensor,
|
||||||
mask: torch.Tensor | None = None,
|
kv_indptr: torch.Tensor,
|
||||||
|
max_seq_len: int,
|
||||||
|
mask: Optional[torch.Tensor] = None,
|
||||||
is_causal: bool = False,
|
is_causal: bool = False,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""Paged GQA decode attention (q_len == 1, direct page-table access).
|
"""SGLang-style paged decode (q_len == 1, flat KV pool).
|
||||||
|
|
||||||
|
Reads K/V directly from a flat pool [size, kv_head, head_dim] via
|
||||||
|
req_to_token indirect indexing. Each request has its own seq_len
|
||||||
|
(from kv_indptr), eliminating padding waste.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
q: [batch, 1, n_heads, head_dim] (blhd, bf16)
|
q: [batch, n_heads, head_dim] (bf16, 3D — no seq dim)
|
||||||
page_table: [batch, max_pages] (int64)
|
k_cache: [pool_size, n_kv_heads, head_dim] (bf16, flat)
|
||||||
k_cache: [n_pages, page_size, n_kv_heads, head_dim] (bf16)
|
|
||||||
v_cache: same as k_cache
|
v_cache: same as k_cache
|
||||||
page_size: tokens per page
|
req_to_token: [num_reqs, max_context_len] (int64) — token -> slot
|
||||||
kv_len: actual sequence length per request
|
req_pool_indices: [batch] (int64) — rows into req_to_token
|
||||||
mask: 2D [batch, kv_len] or 3D [batch, 1, kv_len] (bool, True=keep)
|
kv_indptr: [batch+1] (int32) — prefix sum of per-request seq_lens
|
||||||
|
max_seq_len: max per-request seq_len (Python int, for split computation)
|
||||||
|
mask: 2D [batch, max_seq_len] (bool, True=keep) or None
|
||||||
is_causal: apply causal mask
|
is_causal: apply causal mask
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
[batch, 1, n_heads, head_dim] (blhd, bf16)
|
[batch, n_heads, head_dim] (bf16, 3D)
|
||||||
"""
|
"""
|
||||||
_check_available("attn_paged_decode")
|
_check_available("attn_paged_decode")
|
||||||
causal_offset = (kv_len - 1) if is_causal else -1
|
causal_offset = 0 if is_causal else -1
|
||||||
return _modules["attn_paged_decode"].attn_paged_decode(
|
return _modules["attn_paged_decode"].attn_paged_decode(
|
||||||
q,
|
q,
|
||||||
page_table,
|
|
||||||
k_cache,
|
k_cache,
|
||||||
v_cache,
|
v_cache,
|
||||||
page_size,
|
req_to_token,
|
||||||
kv_len,
|
req_pool_indices,
|
||||||
|
kv_indptr,
|
||||||
|
max_seq_len,
|
||||||
mask=mask,
|
mask=mask,
|
||||||
causal_offset=causal_offset,
|
causal_offset=causal_offset,
|
||||||
layout=1,
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def attn_paged_prefill(
|
||||||
|
q: torch.Tensor,
|
||||||
|
k_cache: torch.Tensor,
|
||||||
|
v_cache: torch.Tensor,
|
||||||
|
req_to_token: torch.Tensor,
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
kv_indptr: torch.Tensor,
|
||||||
|
qo_indptr: torch.Tensor,
|
||||||
|
mask: Optional[torch.Tensor] = None,
|
||||||
|
max_q_len: int = 0,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""SGLang-style paged prefill (ragged batch, flat KV pool).
|
||||||
|
|
||||||
|
Reads K/V directly from a flat pool [size, kv_head, head_dim] via
|
||||||
|
req_to_token. Supports ragged batches: each request has its own
|
||||||
|
q_len and kv_len, addressed via qo_indptr and kv_indptr.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
q: [total_q, n_heads, head_dim] (bf16, 3D — flattened across requests)
|
||||||
|
k_cache: [pool_size, n_kv_heads, head_dim] (bf16, flat)
|
||||||
|
v_cache: same as k_cache
|
||||||
|
req_to_token: [num_reqs, max_context_len] (int64)
|
||||||
|
req_pool_indices: [batch] (int64)
|
||||||
|
kv_indptr: [batch+1] (int32) — prefix sum of per-request kv_lens
|
||||||
|
qo_indptr: [batch+1] (int32) — prefix sum of per-request q_lens
|
||||||
|
mask: 4D [batch, 1, q_len, kv_len] (bool, True=keep) or None
|
||||||
|
max_q_len: max per-request q_len (Python int, for grid computation)
|
||||||
|
is_causal: apply causal mask
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[total_q, n_heads, head_dim] (bf16, 3D)
|
||||||
|
"""
|
||||||
|
_check_available("attn_paged_prefill")
|
||||||
|
causal_offset = 0 if is_causal else -1
|
||||||
|
return _modules["attn_paged_prefill"].attn_paged_prefill(
|
||||||
|
q,
|
||||||
|
k_cache,
|
||||||
|
v_cache,
|
||||||
|
req_to_token,
|
||||||
|
req_pool_indices,
|
||||||
|
kv_indptr,
|
||||||
|
qo_indptr,
|
||||||
|
mask,
|
||||||
|
max_q_len,
|
||||||
|
causal_offset=causal_offset,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -11,7 +11,13 @@ import logging
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
KERNEL_NAMES = ["attn_decode", "attn_prefill", "attn_paged_decode", "rotary_emb"]
|
KERNEL_NAMES = [
|
||||||
|
"attn_decode",
|
||||||
|
"attn_prefill",
|
||||||
|
"attn_paged_decode",
|
||||||
|
"attn_paged_prefill",
|
||||||
|
"rotary_emb",
|
||||||
|
]
|
||||||
|
|
||||||
_available: dict[str, bool] = {}
|
_available: dict[str, bool] = {}
|
||||||
_modules: dict[str, object] = {}
|
_modules: dict[str, object] = {}
|
||||||
|
|||||||
@@ -203,10 +203,8 @@ class KVCache:
|
|||||||
seq_lens: [batch_size] — per-request total sequence lengths
|
seq_lens: [batch_size] — per-request total sequence lengths
|
||||||
out_cache_loc: [batch, new_seq_len] or [batch, 1] — write indices
|
out_cache_loc: [batch, new_seq_len] or [batch, 1] — write indices
|
||||||
max_len: max(seq_lens) as Python int — avoids GPU sync in decode
|
max_len: max(seq_lens) as Python int — avoids GPU sync in decode
|
||||||
page_table: [batch, max_len] — precomputed gather indices for decode;
|
kv_indptr: [batch+1] int32 — prefix sum of seq_lens, precomputed once
|
||||||
None for prefill or when not yet computed.
|
per step so the attention backend avoids rebuilding it per layer.
|
||||||
decode_mask: [batch, max_len] bool — precomputed position validity
|
|
||||||
mask for decode; None for prefill or single-batch decode.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
k_buffer: Tensor
|
k_buffer: Tensor
|
||||||
@@ -216,8 +214,7 @@ class KVCache:
|
|||||||
seq_lens: Tensor
|
seq_lens: Tensor
|
||||||
out_cache_loc: Tensor
|
out_cache_loc: Tensor
|
||||||
max_len: int = 0
|
max_len: int = 0
|
||||||
page_table: Optional[Tensor] = None
|
kv_indptr: Optional[Tensor] = None
|
||||||
decode_mask: Optional[Tensor] = None
|
|
||||||
|
|
||||||
|
|
||||||
class PagePool:
|
class PagePool:
|
||||||
@@ -434,21 +431,14 @@ class PagePool:
|
|||||||
out_cache_loc = self._req_pool.req_to_token[
|
out_cache_loc = self._req_pool.req_to_token[
|
||||||
req_pool_indices, start_pos:seq_len
|
req_pool_indices, start_pos:seq_len
|
||||||
]
|
]
|
||||||
page_table = None
|
|
||||||
decode_mask = None
|
|
||||||
else:
|
else:
|
||||||
write_pos = seq_lens_t - 1
|
write_pos = seq_lens_t - 1
|
||||||
out_cache_loc = self._req_pool.req_to_token[
|
out_cache_loc = self._req_pool.req_to_token[
|
||||||
req_pool_indices, write_pos
|
req_pool_indices, write_pos
|
||||||
].unsqueeze(-1)
|
].unsqueeze(-1)
|
||||||
ml = max(seq_lens)
|
|
||||||
page_table = self._req_pool.req_to_token[req_pool_indices, :ml]
|
kv_indptr = torch.zeros(len(seq_lens) + 1, dtype=torch.int32, device=device)
|
||||||
if len(task_ids) > 1:
|
kv_indptr[1:] = seq_lens_t.cumsum(0).to(torch.int32)
|
||||||
decode_mask = (
|
|
||||||
torch.arange(ml, device=device)[None, :] < seq_lens_t[:, None]
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
decode_mask = None
|
|
||||||
|
|
||||||
return KVCache(
|
return KVCache(
|
||||||
k_buffer=self._storage.k_buffer,
|
k_buffer=self._storage.k_buffer,
|
||||||
@@ -458,8 +448,7 @@ class PagePool:
|
|||||||
seq_lens=seq_lens_t,
|
seq_lens=seq_lens_t,
|
||||||
out_cache_loc=out_cache_loc,
|
out_cache_loc=out_cache_loc,
|
||||||
max_len=max(seq_lens),
|
max_len=max(seq_lens),
|
||||||
page_table=page_table,
|
kv_indptr=kv_indptr,
|
||||||
decode_mask=decode_mask,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# ---- internals ----
|
# ---- internals ----
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ from astrai.model.components.lora import (
|
|||||||
merge_lora,
|
merge_lora,
|
||||||
save_lora,
|
save_lora,
|
||||||
)
|
)
|
||||||
from astrai.model.components.mlp import MLP
|
from astrai.model.components.mlp import MLP, DeepSeekMoE
|
||||||
from astrai.model.components.norm import RMSNorm
|
from astrai.model.components.norm import RMSNorm
|
||||||
from astrai.model.encoder import EmbeddingEncoder
|
from astrai.model.encoder import EmbeddingEncoder
|
||||||
from astrai.model.transformer import AutoRegressiveLM
|
from astrai.model.transformer import AutoRegressiveLM
|
||||||
@@ -19,6 +19,7 @@ __all__ = [
|
|||||||
"Linear",
|
"Linear",
|
||||||
"RMSNorm",
|
"RMSNorm",
|
||||||
"MLP",
|
"MLP",
|
||||||
|
"DeepSeekMoE",
|
||||||
"GQA",
|
"GQA",
|
||||||
"DecoderBlock",
|
"DecoderBlock",
|
||||||
# Models
|
# Models
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ from astrai.model.components.attention import GQA, MLA
|
|||||||
from astrai.model.components.decoder_block import DecoderBlock
|
from astrai.model.components.decoder_block import DecoderBlock
|
||||||
from astrai.model.components.embedding import Embedding
|
from astrai.model.components.embedding import Embedding
|
||||||
from astrai.model.components.linear import Linear
|
from astrai.model.components.linear import Linear
|
||||||
from astrai.model.components.mlp import MLP
|
from astrai.model.components.mlp import MLP, DeepSeekMoE
|
||||||
from astrai.model.components.norm import RMSNorm
|
from astrai.model.components.norm import RMSNorm
|
||||||
from astrai.model.components.rope import (
|
from astrai.model.components.rope import (
|
||||||
RotaryEmbedding,
|
RotaryEmbedding,
|
||||||
@@ -14,6 +14,7 @@ __all__ = [
|
|||||||
"Linear",
|
"Linear",
|
||||||
"RMSNorm",
|
"RMSNorm",
|
||||||
"MLP",
|
"MLP",
|
||||||
|
"DeepSeekMoE",
|
||||||
"Embedding",
|
"Embedding",
|
||||||
"GQA",
|
"GQA",
|
||||||
"MLA",
|
"MLA",
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
from dataclasses import asdict
|
from dataclasses import asdict
|
||||||
from typing import Optional
|
from typing import Optional, TypedDict
|
||||||
|
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
@@ -10,6 +10,11 @@ from astrai.model.components.mlp import FFNFactory
|
|||||||
from astrai.model.components.norm import RMSNorm
|
from astrai.model.components.norm import RMSNorm
|
||||||
|
|
||||||
|
|
||||||
|
class DecoderOutput(TypedDict):
|
||||||
|
hidden_states: Tensor
|
||||||
|
aux_loss: Optional[Tensor]
|
||||||
|
|
||||||
|
|
||||||
class DecoderBlock(nn.Module):
|
class DecoderBlock(nn.Module):
|
||||||
def __init__(self, config, layer_id: int):
|
def __init__(self, config, layer_id: int):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
@@ -26,7 +31,20 @@ class DecoderBlock(nn.Module):
|
|||||||
self.attention = AttnFactory.create(config.attn_type, **cfg, layer_id=layer_id)
|
self.attention = AttnFactory.create(config.attn_type, **cfg, layer_id=layer_id)
|
||||||
self.input_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
self.input_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
||||||
self.post_attention_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
self.post_attention_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
||||||
self.mlp = FFNFactory.create(config.ffn_type, **cfg)
|
ffn_type = self._resolve_ffn_type(config, layer_id)
|
||||||
|
self.mlp = FFNFactory.create(ffn_type, **cfg)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _resolve_ffn_type(config, layer_id: int) -> str:
|
||||||
|
if config.ffn_type != "moe":
|
||||||
|
return config.ffn_type
|
||||||
|
mlp_only = config.mlp_only_layers or []
|
||||||
|
if layer_id in mlp_only:
|
||||||
|
return "mlp"
|
||||||
|
if config.decoder_sparse_step > 1:
|
||||||
|
if (layer_id + 1) % config.decoder_sparse_step != 0:
|
||||||
|
return "mlp"
|
||||||
|
return "moe"
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
@@ -35,7 +53,7 @@ class DecoderBlock(nn.Module):
|
|||||||
attention_mask: Optional[Tensor] = None,
|
attention_mask: Optional[Tensor] = None,
|
||||||
kv_cache: Optional[KVCache] = None,
|
kv_cache: Optional[KVCache] = None,
|
||||||
is_causal: bool = False,
|
is_causal: bool = False,
|
||||||
) -> Tensor:
|
) -> DecoderOutput:
|
||||||
attn_output = self.attention(
|
attn_output = self.attention(
|
||||||
self.input_norm(x),
|
self.input_norm(x),
|
||||||
rotary_emb,
|
rotary_emb,
|
||||||
@@ -44,6 +62,8 @@ class DecoderBlock(nn.Module):
|
|||||||
is_causal,
|
is_causal,
|
||||||
)
|
)
|
||||||
x = attn_output + x
|
x = attn_output + x
|
||||||
x = self.mlp(self.post_attention_norm(x)) + x
|
normalized = self.post_attention_norm(x)
|
||||||
|
mlp_output = self.mlp(normalized)
|
||||||
|
x = mlp_output["hidden_states"] + x
|
||||||
|
|
||||||
return x
|
return {"hidden_states": x, "aux_loss": mlp_output["aux_loss"]}
|
||||||
|
|||||||
@@ -1,3 +1,5 @@
|
|||||||
|
from typing import Optional, TypedDict
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
@@ -11,6 +13,16 @@ class FFNFactory(BaseFactory[nn.Module]):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class FFNOutput(TypedDict):
|
||||||
|
hidden_states: Tensor
|
||||||
|
aux_loss: Optional[Tensor]
|
||||||
|
|
||||||
|
|
||||||
|
class RoutedOutput(TypedDict):
|
||||||
|
hidden_states: Tensor
|
||||||
|
aux_loss: Optional[Tensor]
|
||||||
|
|
||||||
|
|
||||||
@FFNFactory.register("mlp")
|
@FFNFactory.register("mlp")
|
||||||
class MLP(nn.Module):
|
class MLP(nn.Module):
|
||||||
def __init__(self, dim: int, dim_ffn: int, down_init_std: float = 0.02):
|
def __init__(self, dim: int, dim_ffn: int, down_init_std: float = 0.02):
|
||||||
@@ -19,10 +31,10 @@ class MLP(nn.Module):
|
|||||||
self.gate = Linear(dim, dim_ffn)
|
self.gate = Linear(dim, dim_ffn)
|
||||||
self.down = Linear(dim_ffn, dim, init_std=down_init_std)
|
self.down = Linear(dim_ffn, dim, init_std=down_init_std)
|
||||||
|
|
||||||
def forward(self, x: Tensor) -> Tensor:
|
def forward(self, x: Tensor) -> FFNOutput:
|
||||||
gated = self.up(x) * F.silu(self.gate(x))
|
gated = self.up(x) * F.silu(self.gate(x))
|
||||||
out = self.down(gated)
|
out = self.down(gated)
|
||||||
return out
|
return {"hidden_states": out, "aux_loss": None}
|
||||||
|
|
||||||
|
|
||||||
@FFNFactory.register("moe")
|
@FFNFactory.register("moe")
|
||||||
@@ -36,6 +48,9 @@ class DeepSeekMoE(nn.Module):
|
|||||||
n_activated_experts: int = 2,
|
n_activated_experts: int = 2,
|
||||||
topk_method: str = "greedy",
|
topk_method: str = "greedy",
|
||||||
n_layers: int = 1,
|
n_layers: int = 1,
|
||||||
|
moe_intermediate_size: Optional[int] = None,
|
||||||
|
shared_expert_intermediate_size: Optional[int] = None,
|
||||||
|
norm_topk_prob: bool = True,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.dim = dim
|
self.dim = dim
|
||||||
@@ -43,6 +58,16 @@ class DeepSeekMoE(nn.Module):
|
|||||||
self.n_shared_experts = n_shared_experts
|
self.n_shared_experts = n_shared_experts
|
||||||
self.n_activated_experts = n_activated_experts
|
self.n_activated_experts = n_activated_experts
|
||||||
self.topk_method = topk_method
|
self.topk_method = topk_method
|
||||||
|
self.norm_topk_prob = norm_topk_prob
|
||||||
|
|
||||||
|
expert_dim_ffn = (
|
||||||
|
moe_intermediate_size if moe_intermediate_size is not None else dim_ffn
|
||||||
|
)
|
||||||
|
shared_dim_ffn = (
|
||||||
|
shared_expert_intermediate_size
|
||||||
|
if shared_expert_intermediate_size is not None
|
||||||
|
else dim_ffn
|
||||||
|
)
|
||||||
|
|
||||||
self.router = Linear(dim, n_routed_experts, bias=False)
|
self.router = Linear(dim, n_routed_experts, bias=False)
|
||||||
moe_scale = 1 / max(n_shared_experts, 1) + 1 / n_activated_experts
|
moe_scale = 1 / max(n_shared_experts, 1) + 1 / n_activated_experts
|
||||||
@@ -50,33 +75,37 @@ class DeepSeekMoE(nn.Module):
|
|||||||
|
|
||||||
self.shared_experts = nn.ModuleList(
|
self.shared_experts = nn.ModuleList(
|
||||||
[
|
[
|
||||||
MLP(dim, dim_ffn, down_init_std=down_init_std)
|
MLP(dim, shared_dim_ffn, down_init_std=down_init_std)
|
||||||
for _ in range(n_shared_experts)
|
for _ in range(n_shared_experts)
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
self.routed_experts = nn.ModuleList(
|
self.routed_experts = nn.ModuleList(
|
||||||
[
|
[
|
||||||
MLP(dim, dim_ffn, down_init_std=down_init_std)
|
MLP(dim, expert_dim_ffn, down_init_std=down_init_std)
|
||||||
for _ in range(n_routed_experts)
|
for _ in range(n_routed_experts)
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
def forward(self, x: Tensor) -> Tensor:
|
def forward(self, x: Tensor) -> FFNOutput:
|
||||||
|
include_aux_loss = self.training and torch.is_grad_enabled()
|
||||||
bsz, seq_len, dim = x.shape
|
bsz, seq_len, dim = x.shape
|
||||||
x_flat = x.view(-1, dim)
|
x_flat = x.view(-1, dim)
|
||||||
|
|
||||||
shared_out = self._shared_forward(x_flat)
|
shared_out = self._shared_forward(x_flat)
|
||||||
routed_out = self._routed_forward(x_flat)
|
routed_output = self._routed_forward(x_flat, include_aux_loss)
|
||||||
|
|
||||||
out = (shared_out + routed_out).view(bsz, seq_len, dim)
|
out = (shared_out + routed_output["hidden_states"]).view(bsz, seq_len, dim)
|
||||||
return out
|
return {"hidden_states": out, "aux_loss": routed_output["aux_loss"]}
|
||||||
|
|
||||||
def _shared_forward(self, x: Tensor) -> Tensor:
|
def _shared_forward(self, x: Tensor) -> Tensor:
|
||||||
if self.n_shared_experts == 0:
|
if self.n_shared_experts == 0:
|
||||||
return torch.zeros_like(x)
|
return torch.zeros_like(x)
|
||||||
return sum(e(x) for e in self.shared_experts) / self.n_shared_experts
|
return (
|
||||||
|
sum(e(x)["hidden_states"] for e in self.shared_experts)
|
||||||
|
/ self.n_shared_experts
|
||||||
|
)
|
||||||
|
|
||||||
def _routed_forward(self, x: Tensor) -> Tensor:
|
def _routed_forward(self, x: Tensor, include_aux_loss: bool) -> RoutedOutput:
|
||||||
N, D = x.shape
|
N, D = x.shape
|
||||||
K = self.n_activated_experts
|
K = self.n_activated_experts
|
||||||
|
|
||||||
@@ -84,17 +113,29 @@ class DeepSeekMoE(nn.Module):
|
|||||||
router_probs = torch.softmax(router_logits.float(), dim=-1).to(x.dtype)
|
router_probs = torch.softmax(router_logits.float(), dim=-1).to(x.dtype)
|
||||||
|
|
||||||
topk_weights, topk_indices = torch.topk(router_probs, K, dim=-1)
|
topk_weights, topk_indices = torch.topk(router_probs, K, dim=-1)
|
||||||
|
if self.norm_topk_prob:
|
||||||
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
|
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
|
||||||
|
|
||||||
|
aux_loss = None
|
||||||
|
if include_aux_loss:
|
||||||
|
expert_load = F.one_hot(
|
||||||
|
topk_indices, num_classes=self.n_routed_experts
|
||||||
|
).float()
|
||||||
|
expert_load = expert_load.mean(dim=(0, 1))
|
||||||
|
router_prob = router_probs.float().mean(dim=0)
|
||||||
|
aux_loss = self.n_routed_experts * (expert_load * router_prob).sum()
|
||||||
|
|
||||||
output = torch.zeros(N, D, device=x.device, dtype=x.dtype)
|
output = torch.zeros(N, D, device=x.device, dtype=x.dtype)
|
||||||
for expert_idx in range(self.n_routed_experts):
|
for expert_idx in range(self.n_routed_experts):
|
||||||
expert_mask = topk_indices == expert_idx
|
expert_mask = topk_indices == expert_idx
|
||||||
token_idx, k_idx = expert_mask.nonzero(as_tuple=True)
|
token_idx, k_idx = expert_mask.nonzero(as_tuple=True)
|
||||||
if token_idx.numel() == 0:
|
if token_idx.numel() == 0:
|
||||||
continue
|
continue
|
||||||
|
expert = self.routed_experts[expert_idx]
|
||||||
expert_input = x[token_idx]
|
expert_input = x[token_idx]
|
||||||
expert_output = self.routed_experts[expert_idx](expert_input)
|
expert_output = expert(expert_input)["hidden_states"]
|
||||||
|
|
||||||
weights = topk_weights[token_idx, k_idx].unsqueeze(-1)
|
weights = topk_weights[token_idx, k_idx].unsqueeze(-1)
|
||||||
output.index_add_(0, token_idx, expert_output * weights)
|
output.index_add_(0, token_idx, expert_output * weights)
|
||||||
|
|
||||||
return output
|
return {"hidden_states": output, "aux_loss": aux_loss}
|
||||||
|
|||||||
@@ -70,7 +70,7 @@ class EmbeddingEncoder(AutoModel):
|
|||||||
attn_mask = process_attention_mask(input_mask)
|
attn_mask = process_attention_mask(input_mask)
|
||||||
|
|
||||||
for layer in self.layers:
|
for layer in self.layers:
|
||||||
x = layer(x, rotary_emb, attn_mask)
|
x = layer(x, rotary_emb, attn_mask)["hidden_states"]
|
||||||
|
|
||||||
hidden_states = self.norm(x)
|
hidden_states = self.norm(x)
|
||||||
|
|
||||||
|
|||||||
@@ -113,10 +113,23 @@ class AutoRegressiveLM(AutoModel):
|
|||||||
attn_mask = process_attention_mask(input_mask)
|
attn_mask = process_attention_mask(input_mask)
|
||||||
use_sdpa_causal_mask = attn_mask is None
|
use_sdpa_causal_mask = attn_mask is None
|
||||||
|
|
||||||
|
aux_losses = []
|
||||||
for layer in self.layers:
|
for layer in self.layers:
|
||||||
x = layer(x, rotary_emb, attn_mask, kv_cache, use_sdpa_causal_mask)
|
layer_output = layer(
|
||||||
|
x,
|
||||||
|
rotary_emb,
|
||||||
|
attn_mask,
|
||||||
|
kv_cache,
|
||||||
|
use_sdpa_causal_mask,
|
||||||
|
)
|
||||||
|
x = layer_output["hidden_states"]
|
||||||
|
if layer_output["aux_loss"] is not None:
|
||||||
|
aux_losses.append(layer_output["aux_loss"])
|
||||||
|
|
||||||
hidden_states = self.norm(x)
|
hidden_states = self.norm(x)
|
||||||
logits = self.lm_head(hidden_states)
|
logits = self.lm_head(hidden_states)
|
||||||
|
|
||||||
return {"logits": logits, "hidden_states": hidden_states}
|
output = {"logits": logits, "hidden_states": hidden_states}
|
||||||
|
if aux_losses:
|
||||||
|
output["aux_loss"] = torch.stack(aux_losses).mean()
|
||||||
|
return output
|
||||||
|
|||||||
+95
-28
@@ -1,7 +1,7 @@
|
|||||||
"""Training strategy implementations with factory pattern."""
|
"""Training strategy implementations with factory pattern."""
|
||||||
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from typing import Callable, Dict, Union
|
from typing import Callable, Dict, Optional, TypedDict, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
@@ -13,6 +13,16 @@ from astrai.parallel.executor import broadcast_state_dict
|
|||||||
from astrai.trainer.rollout import RolloutResult
|
from astrai.trainer.rollout import RolloutResult
|
||||||
|
|
||||||
|
|
||||||
|
class LossOutput(TypedDict):
|
||||||
|
loss: Tensor
|
||||||
|
metrics: Dict[str, float]
|
||||||
|
|
||||||
|
|
||||||
|
class LogprobsOutput(TypedDict):
|
||||||
|
logprobs: Tensor
|
||||||
|
aux_loss: Optional[Tensor]
|
||||||
|
|
||||||
|
|
||||||
def move_to_device(batch: Dict[str, Tensor], device: str) -> Dict[str, Tensor]:
|
def move_to_device(batch: Dict[str, Tensor], device: str) -> Dict[str, Tensor]:
|
||||||
"""Move batch tensors to specified device with non-blocking transfer."""
|
"""Move batch tensors to specified device with non-blocking transfer."""
|
||||||
return {key: value.to(device, non_blocking=True) for key, value in batch.items()}
|
return {key: value.to(device, non_blocking=True) for key, value in batch.items()}
|
||||||
@@ -24,7 +34,7 @@ def get_logprobs(
|
|||||||
attn_mask: Tensor,
|
attn_mask: Tensor,
|
||||||
loss_mask: Tensor,
|
loss_mask: Tensor,
|
||||||
reduction: str,
|
reduction: str,
|
||||||
) -> Tensor:
|
) -> LogprobsOutput:
|
||||||
"""Compute token-wise log probabilities from model outputs.
|
"""Compute token-wise log probabilities from model outputs.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -46,10 +56,11 @@ def get_logprobs(
|
|||||||
shifted_input_ids = input_ids[:, 1:]
|
shifted_input_ids = input_ids[:, 1:]
|
||||||
shifted_loss_mask = loss_mask[:, 1:]
|
shifted_loss_mask = loss_mask[:, 1:]
|
||||||
|
|
||||||
logits = model(
|
outputs = model(
|
||||||
input_ids[:, :-1],
|
input_ids[:, :-1],
|
||||||
attn_mask[:, :, :-1, :-1] if attn_mask.dim() == 4 else attn_mask[:, :-1],
|
attn_mask[:, :, :-1, :-1] if attn_mask.dim() == 4 else attn_mask[:, :-1],
|
||||||
)["logits"]
|
)
|
||||||
|
logits = outputs["logits"]
|
||||||
log_probs = torch.log_softmax(logits.float(), dim=-1)
|
log_probs = torch.log_softmax(logits.float(), dim=-1)
|
||||||
|
|
||||||
token_logprobs = torch.gather(
|
token_logprobs = torch.gather(
|
||||||
@@ -57,13 +68,14 @@ def get_logprobs(
|
|||||||
).squeeze(-1)
|
).squeeze(-1)
|
||||||
|
|
||||||
if reduction == "mean":
|
if reduction == "mean":
|
||||||
return (token_logprobs * shifted_loss_mask).sum(dim=-1) / shifted_loss_mask.sum(
|
logprobs = (token_logprobs * shifted_loss_mask).sum(
|
||||||
dim=-1
|
dim=-1
|
||||||
).clamp(min=1.0)
|
) / shifted_loss_mask.sum(dim=-1).clamp(min=1.0)
|
||||||
elif reduction == "sum":
|
elif reduction == "sum":
|
||||||
return (token_logprobs * shifted_loss_mask).sum(dim=-1)
|
logprobs = (token_logprobs * shifted_loss_mask).sum(dim=-1)
|
||||||
else:
|
else:
|
||||||
return token_logprobs * shifted_loss_mask
|
logprobs = token_logprobs * shifted_loss_mask
|
||||||
|
return {"logprobs": logprobs, "aux_loss": outputs.get("aux_loss")}
|
||||||
|
|
||||||
|
|
||||||
def make_doc_boundary_mask(position_ids: Tensor) -> Tensor:
|
def make_doc_boundary_mask(position_ids: Tensor) -> Tensor:
|
||||||
@@ -102,6 +114,7 @@ class BaseStrategy(ABC):
|
|||||||
self.model = model
|
self.model = model
|
||||||
self.device = device
|
self.device = device
|
||||||
self.executor = kwargs.pop("executor", None)
|
self.executor = kwargs.pop("executor", None)
|
||||||
|
self.moe_aux_loss_coef = kwargs.pop("moe_aux_loss_coef", 0.01)
|
||||||
self.extra_kwargs = kwargs
|
self.extra_kwargs = kwargs
|
||||||
self._rollout_runner = None
|
self._rollout_runner = None
|
||||||
|
|
||||||
@@ -117,6 +130,33 @@ class BaseStrategy(ABC):
|
|||||||
"""
|
"""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
|
||||||
|
return self._normalize_output(self.compute_loss(batch))
|
||||||
|
|
||||||
|
def _loss_output(
|
||||||
|
self,
|
||||||
|
task_loss: Tensor,
|
||||||
|
metrics: Dict[str, Tensor],
|
||||||
|
aux_loss: Optional[Tensor] = None,
|
||||||
|
) -> LossOutput:
|
||||||
|
total_loss = task_loss
|
||||||
|
if aux_loss is not None:
|
||||||
|
weighted_aux_loss = self.moe_aux_loss_coef * aux_loss
|
||||||
|
total_loss = total_loss + weighted_aux_loss
|
||||||
|
metrics["moe_aux_loss"] = aux_loss
|
||||||
|
metrics["moe_aux_loss_weighted"] = weighted_aux_loss
|
||||||
|
metrics["loss"] = total_loss
|
||||||
|
return {
|
||||||
|
"loss": total_loss,
|
||||||
|
"metrics": {name: value.detach().item() for name, value in metrics.items()},
|
||||||
|
}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _normalize_output(output: Union[LossOutput, Tensor]) -> LossOutput:
|
||||||
|
if isinstance(output, dict):
|
||||||
|
return output
|
||||||
|
return {"loss": output, "metrics": {"loss": output.detach().item()}}
|
||||||
|
|
||||||
def supports_online(self) -> bool:
|
def supports_online(self) -> bool:
|
||||||
"""Whether this strategy can operate with a rollout runner.
|
"""Whether this strategy can operate with a rollout runner.
|
||||||
|
|
||||||
@@ -153,17 +193,17 @@ class BaseStrategy(ABC):
|
|||||||
if self._rollout_runner is not None:
|
if self._rollout_runner is not None:
|
||||||
self._rollout_runner.step()
|
self._rollout_runner.step()
|
||||||
|
|
||||||
def __call__(self, batch: Dict[str, Tensor]) -> Tensor:
|
def __call__(self, batch: Dict[str, Tensor]) -> LossOutput:
|
||||||
"""Run offline or online forward depending on runner injection."""
|
"""Run offline or online forward depending on runner injection."""
|
||||||
if self._rollout_runner is None:
|
if self._rollout_runner is None:
|
||||||
return self.compute_loss(batch)
|
return self.compute_loss_output(batch)
|
||||||
|
|
||||||
result, is_fresh = self._rollout_runner(batch)
|
result, is_fresh = self._rollout_runner(batch)
|
||||||
if is_fresh:
|
if is_fresh:
|
||||||
self._on_rollout_refresh()
|
self._on_rollout_refresh()
|
||||||
|
|
||||||
train_batch = self.prepare_from_rollout(result)
|
train_batch = self.prepare_from_rollout(result)
|
||||||
return self.compute_loss(train_batch)
|
return self.compute_loss_output(train_batch)
|
||||||
|
|
||||||
|
|
||||||
class StrategyFactory(BaseFactory["BaseStrategy"]):
|
class StrategyFactory(BaseFactory["BaseStrategy"]):
|
||||||
@@ -203,9 +243,13 @@ class SEQStrategy(BaseStrategy):
|
|||||||
self.label_smoothing = label_smoothing
|
self.label_smoothing = label_smoothing
|
||||||
|
|
||||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||||
|
return self.compute_loss_output(batch)["loss"]
|
||||||
|
|
||||||
|
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
|
||||||
batch = move_to_device(batch, self.device)
|
batch = move_to_device(batch, self.device)
|
||||||
input_ids, target_ids = batch["input_ids"], batch["target_ids"]
|
input_ids, target_ids = batch["input_ids"], batch["target_ids"]
|
||||||
logits = self.model(input_ids=input_ids)["logits"]
|
outputs = self.model(input_ids=input_ids)
|
||||||
|
logits = outputs["logits"]
|
||||||
|
|
||||||
loss = F.cross_entropy(
|
loss = F.cross_entropy(
|
||||||
input=logits.flatten(0, 1).float(),
|
input=logits.flatten(0, 1).float(),
|
||||||
@@ -213,7 +257,7 @@ class SEQStrategy(BaseStrategy):
|
|||||||
label_smoothing=self.label_smoothing,
|
label_smoothing=self.label_smoothing,
|
||||||
)
|
)
|
||||||
|
|
||||||
return loss
|
return self._loss_output(loss, {"task_loss": loss}, outputs.get("aux_loss"))
|
||||||
|
|
||||||
|
|
||||||
@StrategyFactory.register("sft")
|
@StrategyFactory.register("sft")
|
||||||
@@ -234,6 +278,9 @@ class SFTStrategy(BaseStrategy):
|
|||||||
self.label_smoothing = label_smoothing
|
self.label_smoothing = label_smoothing
|
||||||
|
|
||||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||||
|
return self.compute_loss_output(batch)["loss"]
|
||||||
|
|
||||||
|
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
|
||||||
batch = move_to_device(batch, self.device)
|
batch = move_to_device(batch, self.device)
|
||||||
input_ids, target_ids, position_ids, loss_mask = (
|
input_ids, target_ids, position_ids, loss_mask = (
|
||||||
batch["input_ids"],
|
batch["input_ids"],
|
||||||
@@ -245,9 +292,10 @@ class SFTStrategy(BaseStrategy):
|
|||||||
ignore_index = -100
|
ignore_index = -100
|
||||||
input_mask = make_doc_boundary_mask(position_ids)
|
input_mask = make_doc_boundary_mask(position_ids)
|
||||||
target_ids = target_ids.masked_fill(~loss_mask, ignore_index)
|
target_ids = target_ids.masked_fill(~loss_mask, ignore_index)
|
||||||
logits = self.model(
|
outputs = self.model(
|
||||||
input_ids=input_ids, position_ids=position_ids, input_mask=input_mask
|
input_ids=input_ids, position_ids=position_ids, input_mask=input_mask
|
||||||
)["logits"]
|
)
|
||||||
|
logits = outputs["logits"]
|
||||||
|
|
||||||
loss = F.cross_entropy(
|
loss = F.cross_entropy(
|
||||||
input=logits.flatten(0, 1).float(),
|
input=logits.flatten(0, 1).float(),
|
||||||
@@ -256,7 +304,7 @@ class SFTStrategy(BaseStrategy):
|
|||||||
label_smoothing=self.label_smoothing,
|
label_smoothing=self.label_smoothing,
|
||||||
)
|
)
|
||||||
|
|
||||||
return loss
|
return self._loss_output(loss, {"task_loss": loss}, outputs.get("aux_loss"))
|
||||||
|
|
||||||
|
|
||||||
@StrategyFactory.register("dpo")
|
@StrategyFactory.register("dpo")
|
||||||
@@ -282,6 +330,9 @@ class DPOStrategy(BaseStrategy):
|
|||||||
self.reduction = reduction
|
self.reduction = reduction
|
||||||
|
|
||||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||||
|
return self.compute_loss_output(batch)["loss"]
|
||||||
|
|
||||||
|
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
|
||||||
batch = move_to_device(batch, self.device)
|
batch = move_to_device(batch, self.device)
|
||||||
chosen_ids, rejected_ids = batch["chosen"], batch["rejected"]
|
chosen_ids, rejected_ids = batch["chosen"], batch["rejected"]
|
||||||
chosen_mask, rejected_mask = batch["chosen_mask"], batch["rejected_mask"]
|
chosen_mask, rejected_mask = batch["chosen_mask"], batch["rejected_mask"]
|
||||||
@@ -297,22 +348,25 @@ class DPOStrategy(BaseStrategy):
|
|||||||
)[None, None, :, :] # [1, 1, S, S]
|
)[None, None, :, :] # [1, 1, S, S]
|
||||||
full_mask = key_pad & causal # [B*2, 1, S, S] — composed
|
full_mask = key_pad & causal # [B*2, 1, S, S] — composed
|
||||||
|
|
||||||
log_pi = get_logprobs(
|
policy_output = get_logprobs(
|
||||||
self.model,
|
self.model,
|
||||||
concat_ids,
|
concat_ids,
|
||||||
full_mask,
|
full_mask,
|
||||||
concat_loss_mask,
|
concat_loss_mask,
|
||||||
self.reduction,
|
self.reduction,
|
||||||
)
|
)
|
||||||
|
log_pi = policy_output["logprobs"]
|
||||||
|
aux_loss = policy_output["aux_loss"]
|
||||||
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
log_ref = get_logprobs(
|
ref_output = get_logprobs(
|
||||||
self.ref_model,
|
self.ref_model,
|
||||||
concat_ids,
|
concat_ids,
|
||||||
full_mask,
|
full_mask,
|
||||||
concat_loss_mask,
|
concat_loss_mask,
|
||||||
self.reduction,
|
self.reduction,
|
||||||
)
|
)
|
||||||
|
log_ref = ref_output["logprobs"]
|
||||||
|
|
||||||
log_pi_chosen = log_pi[: chosen_ids.shape[0]]
|
log_pi_chosen = log_pi[: chosen_ids.shape[0]]
|
||||||
log_pi_rejected = log_pi[chosen_ids.shape[0] :]
|
log_pi_rejected = log_pi[chosen_ids.shape[0] :]
|
||||||
@@ -325,7 +379,7 @@ class DPOStrategy(BaseStrategy):
|
|||||||
ratio_diff = pi_log_ratio - ref_log_ratio
|
ratio_diff = pi_log_ratio - ref_log_ratio
|
||||||
dpo_loss = -F.logsigmoid(self.beta * ratio_diff).mean()
|
dpo_loss = -F.logsigmoid(self.beta * ratio_diff).mean()
|
||||||
|
|
||||||
return dpo_loss
|
return self._loss_output(dpo_loss, {"dpo_loss": dpo_loss}, aux_loss)
|
||||||
|
|
||||||
def supports_online(self) -> bool:
|
def supports_online(self) -> bool:
|
||||||
return True
|
return True
|
||||||
@@ -398,6 +452,9 @@ class GRPOStrategy(BaseStrategy):
|
|||||||
self.old_model.load_state_dict(state_dict)
|
self.old_model.load_state_dict(state_dict)
|
||||||
|
|
||||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||||
|
return self.compute_loss_output(batch)["loss"]
|
||||||
|
|
||||||
|
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
|
||||||
batch = move_to_device(batch, self.device)
|
batch = move_to_device(batch, self.device)
|
||||||
prompts = batch["prompts"]
|
prompts = batch["prompts"]
|
||||||
responses = batch["responses"]
|
responses = batch["responses"]
|
||||||
@@ -438,16 +495,23 @@ class GRPOStrategy(BaseStrategy):
|
|||||||
# get_logprobs returns [B*G, S-1] (S = prompt_len + response_len).
|
# get_logprobs returns [B*G, S-1] (S = prompt_len + response_len).
|
||||||
# Response token logprobs occupy the last ``response_len`` positions
|
# Response token logprobs occupy the last ``response_len`` positions
|
||||||
# (the first response token is predicted from the last prompt token).
|
# (the first response token is predicted from the last prompt token).
|
||||||
token_log_probs_policy = get_logprobs(
|
policy_output = get_logprobs(
|
||||||
self.model, full_sequences, attn_mask, full_masks, "none"
|
self.model, full_sequences, attn_mask, full_masks, "none"
|
||||||
)[:, prompt_len - 1 :]
|
)
|
||||||
|
token_log_probs_policy = policy_output["logprobs"]
|
||||||
|
aux_loss = policy_output["aux_loss"]
|
||||||
|
token_log_probs_policy = token_log_probs_policy[:, prompt_len - 1 :]
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
token_log_probs_old = get_logprobs(
|
old_output = get_logprobs(
|
||||||
self.old_model, full_sequences, attn_mask, full_masks, "none"
|
self.old_model, full_sequences, attn_mask, full_masks, "none"
|
||||||
)[:, prompt_len - 1 :]
|
)
|
||||||
token_log_probs_ref = get_logprobs(
|
token_log_probs_old = old_output["logprobs"]
|
||||||
|
token_log_probs_old = token_log_probs_old[:, prompt_len - 1 :]
|
||||||
|
ref_output = get_logprobs(
|
||||||
self.ref_model, full_sequences, attn_mask, full_masks, "none"
|
self.ref_model, full_sequences, attn_mask, full_masks, "none"
|
||||||
)[:, prompt_len - 1 :]
|
)
|
||||||
|
token_log_probs_ref = ref_output["logprobs"]
|
||||||
|
token_log_probs_ref = token_log_probs_ref[:, prompt_len - 1 :]
|
||||||
|
|
||||||
# Reshape to [B, G, response_len]
|
# Reshape to [B, G, response_len]
|
||||||
token_log_probs_policy = token_log_probs_policy.view(batch_size, group_size, -1)
|
token_log_probs_policy = token_log_probs_policy.view(batch_size, group_size, -1)
|
||||||
@@ -480,9 +544,12 @@ class GRPOStrategy(BaseStrategy):
|
|||||||
kl_per_token = r - torch.log(r + eps) - 1.0
|
kl_per_token = r - torch.log(r + eps) - 1.0
|
||||||
kl_penalty = self.kl_coef * (kl_per_token * token_masks).sum() / token_count
|
kl_penalty = self.kl_coef * (kl_per_token * token_masks).sum() / token_count
|
||||||
|
|
||||||
total_loss = policy_loss + kl_penalty
|
task_loss = policy_loss + kl_penalty
|
||||||
|
return self._loss_output(
|
||||||
return total_loss
|
task_loss,
|
||||||
|
{"policy_loss": policy_loss, "kl_loss": kl_penalty},
|
||||||
|
aux_loss,
|
||||||
|
)
|
||||||
|
|
||||||
def supports_online(self) -> bool:
|
def supports_online(self) -> bool:
|
||||||
return True
|
return True
|
||||||
|
|||||||
@@ -260,11 +260,28 @@ class MetricCallback(TrainCallback):
|
|||||||
}
|
}
|
||||||
|
|
||||||
def _metrics(self, context: TrainContext, names):
|
def _metrics(self, context: TrainContext, names):
|
||||||
return {
|
metrics = dict(context.metrics)
|
||||||
m: self._metric_funcs[m](context)
|
for name in names:
|
||||||
for m in names
|
metric_fn = self._metric_funcs.get(name)
|
||||||
if self._metric_funcs[m](context) is not None
|
if metric_fn is None:
|
||||||
}
|
continue
|
||||||
|
value = metric_fn(context)
|
||||||
|
if value is not None:
|
||||||
|
metrics[name] = value
|
||||||
|
selected = set(context.metrics) | set(names)
|
||||||
|
selected.discard("*")
|
||||||
|
result = {name: metrics[name] for name in selected if name in metrics}
|
||||||
|
if context.world_size > 1 and dist.is_initialized() and result:
|
||||||
|
metric_names = sorted(result)
|
||||||
|
values = torch.tensor(
|
||||||
|
[result[name] for name in metric_names],
|
||||||
|
dtype=torch.float32,
|
||||||
|
device=get_current_device(),
|
||||||
|
)
|
||||||
|
dist.all_reduce(values, op=dist.ReduceOp.SUM)
|
||||||
|
values /= context.world_size
|
||||||
|
result.update(zip(metric_names, values.tolist()))
|
||||||
|
return result
|
||||||
|
|
||||||
@only_on_rank(0)
|
@only_on_rank(0)
|
||||||
def _append(self, event_type: str, context: TrainContext, **extra):
|
def _append(self, event_type: str, context: TrainContext, **extra):
|
||||||
@@ -286,8 +303,8 @@ class MetricCallback(TrainCallback):
|
|||||||
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
for batch in context.val_dataloader:
|
for batch in context.val_dataloader:
|
||||||
loss = context.strategy(batch)
|
loss_output = context.strategy(batch)
|
||||||
total_loss += loss.item()
|
total_loss += loss_output["loss"].item()
|
||||||
num_batches += 1
|
num_batches += 1
|
||||||
|
|
||||||
if context.world_size > 1 and dist.is_initialized():
|
if context.world_size > 1 and dist.is_initialized():
|
||||||
|
|||||||
@@ -38,6 +38,7 @@ class TrainContext:
|
|||||||
epoch: int = field(default=0)
|
epoch: int = field(default=0)
|
||||||
consumed_samples: int = field(default=0)
|
consumed_samples: int = field(default=0)
|
||||||
loss: float = field(default=0.0)
|
loss: float = field(default=0.0)
|
||||||
|
metrics: Dict[str, float] = field(default_factory=dict)
|
||||||
grad_norm: Optional[float] = field(default=None)
|
grad_norm: Optional[float] = field(default=None)
|
||||||
grad_snr_tracker: GradSNRTracker = field(default_factory=GradSNRTracker)
|
grad_snr_tracker: GradSNRTracker = field(default_factory=GradSNRTracker)
|
||||||
val_dataloader: Optional[DataLoader] = field(default=None)
|
val_dataloader: Optional[DataLoader] = field(default=None)
|
||||||
@@ -221,6 +222,7 @@ class TrainContextBuilder:
|
|||||||
obj.load_state_dict(extra[name])
|
obj.load_state_dict(extra[name])
|
||||||
|
|
||||||
strategy_kwargs = dict(cfg.extra_kwargs)
|
strategy_kwargs = dict(cfg.extra_kwargs)
|
||||||
|
strategy_kwargs.setdefault("moe_aux_loss_coef", cfg.moe_aux_loss_coef)
|
||||||
|
|
||||||
needs_ref = cfg.strategy in (
|
needs_ref = cfg.strategy in (
|
||||||
"dpo",
|
"dpo",
|
||||||
|
|||||||
@@ -82,9 +82,10 @@ class Trainer:
|
|||||||
break
|
break
|
||||||
with executor.accumulate(context.model):
|
with executor.accumulate(context.model):
|
||||||
self._call_callbacks("on_batch_begin", context)
|
self._call_callbacks("on_batch_begin", context)
|
||||||
loss = context.strategy(batch)
|
loss_output = context.strategy(batch)
|
||||||
context.loss = loss.item()
|
context.loss = loss_output["loss"].item()
|
||||||
stand_loss = loss / executor.grad_accum_steps
|
context.metrics = loss_output["metrics"]
|
||||||
|
stand_loss = loss_output["loss"] / executor.grad_accum_steps
|
||||||
executor.backward(stand_loss)
|
executor.backward(stand_loss)
|
||||||
context.consumed_samples += (
|
context.consumed_samples += (
|
||||||
context.config.batch_per_device * context.world_size
|
context.config.batch_per_device * context.world_size
|
||||||
|
|||||||
@@ -72,4 +72,5 @@ def register(name: str, sources: list[str] | None = None, **kwargs):
|
|||||||
register("attn_decode")
|
register("attn_decode")
|
||||||
register("attn_prefill")
|
register("attn_prefill")
|
||||||
register("attn_paged_decode")
|
register("attn_paged_decode")
|
||||||
|
register("attn_paged_prefill")
|
||||||
register("rotary_emb")
|
register("rotary_emb")
|
||||||
|
|||||||
+42
-14
@@ -1,5 +1,13 @@
|
|||||||
#pragma once
|
#pragma once
|
||||||
|
|
||||||
|
// Tensor layout for Q/K/V tensors passed to attention kernels.
|
||||||
|
// Internally, kernels always operate on BHLD [batch, n_heads, seq_len, head_dim].
|
||||||
|
// When the caller passes BLHD, dims 1 and 2 are transposed at entry.
|
||||||
|
enum TensorLayout : int {
|
||||||
|
BHLD = 0, // [batch, n_heads, seq_len, head_dim]
|
||||||
|
BLHD = 1, // [batch, seq_len, n_heads, head_dim]
|
||||||
|
};
|
||||||
|
|
||||||
|
|
||||||
template<typename T, typename AT = float>
|
template<typename T, typename AT = float>
|
||||||
struct AttentionParams {
|
struct AttentionParams {
|
||||||
@@ -35,35 +43,55 @@ struct AttentionParams {
|
|||||||
AT* __restrict__ ml_part;
|
AT* __restrict__ ml_part;
|
||||||
};
|
};
|
||||||
|
|
||||||
|
// ---- PagedAttentionParams ----
|
||||||
|
// SGLang-style indirect params over a shared KV pool.
|
||||||
|
// k_cache/v_cache: [size, kv_head, head_dim] (bare buffers, no gather).
|
||||||
|
// req_to_token: [num_reqs, max_context_len] token -> slot.
|
||||||
|
// req_pool_indices:[batch] rows of the current batch into req_to_token.
|
||||||
|
// kv_indptr: [batch+1] prefix sum of per-request seq_lens (device).
|
||||||
|
// qo_indptr: [batch+1] prefix sum of per-request q_len (prefill) or
|
||||||
|
// nullptr for decode (q_len == 1 everywhere).
|
||||||
template<typename T, typename AT = float>
|
template<typename T, typename AT = float>
|
||||||
struct PagedAttentionParams {
|
struct PagedAttentionParams {
|
||||||
int batch;
|
int batch;
|
||||||
int q_head;
|
int q_head;
|
||||||
int kv_head;
|
int kv_head;
|
||||||
int q_len;
|
|
||||||
int kv_len;
|
|
||||||
int head_dim;
|
int head_dim;
|
||||||
|
int num_splits;
|
||||||
int use_mask;
|
int use_mask;
|
||||||
int causal_offset;
|
int causal_offset; // -1 = non-causal; >=0 = causal (per-request offset
|
||||||
|
// computed inside kernel from kv_indptr/qo_indptr)
|
||||||
float scale;
|
float scale;
|
||||||
|
|
||||||
int num_splits;
|
// Q: [total_q, q_head, head_dim] (3D flattened — no batch dim).
|
||||||
int page_size;
|
// For decode total_q == batch (q_len=1 per request).
|
||||||
int max_pages;
|
// For prefill total_q == qo_indptr[batch].
|
||||||
|
int q_stride_l, q_stride_h, q_stride_d;
|
||||||
|
|
||||||
// Q strides (layout-agnostic)
|
// Q: [total_q, q_head, head_dim]
|
||||||
int q_stride_b, q_stride_h, q_stride_l, q_stride_d;
|
const T* __restrict__ q;
|
||||||
|
|
||||||
// Mask strides (2D, 3D, or 4D)
|
// Flat KV pool: [size, kv_head, head_dim]
|
||||||
|
const T* __restrict__ k_cache;
|
||||||
|
const T* __restrict__ v_cache;
|
||||||
|
|
||||||
|
// Indexing
|
||||||
|
const int64_t* __restrict__ req_to_token; // [num_reqs, max_context_len]
|
||||||
|
const int64_t* __restrict__ req_pool_indices; // [batch]
|
||||||
|
const int* __restrict__ kv_indptr; // [batch+1]
|
||||||
|
const int* __restrict__ qo_indptr; // [batch+1] or nullptr (decode)
|
||||||
|
int max_context_len; // req_to_token stride (dim 1)
|
||||||
|
int max_seq_len; // max per-request seq_len (host-side, for split computation)
|
||||||
|
int total_q; // total Q tokens across all requests (host-side, for grid)
|
||||||
|
int max_q_len; // max per-request q_len (host-side, for prefill grid)
|
||||||
|
|
||||||
|
// Mask: [batch, max_seq_len] (decode) or [batch, 1, q_len, kv_len]
|
||||||
|
// (prefill, optional). mask_h_stride/mask_q_stride are 0 when those
|
||||||
|
// dims are size 1 (broadcast).
|
||||||
int mask_b_stride;
|
int mask_b_stride;
|
||||||
int mask_h_stride;
|
int mask_h_stride;
|
||||||
int mask_q_stride;
|
int mask_q_stride;
|
||||||
|
|
||||||
const T* __restrict__ q;
|
|
||||||
const T* __restrict__ k_cache;
|
|
||||||
const T* __restrict__ v_cache;
|
|
||||||
const bool* __restrict__ mask;
|
const bool* __restrict__ mask;
|
||||||
const int64_t* __restrict__ page_table;
|
|
||||||
|
|
||||||
T* __restrict__ o;
|
T* __restrict__ o;
|
||||||
AT* __restrict__ o_part;
|
AT* __restrict__ o_part;
|
||||||
|
|||||||
@@ -10,17 +10,21 @@ torch::Tensor attn_decode(
|
|||||||
double scale,
|
double scale,
|
||||||
int64_t layout
|
int64_t layout
|
||||||
) {
|
) {
|
||||||
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||||
|
auto stream = at::cuda::getCurrentCUDAStream();
|
||||||
|
|
||||||
AttentionParams<bf16> p;
|
AttentionParams<bf16> p;
|
||||||
attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p);
|
attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p);
|
||||||
TORCH_CHECK(p.q_len == 1, "Q seq_len must be 1");
|
TORCH_CHECK(p.q_len == 1, "Q seq_len must be 1");
|
||||||
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
|
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
|
||||||
|
|
||||||
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
|
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
|
||||||
auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
|
auto O_view = (layout == BLHD) ? O.transpose(1, 2) : O;
|
||||||
p.o = (bf16*)O_view.data_ptr();
|
p.o = (bf16*)O_view.data_ptr();
|
||||||
|
|
||||||
alloc_split_partials(p);
|
alloc_split_partials(p);
|
||||||
DISPATCH_HEAD_DIM(p.head_dim, dispatch_decode, p);
|
DISPATCH_HEAD_DIM(p.head_dim, dispatch_decode, p, stream);
|
||||||
|
C10_CUDA_CHECK(cudaGetLastError());
|
||||||
return O;
|
return O;
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -32,6 +36,6 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
|||||||
py::arg("mask") = py::none(),
|
py::arg("mask") = py::none(),
|
||||||
py::arg("causal_offset") = -1,
|
py::arg("causal_offset") = -1,
|
||||||
py::arg("scale") = 0.0,
|
py::arg("scale") = 0.0,
|
||||||
py::arg("layout") = 0,
|
py::arg("layout") = (int64_t)BHLD,
|
||||||
"GQA decode (tensor-core head-packing on sm_80+, scalar fallback)");
|
"GQA decode (tensor-core head-packing on sm_80+, scalar fallback)");
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -8,69 +8,82 @@
|
|||||||
#include "attn_prefill_split_q.cuh"
|
#include "attn_prefill_split_q.cuh"
|
||||||
#include "attn_decode_split_kv.cuh"
|
#include "attn_decode_split_kv.cuh"
|
||||||
#include "attn_paged_decode_split_kv.cuh"
|
#include "attn_paged_decode_split_kv.cuh"
|
||||||
|
#include "attn_paged_prefill_split_q.cuh"
|
||||||
#ifndef ASTRAI_NO_MMA
|
#ifndef ASTRAI_NO_MMA
|
||||||
#include "attn_prefill_split_q_mma.cuh"
|
#include "attn_prefill_split_q_mma.cuh"
|
||||||
#include "attn_decode_split_kv_mma.cuh"
|
#include "attn_decode_split_kv_mma.cuh"
|
||||||
#include "attn_paged_decode_split_kv_mma.cuh"
|
#include "attn_paged_decode_split_kv_mma.cuh"
|
||||||
|
#include "attn_paged_prefill_split_q_mma.cuh"
|
||||||
#endif
|
#endif
|
||||||
|
|
||||||
// Split-KV: compute number of splits to fill all SMs for small-batch decode.
|
// Split-KV: compute number of splits to fill all SMs for small-batch decode.
|
||||||
// Caps splits so each split processes at least `min_tiles_per_split` tiles,
|
// Caps splits so each split processes at least `min_tiles_per_split` tiles,
|
||||||
// avoiding excessive loop/prologue overhead when tiles are small.
|
// avoiding excessive loop/prologue overhead when tiles are small.
|
||||||
|
//
|
||||||
|
// Target total grid blocks (`TARGET_BLOCKS`) rather than scaling splits by SM
|
||||||
|
// count. Decode blocks are single-warp (32 threads) and a SM hosts ~11 of
|
||||||
|
// them, so the old `2*sm/base` cap badly undersplit at large batch (B=16 got
|
||||||
|
// 3 splits, optimal ~8). Measured (L20, grid search): bandwidth saturates
|
||||||
|
// near 256-512 total blocks; 512 minimizes worst-case latency across the
|
||||||
|
// B x kv grid; more is pure oversplit overhead.
|
||||||
|
constexpr int DECODE_TARGET_BLOCKS = 512;
|
||||||
inline int compute_num_splits(int base_blocks, int tiles_total,
|
inline int compute_num_splits(int base_blocks, int tiles_total,
|
||||||
int min_tiles_per_split = 1) {
|
int min_tiles_per_split = 1) {
|
||||||
int sm_count = 0;
|
int n = (DECODE_TARGET_BLOCKS + base_blocks - 1) / base_blocks;
|
||||||
cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0);
|
|
||||||
int n = (2 * sm_count + base_blocks - 1) / base_blocks;
|
|
||||||
int max_by_work = tiles_total / min_tiles_per_split;
|
int max_by_work = tiles_total / min_tiles_per_split;
|
||||||
return std::max(1, std::min(n, std::min(max_by_work, MAX_SPLITS)));
|
return std::max(1, std::min(n, std::min(max_by_work, MAX_SPLITS)));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Dispatch IsCausal × HasMask — eliminates the duplicated 4-way if/else
|
||||||
|
// ladder that appeared in each dispatch_* function. FN must be a function
|
||||||
|
// template <int HEAD_DIM, bool IsCausal, bool HasMask>; HEAD_DIM is forwarded
|
||||||
|
// as the first template argument so callers only spell it once.
|
||||||
|
//
|
||||||
|
// Usage: DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_decode_mma, HEAD_DIM, p, group_size);
|
||||||
|
#define DISPATCH_CAUSAL_MASK(is_causal, has_mask, FN, HEAD_DIM, ...) \
|
||||||
|
do { \
|
||||||
|
if (is_causal) { \
|
||||||
|
if (has_mask) FN<HEAD_DIM, true, true>(__VA_ARGS__); \
|
||||||
|
else FN<HEAD_DIM, true, false>(__VA_ARGS__); \
|
||||||
|
} else { \
|
||||||
|
if (has_mask) FN<HEAD_DIM, false, true>(__VA_ARGS__); \
|
||||||
|
else FN<HEAD_DIM, false, false>(__VA_ARGS__); \
|
||||||
|
} \
|
||||||
|
} while (0)
|
||||||
|
|
||||||
// ======================================================================
|
// ======================================================================
|
||||||
// Prefill
|
// Prefill
|
||||||
// ======================================================================
|
// ======================================================================
|
||||||
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
#ifndef ASTRAI_NO_MMA
|
||||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
static inline void launch_prefill_mma(AttentionParams<bf16>& p) {
|
static inline void launch_prefill_mma(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||||
constexpr int WARPS = 4;
|
constexpr int WARPS = 4;
|
||||||
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
|
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
|
||||||
using Traits = KernelTraits<HEAD_DIM, BC, WARPS, 2>;
|
using Traits = KernelTraits<HEAD_DIM, BC, WARPS, 2>;
|
||||||
dim3 grid((p.q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS), p.q_head, p.batch);
|
dim3 grid((p.q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS), p.q_head, p.batch);
|
||||||
dim3 block(Traits::NUM_THREADS);
|
dim3 block(Traits::NUM_THREADS);
|
||||||
attn_prefill_split_q_mma_kernel<Traits, IsCausal, HasMask><<<grid, block>>>(p);
|
attn_prefill_split_q_mma_kernel<Traits, IsCausal, HasMask><<<grid, block, 0, stream>>>(p);
|
||||||
}
|
}
|
||||||
#endif
|
#endif
|
||||||
|
|
||||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
static inline void launch_prefill_scalar(AttentionParams<bf16>& p) {
|
static inline void launch_prefill_scalar(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||||
constexpr int G = 8, ROWS = 32, P_BC = 32;
|
constexpr int G = 8, ROWS = 32, P_BC = 32;
|
||||||
dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch);
|
dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch);
|
||||||
dim3 block(G, ROWS);
|
dim3 block(G, ROWS);
|
||||||
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC, IsCausal, HasMask><<<grid, block>>>(p);
|
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC, IsCausal, HasMask><<<grid, block, 0, stream>>>(p);
|
||||||
}
|
}
|
||||||
|
|
||||||
template <int HEAD_DIM>
|
template <int HEAD_DIM>
|
||||||
static inline void dispatch_prefill(AttentionParams<bf16>& p) {
|
static inline void dispatch_prefill(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||||
bool is_causal = (p.causal_offset >= 0);
|
bool is_causal = (p.causal_offset >= 0);
|
||||||
bool has_mask = (p.use_mask && p.mask);
|
bool has_mask = (p.use_mask && p.mask);
|
||||||
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
#ifndef ASTRAI_NO_MMA
|
||||||
if (is_causal) {
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_prefill_mma, HEAD_DIM, p, stream);
|
||||||
if (has_mask) launch_prefill_mma<HEAD_DIM, true, true>(p);
|
|
||||||
else launch_prefill_mma<HEAD_DIM, true, false>(p);
|
|
||||||
} else {
|
|
||||||
if (has_mask) launch_prefill_mma<HEAD_DIM, false, true>(p);
|
|
||||||
else launch_prefill_mma<HEAD_DIM, false, false>(p);
|
|
||||||
}
|
|
||||||
#else
|
#else
|
||||||
if (is_causal) {
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_prefill_scalar, HEAD_DIM, p, stream);
|
||||||
if (has_mask) launch_prefill_scalar<HEAD_DIM, true, true>(p);
|
|
||||||
else launch_prefill_scalar<HEAD_DIM, true, false>(p);
|
|
||||||
} else {
|
|
||||||
if (has_mask) launch_prefill_scalar<HEAD_DIM, false, true>(p);
|
|
||||||
else launch_prefill_scalar<HEAD_DIM, false, false>(p);
|
|
||||||
}
|
|
||||||
#endif
|
#endif
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -84,112 +97,127 @@ static inline void dispatch_prefill(AttentionParams<bf16>& p) {
|
|||||||
// enabling STAGES=2 (double-buffer) within the 32KB smem budget — eliminates
|
// enabling STAGES=2 (double-buffer) within the 32KB smem budget — eliminates
|
||||||
// the 176-byte spill that STAGES=1+BC=32 suffered.
|
// the 176-byte spill that STAGES=1+BC=32 suffered.
|
||||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
static inline void launch_decode_mma(AttentionParams<bf16>& p, int group_size) {
|
static inline void launch_decode_mma(AttentionParams<bf16>& p, int group_size, cudaStream_t stream) {
|
||||||
int G = p.q_head / p.kv_head;
|
int G = p.q_head / p.kv_head;
|
||||||
constexpr int MAX_G = 16;
|
constexpr int MAX_G = 16;
|
||||||
int num_passes = (G + MAX_G - 1) / MAX_G;
|
int num_passes = (G + MAX_G - 1) / MAX_G;
|
||||||
constexpr int BC = 16;
|
constexpr int BC = 16;
|
||||||
int tiles_total = (p.kv_len + BC - 1) / BC;
|
int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total, 2);
|
p.num_splits = compute_num_splits(p.batch * p.kv_head * num_passes, tiles_total, 2);
|
||||||
constexpr int STAGES = 2;
|
constexpr int STAGES = 2;
|
||||||
using Traits = KernelTraits<HEAD_DIM, BC, 1, STAGES>;
|
using Traits = KernelTraits<HEAD_DIM, BC, 1, STAGES>;
|
||||||
dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
|
dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
|
||||||
attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask><<<grid, 32>>>(p);
|
attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask><<<grid, 32, 0, stream>>>(p);
|
||||||
}
|
}
|
||||||
#endif
|
#endif
|
||||||
|
|
||||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
static inline void launch_decode_scalar(AttentionParams<bf16>& p, int group_size) {
|
static inline void launch_decode_scalar(AttentionParams<bf16>& p, int group_size, cudaStream_t stream) {
|
||||||
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
|
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
|
||||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||||
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
|
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
|
||||||
int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit
|
int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit
|
||||||
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
|
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
|
||||||
dim3 block(32, g);
|
dim3 block(32, g);
|
||||||
attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
|
attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem, stream>>>(p);
|
||||||
}
|
}
|
||||||
|
|
||||||
template <int HEAD_DIM>
|
template <int HEAD_DIM>
|
||||||
static inline void dispatch_decode(AttentionParams<bf16>& p) {
|
static inline void dispatch_decode(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||||
bool is_causal = (p.causal_offset >= 0);
|
bool is_causal = (p.causal_offset >= 0);
|
||||||
bool has_mask = (p.use_mask && p.mask);
|
bool has_mask = (p.use_mask && p.mask);
|
||||||
int group_size = p.q_head / p.kv_head;
|
int group_size = p.q_head / p.kv_head;
|
||||||
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
#ifndef ASTRAI_NO_MMA
|
||||||
if (is_causal) {
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_decode_mma, HEAD_DIM, p, group_size, stream);
|
||||||
if (has_mask) launch_decode_mma<HEAD_DIM, true, true>(p, group_size);
|
|
||||||
else launch_decode_mma<HEAD_DIM, true, false>(p, group_size);
|
|
||||||
} else {
|
|
||||||
if (has_mask) launch_decode_mma<HEAD_DIM, false, true>(p, group_size);
|
|
||||||
else launch_decode_mma<HEAD_DIM, false, false>(p, group_size);
|
|
||||||
}
|
|
||||||
#else
|
#else
|
||||||
if (is_causal) {
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_decode_scalar, HEAD_DIM, p, group_size, stream);
|
||||||
if (has_mask) launch_decode_scalar<HEAD_DIM, true, true>(p, group_size);
|
|
||||||
else launch_decode_scalar<HEAD_DIM, true, false>(p, group_size);
|
|
||||||
} else {
|
|
||||||
if (has_mask) launch_decode_scalar<HEAD_DIM, false, true>(p, group_size);
|
|
||||||
else launch_decode_scalar<HEAD_DIM, false, false>(p, group_size);
|
|
||||||
}
|
|
||||||
#endif
|
#endif
|
||||||
|
|
||||||
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p);
|
||||||
}
|
}
|
||||||
|
|
||||||
// ======================================================================
|
// ======================================================================
|
||||||
// Paged Decode
|
// Paged Decode (SGLang-style: flat pool + req_to_token + kv_indptr)
|
||||||
// ======================================================================
|
// ======================================================================
|
||||||
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
#ifndef ASTRAI_NO_MMA
|
||||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
static inline void launch_paged_decode_mma(PagedAttentionParams<bf16>& p, int group_size) {
|
static inline void launch_paged_decode_mma(PagedAttentionParams<bf16>& p, cudaStream_t stream) {
|
||||||
int G = p.q_head / p.kv_head;
|
int G = p.q_head / p.kv_head;
|
||||||
constexpr int MAX_G = 16;
|
constexpr int MAX_G = 16;
|
||||||
constexpr int BC = 16;
|
constexpr int BC = 16;
|
||||||
int num_passes = (G + MAX_G - 1) / MAX_G;
|
int num_passes = (G + MAX_G - 1) / MAX_G;
|
||||||
int tiles_total = (p.kv_len + BC - 1) / BC;
|
int tiles_total = (p.max_seq_len + BC - 1) / BC;
|
||||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total, 2);
|
p.num_splits = compute_num_splits(p.batch * p.kv_head * num_passes, tiles_total, 2);
|
||||||
constexpr int STAGES = 2;
|
constexpr int STAGES = 2;
|
||||||
using Traits = KernelTraits<HEAD_DIM, BC, 1, STAGES>;
|
using Traits = KernelTraits<HEAD_DIM, BC, 1, STAGES>;
|
||||||
dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
|
dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
|
||||||
paged_attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask> <<<grid, 32>>>(p);
|
paged_attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask> <<<grid, 32, 0, stream>>>(p);
|
||||||
}
|
}
|
||||||
#endif
|
#endif
|
||||||
|
|
||||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
static inline void launch_paged_decode_scalar(PagedAttentionParams<bf16>& p, int group_size) {
|
static inline void launch_paged_decode_scalar(PagedAttentionParams<bf16>& p, int group_size, cudaStream_t stream) {
|
||||||
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
|
int chunks_total = (p.max_seq_len + PDC_CHUNK - 1) / PDC_CHUNK;
|
||||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||||
size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16);
|
size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16);
|
||||||
int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit
|
int g = min(group_size, 32);
|
||||||
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
|
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
|
||||||
dim3 block(32, g);
|
dim3 block(32, g);
|
||||||
paged_attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
|
paged_attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem, stream>>>(p);
|
||||||
}
|
}
|
||||||
|
|
||||||
template <int HEAD_DIM>
|
template <int HEAD_DIM>
|
||||||
static inline void dispatch_paged_decode(PagedAttentionParams<bf16>& p) {
|
static inline void dispatch_paged_decode(PagedAttentionParams<bf16>& p, cudaStream_t stream) {
|
||||||
bool is_causal = (p.causal_offset >= 0);
|
bool is_causal = (p.causal_offset >= 0);
|
||||||
bool has_mask = (p.use_mask && p.mask);
|
bool has_mask = (p.use_mask && p.mask);
|
||||||
int group_size = p.q_head / p.kv_head;
|
int group_size = p.q_head / p.kv_head;
|
||||||
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
#ifndef ASTRAI_NO_MMA
|
||||||
if (is_causal) {
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_decode_mma, HEAD_DIM, p, stream);
|
||||||
if (has_mask) launch_paged_decode_mma<HEAD_DIM, true, true>(p, group_size);
|
|
||||||
else launch_paged_decode_mma<HEAD_DIM, true, false>(p, group_size);
|
|
||||||
} else {
|
|
||||||
if (has_mask) launch_paged_decode_mma<HEAD_DIM, false, true>(p, group_size);
|
|
||||||
else launch_paged_decode_mma<HEAD_DIM, false, false>(p, group_size);
|
|
||||||
}
|
|
||||||
#else
|
#else
|
||||||
if (is_causal) {
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_decode_scalar, HEAD_DIM, p, group_size, stream);
|
||||||
if (has_mask) launch_paged_decode_scalar<HEAD_DIM, true, true>(p, group_size);
|
#endif
|
||||||
else launch_paged_decode_scalar<HEAD_DIM, true, false>(p, group_size);
|
|
||||||
} else {
|
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p);
|
||||||
if (has_mask) launch_paged_decode_scalar<HEAD_DIM, false, true>(p, group_size);
|
}
|
||||||
else launch_paged_decode_scalar<HEAD_DIM, false, false>(p, group_size);
|
|
||||||
|
// ======================================================================
|
||||||
|
// Paged Prefill (SGLang-style: flat pool + ragged batch)
|
||||||
|
// ======================================================================
|
||||||
|
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
|
static inline void launch_paged_prefill_mma(PagedAttentionParams<bf16>& p, cudaStream_t stream) {
|
||||||
|
constexpr int WARPS = 4;
|
||||||
|
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
|
||||||
|
using Traits = KernelTraits<HEAD_DIM, BC, WARPS, 2>;
|
||||||
|
int max_q_tiles = (p.max_q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS);
|
||||||
|
dim3 grid(max_q_tiles, p.q_head, p.batch);
|
||||||
|
dim3 block(Traits::NUM_THREADS);
|
||||||
|
paged_attn_prefill_split_q_mma_kernel<Traits, IsCausal, HasMask><<<grid, block, 0, stream>>>(p);
|
||||||
}
|
}
|
||||||
#endif
|
#endif
|
||||||
|
|
||||||
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
|
static inline void launch_paged_prefill_scalar(PagedAttentionParams<bf16>& p, cudaStream_t stream) {
|
||||||
|
constexpr int G = 8, ROWS = 32, P_BC = 32;
|
||||||
|
int max_q_tiles = (p.max_q_len + ROWS - 1) / ROWS;
|
||||||
|
dim3 grid(max_q_tiles, p.q_head, p.batch);
|
||||||
|
dim3 block(G, ROWS);
|
||||||
|
paged_attn_prefill_split_q_kernel<HEAD_DIM, G, ROWS, P_BC, IsCausal, HasMask>
|
||||||
|
<<<grid, block, 0, stream>>>(p);
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
static inline void dispatch_paged_prefill(PagedAttentionParams<bf16>& p, cudaStream_t stream) {
|
||||||
|
bool is_causal = (p.causal_offset >= 0);
|
||||||
|
bool has_mask = (p.use_mask && p.mask);
|
||||||
|
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_prefill_mma, HEAD_DIM, p, stream);
|
||||||
|
#else
|
||||||
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_prefill_scalar, HEAD_DIM, p, stream);
|
||||||
|
#endif
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -8,14 +8,14 @@
|
|||||||
using bf16 = __nv_bfloat16;
|
using bf16 = __nv_bfloat16;
|
||||||
|
|
||||||
// Dispatch head_dim: shared macro — avoids C++20 lambda template syntax.
|
// Dispatch head_dim: shared macro — avoids C++20 lambda template syntax.
|
||||||
// Usage: DISPATCH_HEAD_DIM(hd, fn, arg)
|
// Usage: DISPATCH_HEAD_DIM(hd, fn, args...)
|
||||||
// Expands to: fn<32>(arg); fn<64>(arg); etc.
|
// Expands to: fn<32>(args...); fn<64>(args...); etc.
|
||||||
#define DISPATCH_HEAD_DIM(hd, fn, arg) \
|
#define DISPATCH_HEAD_DIM(hd, fn, ...) \
|
||||||
switch (hd) { \
|
switch (hd) { \
|
||||||
case 32: fn<32>(arg); break; \
|
case 32: fn<32>(__VA_ARGS__); break; \
|
||||||
case 64: fn<64>(arg); break; \
|
case 64: fn<64>(__VA_ARGS__); break; \
|
||||||
case 128: fn<128>(arg); break; \
|
case 128: fn<128>(__VA_ARGS__); break; \
|
||||||
case 256: fn<256>(arg); break; \
|
case 256: fn<256>(__VA_ARGS__); break; \
|
||||||
default: \
|
default: \
|
||||||
TORCH_CHECK(false, "unsupported head_dim ", hd, \
|
TORCH_CHECK(false, "unsupported head_dim ", hd, \
|
||||||
" (supported: 32, 64, 128, 256)"); \
|
" (supported: 32, 64, 128, 256)"); \
|
||||||
@@ -37,7 +37,7 @@ inline void alloc_split_partials(P& p) {
|
|||||||
// ---- Shared Q-dims + strides extraction ----
|
// ---- Shared Q-dims + strides extraction ----
|
||||||
template <typename P>
|
template <typename P>
|
||||||
inline void extract_q_dims_and_strides(torch::Tensor& q, int64_t layout, P& p) {
|
inline void extract_q_dims_and_strides(torch::Tensor& q, int64_t layout, P& p) {
|
||||||
if (layout == 1) q = q.transpose(1, 2);
|
if (layout == BLHD) q = q.transpose(1, 2);
|
||||||
p.batch = (int)q.size(0);
|
p.batch = (int)q.size(0);
|
||||||
p.q_head = (int)q.size(1);
|
p.q_head = (int)q.size(1);
|
||||||
p.q_len = (int)q.size(2);
|
p.q_len = (int)q.size(2);
|
||||||
@@ -109,7 +109,7 @@ inline void attn_pack_params(
|
|||||||
|
|
||||||
extract_q_dims_and_strides(q, layout, p);
|
extract_q_dims_and_strides(q, layout, p);
|
||||||
|
|
||||||
if (layout == 1) k = k.transpose(1, 2), v = v.transpose(1, 2);
|
if (layout == BLHD) k = k.transpose(1, 2), v = v.transpose(1, 2);
|
||||||
|
|
||||||
p.kv_head = (int)k.size(1);
|
p.kv_head = (int)k.size(1);
|
||||||
p.kv_len = (int)k.size(2);
|
p.kv_len = (int)k.size(2);
|
||||||
@@ -134,54 +134,178 @@ inline void attn_pack_params(
|
|||||||
pack_mask(mask, p);
|
pack_mask(mask, p);
|
||||||
}
|
}
|
||||||
|
|
||||||
// ---- attn_pack_paged_params ----
|
// ---- attn_pack_paged_decode_params ----
|
||||||
|
// SGLang-style: flat KV pool + req_to_token indexing + variable
|
||||||
|
// seq_lens via kv_indptr. Q is [batch, q_head, head_dim] (q_len=1 per req).
|
||||||
template<typename T>
|
template<typename T>
|
||||||
inline void attn_pack_paged_params(
|
inline void attn_pack_paged_decode_params(
|
||||||
torch::Tensor q,
|
torch::Tensor q,
|
||||||
torch::Tensor page_table,
|
|
||||||
torch::Tensor k_cache,
|
torch::Tensor k_cache,
|
||||||
torch::Tensor v_cache,
|
torch::Tensor v_cache,
|
||||||
int64_t page_size,
|
torch::Tensor req_to_token,
|
||||||
int64_t kv_len,
|
torch::Tensor req_pool_indices,
|
||||||
|
torch::Tensor kv_indptr,
|
||||||
|
int64_t max_seq_len,
|
||||||
c10::optional<torch::Tensor> mask,
|
c10::optional<torch::Tensor> mask,
|
||||||
int64_t causal_offset,
|
int64_t causal_offset,
|
||||||
double scale,
|
double scale,
|
||||||
int64_t layout,
|
|
||||||
PagedAttentionParams<T>& p
|
PagedAttentionParams<T>& p
|
||||||
) {
|
) {
|
||||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||||
|
|
||||||
TORCH_CHECK(q.is_cuda() && page_table.is_cuda() && k_cache.is_cuda() && v_cache.is_cuda());
|
TORCH_CHECK(q.is_cuda() && k_cache.is_cuda() && v_cache.is_cuda());
|
||||||
|
TORCH_CHECK(req_to_token.is_cuda() && req_pool_indices.is_cuda() && kv_indptr.is_cuda());
|
||||||
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16");
|
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16");
|
||||||
TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bf16");
|
TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bf16");
|
||||||
TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bf16");
|
TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bf16");
|
||||||
TORCH_CHECK(page_table.dtype() == torch::kLong, "page_table must be int64");
|
TORCH_CHECK(req_to_token.dtype() == torch::kLong, "req_to_token must be int64");
|
||||||
TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must have identical shapes");
|
TORCH_CHECK(req_pool_indices.dtype() == torch::kLong, "req_pool_indices must be int64");
|
||||||
|
TORCH_CHECK(kv_indptr.dtype() == torch::kInt32, "kv_indptr must be int32");
|
||||||
|
TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must match");
|
||||||
|
TORCH_CHECK(k_cache.dim() == 3, "k_cache must be 3D [size, kv_head, head_dim]");
|
||||||
|
TORCH_CHECK(q.dim() == 3, "q must be 3D [batch, q_head, head_dim]");
|
||||||
|
|
||||||
extract_q_dims_and_strides(q, layout, p);
|
p.batch = (int)q.size(0);
|
||||||
|
p.q_head = (int)q.size(1);
|
||||||
p.kv_head = (int)k_cache.size(2);
|
p.head_dim = (int)q.size(2);
|
||||||
p.kv_len = (int)kv_len;
|
p.kv_head = (int)k_cache.size(1);
|
||||||
p.page_size = (int)page_size;
|
TORCH_CHECK(k_cache.size(2) == p.head_dim, "k_cache head_dim mismatch");
|
||||||
p.max_pages = (int)page_table.size(1);
|
|
||||||
|
|
||||||
TORCH_CHECK(q.size(2) == 1, "Q seq_len must be 1 (decode)");
|
|
||||||
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
|
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
|
||||||
TORCH_CHECK(k_cache.size(1) == page_size,
|
TORCH_CHECK(p.q_head % p.kv_head == 0, "q_head must be divisible by kv_head");
|
||||||
"k_cache dim 1 must equal page_size, got ",
|
|
||||||
k_cache.size(1), " vs ", page_size);
|
p.q_stride_l = (int)q.stride(0);
|
||||||
|
p.q_stride_h = (int)q.stride(1);
|
||||||
|
p.q_stride_d = (int)q.stride(2);
|
||||||
|
|
||||||
|
p.k_cache = (const T*)k_cache.data_ptr();
|
||||||
|
p.v_cache = (const T*)v_cache.data_ptr();
|
||||||
|
p.q = (const T*)q.data_ptr();
|
||||||
|
p.req_to_token = req_to_token.data_ptr<int64_t>();
|
||||||
|
p.req_pool_indices = req_pool_indices.data_ptr<int64_t>();
|
||||||
|
p.kv_indptr = kv_indptr.data_ptr<int>();
|
||||||
|
p.qo_indptr = nullptr;
|
||||||
|
p.max_context_len = (int)req_to_token.size(1);
|
||||||
|
p.max_seq_len = (int)max_seq_len;
|
||||||
|
p.total_q = p.batch; // decode: 1 Q token per request
|
||||||
|
p.max_q_len = 1;
|
||||||
|
|
||||||
p.causal_offset = (int)causal_offset;
|
p.causal_offset = (int)causal_offset;
|
||||||
p.use_mask = (mask.has_value() && mask.value().defined()) ? 1 : 0;
|
p.use_mask = (mask.has_value() && mask.value().defined()) ? 1 : 0;
|
||||||
p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
|
p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
|
||||||
|
|
||||||
p.page_table = page_table.data_ptr<int64_t>();
|
if (p.use_mask) {
|
||||||
|
auto m = mask.value();
|
||||||
|
TORCH_CHECK(m.is_cuda() && m.dtype() == torch::kBool, "mask must be bool CUDA");
|
||||||
|
TORCH_CHECK(m.size(0) == p.batch, "mask batch mismatch");
|
||||||
|
p.mask_b_stride = (int)m.stride(0);
|
||||||
|
p.mask_h_stride = 0;
|
||||||
|
p.mask_q_stride = 0;
|
||||||
|
p.mask = m.data_ptr<bool>();
|
||||||
|
} else {
|
||||||
|
p.mask = nullptr;
|
||||||
|
p.mask_b_stride = 0;
|
||||||
|
p.mask_h_stride = 0;
|
||||||
|
p.mask_q_stride = 0;
|
||||||
|
}
|
||||||
|
|
||||||
|
p.o = nullptr;
|
||||||
|
p.o_part = nullptr;
|
||||||
|
p.ml_part = nullptr;
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- attn_pack_paged_prefill_params ----
|
||||||
|
// SGLang-style: flat KV pool + req_to_token + ragged batch via qo_indptr.
|
||||||
|
// Q is [total_q, q_head, head_dim] (flattened across all requests).
|
||||||
|
template<typename T>
|
||||||
|
inline void attn_pack_paged_prefill_params(
|
||||||
|
torch::Tensor q,
|
||||||
|
torch::Tensor k_cache,
|
||||||
|
torch::Tensor v_cache,
|
||||||
|
torch::Tensor req_to_token,
|
||||||
|
torch::Tensor req_pool_indices,
|
||||||
|
torch::Tensor kv_indptr,
|
||||||
|
torch::Tensor qo_indptr,
|
||||||
|
c10::optional<torch::Tensor> mask,
|
||||||
|
int64_t max_q_len,
|
||||||
|
int64_t causal_offset,
|
||||||
|
double scale,
|
||||||
|
PagedAttentionParams<T>& p
|
||||||
|
) {
|
||||||
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||||
|
|
||||||
|
TORCH_CHECK(q.is_cuda() && k_cache.is_cuda() && v_cache.is_cuda());
|
||||||
|
TORCH_CHECK(req_to_token.is_cuda() && req_pool_indices.is_cuda());
|
||||||
|
TORCH_CHECK(kv_indptr.is_cuda() && qo_indptr.is_cuda());
|
||||||
|
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16");
|
||||||
|
TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bf16");
|
||||||
|
TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bf16");
|
||||||
|
TORCH_CHECK(req_to_token.dtype() == torch::kLong, "req_to_token must be int64");
|
||||||
|
TORCH_CHECK(req_pool_indices.dtype() == torch::kLong, "req_pool_indices must be int64");
|
||||||
|
TORCH_CHECK(kv_indptr.dtype() == torch::kInt32, "kv_indptr must be int32");
|
||||||
|
TORCH_CHECK(qo_indptr.dtype() == torch::kInt32, "qo_indptr must be int32");
|
||||||
|
TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must match");
|
||||||
|
TORCH_CHECK(k_cache.dim() == 3, "k_cache must be 3D [size, kv_head, head_dim]");
|
||||||
|
TORCH_CHECK(q.dim() == 3, "q must be 3D [total_q, q_head, head_dim]");
|
||||||
|
|
||||||
|
p.q_head = (int)q.size(1);
|
||||||
|
p.head_dim = (int)q.size(2);
|
||||||
|
p.kv_head = (int)k_cache.size(1);
|
||||||
|
p.batch = (int)req_pool_indices.size(0);
|
||||||
|
TORCH_CHECK(k_cache.size(2) == p.head_dim, "k_cache head_dim mismatch");
|
||||||
|
TORCH_CHECK(p.head_dim % 16 == 0, "head_dim must be multiple of 16");
|
||||||
|
TORCH_CHECK(p.q_head % p.kv_head == 0, "q_head must be divisible by kv_head");
|
||||||
|
TORCH_CHECK(kv_indptr.size(0) == p.batch + 1, "kv_indptr must be [batch+1]");
|
||||||
|
TORCH_CHECK(qo_indptr.size(0) == p.batch + 1, "qo_indptr must be [batch+1]");
|
||||||
|
|
||||||
|
p.q_stride_l = (int)q.stride(0);
|
||||||
|
p.q_stride_h = (int)q.stride(1);
|
||||||
|
p.q_stride_d = (int)q.stride(2);
|
||||||
|
|
||||||
p.k_cache = (const T*)k_cache.data_ptr();
|
p.k_cache = (const T*)k_cache.data_ptr();
|
||||||
p.v_cache = (const T*)v_cache.data_ptr();
|
p.v_cache = (const T*)v_cache.data_ptr();
|
||||||
p.q = (const T*)q.data_ptr();
|
p.q = (const T*)q.data_ptr();
|
||||||
|
p.req_to_token = req_to_token.data_ptr<int64_t>();
|
||||||
|
p.req_pool_indices = req_pool_indices.data_ptr<int64_t>();
|
||||||
|
p.kv_indptr = kv_indptr.data_ptr<int>();
|
||||||
|
p.qo_indptr = qo_indptr.data_ptr<int>();
|
||||||
|
p.max_context_len = (int)req_to_token.size(1);
|
||||||
|
p.total_q = (int)q.size(0); // prefill: flattened Q across all requests
|
||||||
|
p.max_q_len = (int)max_q_len;
|
||||||
|
// max_seq_len is unused by the prefill path (decode uses it for split
|
||||||
|
// computation); fill with max_q_len only to keep the POD struct defined.
|
||||||
|
p.max_seq_len = p.max_q_len;
|
||||||
|
|
||||||
|
p.causal_offset = (int)causal_offset;
|
||||||
|
p.use_mask = (mask.has_value() && mask.value().defined()) ? 1 : 0;
|
||||||
|
if (p.use_mask) {
|
||||||
|
auto m = mask.value();
|
||||||
|
TORCH_CHECK(m.is_cuda() && m.dtype() == torch::kBool, "mask must be bool CUDA");
|
||||||
|
TORCH_CHECK(m.size(0) == p.batch, "mask batch mismatch");
|
||||||
|
if (m.dim() == 2) {
|
||||||
|
TORCH_CHECK(m.size(1) <= p.max_context_len, "mask kv_len mismatch");
|
||||||
|
p.mask_b_stride = (int)m.stride(0);
|
||||||
|
p.mask_h_stride = 0;
|
||||||
|
p.mask_q_stride = 0;
|
||||||
|
} else if (m.dim() == 4) {
|
||||||
|
TORCH_CHECK(m.size(1) == 1 || m.size(1) == p.q_head, "mask head mismatch");
|
||||||
|
TORCH_CHECK(m.size(2) == 1 || m.size(2) == p.max_q_len, "mask q_len mismatch");
|
||||||
|
TORCH_CHECK(m.size(3) <= p.max_context_len, "mask kv_len mismatch");
|
||||||
|
p.mask_b_stride = (int)m.stride(0);
|
||||||
|
p.mask_h_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1);
|
||||||
|
p.mask_q_stride = (m.size(2) == 1) ? 0 : (int)m.stride(2);
|
||||||
|
} else {
|
||||||
|
TORCH_CHECK(false, "mask must be 2D or 4D");
|
||||||
|
}
|
||||||
|
p.mask = m.data_ptr<bool>();
|
||||||
|
} else {
|
||||||
|
p.mask = nullptr;
|
||||||
|
p.mask_b_stride = 0;
|
||||||
|
p.mask_h_stride = 0;
|
||||||
|
p.mask_q_stride = 0;
|
||||||
|
}
|
||||||
|
p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
|
||||||
|
|
||||||
p.o = nullptr;
|
p.o = nullptr;
|
||||||
p.o_part = nullptr;
|
p.o_part = nullptr;
|
||||||
p.ml_part = nullptr;
|
p.ml_part = nullptr;
|
||||||
|
|
||||||
pack_mask(mask, p);
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -3,40 +3,44 @@
|
|||||||
|
|
||||||
torch::Tensor attn_paged_decode(
|
torch::Tensor attn_paged_decode(
|
||||||
torch::Tensor q,
|
torch::Tensor q,
|
||||||
torch::Tensor page_table,
|
|
||||||
torch::Tensor k_cache,
|
torch::Tensor k_cache,
|
||||||
torch::Tensor v_cache,
|
torch::Tensor v_cache,
|
||||||
int64_t page_size,
|
torch::Tensor req_to_token,
|
||||||
int64_t kv_len,
|
torch::Tensor req_pool_indices,
|
||||||
|
torch::Tensor kv_indptr,
|
||||||
|
int64_t max_seq_len,
|
||||||
c10::optional<torch::Tensor> mask,
|
c10::optional<torch::Tensor> mask,
|
||||||
int64_t causal_offset,
|
int64_t causal_offset,
|
||||||
double scale,
|
double scale
|
||||||
int64_t layout
|
|
||||||
) {
|
) {
|
||||||
PagedAttentionParams<bf16> p;
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||||
attn_pack_paged_params(q, page_table, k_cache, v_cache,
|
auto stream = at::cuda::getCurrentCUDAStream();
|
||||||
page_size, kv_len, mask, causal_offset, scale, layout, p);
|
|
||||||
|
|
||||||
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
|
PagedAttentionParams<bf16> p;
|
||||||
auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
|
attn_pack_paged_decode_params(q, k_cache, v_cache,
|
||||||
p.o = (bf16*)O_view.data_ptr();
|
req_to_token, req_pool_indices, kv_indptr,
|
||||||
|
max_seq_len, mask, causal_offset, scale, p);
|
||||||
|
|
||||||
|
auto O = torch::empty({q.size(0), q.size(1), q.size(2)}, q.options());
|
||||||
|
p.o = (bf16*)O.data_ptr();
|
||||||
|
|
||||||
alloc_split_partials(p);
|
alloc_split_partials(p);
|
||||||
DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_decode, p);
|
DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_decode, p, stream);
|
||||||
|
C10_CUDA_CHECK(cudaGetLastError());
|
||||||
return O;
|
return O;
|
||||||
}
|
}
|
||||||
|
|
||||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||||
m.def("attn_paged_decode", &attn_paged_decode,
|
m.def("attn_paged_decode", &attn_paged_decode,
|
||||||
py::arg("q"),
|
py::arg("q"),
|
||||||
py::arg("page_table"),
|
|
||||||
py::arg("k_cache"),
|
py::arg("k_cache"),
|
||||||
py::arg("v_cache"),
|
py::arg("v_cache"),
|
||||||
py::arg("page_size"),
|
py::arg("req_to_token"),
|
||||||
py::arg("kv_len"),
|
py::arg("req_pool_indices"),
|
||||||
|
py::arg("kv_indptr"),
|
||||||
|
py::arg("max_seq_len"),
|
||||||
py::arg("mask") = py::none(),
|
py::arg("mask") = py::none(),
|
||||||
py::arg("causal_offset") = -1,
|
py::arg("causal_offset") = -1,
|
||||||
py::arg("scale") = 0.0,
|
py::arg("scale") = 0.0,
|
||||||
py::arg("layout") = 0,
|
"SGLang-style paged decode: flat KV pool + req_to_token + kv_indptr.");
|
||||||
"Paged GQA decode — split-KV with direct page-table access.");
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -5,6 +5,8 @@
|
|||||||
#include "attn_warp_utils.cuh"
|
#include "attn_warp_utils.cuh"
|
||||||
constexpr int PDC_CHUNK = 64;
|
constexpr int PDC_CHUNK = 64;
|
||||||
|
|
||||||
|
// Scalar paged decode (fallback for sm < 80, no tensor cores).
|
||||||
|
// Reads K/V from flat pool via req_to_token indexing.
|
||||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
__global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p) {
|
__global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p) {
|
||||||
int batch = blockIdx.x / p.kv_head;
|
int batch = blockIdx.x / p.kv_head;
|
||||||
@@ -15,8 +17,11 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
|
|||||||
int lane = threadIdx.x;
|
int lane = threadIdx.x;
|
||||||
int hd_per_thread = p.head_dim / 32;
|
int hd_per_thread = p.head_dim / 32;
|
||||||
|
|
||||||
|
const int seq_len = p.kv_indptr[batch + 1] - p.kv_indptr[batch];
|
||||||
|
const int64_t req_idx = p.req_pool_indices[batch];
|
||||||
|
|
||||||
float q_reg[8];
|
float q_reg[8];
|
||||||
int q_off = batch * p.q_stride_b + q_head * p.q_stride_h
|
int q_off = batch * p.q_stride_l + q_head * p.q_stride_h
|
||||||
+ lane * hd_per_thread * p.q_stride_d;
|
+ lane * hd_per_thread * p.q_stride_d;
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int i = 0; i < hd_per_thread; i++)
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
@@ -26,16 +31,19 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
|
|||||||
|
|
||||||
extern __shared__ __align__(16) bf16 k_smem[];
|
extern __shared__ __align__(16) bf16 k_smem[];
|
||||||
|
|
||||||
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
|
int chunks_total = (seq_len + PDC_CHUNK - 1) / PDC_CHUNK;
|
||||||
int chunks_per_split = (chunks_total + p.num_splits - 1) / p.num_splits;
|
int chunks_per_split = (chunks_total + p.num_splits - 1) / p.num_splits;
|
||||||
int ch_begin = split * chunks_per_split;
|
int ch_begin = split * chunks_per_split;
|
||||||
int ch_end = min(chunks_total, ch_begin + chunks_per_split);
|
int ch_end = min(chunks_total, ch_begin + chunks_per_split);
|
||||||
|
|
||||||
const int mask_base = batch * p.mask_b_stride + q_head * p.mask_h_stride;
|
const int mask_base = batch * p.mask_b_stride;
|
||||||
|
const int64_t pool_stride = (int64_t)p.kv_head * p.head_dim;
|
||||||
|
const int64_t head_off = (int64_t)kv_head * p.head_dim;
|
||||||
|
const int64_t rtt_stride = (int64_t)p.max_context_len;
|
||||||
|
|
||||||
for (int ci = ch_begin; ci < ch_end; ci++) {
|
for (int ci = ch_begin; ci < ch_end; ci++) {
|
||||||
int chunk_start = ci * PDC_CHUNK;
|
int chunk_start = ci * PDC_CHUNK;
|
||||||
int this_chunk = min(PDC_CHUNK, p.kv_len - chunk_start);
|
int this_chunk = min(PDC_CHUNK, seq_len - chunk_start);
|
||||||
|
|
||||||
int total = this_chunk * p.head_dim;
|
int total = this_chunk * p.head_dim;
|
||||||
for (int i = threadIdx.y * 32 + lane; i < total;
|
for (int i = threadIdx.y * 32 + lane; i < total;
|
||||||
@@ -43,14 +51,9 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
|
|||||||
int s = i / p.head_dim;
|
int s = i / p.head_dim;
|
||||||
int d_dim = i % p.head_dim;
|
int d_dim = i % p.head_dim;
|
||||||
int pos = chunk_start + s;
|
int pos = chunk_start + s;
|
||||||
int logical_page = pos / p.page_size;
|
int64_t slot = p.req_to_token[req_idx * rtt_stride + pos];
|
||||||
int page_offset = pos % p.page_size;
|
if (slot >= 0) {
|
||||||
int phys_page = p.page_table[batch * p.max_pages + logical_page];
|
int64_t off = slot * pool_stride + head_off + d_dim;
|
||||||
if (phys_page >= 0) {
|
|
||||||
int64_t off = (int64_t)phys_page * p.page_size * p.kv_head * p.head_dim
|
|
||||||
+ (int64_t)page_offset * p.kv_head * p.head_dim
|
|
||||||
+ (int64_t)kv_head * p.head_dim
|
|
||||||
+ d_dim;
|
|
||||||
k_smem[i] = p.k_cache[off];
|
k_smem[i] = p.k_cache[off];
|
||||||
} else {
|
} else {
|
||||||
k_smem[i] = __float2bfloat16(0.0f);
|
k_smem[i] = __float2bfloat16(0.0f);
|
||||||
@@ -72,10 +75,9 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
|
|||||||
if (!p.mask[mask_base + kv_idx])
|
if (!p.mask[mask_base + kv_idx])
|
||||||
masked = true;
|
masked = true;
|
||||||
}
|
}
|
||||||
if constexpr (IsCausal) {
|
// Decode: the query is the last token, so its valid range [0,
|
||||||
if (kv_idx > p.causal_offset)
|
// seq_len) IS the causal range. IsCausal is accepted for dispatch
|
||||||
masked = true;
|
// uniformity but must not apply causal_offset masking here.
|
||||||
}
|
|
||||||
if (masked)
|
if (masked)
|
||||||
partial = -FLT_MAX;
|
partial = -FLT_MAX;
|
||||||
|
|
||||||
@@ -85,17 +87,13 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
|
|||||||
d = d * alpha + beta;
|
d = d * alpha + beta;
|
||||||
|
|
||||||
int pos = chunk_start + s;
|
int pos = chunk_start + s;
|
||||||
int logical_page = pos / p.page_size;
|
int64_t slot = p.req_to_token[req_idx * rtt_stride + pos];
|
||||||
int page_offset = pos % p.page_size;
|
|
||||||
int phys_page = p.page_table[batch * p.max_pages + logical_page];
|
|
||||||
if (masked) {
|
if (masked) {
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int i = 0; i < hd_per_thread; i++)
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
acc_reg[i] = fmaf(acc_reg[i], alpha, 0.0f);
|
acc_reg[i] = fmaf(acc_reg[i], alpha, 0.0f);
|
||||||
} else if (phys_page >= 0) {
|
} else if (slot >= 0) {
|
||||||
int64_t v_base = (int64_t)phys_page * p.page_size * p.kv_head * p.head_dim
|
int64_t v_base = slot * pool_stride + head_off;
|
||||||
+ (int64_t)page_offset * p.kv_head * p.head_dim
|
|
||||||
+ (int64_t)kv_head * p.head_dim;
|
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int i = 0; i < hd_per_thread; i++)
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
acc_reg[i] = fmaf(acc_reg[i], alpha,
|
acc_reg[i] = fmaf(acc_reg[i], alpha,
|
||||||
@@ -148,6 +146,6 @@ __global__ void paged_attn_decode_combine_kernel(PagedAttentionParams<bf16> p) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
|
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
|
||||||
int o_off = batch * p.q_stride_b + q_head * p.q_stride_h + d * p.q_stride_d;
|
int o_off = batch * p.q_stride_l + q_head * p.q_stride_h + d * p.q_stride_d;
|
||||||
p.o[o_off] = __float2bfloat16(acc * inv);
|
p.o[o_off] = __float2bfloat16(acc * inv);
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -5,12 +5,16 @@
|
|||||||
#include "attn_mma_utils.cuh"
|
#include "attn_mma_utils.cuh"
|
||||||
#include "attn_warp_utils.cuh"
|
#include "attn_warp_utils.cuh"
|
||||||
|
|
||||||
// Paged split-KV tensor-core decode via GQA head-packing.
|
// SGLang-style split-KV tensor-core decode.
|
||||||
// Reads K/V directly from the page pool through a page table — one tile
|
|
||||||
// (BC=32) fits within a single page (page_size >= 32), so the page-table
|
|
||||||
// lookup happens once per tile for cp.async.
|
|
||||||
//
|
//
|
||||||
// IsCausal and HasMask are compile-time bools.
|
// Reads K/V directly from a flat pool [size, kv_head, head_dim] via
|
||||||
|
// req_to_token indexing — no gather, no page-table dimension.
|
||||||
|
// Each batch element has its own seq_len (from kv_indptr), eliminating
|
||||||
|
// padding waste: short sequences only process the tiles they own.
|
||||||
|
//
|
||||||
|
// For decode (q_len=1), causal masking is implicit — each request attends
|
||||||
|
// to [0, seq_len) which is exactly its valid range. The IsCausal flag
|
||||||
|
// is accepted for dispatch uniformity but does not change maxc.
|
||||||
template <typename Traits, bool IsCausal, bool HasMask>
|
template <typename Traits, bool IsCausal, bool HasMask>
|
||||||
__global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16> p) {
|
__global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16> p) {
|
||||||
const int lane = threadIdx.x;
|
const int lane = threadIdx.x;
|
||||||
@@ -22,6 +26,10 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
|
|||||||
const int batch = blockIdx.y;
|
const int batch = blockIdx.y;
|
||||||
const int split = blockIdx.z;
|
const int split = blockIdx.z;
|
||||||
|
|
||||||
|
// Per-request seq_len from device-side kv_indptr — no padding.
|
||||||
|
const int seq_len = p.kv_indptr[batch + 1] - p.kv_indptr[batch];
|
||||||
|
const int64_t req_idx = p.req_pool_indices[batch];
|
||||||
|
|
||||||
constexpr int MAX_G = 16;
|
constexpr int MAX_G = 16;
|
||||||
const int G_total = p.q_head / p.kv_head;
|
const int G_total = p.q_head / p.kv_head;
|
||||||
const int g_begin = pass * MAX_G;
|
const int g_begin = pass * MAX_G;
|
||||||
@@ -31,19 +39,13 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
|
|||||||
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
|
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
|
||||||
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
|
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
|
||||||
|
|
||||||
#pragma unroll
|
const int q_base = batch * p.q_stride_l + q_head0 * p.q_stride_h;
|
||||||
for (int i = lane; i < Traits::STAGES * Traits::BC * Traits::LD; i += 32) {
|
|
||||||
sK[i] = __float2bfloat16(0.0f);
|
|
||||||
sV[i] = __float2bfloat16(0.0f);
|
|
||||||
}
|
|
||||||
__syncwarp();
|
|
||||||
|
|
||||||
const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h;
|
|
||||||
const int qra = gid;
|
const int qra = gid;
|
||||||
const int qrb = gid + 8;
|
const int qrb = gid + 8;
|
||||||
const bool va = qra < G, vb = qrb < G;
|
const bool va = qra < G, vb = qrb < G;
|
||||||
unsigned Qa[Traits::KD][4];
|
unsigned Qa[Traits::KD][4];
|
||||||
load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
|
load_q_mma_frags<Traits::KD>(p.q + q_base,
|
||||||
|
p.q_stride_h, p.q_stride_d,
|
||||||
qra, qrb, va, vb, tid4, Qa);
|
qra, qrb, va, vb, tid4, Qa);
|
||||||
|
|
||||||
float Oacc[Traits::DN8][4];
|
float Oacc[Traits::DN8][4];
|
||||||
@@ -52,19 +54,19 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
|
|||||||
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
||||||
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
||||||
|
|
||||||
const int tiles_total = (p.kv_len + Traits::BC - 1) / Traits::BC;
|
const int tiles_total = (seq_len + Traits::BC - 1) / Traits::BC;
|
||||||
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
|
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
|
||||||
const int ti_begin = split * tiles_per_split;
|
const int ti_begin = split * tiles_per_split;
|
||||||
const int ti_end = min(tiles_total, ti_begin + tiles_per_split);
|
const int ti_end = min(tiles_total, ti_begin + tiles_per_split);
|
||||||
|
|
||||||
const int64_t page_stride = (int64_t)p.page_size * p.kv_head * Traits::HEAD_DIM;
|
// Flat pool stride: [size, kv_head, head_dim] — contiguous.
|
||||||
const int64_t pos_stride = (int64_t)p.kv_head * Traits::HEAD_DIM;
|
const int64_t pool_stride = (int64_t)p.kv_head * Traits::HEAD_DIM;
|
||||||
const int64_t head_off = (int64_t)kv_head * Traits::HEAD_DIM;
|
const int64_t head_off = (int64_t)kv_head * Traits::HEAD_DIM;
|
||||||
|
const int64_t rtt_stride = (int64_t)p.max_context_len;
|
||||||
|
|
||||||
// ---- Load tile lambda: paged addressing ----
|
// ---- Load tile lambda: SGLang addressing ----
|
||||||
// Unified per-element page-table lookup. When page_size >= BC, all
|
// slot = req_to_token[req_idx * max_context_len + kc]
|
||||||
// elements in a tile share the same page, so the lookup is redundant
|
// gmem = k_cache[slot * pool_stride + head_off + d]
|
||||||
// but harmless (L1-cached). This avoids a branch on page_size.
|
|
||||||
auto load_tile = [&](int ti, int buf) {
|
auto load_tile = [&](int ti, int buf) {
|
||||||
int kv0 = ti * Traits::BC;
|
int kv0 = ti * Traits::BC;
|
||||||
bf16* dK = sK + buf * Traits::BC * Traits::LD;
|
bf16* dK = sK + buf * Traits::BC * Traits::LD;
|
||||||
@@ -74,16 +76,13 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
|
|||||||
i += Traits::NUM_THREADS * Traits::VEC) {
|
i += Traits::NUM_THREADS * Traits::VEC) {
|
||||||
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
|
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
|
||||||
int kc = kv0 + r;
|
int kc = kv0 + r;
|
||||||
bool valid = (kc < p.kv_len);
|
bool valid = (kc < seq_len);
|
||||||
if constexpr (HasMask) {
|
if constexpr (HasMask) {
|
||||||
valid = valid && p.mask[batch * p.mask_b_stride + kc];
|
valid = valid && p.mask[batch * p.mask_b_stride + kc];
|
||||||
}
|
}
|
||||||
int phys_page = valid ? p.page_table[batch * p.max_pages + kc] : 0;
|
int64_t slot = valid ? p.req_to_token[req_idx * rtt_stride + kc] : 0;
|
||||||
valid = valid && (phys_page >= 0);
|
valid = valid && (slot >= 0);
|
||||||
int page_off = kc % p.page_size;
|
int64_t gmem_base = slot * pool_stride + head_off;
|
||||||
int64_t gmem_base = (int64_t)phys_page * page_stride
|
|
||||||
+ (int64_t)page_off * pos_stride
|
|
||||||
+ head_off;
|
|
||||||
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
|
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
|
||||||
cp_async_16_pred(&dK[off], &p.k_cache[gmem_base + d], valid);
|
cp_async_16_pred(&dK[off], &p.k_cache[gmem_base + d], valid);
|
||||||
cp_async_16_pred(&dV[off], &p.v_cache[gmem_base + d], valid);
|
cp_async_16_pred(&dV[off], &p.v_cache[gmem_base + d], valid);
|
||||||
@@ -91,10 +90,6 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
|
|||||||
cp_async_commit();
|
cp_async_commit();
|
||||||
};
|
};
|
||||||
|
|
||||||
// ---- Multi-stage cp.async pipeline ----
|
|
||||||
// Prologue loads STAGES tiles; each loop iteration waits only for the
|
|
||||||
// oldest outstanding group (wait_group<STAGES-1>) so the STAGES-1 newer
|
|
||||||
// tile loads stay in flight and overlap with the current tile's compute.
|
|
||||||
constexpr int STAGES = Traits::STAGES;
|
constexpr int STAGES = Traits::STAGES;
|
||||||
const int ntiles = ti_end - ti_begin;
|
const int ntiles = ti_end - ti_begin;
|
||||||
|
|
||||||
@@ -111,8 +106,9 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
|
|||||||
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
||||||
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
||||||
|
|
||||||
int maxc = IsCausal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
|
// For decode, maxc = seq_len regardless of IsCausal — the valid
|
||||||
mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc,
|
// range [0, seq_len) IS the causal range (query is the last token).
|
||||||
|
mma_softmax_tile<Traits, HasMask>(kv0, seq_len, seq_len,
|
||||||
0, 0,
|
0, 0,
|
||||||
p.mask_b_stride, 0, 0,
|
p.mask_b_stride, 0, 0,
|
||||||
batch, 0,
|
batch, 0,
|
||||||
@@ -136,7 +132,6 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
|
|||||||
load_tile(ti_begin + it + STAGES, (it + STAGES) & (STAGES - 1));
|
load_tile(ti_begin + it + STAGES, (it + STAGES) & (STAGES - 1));
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
// Fewer tiles than stages: load all, wait for all, process.
|
|
||||||
for (int i = 0; i < ntiles; i++)
|
for (int i = 0; i < ntiles; i++)
|
||||||
load_tile(ti_begin + i, i);
|
load_tile(ti_begin + i, i);
|
||||||
cp_async_wait_group<0>();
|
cp_async_wait_group<0>();
|
||||||
@@ -145,6 +140,7 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
|
|||||||
process_tile(it, it);
|
process_tile(it, it);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ---- write partials ----
|
||||||
auto split_slot = [&](int h) -> size_t {
|
auto split_slot = [&](int h) -> size_t {
|
||||||
size_t bh = (size_t)batch * p.q_head + h;
|
size_t bh = (size_t)batch * p.q_head + h;
|
||||||
return bh * MAX_SPLITS + split;
|
return bh * MAX_SPLITS + split;
|
||||||
|
|||||||
@@ -0,0 +1,48 @@
|
|||||||
|
#include "attn_dispatchers.cuh"
|
||||||
|
#include "attn_entry_utils.cuh"
|
||||||
|
|
||||||
|
torch::Tensor attn_paged_prefill(
|
||||||
|
torch::Tensor q,
|
||||||
|
torch::Tensor k_cache,
|
||||||
|
torch::Tensor v_cache,
|
||||||
|
torch::Tensor req_to_token,
|
||||||
|
torch::Tensor req_pool_indices,
|
||||||
|
torch::Tensor kv_indptr,
|
||||||
|
torch::Tensor qo_indptr,
|
||||||
|
c10::optional<torch::Tensor> mask,
|
||||||
|
int64_t max_q_len,
|
||||||
|
int64_t causal_offset,
|
||||||
|
double scale
|
||||||
|
) {
|
||||||
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||||
|
auto stream = at::cuda::getCurrentCUDAStream();
|
||||||
|
|
||||||
|
PagedAttentionParams<bf16> p;
|
||||||
|
attn_pack_paged_prefill_params(q, k_cache, v_cache,
|
||||||
|
req_to_token, req_pool_indices,
|
||||||
|
kv_indptr, qo_indptr, mask,
|
||||||
|
max_q_len, causal_offset, scale, p);
|
||||||
|
|
||||||
|
auto O = torch::empty({q.size(0), q.size(1), q.size(2)}, q.options());
|
||||||
|
p.o = (bf16*)O.data_ptr();
|
||||||
|
|
||||||
|
DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_prefill, p, stream);
|
||||||
|
C10_CUDA_CHECK(cudaGetLastError());
|
||||||
|
return O;
|
||||||
|
}
|
||||||
|
|
||||||
|
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||||
|
m.def("attn_paged_prefill", &attn_paged_prefill,
|
||||||
|
py::arg("q"),
|
||||||
|
py::arg("k_cache"),
|
||||||
|
py::arg("v_cache"),
|
||||||
|
py::arg("req_to_token"),
|
||||||
|
py::arg("req_pool_indices"),
|
||||||
|
py::arg("kv_indptr"),
|
||||||
|
py::arg("qo_indptr"),
|
||||||
|
py::arg("mask") = py::none(),
|
||||||
|
py::arg("max_q_len"),
|
||||||
|
py::arg("causal_offset") = -1,
|
||||||
|
py::arg("scale") = 0.0,
|
||||||
|
"SGLang-style paged prefill: flat KV pool + ragged batch.");
|
||||||
|
}
|
||||||
@@ -0,0 +1,126 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
#include <float.h>
|
||||||
|
#include "attn_common.h"
|
||||||
|
|
||||||
|
using bf16 = __nv_bfloat16;
|
||||||
|
|
||||||
|
// Scalar paged prefill (fallback for sm < 80, no tensor cores).
|
||||||
|
// Reads K/V from a flat pool via req_to_token, supports ragged batches
|
||||||
|
// via qo_indptr + kv_indptr. Mirrors the split-Q MMA kernel's indexing:
|
||||||
|
// grid (max_q_tiles, q_head, batch), block (G, ROWS).
|
||||||
|
//
|
||||||
|
// HasMask: 4D mask [batch, 1, q_len, kv_len] (True=keep), columns are
|
||||||
|
// request-local kv positions. q_head is the q-index (mask_h broadcast).
|
||||||
|
//
|
||||||
|
// group_reduce_sum<G> is provided by attn_prefill_split_q.cuh (already
|
||||||
|
// included via the dispatcher).
|
||||||
|
template <int HEAD_DIM, int G, int ROWS, int P_BC, bool IsCausal, bool HasMask>
|
||||||
|
__global__ void paged_attn_prefill_split_q_kernel(PagedAttentionParams<bf16> p) {
|
||||||
|
constexpr int DPT = HEAD_DIM / G;
|
||||||
|
|
||||||
|
const int q_tile = blockIdx.x;
|
||||||
|
const int q_head = blockIdx.y;
|
||||||
|
const int req_b = blockIdx.z;
|
||||||
|
const int gpos = threadIdx.x; // 0..G-1 (d-chunk)
|
||||||
|
const int row = threadIdx.y; // 0..ROWS-1 (q row within tile)
|
||||||
|
const int q_row = q_tile * ROWS + row;
|
||||||
|
|
||||||
|
const int seq_len = p.kv_indptr[req_b + 1] - p.kv_indptr[req_b];
|
||||||
|
const int q_len = p.qo_indptr[req_b + 1] - p.qo_indptr[req_b];
|
||||||
|
const int causal_off = seq_len - q_len;
|
||||||
|
const int64_t req_idx = p.req_pool_indices[req_b];
|
||||||
|
const int kv_head = q_head / (p.q_head / p.kv_head);
|
||||||
|
|
||||||
|
__shared__ __align__(16) bf16 sK[P_BC * HEAD_DIM];
|
||||||
|
__shared__ __align__(16) bf16 sV[P_BC * HEAD_DIM];
|
||||||
|
|
||||||
|
// Q base: absolute token = qo_indptr[req_b] + q_row.
|
||||||
|
float qreg[DPT];
|
||||||
|
if (q_row < q_len) {
|
||||||
|
int q_off = (p.qo_indptr[req_b] + q_row) * p.q_stride_l
|
||||||
|
+ q_head * p.q_stride_h + gpos * DPT * p.q_stride_d;
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = 0; i < DPT; i++)
|
||||||
|
qreg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]);
|
||||||
|
}
|
||||||
|
|
||||||
|
float m = -FLT_MAX, l = 0.0f, acc[DPT];
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = 0; i < DPT; i++) acc[i] = 0.0f;
|
||||||
|
|
||||||
|
const int64_t pool_stride = (int64_t)p.kv_head * p.head_dim;
|
||||||
|
const int64_t head_off = (int64_t)kv_head * p.head_dim;
|
||||||
|
const int64_t rtt_stride = (int64_t)p.max_context_len;
|
||||||
|
const int mask_base = req_b * p.mask_b_stride + q_head * p.mask_h_stride
|
||||||
|
+ q_row * p.mask_q_stride;
|
||||||
|
|
||||||
|
int tiles = (seq_len + P_BC - 1) / P_BC;
|
||||||
|
int tt = G * ROWS;
|
||||||
|
int lid = row * G + gpos;
|
||||||
|
|
||||||
|
// Each warp holds (32/G) q-rows; reduce only within this row's G lanes.
|
||||||
|
int lane_in_warp = lid & 31;
|
||||||
|
unsigned gmask = (G == 32) ? 0xFFFFFFFFu
|
||||||
|
: (((1u << G) - 1u) << (lane_in_warp & ~(G - 1)));
|
||||||
|
|
||||||
|
for (int ti = 0; ti < tiles; ti++) {
|
||||||
|
int kv0 = ti * P_BC;
|
||||||
|
int tlen = min(P_BC, seq_len - kv0);
|
||||||
|
|
||||||
|
// Load K/V tile into shared memory via req_to_token (request-local pos).
|
||||||
|
for (int i = lid; i < tlen * HEAD_DIM; i += tt) {
|
||||||
|
int s = i / HEAD_DIM, d_dim = i % HEAD_DIM;
|
||||||
|
int pos = kv0 + s;
|
||||||
|
int64_t slot = p.req_to_token[req_idx * rtt_stride + pos];
|
||||||
|
int64_t off = slot * pool_stride + head_off + d_dim;
|
||||||
|
sK[i] = (slot >= 0) ? p.k_cache[off] : __float2bfloat16(0.0f);
|
||||||
|
sV[i] = (slot >= 0) ? p.v_cache[off] : __float2bfloat16(0.0f);
|
||||||
|
}
|
||||||
|
__syncthreads();
|
||||||
|
|
||||||
|
int lim = tlen;
|
||||||
|
if constexpr (IsCausal) {
|
||||||
|
if (q_row < q_len) {
|
||||||
|
int ep = causal_off + q_row + 1;
|
||||||
|
if (kv0 >= ep)
|
||||||
|
lim = 0;
|
||||||
|
else if (kv0 + tlen > ep)
|
||||||
|
lim = ep - kv0;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for (int s = 0; s < lim; s++) {
|
||||||
|
bool keep = true;
|
||||||
|
if constexpr (HasMask) {
|
||||||
|
if (q_row < q_len && !p.mask[mask_base + kv0 + s])
|
||||||
|
keep = false;
|
||||||
|
}
|
||||||
|
float w = 0.0f;
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = 0; i < DPT; i++)
|
||||||
|
w += qreg[i] * __bfloat162float(sK[s * HEAD_DIM + gpos * DPT + i]);
|
||||||
|
w = group_reduce_sum<G>(w, gmask) * p.scale;
|
||||||
|
if (!keep) w = -FLT_MAX;
|
||||||
|
|
||||||
|
float nm = fmaxf(m, w);
|
||||||
|
float alpha = __expf(m - nm);
|
||||||
|
float beta = __expf(w - nm);
|
||||||
|
l = l * alpha + beta;
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = 0; i < DPT; i++)
|
||||||
|
acc[i] = acc[i] * alpha
|
||||||
|
+ __bfloat162float(sV[s * HEAD_DIM + gpos * DPT + i]) * beta;
|
||||||
|
m = nm;
|
||||||
|
}
|
||||||
|
__syncthreads();
|
||||||
|
}
|
||||||
|
|
||||||
|
if (q_row >= q_len) return;
|
||||||
|
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
|
||||||
|
int o_off = (p.qo_indptr[req_b] + q_row) * p.q_stride_l
|
||||||
|
+ q_head * p.q_stride_h + gpos * DPT * p.q_stride_d;
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = 0; i < DPT; i++)
|
||||||
|
p.o[o_off + i * p.q_stride_d] = __float2bfloat16(acc[i] * inv);
|
||||||
|
}
|
||||||
@@ -0,0 +1,164 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <cfloat>
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
#include "attn_common.h"
|
||||||
|
#include "attn_mma_utils.cuh"
|
||||||
|
|
||||||
|
// SGLang-style split-Q tensor-core prefill.
|
||||||
|
//
|
||||||
|
// Reads K/V directly from a flat pool [size, kv_head, head_dim] via
|
||||||
|
// req_to_token — no gather, no temporary tensor. Supports ragged batches:
|
||||||
|
// each request has its own q_len and kv_len, addressed via qo_indptr and
|
||||||
|
// kv_indptr.
|
||||||
|
//
|
||||||
|
// Grid: (max_q_tiles, q_head, batch) — one batch element per blockIdx.z.
|
||||||
|
// Blocks beyond a request's q_len exit early after writing sentinel-free
|
||||||
|
// no-ops. This avoids the binary-search approach and guarantees every Q
|
||||||
|
// token is covered, even when q_len < BR*WARPS (e.g. decode-like prefill).
|
||||||
|
//
|
||||||
|
// Q layout: [total_q, q_head, head_dim] (3D, flattened across requests).
|
||||||
|
// O layout: same as Q.
|
||||||
|
//
|
||||||
|
// IsCausal is a compile-time bool. When true, each Q row qi (within its
|
||||||
|
// request) attends to [0, causal_offset_b + qi + 1) where
|
||||||
|
// causal_offset_b = kv_len_b - q_len_b (position of first Q token).
|
||||||
|
template <typename Traits, bool IsCausal, bool HasMask>
|
||||||
|
__global__ void paged_attn_prefill_split_q_mma_kernel(PagedAttentionParams<bf16> p) {
|
||||||
|
const int warp = threadIdx.x / 32;
|
||||||
|
const int lane = threadIdx.x % 32;
|
||||||
|
const int gid = lane >> 2;
|
||||||
|
const int tid4 = lane & 3;
|
||||||
|
|
||||||
|
const int q_head = blockIdx.y;
|
||||||
|
const int req_b = blockIdx.z;
|
||||||
|
const int qrow0 = (blockIdx.x * Traits::WARPS + warp) * Traits::BR;
|
||||||
|
|
||||||
|
const int seq_len = p.kv_indptr[req_b + 1] - p.kv_indptr[req_b];
|
||||||
|
const int q_len = p.qo_indptr[req_b + 1] - p.qo_indptr[req_b];
|
||||||
|
const int causal_off = seq_len - q_len;
|
||||||
|
const int64_t req_idx = p.req_pool_indices[req_b];
|
||||||
|
|
||||||
|
// No per-warp early exit — all warps must participate in __syncthreads.
|
||||||
|
// Warps beyond q_len get zero-filled Q frags (va=vb=false) and skip output.
|
||||||
|
const int kv_head = q_head / (p.q_head / p.kv_head);
|
||||||
|
|
||||||
|
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
|
||||||
|
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
|
||||||
|
|
||||||
|
// Q base: offset by qo_indptr[req_b] to get absolute token address.
|
||||||
|
const int q_base = p.qo_indptr[req_b] * p.q_stride_l + q_head * p.q_stride_h;
|
||||||
|
const int qra = qrow0 + gid;
|
||||||
|
const int qrb = qrow0 + gid + 8;
|
||||||
|
const bool va = qra < q_len, vb = qrb < q_len;
|
||||||
|
unsigned Qa[Traits::KD][4];
|
||||||
|
load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_l, p.q_stride_d,
|
||||||
|
qra, qrb, va, vb, tid4, Qa);
|
||||||
|
|
||||||
|
float Oacc[Traits::DN8][4];
|
||||||
|
#pragma unroll
|
||||||
|
for (int j = 0; j < Traits::DN8; j++)
|
||||||
|
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
||||||
|
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
||||||
|
|
||||||
|
const int64_t pool_stride = (int64_t)p.kv_head * Traits::HEAD_DIM;
|
||||||
|
const int64_t head_off = (int64_t)kv_head * Traits::HEAD_DIM;
|
||||||
|
const int64_t rtt_stride = (int64_t)p.max_context_len;
|
||||||
|
|
||||||
|
const int tiles = (seq_len + Traits::BC - 1) / Traits::BC;
|
||||||
|
const int qr0 = qrow0 + gid;
|
||||||
|
const int qr1 = qrow0 + gid + 8;
|
||||||
|
|
||||||
|
// Causal tile-skip (dead code when IsCausal == false)
|
||||||
|
const int max_kv = qrow0 + Traits::BR - 1 + causal_off;
|
||||||
|
const int block_max_kv =
|
||||||
|
blockIdx.x * Traits::WARPS * Traits::BR + Traits::WARPS * Traits::BR - 1
|
||||||
|
+ causal_off;
|
||||||
|
|
||||||
|
int t_end = tiles - 1;
|
||||||
|
if constexpr (IsCausal) {
|
||||||
|
int bt = block_max_kv / Traits::BC;
|
||||||
|
if (bt < t_end) t_end = bt;
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- Load tile lambda: SGLang addressing ----
|
||||||
|
auto load_tile = [&](int ti, int buf) {
|
||||||
|
int kv0 = ti * Traits::BC;
|
||||||
|
bf16* dK = sK + buf * Traits::BC * Traits::LD;
|
||||||
|
bf16* dV = sV + buf * Traits::BC * Traits::LD;
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = threadIdx.x * Traits::VEC; i < Traits::TOTAL;
|
||||||
|
i += Traits::NUM_THREADS * Traits::VEC) {
|
||||||
|
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
|
||||||
|
int kc = kv0 + r;
|
||||||
|
bool valid = kc < seq_len;
|
||||||
|
int64_t slot = valid ? p.req_to_token[req_idx * rtt_stride + kc] : 0;
|
||||||
|
valid = valid && (slot >= 0);
|
||||||
|
int64_t gmem_base = slot * pool_stride + head_off;
|
||||||
|
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
|
||||||
|
cp_async_16_pred(&dK[off], &p.k_cache[gmem_base + d], valid);
|
||||||
|
cp_async_16_pred(&dV[off], &p.v_cache[gmem_base + d], valid);
|
||||||
|
}
|
||||||
|
cp_async_commit();
|
||||||
|
};
|
||||||
|
|
||||||
|
// ---- Prologue + main loop (FA2-style double-buffer) ----
|
||||||
|
load_tile(0, 0);
|
||||||
|
|
||||||
|
for (int ti = 0; ti <= t_end; ti++) {
|
||||||
|
int buf = ti & 1;
|
||||||
|
|
||||||
|
cp_async_wait_group<0>();
|
||||||
|
__syncthreads();
|
||||||
|
if (ti < t_end) load_tile(ti + 1, (ti + 1) & 1);
|
||||||
|
|
||||||
|
const bf16* bK = sK + buf * Traits::BC * Traits::LD;
|
||||||
|
const bf16* bV = sV + buf * Traits::BC * Traits::LD;
|
||||||
|
int kv0 = ti * Traits::BC;
|
||||||
|
|
||||||
|
if (!IsCausal || kv0 <= max_kv) {
|
||||||
|
float Sacc[Traits::NC8][4];
|
||||||
|
mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
|
||||||
|
|
||||||
|
#pragma unroll
|
||||||
|
for (int n8 = 0; n8 < Traits::NC8; n8++)
|
||||||
|
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
||||||
|
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
||||||
|
|
||||||
|
int maxc0 = IsCausal ? min(seq_len, causal_off + qr0 + 1)
|
||||||
|
: seq_len;
|
||||||
|
int maxc1 = IsCausal ? min(seq_len, causal_off + qr1 + 1)
|
||||||
|
: seq_len;
|
||||||
|
// HasMask: mask[batch, q_head, qi, kc] — kc is request-local.
|
||||||
|
mma_softmax_tile<Traits, HasMask>(kv0, maxc0, maxc1,
|
||||||
|
qr0, qr1,
|
||||||
|
p.mask_b_stride, p.mask_h_stride,
|
||||||
|
p.mask_q_stride,
|
||||||
|
req_b, q_head,
|
||||||
|
p.mask,
|
||||||
|
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||||
|
|
||||||
|
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- write output: packed bf16x2 stores ----
|
||||||
|
float rl0 = (l0 > 1e-20f) ? (1.0f / l0) : 0.0f;
|
||||||
|
float rl1 = (l1 > 1e-20f) ? (1.0f / l1) : 0.0f;
|
||||||
|
const int o_base = p.qo_indptr[req_b] * p.q_stride_l + q_head * p.q_stride_h;
|
||||||
|
#pragma unroll
|
||||||
|
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
|
||||||
|
int d = dn8 * 8 + 2 * tid4;
|
||||||
|
if (qr0 < q_len) {
|
||||||
|
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][0] * rl0,
|
||||||
|
Oacc[dn8][1] * rl0);
|
||||||
|
*reinterpret_cast<__nv_bfloat162*>(
|
||||||
|
&p.o[o_base + qr0 * p.q_stride_l + d * p.q_stride_d]) = v;
|
||||||
|
}
|
||||||
|
if (qr1 < q_len) {
|
||||||
|
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][2] * rl1,
|
||||||
|
Oacc[dn8][3] * rl1);
|
||||||
|
*reinterpret_cast<__nv_bfloat162*>(
|
||||||
|
&p.o[o_base + qr1 * p.q_stride_l + d * p.q_stride_d]) = v;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -10,15 +10,19 @@ torch::Tensor attn_prefill(
|
|||||||
double scale,
|
double scale,
|
||||||
int64_t layout
|
int64_t layout
|
||||||
) {
|
) {
|
||||||
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||||
|
auto stream = at::cuda::getCurrentCUDAStream();
|
||||||
|
|
||||||
AttentionParams<bf16> p;
|
AttentionParams<bf16> p;
|
||||||
attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p);
|
attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p);
|
||||||
TORCH_CHECK(p.head_dim % 16 == 0, "head_dim must be multiple of 16");
|
TORCH_CHECK(p.head_dim % 16 == 0, "head_dim must be multiple of 16");
|
||||||
|
|
||||||
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
|
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
|
||||||
auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
|
auto O_view = (layout == BLHD) ? O.transpose(1, 2) : O;
|
||||||
p.o = (bf16*)O_view.data_ptr();
|
p.o = (bf16*)O_view.data_ptr();
|
||||||
|
|
||||||
DISPATCH_HEAD_DIM(p.head_dim, dispatch_prefill, p);
|
DISPATCH_HEAD_DIM(p.head_dim, dispatch_prefill, p, stream);
|
||||||
|
C10_CUDA_CHECK(cudaGetLastError());
|
||||||
return O;
|
return O;
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -30,6 +34,6 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
|||||||
py::arg("mask") = py::none(),
|
py::arg("mask") = py::none(),
|
||||||
py::arg("causal_offset") = -1,
|
py::arg("causal_offset") = -1,
|
||||||
py::arg("scale") = 0.0,
|
py::arg("scale") = 0.0,
|
||||||
py::arg("layout") = 0,
|
py::arg("layout") = (int64_t)BHLD,
|
||||||
"GQA prefill (tensor-core mma on sm_80+, scalar fallback)");
|
"GQA prefill (tensor-core mma on sm_80+, scalar fallback)");
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,4 +1,6 @@
|
|||||||
#include <torch/extension.h>
|
#include <torch/extension.h>
|
||||||
|
#include <c10/cuda/CUDAGuard.h>
|
||||||
|
#include <c10/cuda/CUDAException.h>
|
||||||
#include <cuda_bf16.h>
|
#include <cuda_bf16.h>
|
||||||
|
|
||||||
__global__ void rotary_emb_kernel(
|
__global__ void rotary_emb_kernel(
|
||||||
@@ -46,6 +48,9 @@ torch::Tensor rotary_emb(
|
|||||||
torch::Tensor x,
|
torch::Tensor x,
|
||||||
torch::Tensor freqs_cis
|
torch::Tensor freqs_cis
|
||||||
) {
|
) {
|
||||||
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(x));
|
||||||
|
auto stream = at::cuda::getCurrentCUDAStream();
|
||||||
|
|
||||||
TORCH_CHECK(x.is_cuda(), "x must be on CUDA");
|
TORCH_CHECK(x.is_cuda(), "x must be on CUDA");
|
||||||
TORCH_CHECK(freqs_cis.is_cuda(), "freqs_cis must be on CUDA");
|
TORCH_CHECK(freqs_cis.is_cuda(), "freqs_cis must be on CUDA");
|
||||||
TORCH_CHECK(x.scalar_type() == torch::kBFloat16, "x must be bf16");
|
TORCH_CHECK(x.scalar_type() == torch::kBFloat16, "x must be bf16");
|
||||||
@@ -53,6 +58,7 @@ torch::Tensor rotary_emb(
|
|||||||
TORCH_CHECK(x.is_contiguous(), "x must be contiguous");
|
TORCH_CHECK(x.is_contiguous(), "x must be contiguous");
|
||||||
TORCH_CHECK(freqs_cis.dim() == 4, "freqs_cis must be 4D [batch, seq_len, dim/2, 2]");
|
TORCH_CHECK(freqs_cis.dim() == 4, "freqs_cis must be 4D [batch, seq_len, dim/2, 2]");
|
||||||
TORCH_CHECK(freqs_cis.is_contiguous(), "freqs_cis must be contiguous");
|
TORCH_CHECK(freqs_cis.is_contiguous(), "freqs_cis must be contiguous");
|
||||||
|
TORCH_CHECK(freqs_cis.scalar_type() == torch::kFloat32, "freqs_cis must be f32");
|
||||||
|
|
||||||
int batch = x.size(0);
|
int batch = x.size(0);
|
||||||
int seq_len = x.size(1);
|
int seq_len = x.size(1);
|
||||||
@@ -60,6 +66,10 @@ torch::Tensor rotary_emb(
|
|||||||
int head_dim = x.size(3);
|
int head_dim = x.size(3);
|
||||||
|
|
||||||
TORCH_CHECK(head_dim % 2 == 0, "head_dim must be even");
|
TORCH_CHECK(head_dim % 2 == 0, "head_dim must be even");
|
||||||
|
TORCH_CHECK(freqs_cis.size(0) == batch, "freqs_cis batch mismatch");
|
||||||
|
TORCH_CHECK(freqs_cis.size(1) == seq_len, "freqs_cis seq_len mismatch");
|
||||||
|
TORCH_CHECK(freqs_cis.size(2) == head_dim / 2, "freqs_cis dim/2 mismatch");
|
||||||
|
TORCH_CHECK(freqs_cis.size(3) == 2, "freqs_cis last dim must be 2 [cos, sin]");
|
||||||
|
|
||||||
auto out = torch::empty_like(x);
|
auto out = torch::empty_like(x);
|
||||||
|
|
||||||
@@ -68,12 +78,13 @@ torch::Tensor rotary_emb(
|
|||||||
int block = 256;
|
int block = 256;
|
||||||
int grid = std::min((total + block - 1) / block, 1024);
|
int grid = std::min((total + block - 1) / block, 1024);
|
||||||
|
|
||||||
rotary_emb_kernel<<<grid, block>>>(
|
rotary_emb_kernel<<<grid, block, 0, stream>>>(
|
||||||
reinterpret_cast<const __nv_bfloat16*>(x.data_ptr()),
|
reinterpret_cast<const __nv_bfloat16*>(x.data_ptr()),
|
||||||
freqs_cis.data_ptr<float>(),
|
freqs_cis.data_ptr<float>(),
|
||||||
reinterpret_cast<__nv_bfloat16*>(out.data_ptr()),
|
reinterpret_cast<__nv_bfloat16*>(out.data_ptr()),
|
||||||
batch, seq_len, n_heads, head_dim
|
batch, seq_len, n_heads, head_dim
|
||||||
);
|
);
|
||||||
|
C10_CUDA_CHECK(cudaGetLastError());
|
||||||
|
|
||||||
return out;
|
return out;
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,185 +0,0 @@
|
|||||||
/*
|
|
||||||
Pure-C test — uses shared dispatcher.
|
|
||||||
nvcc -I csrc -arch=sm_89 -O3 \
|
|
||||||
--use_fast_math --ptxas-options=-O3 --extra-device-vectorization \
|
|
||||||
csrc/tests/attn_decode_test.cu -o test && ./test
|
|
||||||
*/
|
|
||||||
|
|
||||||
#include "test_utils.cuh"
|
|
||||||
#include "../kernels/attn_dispatchers.cuh"
|
|
||||||
|
|
||||||
// Split-K scratch (torch-free)
|
|
||||||
struct DecodeScratch {
|
|
||||||
float* o_part = nullptr;
|
|
||||||
float* ml_part = nullptr;
|
|
||||||
};
|
|
||||||
|
|
||||||
static void setup_scratch(AttentionParams<bf16>& p, DecodeScratch& sc) {
|
|
||||||
int max_splits = 32;
|
|
||||||
cudaMalloc(&sc.o_part, (size_t)p.batch * p.q_head * max_splits * p.head_dim * sizeof(float));
|
|
||||||
cudaMalloc(&sc.ml_part, (size_t)p.batch * p.q_head * max_splits * 2 * sizeof(float));
|
|
||||||
}
|
|
||||||
|
|
||||||
static void free_scratch(DecodeScratch& sc) {
|
|
||||||
cudaFree(sc.o_part); cudaFree(sc.ml_part);
|
|
||||||
}
|
|
||||||
|
|
||||||
// Warmed-up, CUDA-event timed sweep over the production decode MMA path.
|
|
||||||
static void bench() {
|
|
||||||
const int cfgs[][5] = {
|
|
||||||
{1, 32, 4, 512, 128},
|
|
||||||
{1, 32, 4, 1024, 128},
|
|
||||||
{1, 32, 4, 2048, 128},
|
|
||||||
{1, 32, 4, 4096, 128},
|
|
||||||
{16, 32, 4, 2048, 128},
|
|
||||||
{32, 32, 4, 1024, 128},
|
|
||||||
};
|
|
||||||
const int WARMUP = 10, ITERS = 100;
|
|
||||||
printf("\n===== DECODE BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS);
|
|
||||||
print_bench_header();
|
|
||||||
|
|
||||||
for (int ci = 0; ci < 6; ci++) {
|
|
||||||
int B = cfgs[ci][0], Hq = cfgs[ci][1], Hk = cfgs[ci][2];
|
|
||||||
int sl = cfgs[ci][3], D = cfgs[ci][4];
|
|
||||||
size_t nQ = (size_t)B * Hq * D;
|
|
||||||
size_t nKV = (size_t)B * Hk * sl * D;
|
|
||||||
|
|
||||||
bf16 *dQ, *dK, *dV, *dO;
|
|
||||||
cudaMalloc(&dQ, nQ*2); cudaMalloc(&dK, nKV*2);
|
|
||||||
cudaMalloc(&dV, nKV*2); cudaMalloc(&dO, nQ*2);
|
|
||||||
size_t big = nQ > nKV ? nQ : nKV; bf16* tmp = new bf16[big];
|
|
||||||
for (size_t i = 0; i < nQ; i++) tmp[i] = f2bf(randf());
|
|
||||||
cudaMemcpy(dQ, tmp, nQ*2, cudaMemcpyHostToDevice);
|
|
||||||
for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf());
|
|
||||||
cudaMemcpy(dK, tmp, nKV*2, cudaMemcpyHostToDevice);
|
|
||||||
for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf());
|
|
||||||
cudaMemcpy(dV, tmp, nKV*2, cudaMemcpyHostToDevice);
|
|
||||||
delete[] tmp;
|
|
||||||
|
|
||||||
AttentionParams<bf16> p;
|
|
||||||
p.batch = B; p.q_head = Hq; p.kv_head = Hk; p.q_len = 1; p.kv_len = sl;
|
|
||||||
p.head_dim = D; p.use_mask = 0; p.causal_offset = -1;
|
|
||||||
p.scale = 1.0f / sqrtf((float)D);
|
|
||||||
set_default_strides(p);
|
|
||||||
p.q = dQ; p.k = dK; p.v = dV; p.mask = nullptr; p.o = dO;
|
|
||||||
|
|
||||||
DecodeScratch sc;
|
|
||||||
setup_scratch(p, sc);
|
|
||||||
p.o_part = sc.o_part; p.ml_part = sc.ml_part;
|
|
||||||
|
|
||||||
auto launch = [&]() { dispatch_by_head_dim(D, [&]<int H>() { dispatch_decode<H>(p); }); };
|
|
||||||
double flops = 4.0 * B * Hq * (double)sl * D;
|
|
||||||
double bytes = 2.0 * (2.0 * nKV * sizeof(bf16));
|
|
||||||
BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops, bytes);
|
|
||||||
|
|
||||||
char cfg[64];
|
|
||||||
snprintf(cfg, sizeof(cfg),
|
|
||||||
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d causal=%d",
|
|
||||||
B, Hq, Hk, 1, sl, D, 0);
|
|
||||||
print_bench_row(cfg, r);
|
|
||||||
|
|
||||||
cudaFree(dQ); cudaFree(dK); cudaFree(dV); cudaFree(dO);
|
|
||||||
free_scratch(sc);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
static int run_test(int B, int Hq, int Hk, int sl, int D, int causal) {
|
|
||||||
int gs = Hq / Hk;
|
|
||||||
printf("=== B=%d Hq=%d Hk=%d seq=%d D=%d gs=%d causal=%d ===\n",
|
|
||||||
B,Hq,Hk,sl,D,gs,causal);
|
|
||||||
|
|
||||||
size_t nQ = B*Hq*1*D, nKV = B*Hk*sl*D;
|
|
||||||
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
|
|
||||||
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
|
|
||||||
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
|
|
||||||
|
|
||||||
bool* hMask=new bool[B*sl];
|
|
||||||
for (int i=0;i<B*sl;i++) hMask[i]=true;
|
|
||||||
|
|
||||||
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
|
||||||
bool* dMask;
|
|
||||||
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
|
||||||
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
|
||||||
cudaMalloc(&dMask,B*sl);
|
|
||||||
|
|
||||||
tmp=new bf16[max(nQ,nKV)];
|
|
||||||
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
|
|
||||||
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
|
||||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
|
|
||||||
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
|
||||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
|
|
||||||
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
|
||||||
cudaMemcpy(dMask,hMask,B*sl,cudaMemcpyHostToDevice);
|
|
||||||
|
|
||||||
AttentionParams<bf16> p;
|
|
||||||
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=1; p.kv_len=sl; p.head_dim=D;
|
|
||||||
p.use_mask=0; p.causal_offset=causal?0:-1;
|
|
||||||
p.scale=1.0f/sqrtf((float)D);
|
|
||||||
set_default_strides(p);
|
|
||||||
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
|
||||||
|
|
||||||
DecodeScratch sc;
|
|
||||||
setup_scratch(p, sc);
|
|
||||||
p.o_part = sc.o_part; p.ml_part = sc.ml_part;
|
|
||||||
|
|
||||||
double t0=now_ms();
|
|
||||||
dispatch_by_head_dim(D, [&]<int H>() { dispatch_decode<H>(p); });
|
|
||||||
cudaDeviceSynchronize();
|
|
||||||
double kms=now_ms()-t0;
|
|
||||||
cudaError_t err=cudaGetLastError();
|
|
||||||
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
|
|
||||||
|
|
||||||
bf16* hOut=new bf16[nQ];
|
|
||||||
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
|
|
||||||
|
|
||||||
float* ref=new float[nQ];
|
|
||||||
cpu_attention_ref(hQ, hK, hV, hMask, ref, B, Hq, Hk, 1, sl, D, causal ? 0 : -1);
|
|
||||||
|
|
||||||
float max_abs_err=0, max_rel_err=0;
|
|
||||||
for (size_t i=0;i<nQ;i++){
|
|
||||||
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
|
||||||
if(err>max_abs_err) max_abs_err=err;
|
|
||||||
float rel=err/fmaxf(fabsf(ref[i]), 1e-8f);
|
|
||||||
if(rel>max_rel_err) max_rel_err=rel;
|
|
||||||
}
|
|
||||||
const float atol=0.01f, rtol=0.01f;
|
|
||||||
bool pass=true;
|
|
||||||
for (size_t i=0;i<nQ;i++){
|
|
||||||
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
|
||||||
if (err > atol + rtol * fabsf(ref[i])) { pass=false; break; }
|
|
||||||
}
|
|
||||||
printf("kernel: %.3f ms max_abs_err: %.6e max_rel_err: %.6e %s\n\n",
|
|
||||||
kms, max_abs_err, max_rel_err, pass?"PASS":"FAIL");
|
|
||||||
|
|
||||||
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);cudaFree(dMask);
|
|
||||||
free_scratch(sc);
|
|
||||||
delete[]hQ;delete[]hK;delete[]hV;delete[]hMask;delete[]hOut;delete[]ref;delete[]tmp;
|
|
||||||
|
|
||||||
return pass ? 0 : 1;
|
|
||||||
}
|
|
||||||
|
|
||||||
int main() {
|
|
||||||
const int configs[][6] = {
|
|
||||||
{1, 2, 1, 64, 32, 0},
|
|
||||||
{1, 32, 4, 512, 128, 0},
|
|
||||||
{1, 32, 4, 1024, 128, 0},
|
|
||||||
{1, 32, 4, 512, 128, 1},
|
|
||||||
};
|
|
||||||
int n_cfgs = sizeof(configs) / sizeof(configs[0]);
|
|
||||||
int fail = 0;
|
|
||||||
|
|
||||||
for (int ci = 0; ci < n_cfgs; ci++) {
|
|
||||||
int B = configs[ci][0], Hq = configs[ci][1], Hk = configs[ci][2];
|
|
||||||
int sl = configs[ci][3], D = configs[ci][4], causal = configs[ci][5];
|
|
||||||
fail += run_test(B, Hq, Hk, sl, D, causal);
|
|
||||||
if (fail) break;
|
|
||||||
}
|
|
||||||
|
|
||||||
if (fail) {
|
|
||||||
printf("FAILED\n");
|
|
||||||
return fail;
|
|
||||||
}
|
|
||||||
printf("All tests passed!\n");
|
|
||||||
bench();
|
|
||||||
return 0;
|
|
||||||
}
|
|
||||||
@@ -1,308 +0,0 @@
|
|||||||
// Compile:
|
|
||||||
// nvcc -I csrc -arch=sm_89 -O3 --use_fast_math --ptxas-options=-O3 \
|
|
||||||
// --extra-device-vectorization csrc/tests/attn_paged_decode_test.cu \
|
|
||||||
// -o /tmp/test_paged && /tmp/test_paged
|
|
||||||
|
|
||||||
#include <cstring>
|
|
||||||
#include "test_utils.cuh"
|
|
||||||
#include "../kernels/attn_dispatchers.cuh"
|
|
||||||
|
|
||||||
static void gather_kv_cpu(
|
|
||||||
const bf16* h_k_pool, const bf16* h_v_pool,
|
|
||||||
const int64_t* h_pt, int B, int Hkv, int kv_len,
|
|
||||||
int page_size, int head_dim,
|
|
||||||
bf16* h_k, bf16* h_v)
|
|
||||||
{
|
|
||||||
int max_pages = (kv_len + page_size - 1) / page_size;
|
|
||||||
size_t page_stride = (size_t)page_size * Hkv * head_dim;
|
|
||||||
for (int b = 0; b < B; b++) {
|
|
||||||
for (int pos = 0; pos < kv_len; pos++) {
|
|
||||||
int log_pg = pos / page_size;
|
|
||||||
int pg_off = pos % page_size;
|
|
||||||
int phys = (int)h_pt[b * max_pages + log_pg];
|
|
||||||
for (int h = 0; h < Hkv; h++) {
|
|
||||||
size_t src_base = (size_t)phys * page_stride
|
|
||||||
+ (size_t)pg_off * Hkv * head_dim
|
|
||||||
+ h * head_dim;
|
|
||||||
size_t dst_base = ((size_t)b * Hkv + h) * kv_len * head_dim
|
|
||||||
+ (size_t)pos * head_dim;
|
|
||||||
memcpy(h_k + dst_base, h_k_pool + src_base, head_dim * sizeof(bf16));
|
|
||||||
memcpy(h_v + dst_base, h_v_pool + src_base, head_dim * sizeof(bf16));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
template <int HEAD_DIM>
|
|
||||||
static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int causal, int seed) {
|
|
||||||
printf("B=%d Hq=%d Hkv=%d kv_len=%d page_sz=%d head_dim=%d causal=%d ... ",
|
|
||||||
B, Hq, Hkv, kv_len, page_size, HEAD_DIM, causal);
|
|
||||||
fflush(stdout);
|
|
||||||
|
|
||||||
int max_pages = (kv_len + page_size - 1) / page_size;
|
|
||||||
int n_phys_pages = B * max_pages;
|
|
||||||
int max_splits = 32;
|
|
||||||
|
|
||||||
size_t sz_q = (size_t)B * Hq * 1 * HEAD_DIM * sizeof(bf16);
|
|
||||||
size_t sz_o = sz_q;
|
|
||||||
size_t sz_kv = (size_t)n_phys_pages * page_size * Hkv * HEAD_DIM * sizeof(bf16);
|
|
||||||
size_t sz_pt = (size_t)B * max_pages * sizeof(int64_t);
|
|
||||||
size_t sz_op = (size_t)B * Hq * max_splits * HEAD_DIM * sizeof(float);
|
|
||||||
size_t sz_ml = (size_t)B * Hq * max_splits * 2 * sizeof(float);
|
|
||||||
|
|
||||||
bf16 *d_q, *d_o_paged;
|
|
||||||
bf16 *d_k_pool, *d_v_pool;
|
|
||||||
int64_t* d_pt;
|
|
||||||
float *d_op, *d_ml;
|
|
||||||
|
|
||||||
cudaMalloc(&d_q, sz_q);
|
|
||||||
cudaMalloc(&d_o_paged, sz_o);
|
|
||||||
cudaMalloc(&d_k_pool, sz_kv);
|
|
||||||
cudaMalloc(&d_v_pool, sz_kv);
|
|
||||||
cudaMalloc(&d_pt, sz_pt);
|
|
||||||
cudaMalloc(&d_op, sz_op);
|
|
||||||
cudaMalloc(&d_ml, sz_ml);
|
|
||||||
|
|
||||||
srand(seed);
|
|
||||||
auto rnd = [&]() { return (rand() / (float)RAND_MAX) * 2.0f - 1.0f; };
|
|
||||||
|
|
||||||
bf16* h_q = (bf16*)malloc(sz_q);
|
|
||||||
for (int i = 0; i < B * Hq * HEAD_DIM; i++)
|
|
||||||
h_q[i] = __float2bfloat16(rnd());
|
|
||||||
cudaMemcpy(d_q, h_q, sz_q, cudaMemcpyHostToDevice);
|
|
||||||
|
|
||||||
bf16* h_k_pool = (bf16*)malloc(sz_kv);
|
|
||||||
bf16* h_v_pool = (bf16*)malloc(sz_kv);
|
|
||||||
size_t ps = (size_t)page_size * Hkv * HEAD_DIM;
|
|
||||||
for (int pg = 0; pg < n_phys_pages; pg++) {
|
|
||||||
for (int off = 0; off < page_size; off++) {
|
|
||||||
for (int h = 0; h < Hkv; h++) {
|
|
||||||
for (int d = 0; d < HEAD_DIM; d++) {
|
|
||||||
float v = sinf((float)(pg * 7919 + off * 1049 + h * 331 + d));
|
|
||||||
size_t idx = (size_t)pg * ps + (size_t)off * Hkv * HEAD_DIM
|
|
||||||
+ h * HEAD_DIM + d;
|
|
||||||
h_k_pool[idx] = __float2bfloat16(v);
|
|
||||||
h_v_pool[idx] = __float2bfloat16(v * 0.3f);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
|
|
||||||
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
|
|
||||||
|
|
||||||
int64_t* h_pt = (int64_t*)malloc(sz_pt);
|
|
||||||
int next_pg = 0;
|
|
||||||
for (int b = 0; b < B; b++)
|
|
||||||
for (int p = 0; p < max_pages; p++)
|
|
||||||
h_pt[b * max_pages + p] = next_pg++;
|
|
||||||
cudaMemcpy(d_pt, h_pt, sz_pt, cudaMemcpyHostToDevice);
|
|
||||||
|
|
||||||
bf16* h_k_cont = (bf16*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(bf16));
|
|
||||||
bf16* h_v_cont = (bf16*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(bf16));
|
|
||||||
gather_kv_cpu(h_k_pool, h_v_pool, h_pt, B, Hkv, kv_len, page_size, HEAD_DIM, h_k_cont, h_v_cont);
|
|
||||||
|
|
||||||
float* h_q_f = (float*)malloc((size_t)B * Hq * HEAD_DIM * sizeof(float));
|
|
||||||
float* h_k_f = (float*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(float));
|
|
||||||
float* h_v_f = (float*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(float));
|
|
||||||
for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_q_f[i] = bf2f(h_q[i]);
|
|
||||||
for (int i = 0; i < B * kv_len * Hkv * HEAD_DIM; i++) {
|
|
||||||
h_k_f[i] = bf2f(h_k_cont[i]);
|
|
||||||
h_v_f[i] = bf2f(h_v_cont[i]);
|
|
||||||
}
|
|
||||||
|
|
||||||
float* h_o_ref = (float*)calloc(B * Hq * HEAD_DIM, sizeof(float));
|
|
||||||
cpu_attention_ref(h_q_f, h_k_f, h_v_f, nullptr, h_o_ref, B, Hq, Hkv,
|
|
||||||
1, kv_len, HEAD_DIM, causal ? 0 : -1);
|
|
||||||
|
|
||||||
PagedAttentionParams<bf16> p;
|
|
||||||
p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.q_len = 1;
|
|
||||||
p.kv_len = kv_len; p.head_dim = HEAD_DIM;
|
|
||||||
p.use_mask = 0; p.causal_offset = causal ? 0 : -1;
|
|
||||||
set_default_paged_strides(p);
|
|
||||||
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
|
|
||||||
p.page_size = page_size; p.max_pages = max_pages;
|
|
||||||
p.page_table = d_pt;
|
|
||||||
p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
|
||||||
p.q = d_q; p.mask = nullptr; p.o = d_o_paged;
|
|
||||||
p.o_part = d_op; p.ml_part = d_ml;
|
|
||||||
|
|
||||||
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_decode<H>(p); });
|
|
||||||
cudaDeviceSynchronize();
|
|
||||||
|
|
||||||
bf16* h_o_bf16 = (bf16*)malloc(sz_o);
|
|
||||||
cudaMemcpy(h_o_bf16, d_o_paged, sz_o, cudaMemcpyDeviceToHost);
|
|
||||||
float* h_o_paged = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float));
|
|
||||||
for (int i = 0; i < B * Hq * HEAD_DIM; i++)
|
|
||||||
h_o_paged[i] = __bfloat162float(h_o_bf16[i]);
|
|
||||||
|
|
||||||
float max_abs_err = 0.0f, max_rel_err = 0.0f;
|
|
||||||
int bad_idx = -1;
|
|
||||||
for (int i = 0; i < B * Hq * HEAD_DIM; i++) {
|
|
||||||
float e = fabsf(h_o_paged[i] - h_o_ref[i]);
|
|
||||||
if (e > max_abs_err) { max_abs_err = e; bad_idx = i; }
|
|
||||||
float rel = e / fmaxf(fabsf(h_o_ref[i]), 1e-8f);
|
|
||||||
if (rel > max_rel_err) max_rel_err = rel;
|
|
||||||
}
|
|
||||||
|
|
||||||
const float atol = 0.01f, rtol = 0.01f;
|
|
||||||
bool pass = true;
|
|
||||||
for (int i = 0; i < B * Hq * HEAD_DIM; i++) {
|
|
||||||
float e = fabsf(h_o_paged[i] - h_o_ref[i]);
|
|
||||||
if (e > atol + rtol * fabsf(h_o_ref[i])) { pass = false; break; }
|
|
||||||
}
|
|
||||||
|
|
||||||
if (pass) {
|
|
||||||
printf("PASS (max_abs_err=%.4e max_rel_err=%.4e)\n", max_abs_err, max_rel_err);
|
|
||||||
} else {
|
|
||||||
int b = bad_idx / (Hq * HEAD_DIM);
|
|
||||||
int h = (bad_idx / HEAD_DIM) % Hq;
|
|
||||||
int d = bad_idx % HEAD_DIM;
|
|
||||||
printf("FAIL (max_abs_err=%.4e max_rel_err=%.4e at [%d,%d,%d]: ref=%.4f got=%.4f)\n",
|
|
||||||
max_abs_err, max_rel_err, b, h, d, h_o_ref[bad_idx], h_o_paged[bad_idx]);
|
|
||||||
printf(" ref[0..7]:");
|
|
||||||
for (int i = 0; i < 8 && i < HEAD_DIM; i++)
|
|
||||||
printf(" %.4f", h_o_ref[i]);
|
|
||||||
printf("\n got[0..7]:");
|
|
||||||
for (int i = 0; i < 8 && i < HEAD_DIM; i++)
|
|
||||||
printf(" %.4f", h_o_paged[i]);
|
|
||||||
printf("\n");
|
|
||||||
}
|
|
||||||
|
|
||||||
free(h_q); free(h_k_pool); free(h_v_pool); free(h_pt);
|
|
||||||
free(h_k_cont); free(h_v_cont);
|
|
||||||
free(h_q_f); free(h_k_f); free(h_v_f);
|
|
||||||
free(h_o_ref); free(h_o_bf16); free(h_o_paged);
|
|
||||||
cudaFree(d_q); cudaFree(d_o_paged);
|
|
||||||
cudaFree(d_k_pool); cudaFree(d_v_pool); cudaFree(d_pt);
|
|
||||||
cudaFree(d_op); cudaFree(d_ml);
|
|
||||||
|
|
||||||
return pass ? 0 : 1;
|
|
||||||
}
|
|
||||||
|
|
||||||
struct TestCase {
|
|
||||||
int head_dim;
|
|
||||||
int B, Hq, Hkv, kv_len, page_size, causal, seed;
|
|
||||||
};
|
|
||||||
|
|
||||||
static const TestCase TESTS[] = {
|
|
||||||
{128, 1, 1, 1, 8, 128, 0, 1},
|
|
||||||
{128, 1, 4, 4, 128, 128, 0, 2},
|
|
||||||
{128, 2, 4, 4, 256, 128, 0, 3},
|
|
||||||
{128, 1, 4, 1, 64, 64, 0, 4},
|
|
||||||
{128, 1, 8, 2, 64, 128, 0, 5},
|
|
||||||
{128, 2, 16, 4, 128, 128, 0, 6},
|
|
||||||
{64, 1, 4, 2, 32, 128, 0, 7},
|
|
||||||
{256, 1, 2, 1, 16, 128, 0, 8},
|
|
||||||
{32, 1, 4, 2, 32, 64, 0, 9},
|
|
||||||
{128, 3, 8, 2, 256, 128, 0, 10},
|
|
||||||
{128, 2, 32, 8, 512, 128, 0, 11},
|
|
||||||
{128, 1, 16, 2, 256, 128, 0, 12},
|
|
||||||
{128, 2, 32, 4, 512, 128, 0, 13},
|
|
||||||
{128, 2, 8, 2, 128, 128, 1, 14}, // causal
|
|
||||||
};
|
|
||||||
|
|
||||||
static int dispatch_test(const TestCase& tc) {
|
|
||||||
int r = 0;
|
|
||||||
dispatch_by_head_dim(tc.head_dim, [&]<int D>() {
|
|
||||||
r = run_test<D>(tc.B, tc.Hq, tc.Hkv, tc.kv_len, tc.page_size, tc.causal, tc.seed);
|
|
||||||
});
|
|
||||||
return r;
|
|
||||||
}
|
|
||||||
|
|
||||||
template <int HEAD_DIM>
|
|
||||||
static void bench_config(int B, int Hq, int Hkv, int kv_len, int page_size) {
|
|
||||||
int max_pages = (kv_len + page_size - 1) / page_size;
|
|
||||||
int n_phys_pages = B * max_pages;
|
|
||||||
int max_splits = 32;
|
|
||||||
|
|
||||||
size_t sz_q = (size_t)B * Hq * 1 * HEAD_DIM * sizeof(bf16);
|
|
||||||
size_t sz_kv = (size_t)n_phys_pages * page_size * Hkv * HEAD_DIM * sizeof(bf16);
|
|
||||||
size_t sz_pt = (size_t)B * max_pages * sizeof(int64_t);
|
|
||||||
size_t sz_op = (size_t)B * Hq * max_splits * HEAD_DIM * sizeof(float);
|
|
||||||
size_t sz_ml = (size_t)B * Hq * max_splits * 2 * sizeof(float);
|
|
||||||
|
|
||||||
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
|
|
||||||
int64_t* d_pt;
|
|
||||||
float *d_op, *d_ml;
|
|
||||||
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
|
|
||||||
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
|
|
||||||
cudaMalloc(&d_pt, sz_pt);
|
|
||||||
cudaMalloc(&d_op, sz_op); cudaMalloc(&d_ml, sz_ml);
|
|
||||||
|
|
||||||
bf16* tmp = (bf16*)malloc(sz_kv > sz_q ? sz_kv : sz_q);
|
|
||||||
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) tmp[i] = f2bf(randf());
|
|
||||||
cudaMemcpy(d_q, tmp, sz_q, cudaMemcpyHostToDevice);
|
|
||||||
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) tmp[i] = f2bf(randf());
|
|
||||||
cudaMemcpy(d_k_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
|
|
||||||
cudaMemcpy(d_v_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
|
|
||||||
|
|
||||||
int64_t* h_pt = (int64_t*)malloc(sz_pt);
|
|
||||||
int next_pg = 0;
|
|
||||||
for (int b = 0; b < B; b++)
|
|
||||||
for (int p = 0; p < max_pages; p++)
|
|
||||||
h_pt[b * max_pages + p] = next_pg++;
|
|
||||||
cudaMemcpy(d_pt, h_pt, sz_pt, cudaMemcpyHostToDevice);
|
|
||||||
free(h_pt);
|
|
||||||
|
|
||||||
PagedAttentionParams<bf16> pa;
|
|
||||||
pa.batch = B; pa.q_head = Hq; pa.kv_head = Hkv; pa.q_len = 1;
|
|
||||||
pa.kv_len = kv_len; pa.head_dim = HEAD_DIM;
|
|
||||||
pa.use_mask = 0; pa.causal_offset = -1;
|
|
||||||
set_default_paged_strides(pa);
|
|
||||||
pa.scale = 1.0f / sqrtf((float)HEAD_DIM);
|
|
||||||
pa.page_size = page_size; pa.max_pages = max_pages;
|
|
||||||
pa.page_table = d_pt;
|
|
||||||
pa.k_cache = d_k_pool; pa.v_cache = d_v_pool;
|
|
||||||
pa.q = d_q; pa.mask = nullptr; pa.o = d_o;
|
|
||||||
pa.o_part = d_op; pa.ml_part = d_ml;
|
|
||||||
|
|
||||||
const int WARMUP = 10, ITERS = 100;
|
|
||||||
auto launch = [&]() {
|
|
||||||
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_decode<H>(pa); });
|
|
||||||
};
|
|
||||||
double flops = 4.0 * B * Hq * (double)kv_len * HEAD_DIM;
|
|
||||||
size_t nKV = (size_t)B * Hkv * kv_len * HEAD_DIM;
|
|
||||||
double bytes = 2.0 * (2.0 * nKV * sizeof(bf16));
|
|
||||||
BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops, bytes);
|
|
||||||
|
|
||||||
char cfg[64];
|
|
||||||
snprintf(cfg, sizeof(cfg),
|
|
||||||
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d page=%3d",
|
|
||||||
B, Hq, Hkv, 1, kv_len, HEAD_DIM, page_size);
|
|
||||||
print_bench_row(cfg, r);
|
|
||||||
|
|
||||||
free(tmp);
|
|
||||||
cudaFree(d_q); cudaFree(d_o);
|
|
||||||
cudaFree(d_k_pool); cudaFree(d_v_pool); cudaFree(d_pt);
|
|
||||||
cudaFree(d_op); cudaFree(d_ml);
|
|
||||||
}
|
|
||||||
|
|
||||||
static void bench() {
|
|
||||||
printf("\n===== PAGED DECODE BENCH =====\n");
|
|
||||||
print_bench_header();
|
|
||||||
bench_config<128>(1, 32, 4, 512, 128);
|
|
||||||
bench_config<128>(1, 32, 4, 1024, 128);
|
|
||||||
bench_config<128>(1, 32, 4, 2048, 128);
|
|
||||||
bench_config<128>(1, 32, 4, 4096, 128);
|
|
||||||
bench_config<128>(16, 32, 4, 2048, 128);
|
|
||||||
bench_config<128>(32, 32, 4, 1024, 128);
|
|
||||||
}
|
|
||||||
|
|
||||||
int main() {
|
|
||||||
int n = sizeof(TESTS) / sizeof(TESTS[0]);
|
|
||||||
int fail = 0;
|
|
||||||
printf("=== Paged Decode vs CPU reference (%d cases) ===\n\n", n);
|
|
||||||
|
|
||||||
for (int i = 0; i < n; i++) {
|
|
||||||
fail += dispatch_test(TESTS[i]);
|
|
||||||
if (fail) break;
|
|
||||||
}
|
|
||||||
|
|
||||||
if (fail) {
|
|
||||||
printf("\nFAILED (%d/%d tests failed)\n", fail, n);
|
|
||||||
return fail;
|
|
||||||
}
|
|
||||||
printf("\nAll %d tests passed!\n", n);
|
|
||||||
bench();
|
|
||||||
return 0;
|
|
||||||
}
|
|
||||||
@@ -0,0 +1,945 @@
|
|||||||
|
// Compile:
|
||||||
|
// nvcc -I csrc -arch=sm_89 -O3 --use_fast_math --ptxas-options=-O3 \
|
||||||
|
// --extra-device-vectorization -Xcompiler -fopenmp \
|
||||||
|
// csrc/tests/attn_paged_test.cu \
|
||||||
|
// -o /tmp/test_paged && /tmp/test_paged
|
||||||
|
|
||||||
|
#include <cstring>
|
||||||
|
#include <vector>
|
||||||
|
#include "test_utils.cuh"
|
||||||
|
#include "../kernels/attn_dispatchers.cuh"
|
||||||
|
|
||||||
|
// ---- CPU reference: paged decode with variable seq_lens ----
|
||||||
|
// Q: [B, Hq, D], K/V pool: [pool_size, Hkv, D]
|
||||||
|
// req_to_token: [num_reqs, max_ctx_len], req_pool_indices: [B]
|
||||||
|
// kv_indptr: [B+1]. mask: [B, max_seq_len] bool (True=keep) or NULL.
|
||||||
|
static void cpu_paged_decode_ref(
|
||||||
|
const float* Q, const float* K_pool, const float* V_pool,
|
||||||
|
const int64_t* req_to_token, const int64_t* req_pool_indices,
|
||||||
|
const int* kv_indptr, const bool* mask, int mask_b_stride,
|
||||||
|
int B, int Hq, int Hkv, int D, int max_ctx_len,
|
||||||
|
float* O)
|
||||||
|
{
|
||||||
|
float scale = 1.0f / sqrtf((float)D);
|
||||||
|
int n_rep = Hq / Hkv;
|
||||||
|
for (int b = 0; b < B; b++) {
|
||||||
|
int seq_len = kv_indptr[b + 1] - kv_indptr[b];
|
||||||
|
int64_t req_idx = req_pool_indices[b];
|
||||||
|
#pragma omp parallel for schedule(dynamic)
|
||||||
|
for (int h = 0; h < Hq; h++) {
|
||||||
|
int kv_h = h / n_rep;
|
||||||
|
float mv = -INFINITY, sv = 0.0f;
|
||||||
|
float accum[256] = {0.0f};
|
||||||
|
for (int kj = 0; kj < seq_len; kj++) {
|
||||||
|
if (mask && !mask[b * mask_b_stride + kj]) continue;
|
||||||
|
int64_t slot = req_to_token[req_idx * max_ctx_len + kj];
|
||||||
|
float dot = 0.0f;
|
||||||
|
for (int d = 0; d < D; d++)
|
||||||
|
dot += Q[(b * Hq + h) * D + d] *
|
||||||
|
K_pool[slot * Hkv * D + kv_h * D + d];
|
||||||
|
dot *= scale;
|
||||||
|
float nm = fmaxf(mv, dot);
|
||||||
|
float a = expf(mv - nm);
|
||||||
|
float be = expf(dot - nm);
|
||||||
|
sv = sv * a + be;
|
||||||
|
for (int d = 0; d < D; d++)
|
||||||
|
accum[d] = accum[d] * a +
|
||||||
|
V_pool[slot * Hkv * D + kv_h * D + d] * be;
|
||||||
|
mv = nm;
|
||||||
|
}
|
||||||
|
float inv = 1.0f / sv;
|
||||||
|
for (int d = 0; d < D; d++)
|
||||||
|
O[(b * Hq + h) * D + d] = accum[d] * inv;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- CPU reference: paged prefill with ragged batch ----
|
||||||
|
// Q: [total_q, Hq, D], K/V pool: [pool_size, Hkv, D]
|
||||||
|
// req_to_token: [num_reqs, max_ctx_len], req_pool_indices: [B]
|
||||||
|
// kv_indptr: [B+1], qo_indptr: [B+1].
|
||||||
|
// mask: [B, max_q_len, max_seq_len] bool (True=keep, q-local + kv-local
|
||||||
|
// positions) or NULL. Used only when causal==0 to apply an arbitrary
|
||||||
|
// attention mask on top of the (unused) causal logic.
|
||||||
|
static void cpu_paged_prefill_ref(
|
||||||
|
const float* Q, const float* K_pool, const float* V_pool,
|
||||||
|
const int64_t* req_to_token, const int64_t* req_pool_indices,
|
||||||
|
const int* kv_indptr, const int* qo_indptr,
|
||||||
|
const bool* mask, int mask_q_stride, int mask_kv_stride,
|
||||||
|
int B, int Hq, int Hkv, int D, int max_ctx_len, int causal,
|
||||||
|
float* O)
|
||||||
|
{
|
||||||
|
float scale = 1.0f / sqrtf((float)D);
|
||||||
|
int n_rep = Hq / Hkv;
|
||||||
|
for (int b = 0; b < B; b++) {
|
||||||
|
int seq_len = kv_indptr[b + 1] - kv_indptr[b];
|
||||||
|
int q_len = qo_indptr[b + 1] - qo_indptr[b];
|
||||||
|
int causal_off = seq_len - q_len;
|
||||||
|
int64_t req_idx = req_pool_indices[b];
|
||||||
|
#pragma omp parallel for collapse(2) schedule(dynamic)
|
||||||
|
for (int h = 0; h < Hq; h++) {
|
||||||
|
for (int qi = 0; qi < q_len; qi++) {
|
||||||
|
int kv_h = h / n_rep;
|
||||||
|
float mv = -INFINITY, sv = 0.0f;
|
||||||
|
float accum[256] = {0.0f};
|
||||||
|
int lim = causal ? min(seq_len, causal_off + qi + 1) : seq_len;
|
||||||
|
for (int kj = 0; kj < lim; kj++) {
|
||||||
|
if (mask && !mask[b * mask_q_stride * mask_kv_stride
|
||||||
|
+ qi * mask_kv_stride + kj]) continue;
|
||||||
|
int64_t slot = req_to_token[req_idx * max_ctx_len + kj];
|
||||||
|
float dot = 0.0f;
|
||||||
|
for (int d = 0; d < D; d++)
|
||||||
|
dot += Q[(qo_indptr[b] + qi) * Hq * D + h * D + d] *
|
||||||
|
K_pool[slot * Hkv * D + kv_h * D + d];
|
||||||
|
dot *= scale;
|
||||||
|
float nm = fmaxf(mv, dot);
|
||||||
|
float a = expf(mv - nm);
|
||||||
|
float be = expf(dot - nm);
|
||||||
|
sv = sv * a + be;
|
||||||
|
for (int d = 0; d < D; d++)
|
||||||
|
accum[d] = accum[d] * a +
|
||||||
|
V_pool[slot * Hkv * D + kv_h * D + d] * be;
|
||||||
|
mv = nm;
|
||||||
|
}
|
||||||
|
float inv = 1.0f / sv;
|
||||||
|
for (int d = 0; d < D; d++)
|
||||||
|
O[(qo_indptr[b] + qi) * Hq * D + h * D + d] = accum[d] * inv;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- paged validation table (kernel vs CPU ref, abs error only) ----
|
||||||
|
inline void print_paged_header() {
|
||||||
|
printf("%-58s | %11s | %6s\n",
|
||||||
|
"config", "max_err", "result");
|
||||||
|
printf("----------------------------------------------------------------"
|
||||||
|
"--------------------------------\n");
|
||||||
|
}
|
||||||
|
|
||||||
|
inline void print_paged_row(const char* cfg, float max_err, bool pass) {
|
||||||
|
printf("%-58s | %11.3e | %s\n",
|
||||||
|
cfg, max_err, pass ? "PASS" : "FAIL");
|
||||||
|
}
|
||||||
|
|
||||||
|
// ======================================================================
|
||||||
|
// DECODE TEST
|
||||||
|
// ======================================================================
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
static int run_decode_test(int B, int Hq, int Hkv, int max_seq,
|
||||||
|
int causal, int seed) {
|
||||||
|
// Variable seq_lens per request
|
||||||
|
srand(seed);
|
||||||
|
std::vector<int> seq_lens(B);
|
||||||
|
for (int b = 0; b < B; b++)
|
||||||
|
seq_lens[b] = 8 + rand() % (max_seq - 8);
|
||||||
|
int max_sl = *std::max_element(seq_lens.begin(), seq_lens.end());
|
||||||
|
int max_ctx = max_sl + 16;
|
||||||
|
|
||||||
|
int pool_size = B * max_ctx;
|
||||||
|
int num_reqs = B + 4;
|
||||||
|
|
||||||
|
char cfg[80];
|
||||||
|
snprintf(cfg, sizeof(cfg), "DECODE B=%d Hq=%d Hkv=%d D=%d max_sl=%d causal=%d",
|
||||||
|
B, Hq, Hkv, HEAD_DIM, max_sl, causal);
|
||||||
|
|
||||||
|
size_t sz_q = (size_t)B * Hq * HEAD_DIM * sizeof(bf16);
|
||||||
|
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
|
||||||
|
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t);
|
||||||
|
size_t sz_rpi = (size_t)B * sizeof(int64_t);
|
||||||
|
size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
|
||||||
|
size_t sz_op = (size_t)B * Hq * MAX_SPLITS * HEAD_DIM * sizeof(float);
|
||||||
|
size_t sz_ml = (size_t)B * Hq * MAX_SPLITS * 2 * sizeof(float);
|
||||||
|
|
||||||
|
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
|
||||||
|
int64_t *d_rtt, *d_rpi;
|
||||||
|
int *d_kvi;
|
||||||
|
float *d_op, *d_ml;
|
||||||
|
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
|
||||||
|
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
|
||||||
|
cudaMalloc(&d_rtt, sz_rtt); cudaMalloc(&d_rpi, sz_rpi);
|
||||||
|
cudaMalloc(&d_kvi, sz_kvi);
|
||||||
|
cudaMalloc(&d_op, sz_op); cudaMalloc(&d_ml, sz_ml);
|
||||||
|
|
||||||
|
auto rnd = [&]() { return (rand() / (float)RAND_MAX) * 2.0f - 1.0f; };
|
||||||
|
|
||||||
|
bf16* h_q = (bf16*)malloc(sz_q);
|
||||||
|
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) h_q[i] = f2bf(rnd());
|
||||||
|
cudaMemcpy(d_q, h_q, sz_q, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
bf16* h_k_pool = (bf16*)malloc(sz_kv);
|
||||||
|
bf16* h_v_pool = (bf16*)malloc(sz_kv);
|
||||||
|
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) {
|
||||||
|
h_k_pool[i] = f2bf(rnd());
|
||||||
|
h_v_pool[i] = f2bf(rnd());
|
||||||
|
}
|
||||||
|
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
|
||||||
|
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
// req_to_token: assign unique slots per request (scattered, not contiguous)
|
||||||
|
int64_t* h_rtt = (int64_t*)malloc(sz_rtt);
|
||||||
|
int next_slot = 0;
|
||||||
|
for (int r = 0; r < num_reqs; r++)
|
||||||
|
for (int p = 0; p < max_ctx; p++) {
|
||||||
|
h_rtt[r * max_ctx + p] = next_slot % pool_size;
|
||||||
|
next_slot++;
|
||||||
|
}
|
||||||
|
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
// req_pool_indices: pick B random request rows
|
||||||
|
int64_t* h_rpi = (int64_t*)malloc(sz_rpi);
|
||||||
|
for (int b = 0; b < B; b++) h_rpi[b] = b;
|
||||||
|
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
// kv_indptr: prefix sum of seq_lens
|
||||||
|
int* h_kvi = (int*)malloc(sz_kvi);
|
||||||
|
h_kvi[0] = 0;
|
||||||
|
for (int b = 0; b < B; b++) h_kvi[b + 1] = h_kvi[b] + seq_lens[b];
|
||||||
|
cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
// CPU reference
|
||||||
|
float* h_q_f = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float));
|
||||||
|
float* h_k_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
|
||||||
|
float* h_v_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
|
||||||
|
for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_q_f[i] = bf2f(h_q[i]);
|
||||||
|
for (int i = 0; i < pool_size * Hkv * HEAD_DIM; i++) {
|
||||||
|
h_k_f[i] = bf2f(h_k_pool[i]);
|
||||||
|
h_v_f[i] = bf2f(h_v_pool[i]);
|
||||||
|
}
|
||||||
|
float* h_o_ref = (float*)calloc(B * Hq * HEAD_DIM, sizeof(float));
|
||||||
|
cpu_paged_decode_ref(h_q_f, h_k_f, h_v_f, h_rtt, h_rpi, h_kvi,
|
||||||
|
nullptr, 0,
|
||||||
|
B, Hq, Hkv, HEAD_DIM, max_ctx, h_o_ref);
|
||||||
|
|
||||||
|
// Kernel launch
|
||||||
|
PagedAttentionParams<bf16> p;
|
||||||
|
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
|
||||||
|
p.head_dim = HEAD_DIM; p.total_q = B;
|
||||||
|
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
||||||
|
p.max_context_len = max_ctx; p.max_seq_len = max_sl;
|
||||||
|
p.causal_offset = causal ? 0 : -1; p.use_mask = 0;
|
||||||
|
p.mask = nullptr; p.mask_b_stride = 0;
|
||||||
|
p.mask_h_stride = 0; p.mask_q_stride = 0;
|
||||||
|
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
|
||||||
|
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
||||||
|
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
|
||||||
|
p.kv_indptr = d_kvi; p.qo_indptr = nullptr;
|
||||||
|
p.o = d_o; p.o_part = d_op; p.ml_part = d_ml;
|
||||||
|
|
||||||
|
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_decode<H>(p, 0); });
|
||||||
|
cudaDeviceSynchronize();
|
||||||
|
|
||||||
|
bf16* h_o_bf = (bf16*)malloc(sz_q);
|
||||||
|
cudaMemcpy(h_o_bf, d_o, sz_q, cudaMemcpyDeviceToHost);
|
||||||
|
float* h_o_got = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float));
|
||||||
|
for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]);
|
||||||
|
|
||||||
|
const float atol = 0.02f, rtol = 0.02f;
|
||||||
|
bool pass = true;
|
||||||
|
float max_err = 0.0f;
|
||||||
|
for (int i = 0; i < B * Hq * HEAD_DIM; i++) {
|
||||||
|
float e = fabsf(h_o_got[i] - h_o_ref[i]);
|
||||||
|
if (e > max_err) max_err = e;
|
||||||
|
if (e > atol + rtol * fabsf(h_o_ref[i])) { pass = false; break; }
|
||||||
|
}
|
||||||
|
|
||||||
|
print_paged_row(cfg, max_err, pass);
|
||||||
|
|
||||||
|
free(h_q); free(h_k_pool); free(h_v_pool); free(h_rtt); free(h_rpi);
|
||||||
|
free(h_kvi); free(h_q_f); free(h_k_f); free(h_v_f);
|
||||||
|
free(h_o_ref); free(h_o_bf); free(h_o_got);
|
||||||
|
cudaFree(d_q); cudaFree(d_o); cudaFree(d_k_pool); cudaFree(d_v_pool);
|
||||||
|
cudaFree(d_rtt); cudaFree(d_rpi); cudaFree(d_kvi); cudaFree(d_op); cudaFree(d_ml);
|
||||||
|
return pass ? 0 : 1;
|
||||||
|
}
|
||||||
|
|
||||||
|
// ======================================================================
|
||||||
|
// DECODE WITH MASK TEST (regression: 2D mask on mixed seq_lens)
|
||||||
|
// ======================================================================
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq,
|
||||||
|
int seed) {
|
||||||
|
srand(seed);
|
||||||
|
std::vector<int> seq_lens(B);
|
||||||
|
for (int b = 0; b < B; b++)
|
||||||
|
seq_lens[b] = 8 + rand() % (max_seq - 8);
|
||||||
|
int max_sl = *std::max_element(seq_lens.begin(), seq_lens.end());
|
||||||
|
int max_ctx = max_sl + 16;
|
||||||
|
int pool_size = B * max_ctx;
|
||||||
|
int num_reqs = B + 4;
|
||||||
|
|
||||||
|
char cfg[80];
|
||||||
|
snprintf(cfg, sizeof(cfg), "DECODE-MASK B=%d Hq=%d Hkv=%d D=%d max_sl=%d",
|
||||||
|
B, Hq, Hkv, HEAD_DIM, max_sl);
|
||||||
|
|
||||||
|
size_t sz_q = (size_t)B * Hq * HEAD_DIM * sizeof(bf16);
|
||||||
|
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
|
||||||
|
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t);
|
||||||
|
size_t sz_rpi = (size_t)B * sizeof(int64_t);
|
||||||
|
size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
|
||||||
|
size_t sz_mask = (size_t)B * max_sl * sizeof(bool);
|
||||||
|
size_t sz_op = (size_t)B * Hq * MAX_SPLITS * HEAD_DIM * sizeof(float);
|
||||||
|
size_t sz_ml = (size_t)B * Hq * MAX_SPLITS * 2 * sizeof(float);
|
||||||
|
|
||||||
|
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
|
||||||
|
int64_t *d_rtt, *d_rpi;
|
||||||
|
int *d_kvi;
|
||||||
|
bool *d_mask;
|
||||||
|
float *d_op, *d_ml;
|
||||||
|
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
|
||||||
|
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
|
||||||
|
cudaMalloc(&d_rtt, sz_rtt); cudaMalloc(&d_rpi, sz_rpi);
|
||||||
|
cudaMalloc(&d_kvi, sz_kvi);
|
||||||
|
cudaMalloc(&d_mask, sz_mask);
|
||||||
|
cudaMalloc(&d_op, sz_op); cudaMalloc(&d_ml, sz_ml);
|
||||||
|
|
||||||
|
auto rnd = [&]() { return (rand() / (float)RAND_MAX) * 2.0f - 1.0f; };
|
||||||
|
|
||||||
|
bf16* h_q = (bf16*)malloc(sz_q);
|
||||||
|
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) h_q[i] = f2bf(rnd());
|
||||||
|
cudaMemcpy(d_q, h_q, sz_q, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
bf16* h_k_pool = (bf16*)malloc(sz_kv);
|
||||||
|
bf16* h_v_pool = (bf16*)malloc(sz_kv);
|
||||||
|
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) {
|
||||||
|
h_k_pool[i] = f2bf(rnd());
|
||||||
|
h_v_pool[i] = f2bf(rnd());
|
||||||
|
}
|
||||||
|
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
|
||||||
|
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
int64_t* h_rtt = (int64_t*)malloc(sz_rtt);
|
||||||
|
int next_slot = 0;
|
||||||
|
for (int r = 0; r < num_reqs; r++)
|
||||||
|
for (int p = 0; p < max_ctx; p++) {
|
||||||
|
h_rtt[r * max_ctx + p] = next_slot % pool_size;
|
||||||
|
next_slot++;
|
||||||
|
}
|
||||||
|
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
int64_t* h_rpi = (int64_t*)malloc(sz_rpi);
|
||||||
|
for (int b = 0; b < B; b++) h_rpi[b] = b;
|
||||||
|
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
int* h_kvi = (int*)malloc(sz_kvi);
|
||||||
|
h_kvi[0] = 0;
|
||||||
|
for (int b = 0; b < B; b++) h_kvi[b + 1] = h_kvi[b] + seq_lens[b];
|
||||||
|
cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
// Mask: keep first half of each request's kv range, drop the rest —
|
||||||
|
// exercises the HasMask path with per-request seq_len.
|
||||||
|
bool* h_mask = (bool*)malloc(sz_mask);
|
||||||
|
for (int b = 0; b < B; b++)
|
||||||
|
for (int k = 0; k < max_sl; k++)
|
||||||
|
h_mask[b * max_sl + k] = (k < seq_lens[b]) && (k % 2 == 0);
|
||||||
|
cudaMemcpy(d_mask, h_mask, sz_mask, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
float* h_q_f = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float));
|
||||||
|
float* h_k_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
|
||||||
|
float* h_v_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
|
||||||
|
for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_q_f[i] = bf2f(h_q[i]);
|
||||||
|
for (int i = 0; i < pool_size * Hkv * HEAD_DIM; i++) {
|
||||||
|
h_k_f[i] = bf2f(h_k_pool[i]);
|
||||||
|
h_v_f[i] = bf2f(h_v_pool[i]);
|
||||||
|
}
|
||||||
|
float* h_o_ref = (float*)calloc(B * Hq * HEAD_DIM, sizeof(float));
|
||||||
|
cpu_paged_decode_ref(h_q_f, h_k_f, h_v_f, h_rtt, h_rpi, h_kvi,
|
||||||
|
h_mask, max_sl,
|
||||||
|
B, Hq, Hkv, HEAD_DIM, max_ctx, h_o_ref);
|
||||||
|
|
||||||
|
PagedAttentionParams<bf16> p;
|
||||||
|
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
|
||||||
|
p.head_dim = HEAD_DIM; p.total_q = B;
|
||||||
|
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
||||||
|
p.max_context_len = max_ctx; p.max_seq_len = max_sl;
|
||||||
|
p.causal_offset = -1; p.use_mask = 1;
|
||||||
|
p.mask = d_mask; p.mask_b_stride = max_sl;
|
||||||
|
p.mask_h_stride = 0; p.mask_q_stride = 0;
|
||||||
|
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
|
||||||
|
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
||||||
|
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
|
||||||
|
p.kv_indptr = d_kvi; p.qo_indptr = nullptr;
|
||||||
|
p.o = d_o; p.o_part = d_op; p.ml_part = d_ml;
|
||||||
|
|
||||||
|
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_decode<H>(p, 0); });
|
||||||
|
cudaDeviceSynchronize();
|
||||||
|
|
||||||
|
bf16* h_o_bf = (bf16*)malloc(sz_q);
|
||||||
|
cudaMemcpy(h_o_bf, d_o, sz_q, cudaMemcpyDeviceToHost);
|
||||||
|
float* h_o_got = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float));
|
||||||
|
for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]);
|
||||||
|
|
||||||
|
const float atol = 0.02f, rtol = 0.02f;
|
||||||
|
bool pass = true;
|
||||||
|
float max_err = 0.0f;
|
||||||
|
for (int i = 0; i < B * Hq * HEAD_DIM; i++) {
|
||||||
|
float e = fabsf(h_o_got[i] - h_o_ref[i]);
|
||||||
|
if (e > max_err) max_err = e;
|
||||||
|
if (e > atol + rtol * fabsf(h_o_ref[i])) { pass = false; break; }
|
||||||
|
}
|
||||||
|
|
||||||
|
print_paged_row(cfg, max_err, pass);
|
||||||
|
|
||||||
|
free(h_q); free(h_k_pool); free(h_v_pool); free(h_rtt); free(h_rpi);
|
||||||
|
free(h_kvi); free(h_mask); free(h_q_f); free(h_k_f); free(h_v_f);
|
||||||
|
free(h_o_ref); free(h_o_bf); free(h_o_got);
|
||||||
|
cudaFree(d_q); cudaFree(d_o); cudaFree(d_k_pool); cudaFree(d_v_pool);
|
||||||
|
cudaFree(d_rtt); cudaFree(d_rpi); cudaFree(d_kvi); cudaFree(d_mask);
|
||||||
|
cudaFree(d_op); cudaFree(d_ml);
|
||||||
|
return pass ? 0 : 1;
|
||||||
|
}
|
||||||
|
|
||||||
|
// ======================================================================
|
||||||
|
// PREFILL TEST
|
||||||
|
// ======================================================================
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
static int run_prefill_test(int B, int Hq, int Hkv,
|
||||||
|
std::vector<int>& q_lens,
|
||||||
|
std::vector<int>& kv_lens,
|
||||||
|
int causal, int seed) {
|
||||||
|
int total_q = 0;
|
||||||
|
int max_sl = 0;
|
||||||
|
for (int b = 0; b < B; b++) {
|
||||||
|
total_q += q_lens[b];
|
||||||
|
max_sl = max(max_sl, kv_lens[b]);
|
||||||
|
}
|
||||||
|
int max_ctx = max_sl + 16;
|
||||||
|
int pool_size = B * max_ctx;
|
||||||
|
int num_reqs = B + 4;
|
||||||
|
|
||||||
|
char cfg[80];
|
||||||
|
snprintf(cfg, sizeof(cfg), "PREFILL B=%d Hq=%d Hkv=%d D=%d max_sl=%d causal=%d",
|
||||||
|
B, Hq, Hkv, HEAD_DIM, max_sl, causal);
|
||||||
|
|
||||||
|
size_t sz_q = (size_t)total_q * Hq * HEAD_DIM * sizeof(bf16);
|
||||||
|
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
|
||||||
|
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t);
|
||||||
|
size_t sz_rpi = (size_t)B * sizeof(int64_t);
|
||||||
|
size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
|
||||||
|
size_t sz_qoi = (size_t)(B + 1) * sizeof(int);
|
||||||
|
|
||||||
|
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
|
||||||
|
int64_t *d_rtt, *d_rpi;
|
||||||
|
int *d_kvi, *d_qoi;
|
||||||
|
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
|
||||||
|
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
|
||||||
|
cudaMalloc(&d_rtt, sz_rtt); cudaMalloc(&d_rpi, sz_rpi);
|
||||||
|
cudaMalloc(&d_kvi, sz_kvi); cudaMalloc(&d_qoi, sz_qoi);
|
||||||
|
|
||||||
|
srand(seed);
|
||||||
|
auto rnd = [&]() { return (rand() / (float)RAND_MAX) * 2.0f - 1.0f; };
|
||||||
|
|
||||||
|
bf16* h_q = (bf16*)malloc(sz_q);
|
||||||
|
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) h_q[i] = f2bf(rnd());
|
||||||
|
cudaMemcpy(d_q, h_q, sz_q, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
bf16* h_k_pool = (bf16*)malloc(sz_kv);
|
||||||
|
bf16* h_v_pool = (bf16*)malloc(sz_kv);
|
||||||
|
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) {
|
||||||
|
h_k_pool[i] = f2bf(rnd());
|
||||||
|
h_v_pool[i] = f2bf(rnd());
|
||||||
|
}
|
||||||
|
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
|
||||||
|
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
int64_t* h_rtt = (int64_t*)malloc(sz_rtt);
|
||||||
|
int next_slot = 0;
|
||||||
|
for (int r = 0; r < num_reqs; r++)
|
||||||
|
for (int p = 0; p < max_ctx; p++) {
|
||||||
|
h_rtt[r * max_ctx + p] = next_slot % pool_size;
|
||||||
|
next_slot++;
|
||||||
|
}
|
||||||
|
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
int64_t* h_rpi = (int64_t*)malloc(sz_rpi);
|
||||||
|
for (int b = 0; b < B; b++) h_rpi[b] = b;
|
||||||
|
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
int* h_kvi = (int*)malloc(sz_kvi);
|
||||||
|
h_kvi[0] = 0;
|
||||||
|
for (int b = 0; b < B; b++) h_kvi[b + 1] = h_kvi[b] + kv_lens[b];
|
||||||
|
cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
int* h_qoi = (int*)malloc(sz_qoi);
|
||||||
|
h_qoi[0] = 0;
|
||||||
|
for (int b = 0; b < B; b++) h_qoi[b + 1] = h_qoi[b] + q_lens[b];
|
||||||
|
cudaMemcpy(d_qoi, h_qoi, sz_qoi, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
// CPU reference
|
||||||
|
float* h_q_f = (float*)malloc(total_q * Hq * HEAD_DIM * sizeof(float));
|
||||||
|
float* h_k_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
|
||||||
|
float* h_v_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
|
||||||
|
for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) h_q_f[i] = bf2f(h_q[i]);
|
||||||
|
for (int i = 0; i < pool_size * Hkv * HEAD_DIM; i++) {
|
||||||
|
h_k_f[i] = bf2f(h_k_pool[i]);
|
||||||
|
h_v_f[i] = bf2f(h_v_pool[i]);
|
||||||
|
}
|
||||||
|
float* h_o_ref = (float*)calloc(total_q * Hq * HEAD_DIM, sizeof(float));
|
||||||
|
cpu_paged_prefill_ref(h_q_f, h_k_f, h_v_f, h_rtt, h_rpi, h_kvi, h_qoi,
|
||||||
|
nullptr, 0, 0,
|
||||||
|
B, Hq, Hkv, HEAD_DIM, max_ctx, causal, h_o_ref);
|
||||||
|
|
||||||
|
// Kernel launch
|
||||||
|
PagedAttentionParams<bf16> p;
|
||||||
|
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
|
||||||
|
p.head_dim = HEAD_DIM; p.total_q = total_q;
|
||||||
|
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
||||||
|
p.max_context_len = max_ctx; p.max_seq_len = max_sl;
|
||||||
|
int max_ql = 0;
|
||||||
|
for (int b = 0; b < B; b++) max_ql = max(max_ql, q_lens[b]);
|
||||||
|
p.max_q_len = max_ql;
|
||||||
|
p.causal_offset = causal ? 0 : -1; p.use_mask = 0;
|
||||||
|
p.mask = nullptr; p.mask_b_stride = 0;
|
||||||
|
p.mask_h_stride = 0; p.mask_q_stride = 0;
|
||||||
|
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
|
||||||
|
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
||||||
|
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
|
||||||
|
p.kv_indptr = d_kvi; p.qo_indptr = d_qoi;
|
||||||
|
p.o = d_o; p.o_part = nullptr; p.ml_part = nullptr;
|
||||||
|
|
||||||
|
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_prefill<H>(p, 0); });
|
||||||
|
cudaDeviceSynchronize();
|
||||||
|
|
||||||
|
bf16* h_o_bf = (bf16*)malloc(sz_q);
|
||||||
|
cudaMemcpy(h_o_bf, d_o, sz_q, cudaMemcpyDeviceToHost);
|
||||||
|
float* h_o_got = (float*)malloc(total_q * Hq * HEAD_DIM * sizeof(float));
|
||||||
|
for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]);
|
||||||
|
|
||||||
|
const float atol = 0.02f, rtol = 0.02f;
|
||||||
|
bool pass = true;
|
||||||
|
float max_err = 0.0f;
|
||||||
|
for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) {
|
||||||
|
float e = fabsf(h_o_got[i] - h_o_ref[i]);
|
||||||
|
if (e > max_err) max_err = e;
|
||||||
|
if (e > atol + rtol * fabsf(h_o_ref[i])) { pass = false; break; }
|
||||||
|
}
|
||||||
|
|
||||||
|
print_paged_row(cfg, max_err, pass);
|
||||||
|
|
||||||
|
free(h_q); free(h_k_pool); free(h_v_pool); free(h_rtt); free(h_rpi);
|
||||||
|
free(h_kvi); free(h_qoi); free(h_q_f); free(h_k_f); free(h_v_f);
|
||||||
|
free(h_o_ref); free(h_o_bf); free(h_o_got);
|
||||||
|
cudaFree(d_q); cudaFree(d_o); cudaFree(d_k_pool); cudaFree(d_v_pool);
|
||||||
|
cudaFree(d_rtt); cudaFree(d_rpi); cudaFree(d_kvi); cudaFree(d_qoi);
|
||||||
|
return pass ? 0 : 1;
|
||||||
|
}
|
||||||
|
|
||||||
|
// ======================================================================
|
||||||
|
// PREFILL WITH MASK TEST (regression: 4D causal mask on single request)
|
||||||
|
// ======================================================================
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) {
|
||||||
|
srand(seed);
|
||||||
|
int B = 1;
|
||||||
|
int total_q = q_len;
|
||||||
|
int seq_len = q_len; // pure prefill: kv_len == q_len
|
||||||
|
int max_ctx = seq_len + 16;
|
||||||
|
int pool_size = B * max_ctx;
|
||||||
|
int num_reqs = B + 4;
|
||||||
|
|
||||||
|
char cfg[80];
|
||||||
|
snprintf(cfg, sizeof(cfg), "PREFILL-MASK Hq=%d Hkv=%d D=%d q_len=%d",
|
||||||
|
Hq, Hkv, HEAD_DIM, q_len);
|
||||||
|
fflush(stdout);
|
||||||
|
|
||||||
|
size_t sz_q = (size_t)total_q * Hq * HEAD_DIM * sizeof(bf16);
|
||||||
|
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
|
||||||
|
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t);
|
||||||
|
size_t sz_rpi = (size_t)B * sizeof(int64_t);
|
||||||
|
size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
|
||||||
|
size_t sz_qoi = (size_t)(B + 1) * sizeof(int);
|
||||||
|
size_t sz_mask = (size_t)B * q_len * q_len * sizeof(bool);
|
||||||
|
|
||||||
|
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
|
||||||
|
int64_t *d_rtt, *d_rpi;
|
||||||
|
int *d_kvi, *d_qoi;
|
||||||
|
bool *d_mask;
|
||||||
|
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
|
||||||
|
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
|
||||||
|
cudaMalloc(&d_rtt, sz_rtt); cudaMalloc(&d_rpi, sz_rpi);
|
||||||
|
cudaMalloc(&d_kvi, sz_kvi); cudaMalloc(&d_qoi, sz_qoi);
|
||||||
|
cudaMalloc(&d_mask, sz_mask);
|
||||||
|
|
||||||
|
auto rnd = [&]() { return (rand() / (float)RAND_MAX) * 2.0f - 1.0f; };
|
||||||
|
|
||||||
|
bf16* h_q = (bf16*)malloc(sz_q);
|
||||||
|
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) h_q[i] = f2bf(rnd());
|
||||||
|
cudaMemcpy(d_q, h_q, sz_q, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
bf16* h_k_pool = (bf16*)malloc(sz_kv);
|
||||||
|
bf16* h_v_pool = (bf16*)malloc(sz_kv);
|
||||||
|
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) {
|
||||||
|
h_k_pool[i] = f2bf(rnd());
|
||||||
|
h_v_pool[i] = f2bf(rnd());
|
||||||
|
}
|
||||||
|
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
|
||||||
|
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
int64_t* h_rtt = (int64_t*)malloc(sz_rtt);
|
||||||
|
int next_slot = 0;
|
||||||
|
for (int r = 0; r < num_reqs; r++)
|
||||||
|
for (int p = 0; p < max_ctx; p++) {
|
||||||
|
h_rtt[r * max_ctx + p] = next_slot % pool_size;
|
||||||
|
next_slot++;
|
||||||
|
}
|
||||||
|
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
int64_t* h_rpi = (int64_t*)malloc(sz_rpi);
|
||||||
|
h_rpi[0] = 0;
|
||||||
|
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
int* h_kvi = (int*)malloc(sz_kvi);
|
||||||
|
h_kvi[0] = 0; h_kvi[1] = seq_len;
|
||||||
|
cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
int* h_qoi = (int*)malloc(sz_qoi);
|
||||||
|
h_qoi[0] = 0; h_qoi[1] = q_len;
|
||||||
|
cudaMemcpy(d_qoi, h_qoi, sz_qoi, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
// 4D causal mask [B, 1, q_len, q_len], True=keep.
|
||||||
|
bool* h_mask = (bool*)malloc(sz_mask);
|
||||||
|
for (int qi = 0; qi < q_len; qi++)
|
||||||
|
for (int kj = 0; kj < q_len; kj++)
|
||||||
|
h_mask[qi * q_len + kj] = (kj <= qi);
|
||||||
|
cudaMemcpy(d_mask, h_mask, sz_mask, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
float* h_q_f = (float*)malloc(total_q * Hq * HEAD_DIM * sizeof(float));
|
||||||
|
float* h_k_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
|
||||||
|
float* h_v_f = (float*)malloc(pool_size * Hkv * HEAD_DIM * sizeof(float));
|
||||||
|
for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) h_q_f[i] = bf2f(h_q[i]);
|
||||||
|
for (int i = 0; i < pool_size * Hkv * HEAD_DIM; i++) {
|
||||||
|
h_k_f[i] = bf2f(h_k_pool[i]);
|
||||||
|
h_v_f[i] = bf2f(h_v_pool[i]);
|
||||||
|
}
|
||||||
|
float* h_o_ref = (float*)calloc(total_q * Hq * HEAD_DIM, sizeof(float));
|
||||||
|
// CPU ref with causal=0 so it consults the mask (not the causal flag).
|
||||||
|
cpu_paged_prefill_ref(h_q_f, h_k_f, h_v_f, h_rtt, h_rpi, h_kvi, h_qoi,
|
||||||
|
h_mask, q_len, q_len,
|
||||||
|
B, Hq, Hkv, HEAD_DIM, max_ctx, 0, h_o_ref);
|
||||||
|
|
||||||
|
PagedAttentionParams<bf16> p;
|
||||||
|
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
|
||||||
|
p.head_dim = HEAD_DIM; p.total_q = total_q;
|
||||||
|
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
||||||
|
p.max_context_len = max_ctx; p.max_seq_len = q_len;
|
||||||
|
p.max_q_len = q_len;
|
||||||
|
p.causal_offset = -1; p.use_mask = 1;
|
||||||
|
p.mask = d_mask; p.mask_b_stride = q_len * q_len;
|
||||||
|
p.mask_h_stride = 0; p.mask_q_stride = q_len;
|
||||||
|
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
|
||||||
|
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
||||||
|
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
|
||||||
|
p.kv_indptr = d_kvi; p.qo_indptr = d_qoi;
|
||||||
|
p.o = d_o; p.o_part = nullptr; p.ml_part = nullptr;
|
||||||
|
|
||||||
|
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_prefill<H>(p, 0); });
|
||||||
|
cudaDeviceSynchronize();
|
||||||
|
|
||||||
|
bf16* h_o_bf = (bf16*)malloc(sz_q);
|
||||||
|
cudaMemcpy(h_o_bf, d_o, sz_q, cudaMemcpyDeviceToHost);
|
||||||
|
float* h_o_got = (float*)malloc(total_q * Hq * HEAD_DIM * sizeof(float));
|
||||||
|
for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) h_o_got[i] = bf2f(h_o_bf[i]);
|
||||||
|
|
||||||
|
const float atol = 0.02f, rtol = 0.02f;
|
||||||
|
bool pass = true;
|
||||||
|
float max_err = 0.0f;
|
||||||
|
for (int i = 0; i < total_q * Hq * HEAD_DIM; i++) {
|
||||||
|
float e = fabsf(h_o_got[i] - h_o_ref[i]);
|
||||||
|
if (e > max_err) max_err = e;
|
||||||
|
if (e > atol + rtol * fabsf(h_o_ref[i])) { pass = false; break; }
|
||||||
|
}
|
||||||
|
|
||||||
|
print_paged_row(cfg, max_err, pass);
|
||||||
|
|
||||||
|
free(h_q); free(h_k_pool); free(h_v_pool); free(h_rtt); free(h_rpi);
|
||||||
|
free(h_kvi); free(h_qoi); free(h_mask); free(h_q_f); free(h_k_f); free(h_v_f);
|
||||||
|
free(h_o_ref); free(h_o_bf); free(h_o_got);
|
||||||
|
cudaFree(d_q); cudaFree(d_o); cudaFree(d_k_pool); cudaFree(d_v_pool);
|
||||||
|
cudaFree(d_rtt); cudaFree(d_rpi); cudaFree(d_kvi); cudaFree(d_qoi);
|
||||||
|
cudaFree(d_mask);
|
||||||
|
return pass ? 0 : 1;
|
||||||
|
}
|
||||||
|
|
||||||
|
// ======================================================================
|
||||||
|
// BENCH
|
||||||
|
// ======================================================================
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
static void bench_decode(int B, int Hq, int Hkv, int seq_len) {
|
||||||
|
int max_ctx = seq_len + 16;
|
||||||
|
int pool_size = B * max_ctx;
|
||||||
|
int num_reqs = B;
|
||||||
|
|
||||||
|
size_t sz_q = (size_t)B * Hq * HEAD_DIM * sizeof(bf16);
|
||||||
|
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
|
||||||
|
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t);
|
||||||
|
size_t sz_rpi = (size_t)B * sizeof(int64_t);
|
||||||
|
size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
|
||||||
|
size_t sz_op = (size_t)B * Hq * MAX_SPLITS * HEAD_DIM * sizeof(float);
|
||||||
|
size_t sz_ml = (size_t)B * Hq * MAX_SPLITS * 2 * sizeof(float);
|
||||||
|
|
||||||
|
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
|
||||||
|
int64_t *d_rtt, *d_rpi;
|
||||||
|
int *d_kvi;
|
||||||
|
float *d_op, *d_ml;
|
||||||
|
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
|
||||||
|
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
|
||||||
|
cudaMalloc(&d_rtt, sz_rtt); cudaMalloc(&d_rpi, sz_rpi);
|
||||||
|
cudaMalloc(&d_kvi, sz_kvi);
|
||||||
|
cudaMalloc(&d_op, sz_op); cudaMalloc(&d_ml, sz_ml);
|
||||||
|
|
||||||
|
bf16* tmp = (bf16*)malloc(sz_kv > sz_q ? sz_kv : sz_q);
|
||||||
|
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) tmp[i] = f2bf(randf());
|
||||||
|
cudaMemcpy(d_q, tmp, sz_q, cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) tmp[i] = f2bf(randf());
|
||||||
|
cudaMemcpy(d_k_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
|
||||||
|
cudaMemcpy(d_v_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
int64_t* h_rtt = (int64_t*)malloc(sz_rtt);
|
||||||
|
for (int r = 0; r < num_reqs; r++)
|
||||||
|
for (int p = 0; p < max_ctx; p++)
|
||||||
|
h_rtt[r * max_ctx + p] = (r * max_ctx + p) % pool_size;
|
||||||
|
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
|
||||||
|
int64_t* h_rpi = (int64_t*)malloc(sz_rpi);
|
||||||
|
for (int b = 0; b < B; b++) h_rpi[b] = b;
|
||||||
|
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
|
||||||
|
int* h_kvi = (int*)malloc(sz_kvi);
|
||||||
|
h_kvi[0] = 0;
|
||||||
|
for (int b = 0; b < B; b++) h_kvi[b + 1] = h_kvi[b] + seq_len;
|
||||||
|
cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
PagedAttentionParams<bf16> p;
|
||||||
|
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
|
||||||
|
p.head_dim = HEAD_DIM; p.total_q = B;
|
||||||
|
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
||||||
|
p.max_context_len = max_ctx; p.max_seq_len = seq_len;
|
||||||
|
p.causal_offset = 0; p.use_mask = 0;
|
||||||
|
p.mask = nullptr; p.mask_b_stride = 0;
|
||||||
|
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
|
||||||
|
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
||||||
|
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
|
||||||
|
p.kv_indptr = d_kvi; p.qo_indptr = nullptr;
|
||||||
|
p.o = d_o; p.o_part = d_op; p.ml_part = d_ml;
|
||||||
|
|
||||||
|
auto launch = [&]() {
|
||||||
|
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_decode<H>(p, 0); });
|
||||||
|
};
|
||||||
|
// Decode: q_len=1, query is the last token → attends to all [0, seq_len).
|
||||||
|
// FLOPs = 2 * (QK^T + PV) = 4 * B * Hq * seq_len * D.
|
||||||
|
double flops = 4.0 * B * Hq * (double)seq_len * HEAD_DIM;
|
||||||
|
BenchResult r = bench_kernel(launch, 3, 10, flops);
|
||||||
|
|
||||||
|
char cfg[64];
|
||||||
|
snprintf(cfg, sizeof(cfg), "DEC B=%2d Hq=%2d Hk=%d kv=%4d D=%3d",
|
||||||
|
B, Hq, Hkv, seq_len, HEAD_DIM);
|
||||||
|
print_bench_row(cfg, r);
|
||||||
|
|
||||||
|
free(tmp); free(h_rtt); free(h_rpi); free(h_kvi);
|
||||||
|
cudaFree(d_q); cudaFree(d_o); cudaFree(d_k_pool); cudaFree(d_v_pool);
|
||||||
|
cudaFree(d_rtt); cudaFree(d_rpi); cudaFree(d_kvi); cudaFree(d_op); cudaFree(d_ml);
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
static void bench_prefill(int B, int Hq, int Hkv, int q_len, int kv_len, int causal) {
|
||||||
|
int total_q = B * q_len;
|
||||||
|
int max_ctx = kv_len + 16;
|
||||||
|
int pool_size = B * max_ctx;
|
||||||
|
int num_reqs = B;
|
||||||
|
|
||||||
|
size_t sz_q = (size_t)total_q * Hq * HEAD_DIM * sizeof(bf16);
|
||||||
|
size_t sz_kv = (size_t)pool_size * Hkv * HEAD_DIM * sizeof(bf16);
|
||||||
|
size_t sz_rtt = (size_t)num_reqs * max_ctx * sizeof(int64_t);
|
||||||
|
size_t sz_rpi = (size_t)B * sizeof(int64_t);
|
||||||
|
size_t sz_kvi = (size_t)(B + 1) * sizeof(int);
|
||||||
|
size_t sz_qoi = (size_t)(B + 1) * sizeof(int);
|
||||||
|
|
||||||
|
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
|
||||||
|
int64_t *d_rtt, *d_rpi;
|
||||||
|
int *d_kvi, *d_qoi;
|
||||||
|
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
|
||||||
|
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
|
||||||
|
cudaMalloc(&d_rtt, sz_rtt); cudaMalloc(&d_rpi, sz_rpi);
|
||||||
|
cudaMalloc(&d_kvi, sz_kvi); cudaMalloc(&d_qoi, sz_qoi);
|
||||||
|
|
||||||
|
bf16* tmp = (bf16*)malloc(sz_kv > sz_q ? sz_kv : sz_q);
|
||||||
|
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) tmp[i] = f2bf(randf());
|
||||||
|
cudaMemcpy(d_q, tmp, sz_q, cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) tmp[i] = f2bf(randf());
|
||||||
|
cudaMemcpy(d_k_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
|
||||||
|
cudaMemcpy(d_v_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
int64_t* h_rtt = (int64_t*)malloc(sz_rtt);
|
||||||
|
for (int r = 0; r < num_reqs; r++)
|
||||||
|
for (int p = 0; p < max_ctx; p++)
|
||||||
|
h_rtt[r * max_ctx + p] = (r * max_ctx + p) % pool_size;
|
||||||
|
cudaMemcpy(d_rtt, h_rtt, sz_rtt, cudaMemcpyHostToDevice);
|
||||||
|
int64_t* h_rpi = (int64_t*)malloc(sz_rpi);
|
||||||
|
for (int b = 0; b < B; b++) h_rpi[b] = b;
|
||||||
|
cudaMemcpy(d_rpi, h_rpi, sz_rpi, cudaMemcpyHostToDevice);
|
||||||
|
int* h_kvi = (int*)malloc(sz_kvi);
|
||||||
|
h_kvi[0] = 0;
|
||||||
|
for (int b = 0; b < B; b++) h_kvi[b + 1] = h_kvi[b] + kv_len;
|
||||||
|
cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice);
|
||||||
|
int* h_qoi = (int*)malloc(sz_qoi);
|
||||||
|
h_qoi[0] = 0;
|
||||||
|
for (int b = 0; b < B; b++) h_qoi[b + 1] = h_qoi[b] + q_len;
|
||||||
|
cudaMemcpy(d_qoi, h_qoi, sz_qoi, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
PagedAttentionParams<bf16> p;
|
||||||
|
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
|
||||||
|
p.head_dim = HEAD_DIM; p.total_q = total_q;
|
||||||
|
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
||||||
|
p.max_context_len = max_ctx; p.max_seq_len = kv_len;
|
||||||
|
p.total_q = total_q; p.max_q_len = q_len;
|
||||||
|
p.causal_offset = causal ? 0 : -1; p.use_mask = 0;
|
||||||
|
p.mask = nullptr; p.mask_b_stride = 0;
|
||||||
|
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
|
||||||
|
p.q = d_q; p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
||||||
|
p.req_to_token = d_rtt; p.req_pool_indices = d_rpi;
|
||||||
|
p.kv_indptr = d_kvi; p.qo_indptr = d_qoi;
|
||||||
|
p.o = d_o; p.o_part = nullptr; p.ml_part = nullptr;
|
||||||
|
|
||||||
|
auto launch = [&]() {
|
||||||
|
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_prefill<H>(p, 0); });
|
||||||
|
};
|
||||||
|
// FLOPs = 2 * (QK^T + PV) = 4 * effective_qk_pairs * Hq * D.
|
||||||
|
// Non-causal: effective = q_len * kv_len.
|
||||||
|
// Causal: Q row qi attends to [0, causal_off + qi + 1) where
|
||||||
|
// causal_off = kv_len - q_len. Total KV accesses per request:
|
||||||
|
// sum_{qi=0}^{q_len-1} (kv_len - q_len + qi + 1)
|
||||||
|
// = q_len * (kv_len - q_len) + q_len * (q_len + 1) / 2.
|
||||||
|
double eff_kv;
|
||||||
|
if (causal) {
|
||||||
|
eff_kv = (double)q_len * (kv_len - q_len)
|
||||||
|
+ (double)q_len * (q_len + 1) / 2.0;
|
||||||
|
} else {
|
||||||
|
eff_kv = (double)q_len * kv_len;
|
||||||
|
}
|
||||||
|
double flops = 4.0 * B * Hq * eff_kv * HEAD_DIM;
|
||||||
|
BenchResult r = bench_kernel(launch, 3, 10, flops);
|
||||||
|
|
||||||
|
char cfg[80];
|
||||||
|
snprintf(cfg, sizeof(cfg), "PRE B=%d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d c=%d",
|
||||||
|
B, Hq, Hkv, q_len, kv_len, HEAD_DIM, causal);
|
||||||
|
print_bench_row(cfg, r);
|
||||||
|
|
||||||
|
free(tmp); free(h_rtt); free(h_rpi); free(h_kvi); free(h_qoi);
|
||||||
|
cudaFree(d_q); cudaFree(d_o); cudaFree(d_k_pool); cudaFree(d_v_pool);
|
||||||
|
cudaFree(d_rtt); cudaFree(d_rpi); cudaFree(d_kvi); cudaFree(d_qoi);
|
||||||
|
}
|
||||||
|
|
||||||
|
int main() {
|
||||||
|
int fail = 0;
|
||||||
|
|
||||||
|
// ===== DECODE TESTS =====
|
||||||
|
printf("=== Paged Decode Tests ===\n");
|
||||||
|
print_paged_header();
|
||||||
|
fail += run_decode_test<128>(1, 32, 4, 512, 0, 1);
|
||||||
|
fail += run_decode_test<128>(1, 32, 4, 1024, 0, 2);
|
||||||
|
fail += run_decode_test<128>(4, 32, 4, 512, 0, 3);
|
||||||
|
fail += run_decode_test<128>(8, 32, 4, 1024, 0, 4);
|
||||||
|
fail += run_decode_test<128>(4, 32, 8, 2048, 0, 5);
|
||||||
|
fail += run_decode_test<128>(1, 16, 1, 256, 0, 6);
|
||||||
|
fail += run_decode_test<128>(2, 8, 2, 512, 1, 7);
|
||||||
|
fail += run_decode_test<64>(1, 4, 2, 256, 0, 8);
|
||||||
|
fail += run_decode_test<256>(1, 2, 1, 256, 0, 9);
|
||||||
|
fail += run_decode_test<128>(16, 32, 4, 2048, 0, 10);
|
||||||
|
fail += run_decode_test<128>(32, 32, 4, 1024, 0, 11);
|
||||||
|
|
||||||
|
// Decode with 2D mask (regression: mixed seq_lens + HasMask)
|
||||||
|
fail += run_decode_mask_test<128>(2, 8, 2, 256, 30);
|
||||||
|
fail += run_decode_mask_test<128>(4, 32, 4, 512, 31);
|
||||||
|
fail += run_decode_mask_test<64>(2, 4, 2, 128, 32);
|
||||||
|
|
||||||
|
if (fail) { printf("\nFAILED decode tests\n"); return fail; }
|
||||||
|
|
||||||
|
// ===== PREFILL TESTS =====
|
||||||
|
printf("\n=== Paged Prefill Tests ===\n");
|
||||||
|
print_paged_header();
|
||||||
|
// Single request, pure prefill (q_len == kv_len)
|
||||||
|
{
|
||||||
|
std::vector<int> ql = {512};
|
||||||
|
std::vector<int> kl = {512};
|
||||||
|
fail += run_prefill_test<128>(1, 32, 4, ql, kl, 1, 20);
|
||||||
|
}
|
||||||
|
{
|
||||||
|
std::vector<int> ql = {1024};
|
||||||
|
std::vector<int> kl = {1024};
|
||||||
|
fail += run_prefill_test<128>(1, 32, 4, ql, kl, 1, 21);
|
||||||
|
}
|
||||||
|
{
|
||||||
|
std::vector<int> ql = {2048};
|
||||||
|
std::vector<int> kl = {2048};
|
||||||
|
fail += run_prefill_test<128>(1, 32, 4, ql, kl, 1, 22);
|
||||||
|
}
|
||||||
|
// Ragged batch: different q_lens and kv_lens
|
||||||
|
{
|
||||||
|
std::vector<int> ql = {128, 256, 64};
|
||||||
|
std::vector<int> kl = {128, 256, 64};
|
||||||
|
fail += run_prefill_test<128>(3, 32, 4, ql, kl, 1, 23);
|
||||||
|
}
|
||||||
|
{
|
||||||
|
std::vector<int> ql = {64, 128, 256, 32};
|
||||||
|
std::vector<int> kl = {64, 128, 256, 32};
|
||||||
|
fail += run_prefill_test<128>(4, 32, 4, ql, kl, 1, 24);
|
||||||
|
}
|
||||||
|
// Extend: kv_len > q_len (append to existing cache)
|
||||||
|
{
|
||||||
|
std::vector<int> ql = {64, 128};
|
||||||
|
std::vector<int> kl = {256, 512};
|
||||||
|
fail += run_prefill_test<128>(2, 32, 4, ql, kl, 1, 25);
|
||||||
|
}
|
||||||
|
// Non-causal
|
||||||
|
{
|
||||||
|
std::vector<int> ql = {256, 128};
|
||||||
|
std::vector<int> kl = {256, 128};
|
||||||
|
fail += run_prefill_test<128>(2, 32, 4, ql, kl, 0, 26);
|
||||||
|
}
|
||||||
|
// Single token (q_len=1 per request, like decode but via prefill path)
|
||||||
|
{
|
||||||
|
std::vector<int> ql = {1, 1, 1, 1};
|
||||||
|
std::vector<int> kl = {128, 256, 64, 512};
|
||||||
|
fail += run_prefill_test<128>(4, 32, 4, ql, kl, 1, 27);
|
||||||
|
}
|
||||||
|
// D=64
|
||||||
|
{
|
||||||
|
std::vector<int> ql = {128, 64};
|
||||||
|
std::vector<int> kl = {128, 64};
|
||||||
|
fail += run_prefill_test<64>(2, 4, 2, ql, kl, 1, 28);
|
||||||
|
}
|
||||||
|
// D=256
|
||||||
|
{
|
||||||
|
std::vector<int> ql = {128, 64};
|
||||||
|
std::vector<int> kl = {128, 64};
|
||||||
|
fail += run_prefill_test<256>(2, 2, 1, ql, kl, 1, 29);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Prefill with 4D causal mask (regression: single-request mask path)
|
||||||
|
fail += run_prefill_mask_test<128>(32, 4, 512, 40);
|
||||||
|
fail += run_prefill_mask_test<128>(32, 4, 1024, 41);
|
||||||
|
fail += run_prefill_mask_test<64>(4, 2, 256, 42);
|
||||||
|
|
||||||
|
if (fail) { printf("\nFAILED prefill tests\n"); return fail; }
|
||||||
|
printf("\nAll tests passed!\n");
|
||||||
|
|
||||||
|
// ===== BENCH =====
|
||||||
|
printf("\n===== PAGED DECODE BENCH =====\n");
|
||||||
|
print_bench_header();
|
||||||
|
bench_decode<128>(1, 32, 4, 512);
|
||||||
|
bench_decode<128>(1, 32, 4, 1024);
|
||||||
|
bench_decode<128>(1, 32, 4, 2048);
|
||||||
|
bench_decode<128>(1, 32, 4, 4096);
|
||||||
|
bench_decode<128>(4, 32, 4, 2048);
|
||||||
|
bench_decode<128>(16, 32, 4, 2048);
|
||||||
|
bench_decode<128>(32, 32, 4, 1024);
|
||||||
|
|
||||||
|
printf("\n===== PAGED PREFILL BENCH =====\n");
|
||||||
|
print_bench_header();
|
||||||
|
bench_prefill<128>(1, 32, 4, 512, 512, 0);
|
||||||
|
bench_prefill<128>(1, 32, 4, 1024, 1024, 0);
|
||||||
|
bench_prefill<128>(1, 32, 4, 2048, 2048, 0);
|
||||||
|
bench_prefill<128>(1, 32, 4, 2048, 2048, 1);
|
||||||
|
bench_prefill<128>(4, 32, 4, 2048, 2048, 1);
|
||||||
|
bench_prefill<128>(1, 32, 4, 4096, 4096, 1);
|
||||||
|
|
||||||
|
return 0;
|
||||||
|
}
|
||||||
@@ -1,169 +0,0 @@
|
|||||||
/*
|
|
||||||
Pure-C test — uses shared dispatcher.
|
|
||||||
nvcc -I csrc -arch=sm_89 -O3 \
|
|
||||||
--use_fast_math --ptxas-options=-O3 --extra-device-vectorization \
|
|
||||||
csrc/tests/attn_prefill_test.cu -o test && ./test
|
|
||||||
*/
|
|
||||||
|
|
||||||
#include "test_utils.cuh"
|
|
||||||
#include "../kernels/attn_dispatchers.cuh"
|
|
||||||
|
|
||||||
// Warmed-up, CUDA-event timed throughput sweep over the production MMA path.
|
|
||||||
static void bench() {
|
|
||||||
const int cfgs[][7] = {
|
|
||||||
{1,32,4,512,512,128,0},
|
|
||||||
{1,32,4,1024,1024,128,0},
|
|
||||||
{1,32,4,2048,2048,128,0},
|
|
||||||
{1,32,4,2048,2048,128,1},
|
|
||||||
{4,32,4,2048,2048,128,1},
|
|
||||||
{1,32,4,4096,4096,128,1},
|
|
||||||
};
|
|
||||||
int n = sizeof(cfgs)/sizeof(cfgs[0]);
|
|
||||||
const int WARMUP = 10, ITERS = 50;
|
|
||||||
printf("\n===== PREFILL BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS);
|
|
||||||
printf("%-46s | %10s | %10s | %10s\n",
|
|
||||||
"config", "latency", "bandwidth", "throughput");
|
|
||||||
printf("---------------------------------------------------------------"
|
|
||||||
"----------------------------\n");
|
|
||||||
|
|
||||||
for (int ci = 0; ci < n; ci++) {
|
|
||||||
int B=cfgs[ci][0], Hq=cfgs[ci][1], Hk=cfgs[ci][2];
|
|
||||||
int ql=cfgs[ci][3], kl=cfgs[ci][4], D=cfgs[ci][5], causal=cfgs[ci][6];
|
|
||||||
size_t nQ=(size_t)B*Hq*ql*D, nKV=(size_t)B*Hk*kl*D;
|
|
||||||
|
|
||||||
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
|
||||||
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
|
||||||
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
|
||||||
size_t big = nQ>nKV?nQ:nKV; tmp=new bf16[big];
|
|
||||||
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(randf());
|
|
||||||
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
|
||||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(randf());
|
|
||||||
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
|
||||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(randf());
|
|
||||||
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
|
||||||
|
|
||||||
AttentionParams<bf16> p;
|
|
||||||
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D;
|
|
||||||
p.use_mask=0; p.causal_offset=causal?0:-1;
|
|
||||||
set_default_strides(p);
|
|
||||||
p.scale=1.0f/sqrtf((float)D);
|
|
||||||
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
|
||||||
|
|
||||||
auto launch = [&]() { dispatch_by_head_dim(D, [&]<int H>() { dispatch_prefill<H>(p); }); };
|
|
||||||
for (int i=0;i<WARMUP;i++) launch();
|
|
||||||
cudaDeviceSynchronize();
|
|
||||||
cudaError_t err=cudaGetLastError();
|
|
||||||
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return;}
|
|
||||||
|
|
||||||
cudaEvent_t s,e; cudaEventCreate(&s); cudaEventCreate(&e);
|
|
||||||
cudaEventRecord(s);
|
|
||||||
for (int i=0;i<ITERS;i++) launch();
|
|
||||||
cudaEventRecord(e); cudaEventSynchronize(e);
|
|
||||||
float ms=0; cudaEventElapsedTime(&ms,s,e); ms/=ITERS;
|
|
||||||
|
|
||||||
double flops = 4.0*B*Hq*(double)ql*kl*D;
|
|
||||||
if (causal) flops *= 0.5;
|
|
||||||
double tflops = flops/(ms*1e-3)/1e12;
|
|
||||||
double bytes = 2.0 * (2.0*nQ + 2.0*nKV);
|
|
||||||
double gbps = bytes/(ms*1e-3)/1e9;
|
|
||||||
|
|
||||||
char cfg[64];
|
|
||||||
snprintf(cfg, sizeof(cfg),
|
|
||||||
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d causal=%d",
|
|
||||||
B,Hq,Hk,ql,kl,D,causal);
|
|
||||||
printf("%-46s | %7.4f ms | %7.1f GB/s | %6.2f TFLOP/s\n",
|
|
||||||
cfg, ms, gbps, tflops);
|
|
||||||
|
|
||||||
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
|
|
||||||
delete[]tmp; cudaEventDestroy(s); cudaEventDestroy(e);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
static int run_test(int B, int Hq, int Hk, int ql, int kl, int D, int causal) {
|
|
||||||
printf("=== B=%d Hq=%d Hk=%d q=%d kv=%d D=%d causal=%d ===\n",
|
|
||||||
B,Hq,Hk,ql,kl,D,causal);
|
|
||||||
|
|
||||||
size_t nQ = B*Hq*ql*D, nKV = B*Hk*kl*D;
|
|
||||||
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
|
|
||||||
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
|
|
||||||
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
|
|
||||||
|
|
||||||
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
|
||||||
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
|
||||||
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
|
||||||
tmp=new bf16[max(nQ,nKV)];
|
|
||||||
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
|
|
||||||
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
|
||||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
|
|
||||||
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
|
||||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
|
|
||||||
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
|
||||||
|
|
||||||
AttentionParams<bf16> p;
|
|
||||||
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D;
|
|
||||||
p.use_mask=0; p.causal_offset=causal?0:-1;
|
|
||||||
set_default_strides(p);
|
|
||||||
p.scale=1.0f/sqrtf((float)D);
|
|
||||||
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
|
||||||
|
|
||||||
double t0=now_ms();
|
|
||||||
dispatch_by_head_dim(D, [&]<int H>() { dispatch_prefill<H>(p); });
|
|
||||||
cudaDeviceSynchronize();
|
|
||||||
double kms=now_ms()-t0;
|
|
||||||
cudaError_t err=cudaGetLastError();
|
|
||||||
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
|
|
||||||
|
|
||||||
bf16* hOut=new bf16[nQ];
|
|
||||||
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
|
|
||||||
|
|
||||||
float* ref=new float[nQ];
|
|
||||||
cpu_attention_ref(hQ, hK, hV, nullptr, ref, B, Hq, Hk, ql, kl, D, causal ? 0 : -1);
|
|
||||||
|
|
||||||
float max_abs_err=0, max_rel_err=0;
|
|
||||||
for (size_t i=0;i<nQ;i++) {
|
|
||||||
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
|
||||||
if(err>max_abs_err) max_abs_err=err;
|
|
||||||
float rel=err/fmaxf(fabsf(ref[i]), 1e-8f);
|
|
||||||
if(rel>max_rel_err) max_rel_err=rel;
|
|
||||||
}
|
|
||||||
const float atol=0.01f, rtol=0.01f;
|
|
||||||
bool pass=true;
|
|
||||||
for (size_t i=0;i<nQ;i++) {
|
|
||||||
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
|
||||||
if (err > atol + rtol * fabsf(ref[i])) { pass=false; break; }
|
|
||||||
}
|
|
||||||
printf("kernel: %.3f ms max_abs_err: %.6e max_rel_err: %.6e %s\n\n",
|
|
||||||
kms, max_abs_err, max_rel_err, pass?"PASS":"FAIL");
|
|
||||||
|
|
||||||
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
|
|
||||||
delete[]hQ;delete[]hK;delete[]hV;delete[]hOut;delete[]ref;delete[]tmp;
|
|
||||||
|
|
||||||
return pass ? 0 : 1;
|
|
||||||
}
|
|
||||||
|
|
||||||
int main() {
|
|
||||||
const int configs[][7] = {
|
|
||||||
{1,2,1,64,128,64,0}, // tiny: B,Hq,Hk,q,kv,D,causal
|
|
||||||
{1,32,4,512,512,128,0}, // standard
|
|
||||||
{1,32,4,128,256,128,0}, // medium
|
|
||||||
{1,4,2,256,256,128,1}, // causal
|
|
||||||
};
|
|
||||||
int n_configs = sizeof(configs) / sizeof(configs[0]);
|
|
||||||
int fail = 0;
|
|
||||||
|
|
||||||
for (int ci = 0; ci < n_configs; ci++) {
|
|
||||||
int B=configs[ci][0], Hq=configs[ci][1], Hk=configs[ci][2];
|
|
||||||
int ql=configs[ci][3], kl=configs[ci][4], D=configs[ci][5];
|
|
||||||
int causal=configs[ci][6];
|
|
||||||
fail += run_test(B, Hq, Hk, ql, kl, D, causal);
|
|
||||||
if (fail) break;
|
|
||||||
}
|
|
||||||
|
|
||||||
if (fail) {
|
|
||||||
printf("FAILED\n");
|
|
||||||
return fail;
|
|
||||||
}
|
|
||||||
printf("All tests passed!\n");
|
|
||||||
bench();
|
|
||||||
return 0;
|
|
||||||
}
|
|
||||||
@@ -0,0 +1,346 @@
|
|||||||
|
/*
|
||||||
|
Pure-C test — uses shared dispatcher. Combines the decode (split-KV) and
|
||||||
|
prefill (split-Q) correctness checks + benchmarks into one binary.
|
||||||
|
nvcc -I csrc -arch=sm_89 -O3 \
|
||||||
|
--use_fast_math --ptxas-options=-O3 --extra-device-vectorization \
|
||||||
|
-Xcompiler -fopenmp csrc/tests/attn_test.cu -o test && ./test
|
||||||
|
*/
|
||||||
|
|
||||||
|
#include "test_utils.cuh"
|
||||||
|
#include "../kernels/attn_dispatchers.cuh"
|
||||||
|
|
||||||
|
// Split-K scratch (torch-free)
|
||||||
|
struct DecodeScratch {
|
||||||
|
float* o_part = nullptr;
|
||||||
|
float* ml_part = nullptr;
|
||||||
|
};
|
||||||
|
|
||||||
|
static void setup_scratch(AttentionParams<bf16>& p, DecodeScratch& sc) {
|
||||||
|
int max_splits = 32;
|
||||||
|
cudaMalloc(&sc.o_part, (size_t)p.batch * p.q_head * max_splits * p.head_dim * sizeof(float));
|
||||||
|
cudaMalloc(&sc.ml_part, (size_t)p.batch * p.q_head * max_splits * 2 * sizeof(float));
|
||||||
|
}
|
||||||
|
|
||||||
|
static void free_scratch(DecodeScratch& sc) {
|
||||||
|
cudaFree(sc.o_part); cudaFree(sc.ml_part);
|
||||||
|
}
|
||||||
|
|
||||||
|
// ======================================================================
|
||||||
|
// DECODE
|
||||||
|
// ======================================================================
|
||||||
|
|
||||||
|
static int run_decode_test(int B, int Hq, int Hk, int sl, int D, int causal) {
|
||||||
|
int gs = Hq / Hk;
|
||||||
|
|
||||||
|
size_t nQ = B*Hq*1*D, nKV = B*Hk*sl*D;
|
||||||
|
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
|
||||||
|
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
|
||||||
|
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
|
||||||
|
|
||||||
|
bool* hMask=new bool[B*sl];
|
||||||
|
for (int i=0;i<B*sl;i++) hMask[i]=true;
|
||||||
|
|
||||||
|
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
||||||
|
bool* dMask;
|
||||||
|
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
||||||
|
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
||||||
|
cudaMalloc(&dMask,B*sl);
|
||||||
|
|
||||||
|
tmp=new bf16[max(nQ,nKV)];
|
||||||
|
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
|
||||||
|
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
|
||||||
|
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
|
||||||
|
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||||
|
cudaMemcpy(dMask,hMask,B*sl,cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
AttentionParams<bf16> p;
|
||||||
|
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=1; p.kv_len=sl; p.head_dim=D;
|
||||||
|
p.use_mask=0; p.causal_offset=causal?0:-1;
|
||||||
|
p.scale=1.0f/sqrtf((float)D);
|
||||||
|
set_default_strides(p);
|
||||||
|
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||||
|
|
||||||
|
DecodeScratch sc;
|
||||||
|
setup_scratch(p, sc);
|
||||||
|
p.o_part = sc.o_part; p.ml_part = sc.ml_part;
|
||||||
|
|
||||||
|
double t0=now_ms();
|
||||||
|
dispatch_by_head_dim(D, [&]<int H>() { dispatch_decode<H>(p, 0); });
|
||||||
|
cudaDeviceSynchronize();
|
||||||
|
(void)t0;
|
||||||
|
cudaError_t err=cudaGetLastError();
|
||||||
|
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
|
||||||
|
|
||||||
|
bf16* hOut=new bf16[nQ];
|
||||||
|
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
|
||||||
|
|
||||||
|
float* ref=new float[nQ];
|
||||||
|
cpu_attention_ref(hQ, hK, hV, hMask, ref, B, Hq, Hk, 1, sl, D, causal ? 0 : -1);
|
||||||
|
|
||||||
|
float max_abs_err=0, max_rel_err=0;
|
||||||
|
for (size_t i=0;i<nQ;i++){
|
||||||
|
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
||||||
|
if(err>max_abs_err) max_abs_err=err;
|
||||||
|
float rel=err/fmaxf(fabsf(ref[i]), 1e-8f);
|
||||||
|
if(rel>max_rel_err) max_rel_err=rel;
|
||||||
|
}
|
||||||
|
const float atol=0.01f, rtol=0.01f;
|
||||||
|
bool pass=true;
|
||||||
|
for (size_t i=0;i<nQ;i++){
|
||||||
|
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
||||||
|
if (err > atol + rtol * fabsf(ref[i])) { pass=false; break; }
|
||||||
|
}
|
||||||
|
char cfg[64];
|
||||||
|
snprintf(cfg, sizeof(cfg), "B=%2d Hq=%2d Hk=%d seq=%4d D=%3d causal=%d",
|
||||||
|
B, Hq, Hk, sl, D, causal);
|
||||||
|
print_test_row(cfg, max_abs_err, max_rel_err, pass);
|
||||||
|
|
||||||
|
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);cudaFree(dMask);
|
||||||
|
free_scratch(sc);
|
||||||
|
delete[]hQ;delete[]hK;delete[]hV;delete[]hMask;delete[]hOut;delete[]ref;delete[]tmp;
|
||||||
|
|
||||||
|
return pass ? 0 : 1;
|
||||||
|
}
|
||||||
|
|
||||||
|
static void bench_decode() {
|
||||||
|
const int cfgs[][5] = {
|
||||||
|
{1, 32, 4, 512, 128},
|
||||||
|
{1, 32, 4, 1024, 128},
|
||||||
|
{1, 32, 4, 2048, 128},
|
||||||
|
{1, 32, 4, 4096, 128},
|
||||||
|
{16, 32, 4, 2048, 128},
|
||||||
|
{32, 32, 4, 1024, 128},
|
||||||
|
};
|
||||||
|
const int WARMUP = 3, ITERS = 10;
|
||||||
|
printf("\n===== DECODE BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS);
|
||||||
|
print_bench_header();
|
||||||
|
|
||||||
|
for (int ci = 0; ci < 6; ci++) {
|
||||||
|
int B = cfgs[ci][0], Hq = cfgs[ci][1], Hk = cfgs[ci][2];
|
||||||
|
int sl = cfgs[ci][3], D = cfgs[ci][4];
|
||||||
|
size_t nQ = (size_t)B * Hq * D;
|
||||||
|
size_t nKV = (size_t)B * Hk * sl * D;
|
||||||
|
|
||||||
|
bf16 *dQ, *dK, *dV, *dO;
|
||||||
|
cudaMalloc(&dQ, nQ*2); cudaMalloc(&dK, nKV*2);
|
||||||
|
cudaMalloc(&dV, nKV*2); cudaMalloc(&dO, nQ*2);
|
||||||
|
size_t big = nQ > nKV ? nQ : nKV; bf16* tmp = new bf16[big];
|
||||||
|
for (size_t i = 0; i < nQ; i++) tmp[i] = f2bf(randf());
|
||||||
|
cudaMemcpy(dQ, tmp, nQ*2, cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf());
|
||||||
|
cudaMemcpy(dK, tmp, nKV*2, cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf());
|
||||||
|
cudaMemcpy(dV, tmp, nKV*2, cudaMemcpyHostToDevice);
|
||||||
|
delete[] tmp;
|
||||||
|
|
||||||
|
AttentionParams<bf16> p;
|
||||||
|
p.batch = B; p.q_head = Hq; p.kv_head = Hk; p.q_len = 1; p.kv_len = sl;
|
||||||
|
p.head_dim = D; p.use_mask = 0; p.causal_offset = -1;
|
||||||
|
p.scale = 1.0f / sqrtf((float)D);
|
||||||
|
set_default_strides(p);
|
||||||
|
p.q = dQ; p.k = dK; p.v = dV; p.mask = nullptr; p.o = dO;
|
||||||
|
|
||||||
|
DecodeScratch sc;
|
||||||
|
setup_scratch(p, sc);
|
||||||
|
p.o_part = sc.o_part; p.ml_part = sc.ml_part;
|
||||||
|
|
||||||
|
auto launch = [&]() { dispatch_by_head_dim(D, [&]<int H>() { dispatch_decode<H>(p, 0); }); };
|
||||||
|
double flops = 4.0 * B * Hq * (double)sl * D;
|
||||||
|
BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops);
|
||||||
|
|
||||||
|
char cfg[64];
|
||||||
|
snprintf(cfg, sizeof(cfg),
|
||||||
|
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d causal=%d",
|
||||||
|
B, Hq, Hk, 1, sl, D, 0);
|
||||||
|
print_bench_row(cfg, r);
|
||||||
|
|
||||||
|
cudaFree(dQ); cudaFree(dK); cudaFree(dV); cudaFree(dO);
|
||||||
|
free_scratch(sc);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ======================================================================
|
||||||
|
// PREFILL
|
||||||
|
// ======================================================================
|
||||||
|
|
||||||
|
static int run_prefill_test(int B, int Hq, int Hk, int ql, int kl, int D, int causal) {
|
||||||
|
size_t nQ = B*Hq*ql*D, nKV = B*Hk*kl*D;
|
||||||
|
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
|
||||||
|
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
|
||||||
|
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
|
||||||
|
|
||||||
|
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
||||||
|
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
||||||
|
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
||||||
|
tmp=new bf16[max(nQ,nKV)];
|
||||||
|
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
|
||||||
|
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
|
||||||
|
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
|
||||||
|
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
AttentionParams<bf16> p;
|
||||||
|
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D;
|
||||||
|
p.use_mask=0; p.causal_offset=causal?0:-1;
|
||||||
|
set_default_strides(p);
|
||||||
|
p.scale=1.0f/sqrtf((float)D);
|
||||||
|
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||||
|
|
||||||
|
double t0=now_ms();
|
||||||
|
dispatch_by_head_dim(D, [&]<int H>() { dispatch_prefill<H>(p, 0); });
|
||||||
|
cudaDeviceSynchronize();
|
||||||
|
(void)t0;
|
||||||
|
cudaError_t err=cudaGetLastError();
|
||||||
|
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
|
||||||
|
|
||||||
|
bf16* hOut=new bf16[nQ];
|
||||||
|
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
|
||||||
|
|
||||||
|
float* ref=new float[nQ];
|
||||||
|
cpu_attention_ref(hQ, hK, hV, nullptr, ref, B, Hq, Hk, ql, kl, D, causal ? 0 : -1);
|
||||||
|
|
||||||
|
float max_abs_err=0, max_rel_err=0;
|
||||||
|
for (size_t i=0;i<nQ;i++) {
|
||||||
|
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
||||||
|
if(err>max_abs_err) max_abs_err=err;
|
||||||
|
float rel=err/fmaxf(fabsf(ref[i]), 1e-8f);
|
||||||
|
if(rel>max_rel_err) max_rel_err=rel;
|
||||||
|
}
|
||||||
|
const float atol=0.01f, rtol=0.01f;
|
||||||
|
bool pass=true;
|
||||||
|
for (size_t i=0;i<nQ;i++) {
|
||||||
|
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
||||||
|
if (err > atol + rtol * fabsf(ref[i])) { pass=false; break; }
|
||||||
|
}
|
||||||
|
char cfg[64];
|
||||||
|
snprintf(cfg, sizeof(cfg), "B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d causal=%d",
|
||||||
|
B, Hq, Hk, ql, kl, D, causal);
|
||||||
|
print_test_row(cfg, max_abs_err, max_rel_err, pass);
|
||||||
|
|
||||||
|
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
|
||||||
|
delete[]hQ;delete[]hK;delete[]hV;delete[]hOut;delete[]ref;delete[]tmp;
|
||||||
|
|
||||||
|
return pass ? 0 : 1;
|
||||||
|
}
|
||||||
|
|
||||||
|
static void bench_prefill() {
|
||||||
|
const int cfgs[][7] = {
|
||||||
|
{1,32,4,512,512,128,0},
|
||||||
|
{1,32,4,1024,1024,128,0},
|
||||||
|
{1,32,4,2048,2048,128,0},
|
||||||
|
{1,32,4,2048,2048,128,1},
|
||||||
|
{4,32,4,2048,2048,128,1},
|
||||||
|
{1,32,4,4096,4096,128,1},
|
||||||
|
};
|
||||||
|
int n = sizeof(cfgs)/sizeof(cfgs[0]);
|
||||||
|
const int WARMUP = 3, ITERS = 10;
|
||||||
|
printf("\n===== PREFILL BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS);
|
||||||
|
print_bench_header();
|
||||||
|
|
||||||
|
for (int ci = 0; ci < n; ci++) {
|
||||||
|
int B=cfgs[ci][0], Hq=cfgs[ci][1], Hk=cfgs[ci][2];
|
||||||
|
int ql=cfgs[ci][3], kl=cfgs[ci][4], D=cfgs[ci][5], causal=cfgs[ci][6];
|
||||||
|
size_t nQ=(size_t)B*Hq*ql*D, nKV=(size_t)B*Hk*kl*D;
|
||||||
|
|
||||||
|
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
||||||
|
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
||||||
|
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
||||||
|
size_t big = nQ>nKV?nQ:nKV; tmp=new bf16[big];
|
||||||
|
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(randf());
|
||||||
|
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(randf());
|
||||||
|
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(randf());
|
||||||
|
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
AttentionParams<bf16> p;
|
||||||
|
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D;
|
||||||
|
p.use_mask=0; p.causal_offset=causal?0:-1;
|
||||||
|
set_default_strides(p);
|
||||||
|
p.scale=1.0f/sqrtf((float)D);
|
||||||
|
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||||
|
|
||||||
|
auto launch = [&]() { dispatch_by_head_dim(D, [&]<int H>() { dispatch_prefill<H>(p, 0); }); };
|
||||||
|
for (int i=0;i<WARMUP;i++) launch();
|
||||||
|
cudaDeviceSynchronize();
|
||||||
|
cudaError_t err=cudaGetLastError();
|
||||||
|
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return;}
|
||||||
|
|
||||||
|
cudaEvent_t s,e; cudaEventCreate(&s); cudaEventCreate(&e);
|
||||||
|
cudaEventRecord(s);
|
||||||
|
for (int i=0;i<ITERS;i++) launch();
|
||||||
|
cudaEventRecord(e); cudaEventSynchronize(e);
|
||||||
|
float ms=0; cudaEventElapsedTime(&ms,s,e); ms/=ITERS;
|
||||||
|
|
||||||
|
double flops = 4.0*B*Hq*(double)ql*kl*D;
|
||||||
|
if (causal) flops *= 0.5;
|
||||||
|
double tflops = flops/(ms*1e-3)/1e12;
|
||||||
|
BenchResult r{ms, tflops};
|
||||||
|
|
||||||
|
char cfg[64];
|
||||||
|
snprintf(cfg, sizeof(cfg),
|
||||||
|
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d causal=%d",
|
||||||
|
B,Hq,Hk,ql,kl,D,causal);
|
||||||
|
print_bench_row(cfg, r);
|
||||||
|
|
||||||
|
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
|
||||||
|
delete[]tmp; cudaEventDestroy(s); cudaEventDestroy(e);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ======================================================================
|
||||||
|
// MAIN
|
||||||
|
// ======================================================================
|
||||||
|
|
||||||
|
int main() {
|
||||||
|
int fail = 0;
|
||||||
|
|
||||||
|
// ---- DECODE ----
|
||||||
|
{
|
||||||
|
const int configs[][6] = {
|
||||||
|
{1, 2, 1, 64, 32, 0},
|
||||||
|
{1, 32, 4, 512, 128, 0},
|
||||||
|
{1, 32, 4, 1024, 128, 0},
|
||||||
|
{1, 32, 4, 512, 128, 1},
|
||||||
|
};
|
||||||
|
int n_cfgs = sizeof(configs) / sizeof(configs[0]);
|
||||||
|
printf("=== DECODE TESTS ===\n");
|
||||||
|
print_test_header();
|
||||||
|
for (int ci = 0; ci < n_cfgs; ci++) {
|
||||||
|
int B = configs[ci][0], Hq = configs[ci][1], Hk = configs[ci][2];
|
||||||
|
int sl = configs[ci][3], D = configs[ci][4], causal = configs[ci][5];
|
||||||
|
fail += run_decode_test(B, Hq, Hk, sl, D, causal);
|
||||||
|
if (fail) break;
|
||||||
|
}
|
||||||
|
if (fail) { printf("FAILED decode tests\n"); return fail; }
|
||||||
|
bench_decode();
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- PREFILL ----
|
||||||
|
{
|
||||||
|
const int configs[][7] = {
|
||||||
|
{1,2,1,64,128,64,0}, // tiny: B,Hq,Hk,q,kv,D,causal
|
||||||
|
{1,32,4,512,512,128,0}, // standard
|
||||||
|
{1,32,4,128,256,128,0}, // medium
|
||||||
|
{1,4,2,256,256,128,1}, // causal
|
||||||
|
};
|
||||||
|
int n_configs = sizeof(configs) / sizeof(configs[0]);
|
||||||
|
printf("\n=== PREFILL TESTS ===\n");
|
||||||
|
print_test_header();
|
||||||
|
for (int ci = 0; ci < n_configs; ci++) {
|
||||||
|
int B=configs[ci][0], Hq=configs[ci][1], Hk=configs[ci][2];
|
||||||
|
int ql=configs[ci][3], kl=configs[ci][4], D=configs[ci][5];
|
||||||
|
int causal=configs[ci][6];
|
||||||
|
fail += run_prefill_test(B, Hq, Hk, ql, kl, D, causal);
|
||||||
|
if (fail) break;
|
||||||
|
}
|
||||||
|
if (fail) { printf("FAILED prefill tests\n"); return fail; }
|
||||||
|
bench_prefill();
|
||||||
|
}
|
||||||
|
|
||||||
|
printf("\nAll tests passed!\n");
|
||||||
|
return 0;
|
||||||
|
}
|
||||||
@@ -29,19 +29,18 @@ inline double now_ms() {
|
|||||||
|
|
||||||
struct BenchResult {
|
struct BenchResult {
|
||||||
float ms;
|
float ms;
|
||||||
double gbps;
|
|
||||||
double tflops;
|
double tflops;
|
||||||
};
|
};
|
||||||
|
|
||||||
template <typename Fn>
|
template <typename Fn>
|
||||||
BenchResult bench_kernel(Fn launch, int warmup, int iters,
|
BenchResult bench_kernel(Fn launch, int warmup, int iters,
|
||||||
double flops, double bytes) {
|
double flops) {
|
||||||
for (int i = 0; i < warmup; i++) launch();
|
for (int i = 0; i < warmup; i++) launch();
|
||||||
cudaDeviceSynchronize();
|
cudaDeviceSynchronize();
|
||||||
cudaError_t err = cudaGetLastError();
|
cudaError_t err = cudaGetLastError();
|
||||||
if (err != cudaSuccess) {
|
if (err != cudaSuccess) {
|
||||||
printf("CUDA error before bench: %s\n", cudaGetErrorString(err));
|
printf("CUDA error before bench: %s\n", cudaGetErrorString(err));
|
||||||
return {0, 0, 0};
|
return {0, 0};
|
||||||
}
|
}
|
||||||
|
|
||||||
cudaEvent_t s, e;
|
cudaEvent_t s, e;
|
||||||
@@ -52,19 +51,33 @@ BenchResult bench_kernel(Fn launch, int warmup, int iters,
|
|||||||
float ms = 0; cudaEventElapsedTime(&ms, s, e); ms /= iters;
|
float ms = 0; cudaEventElapsedTime(&ms, s, e); ms /= iters;
|
||||||
cudaEventDestroy(s); cudaEventDestroy(e);
|
cudaEventDestroy(s); cudaEventDestroy(e);
|
||||||
|
|
||||||
return {ms, bytes / (ms * 1e-3) / 1e9, flops / (ms * 1e-3) / 1e12};
|
return {ms, flops / (ms * 1e-3) / 1e12};
|
||||||
}
|
}
|
||||||
|
|
||||||
inline void print_bench_header() {
|
inline void print_bench_header() {
|
||||||
printf("%-46s | %10s | %10s | %10s\n",
|
printf("%-46s | %10s | %10s\n",
|
||||||
"config", "latency", "bandwidth", "throughput");
|
"config", "latency", "TFLOP/s");
|
||||||
printf("---------------------------------------------------------------"
|
printf("---------------------------------------------------------------"
|
||||||
"----------------------------\n");
|
"----------------------------\n");
|
||||||
}
|
}
|
||||||
|
|
||||||
inline void print_bench_row(const char* cfg, const BenchResult& r) {
|
inline void print_bench_row(const char* cfg, const BenchResult& r) {
|
||||||
printf("%-46s | %7.4f ms | %7.1f GB/s | %6.2f TFLOP/s\n",
|
printf("%-46s | %7.4f ms | %6.2f\n",
|
||||||
cfg, r.ms, r.gbps, r.tflops);
|
cfg, r.ms, r.tflops);
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- validation table (kernel vs CPU reference) ----
|
||||||
|
inline void print_test_header() {
|
||||||
|
printf("%-46s | %11s | %11s | %6s\n",
|
||||||
|
"config", "max_abs_err", "max_rel_err", "result");
|
||||||
|
printf("----------------------------------------------------------------"
|
||||||
|
"----------------------------\n");
|
||||||
|
}
|
||||||
|
|
||||||
|
inline void print_test_row(const char* cfg, float max_abs_err,
|
||||||
|
float max_rel_err, bool pass) {
|
||||||
|
printf("%-46s | %11.3e | %11.3e | %s\n",
|
||||||
|
cfg, max_abs_err, max_rel_err, pass ? "PASS" : "FAIL");
|
||||||
}
|
}
|
||||||
|
|
||||||
template <int... Ds>
|
template <int... Ds>
|
||||||
@@ -135,9 +148,10 @@ static void cpu_attention_ref(
|
|||||||
float scale = 1.0f / sqrtf((float)D);
|
float scale = 1.0f / sqrtf((float)D);
|
||||||
int n_rep = Hq / Hk;
|
int n_rep = Hq / Hk;
|
||||||
for (int b = 0; b < B; b++) {
|
for (int b = 0; b < B; b++) {
|
||||||
|
#pragma omp parallel for collapse(2) schedule(dynamic)
|
||||||
for (int h = 0; h < Hq; h++) {
|
for (int h = 0; h < Hq; h++) {
|
||||||
int kv_h = h / n_rep;
|
|
||||||
for (int qi = 0; qi < q_len; qi++) {
|
for (int qi = 0; qi < q_len; qi++) {
|
||||||
|
int kv_h = h / n_rep;
|
||||||
float mv = -INFINITY, sv = 0.0f;
|
float mv = -INFINITY, sv = 0.0f;
|
||||||
float accum[256] = {0.0f};
|
float accum[256] = {0.0f};
|
||||||
int lim = kv_len;
|
int lim = kv_len;
|
||||||
|
|||||||
@@ -62,6 +62,8 @@
|
|||||||
|
|
||||||
**1. 安装**
|
**1. 安装**
|
||||||
|
|
||||||
|
AstrAI 需要 Python 3.12+,并精确固定 PyTorch 版本为 `2.11.0`。训练、`scripts/tools/generate.py`、生成式评估和生成演示需要 CUDA;CPU 支持仅适用于提供明确 CPU 设备路径的组件,例如 HTTP 服务和直接打分评估。
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
git clone https://github.com/ViperEkura/AstrAI.git
|
git clone https://github.com/ViperEkura/AstrAI.git
|
||||||
cd AstrAI
|
cd AstrAI
|
||||||
@@ -138,7 +140,7 @@ curl http://localhost:8000/v1/chat/completions \
|
|||||||
# 下载模型权重(运行演示前必需)
|
# 下载模型权重(运行演示前必需)
|
||||||
python scripts/demo/download.py # model → params/
|
python scripts/demo/download.py # model → params/
|
||||||
|
|
||||||
# 交互式流式聊天(多轮对话,保持历史记录)
|
# 单轮交互式流式提示循环(不保留对话历史)
|
||||||
python scripts/demo/stream_chat.py
|
python scripts/demo/stream_chat.py
|
||||||
# 在 >> 后输入消息,输入 !exit 退出
|
# 在 >> 后输入消息,输入 !exit 退出
|
||||||
|
|
||||||
@@ -189,7 +191,7 @@ docker run --gpus all -v /path/to/data:/data -it astrai:latest
|
|||||||
# Docker Compose(GPU,默认)
|
# Docker Compose(GPU,默认)
|
||||||
docker compose up -d
|
docker compose up -d
|
||||||
|
|
||||||
# Docker Compose(仅 CPU)
|
# Docker Compose CPU 服务配置(不支持仅限 CUDA 的生成脚本和演示)
|
||||||
docker compose --profile cpu up -d
|
docker compose --profile cpu up -d
|
||||||
```
|
```
|
||||||
|
|
||||||
@@ -239,7 +241,7 @@ SSE 流式格式、错误码和统计端点详见[推理文档](guides/inference
|
|||||||
|
|
||||||
### 贡献
|
### 贡献
|
||||||
|
|
||||||
我们欢迎贡献!请参阅[贡献指南](../../CONTRIBUTING.md)了解详情。
|
我们欢迎贡献!请参阅[贡献指南](../CONTRIBUTING.md)了解详情。
|
||||||
|
|
||||||
1. Fork 本仓库。
|
1. Fork 本仓库。
|
||||||
2. 创建功能分支。
|
2. 创建功能分支。
|
||||||
@@ -256,7 +258,7 @@ SSE 流式格式、错误码和统计端点详见[推理文档](guides/inference
|
|||||||
|
|
||||||
### 许可证
|
### 许可证
|
||||||
|
|
||||||
本项目采用 [GPL-3.0 许可证](../../LICENSE)。
|
本项目采用 [GPL-3.0 许可证](../LICENSE)。
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
|||||||
+110
-104
@@ -4,7 +4,7 @@
|
|||||||
|
|
||||||
- [Class Diagram](#class-diagram) — Full Mermaid class diagram across 10+ namespaces
|
- [Class Diagram](#class-diagram) — Full Mermaid class diagram across 10+ namespaces
|
||||||
- [Module Overview](#module-overview) — Component inventory per module
|
- [Module Overview](#module-overview) — Component inventory per module
|
||||||
- [Design Patterns](#design-patterns) — 13 documented patterns with classes
|
- [Design Patterns](#design-patterns) — 15 documented patterns with classes
|
||||||
- [Core Relationships](#core-relationships) — 11 key inter-component relationships
|
- [Core Relationships](#core-relationships) — 11 key inter-component relationships
|
||||||
|
|
||||||
## Class Diagram
|
## Class Diagram
|
||||||
@@ -49,6 +49,11 @@ classDiagram
|
|||||||
+Optional[int] n_shared_experts
|
+Optional[int] n_shared_experts
|
||||||
+Optional[int] n_activated_experts
|
+Optional[int] n_activated_experts
|
||||||
+Optional[str] topk_method
|
+Optional[str] topk_method
|
||||||
|
+Optional[int] moe_intermediate_size
|
||||||
|
+Optional[int] shared_expert_intermediate_size
|
||||||
|
+bool norm_topk_prob
|
||||||
|
+int decoder_sparse_step
|
||||||
|
+Optional[List[int]] mlp_only_layers
|
||||||
}
|
}
|
||||||
|
|
||||||
class EncoderConfig {
|
class EncoderConfig {
|
||||||
@@ -63,6 +68,7 @@ classDiagram
|
|||||||
+Optional[int] num_attention_heads
|
+Optional[int] num_attention_heads
|
||||||
+Optional[int] num_key_value_heads
|
+Optional[int] num_key_value_heads
|
||||||
+Optional[bool] use_qk_norm
|
+Optional[bool] use_qk_norm
|
||||||
|
+Optional[bool] use_gated_attention
|
||||||
+str ffn_type
|
+str ffn_type
|
||||||
+Optional[dict] rope_scaling
|
+Optional[dict] rope_scaling
|
||||||
+Optional[str] pooling_type
|
+Optional[str] pooling_type
|
||||||
@@ -114,22 +120,25 @@ classDiagram
|
|||||||
+Dataset dataset
|
+Dataset dataset
|
||||||
+Callable optimizer_fn
|
+Callable optimizer_fn
|
||||||
+Callable scheduler_fn
|
+Callable scheduler_fn
|
||||||
|
+Optional[str] optimizer_name
|
||||||
|
+Dict[str, Any] optimizer_hyperparameters
|
||||||
+int n_epoch
|
+int n_epoch
|
||||||
+int batch_per_device
|
+int batch_per_device
|
||||||
+int grad_accum_steps
|
+int grad_accum_steps
|
||||||
+Optional[float] max_grad_norm
|
+Optional[float] max_grad_norm
|
||||||
+list gradient_checkpointing_modules
|
+list gradient_checkpointing_modules
|
||||||
|
+Optional[str] compile_mode
|
||||||
+int start_epoch
|
+int start_epoch
|
||||||
+int start_samples
|
+int start_samples
|
||||||
+str ckpt_dir
|
+str ckpt_dir
|
||||||
+int ckpt_interval
|
+int ckpt_interval
|
||||||
+str log_dir
|
|
||||||
+List[str] metrics
|
+List[str] metrics
|
||||||
+Optional[LoRAConfig] lora
|
+Optional[LoRAConfig] lora
|
||||||
+int random_seed
|
+int random_seed
|
||||||
+int num_workers
|
+int num_workers
|
||||||
+Optional[int] prefetch_factor
|
+Optional[int] prefetch_factor
|
||||||
+bool pin_memory
|
+bool pin_memory
|
||||||
|
+Optional[Callable] collate_fn
|
||||||
+int nprocs
|
+int nprocs
|
||||||
+str backend
|
+str backend
|
||||||
+str master_addr
|
+str master_addr
|
||||||
@@ -140,6 +149,7 @@ classDiagram
|
|||||||
+Optional[float] val_split
|
+Optional[float] val_split
|
||||||
+int val_step
|
+int val_step
|
||||||
+float neftune_alpha
|
+float neftune_alpha
|
||||||
|
+float moe_aux_loss_coef
|
||||||
+str parallel_mode
|
+str parallel_mode
|
||||||
+int rollout_interval
|
+int rollout_interval
|
||||||
+float rollout_temperature
|
+float rollout_temperature
|
||||||
@@ -149,7 +159,6 @@ classDiagram
|
|||||||
+Optional[Callable] reward_model_fn
|
+Optional[Callable] reward_model_fn
|
||||||
+dict executor_kwargs
|
+dict executor_kwargs
|
||||||
+dict extra_kwargs
|
+dict extra_kwargs
|
||||||
+validate()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
}
|
}
|
||||||
@@ -205,10 +214,6 @@ classDiagram
|
|||||||
-_fetch_record_key(key, index) Tensor
|
-_fetch_record_key(key, index) Tensor
|
||||||
}
|
}
|
||||||
|
|
||||||
class H5Store {
|
|
||||||
+load(path)
|
|
||||||
}
|
|
||||||
|
|
||||||
class MmapStore {
|
class MmapStore {
|
||||||
+List _mmap_refs
|
+List _mmap_refs
|
||||||
+load(path)
|
+load(path)
|
||||||
@@ -260,11 +265,15 @@ classDiagram
|
|||||||
}
|
}
|
||||||
|
|
||||||
namespace model {
|
namespace model {
|
||||||
class AutoModel {
|
class ModelFactory {
|
||||||
+BaseModelConfig config
|
|
||||||
+Dict _entries
|
+Dict _entries
|
||||||
+register(name) decorator
|
+register(name) decorator
|
||||||
+get_component_class(name) Type
|
+get_component_class(name) Type
|
||||||
|
}
|
||||||
|
|
||||||
|
class AutoModel {
|
||||||
|
<<nn.Module>>
|
||||||
|
+BaseModelConfig config
|
||||||
+from_pretrained(path, disable_random_init, strict) nn.Module
|
+from_pretrained(path, disable_random_init, strict) nn.Module
|
||||||
+save_pretrained(save_directory)
|
+save_pretrained(save_directory)
|
||||||
+to(*args, **kwargs) Self
|
+to(*args, **kwargs) Self
|
||||||
@@ -299,7 +308,13 @@ classDiagram
|
|||||||
+RMSNorm input_norm
|
+RMSNorm input_norm
|
||||||
+nn.Module mlp # MLP or DeepSeekMoE via FFNFactory
|
+nn.Module mlp # MLP or DeepSeekMoE via FFNFactory
|
||||||
+RMSNorm post_attention_norm
|
+RMSNorm post_attention_norm
|
||||||
+forward(x, rotary_emb, attention_mask, kv_cache) Tensor
|
+forward(x, rotary_emb, attention_mask, kv_cache, is_causal) DecoderOutput
|
||||||
|
}
|
||||||
|
|
||||||
|
class DecoderOutput {
|
||||||
|
<<TypedDict>>
|
||||||
|
+Tensor hidden_states
|
||||||
|
+Optional[Tensor] aux_loss
|
||||||
}
|
}
|
||||||
|
|
||||||
class GQA {
|
class GQA {
|
||||||
@@ -314,7 +329,7 @@ classDiagram
|
|||||||
+Linear q_proj, k_proj, v_proj, o_proj
|
+Linear q_proj, k_proj, v_proj, o_proj
|
||||||
+Linear gate # only if use_gated_attention
|
+Linear gate # only if use_gated_attention
|
||||||
+RMSNorm q_norm, k_norm # only if use_qk_norm
|
+RMSNorm q_norm, k_norm # only if use_qk_norm
|
||||||
+forward(x, rotary_emb, attn_mask, kv_cache) Tensor
|
+forward(x, rotary_emb, attn_mask, kv_cache, is_causal) Tensor
|
||||||
}
|
}
|
||||||
|
|
||||||
class MLA {
|
class MLA {
|
||||||
@@ -334,12 +349,18 @@ classDiagram
|
|||||||
+Linear gate # only if use_gated_attention
|
+Linear gate # only if use_gated_attention
|
||||||
+RMSNorm kv_norm
|
+RMSNorm kv_norm
|
||||||
+RMSNorm q_norm, k_norm # only if use_qk_norm
|
+RMSNorm q_norm, k_norm # only if use_qk_norm
|
||||||
+forward(x, rotary_emb, attn_mask, kv_cache) Tensor
|
+forward(x, rotary_emb, attn_mask, kv_cache, is_causal) Tensor
|
||||||
}
|
}
|
||||||
|
|
||||||
class MLP {
|
class MLP {
|
||||||
+Linear up, gate, down
|
+Linear up, gate, down
|
||||||
+forward(x) Tensor
|
+forward(x) FFNOutput
|
||||||
|
}
|
||||||
|
|
||||||
|
class FFNOutput {
|
||||||
|
<<TypedDict>>
|
||||||
|
+Tensor hidden_states
|
||||||
|
+Optional[Tensor] aux_loss
|
||||||
}
|
}
|
||||||
|
|
||||||
class DeepSeekMoE {
|
class DeepSeekMoE {
|
||||||
@@ -351,7 +372,7 @@ classDiagram
|
|||||||
+Linear router
|
+Linear router
|
||||||
+ModuleList shared_experts
|
+ModuleList shared_experts
|
||||||
+ModuleList routed_experts
|
+ModuleList routed_experts
|
||||||
+forward(x) Tensor
|
+forward(x) FFNOutput
|
||||||
}
|
}
|
||||||
|
|
||||||
class AttnFactory {
|
class AttnFactory {
|
||||||
@@ -380,9 +401,8 @@ classDiagram
|
|||||||
+int max_len
|
+int max_len
|
||||||
+float base
|
+float base
|
||||||
+Optional[Dict] rope_scaling
|
+Optional[Dict] rope_scaling
|
||||||
+Tensor cos_table
|
+Tensor freqs_cis
|
||||||
+Tensor sin_table
|
+forward(x, position_ids=None) Tensor
|
||||||
+forward(x, position_ids=None) Tuple[Tensor, Tensor]
|
|
||||||
}
|
}
|
||||||
|
|
||||||
class Embedding {
|
class Embedding {
|
||||||
@@ -486,10 +506,6 @@ classDiagram
|
|||||||
+save(output_dir, domain, shard_idx, tensors)
|
+save(output_dir, domain, shard_idx, tensors)
|
||||||
}
|
}
|
||||||
|
|
||||||
class H5Writer {
|
|
||||||
+save(output_dir, domain, shard_idx, tensors)
|
|
||||||
}
|
|
||||||
|
|
||||||
class Pipeline {
|
class Pipeline {
|
||||||
+PipelineConfig config
|
+PipelineConfig config
|
||||||
+List[str] paths
|
+List[str] paths
|
||||||
@@ -559,7 +575,7 @@ classDiagram
|
|||||||
class Trainer {
|
class Trainer {
|
||||||
+TrainConfig train_config
|
+TrainConfig train_config
|
||||||
+List[TrainCallback] callbacks
|
+List[TrainCallback] callbacks
|
||||||
+train(resume_dir)
|
+train(param_path=None, resume=False)
|
||||||
-_get_default_callbacks() List[TrainCallback]
|
-_get_default_callbacks() List[TrainCallback]
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -576,13 +592,17 @@ classDiagram
|
|||||||
+int epoch
|
+int epoch
|
||||||
+int consumed_samples
|
+int consumed_samples
|
||||||
+float loss
|
+float loss
|
||||||
+float grad_norm
|
+Dict[str, float] metrics
|
||||||
|
+Optional[float] grad_norm
|
||||||
|
+GradSNRTracker grad_snr_tracker
|
||||||
+DataLoader val_dataloader
|
+DataLoader val_dataloader
|
||||||
+float val_loss
|
+Optional[float] val_loss
|
||||||
+int world_size
|
+int world_size
|
||||||
+int rank
|
+int rank
|
||||||
+dict kwargs
|
+dict kwargs
|
||||||
+optimizer_step() int
|
+stop_requested (property) bool
|
||||||
|
+optimizer_step (property) int
|
||||||
|
+request_stop()
|
||||||
}
|
}
|
||||||
|
|
||||||
class TrainContextBuilder {
|
class TrainContextBuilder {
|
||||||
@@ -594,11 +614,22 @@ classDiagram
|
|||||||
class BaseStrategy {
|
class BaseStrategy {
|
||||||
+Callable model
|
+Callable model
|
||||||
+Optional[BaseExecutor] executor
|
+Optional[BaseExecutor] executor
|
||||||
+Optional[Callable] model_fn
|
+float moe_aux_loss_coef
|
||||||
+dict extra_kwargs
|
+dict extra_kwargs
|
||||||
+str device
|
+str device
|
||||||
+__call__(batch) Tensor
|
+__call__(batch) LossOutput
|
||||||
+compute_loss(batch) Tensor
|
+compute_loss(batch) Tensor
|
||||||
|
+compute_loss_output(batch) LossOutput
|
||||||
|
+supports_online() bool
|
||||||
|
+set_rollout_runner(runner)
|
||||||
|
+prepare_from_rollout(result) Dict
|
||||||
|
+on_optimizer_step()
|
||||||
|
}
|
||||||
|
|
||||||
|
class LossOutput {
|
||||||
|
<<TypedDict>>
|
||||||
|
+Tensor loss
|
||||||
|
+Dict[str, float] metrics
|
||||||
}
|
}
|
||||||
|
|
||||||
class StrategyFactory {
|
class StrategyFactory {
|
||||||
@@ -636,9 +667,12 @@ classDiagram
|
|||||||
|
|
||||||
class RawRollout {
|
class RawRollout {
|
||||||
+Tensor prompts
|
+Tensor prompts
|
||||||
|
+Tensor prompt_mask
|
||||||
+Tensor responses
|
+Tensor responses
|
||||||
+Tensor response_mask
|
+Tensor response_mask
|
||||||
+Tensor logprobs_old
|
+Tensor logprobs_old
|
||||||
|
+List[str] prompt_texts
|
||||||
|
+List[List[str]] response_texts
|
||||||
}
|
}
|
||||||
|
|
||||||
class RolloutResult {
|
class RolloutResult {
|
||||||
@@ -647,10 +681,18 @@ classDiagram
|
|||||||
|
|
||||||
class BaseRewardModel {
|
class BaseRewardModel {
|
||||||
<<abstract>>
|
<<abstract>>
|
||||||
+score(prompts, responses) Tensor
|
+score(List[str] prompts, List[List[str]] responses) Tensor
|
||||||
}
|
}
|
||||||
|
|
||||||
class RolloutGenerator {
|
class RolloutGenerator {
|
||||||
|
+InferenceScheduler scheduler
|
||||||
|
+int max_tokens
|
||||||
|
+int group_size
|
||||||
|
+float temperature
|
||||||
|
+int top_k
|
||||||
|
+float top_p
|
||||||
|
+float frequency_penalty
|
||||||
|
+int rep_window
|
||||||
+generate(batch) RawRollout
|
+generate(batch) RawRollout
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -740,7 +782,7 @@ classDiagram
|
|||||||
}
|
}
|
||||||
|
|
||||||
class MetricCallback {
|
class MetricCallback {
|
||||||
+Path log_dir
|
+Path ckpt_dir
|
||||||
+int save_interval
|
+int save_interval
|
||||||
+List[str] metrics
|
+List[str] metrics
|
||||||
+int val_step
|
+int val_step
|
||||||
@@ -764,9 +806,9 @@ classDiagram
|
|||||||
+nn.Module model
|
+nn.Module model
|
||||||
+AutoTokenizer tokenizer
|
+AutoTokenizer tokenizer
|
||||||
+InferenceScheduler scheduler
|
+InferenceScheduler scheduler
|
||||||
+generate(prompt, stream, max_tokens, temperature, top_p, top_k) Union[Generator, str, List[str]]
|
+generate(prompt, stream, max_tokens, temperature, top_p, top_k, frequency_penalty, rep_window) Union[Generator, str, List[str]]
|
||||||
+generate_with_request(request) Union[Generator, str, List[str]]
|
+generate_with_request(request) Union[Generator, str, List[str]]
|
||||||
+generate_async(prompt, max_tokens, temperature, top_p, top_k) AsyncGenerator
|
+generate_async(prompt, max_tokens, temperature, top_p, top_k, frequency_penalty, rep_window) AsyncGenerator
|
||||||
+get_stats() Dict
|
+get_stats() Dict
|
||||||
+shutdown()
|
+shutdown()
|
||||||
}
|
}
|
||||||
@@ -774,18 +816,18 @@ classDiagram
|
|||||||
class Executor {
|
class Executor {
|
||||||
+AutoModel model
|
+AutoModel model
|
||||||
+AutoTokenizer tokenizer
|
+AutoTokenizer tokenizer
|
||||||
+KVCache page_cache
|
+PagePool kv_cache
|
||||||
+Optional[str] device
|
+Optional[str] device
|
||||||
+Optional[torch.dtype] dtype
|
+Optional[torch.dtype] dtype
|
||||||
+execute_prefill(tasks, prompt_len, start_pos)
|
+execute_prefill(tasks, prompt_len, start_pos)
|
||||||
+execute_decode(tasks) List[int]
|
+execute_decode(tasks, return_logprobs=False) Union[List[int], List[Tuple[int, float]]]
|
||||||
}
|
}
|
||||||
|
|
||||||
class InferenceScheduler {
|
class InferenceScheduler {
|
||||||
+KVCache _page_cache
|
+PagePool _cache
|
||||||
+Executor _executor
|
+Executor _executor
|
||||||
+TaskManager _task_mgr
|
+TaskManager _task_mgr
|
||||||
+bool _running
|
+Event _stop_event
|
||||||
+Thread _loop_thread
|
+Thread _loop_thread
|
||||||
+int max_seq_len
|
+int max_seq_len
|
||||||
+str device
|
+str device
|
||||||
@@ -795,6 +837,7 @@ classDiagram
|
|||||||
+start()
|
+start()
|
||||||
+stop()
|
+stop()
|
||||||
+get_stats() Dict
|
+get_stats() Dict
|
||||||
|
+run_batch(prompt_ids_list, max_tokens, temperature, top_p, top_k, frequency_penalty, rep_window, return_logprobs) Union[List[List[int]], List[Tuple[List[int], List[float]]]]
|
||||||
}
|
}
|
||||||
|
|
||||||
class Allocator {
|
class Allocator {
|
||||||
@@ -816,16 +859,6 @@ classDiagram
|
|||||||
+record(page_idx, token_ids, logical_page_idx)
|
+record(page_idx, token_ids, logical_page_idx)
|
||||||
}
|
}
|
||||||
|
|
||||||
class PagePool {
|
|
||||||
-Allocator _alloc
|
|
||||||
-PrefixCache _prefix
|
|
||||||
+alloc() int
|
|
||||||
+free(idx)
|
|
||||||
+inc_ref(idx)
|
|
||||||
+lookup(token_ids) List[int]
|
|
||||||
+record(page_idx, token_ids, logical_page_idx)
|
|
||||||
}
|
|
||||||
|
|
||||||
class KVStorage {
|
class KVStorage {
|
||||||
+int size
|
+int size
|
||||||
+Tensor k_buffer
|
+Tensor k_buffer
|
||||||
@@ -852,8 +885,7 @@ classDiagram
|
|||||||
+Tensor seq_lens
|
+Tensor seq_lens
|
||||||
+Tensor out_cache_loc
|
+Tensor out_cache_loc
|
||||||
+int max_len
|
+int max_len
|
||||||
+Optional[Tensor] page_table
|
+Optional[Tensor] kv_indptr
|
||||||
+Optional[Tensor] decode_mask
|
|
||||||
}
|
}
|
||||||
|
|
||||||
class PagePool {
|
class PagePool {
|
||||||
@@ -878,6 +910,8 @@ classDiagram
|
|||||||
+float temperature
|
+float temperature
|
||||||
+float top_p
|
+float top_p
|
||||||
+int top_k
|
+int top_k
|
||||||
|
+float frequency_penalty
|
||||||
|
+int rep_window
|
||||||
+TaskStatus status
|
+TaskStatus status
|
||||||
+List output_ids
|
+List output_ids
|
||||||
+int input_tokens
|
+int input_tokens
|
||||||
@@ -923,27 +957,29 @@ classDiagram
|
|||||||
+float top_p
|
+float top_p
|
||||||
+float temperature
|
+float temperature
|
||||||
+Optional[int] max_tokens
|
+Optional[int] max_tokens
|
||||||
|
+float frequency_penalty
|
||||||
|
+int rep_window
|
||||||
+bool stream
|
+bool stream
|
||||||
}
|
}
|
||||||
|
|
||||||
class BaseSamplingStrategy {
|
class BaseSamplingStrategy {
|
||||||
<<abstract>>
|
<<abstract>>
|
||||||
+apply(logits, filter_value) Tensor
|
+apply(logits, filter_value, input_ids, input_mask) Tensor
|
||||||
}
|
}
|
||||||
|
|
||||||
class TemperatureStrategy {
|
class TemperatureStrategy {
|
||||||
+float temperature
|
+float temperature
|
||||||
+apply(logits, filter_value) Tensor
|
+apply(logits, filter_value, input_ids, input_mask) Tensor
|
||||||
}
|
}
|
||||||
|
|
||||||
class TopKStrategy {
|
class TopKStrategy {
|
||||||
+int top_k
|
+int top_k
|
||||||
+apply(logits, filter_value) Tensor
|
+apply(logits, filter_value, input_ids, input_mask) Tensor
|
||||||
}
|
}
|
||||||
|
|
||||||
class TopPStrategy {
|
class TopPStrategy {
|
||||||
+float top_p
|
+float top_p
|
||||||
+apply(logits, filter_value) Tensor
|
+apply(logits, filter_value, input_ids, input_mask) Tensor
|
||||||
}
|
}
|
||||||
|
|
||||||
class FrequencyPenaltyStrategy {
|
class FrequencyPenaltyStrategy {
|
||||||
@@ -953,8 +989,8 @@ classDiagram
|
|||||||
|
|
||||||
class SamplingPipeline {
|
class SamplingPipeline {
|
||||||
+List[BaseSamplingStrategy] strategies
|
+List[BaseSamplingStrategy] strategies
|
||||||
+apply(logits, filter_value) Tensor
|
+apply(logits, filter_value, input_ids, input_mask) Tensor
|
||||||
+sample(logits, filter_value) Tensor
|
+sample(logits, filter_value, input_ids, input_mask, return_logprobs) Union[Tensor, Tuple[Tensor, Tensor]]
|
||||||
}
|
}
|
||||||
|
|
||||||
class StreamDecoder {
|
class StreamDecoder {
|
||||||
@@ -1029,7 +1065,7 @@ classDiagram
|
|||||||
<<abstract>>
|
<<abstract>>
|
||||||
+prepare(request, engine) Tuple[str, GenContext, List[str]]
|
+prepare(request, engine) Tuple[str, GenContext, List[str]]
|
||||||
+format_stream_start(ctx) List[str]
|
+format_stream_start(ctx) List[str]
|
||||||
+format_chunk(token) List[str]
|
+format_chunk(token, **kwargs) List[str]
|
||||||
+format_stream_end(ctx, stop) List[str]
|
+format_stream_end(ctx, stop) List[str]
|
||||||
+format_response(ctx, content, stop) Dict
|
+format_response(ctx, content, stop) Dict
|
||||||
}
|
}
|
||||||
@@ -1037,7 +1073,7 @@ classDiagram
|
|||||||
class OpenAIResponseBuilder {
|
class OpenAIResponseBuilder {
|
||||||
+prepare(request, engine) Tuple
|
+prepare(request, engine) Tuple
|
||||||
+format_stream_start(ctx) List[str]
|
+format_stream_start(ctx) List[str]
|
||||||
+format_chunk(token) List[str]
|
+format_chunk(token, **kwargs) List[str]
|
||||||
+format_stream_end(ctx, stop) List[str]
|
+format_stream_end(ctx, stop) List[str]
|
||||||
+format_response(ctx, content, stop) Dict
|
+format_response(ctx, content, stop) Dict
|
||||||
}
|
}
|
||||||
@@ -1045,7 +1081,7 @@ classDiagram
|
|||||||
class AnthropicResponseBuilder {
|
class AnthropicResponseBuilder {
|
||||||
+prepare(request, engine) Tuple
|
+prepare(request, engine) Tuple
|
||||||
+format_stream_start(ctx) List[str]
|
+format_stream_start(ctx) List[str]
|
||||||
+format_chunk(token) List[str]
|
+format_chunk(token, **kwargs) List[str]
|
||||||
+format_stream_end(ctx, stop) List[str]
|
+format_stream_end(ctx, stop) List[str]
|
||||||
+format_response(ctx, content, stop) Dict
|
+format_response(ctx, content, stop) Dict
|
||||||
}
|
}
|
||||||
@@ -1153,10 +1189,13 @@ classDiagram
|
|||||||
|
|
||||||
class BaseExecutor {
|
class BaseExecutor {
|
||||||
+GradientState gradient_state
|
+GradientState gradient_state
|
||||||
+prepare(model_fn, optimizer_fn, scheduler_fn, before_wrap) tuple
|
+prepare(model_fn, optimizer_fn, scheduler_fn, before_wrap, after_wrap) tuple
|
||||||
+accumulate(model) context manager
|
+accumulate(model) context manager
|
||||||
+backward(loss)
|
+backward(loss)
|
||||||
+unwrap_model(model) dict
|
+unwrap_model(model) dict
|
||||||
|
+checkpoint_context(model) context manager
|
||||||
|
+clip_grad_norm(model, max_norm) float
|
||||||
|
+use_distributed (property) bool
|
||||||
+sync_gradients (property) bool
|
+sync_gradients (property) bool
|
||||||
+grad_accum_steps (property) int
|
+grad_accum_steps (property) int
|
||||||
}
|
}
|
||||||
@@ -1173,7 +1212,8 @@ classDiagram
|
|||||||
class FSDPExecutor {
|
class FSDPExecutor {
|
||||||
-_prepare_model(model) nn.Module
|
-_prepare_model(model) nn.Module
|
||||||
-_no_sync(model) context manager
|
-_no_sync(model) context manager
|
||||||
+unwrap_model(model) dict
|
+unwrap_model(model) Optional[dict]
|
||||||
|
+clip_grad_norm(model, max_norm) float
|
||||||
}
|
}
|
||||||
|
|
||||||
class ExecutorFactory {
|
class ExecutorFactory {
|
||||||
@@ -1182,33 +1222,6 @@ classDiagram
|
|||||||
+create(parallel_mode, **kwargs) BaseExecutor
|
+create(parallel_mode, **kwargs) BaseExecutor
|
||||||
}
|
}
|
||||||
|
|
||||||
class ParallelModel {
|
|
||||||
+dist.ProcessGroup process_group
|
|
||||||
+int rank
|
|
||||||
+int world_size
|
|
||||||
}
|
|
||||||
|
|
||||||
class ColumnParallelLinear {
|
|
||||||
+int in_features
|
|
||||||
+int out_features
|
|
||||||
+int out_features_per_rank
|
|
||||||
+bool gather_results
|
|
||||||
+Parameter weight
|
|
||||||
+Optional[Parameter] bias
|
|
||||||
+forward(x) Tensor
|
|
||||||
+load_state_dict(state_dict)
|
|
||||||
}
|
|
||||||
|
|
||||||
class RowParallelLinear {
|
|
||||||
+int in_features
|
|
||||||
+int out_features
|
|
||||||
+int in_features_per_rank
|
|
||||||
+bool reduce_results
|
|
||||||
+Parameter weight
|
|
||||||
+Optional[Parameter] bias
|
|
||||||
+forward(x) Tensor
|
|
||||||
+load_state_dict(state_dict)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
%% Relationships — UML notation: <|-- generalization, *-- composition, o-- aggregation, --> association, ..> dependency
|
%% Relationships — UML notation: <|-- generalization, *-- composition, o-- aggregation, --> association, ..> dependency
|
||||||
@@ -1230,11 +1243,8 @@ classDiagram
|
|||||||
BaseDataset <|-- SFTDataset
|
BaseDataset <|-- SFTDataset
|
||||||
BaseDataset <|-- DPODataset
|
BaseDataset <|-- DPODataset
|
||||||
BaseDataset <|-- GRPODataset
|
BaseDataset <|-- GRPODataset
|
||||||
Store <|-- H5Store
|
|
||||||
Store <|-- MmapStore
|
Store <|-- MmapStore
|
||||||
Store <|-- JsonlStore
|
Store <|-- JsonlStore
|
||||||
H5Store --|> Streamable
|
|
||||||
H5Store --|> Recordable
|
|
||||||
MmapStore --|> Streamable
|
MmapStore --|> Streamable
|
||||||
MmapStore --|> Recordable
|
MmapStore --|> Recordable
|
||||||
JsonlStore --|> Streamable
|
JsonlStore --|> Streamable
|
||||||
@@ -1243,8 +1253,6 @@ classDiagram
|
|||||||
BaseSamplingStrategy <|-- TopKStrategy
|
BaseSamplingStrategy <|-- TopKStrategy
|
||||||
BaseSamplingStrategy <|-- TopPStrategy
|
BaseSamplingStrategy <|-- TopPStrategy
|
||||||
BaseSamplingStrategy <|-- FrequencyPenaltyStrategy
|
BaseSamplingStrategy <|-- FrequencyPenaltyStrategy
|
||||||
ParallelModel <|-- RowParallelLinear
|
|
||||||
ParallelModel <|-- ColumnParallelLinear
|
|
||||||
AutoModel <|-- AutoRegressiveLM
|
AutoModel <|-- AutoRegressiveLM
|
||||||
AutoModel <|-- EmbeddingEncoder
|
AutoModel <|-- EmbeddingEncoder
|
||||||
BaseConfig <|-- BaseModelConfig
|
BaseConfig <|-- BaseModelConfig
|
||||||
@@ -1255,7 +1263,7 @@ classDiagram
|
|||||||
BaseConfig <|-- PipelineConfig
|
BaseConfig <|-- PipelineConfig
|
||||||
BaseModelConfig <|-- AutoRegressiveLMConfig
|
BaseModelConfig <|-- AutoRegressiveLMConfig
|
||||||
BaseModelConfig <|-- EncoderConfig
|
BaseModelConfig <|-- EncoderConfig
|
||||||
BaseFactory <|-- AutoModel
|
BaseFactory <|-- ModelFactory
|
||||||
BaseFactory <|-- AttnFactory
|
BaseFactory <|-- AttnFactory
|
||||||
BaseFactory <|-- FFNFactory
|
BaseFactory <|-- FFNFactory
|
||||||
BaseFactory <|-- DatasetFactory
|
BaseFactory <|-- DatasetFactory
|
||||||
@@ -1286,7 +1294,6 @@ classDiagram
|
|||||||
PositionIdStrategy <|-- DocResetPositionId
|
PositionIdStrategy <|-- DocResetPositionId
|
||||||
PositionIdStrategy <|-- ContinuousPositionId
|
PositionIdStrategy <|-- ContinuousPositionId
|
||||||
StoreWriter <|-- BinWriter
|
StoreWriter <|-- BinWriter
|
||||||
StoreWriter <|-- H5Writer
|
|
||||||
RawRollout <|-- RolloutResult
|
RawRollout <|-- RolloutResult
|
||||||
LaunchStrategy <|-- TorchrunStrategy
|
LaunchStrategy <|-- TorchrunStrategy
|
||||||
LaunchStrategy <|-- LocalStrategy
|
LaunchStrategy <|-- LocalStrategy
|
||||||
@@ -1317,8 +1324,6 @@ classDiagram
|
|||||||
%% --- Aggregation (weak ownership) ---
|
%% --- Aggregation (weak ownership) ---
|
||||||
AutoModel o-- BaseModelConfig
|
AutoModel o-- BaseModelConfig
|
||||||
AutoTokenizer o-- ChatTemplate
|
AutoTokenizer o-- ChatTemplate
|
||||||
PagePool o-- Allocator
|
|
||||||
PagePool o-- PrefixCache
|
|
||||||
Trainer o-- TrainCallback
|
Trainer o-- TrainCallback
|
||||||
TrainContext o-- BaseStrategy
|
TrainContext o-- BaseStrategy
|
||||||
TrainContext o-- BaseScheduler
|
TrainContext o-- BaseScheduler
|
||||||
@@ -1352,11 +1357,12 @@ classDiagram
|
|||||||
FFNFactory ..> DeepSeekMoE : creates
|
FFNFactory ..> DeepSeekMoE : creates
|
||||||
DecoderBlock ..> AttnFactory : uses
|
DecoderBlock ..> AttnFactory : uses
|
||||||
DecoderBlock ..> FFNFactory : uses
|
DecoderBlock ..> FFNFactory : uses
|
||||||
StoreFactory ..> H5Store : creates
|
|
||||||
StoreFactory ..> MmapStore : creates
|
StoreFactory ..> MmapStore : creates
|
||||||
StoreFactory ..> JsonlStore : creates
|
StoreFactory ..> JsonlStore : creates
|
||||||
ConfigFactory ..> AutoRegressiveLMConfig : creates
|
ConfigFactory ..> AutoRegressiveLMConfig : creates
|
||||||
ConfigFactory ..> EncoderConfig : creates
|
ConfigFactory ..> EncoderConfig : creates
|
||||||
|
ModelFactory ..> AutoRegressiveLM : creates
|
||||||
|
ModelFactory ..> EmbeddingEncoder : creates
|
||||||
ExecutorFactory ..> NoneExecutor : creates
|
ExecutorFactory ..> NoneExecutor : creates
|
||||||
ExecutorFactory ..> DDPExecutor : creates
|
ExecutorFactory ..> DDPExecutor : creates
|
||||||
ExecutorFactory ..> FSDPExecutor : creates
|
ExecutorFactory ..> FSDPExecutor : creates
|
||||||
@@ -1399,10 +1405,10 @@ classDiagram
|
|||||||
| Module | Components | Description |
|
| Module | Components | Description |
|
||||||
|--------|------------|-------------|
|
|--------|------------|-------------|
|
||||||
| **astrai.config** | BaseConfig, BaseModelConfig, AutoRegressiveLMConfig, EncoderConfig, ConfigFactory, TrainConfig, PipelineConfig, InputConfig, ProcessingConfig, OutputConfig | Configuration management (to_dict/from_dict, to_file/from_file) |
|
| **astrai.config** | BaseConfig, BaseModelConfig, AutoRegressiveLMConfig, EncoderConfig, ConfigFactory, TrainConfig, PipelineConfig, InputConfig, ProcessingConfig, OutputConfig | Configuration management (to_dict/from_dict, to_file/from_file) |
|
||||||
| **astrai.preprocessing** | SectionRenderer, BaseMaskBuilder, MaskBuilderFactory, SectionedMaskBuilder, SingleOutputMaskBuilder, MultiOutputMaskBuilder, Pipeline, TokenizeTransform, PackingStrategy, PackingStrategyFactory, SimplePacking, BFDPacking, BFDSplitPacking, PositionIdStrategy, PositionIdStrategyFactory, NoPositionId, DocResetPositionId, ContinuousPositionId, StoreWriter, StoreWriterFactory, BinWriter, H5Writer | Declarative JSON-driven data preprocessing |
|
| **astrai.preprocessing** | SectionRenderer, BaseMaskBuilder, MaskBuilderFactory, SectionedMaskBuilder, SingleOutputMaskBuilder, MultiOutputMaskBuilder, Pipeline, TokenizeTransform, PackingStrategy, PackingStrategyFactory, SimplePacking, BFDPacking, BFDSplitPacking, PositionIdStrategy, PositionIdStrategyFactory, NoPositionId, DocResetPositionId, ContinuousPositionId, StoreWriter, StoreWriterFactory, BinWriter | Declarative JSON-driven data preprocessing |
|
||||||
| **astrai.dataset** | BaseDataset, SEQDataset, SFTDataset, DPODataset, GRPODataset, Store, Streamable, Recordable, H5Store, MmapStore, JsonlSource, JsonlStore, StoreFactory, RDSampler, DatasetFactory | Dataset loading and management |
|
| **astrai.dataset** | BaseDataset, SEQDataset, SFTDataset, DPODataset, GRPODataset, Store, Streamable, Recordable, MmapStore, JsonlSource, JsonlStore, StoreFactory, RDSampler, DatasetFactory | Dataset loading and management |
|
||||||
| **astrai.serialization** | Checkpoint | Model serialization |
|
| **astrai.serialization** | Checkpoint | Model serialization |
|
||||||
| **astrai.model** | AutoModel, AutoRegressiveLM, EmbeddingEncoder, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, RMSNorm, Linear, LoRAConfig, LoRALinear, RotaryEmbedding, Embedding | Neural network model |
|
| **astrai.model** | ModelFactory, AutoModel, AutoRegressiveLM, EmbeddingEncoder, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, RMSNorm, Linear, LoRAConfig, LoRALinear, RotaryEmbedding, Embedding | Neural network model |
|
||||||
| **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
|
| **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
|
||||||
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy–GRPOStrategy, StrategyFactory, BaseScheduler–WSDScheduler, SchedulerFactory, TrainCallback(Protocol)–MetricCallback, CallbackFactory, RawRollout, RolloutResult, BaseRewardModel, RolloutGenerator, RolloutRunner | Training workflow |
|
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy–GRPOStrategy, StrategyFactory, BaseScheduler–WSDScheduler, SchedulerFactory, TrainCallback(Protocol)–MetricCallback, CallbackFactory, RawRollout, RolloutResult, BaseRewardModel, RolloutGenerator, RolloutRunner | Training workflow |
|
||||||
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, PagePool, KVStorage, ReqToTokenPool, KVCache, Allocator, PrefixCache, Task, TaskManager, TaskStatus, StreamDecoder, GenerationRequest, GenerateResult, BaseSamplingStrategy–SamplingPipeline, FrequencyPenaltyStrategy, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, StopInfo, ChatMessage, FunctionDef, ToolDef, ChatCompletionRequest, AnthropicMessage, MessagesRequest, BaseToolParser, ToolParserFactory, SimpleJsonToolParser | Inference service |
|
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, PagePool, KVStorage, ReqToTokenPool, KVCache, Allocator, PrefixCache, Task, TaskManager, TaskStatus, StreamDecoder, GenerationRequest, GenerateResult, BaseSamplingStrategy–SamplingPipeline, FrequencyPenaltyStrategy, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, StopInfo, ChatMessage, FunctionDef, ToolDef, ChatCompletionRequest, AnthropicMessage, MessagesRequest, BaseToolParser, ToolParserFactory, SimpleJsonToolParser | Inference service |
|
||||||
@@ -1415,7 +1421,7 @@ classDiagram
|
|||||||
|
|
||||||
| Pattern | Classes | Purpose |
|
| Pattern | Classes | Purpose |
|
||||||
|---------|---------|---------|
|
|---------|---------|---------|
|
||||||
| **Factory** | `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StoreFactory`, `ConfigFactory`, `ExecutorFactory`, `MaskBuilderFactory`, `StoreWriterFactory`, `PackingStrategyFactory`, `PositionIdStrategyFactory`, `ToolParserFactory` | Decorator-based component creation |
|
| **Factory** | `ModelFactory`, `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StoreFactory`, `ConfigFactory`, `ExecutorFactory`, `MaskBuilderFactory`, `StoreWriterFactory`, `PackingStrategyFactory`, `PositionIdStrategyFactory`, `ToolParserFactory` | Decorator-based component creation |
|
||||||
| **Registry** | `BaseFactory` | Component registration |
|
| **Registry** | `BaseFactory` | Component registration |
|
||||||
| **Strategy** | `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy` | Training strategy switching |
|
| **Strategy** | `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy` | Training strategy switching |
|
||||||
| **Strategy (Sampling)** | `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations |
|
| **Strategy (Sampling)** | `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations |
|
||||||
@@ -1427,9 +1433,9 @@ classDiagram
|
|||||||
| **Strategy (Attention)** | `AttentionBackend`, `TorchNativeBackend`, `CudaBackend` | Attention computation backend switching via context manager |
|
| **Strategy (Attention)** | `AttentionBackend`, `TorchNativeBackend`, `CudaBackend` | Attention computation backend switching via context manager |
|
||||||
| **Auto-dispatch (Rotary)** | `apply_rotary_emb`, `rotary_backend.py`, `rotary_ops.py` | Rotary embedding CUDA kernel auto-dispatch with torch fallback |
|
| **Auto-dispatch (Rotary)** | `apply_rotary_emb`, `rotary_backend.py`, `rotary_ops.py` | Rotary embedding CUDA kernel auto-dispatch with torch fallback |
|
||||||
| **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor`, `FSDPExecutor` | Gradient accumulation & model distribution |
|
| **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor`, `FSDPExecutor` | Gradient accumulation & model distribution |
|
||||||
| **Storage** | `Store`, `H5Store`, `MmapStore`, `JsonlStore` | Format-agnostic data access with multi-segment support |
|
| **Storage** | `Store`, `MmapStore`, `JsonlStore` | Format-agnostic data access with multi-segment support |
|
||||||
| **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching |
|
| **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching |
|
||||||
| **AutoModel Registry** | `AutoModel`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading |
|
| **Model Registry** | `ModelFactory`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading |
|
||||||
|
|
||||||
## Core Relationships
|
## Core Relationships
|
||||||
|
|
||||||
@@ -1439,10 +1445,10 @@ classDiagram
|
|||||||
4. **Executor Selection**: `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)` → `NoneExecutor` / `DDPExecutor` / `FSDPExecutor`
|
4. **Executor Selection**: `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)` → `NoneExecutor` / `DDPExecutor` / `FSDPExecutor`
|
||||||
5. **Inference Flow**: `InferenceEngine` → `InferenceScheduler` → `AutoRegressiveLM`, backed by `PagePool` + `KVCache` + `SamplingPipeline`. Attention backend selected via `attn_backend()` context manager (`TorchNativeBackend` default, `CudaBackend` for CUDA kernels). Rotary embedding auto-dispatches to CUDA kernel when available (inference mode), else torch complex multiply (training).
|
5. **Inference Flow**: `InferenceEngine` → `InferenceScheduler` → `AutoRegressiveLM`, backed by `PagePool` + `KVCache` + `SamplingPipeline`. Attention backend selected via `attn_backend()` context manager (`TorchNativeBackend` default, `CudaBackend` for CUDA kernels). Rotary embedding auto-dispatches to CUDA kernel when available (inference mode), else torch complex multiply (training).
|
||||||
6. **Distributed**: `spawn_parallel_fn` + `setup_parallel` for multi-process DDP
|
6. **Distributed**: `spawn_parallel_fn` + `setup_parallel` for multi-process DDP
|
||||||
7. **Dataset Loading**: `DatasetFactory` creates datasets, `Store` (H5Store/MmapStore/JsonlStore) loads data with explicit `_length` and multi-segment `_data`
|
7. **Dataset Loading**: `DatasetFactory` creates datasets, `Store` (`MmapStore`/`JsonlStore`) loads data with explicit `_length` and multi-segment `_data`
|
||||||
8. **Checkpoint**: `Checkpoint` saves/loads safetensors + metadata (rank-0 only), extra state saved as `{key}.pt`
|
8. **Checkpoint**: `Checkpoint` saves/loads safetensors + metadata; `CheckpointCallback` performs rank-0 training saves, with extra state saved as `{key}.pt`
|
||||||
9. **Scheduler**: `SchedulerFactory` creates `CosineScheduler`/`SGDRScheduler`/`WSDScheduler`
|
9. **Scheduler**: `SchedulerFactory` creates `CosineScheduler`/`SGDRScheduler`/`WSDScheduler`
|
||||||
10. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops
|
10. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops
|
||||||
11. **Protocols**: `OptimizerProtocol` / `SchedulerProtocol` — structural subtyping for `AccumOptimizer` / `AccumScheduler` wrappers
|
11. **Protocols**: `OptimizerProtocol` / `SchedulerProtocol` — structural subtyping for `AccumOptimizer` / `AccumScheduler` wrappers
|
||||||
|
|
||||||
> Document Update Time: 2026-07-31
|
> Document Update Time: 2026-08-02
|
||||||
|
|||||||
@@ -117,23 +117,22 @@ Each `csrc/tests/*.cu` file has the `nvcc` compile command in its header comment
|
|||||||
```bash
|
```bash
|
||||||
nvcc -I csrc -arch=sm_89 -O3 --use_fast_math \
|
nvcc -I csrc -arch=sm_89 -O3 --use_fast_math \
|
||||||
--ptxas-options=-O3,-v --extra-device-vectorization \
|
--ptxas-options=-O3,-v --extra-device-vectorization \
|
||||||
csrc/tests/attn_decode_test.cu -o /tmp/test && /tmp/test
|
-Xcompiler -fopenmp csrc/tests/attn_test.cu -o /tmp/test && /tmp/test
|
||||||
```
|
```
|
||||||
|
|
||||||
Test files:
|
Test files:
|
||||||
- `attn_decode_test.cu` — basic decode kernel
|
- `attn_test.cu` — decode + prefill kernels (correctness tables + benchmarks)
|
||||||
- `attn_paged_decode_test.cu` — paged decode kernel
|
- `attn_paged_test.cu` — paged decode/prefill kernels
|
||||||
- `attn_prefill_test.cu` — prefill kernel
|
|
||||||
|
|
||||||
## Benchmarks
|
## Benchmarks
|
||||||
|
|
||||||
Hardware: NVIDIA L20 (sm_89, 46 GB), CUDA 12.8, driver 570.86.
|
Hardware: NVIDIA L20 (sm_89, 46 GB), CUDA 12.8, driver 570.86.
|
||||||
|
|
||||||
Reproduce:
|
Reproduce (decode + prefill in `attn_test.cu`, paged in `attn_paged_test.cu`):
|
||||||
```bash
|
```bash
|
||||||
nvcc -I csrc -arch=sm_89 -O3 --use_fast_math \
|
nvcc -I csrc -arch=sm_89 -O3 --use_fast_math \
|
||||||
--ptxas-options=-O3,-v --extra-device-vectorization \
|
--ptxas-options=-O3,-v --extra-device-vectorization \
|
||||||
csrc/tests/attn_<name>_test.cu -o /tmp/test && /tmp/test
|
-Xcompiler -fopenmp csrc/tests/attn_test.cu -o /tmp/test && /tmp/test
|
||||||
```
|
```
|
||||||
|
|
||||||
## Known Optimization Targets
|
## Known Optimization Targets
|
||||||
@@ -165,9 +164,8 @@ csrc/
|
|||||||
│ └── attn_warp_utils.cuh # Warp-level utilities
|
│ └── attn_warp_utils.cuh # Warp-level utilities
|
||||||
└── tests/
|
└── tests/
|
||||||
├── test_utils.cuh # Shared test utilities
|
├── test_utils.cuh # Shared test utilities
|
||||||
├── attn_decode_test.cu # Decode kernel test
|
├── attn_test.cu # Decode + prefill kernels
|
||||||
├── attn_paged_decode_test.cu # Paged decode test
|
└── attn_paged_test.cu # Paged decode/prefill kernels
|
||||||
└── attn_prefill_test.cu # Prefill kernel test
|
|
||||||
```
|
```
|
||||||
|
|
||||||
Compiled `.so` files are placed in `astrai/extension/lib/`, separate from Python source files.
|
Compiled `.so` files are placed in `astrai/extension/lib/`, separate from Python source files.
|
||||||
|
|||||||
+69
-35
@@ -14,26 +14,30 @@ This document describes the data pipeline: from raw text to model input tensors.
|
|||||||
## Overview
|
## Overview
|
||||||
|
|
||||||
```
|
```
|
||||||
JSONL Lines → Pipeline (mask builder) → Tokenized Tensors
|
JSON / JSONL Records → Pipeline (mask builder) → Tokenized Tensors
|
||||||
↓
|
↓
|
||||||
.h5 or .bin storage
|
.bin storage
|
||||||
↓
|
↓
|
||||||
Store.load()
|
Store.load()
|
||||||
↓
|
↓
|
||||||
Store.fetch(begin, end, keys)
|
Store.fetch(begin, end, keys)
|
||||||
↓
|
↓
|
||||||
BaseDataset.__getitem__(idx)
|
Dataset.__getitem__(idx)
|
||||||
↓
|
↓
|
||||||
Sampler → DataLoader → Training / Inference
|
RDSampler → DataLoader → Training
|
||||||
```
|
```
|
||||||
|
|
||||||
## Data Preparation
|
## Data Preparation
|
||||||
|
|
||||||
Raw text is tokenized via `AutoTokenizer.encode()` and saved as HDF5 (`.h5`) or binary (`.bin` + `meta.json`) files with keyed tensor groups.
|
The offline `Pipeline` accepts `.jsonl` records and `.json` files containing one
|
||||||
|
object or a list of objects. It tokenizes them and writes binary shards (`.bin`
|
||||||
|
plus `meta.json`) with keyed tensor groups. Binary is the only registered output
|
||||||
|
writer; the pipeline cannot emit JSONL.
|
||||||
|
|
||||||
### Tokenization
|
### Tokenization
|
||||||
|
|
||||||
The `Pipeline` reads JSONL lines, applies the mask builder (see [Preprocessing](../guides/preprocessing.md)), and produces flat token sequences:
|
The `Pipeline` reads JSON/JSONL records, applies the mask builder (see
|
||||||
|
[Preprocessing](../guides/preprocessing.md)), and produces token sequences:
|
||||||
|
|
||||||
```python
|
```python
|
||||||
# Per JSONL line: messages → chat template → token IDs + loss mask
|
# Per JSONL line: messages → chat template → token IDs + loss mask
|
||||||
@@ -42,84 +46,114 @@ loss_mask = [0, 0, 0, 1, 1, 1, 1, 1, 1] # 0=masked, 1=train
|
|||||||
# Stored as flat tensors, packed with other lines by packing strategy
|
# Stored as flat tensors, packed with other lines by packing strategy
|
||||||
```
|
```
|
||||||
|
|
||||||
The output `meta.json` records the storage format, key names, dtype, total token count, and tensor shapes for each shard.
|
For default single-output preprocessing, the stored keys are `sequence` and
|
||||||
|
`position_ids`, plus `loss_mask` when masking is required. Packing is supported
|
||||||
|
for single-output data with a `sequence` key. Shard flushing counts the primary
|
||||||
|
flat sequence for each record: `sequence` in single-output mode, otherwise the
|
||||||
|
first flat source output.
|
||||||
|
|
||||||
|
The exact shard `meta.json` schema is a top-level mapping from key to tensor
|
||||||
|
metadata. It does not contain a storage-format or total-token field:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"sequence": {"shape": [123456], "dtype": "int32"},
|
||||||
|
"loss_mask": {"shape": [123456], "dtype": "bool"},
|
||||||
|
"position_ids": {"shape": [123456], "dtype": "int32"}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Record-aware binary data may also include `"offsets": [0, ...]` inside a key's
|
||||||
|
metadata, but the preprocessing `BinWriter` currently does not write offsets.
|
||||||
|
|
||||||
### Format Detection
|
### Format Detection
|
||||||
|
|
||||||
`detect_format(load_path)` inspects the path:
|
`detect_format(load_path)` inspects the path:
|
||||||
|
|
||||||
- If `load_path` is a file: checks suffix — `.h5`/`.hdf5` → `"h5"`, `.jsonl` → `"jsonl"`, unknown suffix raises `ValueError`
|
- If `load_path` is a file: `.jsonl` selects `"jsonl"`; other suffixes raise `ValueError`.
|
||||||
- If `load_path` is a directory: recursively globs for `*.h5`/`*.hdf5` files → `"h5"`, `*.bin` + `**/meta.json` → `"bin"`, or `*.jsonl` + `dataset_config.json` → `"jsonl"`
|
- If `load_path` is a directory: any recursive `*.bin` plus a `meta.json` selects `"bin"`; otherwise any recursive `*.jsonl` selects `"jsonl"`.
|
||||||
|
- Detection does not require `dataset_config.json`; configuration is selected later when `JsonlStore.load()` chooses a transform.
|
||||||
|
|
||||||
### Store Backends
|
### Store Backends
|
||||||
|
|
||||||
Storage format is auto-detected by `detect_format()`; backends are dispatched via registry:
|
Storage format is auto-detected by `detect_format()`; backends are dispatched via registry:
|
||||||
|
|
||||||
```
|
```
|
||||||
StoreFactory.create("h5") → H5Store
|
|
||||||
StoreFactory.create("bin") → MmapStore
|
StoreFactory.create("bin") → MmapStore
|
||||||
StoreFactory.create("jsonl") → JsonlStore
|
StoreFactory.create("jsonl") → JsonlStore
|
||||||
```
|
```
|
||||||
|
|
||||||
All three inherit `Store` (base, owns `_data`/`_cum`/`_offsets`/`_normalize`) plus the `Streamable` and `Recordable` mixins, so every backend supports both `fetch(begin, end, keys)` (stream) and `fetch_record(index, keys)` (record) APIs.
|
Both stores inherit `Store` and compose the `Streamable` and `Recordable`
|
||||||
|
access methods.
|
||||||
**H5Store**: Reads HDF5 files. Tensors are loaded into host memory and normalized into segmented storage. `segments_are_records=True` — each `data_i` dataset is one record.
|
|
||||||
|
|
||||||
**MmapStore**: Memory-maps `.bin` files. OS page cache sharing is native — no explicit `share_memory_()` needed. Uses `torch.from_numpy(np.memmap(...))`. `segments_are_records=False` — bin segments are contiguous streams; record access is driven by `_offsets` (written when `save_bin(..., record_keys=...)` was used at preprocessing time).
|
**MmapStore**: Memory-maps `.bin` files. OS page cache sharing is native — no explicit `share_memory_()` needed. Uses `torch.from_numpy(np.memmap(...))`. `segments_are_records=False` — bin segments are contiguous streams; record access is driven by `_offsets` (written when `save_bin(..., record_keys=...)` was used at preprocessing time).
|
||||||
|
|
||||||
**JsonlStore**: On-the-fly tokenization of raw JSONL files at load time. Requires a `dataset_config.json` alongside the `.jsonl` files following the same `PipelineConfig` schema with an additional `tokenizer_path` field. Two modes: eager (default, applies `TokenizeTransform` to all records at load) and lazy (`processor=fn` given, defers tokenisation to `fetch_record` — used by DPO/GRPO).
|
**JsonlStore**: Reads a `.jsonl` file or the sorted top-level `*.jsonl` files in
|
||||||
|
a directory. Eager transform selection uses the first available route:
|
||||||
|
|
||||||
All backends normalise tensors into `Store._data[Dict[str, List[Tensor]]]` + `Store._cum[Dict[str, List[int]]]` (cumulative lengths for bisect-based stream indexing) + `Store._offsets[Dict[str, List[int]]]` (per-record offsets for record-mode indexing). Nested keys (GRPO `responses`/`masks` as `List[List[Tensor]]`) are stored as-is and excluded from both bookkeepings — they are only accessed record-by-record.
|
1. An explicit `transform=` argument.
|
||||||
|
2. `dataset_config.json` in the JSONL directory. It follows `PipelineConfig` and may add `tokenizer_path`; when omitted, the config directory is used.
|
||||||
|
3. The built-in `messages` transform when `tokenizer_path=` is supplied. It masks system/user turns, trains assistant turns, and emits document-reset position IDs.
|
||||||
|
|
||||||
|
Only DPO gets an automatic lazy route from `DatasetFactory`: raw JSONL plus
|
||||||
|
`tokenizer_path` installs `dpo_processor` and tokenizes each record in
|
||||||
|
`fetch_record`. GRPO does not currently have an automatic lazy processor.
|
||||||
|
|
||||||
|
Eager-loaded stores normalize tensors into `Store._data[Dict[str, List[Tensor]]]` + `Store._cum[Dict[str, List[int]]]` (cumulative lengths for stream indexing) + `Store._offsets[Dict[str, List[int]]]` (per-record offsets for record indexing). Nested JSONL keys such as GRPO `responses`/`masks` are kept as record values and excluded from stream bookkeeping. Lazy DPO instead retains raw records and processes them in `fetch_record`.
|
||||||
|
|
||||||
## Data Keys by Training Type
|
## Data Keys by Training Type
|
||||||
|
|
||||||
| Type | Storage Keys | Access Mode |
|
| Type | Storage Keys | Access Mode |
|
||||||
|------|-------------|-------------|
|
|------|-------------|-------------|
|
||||||
| `seq` | `sequence` (→ input_ids, target_ids via offset-by-1) | stream (`fetch`) |
|
| `seq` | `sequence`, `position_ids` by default (`SEQDataset` consumes only `sequence`) | stream (`fetch`) |
|
||||||
| `sft` | `sequence`, `loss_mask`, `position_ids` | stream (`fetch`) |
|
| `sft` | `sequence`, `loss_mask`, `position_ids` | stream (`fetch`) |
|
||||||
| `dpo` | `chosen`, `rejected`, `chosen_mask`, `rejected_mask` | record (`fetch_record`) |
|
| `dpo` | `chosen`, `rejected`, `chosen_mask`, `rejected_mask` | record (`fetch_record`) |
|
||||||
| `grpo` | `prompts`, `responses`, `masks`, `rewards` | record (`fetch_record`) |
|
| `grpo` | `prompts`, `responses`, `masks`, `rewards` | record (`fetch_record`) |
|
||||||
|
|
||||||
|
Offline `.bin` output from DPO/GRPO preprocessing is not currently loadable for
|
||||||
|
training. DPO shards are written without record offsets, while GRPO response
|
||||||
|
groups are flattened without preserving record/group boundaries. Supported raw
|
||||||
|
routes are eager JSONL for SEQ/SFT and automatic lazy JSONL for DPO. GRPO
|
||||||
|
requires a caller-built, already-loaded record store.
|
||||||
|
|
||||||
## Dataset Architecture
|
## Dataset Architecture
|
||||||
|
|
||||||
```
|
```
|
||||||
DatasetFactory.load(
|
DatasetFactory.load(...)
|
||||||
train_type, load_path=None, window_size=0, stride=None,
|
|
||||||
storage_type=None, tokenizer_path=None,
|
|
||||||
max_len=2048, store=None
|
|
||||||
)
|
|
||||||
→ BaseDataset.load(load_path, storage_type=None)
|
|
||||||
→ detect_format(load_path)
|
→ detect_format(load_path)
|
||||||
→ StoreFactory.create(storage_type)
|
→ optionally build dpo_processor for raw JSONL
|
||||||
→ Store.load(load_path)
|
→ StoreFactory.create(storage_type, window_size, stride)
|
||||||
→ _normalize(raw) # base Store, shared by both backends
|
→ Store.load(load_path, transform=... or processor=...)
|
||||||
→ Store._data[Dict[str, List[Tensor]]]
|
→ DatasetFactory.create(train_type, store=store)
|
||||||
+ _cum[Dict[str, List[int]]] (stream mode)
|
|
||||||
+ _offsets[Dict[str, List[int]]] (record mode)
|
|
||||||
|
|
||||||
Stream datasets (SEQ/SFT):
|
Stream datasets (SEQ/SFT):
|
||||||
BaseDataset.__getitem__(idx)
|
BaseDataset.__getitem__(idx)
|
||||||
→ get_index(idx) → [begin, end)
|
→ Store.sample_window(idx) → [begin, end)
|
||||||
→ Store.fetch(begin, end, keys) → Tensor / Dict[str, Tensor]
|
→ Store.fetch(begin, end, keys) → Tensor / Dict[str, Tensor]
|
||||||
|
|
||||||
Record datasets (DPO/GRPO via RecordDataset):
|
Record datasets (DPO/GRPO):
|
||||||
RecordDataset.__getitem__(idx)
|
DPODataset/GRPODataset.__getitem__(idx)
|
||||||
→ Store.fetch_record(idx, keys) → Tensor / Dict[str, Tensor]
|
→ Store.fetch_record(idx, keys) → Tensor / Dict[str, Tensor]
|
||||||
```
|
```
|
||||||
|
|
||||||
Class hierarchy: `BaseDataset` ← `SEQDataset` / `SFTDataset` (stream); `BaseDataset` ← `RecordDataset` ← `DPODataset` / `GRPODataset` (record).
|
Class hierarchy: `BaseDataset` is the direct base of `SEQDataset`, `SFTDataset`,
|
||||||
|
`DPODataset`, and `GRPODataset`. There is no `RecordDataset` class.
|
||||||
|
|
||||||
`window_size` = max input length, `stride` = step between consecutive samples (defaults to `window_size`, optional). Only meaningful for stream datasets — record datasets ignore both. `storage_type` defaults to `None` (auto-detect via `detect_format`).
|
`window_size` = max input length, `stride` = step between consecutive samples (defaults to `window_size`, optional). Only meaningful for stream datasets — record datasets ignore both. `storage_type` defaults to `None` (auto-detect via `detect_format`).
|
||||||
|
|
||||||
`tokenizer_path` triggers lazy on-the-fly tokenisation for record datasets on raw JSONL (DPO builds a `dpo_processor`; SEQ/SFT/pre-tokenised backends ignore it). `store` (pre-built `Store`) bypasses `load_path`/`storage_type`/`tokenizer_path` entirely — the caller controls Store construction.
|
For raw JSONL, `tokenizer_path` builds the lazy processor only for DPO. For
|
||||||
|
SEQ/SFT it is forwarded to `JsonlStore` so the built-in eager `messages`
|
||||||
|
transform can be selected when no `dataset_config.json` exists. GRPO receives no
|
||||||
|
automatic processor. A pre-built `store` bypasses path, format, tokenizer,
|
||||||
|
window, and stride setup entirely.
|
||||||
|
|
||||||
`Store.fetch(begin, end, keys)` (stream mode, on `Streamable`): accepts a single key (`str`) returning a `Tensor`, or a list of keys returning `Dict[str, Tensor]`. Internally uses `bisect` across multi-segment tensors. Raises `RuntimeError("Store not loaded")` if called before `load()`.
|
`Store.fetch(begin, end, keys)` (stream mode, on `Streamable`): accepts a single key (`str`) returning a `Tensor`, or a list of keys returning `Dict[str, Tensor]`. Internally uses `bisect` across multi-segment tensors. Raises `RuntimeError("Store not loaded")` if called before `load()`.
|
||||||
|
|
||||||
`Store.fetch_record(index, keys)` (record mode, on `Recordable`): same key API. Uses `_offsets[key]` when present (bin layout with per-record offsets), otherwise indexes `_data[key]` directly (H5/JSONL where each segment is one record).
|
`Store.fetch_record(index, keys)` (record mode, on `Recordable`): same key API. Uses `_offsets[key]` when present for binary record layouts; otherwise it indexes per-record JSONL tensors directly.
|
||||||
|
|
||||||
## Sampler
|
## Sampler
|
||||||
|
|
||||||
`ResumableDistributedSampler` supports checkpoint-aware distributed sampling:
|
`RDSampler` supports checkpoint-aware distributed sampling:
|
||||||
|
|
||||||
- Tracks `start_epoch` / `start_iter` for resume
|
- Tracks `start_epoch` / `start_iter` for resume
|
||||||
- Shuffle via `torch.Generator(seed + epoch)`
|
- Shuffle via `torch.Generator(seed + epoch)`
|
||||||
|
|||||||
+30
-13
@@ -41,7 +41,12 @@ 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 $$
|
$$ 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 `cos_table` and `sin_table` (f32, `[max_len, dim/2]`). `forward()` returns a `(cos, sin)` tuple 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 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.
|
||||||
|
|
||||||
**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.
|
**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.
|
||||||
|
|
||||||
@@ -51,13 +56,13 @@ $$ 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_{
|
|||||||
|
|
||||||
Next-token cross-entropy with optional label smoothing:
|
Next-token cross-entropy with optional label smoothing:
|
||||||
|
|
||||||
$$ L_{\text{PT}} = -\sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta) $$
|
$$ L_{\text{PT}} = -\frac{1}{T}\sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta) $$
|
||||||
|
|
||||||
### SFT (Supervised Fine-Tuning)
|
### SFT (Supervised Fine-Tuning)
|
||||||
|
|
||||||
Masked cross-entropy (`ignore_index=-100`) over response tokens only:
|
Masked cross-entropy (`ignore_index=-100`) over response tokens only:
|
||||||
|
|
||||||
$$ L_{\text{SFT}} = -\sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta) $$
|
$$ L_{\text{SFT}} = -\frac{1}{L}\sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta) $$
|
||||||
|
|
||||||
Prompt tokens are masked out via `loss_mask`; only response tokens contribute to the loss.
|
Prompt tokens are masked out via `loss_mask`; only response tokens contribute to the loss.
|
||||||
|
|
||||||
@@ -81,6 +86,14 @@ Where $\rho_t = \pi_\theta(a_t|s_t) / \pi_{\text{old}}(a_t|s_t)$ is the per-toke
|
|||||||
|
|
||||||
Parameters: `group_size=4`, `clip_eps=0.2`, `kl_coef=0.01`.
|
Parameters: `group_size=4`, `clip_eps=0.2`, `kl_coef=0.01`.
|
||||||
|
|
||||||
|
### MoE Load Balancing
|
||||||
|
|
||||||
|
MoE layers add a differentiable load-balancing term based on mean router probabilities and top-k expert assignment frequency. The training objective is:
|
||||||
|
|
||||||
|
$$ L = L_{\text{task}} + \lambda_{\text{MoE}} L_{\text{aux}} $$
|
||||||
|
|
||||||
|
`TrainConfig.moe_aux_loss_coef` controls $\lambda_{\text{MoE}}$ (default `0.01`). The unweighted and weighted auxiliary losses are logged separately.
|
||||||
|
|
||||||
## Training Loop Internals
|
## Training Loop Internals
|
||||||
|
|
||||||
Two-level loop: **epoch** → **batch**. Optimizer step fires every `grad_accum_steps` batches.
|
Two-level loop: **epoch** → **batch**. Optimizer step fires every `grad_accum_steps` batches.
|
||||||
@@ -90,11 +103,12 @@ on_train_begin
|
|||||||
model.train()
|
model.train()
|
||||||
on_epoch_begin
|
on_epoch_begin
|
||||||
for batch in dataloader:
|
for batch in dataloader:
|
||||||
on_batch_begin
|
|
||||||
with executor.accumulate(model):
|
with executor.accumulate(model):
|
||||||
loss = strategy.compute_loss(batch)
|
on_batch_begin
|
||||||
context.loss = loss.item()
|
loss_output = strategy(batch)
|
||||||
stand_loss = loss / executor.grad_accum_steps
|
context.loss = loss_output["loss"].item()
|
||||||
|
context.metrics = loss_output["metrics"]
|
||||||
|
stand_loss = loss_output["loss"] / executor.grad_accum_steps
|
||||||
executor.backward(stand_loss)
|
executor.backward(stand_loss)
|
||||||
context.consumed_samples += (
|
context.consumed_samples += (
|
||||||
context.config.batch_per_device * context.world_size
|
context.config.batch_per_device * context.world_size
|
||||||
@@ -104,6 +118,7 @@ on_train_begin
|
|||||||
if executor.sync_gradients:
|
if executor.sync_gradients:
|
||||||
on_optimizer_step
|
on_optimizer_step
|
||||||
optimizer.step()
|
optimizer.step()
|
||||||
|
strategy.on_optimizer_step()
|
||||||
optimizer.zero_grad()
|
optimizer.zero_grad()
|
||||||
if scheduler:
|
if scheduler:
|
||||||
scheduler.step()
|
scheduler.step()
|
||||||
@@ -112,21 +127,23 @@ on_train_end
|
|||||||
```
|
```
|
||||||
|
|
||||||
The loss is divided by `grad_accum_steps` before `backward()`, so accumulated gradients sum to the correct mean.
|
The loss is divided by `grad_accum_steps` before `backward()`, so accumulated gradients sum to the correct mean.
|
||||||
|
Strategy metrics are detached and converted to Python `float` values before the
|
||||||
|
`LossOutput` is returned; only `LossOutput.loss` remains a differentiable tensor.
|
||||||
|
|
||||||
## Callback Lifecycle
|
## Callback Lifecycle
|
||||||
|
|
||||||
| Hook | Fires | Default callback |
|
| Hook | Fires | Default callback |
|
||||||
|------|-------|-----------------|
|
|------|-------|-----------------|
|
||||||
| `on_train_begin` | Before training starts | `GradientCheckpointingCallback` |
|
| `on_train_begin` | Before training starts | `GradientCheckpointingCallback`, `CheckpointCallback`, `MetricCallback` |
|
||||||
| `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` |
|
| `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` |
|
||||||
| `on_batch_begin` | Every batch | — |
|
| `on_batch_begin` | Every batch | — |
|
||||||
| `on_optimizer_step` | Every accumulation window | `GradientClippingCallback`, `MetricCallback`, `ProgressBarCallback` |
|
| `on_optimizer_step` | Every accumulation window | `MetricCallback`, `ProgressBarCallback`, `GradientClippingCallback` |
|
||||||
| `on_batch_end` | Every batch | `CheckpointCallback` |
|
| `on_batch_end` | Every batch | `CheckpointCallback` |
|
||||||
| `on_epoch_end` | End of each epoch | `MetricCallback`, `ProgressBarCallback` |
|
| `on_epoch_end` | End of each epoch | `MetricCallback`, `ProgressBarCallback` |
|
||||||
| `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` |
|
| `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` |
|
||||||
| `on_train_end` | Training ends (always via finally) | `CheckpointCallback`, `MetricCallback`, `GradientCheckpointingCallback` |
|
| `on_train_end` | Training exits after `on_train_begin` completes (via `finally`) | `GradientCheckpointingCallback`, `CheckpointCallback`, `MetricCallback` |
|
||||||
|
|
||||||
Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric` (JSONL + validation, rank-0), `progress_bar` (tqdm), `gradient_clipping` (always registered; computes grad norm, clips only when `max_grad_norm` is not `None`).
|
Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric` (JSONL + validation, rank-0), `progress_bar` (tqdm, rank-0), `gradient_clipping`. The gradient-clipping callback is always registered and always calls `executor.clip_grad_norm()` with the numeric `max_grad_norm` value.
|
||||||
|
|
||||||
## KV Cache Mathematics
|
## KV Cache Mathematics
|
||||||
|
|
||||||
@@ -151,7 +168,7 @@ Three-layer separation (SGLang-inspired):
|
|||||||
- **ReqToTokenPool**: Index table `[req_idx, pos] → physical token slot`, shared across all layers.
|
- **ReqToTokenPool**: Index table `[req_idx, pos] → physical token slot`, shared across all layers.
|
||||||
- **Allocator + PrefixCache**: Paged-mode slot allocation with ref-counting, LRU eviction, and hash-based prefix sharing.
|
- **Allocator + PrefixCache**: Paged-mode slot allocation with ref-counting, LRU eviction, and hash-based prefix sharing.
|
||||||
|
|
||||||
`PagePool` orchestrates all three. In contiguous mode (default), `req_to_token` is a trivial linear mapping. In paged mode, slots are allocated on demand with prefix caching support. `bind_tasks()` returns a `KVCache` dataclass with precomputed `page_table` and `decode_mask` fields (computed once per decode step, shared across all layers). Attention layers access buffers directly — no methods, no abstraction.
|
`PagePool` orchestrates all three. In contiguous mode (default), `req_to_token` is a trivial linear mapping. In paged mode, slots are allocated on demand with prefix caching support. `bind_tasks()` returns a `KVCache` dataclass with `kv_indptr`, a prefix-sum index over sequence lengths computed once per step and shared across layers. Attention layers access buffers directly — no methods, no abstraction.
|
||||||
|
|
||||||
### Attention Backend
|
### Attention Backend
|
||||||
|
|
||||||
@@ -230,4 +247,4 @@ total_steps = (batches_per_replica // grad_accum_steps) * n_epoch
|
|||||||
|
|
||||||
This accounts for data-parallel sharding — each rank processes `1/nprocs` of the dataset.
|
This accounts for data-parallel sharding — each rank processes `1/nprocs` of the dataset.
|
||||||
|
|
||||||
> Document Update Time: 2026-07-31
|
> Document Update Time: 2026-08-02
|
||||||
|
|||||||
+22
-3
@@ -2,11 +2,23 @@
|
|||||||
|
|
||||||
This guide walks you through installing AstrAI, downloading a model, running inference, preprocessing data, and launching your first training job.
|
This guide walks you through installing AstrAI, downloading a model, running inference, preprocessing data, and launching your first training job.
|
||||||
|
|
||||||
|
## Contents
|
||||||
|
|
||||||
|
- [Prerequisites](#prerequisites)
|
||||||
|
- [1. Install](#1-install)
|
||||||
|
- [2. Download Model Weights](#2-download-model-weights)
|
||||||
|
- [3. Run Inference](#3-run-inference)
|
||||||
|
- [4. Preprocess Data](#4-preprocess-data)
|
||||||
|
- [5. Train](#5-train)
|
||||||
|
- [6. Evaluate](#6-evaluate)
|
||||||
|
- [7. Docker](#7-docker)
|
||||||
|
- [Next Steps](#next-steps)
|
||||||
|
|
||||||
## Prerequisites
|
## Prerequisites
|
||||||
|
|
||||||
- **Python 3.12+**
|
- **Python 3.12+**
|
||||||
- **PyTorch 2.11+** (CUDA 12.8 recommended for GPU support)
|
- **PyTorch 2.11.0** (the exact version pinned by AstrAI; CUDA 12.8 build recommended for GPU support)
|
||||||
- NVIDIA GPU with CUDA (optional but recommended; CPU works for inference)
|
- NVIDIA GPU with CUDA for training, `scripts/tools/generate.py`, generation evaluations, and demos. The HTTP server and direct-scoring evaluations can run on CPU where their CLI exposes a CPU device.
|
||||||
|
|
||||||
## 1. Install
|
## 1. Install
|
||||||
|
|
||||||
@@ -55,7 +67,7 @@ python scripts/demo/stream_chat.py
|
|||||||
# Type your message after >>, type !exit to quit
|
# Type your message after >>, type !exit to quit
|
||||||
```
|
```
|
||||||
|
|
||||||
This starts a multi-turn interactive chat session with streaming output.
|
This starts a single-turn interactive prompt loop with streaming output. Each prompt is independent; conversation history is not retained.
|
||||||
|
|
||||||
### Start an HTTP Server
|
### Start an HTTP Server
|
||||||
|
|
||||||
@@ -192,6 +204,13 @@ See [Training Guide](guides/training.md) for loss formulas and strategies. See [
|
|||||||
|
|
||||||
## 6. Evaluate
|
## 6. Evaluate
|
||||||
|
|
||||||
|
HumanEval and MMLU download their benchmark data through HuggingFace
|
||||||
|
`datasets`, which is not part of the base install:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
pip install datasets
|
||||||
|
```
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
# HumanEval (code generation, auto-downloads dataset)
|
# HumanEval (code generation, auto-downloads dataset)
|
||||||
python scripts/eval/evaluate_humaneval.py --param_path ./params --num_samples 20
|
python scripts/eval/evaluate_humaneval.py --param_path ./params --num_samples 20
|
||||||
|
|||||||
+25
-17
@@ -2,6 +2,18 @@
|
|||||||
|
|
||||||
AstrAI supports three parallel modes: **single GPU** (`none`), **Data Parallel** (`ddp`), and **Fully Sharded Data Parallel** (`fsdp`). This guide covers when to use each, how to launch multi-GPU training, and how gradient accumulation works.
|
AstrAI supports three parallel modes: **single GPU** (`none`), **Data Parallel** (`ddp`), and **Fully Sharded Data Parallel** (`fsdp`). This guide covers when to use each, how to launch multi-GPU training, and how gradient accumulation works.
|
||||||
|
|
||||||
|
## Contents
|
||||||
|
|
||||||
|
- [Quick Start](#quick-start)
|
||||||
|
- [Parallel Modes](#parallel-modes)
|
||||||
|
- [Gradient Accumulation](#gradient-accumulation)
|
||||||
|
- [Process Launching](#process-launching)
|
||||||
|
- [NCCL Troubleshooting](#nccl-troubleshooting)
|
||||||
|
- [Checkpoint Saving](#checkpoint-saving)
|
||||||
|
- [Total Steps Calculation](#total-steps-calculation)
|
||||||
|
- [Real Examples](#real-examples)
|
||||||
|
- [CLI Parameters](#cli-parameters)
|
||||||
|
|
||||||
## Quick Start
|
## Quick Start
|
||||||
|
|
||||||
### Single GPU
|
### Single GPU
|
||||||
@@ -21,9 +33,6 @@ python scripts/tools/train.py \
|
|||||||
|
|
||||||
```bash
|
```bash
|
||||||
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||||
export NCCL_P2P_DISABLE=1
|
|
||||||
export NCCL_NET_GDR_LEVEL=0
|
|
||||||
|
|
||||||
python scripts/tools/train.py \
|
python scripts/tools/train.py \
|
||||||
--train_type=sft \
|
--train_type=sft \
|
||||||
--param_path ./params \
|
--param_path ./params \
|
||||||
@@ -38,9 +47,6 @@ python scripts/tools/train.py \
|
|||||||
|
|
||||||
```bash
|
```bash
|
||||||
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||||
export NCCL_P2P_DISABLE=1
|
|
||||||
export NCCL_NET_GDR_LEVEL=0
|
|
||||||
|
|
||||||
python scripts/tools/train.py \
|
python scripts/tools/train.py \
|
||||||
--train_type=sft \
|
--train_type=sft \
|
||||||
--param_path ./params \
|
--param_path ./params \
|
||||||
@@ -110,7 +116,7 @@ AstrAI auto-detects the launch method:
|
|||||||
|
|
||||||
| Detection | Strategy | Use Case |
|
| Detection | Strategy | Use Case |
|
||||||
|-----------|----------|----------|
|
|-----------|----------|----------|
|
||||||
| `torchelastic` / `torchrun` env vars | `TorchrunStrategy` | External orchestrator (torchrun, SLURM, K8s) |
|
| `torchelastic` / `torchrun` env vars | `TorchrunStrategy` | External orchestrator (`torchrun`, K8s) |
|
||||||
| `RANK` + `WORLD_SIZE` env vars | `TorchrunStrategy` | External launch |
|
| `RANK` + `WORLD_SIZE` env vars | `TorchrunStrategy` | External launch |
|
||||||
| Neither | `LocalStrategy` | `python scripts/tools/train.py` (in-process spawn) |
|
| Neither | `LocalStrategy` | `python scripts/tools/train.py` (in-process spawn) |
|
||||||
|
|
||||||
@@ -126,23 +132,28 @@ For multi-node or SLURM environments:
|
|||||||
torchrun --nproc_per_node=4 scripts/tools/train.py \
|
torchrun --nproc_per_node=4 scripts/tools/train.py \
|
||||||
--train_type=sft \
|
--train_type=sft \
|
||||||
--parallel_mode=ddp \
|
--parallel_mode=ddp \
|
||||||
|
--nprocs=4 \
|
||||||
--param_path ./params \
|
--param_path ./params \
|
||||||
--data_root_path ./dataset \
|
--data_root_path ./dataset \
|
||||||
--batch_per_device=4
|
--batch_per_device=4
|
||||||
```
|
```
|
||||||
|
|
||||||
When launched via torchrun, AstrAI reads `RANK`, `WORLD_SIZE`, `LOCAL_RANK` from the environment and uses `TorchrunStrategy`. The `--nprocs` flag is ignored (the orchestrator controls process count).
|
When launched via `torchrun`, the launcher creates the worker processes. AstrAI reads `RANK`, `WORLD_SIZE`, and `LOCAL_RANK` from the environment and uses `TorchrunStrategy`; `--nprocs` does not control process creation in this mode.
|
||||||
|
|
||||||
## NCCL Environment Variables
|
The current training CLI still uses `--nprocs` when calculating scheduler `total_steps`. Set it to the global `WORLD_SIZE` so the step count reflects data-parallel sharding, including multi-node runs.
|
||||||
|
|
||||||
For multi-GPU training, you **must** set these environment variables:
|
Raw Slurm variables such as `SLURM_PROCID`, `SLURM_NTASKS`, and `SLURM_LOCALID` are not recognized automatically. Launch through `torchrun`, or map the scheduler's variables to `RANK`, `WORLD_SIZE`, `LOCAL_RANK`, `MASTER_ADDR`, and `MASTER_PORT` before starting AstrAI. The same requirement applies to launchers that expose only OpenMPI-specific variables.
|
||||||
|
|
||||||
|
## NCCL Troubleshooting
|
||||||
|
|
||||||
|
The following variables are troubleshooting options for hardware or network configurations where NCCL hangs or fails. They are not general requirements and can reduce performance by disabling peer-to-peer or GPUDirect RDMA paths:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
export NCCL_P2P_DISABLE=1
|
export NCCL_P2P_DISABLE=1
|
||||||
export NCCL_NET_GDR_LEVEL=0
|
export NCCL_NET_GDR_LEVEL=0
|
||||||
```
|
```
|
||||||
|
|
||||||
These are required on certain hardware configurations (see `AGENTS.md`). Without them, NCCL may hang or crash during collective operations. These are set in the training shell scripts (`train-seq.sh`, `train-sft.sh`, `train-dpo.sh`) but not in Python code — you must export them before launching.
|
Apply them only after confirming the relevant NCCL transport is the source of the failure. AstrAI does not set them in Python.
|
||||||
|
|
||||||
## Checkpoint Saving
|
## Checkpoint Saving
|
||||||
|
|
||||||
@@ -176,9 +187,6 @@ This ensures the LR schedule is correctly scaled regardless of the number of GPU
|
|||||||
|
|
||||||
```bash
|
```bash
|
||||||
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||||
export NCCL_P2P_DISABLE=1
|
|
||||||
export NCCL_NET_GDR_LEVEL=0
|
|
||||||
|
|
||||||
python scripts/tools/train.py \
|
python scripts/tools/train.py \
|
||||||
--train_type=seq \
|
--train_type=seq \
|
||||||
--param_path ./params \
|
--param_path ./params \
|
||||||
@@ -240,7 +248,7 @@ python scripts/tools/train.py \
|
|||||||
|
|
||||||
| Parameter | Default | Description |
|
| Parameter | Default | Description |
|
||||||
|-----------|---------|-------------|
|
|-----------|---------|-------------|
|
||||||
| `--nprocs` | 1 | Number of GPUs / processes |
|
| `--nprocs` | 1 | Local process count for AstrAI's launcher; under `torchrun`, set it to global `WORLD_SIZE` for total-step calculation |
|
||||||
| `--parallel_mode` | `fsdp` | `none`, `ddp`, or `fsdp` |
|
| `--parallel_mode` | `fsdp` | `none`, `ddp`, or `fsdp` |
|
||||||
| `--start_method` | `spawn` | Multiprocessing start method (`spawn`, `fork`, `forkserver`) |
|
| `--start_method` | `spawn` | Multiprocessing start method (`spawn`, `fork`, `forkserver`) |
|
||||||
| `--backend` | `nccl` | Distributed backend (`nccl`, `gloo`) |
|
| `--backend` | `nccl` | Distributed backend (`nccl`, `gloo`) |
|
||||||
@@ -248,8 +256,8 @@ python scripts/tools/train.py \
|
|||||||
| `--master_port` | `29500` | Master node port |
|
| `--master_port` | `29500` | Master node port |
|
||||||
| `--device_type` | `cuda` | Device type |
|
| `--device_type` | `cuda` | Device type |
|
||||||
|
|
||||||
> `--tp_size` is parsed but **not yet wired** — tensor parallelism is future work. `ColumnParallelLinear` / `RowParallelLinear` exist in `astrai/parallel/module.py` but are not used by the model.
|
> `--tp_size` is accepted by the CLI but discarded before configuration. Tensor parallelism is not implemented, and there is no tensor-parallel module or model integration.
|
||||||
|
|
||||||
Full parameter reference: [CLI Reference](params.md). Training loop and strategies: [Training Guide](training.md).
|
Full parameter reference: [CLI Reference](params.md). Training loop and strategies: [Training Guide](training.md).
|
||||||
|
|
||||||
> Document Update Time: 2026-07-30
|
> Document Update Time: 2026-08-02
|
||||||
|
|||||||
+46
-12
@@ -2,6 +2,29 @@
|
|||||||
|
|
||||||
AstrAI provides 7 evaluation scripts in `scripts/eval/` covering code generation, knowledge QA, perplexity, summarization, data quality, instruction following, and weight analysis.
|
AstrAI provides 7 evaluation scripts in `scripts/eval/` covering code generation, knowledge QA, perplexity, summarization, data quality, instruction following, and weight analysis.
|
||||||
|
|
||||||
|
## Contents
|
||||||
|
|
||||||
|
- [Prerequisites](#prerequisites)
|
||||||
|
- [Overview](#overview)
|
||||||
|
- [HumanEval](#humaneval-code-generation)
|
||||||
|
- [MMLU](#mmlu-knowledge-qa)
|
||||||
|
- [Perplexity](#perplexity-ppl)
|
||||||
|
- [ROUGE](#rouge)
|
||||||
|
- [IFD](#ifd-instruction-following-difficulty)
|
||||||
|
- [IFEval](#ifeval-instruction-following)
|
||||||
|
- [Weight Analysis](#weight-analysis)
|
||||||
|
- [Tips](#tips)
|
||||||
|
|
||||||
|
## Prerequisites
|
||||||
|
|
||||||
|
HumanEval, MMLU, and IFEval import HuggingFace `datasets` to download their benchmark data. This package is not installed by AstrAI's base dependencies, so install it before running those scripts:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
pip install datasets
|
||||||
|
```
|
||||||
|
|
||||||
|
The generation-based scripts require CUDA because they load the model on `cuda` with `bfloat16`. Direct-scoring and metric scripts support the devices shown below.
|
||||||
|
|
||||||
## Overview
|
## Overview
|
||||||
|
|
||||||
| Script | Metric | Model Invocation | External Dataset |
|
| Script | Metric | Model Invocation | External Dataset |
|
||||||
@@ -18,7 +41,15 @@ Two invocation patterns exist:
|
|||||||
- **Generation benchmarks** (HumanEval, IFEval): use `InferenceEngine` to generate responses, then score them.
|
- **Generation benchmarks** (HumanEval, IFEval): use `InferenceEngine` to generate responses, then score them.
|
||||||
- **Scoring benchmarks** (MMLU, PPL, IFD): call `model()` directly under `torch.inference_mode()` for log-likelihood computation.
|
- **Scoring benchmarks** (MMLU, PPL, IFD): call `model()` directly under `torch.inference_mode()` for log-likelihood computation.
|
||||||
|
|
||||||
Common defaults: `--param_path` defaults to `./params`; dtype defaults to `bfloat16` on CUDA, `float32` on CPU.
|
| Script | Device support |
|
||||||
|
|--------|----------------|
|
||||||
|
| HumanEval | CUDA for generation; `--test_only` can score existing completions without loading a model |
|
||||||
|
| IFEval | CUDA only |
|
||||||
|
| MMLU | CUDA or CPU via `--device`; auto-selects CUDA when available |
|
||||||
|
| PPL | CUDA or CPU via `--device`; auto-selects CUDA when available |
|
||||||
|
| IFD | CUDA or CPU via `--device`; auto-selects CUDA when available |
|
||||||
|
| ROUGE | CPU-only metric computation; no model is loaded |
|
||||||
|
| Weight analysis | CUDA by default; CPU supported via `--device cpu` |
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
@@ -30,7 +61,7 @@ Generates completions for 164 programming problems, executes them against hidden
|
|||||||
python scripts/eval/evaluate_humaneval.py \
|
python scripts/eval/evaluate_humaneval.py \
|
||||||
--param_path ./params \
|
--param_path ./params \
|
||||||
--num_samples 20 \
|
--num_samples 20 \
|
||||||
--batch_size 32 \
|
--batch_size 64 \
|
||||||
--max_tokens 512 \
|
--max_tokens 512 \
|
||||||
--output results/humaneval.json
|
--output results/humaneval.json
|
||||||
```
|
```
|
||||||
@@ -47,7 +78,8 @@ python scripts/eval/evaluate_humaneval.py \
|
|||||||
| `--temperature` | 0.8 | Sampling temperature |
|
| `--temperature` | 0.8 | Sampling temperature |
|
||||||
| `--top_p` | 0.95 | Nucleus sampling threshold |
|
| `--top_p` | 0.95 | Nucleus sampling threshold |
|
||||||
| `--top_k` | 50 | Top-k sampling |
|
| `--top_k` | 50 | Top-k sampling |
|
||||||
| `--batch_size` | 32 | Generation batch size |
|
| `--batch_size` | 64 | Generation batch size |
|
||||||
|
| `--max_seq_len` | 4096 | KV cache sequence length |
|
||||||
| `--test_workers` | 8 | ProcessPoolExecutor workers for test execution |
|
| `--test_workers` | 8 | ProcessPoolExecutor workers for test execution |
|
||||||
| `--test_timeout` | 3.0 | Per-subprocess timeout (seconds) |
|
| `--test_timeout` | 3.0 | Per-subprocess timeout (seconds) |
|
||||||
| `--problems` | None | Restrict to specific problem indices |
|
| `--problems` | None | Restrict to specific problem indices |
|
||||||
@@ -66,7 +98,7 @@ python scripts/eval/evaluate_humaneval.py \
|
|||||||
python scripts/eval/evaluate_mmlu.py \
|
python scripts/eval/evaluate_mmlu.py \
|
||||||
--param_path ./params \
|
--param_path ./params \
|
||||||
--n_shot 5 \
|
--n_shot 5 \
|
||||||
--subjects math_algebra history_us \
|
--subjects abstract_algebra high_school_us_history \
|
||||||
--output results/mmlu.json
|
--output results/mmlu.json
|
||||||
```
|
```
|
||||||
|
|
||||||
@@ -82,12 +114,13 @@ python scripts/eval/evaluate_mmlu.py \
|
|||||||
| `--device` | auto | Device (`cuda` / `cpu`) |
|
| `--device` | auto | Device (`cuda` / `cpu`) |
|
||||||
| `--dtype` | auto | `bfloat16` on CUDA, `float32` on CPU |
|
| `--dtype` | auto | `bfloat16` on CUDA, `float32` on CPU |
|
||||||
| `--seed` | 0 | Seed for option permutation (0 = enabled, -1 = disabled) |
|
| `--seed` | 0 | Seed for option permutation (0 = enabled, -1 = disabled) |
|
||||||
|
| `--batch_size` | 4 | Questions per batch; each question produces four choice rows |
|
||||||
|
|
||||||
**How it works**: For each question, builds a prompt with n-shot examples, then scores each choice (A/B/C/D) by computing the summed log-likelihood of the choice token given the context. The choice with the highest log-prob is the prediction.
|
**How it works**: For each question, builds a prompt with n-shot examples, then scores each choice (A/B/C/D) by computing the summed log-likelihood of the choice token given the context. The choice with the highest log-prob is the prediction.
|
||||||
|
|
||||||
**Output**: stdout prints per-subject accuracy and overall. With `--output`, writes per-subject `{accuracy, correct, total}` + `_overall` aggregate.
|
**Output**: stdout prints per-subject accuracy and overall. With `--output`, writes per-subject `{accuracy, correct, total}` + `_overall` aggregate.
|
||||||
|
|
||||||
**Data**: Auto-downloads `cais/mmlu` from HuggingFace. Stored as per-subject CSVs in `<data_dir>/<split>/` and `<data_dir>/dev/` (for few-shot).
|
**Data**: Auto-downloads `cais/mmlu` from HuggingFace. Stored as per-subject CSVs in `<data_dir>/<split>/` and `<data_dir>/dev/` (for few-shot). `--subjects` accepts canonical MMLU names such as `abstract_algebra`, `college_computer_science`, `high_school_us_history`, and `world_religions`.
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
@@ -100,7 +133,7 @@ python scripts/eval/evaluate_ppl.py \
|
|||||||
--param_path ./params \
|
--param_path ./params \
|
||||||
--input_path data.jsonl \
|
--input_path data.jsonl \
|
||||||
--output_dir ppl_results/ \
|
--output_dir ppl_results/ \
|
||||||
--batch_size 4 \
|
--batch_size 64 \
|
||||||
--max_length 2048
|
--max_length 2048
|
||||||
```
|
```
|
||||||
|
|
||||||
@@ -110,7 +143,7 @@ python scripts/eval/evaluate_ppl.py \
|
|||||||
| `--input_path` | required | Input file, glob, or directory |
|
| `--input_path` | required | Input file, glob, or directory |
|
||||||
| `--output_dir` | required | Output directory for `summary.json` + token JSONL |
|
| `--output_dir` | required | Output directory for `summary.json` + token JSONL |
|
||||||
| `--text_key` | `text` | Key for the text field in input data |
|
| `--text_key` | `text` | Key for the text field in input data |
|
||||||
| `--batch_size` | 4 | Batch size |
|
| `--batch_size` | 64 | Batch size |
|
||||||
| `--max_length` | 2048 | Max sequence length (tokens) |
|
| `--max_length` | 2048 | Max sequence length (tokens) |
|
||||||
| `--token_level` | False | Store per-token log_probs + token-type analysis |
|
| `--token_level` | False | Store per-token log_probs + token-type analysis |
|
||||||
| `--max_samples` | None | Random subsample per file |
|
| `--max_samples` | None | Random subsample per file |
|
||||||
@@ -119,7 +152,7 @@ python scripts/eval/evaluate_ppl.py \
|
|||||||
|
|
||||||
**Input**: JSONL or JSON files. Each item must have a field named by `--text_key` (default `text`). If `--input_path` is a directory, recursively collects `*.jsonl` and `*.json`.
|
**Input**: JSONL or JSON files. Each item must have a field named by `--text_key` (default `text`). If `--input_path` is a directory, recursively collects `*.jsonl` and `*.json`.
|
||||||
|
|
||||||
**Output**: `summary.json` with per-file stats (tokens, mean/median loss, perplexity, p50/p90/p95/p99). With `--token_level`, also writes per-token JSONL with token IDs and log-probs.
|
**Output**: `summary.json` with per-file token count, mean loss, perplexity, and p50/p90/p95/p99 loss. Median loss is included only with `--token_level`; that mode also writes per-token JSONL with token IDs and log-probs.
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
@@ -210,7 +243,8 @@ python scripts/eval/evaluate_ifeval.py \
|
|||||||
| `--top_p` | 0.95 | Top-p sampling |
|
| `--top_p` | 0.95 | Top-p sampling |
|
||||||
| `--top_k` | 50 | Top-k sampling |
|
| `--top_k` | 50 | Top-k sampling |
|
||||||
| `--num_samples` | 1 | Samples per problem (best-of-n scoring) |
|
| `--num_samples` | 1 | Samples per problem (best-of-n scoring) |
|
||||||
| `--batch_size` | 1 | Inference batch size |
|
| `--batch_size` | 64 | Inference batch size |
|
||||||
|
| `--max_seq_len` | 4096 | KV cache sequence length |
|
||||||
| `--limit` | None | Limit to first N problems (quick testing) |
|
| `--limit` | None | Limit to first N problems (quick testing) |
|
||||||
| `--dump_responses` | None | Path to dump raw responses as JSONL |
|
| `--dump_responses` | None | Path to dump raw responses as JSONL |
|
||||||
|
|
||||||
@@ -232,7 +266,7 @@ python scripts/eval/analyze_weights.py \
|
|||||||
|
|
||||||
| Parameter | Default | Description |
|
| Parameter | Default | Description |
|
||||||
|-----------|---------|-------------|
|
|-----------|---------|-------------|
|
||||||
| `--ckpt_dir` | required | Checkpoint dir with `model.safetensors` + `config.json` |
|
| `--ckpt_dir` | required | Checkpoint directory containing `model.safetensors` |
|
||||||
| `--compare` | None | Additional checkpoint dirs to compare |
|
| `--compare` | None | Additional checkpoint dirs to compare |
|
||||||
| `--no_svd` | False | Skip SVD; show only weight stats (faster) |
|
| `--no_svd` | False | Skip SVD; show only weight stats (faster) |
|
||||||
| `--output` | None | Save results as JSON |
|
| `--output` | None | Save results as JSON |
|
||||||
@@ -245,8 +279,8 @@ python scripts/eval/analyze_weights.py \
|
|||||||
## Tips
|
## Tips
|
||||||
|
|
||||||
- **Quick test**: Use `--limit` (IFEval) or `--problems` (HumanEval) to run on a small subset first.
|
- **Quick test**: Use `--limit` (IFEval) or `--problems` (HumanEval) to run on a small subset first.
|
||||||
- **Auto-download**: HumanEval, MMLU, and IFEval auto-download their datasets on first run. The other scripts expect user-provided data.
|
- **Auto-download**: After installing `datasets`, HumanEval, MMLU, and IFEval auto-download their datasets on first run. The other scripts expect user-provided data.
|
||||||
- **Output formats**: `--output` writes a single JSON for most scripts. PPL and IFD write an `--output_dir` containing `summary.json` plus per-file artifacts.
|
- **Output formats**: `--output` writes a single JSON for most scripts. PPL and IFD write an `--output_dir` containing `summary.json` plus per-file artifacts.
|
||||||
- **CPU mode**: All scripts auto-detect CUDA. To force CPU, use `--device cpu --dtype float32`.
|
- **CPU mode**: MMLU, PPL, and IFD support `--device cpu --dtype float32`; weight analysis supports `--device cpu`. HumanEval generation and IFEval are CUDA-only.
|
||||||
|
|
||||||
> Document Update Time: 2026-07-30
|
> Document Update Time: 2026-07-30
|
||||||
|
|||||||
+52
-13
@@ -49,8 +49,7 @@ KVCache
|
|||||||
├── seq_lens [batch_size]
|
├── seq_lens [batch_size]
|
||||||
├── out_cache_loc [batch, seq_len] — write indices for this forward
|
├── out_cache_loc [batch, seq_len] — write indices for this forward
|
||||||
├── max_len int — max(seq_lens), avoids GPU sync in decode
|
├── max_len int — max(seq_lens), avoids GPU sync in decode
|
||||||
├── page_table [batch, max_len] — precomputed gather indices for decode (None for prefill)
|
└── kv_indptr [batch + 1] int32 — prefix sum of seq_lens, precomputed once per step
|
||||||
└── decode_mask [batch, max_len] bool — precomputed position validity mask (None for single-batch decode)
|
|
||||||
```
|
```
|
||||||
|
|
||||||
Attention layers do raw buffer indexing: `k_buffer[layer_id, out_cache_loc] = k` to write, `k_buffer[layer_id, indices]` to gather.
|
Attention layers do raw buffer indexing: `k_buffer[layer_id, out_cache_loc] = k` to write, `k_buffer[layer_id, indices]` to gather.
|
||||||
@@ -87,7 +86,9 @@ Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/rotary_b
|
|||||||
- **CUDA kernel** (`rotary_emb.cu`): fused cos/sin lookup + rotation in a single kernel, used when the kernel is available, input is on CUDA, and `torch.is_grad_enabled()` is `False` (inference mode)
|
- **CUDA kernel** (`rotary_emb.cu`): fused cos/sin lookup + rotation in a single kernel, used when the kernel is available, input is on CUDA, and `torch.is_grad_enabled()` is `False` (inference mode)
|
||||||
- **Torch fallback**: complex multiply path (`torch.view_as_complex` → `torch.complex` multiply → `torch.view_as_real`), used during training (supports autograd backward) or when the CUDA kernel is not available
|
- **Torch fallback**: complex multiply path (`torch.view_as_complex` → `torch.complex` multiply → `torch.view_as_real`), used during training (supports autograd backward) or when the CUDA kernel is not available
|
||||||
|
|
||||||
`RotaryEmbedding` stores `cos_table`/`sin_table` as f32 buffers and returns a `(cos, sin)` tuple from `forward()`. Both attention backends share the same rotary dispatch — it is backend-agnostic.
|
`RotaryEmbedding` stores a complex `freqs_cis` buffer and returns a tensor
|
||||||
|
from `forward()`. Both attention backends share the same rotary dispatch — it
|
||||||
|
is backend-agnostic.
|
||||||
|
|
||||||
## Continuous Batching
|
## Continuous Batching
|
||||||
|
|
||||||
@@ -183,22 +184,58 @@ curl -X POST http://localhost:8000/v1/messages \
|
|||||||
-d '{"model":"astrai","system":"You are helpful.","messages":[{"role":"user","content":"Hello"}],"max_tokens":512}'
|
-d '{"model":"astrai","system":"You are helpful.","messages":[{"role":"user","content":"Hello"}],"max_tokens":512}'
|
||||||
```
|
```
|
||||||
|
|
||||||
Supports `stop_sequences` and streaming via `event: content_block_delta`.
|
Supports `stop_sequences` and streaming via `event: content_block_delta`. Anthropic streams also end with the shared `data: [DONE]` sentinel after `event: message_stop`.
|
||||||
|
|
||||||
### GenerationRequest Parameters
|
### Request Parameters
|
||||||
|
|
||||||
|
The HTTP protocols and direct engine API have distinct request models and defaults.
|
||||||
|
|
||||||
|
**OpenAI** (`ChatCompletionRequest`):
|
||||||
|
|
||||||
| Param | Type | Default | Description |
|
| Param | Type | Default | Description |
|
||||||
|-------|------|---------|-------------|
|
|-------|------|---------|-------------|
|
||||||
|
| `model` | str | `"astrai"` | Model name returned in responses |
|
||||||
| `messages` | List[dict] | required | Chat messages (role, content) |
|
| `messages` | List[dict] | required | Chat messages (role, content) |
|
||||||
| `top_k` | int | 50 | Top-k count |
|
| `temperature` | Optional[float] | 1.0 | Sampling temperature (0.0-2.0) |
|
||||||
| `top_p` | float | 1.0 | Nucleus threshold |
|
| `top_p` | Optional[float] | 1.0 | Nucleus threshold (0.0-1.0) |
|
||||||
| `temperature` | float | 1.0 | Sampling temperature (> 0.0) |
|
| `top_k` | Optional[int] | 50 | Top-k count |
|
||||||
| `max_tokens` | Optional[int] | None | Max generation length |
|
| `max_tokens` | Optional[int] | 2048 | Max generation length |
|
||||||
| `stream` | bool | False | Stream output |
|
| `stream` | Optional[bool] | False | Stream output |
|
||||||
| `stop` | Optional[Union[str, List[str]]] | None | Stop sequences |
|
| `stop` | Optional[Union[str, List[str]]] | None | Stop sequences |
|
||||||
| `frequency_penalty` | float | 0.0 | Frequency penalty |
|
| `n` | Optional[int] | 1 | Number of choices requested |
|
||||||
| `tools` | Optional[List[dict]] | None | Tool definitions for function calling |
|
| `presence_penalty` | Optional[float] | 0.0 | Presence penalty (-2.0 to 2.0) |
|
||||||
| `tool_choice` | Optional[str] | None | Tool selection mode |
|
| `frequency_penalty` | Optional[float] | 0.0 | Frequency penalty (-2.0 to 2.0) |
|
||||||
|
| `logit_bias` | Optional[Dict[int, float]] | None | Per-token logit bias |
|
||||||
|
| `user` | Optional[str] | None | End-user identifier |
|
||||||
|
| `tools` | Optional[List[ToolDef]] | None | Tool definitions for function calling |
|
||||||
|
| `tool_choice` | Optional[Union[str, Dict[str, Any]]] | `"auto"` | Tool selection mode or explicit tool choice |
|
||||||
|
|
||||||
|
**Anthropic** (`MessagesRequest`):
|
||||||
|
|
||||||
|
| Param | Type | Default | Description |
|
||||||
|
|-------|------|---------|-------------|
|
||||||
|
| `model` | str | `"astrai"` | Model name returned in responses |
|
||||||
|
| `messages` | List[AnthropicMessage] | required | User/assistant messages |
|
||||||
|
| `system` | Optional[str] | None | System prompt |
|
||||||
|
| `max_tokens` | int | 1024 | Max generation length |
|
||||||
|
| `temperature` | Optional[float] | 1.0 | Sampling temperature (0.0-2.0) |
|
||||||
|
| `top_p` | Optional[float] | 1.0 | Nucleus threshold (0.0-1.0) |
|
||||||
|
| `top_k` | Optional[int] | 50 | Top-k count |
|
||||||
|
| `stream` | Optional[bool] | False | Stream output |
|
||||||
|
| `stop_sequences` | Optional[List[str]] | None | Stop sequences |
|
||||||
|
|
||||||
|
**Engine** (`GenerationRequest`):
|
||||||
|
|
||||||
|
| Param | Type | Default | Description |
|
||||||
|
|-------|------|---------|-------------|
|
||||||
|
| `messages` | List[Dict[str, str]] | required | Messages to format before generation |
|
||||||
|
| `top_k` | int | 50 | Top-k count; 0 disables filtering |
|
||||||
|
| `top_p` | float | 1.0 | Nucleus threshold |
|
||||||
|
| `temperature` | float | 1.0 | Sampling temperature; 0 enables greedy decoding |
|
||||||
|
| `max_tokens` | Optional[int] | None | Max generation length |
|
||||||
|
| `frequency_penalty` | float | 0.0 | Frequency penalty (-2.0 to 2.0) |
|
||||||
|
| `rep_window` | int | 64 | Recent-token window used by the frequency penalty |
|
||||||
|
| `stream` | bool | False | Stream output |
|
||||||
|
|
||||||
### SSE Streaming Format
|
### SSE Streaming Format
|
||||||
|
|
||||||
@@ -240,6 +277,8 @@ data: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":
|
|||||||
|
|
||||||
event: message_stop
|
event: message_stop
|
||||||
data: {"type":"message_stop"}
|
data: {"type":"message_stop"}
|
||||||
|
|
||||||
|
data: [DONE]
|
||||||
```
|
```
|
||||||
|
|
||||||
### Error Responses
|
### Error Responses
|
||||||
|
|||||||
+37
-30
@@ -13,9 +13,11 @@
|
|||||||
|
|
||||||
| Parameter | Description | Default |
|
| Parameter | Description | Default |
|
||||||
|-----------|-------------|---------|
|
|-----------|-------------|---------|
|
||||||
|
| `--config`, `-c` | YAML config file; explicit CLI options override YAML values | None |
|
||||||
| `--train_type` | Training type (`seq`, `sft`, `dpo`, `grpo`, `online_grpo`, `online_dpo`) | required |
|
| `--train_type` | Training type (`seq`, `sft`, `dpo`, `grpo`, `online_grpo`, `online_dpo`) | required |
|
||||||
| `--data_root_path` | Dataset root directory | required |
|
| `--data_root_path` | Dataset root directory | required |
|
||||||
| `--param_path` | Model parameters or checkpoint path | required |
|
| `--param_path` | Model parameters or checkpoint path | required |
|
||||||
|
| `--resume` | Resume training from `--param_path` | False |
|
||||||
| `--n_epoch` | Total training epochs | 1 |
|
| `--n_epoch` | Total training epochs | 1 |
|
||||||
| `--batch_per_device` | Batch size per device | 1 |
|
| `--batch_per_device` | Batch size per device | 1 |
|
||||||
| `--grad_accum_steps` | Gradient accumulation steps between optimizer steps | 1 |
|
| `--grad_accum_steps` | Gradient accumulation steps between optimizer steps | 1 |
|
||||||
@@ -26,7 +28,7 @@
|
|||||||
|-----------|-------------|---------|
|
|-----------|-------------|---------|
|
||||||
| `--warmup_ratio` | Fraction of total steps used for LR warmup | 0.05 |
|
| `--warmup_ratio` | Fraction of total steps used for LR warmup | 0.05 |
|
||||||
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
|
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
|
||||||
| `--max_grad_norm` | Maximum gradient norm for clipping (None disables) | 1.0 |
|
| `--max_grad_norm` | Maximum gradient norm for clipping; the current CLI requires a positive number | 1.0 |
|
||||||
|
|
||||||
### Optimizer
|
### Optimizer
|
||||||
|
|
||||||
@@ -36,9 +38,9 @@ non-matrix parameters through **AdamW** (`fused=True`).
|
|||||||
| Parameter | Description | Default |
|
| Parameter | Description | Default |
|
||||||
|-----------|-------------|---------|
|
|-----------|-------------|---------|
|
||||||
| `--optimizer` | Built-in optimizer (`muon_adamw`, `nora_nadamw`, `mano_adamw`) | `muon_adamw` |
|
| `--optimizer` | Built-in optimizer (`muon_adamw`, `nora_nadamw`, `mano_adamw`) | `muon_adamw` |
|
||||||
| `--weight_decay` | Weight decay (applied to Muon matrix params; non-matrix use 0) | 0.1 |
|
| `--weight_decay` | Weight decay for optimizer parameter groups that are eligible for decay | 0.1 |
|
||||||
| `--muon_momentum` | Muon momentum factor | 0.95 |
|
| `--muon_momentum` | Muon momentum factor | 0.95 |
|
||||||
| `--muon_nesterov` | Enable Nesterov momentum for Muon | True |
|
| `--muon_nesterov`, `--no-muon_nesterov` | Enable or disable Nesterov momentum for Muon | enabled |
|
||||||
| `--muon_ns_steps` | Newton-Schulz iteration steps for Muon | 5 |
|
| `--muon_ns_steps` | Newton-Schulz iteration steps for Muon | 5 |
|
||||||
| `--muon_adjust_lr` | Muon LR adjustment strategy (`original`, `match_rms_adamw`) | `match_rms_adamw` |
|
| `--muon_adjust_lr` | Muon LR adjustment strategy (`original`, `match_rms_adamw`) | `match_rms_adamw` |
|
||||||
|
|
||||||
@@ -56,15 +58,18 @@ under DTensor sharding and rejects layouts sharded along the last dimension.
|
|||||||
| `--nora_weight_decay` | Nora matrix weight decay | 0.0 |
|
| `--nora_weight_decay` | Nora matrix weight decay | 0.0 |
|
||||||
|
|
||||||
`mano_adamw` routes internal `Linear.weight` matrices to **Mano** (manifold
|
`mano_adamw` routes internal `Linear.weight` matrices to **Mano** (manifold
|
||||||
normalized optimizer) and the remaining parameters to **NAdamW**. Mano projects
|
normalized optimizer) and the remaining parameters to **AdamW**. Mano projects
|
||||||
the momentum onto the tangent space of the Oblique manifold and normalizes it,
|
the momentum onto the tangent space of the Oblique manifold and normalizes it,
|
||||||
alternating the projection axis (row/column) each step — replacing Muon's
|
alternating the projection axis (row/column) each step — replacing Muon's
|
||||||
Newton-Schulz iteration with a cheaper normalization.
|
Newton-Schulz iteration with a cheaper normalization.
|
||||||
|
|
||||||
| Parameter | Description | Default |
|
| Parameter | Description | Default |
|
||||||
|-----------|-------------|---------|
|
|-----------|-------------|---------|
|
||||||
| `--mano_momentum` | Mano momentum factor | 0.95 |
|
| `--mano_momentum` | Accepted by the CLI but currently ignored by optimizer construction | 0.95 |
|
||||||
| `--mano_nesterov` | Enable Nesterov momentum for Mano | True |
|
| `--mano_nesterov`, `--no-mano_nesterov` | Accepted by the CLI but currently ignored by optimizer construction | enabled |
|
||||||
|
|
||||||
|
The two Mano-specific flags are reserved for future wiring; do not rely on them
|
||||||
|
to change optimizer behavior in the current release.
|
||||||
|
|
||||||
Optimizer identity and hyperparameters are saved in checkpoint metadata. Optimizer
|
Optimizer identity and hyperparameters are saved in checkpoint metadata. Optimizer
|
||||||
states are intentionally not interchangeable: resume older MuonAdamW checkpoints
|
states are intentionally not interchangeable: resume older MuonAdamW checkpoints
|
||||||
@@ -78,7 +83,7 @@ with `--optimizer=muon_adamw`.
|
|||||||
| `--stride` | Stride for sliding window over sequences | None |
|
| `--stride` | Stride for sliding window over sequences | None |
|
||||||
| `--random_seed` | Random seed for reproducibility | 3407 |
|
| `--random_seed` | Random seed for reproducibility | 3407 |
|
||||||
| `--num_workers` | DataLoader worker processes | 4 |
|
| `--num_workers` | DataLoader worker processes | 4 |
|
||||||
| `--no_pin_memory` | Disable pin_memory (enabled by default) | (flag) |
|
| `--pin_memory`, `--no-pin_memory` | Enable or disable DataLoader pinned memory | enabled |
|
||||||
|
|
||||||
### Checkpoint & Resume
|
### Checkpoint & Resume
|
||||||
|
|
||||||
@@ -100,14 +105,20 @@ with `--optimizer=muon_adamw`.
|
|||||||
|
|
||||||
| Parameter | Description | Default |
|
| Parameter | Description | Default |
|
||||||
|-----------|-------------|---------|
|
|-----------|-------------|---------|
|
||||||
| `--log_dir` | Directory for metric logs | checkpoint/logs |
|
| `--metrics` | Repeatable metric option (for example, `--metrics loss --metrics lr --metrics val_loss`) | `loss`, `lr`, `grad_norm`, `grad_snr` |
|
||||||
| `--metrics` | Metrics to log (e.g. --metrics loss lr val_loss) | ["loss", "lr", "grad_norm"] |
|
|
||||||
|
|
||||||
### Gradient Checkpointing
|
### Gradient Checkpointing
|
||||||
|
|
||||||
| Parameter | Description | Default |
|
| Parameter | Description | Default |
|
||||||
|-----------|-------------|---------|
|
|-----------|-------------|---------|
|
||||||
| `--gradient_checkpointing` | Enable activation checkpointing for DecoderBlock modules | False |
|
| `--gradient_checkpointing`, `--no-gradient_checkpointing` | Enable or disable activation checkpointing for DecoderBlock modules | disabled |
|
||||||
|
|
||||||
|
### Miscellaneous
|
||||||
|
|
||||||
|
| Parameter | Description | Default |
|
||||||
|
|-----------|-------------|---------|
|
||||||
|
| `--compile` | Enable `torch.compile` with mode `default`, `reduce-overhead`, or `max-autotune`; omit to disable | None |
|
||||||
|
| `--dry-run` | Validate the merged configuration and print the training plan without training | False |
|
||||||
|
|
||||||
### Distributed Training
|
### Distributed Training
|
||||||
|
|
||||||
@@ -120,21 +131,25 @@ with `--optimizer=muon_adamw`.
|
|||||||
| `--backend` | Distributed training backend | nccl |
|
| `--backend` | Distributed training backend | nccl |
|
||||||
| `--master_addr` | Master node address | localhost |
|
| `--master_addr` | Master node address | localhost |
|
||||||
| `--master_port` | Master node port | 29500 |
|
| `--master_port` | Master node port | 29500 |
|
||||||
|
| `--tp_size` | Reserved tensor-parallel size; accepted but currently ignored | None |
|
||||||
|
|
||||||
### Strategy-specific
|
### Strategy-specific
|
||||||
|
|
||||||
| Parameter | Description | Default | Used by |
|
| Parameter | Description | Default | Used by |
|
||||||
|-----------|-------------|---------|---------|
|
|-----------|-------------|---------|---------|
|
||||||
| `--dpo_beta` | DPO beta value | 0.1 | `dpo` |
|
| `--dpo_beta` | DPO beta value | 0.1 | `dpo`, `online_dpo` |
|
||||||
| `--label_smoothing` | Label smoothing for cross-entropy loss | 0.0 | `seq`, `sft` |
|
| `--label_smoothing` | Label smoothing for cross-entropy loss | 0.0 | `seq`, `sft` |
|
||||||
| `--group_size` | GRPO group size | 4 | `grpo` |
|
| `--group_size` | GRPO/rollout group size | 4 | `grpo`, `online_grpo`, `online_dpo` |
|
||||||
| `--grpo_clip_eps` | GRPO clipping epsilon | 0.2 | `grpo` |
|
| `--grpo_clip_eps` | GRPO clipping epsilon | 0.2 | `grpo`, `online_grpo` |
|
||||||
| `--grpo_kl_coef` | GRPO KL penalty coefficient | 0.01 | `grpo` |
|
| `--grpo_kl_coef` | GRPO KL penalty coefficient | 0.01 | `grpo`, `online_grpo` |
|
||||||
| `--neftune_alpha` | NEFTune noise alpha (0=disabled, typical: 5.0) | 0.0 | `sft` |
|
| `--neftune_alpha` | NEFTune noise alpha (0=disabled, typical: 5.0) | 0.0 | `sft` |
|
||||||
|
|
||||||
### Online Rollout
|
### Online Rollout
|
||||||
|
|
||||||
These options apply to `online_grpo` and `online_dpo`. Online strategies require
|
`online_grpo` and `online_dpo` are factory aliases for the existing `grpo` and
|
||||||
|
`dpo` strategy classes; online behavior is enabled by rollout components rather
|
||||||
|
than separate strategy subclasses. These options apply to the online aliases.
|
||||||
|
Online strategies require
|
||||||
a `BaseRewardModel` factory in `TrainConfig`; `train.py` does not currently
|
a `BaseRewardModel` factory in `TrainConfig`; `train.py` does not currently
|
||||||
provide a command-line option for configuring one.
|
provide a command-line option for configuring one.
|
||||||
|
|
||||||
@@ -151,7 +166,7 @@ provide a command-line option for configuring one.
|
|||||||
| Parameter | Description | Default |
|
| Parameter | Description | Default |
|
||||||
|-----------|-------------|---------|
|
|-----------|-------------|---------|
|
||||||
| `--schedule_type` | LR scheduler type (`cosine`, `sgdr`, `wsd`) | cosine |
|
| `--schedule_type` | LR scheduler type (`cosine`, `sgdr`, `wsd`) | cosine |
|
||||||
| `--min_rate` | Minimum LR as fraction of base LR | None (scheduler default: 0.05 for cosine/SGDR, 0.0 for WSD) |
|
| `--min_rate` | Minimum LR as fraction of base LR | None (all current schedulers use their effective default of 0.01) |
|
||||||
| `--cycle_length` | SGDR first cycle length in steps | None (total_steps - warmup_steps) |
|
| `--cycle_length` | SGDR first cycle length in steps | None (total_steps - warmup_steps) |
|
||||||
| `--t_mult` | SGDR cycle length multiplier per restart | 2 |
|
| `--t_mult` | SGDR cycle length multiplier per restart | 2 |
|
||||||
| `--stable_steps` | WSD stable plateau steps | None (80% of post-warmup steps) |
|
| `--stable_steps` | WSD stable plateau steps | None (80% of post-warmup steps) |
|
||||||
@@ -204,14 +219,6 @@ python scripts/tools/server.py --param_path ./params --device cuda --dtype bfloa
|
|||||||
|
|
||||||
See [Inference Guide](inference.md) for HTTP API documentation.
|
See [Inference Guide](inference.md) for HTTP API documentation.
|
||||||
|
|
||||||
# Preprocess
|
|
||||||
|
|
||||||
```bash
|
|
||||||
python scripts/tools/preprocess.py data/*.jsonl -o output/ -c config.json
|
|
||||||
```
|
|
||||||
|
|
||||||
See [Preprocessing Guide](preprocessing.md) for config file format and examples.
|
|
||||||
|
|
||||||
## Generate (`generate.py`)
|
## Generate (`generate.py`)
|
||||||
|
|
||||||
| Parameter | Type | Default | Description |
|
| Parameter | Type | Default | Description |
|
||||||
@@ -221,13 +228,12 @@ See [Preprocessing Guide](preprocessing.md) for config file format and examples.
|
|||||||
| `--output_json_file` | str | required | Path to the output JSONL file |
|
| `--output_json_file` | str | required | Path to the output JSONL file |
|
||||||
| `--question_key` | str | `question` | Key for the question in input JSON |
|
| `--question_key` | str | `question` | Key for the question in input JSON |
|
||||||
| `--response_key` | str | `response` | Key for the response in output JSON |
|
| `--response_key` | str | `response` | Key for the response in output JSON |
|
||||||
| `--temperature` | float | `0.60` | Sampling temperature |
|
| `--temperature` | float | `0.8` | Sampling temperature |
|
||||||
| `--top_k` | int | `30` | Top-k filtering |
|
| `--top_k` | int | `50` | Top-k filtering |
|
||||||
| `--top_p` | float | `0.95` | Nucleus sampling threshold |
|
| `--top_p` | float | `0.95` | Nucleus sampling threshold |
|
||||||
| `--batch_size` | int | `1` | Batch size for generation |
|
| `--batch_size` | int | `1` | Batch size for generation |
|
||||||
| `--num_samples` | int | `1` | Responses per prompt |
|
| `--num_samples` | int | `1` | Responses per prompt |
|
||||||
| `--max_tokens` | int | model config `max_position_embeddings` | Maximum tokens to generate |
|
| `--max_seq_len` | int | `2048` | KV cache sequence length |
|
||||||
| `--cache_len` | int | `2048` | KV cache length |
|
|
||||||
| `--frequency_penalty` | float | `0.0` | Frequency penalty |
|
| `--frequency_penalty` | float | `0.0` | Frequency penalty |
|
||||||
| `--rep_window` | int | `64` | Window size for frequency penalty |
|
| `--rep_window` | int | `64` | Window size for frequency penalty |
|
||||||
|
|
||||||
@@ -243,14 +249,15 @@ python scripts/tools/generate.py \
|
|||||||
|
|
||||||
| Parameter | Type | Default | Description |
|
| Parameter | Type | Default | Description |
|
||||||
|-----------|------|---------|-------------|
|
|-----------|------|---------|-------------|
|
||||||
| `input_files` | path(s) | required | Input JSONL file(s), supports glob (`data/*.jsonl`) |
|
| `input_files` | path(s) | required | One or more existing `.jsonl` or `.json` paths. Wildcards work only when expanded by the invoking shell; the CLI does not expand globs itself. |
|
||||||
| `--output_dir`, `-o` | path | required | Output directory for processed data |
|
| `--output_dir`, `-o` | path | required | Output directory for processed data |
|
||||||
| `--config`, `-c` | path | required | Preprocessing pipeline config (JSON) |
|
| `--config`, `-c` | path | required | Preprocessing pipeline config (JSON) |
|
||||||
| `--tokenizer_path` | str | `params` | Path to tokenizer directory |
|
| `--tokenizer_path` | str | `params` | Path to tokenizer directory |
|
||||||
|
| `--batch_size` | int | config value (`256` by default) | Override records processed per batch; must be at least 1 |
|
||||||
|
|
||||||
Usage:
|
Usage:
|
||||||
```bash
|
```bash
|
||||||
python scripts/tools/preprocess.py data/*.jsonl -o output/ -c sft.json
|
python scripts/tools/preprocess.py data/part-000.jsonl data/part-001.jsonl -o output/ -c sft.json
|
||||||
```
|
```
|
||||||
|
|
||||||
See [Preprocessing Guide](preprocessing.md) for config file format and examples.
|
See [Preprocessing Guide](preprocessing.md) for config file format and examples.
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ Declarative JSON-driven data preprocessing. `MaskBuilderFactory` supports three
|
|||||||
- [Configuration Reference](#configuration-reference) — all fields
|
- [Configuration Reference](#configuration-reference) — all fields
|
||||||
- [Mask Algorithm](#mask-algorithm)
|
- [Mask Algorithm](#mask-algorithm)
|
||||||
- [Output Layout](#output-layout)
|
- [Output Layout](#output-layout)
|
||||||
|
- [Training Compatibility](#training-compatibility)
|
||||||
- [CLI](#cli)
|
- [CLI](#cli)
|
||||||
- [Python API](#python-api)
|
- [Python API](#python-api)
|
||||||
|
|
||||||
@@ -40,7 +41,7 @@ A single config file captures the entire pipeline, reusable and version-controll
|
|||||||
| Field | Type | Default | Description |
|
| Field | Type | Default | Description |
|
||||||
|-------|------|---------|-------------|
|
|-------|------|---------|-------------|
|
||||||
| `field` | str | -- | JSONL key to read |
|
| `field` | str | -- | JSONL key to read |
|
||||||
| `action` | str | -- | `"train"` / `"mask"` / `"$role"` |
|
| `action` | str | -- | `"train"` / `"mask"` / `"$role"` / `"value"`; `"value"` copies raw values without tokenization |
|
||||||
| `template` | bool | `false` | Apply `chat_template` per message |
|
| `template` | bool | `false` | Apply `chat_template` per message |
|
||||||
| `add_special_tokens` | bool | `true` for first non-template section | Add special tokens during encode |
|
| `add_special_tokens` | bool | `true` for first non-template section | Add special tokens during encode |
|
||||||
|
|
||||||
@@ -89,7 +90,7 @@ Config:
|
|||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
Output keys: `sequence` (int32), `loss_mask` (bool)
|
Output keys: `sequence` (int32), `loss_mask` (bool), `position_ids` (int32)
|
||||||
|
|
||||||
### SFT Instruction
|
### SFT Instruction
|
||||||
|
|
||||||
@@ -116,7 +117,7 @@ Config:
|
|||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
Output keys: `sequence`, `loss_mask`
|
Output keys: `sequence`, `loss_mask`, `position_ids`
|
||||||
|
|
||||||
### Pretrain
|
### Pretrain
|
||||||
|
|
||||||
@@ -142,7 +143,7 @@ Config:
|
|||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
Output keys: `sequence` (no `loss_mask` — all tokens trained)
|
Output keys: `sequence`, `position_ids` (no `loss_mask` — all tokens trained)
|
||||||
|
|
||||||
### DPO
|
### DPO
|
||||||
|
|
||||||
@@ -180,6 +181,11 @@ Config:
|
|||||||
|
|
||||||
Output keys: `chosen`, `chosen_mask`, `rejected`, `rejected_mask`
|
Output keys: `chosen`, `chosen_mask`, `rejected`, `rejected_mask`
|
||||||
|
|
||||||
|
The offline `Pipeline` can construct these keys, but its `.bin` output is not
|
||||||
|
currently loadable for DPO training because the writer does not preserve
|
||||||
|
per-record offsets. Train DPO directly from raw JSONL instead; see
|
||||||
|
[Training Compatibility](#training-compatibility).
|
||||||
|
|
||||||
### GRPO
|
### GRPO
|
||||||
|
|
||||||
Input JSONL:
|
Input JSONL:
|
||||||
@@ -228,6 +234,11 @@ Output keys: `prompts`, `prompts_mask`, `responses`, `masks`, `rewards` (float32
|
|||||||
- `mask_key: "masks"` — rename the auto-generated mask key (default: `responses_mask`)
|
- `mask_key: "masks"` — rename the auto-generated mask key (default: `responses_mask`)
|
||||||
- `prompts_mask` is auto-generated (all masked) and unused by GRPOStrategy
|
- `prompts_mask` is auto-generated (all masked) and unused by GRPOStrategy
|
||||||
|
|
||||||
|
The offline `Pipeline` flattens GRPO response groups for `.bin` output without
|
||||||
|
preserving their boundaries, and there is no automatic raw-JSONL GRPO processor
|
||||||
|
in `DatasetFactory`. See
|
||||||
|
[Training Compatibility](#training-compatibility) for the supported routes.
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## Configuration Reference
|
## Configuration Reference
|
||||||
@@ -257,7 +268,7 @@ When `sources` is set, `sections` is ignored.
|
|||||||
| `max_chars` | int | `2000000` | Skip text-mode items longer than this |
|
| `max_chars` | int | `2000000` | Skip text-mode items longer than this |
|
||||||
| `max_items` | int or null | `null` | Stop after N documents |
|
| `max_items` | int or null | `null` | Stop after N documents |
|
||||||
| `batch_size` | int | `256` | Records per tokenization batch |
|
| `batch_size` | int | `256` | Records per tokenization batch |
|
||||||
| `packing_strategy` | str | `"simple"` | Packing strategy: `"simple"`, `"bfd"`, `"bfd_split"` |
|
| `packing_strategy` | str | `"simple"` | Packing is supported for single-output data with a `sequence` key: `"simple"`, `"bfd"`, or `"bfd_split"`. Multi-output DPO/GRPO data is not record-preserving packed output. |
|
||||||
| `max_packed_len` | int | `8192` | Maximum length of a packed bin |
|
| `max_packed_len` | int | `8192` | Maximum length of a packed bin |
|
||||||
| `truncation_mode` | str | `"keep_start"` | How to truncate sequences: `"keep_start"` or `"keep_end"` |
|
| `truncation_mode` | str | `"keep_start"` | How to truncate sequences: `"keep_start"` or `"keep_end"` |
|
||||||
|
|
||||||
@@ -266,8 +277,8 @@ When `sources` is set, `sections` is ignored.
|
|||||||
| Field | Type | Default | Description |
|
| Field | Type | Default | Description |
|
||||||
|-------|------|---------|-------------|
|
|-------|------|---------|-------------|
|
||||||
| `domain_key` | str or null | `null` | JSONL key for domain grouping |
|
| `domain_key` | str or null | `null` | JSONL key for domain grouping |
|
||||||
| `storage_format` | str | `"bin"` | `"bin"` (mmap). Reading also supports `"jsonl"` for on-the-fly tokenization |
|
| `storage_format` | str | `"bin"` | Pipeline output format. Only `"bin"` has a registered writer; `"jsonl"` is accepted by config validation but cannot be emitted by `Pipeline`. |
|
||||||
| `max_tokens_per_shard` | int | `100000000` | Flush threshold in cumulative tokens |
|
| `max_tokens_per_shard` | int | `100000000` | Flush threshold counted from each record's primary flat sequence: `sequence` for single-output data, otherwise the first flat source output |
|
||||||
| `dtype` | dict[str, str] | `{}` | Per-key tensor dtype override (e.g. `{"loss_mask": "bool"}`) |
|
| `dtype` | dict[str, str] | `{}` | Per-key tensor dtype override (e.g. `{"loss_mask": "bool"}`) |
|
||||||
| `position_ids_mode` | str | `"doc_reset"` | How to compute position_ids: `"none"`, `"doc_reset"`, `"continuous"` |
|
| `position_ids_mode` | str | `"doc_reset"` | How to compute position_ids: `"none"`, `"doc_reset"`, `"continuous"` |
|
||||||
|
|
||||||
@@ -304,11 +315,13 @@ output/
|
|||||||
meta.json
|
meta.json
|
||||||
sequence.bin
|
sequence.bin
|
||||||
loss_mask.bin
|
loss_mask.bin
|
||||||
|
position_ids.bin
|
||||||
wiki/
|
wiki/
|
||||||
shard_0000/
|
shard_0000/
|
||||||
meta.json
|
meta.json
|
||||||
sequence.bin
|
sequence.bin
|
||||||
loss_mask.bin
|
loss_mask.bin
|
||||||
|
position_ids.bin
|
||||||
```
|
```
|
||||||
|
|
||||||
### Multi-Shard (`bin`)
|
### Multi-Shard (`bin`)
|
||||||
@@ -322,13 +335,44 @@ output/
|
|||||||
meta.json
|
meta.json
|
||||||
sequence.bin
|
sequence.bin
|
||||||
loss_mask.bin
|
loss_mask.bin
|
||||||
|
position_ids.bin
|
||||||
shard_0001/
|
shard_0001/
|
||||||
meta.json
|
meta.json
|
||||||
sequence.bin
|
sequence.bin
|
||||||
loss_mask.bin
|
loss_mask.bin
|
||||||
|
position_ids.bin
|
||||||
```
|
```
|
||||||
|
|
||||||
For `bin` format, `MmapStore` discovers all shards under the domain directory via `rglob("meta.json")`. For `h5` format, `H5Store` discovers `.h5`/`.hdf5` files via recursive glob.
|
`MmapStore` discovers binary shards recursively through their `meta.json` files.
|
||||||
|
Each shard's metadata is a top-level object keyed by tensor name:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"sequence": {"shape": [123456], "dtype": "int32"},
|
||||||
|
"loss_mask": {"shape": [123456], "dtype": "bool"},
|
||||||
|
"position_ids": {"shape": [123456], "dtype": "int32"}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
An optional `offsets` array may appear for record-oriented binary data written
|
||||||
|
through `save_bin(..., record_keys=...)`; the preprocessing `BinWriter` does not
|
||||||
|
currently request those offsets.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Training Compatibility
|
||||||
|
|
||||||
|
| Training type | Supported input route |
|
||||||
|
|---------------|-----------------------|
|
||||||
|
| `seq` | Offline preprocessed `.bin`, or raw `.jsonl` eagerly transformed by `JsonlStore` using `dataset_config.json` or the built-in `messages` config |
|
||||||
|
| `sft` | Offline preprocessed `.bin`, or raw `.jsonl` through the same eager transform routes |
|
||||||
|
| `dpo` | Raw `.jsonl` through the automatic lazy DPO processor selected by `DatasetFactory` when `tokenizer_path` is supplied, or a caller-provided record store |
|
||||||
|
| `grpo` | A caller-provided, already-loaded `Store` with record-shaped `prompts`, `responses`, `masks`, and `rewards`; no automatic raw-JSONL processor is currently wired |
|
||||||
|
|
||||||
|
Offline DPO and GRPO preprocessing configs describe the intended token fields,
|
||||||
|
but their `.bin` output is not currently loadable for training. DPO binary
|
||||||
|
shards lack per-record offsets. GRPO response groups are flattened before the
|
||||||
|
binary writer and their record/group boundaries are not preserved.
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
@@ -336,15 +380,20 @@ For `bin` format, `MmapStore` discovers all shards under the domain directory vi
|
|||||||
|
|
||||||
```bash
|
```bash
|
||||||
# SFT
|
# SFT
|
||||||
python scripts/tools/preprocess.py data/sft/*.jsonl -o output/sft/ -c configs/sft_chat.json
|
python scripts/tools/preprocess.py data/sft/part-000.jsonl -o output/sft/ -c configs/sft_chat.json --batch_size 128
|
||||||
|
|
||||||
# DPO
|
# DPO
|
||||||
python scripts/tools/preprocess.py data/dpo/*.jsonl -o output/dpo/ -c configs/dpo.json --tokenizer_path params
|
python scripts/tools/preprocess.py data/dpo/part-000.jsonl -o output/dpo/ -c configs/dpo.json --tokenizer_path params
|
||||||
|
|
||||||
# GRPO
|
# GRPO
|
||||||
python scripts/tools/preprocess.py data/grpo/*.jsonl -o output/grpo/ -c configs/grpo.json
|
python scripts/tools/preprocess.py data/grpo/part-000.jsonl -o output/grpo/ -c configs/grpo.json
|
||||||
```
|
```
|
||||||
|
|
||||||
|
Inputs may be `.jsonl` files or `.json` files containing one object or a list of
|
||||||
|
objects. Each positional path must exist. A wildcard such as `data/*.jsonl`
|
||||||
|
works only when the invoking shell expands it before Click receives the
|
||||||
|
arguments; otherwise pass the files explicitly.
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## Python API
|
## Python API
|
||||||
|
|||||||
+24
-15
@@ -41,7 +41,10 @@ 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 $$
|
$$ 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 `cos_table` and `sin_table` (f32, `[max_len, dim/2]`). `forward()` returns a `(cos, sin)` tuple 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.
|
`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.
|
||||||
|
|
||||||
## Training Loop
|
## Training Loop
|
||||||
|
|
||||||
@@ -52,11 +55,12 @@ on_train_begin
|
|||||||
model.train()
|
model.train()
|
||||||
on_epoch_begin
|
on_epoch_begin
|
||||||
for batch in dataloader:
|
for batch in dataloader:
|
||||||
on_batch_begin
|
|
||||||
with executor.accumulate(model):
|
with executor.accumulate(model):
|
||||||
loss = strategy.compute_loss(batch)
|
on_batch_begin
|
||||||
context.loss = loss.item()
|
loss_output = strategy(batch)
|
||||||
stand_loss = loss / executor.grad_accum_steps
|
context.loss = loss_output["loss"].item()
|
||||||
|
context.metrics = loss_output["metrics"]
|
||||||
|
stand_loss = loss_output["loss"] / executor.grad_accum_steps
|
||||||
executor.backward(stand_loss)
|
executor.backward(stand_loss)
|
||||||
context.consumed_samples += (
|
context.consumed_samples += (
|
||||||
context.config.batch_per_device * context.world_size
|
context.config.batch_per_device * context.world_size
|
||||||
@@ -66,6 +70,7 @@ on_train_begin
|
|||||||
if executor.sync_gradients:
|
if executor.sync_gradients:
|
||||||
on_optimizer_step
|
on_optimizer_step
|
||||||
optimizer.step()
|
optimizer.step()
|
||||||
|
strategy.on_optimizer_step()
|
||||||
optimizer.zero_grad()
|
optimizer.zero_grad()
|
||||||
if scheduler:
|
if scheduler:
|
||||||
scheduler.step()
|
scheduler.step()
|
||||||
@@ -77,16 +82,18 @@ on_train_end
|
|||||||
|
|
||||||
| Hook | Fires | Default callback |
|
| Hook | Fires | Default callback |
|
||||||
|------|-------|-----------------|
|
|------|-------|-----------------|
|
||||||
| `on_train_begin` | Before training starts | `GradientCheckpointingCallback` |
|
| `on_train_begin` | Before training starts | `GradientCheckpointingCallback`, `CheckpointCallback`, `MetricCallback` |
|
||||||
| `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` |
|
| `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` |
|
||||||
| `on_batch_begin` | Every batch | — |
|
| `on_batch_begin` | Every batch | — |
|
||||||
| `on_optimizer_step` | Every accumulation window | `GradientClippingCallback`, `MetricCallback`, `ProgressBarCallback` |
|
| `on_optimizer_step` | Every accumulation window | `MetricCallback`, `ProgressBarCallback`, `GradientClippingCallback` |
|
||||||
| `on_batch_end` | Every batch | `CheckpointCallback` |
|
| `on_batch_end` | Every batch | `CheckpointCallback` |
|
||||||
| `on_epoch_end` | End of each epoch | `MetricCallback`, `ProgressBarCallback` |
|
| `on_epoch_end` | End of each epoch | `MetricCallback`, `ProgressBarCallback` |
|
||||||
| `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` |
|
| `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` |
|
||||||
| `on_train_end` | Training ends (always via finally) | `CheckpointCallback`, `MetricCallback`, `GradientCheckpointingCallback` |
|
| `on_train_end` | Training exits after `on_train_begin` completes (via `finally`) | `GradientCheckpointingCallback`, `CheckpointCallback`, `MetricCallback` |
|
||||||
|
|
||||||
Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric` (JSONL + validation, rank-0), `progress_bar` (tqdm), `gradient_clipping` (always registered; computes grad norm, clips only when `max_grad_norm` is not `None`).
|
Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric` (JSONL + validation, rank-0), `progress_bar` (tqdm, rank-0), `gradient_clipping`. The gradient-clipping callback is always registered and always calls `executor.clip_grad_norm()` with the numeric `max_grad_norm` value.
|
||||||
|
|
||||||
|
Strategies return `{"loss": Tensor, "metrics": Dict[str, float]}` when called by the trainer. Built-in metrics include the task-specific loss and, for MoE models, `moe_aux_loss` plus `moe_aux_loss_weighted`. Direct `compute_loss(batch)` calls continue to return a single loss tensor.
|
||||||
|
|
||||||
## Strategies
|
## Strategies
|
||||||
|
|
||||||
@@ -95,7 +102,7 @@ Default callbacks (in order): `gradient_checkpointing` (activation checkpointing
|
|||||||
Next-token cross-entropy with optional label smoothing:
|
Next-token cross-entropy with optional label smoothing:
|
||||||
|
|
||||||
$$
|
$$
|
||||||
L_{\text{PT}} = -\sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta)
|
L_{\text{PT}} = -\frac{1}{T}\sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta)
|
||||||
$$
|
$$
|
||||||
|
|
||||||
Keys: `input_ids`, `target_ids`. Optional: `label_smoothing`.
|
Keys: `input_ids`, `target_ids`. Optional: `label_smoothing`.
|
||||||
@@ -105,7 +112,7 @@ Keys: `input_ids`, `target_ids`. Optional: `label_smoothing`.
|
|||||||
Masked cross-entropy (`ignore_index=-100`) over response tokens:
|
Masked cross-entropy (`ignore_index=-100`) over response tokens:
|
||||||
|
|
||||||
$$
|
$$
|
||||||
L_{\text{SFT}} = -\sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta)
|
L_{\text{SFT}} = -\frac{1}{L}\sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta)
|
||||||
$$
|
$$
|
||||||
|
|
||||||
Keys: `input_ids`, `target_ids`, `loss_mask`, `position_ids`. Optional: `label_smoothing`.
|
Keys: `input_ids`, `target_ids`, `loss_mask`, `position_ids`. Optional: `label_smoothing`.
|
||||||
@@ -165,9 +172,9 @@ model factory.
|
|||||||
|------|-------|-------------|
|
|------|-------|-------------|
|
||||||
| Cosine | `CosineScheduler` | Linear warmup → cosine decay to `min_rate` |
|
| Cosine | `CosineScheduler` | Linear warmup → cosine decay to `min_rate` |
|
||||||
| SGDR | `SGDRScheduler` | Cosine annealing with warm restarts (`t_mult=2`) |
|
| SGDR | `SGDRScheduler` | Cosine annealing with warm restarts (`t_mult=2`) |
|
||||||
| WSD | `WSDScheduler` | Warmup-Stable-Decay with sqrt cooldown |
|
| WSD | `WSDScheduler` | Warmup-Stable-Decay with quadratic decay |
|
||||||
|
|
||||||
Created by `SchedulerFactory.create(schedule_type, optimizer, **kwargs)`. Valid types: `"cosine"`, `"sgdr"`, `"wsd"`. Omit to use no scheduler.
|
Created by `SchedulerFactory.create(schedule_type, optimizer, **kwargs)`. Valid types: `"cosine"`, `"sgdr"`, `"wsd"`. The training CLI always creates a scheduler and defaults `--schedule_type` to `"cosine"`.
|
||||||
|
|
||||||
## Gradient Checkpointing
|
## Gradient Checkpointing
|
||||||
|
|
||||||
@@ -185,10 +192,12 @@ Callback wraps each `DecoderBlock.forward` with `torch.utils.checkpoint.checkpoi
|
|||||||
|
|
||||||
```
|
```
|
||||||
Checkpoint(state_dict, epoch, consumed_samples, extra, meta, config)
|
Checkpoint(state_dict, epoch, consumed_samples, extra, meta, config)
|
||||||
├── save(save_dir) rank-0 only: meta.json (epoch/consumed_samples/timestamp) + config.json (model config) + model.safetensors + optional {key}.pt (optimizer.pt, scheduler.pt)
|
├── save(save_dir) meta.json (epoch/consumed_samples/timestamp) + config.json (model config) + model.safetensors + optional {key}.pt (optimizer.pt, scheduler.pt)
|
||||||
└── load(save_dir, broadcast=False) loads from local disk; set broadcast=True to broadcast metadata from rank-0
|
└── load(save_dir, broadcast=False) loads from local disk; set broadcast=True to broadcast metadata from rank-0
|
||||||
```
|
```
|
||||||
|
|
||||||
|
`Checkpoint.save()` writes whenever it is called. During training, `CheckpointCallback` uses the executor checkpoint context so only rank 0 receives a state dict and calls `save()`.
|
||||||
|
|
||||||
Optimizer/scheduler state persisted by default via `Checkpoint.extra`.
|
Optimizer/scheduler state persisted by default via `Checkpoint.extra`.
|
||||||
Model config (`context.model_config`) saved into `config.json` during training via `CheckpointCallback`.
|
Model config (`context.model_config`) saved into `config.json` during training via `CheckpointCallback`.
|
||||||
|
|
||||||
@@ -232,4 +241,4 @@ nohup python scripts/tools/train.py \
|
|||||||
|
|
||||||
Full parameter reference at [params.md](params.md).
|
Full parameter reference at [params.md](params.md).
|
||||||
|
|
||||||
> Document Update Time: 2026-07-31
|
> Document Update Time: 2026-08-02
|
||||||
|
|||||||
+63
-13
@@ -12,8 +12,12 @@ from astrai.model import AutoModel
|
|||||||
|
|
||||||
_DTYPES = ["bfloat16", "float16", "float32"]
|
_DTYPES = ["bfloat16", "float16", "float32"]
|
||||||
_CACHES = ["contiguous", "paged"]
|
_CACHES = ["contiguous", "paged"]
|
||||||
DEFAULT_CKPT = str(Path(__file__).resolve().parents[2] / "ckpt_bucket" / "kami-15bt")
|
_BACKENDS = ["cuda", "torch_native"]
|
||||||
CACHE_MAX_SEQ = 2048
|
|
||||||
|
_BACKEND_MAP = {
|
||||||
|
"cuda": ATTN_BACKEND.CUDA,
|
||||||
|
"torch_native": ATTN_BACKEND.TORCH_NATIVE,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
class BenchmarkResult:
|
class BenchmarkResult:
|
||||||
@@ -42,20 +46,22 @@ class GenerationBenchmark:
|
|||||||
device: str = "cuda",
|
device: str = "cuda",
|
||||||
dtype: torch.dtype = torch.bfloat16,
|
dtype: torch.dtype = torch.bfloat16,
|
||||||
cache_type: str = "contiguous",
|
cache_type: str = "contiguous",
|
||||||
|
backend: ATTN_BACKEND = ATTN_BACKEND.CUDA,
|
||||||
):
|
):
|
||||||
self.device = device
|
self.device = device
|
||||||
self.dtype = dtype
|
self.dtype = dtype
|
||||||
self.cache_type = cache_type
|
self.cache_type = cache_type
|
||||||
self.model = model
|
self.model = model
|
||||||
self.config = config
|
self.config = config
|
||||||
|
self.backend = backend
|
||||||
|
|
||||||
def _make_pool(self, batch_size: int) -> PagePool:
|
def _make_pool(self, batch_size: int, max_seq_len: int) -> PagePool:
|
||||||
return PagePool(
|
return PagePool(
|
||||||
n_layers=self.config.num_hidden_layers,
|
n_layers=self.config.num_hidden_layers,
|
||||||
n_kv_heads=self.config.num_key_value_heads,
|
n_kv_heads=self.config.num_key_value_heads,
|
||||||
head_dim=self.config.hidden_size // self.config.num_attention_heads,
|
head_dim=self.config.hidden_size // self.config.num_attention_heads,
|
||||||
max_batch_size=batch_size,
|
max_batch_size=batch_size,
|
||||||
max_seq_len=CACHE_MAX_SEQ,
|
max_seq_len=max_seq_len,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
dtype=self.dtype,
|
dtype=self.dtype,
|
||||||
page_size=1,
|
page_size=1,
|
||||||
@@ -82,7 +88,7 @@ class GenerationBenchmark:
|
|||||||
kv_cache = pool.bind_tasks(
|
kv_cache = pool.bind_tasks(
|
||||||
task_ids, [prompt_len] * batch_size, self.device, start_pos=0
|
task_ids, [prompt_len] * batch_size, self.device, start_pos=0
|
||||||
)
|
)
|
||||||
with torch.inference_mode(), attn_backend(ATTN_BACKEND.CUDA):
|
with torch.inference_mode(), attn_backend(self.backend):
|
||||||
self.model(
|
self.model(
|
||||||
input_ids,
|
input_ids,
|
||||||
input_mask=input_mask,
|
input_mask=input_mask,
|
||||||
@@ -105,7 +111,7 @@ class GenerationBenchmark:
|
|||||||
total_len, device=self.device
|
total_len, device=self.device
|
||||||
)
|
)
|
||||||
kv_cache = pool.bind_tasks(task_ids, [seq_len + 1] * batch_size, self.device)
|
kv_cache = pool.bind_tasks(task_ids, [seq_len + 1] * batch_size, self.device)
|
||||||
with torch.inference_mode(), attn_backend(ATTN_BACKEND.CUDA):
|
with torch.inference_mode(), attn_backend(self.backend):
|
||||||
self.model(
|
self.model(
|
||||||
input_ids,
|
input_ids,
|
||||||
input_mask=input_mask,
|
input_mask=input_mask,
|
||||||
@@ -121,6 +127,11 @@ class GenerationBenchmark:
|
|||||||
) -> BenchmarkResult:
|
) -> BenchmarkResult:
|
||||||
import time
|
import time
|
||||||
|
|
||||||
|
pool = self._make_pool(batch_size, prompt_length)
|
||||||
|
task_ids = [f"bench_prefill_{i}" for i in range(batch_size)]
|
||||||
|
for tid in task_ids:
|
||||||
|
pool.task_alloc(tid, list(range(prompt_length)))
|
||||||
|
|
||||||
input_ids = torch.randint(
|
input_ids = torch.randint(
|
||||||
0, self.config.vocab_size, (batch_size, prompt_length), device=self.device
|
0, self.config.vocab_size, (batch_size, prompt_length), device=self.device
|
||||||
)
|
)
|
||||||
@@ -129,16 +140,32 @@ class GenerationBenchmark:
|
|||||||
.unsqueeze(0)
|
.unsqueeze(0)
|
||||||
.expand(batch_size, -1)
|
.expand(batch_size, -1)
|
||||||
)
|
)
|
||||||
|
input_mask = position_ids.unsqueeze(-1) >= torch.arange(
|
||||||
|
prompt_length, device=self.device
|
||||||
|
)
|
||||||
|
kv_cache = pool.bind_tasks(
|
||||||
|
task_ids, [prompt_length] * batch_size, self.device, start_pos=0
|
||||||
|
)
|
||||||
|
|
||||||
for _ in range(3):
|
for _ in range(3):
|
||||||
with torch.inference_mode(), attn_backend(ATTN_BACKEND.CUDA):
|
with torch.inference_mode(), attn_backend(self.backend):
|
||||||
self.model(input_ids, position_ids=position_ids)
|
self.model(
|
||||||
|
input_ids,
|
||||||
|
input_mask=input_mask,
|
||||||
|
kv_cache=kv_cache,
|
||||||
|
position_ids=position_ids,
|
||||||
|
)
|
||||||
|
|
||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
t0 = time.perf_counter()
|
t0 = time.perf_counter()
|
||||||
for _ in range(num_trials):
|
for _ in range(num_trials):
|
||||||
with torch.inference_mode(), attn_backend(ATTN_BACKEND.CUDA):
|
with torch.inference_mode(), attn_backend(self.backend):
|
||||||
self.model(input_ids, position_ids=position_ids)
|
self.model(
|
||||||
|
input_ids,
|
||||||
|
input_mask=input_mask,
|
||||||
|
kv_cache=kv_cache,
|
||||||
|
position_ids=position_ids,
|
||||||
|
)
|
||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
elapsed = time.perf_counter() - t0
|
elapsed = time.perf_counter() - t0
|
||||||
tokens = batch_size * prompt_length * num_trials
|
tokens = batch_size * prompt_length * num_trials
|
||||||
@@ -161,7 +188,10 @@ class GenerationBenchmark:
|
|||||||
) -> BenchmarkResult:
|
) -> BenchmarkResult:
|
||||||
import time
|
import time
|
||||||
|
|
||||||
pool = self._make_pool(batch_size)
|
# Decode grows seq_len monotonically up to prompt + 5 + gen*num_trials
|
||||||
|
# (warmup 5 steps, then one step per trial), so size the pool to cover it.
|
||||||
|
max_seq_len = prompt_length + 5 + gen_length * num_trials
|
||||||
|
pool = self._make_pool(batch_size, max_seq_len)
|
||||||
task_ids = self._run_prefill(pool, batch_size, prompt_length)
|
task_ids = self._run_prefill(pool, batch_size, prompt_length)
|
||||||
|
|
||||||
for i in range(5):
|
for i in range(5):
|
||||||
@@ -208,6 +238,17 @@ def print_benchmark_result(result: BenchmarkResult) -> None:
|
|||||||
@click.option(
|
@click.option(
|
||||||
"--cache", type=click.Choice(_CACHES), default="contiguous", help="KV cache type."
|
"--cache", type=click.Choice(_CACHES), default="contiguous", help="KV cache type."
|
||||||
)
|
)
|
||||||
|
@click.option(
|
||||||
|
"--backend",
|
||||||
|
type=click.Choice(_BACKENDS),
|
||||||
|
default="cuda",
|
||||||
|
help="Attention backend.",
|
||||||
|
)
|
||||||
|
@click.option(
|
||||||
|
"--compare",
|
||||||
|
is_flag=True,
|
||||||
|
help="Run both backends and print side-by-side speed comparison.",
|
||||||
|
)
|
||||||
@click.option("--batch_size", type=int, default=4, help="Batch size.")
|
@click.option("--batch_size", type=int, default=4, help="Batch size.")
|
||||||
@click.option("--prompt_length", type=int, default=512, help="Prompt length.")
|
@click.option("--prompt_length", type=int, default=512, help="Prompt length.")
|
||||||
@click.option("--gen_length", type=int, default=128, help="Generation length.")
|
@click.option("--gen_length", type=int, default=128, help="Generation length.")
|
||||||
@@ -216,13 +257,16 @@ def print_benchmark_result(result: BenchmarkResult) -> None:
|
|||||||
@click.option("--decode_only", is_flag=True, help="Decode benchmark only.")
|
@click.option("--decode_only", is_flag=True, help="Decode benchmark only.")
|
||||||
@click.option(
|
@click.option(
|
||||||
"--ckpt",
|
"--ckpt",
|
||||||
default=DEFAULT_CKPT,
|
required=True,
|
||||||
|
type=click.Path(exists=True, file_okay=False, dir_okay=True, path_type=Path),
|
||||||
help="Checkpoint directory.",
|
help="Checkpoint directory.",
|
||||||
)
|
)
|
||||||
def benchmark_command(
|
def benchmark_command(
|
||||||
device: str,
|
device: str,
|
||||||
dtype: str,
|
dtype: str,
|
||||||
cache: str,
|
cache: str,
|
||||||
|
backend: str,
|
||||||
|
compare: bool,
|
||||||
batch_size: int,
|
batch_size: int,
|
||||||
prompt_length: int,
|
prompt_length: int,
|
||||||
gen_length: int,
|
gen_length: int,
|
||||||
@@ -244,15 +288,21 @@ def benchmark_command(
|
|||||||
model.to(device=device, dtype=dtype_map[dtype])
|
model.to(device=device, dtype=dtype_map[dtype])
|
||||||
model.eval()
|
model.eval()
|
||||||
|
|
||||||
|
backends = _BACKENDS if compare else [backend]
|
||||||
|
|
||||||
|
for name in backends:
|
||||||
bench = GenerationBenchmark(
|
bench = GenerationBenchmark(
|
||||||
model=model,
|
model=model,
|
||||||
config=config,
|
config=config,
|
||||||
device=device,
|
device=device,
|
||||||
dtype=dtype_map[dtype],
|
dtype=dtype_map[dtype],
|
||||||
cache_type=cache,
|
cache_type=cache,
|
||||||
|
backend=_BACKEND_MAP[name],
|
||||||
)
|
)
|
||||||
|
|
||||||
click.secho(f"Benchmark: device={device} dtype={dtype}", bold=True)
|
click.secho(
|
||||||
|
f"Benchmark: device={device} dtype={dtype} backend={name}", bold=True
|
||||||
|
)
|
||||||
|
|
||||||
if not decode_only:
|
if not decode_only:
|
||||||
result = bench.run_prefill_benchmark(
|
result = bench.run_prefill_benchmark(
|
||||||
|
|||||||
@@ -1,4 +1,6 @@
|
|||||||
import os
|
import os
|
||||||
|
import shutil
|
||||||
|
import subprocess
|
||||||
import sys
|
import sys
|
||||||
import warnings
|
import warnings
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
@@ -66,7 +68,78 @@ if _should_build():
|
|||||||
extra_link_args=[f"-Wl,-rpath,{_torch_lib}"],
|
extra_link_args=[f"-Wl,-rpath,{_torch_lib}"],
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
cmdclass["build_ext"] = BuildExtension
|
|
||||||
|
# Parallel build — each extension is an independent ninja project, so we
|
||||||
|
# can compile them concurrently. BuildExtension compiles them serially by
|
||||||
|
# default; this subclass dispatches each extension to a subprocess.
|
||||||
|
# Set BUILD_PARALLEL=N to override (default: min(n_exts, 4)).
|
||||||
|
_single_ext = os.environ.get("ASTRAI_BUILD_SINGLE_EXT", "")
|
||||||
|
|
||||||
|
class ParallelBuildExtension(BuildExtension):
|
||||||
|
def build_extensions(self):
|
||||||
|
if _single_ext:
|
||||||
|
self.extensions = [e for e in self.extensions if e.name == _single_ext]
|
||||||
|
if not self.extensions:
|
||||||
|
return
|
||||||
|
super().build_extensions()
|
||||||
|
return
|
||||||
|
|
||||||
|
n = len(self.extensions)
|
||||||
|
max_workers = int(os.environ.get("BUILD_PARALLEL", 8))
|
||||||
|
if max_workers <= 1 or n <= 1:
|
||||||
|
super().build_extensions()
|
||||||
|
return
|
||||||
|
|
||||||
|
# Each subprocess gets its own build-temp / build-lib so the
|
||||||
|
# ninja files (build.ninja, .ninja_log) never race. The built
|
||||||
|
# .so files are then collected into the parent's build_lib so the
|
||||||
|
# normal setuptools copy steps (inplace / editable wheel) work.
|
||||||
|
names = [e.name for e in self.extensions]
|
||||||
|
env = {**os.environ, "BUILD_PARALLEL": "1"}
|
||||||
|
base = os.path.join("build", "parallel")
|
||||||
|
os.makedirs(base, exist_ok=True)
|
||||||
|
procs = {}
|
||||||
|
for i in range(0, len(names), max_workers):
|
||||||
|
batch = names[i : i + max_workers]
|
||||||
|
for name in batch:
|
||||||
|
e = {**env, "ASTRAI_BUILD_SINGLE_EXT": name}
|
||||||
|
tag = name.replace(".", "_")
|
||||||
|
subdir = os.path.join(base, tag)
|
||||||
|
cmd = [
|
||||||
|
sys.executable,
|
||||||
|
__file__,
|
||||||
|
"build_ext",
|
||||||
|
"--build-temp",
|
||||||
|
os.path.join(subdir, "temp"),
|
||||||
|
"--build-lib",
|
||||||
|
os.path.join(subdir, "lib"),
|
||||||
|
]
|
||||||
|
procs[name] = subprocess.Popen(
|
||||||
|
cmd, env=e, stdout=subprocess.PIPE, stderr=subprocess.STDOUT
|
||||||
|
)
|
||||||
|
for name in batch:
|
||||||
|
out, _ = procs[name].communicate()
|
||||||
|
if procs[name].returncode != 0:
|
||||||
|
sys.stdout.write(out.decode())
|
||||||
|
raise RuntimeError(
|
||||||
|
f"parallel build failed for {name} "
|
||||||
|
f"(exit {procs[name].returncode})"
|
||||||
|
)
|
||||||
|
self._collect_extensions(
|
||||||
|
os.path.join(base, name.replace(".", "_"), "lib")
|
||||||
|
)
|
||||||
|
|
||||||
|
def _collect_extensions(self, sub_lib):
|
||||||
|
src = os.path.join(sub_lib, "astrai", "extension", "lib")
|
||||||
|
if not os.path.isdir(src):
|
||||||
|
return
|
||||||
|
dst = os.path.join(self.build_lib, "astrai", "extension", "lib")
|
||||||
|
os.makedirs(dst, exist_ok=True)
|
||||||
|
for f in os.listdir(src):
|
||||||
|
if f.endswith(".so"):
|
||||||
|
shutil.copy2(os.path.join(src, f), os.path.join(dst, f))
|
||||||
|
|
||||||
|
cmdclass["build_ext"] = ParallelBuildExtension
|
||||||
|
|
||||||
if not cmdclass:
|
if not cmdclass:
|
||||||
|
|
||||||
|
|||||||
@@ -217,14 +217,8 @@ def test_unloaded_sample_window_raises():
|
|||||||
store.sample_window(0)
|
store.sample_window(0)
|
||||||
|
|
||||||
|
|
||||||
def test_unloaded_dataset_len():
|
|
||||||
"""__len__ on a store with no data returns 0."""
|
|
||||||
store = MmapStore(window_size=64, stride=64)
|
|
||||||
assert len(store) == 0
|
|
||||||
|
|
||||||
|
|
||||||
def test_store_unloaded_len():
|
def test_store_unloaded_len():
|
||||||
"""Unloaded Store has __len__ == 0"""
|
"""Unloaded Store has __len__ == 0."""
|
||||||
store = MmapStore()
|
store = MmapStore()
|
||||||
assert len(store) == 0
|
assert len(store) == 0
|
||||||
assert store.keys == []
|
assert store.keys == []
|
||||||
@@ -498,21 +492,21 @@ def _write_json_dataset(test_dir, tokenizer_path, records, config_overrides=None
|
|||||||
return data_dir
|
return data_dir
|
||||||
|
|
||||||
|
|
||||||
def test_detect_format_jsonl_dir(base_test_env):
|
@pytest.mark.parametrize(
|
||||||
|
"use_jsonl",
|
||||||
|
[True, False],
|
||||||
|
)
|
||||||
|
def test_detect_format_data_dir(base_test_env, use_jsonl):
|
||||||
|
"""detect_format returns 'jsonl' for dirs of .jsonl or .json files."""
|
||||||
test_dir = base_test_env["test_dir"]
|
test_dir = base_test_env["test_dir"]
|
||||||
tokenizer_path = _save_test_tokenizer(test_dir, base_test_env["tokenizer"])
|
tokenizer_path = _save_test_tokenizer(test_dir, base_test_env["tokenizer"])
|
||||||
|
if use_jsonl:
|
||||||
data_dir = _write_jsonl_dataset(
|
data_dir = _write_jsonl_dataset(
|
||||||
test_dir,
|
test_dir,
|
||||||
tokenizer_path,
|
tokenizer_path,
|
||||||
[{"text": "hello world"}, {"text": "foo bar baz"}],
|
[{"text": "hello world"}, {"text": "foo bar baz"}],
|
||||||
)
|
)
|
||||||
assert detect_format(data_dir) == "jsonl"
|
else:
|
||||||
|
|
||||||
|
|
||||||
def test_detect_format_json_dir(base_test_env):
|
|
||||||
"""detect_format returns 'jsonl' for directory with .json files."""
|
|
||||||
test_dir = base_test_env["test_dir"]
|
|
||||||
tokenizer_path = _save_test_tokenizer(test_dir, base_test_env["tokenizer"])
|
|
||||||
data_dir = _write_json_dataset(
|
data_dir = _write_json_dataset(
|
||||||
test_dir,
|
test_dir,
|
||||||
tokenizer_path,
|
tokenizer_path,
|
||||||
@@ -745,32 +739,6 @@ def test_sft_jsonl_explicit_config_takes_priority(base_test_env):
|
|||||||
assert "loss_mask" in dataset.keys
|
assert "loss_mask" in dataset.keys
|
||||||
|
|
||||||
|
|
||||||
def test_jsonl_store_pipeline_config_roundtrip(base_test_env):
|
|
||||||
test_dir = base_test_env["test_dir"]
|
|
||||||
config_path = os.path.join(test_dir, "dataset_config.json")
|
|
||||||
with open(config_path, "w", encoding="utf-8") as f:
|
|
||||||
json.dump(
|
|
||||||
{
|
|
||||||
"tokenizer_path": os.path.join(test_dir, "tokenizer"),
|
|
||||||
"version": 1,
|
|
||||||
"input": {"sections": [{"field": "text", "action": "train"}]},
|
|
||||||
"mask": {"assistant": "train"},
|
|
||||||
"preprocessing": {"max_seq_len": 64},
|
|
||||||
"output": {"position_ids_mode": "doc_reset"},
|
|
||||||
},
|
|
||||||
f,
|
|
||||||
ensure_ascii=False,
|
|
||||||
indent=2,
|
|
||||||
)
|
|
||||||
|
|
||||||
with open(config_path, "r", encoding="utf-8") as f:
|
|
||||||
raw = json.load(f)
|
|
||||||
raw.pop("tokenizer_path")
|
|
||||||
config = PipelineConfig.from_dict(raw)
|
|
||||||
assert config.output.position_ids_mode == "doc_reset"
|
|
||||||
assert config.preprocessing.max_seq_len == 64
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# GRPO end-to-end: builder → JsonlStore → GRPODataset → collate_fn
|
# GRPO end-to-end: builder → JsonlStore → GRPODataset → collate_fn
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|||||||
@@ -13,18 +13,26 @@ from tests.extension.conftest import D, skip_no_kernel
|
|||||||
|
|
||||||
@skip_no_kernel
|
@skip_no_kernel
|
||||||
def test_training_forward_matches_torch(cuda_model):
|
def test_training_forward_matches_torch(cuda_model):
|
||||||
"""Training forward (kv_cache=None) should produce identical logits."""
|
"""Training forward (kv_cache=None) should produce identical logits.
|
||||||
|
|
||||||
|
CudaBackend is inference-only: it raises when kv_cache is None. Training
|
||||||
|
must use TorchNativeBackend (the default). Verify the torch path is
|
||||||
|
stable and that CudaBackend rejects the training path explicitly.
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
|
||||||
model, _ = cuda_model
|
model, _ = cuda_model
|
||||||
input_ids = torch.randint(0, 1000, (2, 16), device="cuda")
|
input_ids = torch.randint(0, 1000, (2, 16), device="cuda")
|
||||||
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
out_torch = model(input_ids)
|
out_torch = model(input_ids)
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError, match="does not support training"):
|
||||||
with attn_backend(ATTN_BACKEND.CUDA):
|
with attn_backend(ATTN_BACKEND.CUDA):
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
out_cuda = model(input_ids)
|
model(input_ids)
|
||||||
|
|
||||||
diff = (out_torch["logits"].float() - out_cuda["logits"].float()).abs().max().item()
|
assert out_torch["logits"].shape[0] == 2
|
||||||
assert diff == 0.0, f"Training forward diff {diff} should be 0"
|
|
||||||
|
|
||||||
|
|
||||||
@skip_no_kernel
|
@skip_no_kernel
|
||||||
|
|||||||
@@ -177,24 +177,6 @@ def test_scheduler_concurrent_get_stats(mock_model_and_tokenizer):
|
|||||||
assert stats["total_tasks"] >= 0
|
assert stats["total_tasks"] >= 0
|
||||||
|
|
||||||
|
|
||||||
def test_prefill_skips_fully_cached_tasks(mock_model_and_tokenizer):
|
|
||||||
"""Tasks whose entire prompt is cached skip the prefill phase."""
|
|
||||||
mock_model, mock_tokenizer = mock_model_and_tokenizer
|
|
||||||
|
|
||||||
with patch("astrai.inference.core.scheduler.AutoModel"):
|
|
||||||
with patch("astrai.inference.core.scheduler.AutoTokenizer"):
|
|
||||||
scheduler = InferenceScheduler(
|
|
||||||
model=mock_model,
|
|
||||||
tokenizer=mock_tokenizer,
|
|
||||||
max_batch_size=4,
|
|
||||||
device="cpu",
|
|
||||||
)
|
|
||||||
|
|
||||||
task_id = scheduler.add_task("short prompt", stream_callback=lambda t: None)
|
|
||||||
scheduler.stop()
|
|
||||||
assert task_id.startswith("task_")
|
|
||||||
|
|
||||||
|
|
||||||
def _make_real_scheduler(device):
|
def _make_real_scheduler(device):
|
||||||
"""Build a scheduler backed by a tiny real model for run_batch tests."""
|
"""Build a scheduler backed by a tiny real model for run_batch tests."""
|
||||||
cfg = make_rollout_config(max_position_embeddings=64)
|
cfg = make_rollout_config(max_position_embeddings=64)
|
||||||
|
|||||||
@@ -51,7 +51,7 @@ def test_task_manager_add_task():
|
|||||||
assert len(tm.waiting_queue) == 1
|
assert len(tm.waiting_queue) == 1
|
||||||
|
|
||||||
|
|
||||||
def test_task_manager_add_task_too_long_immediate_stop():
|
def test_task_manager_long_prompt_truncated_not_stopped():
|
||||||
t = _make_mock_tokenizer()
|
t = _make_mock_tokenizer()
|
||||||
t.encode.return_value = list(range(9000))
|
t.encode.return_value = list(range(9000))
|
||||||
cb_calls = []
|
cb_calls = []
|
||||||
@@ -60,6 +60,7 @@ def test_task_manager_add_task_too_long_immediate_stop():
|
|||||||
tm.add_task("long", stream_callback=lambda tok: cb_calls.append(tok))
|
tm.add_task("long", stream_callback=lambda tok: cb_calls.append(tok))
|
||||||
assert len(cb_calls) == 0
|
assert len(cb_calls) == 0
|
||||||
assert len(tm.waiting_queue) == 1
|
assert len(tm.waiting_queue) == 1
|
||||||
|
assert len(tm.waiting_queue[0].prompt_ids) == 16
|
||||||
|
|
||||||
|
|
||||||
def test_task_manager_remove_task():
|
def test_task_manager_remove_task():
|
||||||
|
|||||||
@@ -59,14 +59,15 @@ def test_find_multiple_tool_calls():
|
|||||||
assert results[1]["name"] == "f2"
|
assert results[1]["name"] == "f2"
|
||||||
|
|
||||||
|
|
||||||
def test_find_no_tool_call():
|
@pytest.mark.parametrize(
|
||||||
results = _find_tool_calls("Hello, how are you?")
|
"text,expected_count",
|
||||||
assert len(results) == 0
|
[
|
||||||
|
("Hello, how are you?", 0),
|
||||||
|
('{"not_a_tool": true}', 0),
|
||||||
def test_find_non_tool_json_skipped():
|
],
|
||||||
results = _find_tool_calls('{"not_a_tool": true}')
|
)
|
||||||
assert len(results) == 0
|
def test_find_no_tool_call(text, expected_count):
|
||||||
|
assert len(_find_tool_calls(text)) == expected_count
|
||||||
|
|
||||||
|
|
||||||
def test_find_no_arguments_field():
|
def test_find_no_arguments_field():
|
||||||
@@ -76,79 +77,6 @@ def test_find_no_arguments_field():
|
|||||||
assert results[0]["args"] == ""
|
assert results[0]["args"] == ""
|
||||||
|
|
||||||
|
|
||||||
def test_find_deeply_nested_arguments():
|
|
||||||
text = '{"name": "deep", "arguments": {"a": {"b": {"c": {"d": 4}}}}}'
|
|
||||||
results = _find_tool_calls(text)
|
|
||||||
assert len(results) == 1
|
|
||||||
assert results[0]["name"] == "deep"
|
|
||||||
assert '"d": 4' in results[0]["args"]
|
|
||||||
|
|
||||||
|
|
||||||
def test_find_arguments_with_boolean_and_null():
|
|
||||||
text = '{"name": "flags", "arguments": {"active": true, "count": 0, "nick": null}}'
|
|
||||||
results = _find_tool_calls(text)
|
|
||||||
assert len(results) == 1
|
|
||||||
assert results[0]["name"] == "flags"
|
|
||||||
assert "true" in results[0]["args"]
|
|
||||||
assert "null" in results[0]["args"]
|
|
||||||
|
|
||||||
|
|
||||||
def test_find_arguments_with_array():
|
|
||||||
text = '{"name": "add_items", "arguments": {"items": [1, 2, 3], "name": "list"}}'
|
|
||||||
results = _find_tool_calls(text)
|
|
||||||
assert len(results) == 1
|
|
||||||
assert results[0]["name"] == "add_items"
|
|
||||||
assert "[1, 2, 3]" in results[0]["args"]
|
|
||||||
|
|
||||||
|
|
||||||
def test_find_arguments_with_nested_array_of_objects():
|
|
||||||
text = '{"name": "batch", "arguments": {"rows": [{"id": 1, "val": "a"}, {"id": 2, "val": "b"}]}}'
|
|
||||||
results = _find_tool_calls(text)
|
|
||||||
assert len(results) == 1
|
|
||||||
assert '"rows"' in results[0]["args"]
|
|
||||||
assert '"id": 1' in results[0]["args"]
|
|
||||||
|
|
||||||
|
|
||||||
def test_find_arguments_as_string_not_object():
|
|
||||||
text = '{"name": "echo", "arguments": "just a string"}'
|
|
||||||
results = _find_tool_calls(text)
|
|
||||||
assert len(results) == 1
|
|
||||||
assert results[0]["name"] == "echo"
|
|
||||||
assert "just a string" in results[0]["args"]
|
|
||||||
|
|
||||||
|
|
||||||
def test_find_arguments_with_unicode():
|
|
||||||
text = (
|
|
||||||
'{"name": "translate", "arguments": {"text": "\u4f60\u597d\uff0c\u4e16\u754c"}}'
|
|
||||||
)
|
|
||||||
results = _find_tool_calls(text)
|
|
||||||
assert len(results) == 1
|
|
||||||
assert results[0]["name"] == "translate"
|
|
||||||
|
|
||||||
|
|
||||||
def test_find_arguments_with_escaped_quotes():
|
|
||||||
text = '{"name": "format", "arguments": {"template": "he said \\"hello\\""}}'
|
|
||||||
results = _find_tool_calls(text)
|
|
||||||
assert len(results) == 1
|
|
||||||
assert 'he said \\"hello\\"' in results[0]["args"]
|
|
||||||
|
|
||||||
|
|
||||||
def test_find_arguments_with_braces_in_string():
|
|
||||||
text = '{"name": "eval", "arguments": {"code": "function(x) { return x + 1; }"}}'
|
|
||||||
results = _find_tool_calls(text)
|
|
||||||
assert len(results) == 1
|
|
||||||
assert results[0]["name"] == "eval"
|
|
||||||
assert "function(x) { return x + 1; }" in results[0]["args"]
|
|
||||||
|
|
||||||
|
|
||||||
def test_find_many_properties():
|
|
||||||
args = ",".join(f'"{chr(97 + i % 26)}" : {i}' for i in range(20))
|
|
||||||
text = '{"name": "many", "arguments": {' + args + "}}"
|
|
||||||
results = _find_tool_calls(text)
|
|
||||||
assert len(results) == 1
|
|
||||||
assert results[0]["name"] == "many"
|
|
||||||
|
|
||||||
|
|
||||||
def test_find_empty_arguments():
|
def test_find_empty_arguments():
|
||||||
results = _find_tool_calls('{"name": "ping", "arguments": {}}')
|
results = _find_tool_calls('{"name": "ping", "arguments": {}}')
|
||||||
assert len(results) == 1
|
assert len(results) == 1
|
||||||
@@ -164,6 +92,62 @@ def test_find_extracts_correct_arg_start_position():
|
|||||||
assert json_str == text
|
assert json_str == text
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"text,expected_name,arg_substr",
|
||||||
|
[
|
||||||
|
(
|
||||||
|
'{"name": "deep", "arguments": {"a": {"b": {"c": {"d": 4}}}}}',
|
||||||
|
"deep",
|
||||||
|
'"d": 4',
|
||||||
|
),
|
||||||
|
(
|
||||||
|
'{"name": "flags", "arguments": {"active": true, "count": 0, "nick": null}}',
|
||||||
|
"flags",
|
||||||
|
"null",
|
||||||
|
),
|
||||||
|
(
|
||||||
|
'{"name": "add_items", "arguments": {"items": [1, 2, 3], "name": "list"}}',
|
||||||
|
"add_items",
|
||||||
|
"[1, 2, 3]",
|
||||||
|
),
|
||||||
|
(
|
||||||
|
'{"name": "batch", "arguments": {"rows": [{"id": 1, "val": "a"}, {"id": 2, "val": "b"}]}}',
|
||||||
|
"batch",
|
||||||
|
'"id": 1',
|
||||||
|
),
|
||||||
|
('{"name": "echo", "arguments": "just a string"}', "echo", "just a string"),
|
||||||
|
(
|
||||||
|
'{"name": "translate", "arguments": {"text": "\u4f60\u597d\uff0c\u4e16\u754c"}}',
|
||||||
|
"translate",
|
||||||
|
"\u4f60\u597d",
|
||||||
|
),
|
||||||
|
(
|
||||||
|
'{"name": "format", "arguments": {"template": "he said \\"hello\\""}}',
|
||||||
|
"format",
|
||||||
|
'he said \\"hello\\"',
|
||||||
|
),
|
||||||
|
(
|
||||||
|
'{"name": "eval", "arguments": {"code": "function(x) { return x + 1; }"}}',
|
||||||
|
"eval",
|
||||||
|
"function(x) { return x + 1; }",
|
||||||
|
),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_find_arguments_variants(text, expected_name, arg_substr):
|
||||||
|
results = _find_tool_calls(text)
|
||||||
|
assert len(results) == 1
|
||||||
|
assert results[0]["name"] == expected_name
|
||||||
|
assert arg_substr in results[0]["args"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_find_many_properties():
|
||||||
|
args = ",".join(f'"{chr(97 + i % 26)}" : {i}' for i in range(20))
|
||||||
|
text = '{"name": "many", "arguments": {' + args + "}}"
|
||||||
|
results = _find_tool_calls(text)
|
||||||
|
assert len(results) == 1
|
||||||
|
assert results[0]["name"] == "many"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
"text,expected_name,expected_complete",
|
"text,expected_name,expected_complete",
|
||||||
[
|
[
|
||||||
@@ -340,30 +324,21 @@ def test_streaming_multiple_tool_calls_incremental():
|
|||||||
assert "f2" in names
|
assert "f2" in names
|
||||||
|
|
||||||
|
|
||||||
def test_streaming_deeply_nested_args():
|
@pytest.mark.parametrize(
|
||||||
parser = SimpleJsonToolParser()
|
"text,arg_substr",
|
||||||
text = '{"name": "deep", "arguments": {"a": {"b": {"c": 42}}}}'
|
[
|
||||||
_, args_chunks = _simulate_streaming(parser, text)
|
('{"name": "deep", "arguments": {"a": {"b": {"c": 42}}}}', '"c": 42'),
|
||||||
joined = "".join(args_chunks)
|
(
|
||||||
assert '"c": 42' in joined
|
'{"name": "translate", "arguments": {"text": "\u4f60\u597d\uff0c\u4e16\u754c"}}',
|
||||||
|
"\u4f60\u597d",
|
||||||
|
),
|
||||||
def test_streaming_args_with_unicode():
|
('{"name": "add", "arguments": {"items": [1, 2, 3]}}', "[1, 2, 3]"),
|
||||||
parser = SimpleJsonToolParser()
|
],
|
||||||
text = (
|
|
||||||
'{"name": "translate", "arguments": {"text": "\u4f60\u597d\uff0c\u4e16\u754c"}}'
|
|
||||||
)
|
)
|
||||||
_, args_chunks = _simulate_streaming(parser, text)
|
def test_streaming_args_variants(text, arg_substr):
|
||||||
joined = "".join(args_chunks)
|
|
||||||
assert "\u4f60\u597d" in joined
|
|
||||||
|
|
||||||
|
|
||||||
def test_streaming_args_with_array():
|
|
||||||
parser = SimpleJsonToolParser()
|
parser = SimpleJsonToolParser()
|
||||||
text = '{"name": "add", "arguments": {"items": [1, 2, 3]}}'
|
|
||||||
_, args_chunks = _simulate_streaming(parser, text)
|
_, args_chunks = _simulate_streaming(parser, text)
|
||||||
joined = "".join(args_chunks)
|
assert arg_substr in "".join(args_chunks)
|
||||||
assert "[1, 2, 3]" in joined
|
|
||||||
|
|
||||||
|
|
||||||
def test_streaming_empty_arguments():
|
def test_streaming_empty_arguments():
|
||||||
@@ -514,7 +489,6 @@ def test_feed_then_parse_complete_same_instance():
|
|||||||
('{ "name" : "f"}', True),
|
('{ "name" : "f"}', True),
|
||||||
('{"other": 1}', False),
|
('{"other": 1}', False),
|
||||||
('prefix {"name": "f", "args": {}}', True),
|
('prefix {"name": "f", "args": {}}', True),
|
||||||
('{"name": "f"}', True), # match at start
|
|
||||||
(' {"name": "f"}', True),
|
(' {"name": "f"}', True),
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
@@ -526,10 +500,6 @@ def test_pattern_regex(text, matches):
|
|||||||
assert result is None
|
assert result is None
|
||||||
|
|
||||||
|
|
||||||
def test_pattern_name_at_start():
|
|
||||||
assert _TOOL_CALL_HEAD_RE.match('{"name": "f"}')
|
|
||||||
|
|
||||||
|
|
||||||
def test_factory_register_and_create():
|
def test_factory_register_and_create():
|
||||||
parser = ToolParserFactory.create("simple_json")
|
parser = ToolParserFactory.create("simple_json")
|
||||||
assert isinstance(parser, BaseToolParser)
|
assert isinstance(parser, BaseToolParser)
|
||||||
@@ -547,10 +517,6 @@ def test_factory_list_registered():
|
|||||||
assert "simple_json" in ToolParserFactory.list_registered()
|
assert "simple_json" in ToolParserFactory.list_registered()
|
||||||
|
|
||||||
|
|
||||||
def test_factory_create_with_no_extra_kwargs():
|
|
||||||
assert isinstance(ToolParserFactory.create("simple_json"), BaseToolParser)
|
|
||||||
|
|
||||||
|
|
||||||
def test_factory_create_with_tools_only():
|
def test_factory_create_with_tools_only():
|
||||||
tools = [
|
tools = [
|
||||||
{
|
{
|
||||||
@@ -563,13 +529,6 @@ def test_factory_create_with_tools_only():
|
|||||||
assert parser.tool_choice == "auto"
|
assert parser.tool_choice == "auto"
|
||||||
|
|
||||||
|
|
||||||
def test_feed_accepts_token_ids_and_ignores_them():
|
|
||||||
parser = SimpleJsonToolParser()
|
|
||||||
text = '{"name": "get_weather", "arguments": {"city": "Beijing"}}'
|
|
||||||
deltas_with = parser.feed(text, current_token_ids=[123, 456], delta_token_ids=[456])
|
|
||||||
assert len(deltas_with) > 0
|
|
||||||
|
|
||||||
|
|
||||||
def test_feed_token_ids_do_not_affect_parsing():
|
def test_feed_token_ids_do_not_affect_parsing():
|
||||||
parser_no_ids = SimpleJsonToolParser()
|
parser_no_ids = SimpleJsonToolParser()
|
||||||
parser_with_ids = SimpleJsonToolParser()
|
parser_with_ids = SimpleJsonToolParser()
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
import pytest
|
import pytest
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from astrai.model.components.mlp import MLP, DeepSeekMoE
|
||||||
from astrai.model.transformer import AutoRegressiveLM
|
from astrai.model.transformer import AutoRegressiveLM
|
||||||
from tests.helpers import TINY_CONFIG
|
from tests.helpers import TINY_CONFIG
|
||||||
|
|
||||||
@@ -32,6 +33,59 @@ CONFIGS = [
|
|||||||
},
|
},
|
||||||
id="gqa_moe",
|
id="gqa_moe",
|
||||||
),
|
),
|
||||||
|
pytest.param(
|
||||||
|
{
|
||||||
|
**TINY_CONFIG,
|
||||||
|
"attn_type": "gqa",
|
||||||
|
"ffn_type": "moe",
|
||||||
|
"n_routed_experts": 4,
|
||||||
|
"n_shared_experts": 1,
|
||||||
|
"n_activated_experts": 2,
|
||||||
|
"topk_method": "greedy",
|
||||||
|
"mlp_only_layers": [0],
|
||||||
|
},
|
||||||
|
id="gqa_moe_dense_first",
|
||||||
|
),
|
||||||
|
pytest.param(
|
||||||
|
{
|
||||||
|
**TINY_CONFIG,
|
||||||
|
"attn_type": "gqa",
|
||||||
|
"ffn_type": "moe",
|
||||||
|
"n_routed_experts": 4,
|
||||||
|
"n_shared_experts": 1,
|
||||||
|
"n_activated_experts": 2,
|
||||||
|
"topk_method": "greedy",
|
||||||
|
"decoder_sparse_step": 2,
|
||||||
|
},
|
||||||
|
id="gqa_moe_sparse_step",
|
||||||
|
),
|
||||||
|
pytest.param(
|
||||||
|
{
|
||||||
|
**TINY_CONFIG,
|
||||||
|
"attn_type": "gqa",
|
||||||
|
"ffn_type": "moe",
|
||||||
|
"n_routed_experts": 4,
|
||||||
|
"n_shared_experts": 1,
|
||||||
|
"n_activated_experts": 2,
|
||||||
|
"topk_method": "greedy",
|
||||||
|
"norm_topk_prob": True,
|
||||||
|
},
|
||||||
|
id="gqa_moe_norm_topk",
|
||||||
|
),
|
||||||
|
pytest.param(
|
||||||
|
{
|
||||||
|
**TINY_CONFIG,
|
||||||
|
"attn_type": "gqa",
|
||||||
|
"ffn_type": "moe",
|
||||||
|
"n_routed_experts": 4,
|
||||||
|
"n_shared_experts": 1,
|
||||||
|
"n_activated_experts": 2,
|
||||||
|
"topk_method": "greedy",
|
||||||
|
"moe_intermediate_size": 24,
|
||||||
|
"shared_expert_intermediate_size": 20,
|
||||||
|
},
|
||||||
|
id="gqa_moe_custom_intermediate",
|
||||||
|
),
|
||||||
pytest.param(
|
pytest.param(
|
||||||
{
|
{
|
||||||
**TINY_CONFIG,
|
**TINY_CONFIG,
|
||||||
@@ -105,3 +159,168 @@ def test_model_forward_with_padding(config_kwargs, device):
|
|||||||
|
|
||||||
assert output["logits"].shape == (batch_size, seq_len, config.vocab_size)
|
assert output["logits"].shape == (batch_size, seq_len, config.vocab_size)
|
||||||
assert not torch.isnan(output["logits"]).any()
|
assert not torch.isnan(output["logits"]).any()
|
||||||
|
|
||||||
|
|
||||||
|
def test_moe_per_layer_ffn_resolution():
|
||||||
|
"""Verify that mlp_only_layers and decoder_sparse_step resolve FFN types correctly."""
|
||||||
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
|
|
||||||
|
# mlp_only_layers: first layer dense, rest MoE
|
||||||
|
config = AutoRegressiveLMConfig(
|
||||||
|
**{
|
||||||
|
**TINY_CONFIG,
|
||||||
|
"attn_type": "gqa",
|
||||||
|
"ffn_type": "moe",
|
||||||
|
"n_routed_experts": 4,
|
||||||
|
"n_shared_experts": 1,
|
||||||
|
"n_activated_experts": 2,
|
||||||
|
"mlp_only_layers": [0],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
model = AutoRegressiveLM(config)
|
||||||
|
assert isinstance(model.layers[0].mlp, MLP)
|
||||||
|
assert not isinstance(model.layers[0].mlp, DeepSeekMoE)
|
||||||
|
assert isinstance(model.layers[1].mlp, DeepSeekMoE)
|
||||||
|
|
||||||
|
# decoder_sparse_step=2: every other layer is MoE
|
||||||
|
config2 = AutoRegressiveLMConfig(
|
||||||
|
**{
|
||||||
|
**TINY_CONFIG,
|
||||||
|
"attn_type": "gqa",
|
||||||
|
"ffn_type": "moe",
|
||||||
|
"n_routed_experts": 4,
|
||||||
|
"n_shared_experts": 1,
|
||||||
|
"n_activated_experts": 2,
|
||||||
|
"decoder_sparse_step": 2,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
model2 = AutoRegressiveLM(config2)
|
||||||
|
# layer 0 (id=0): (0+1)%2=1 != 0 -> MLP
|
||||||
|
assert isinstance(model2.layers[0].mlp, MLP)
|
||||||
|
assert not isinstance(model2.layers[0].mlp, DeepSeekMoE)
|
||||||
|
# layer 1 (id=1): (1+1)%2=0 -> MoE
|
||||||
|
assert isinstance(model2.layers[1].mlp, DeepSeekMoE)
|
||||||
|
|
||||||
|
# decoder_sparse_step=1 (default): all layers MoE
|
||||||
|
config3 = AutoRegressiveLMConfig(
|
||||||
|
**{
|
||||||
|
**TINY_CONFIG,
|
||||||
|
"attn_type": "gqa",
|
||||||
|
"ffn_type": "moe",
|
||||||
|
"n_routed_experts": 4,
|
||||||
|
"n_shared_experts": 1,
|
||||||
|
"n_activated_experts": 2,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
model3 = AutoRegressiveLM(config3)
|
||||||
|
for layer in model3.layers:
|
||||||
|
assert isinstance(layer.mlp, DeepSeekMoE)
|
||||||
|
|
||||||
|
|
||||||
|
def test_moe_custom_intermediate_shape():
|
||||||
|
"""Verify MoE uses custom intermediate sizes when specified."""
|
||||||
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
|
|
||||||
|
config = AutoRegressiveLMConfig(
|
||||||
|
**{
|
||||||
|
**TINY_CONFIG,
|
||||||
|
"attn_type": "gqa",
|
||||||
|
"ffn_type": "moe",
|
||||||
|
"n_routed_experts": 4,
|
||||||
|
"n_shared_experts": 1,
|
||||||
|
"n_activated_experts": 2,
|
||||||
|
"moe_intermediate_size": 24,
|
||||||
|
"shared_expert_intermediate_size": 20,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
model = AutoRegressiveLM(config)
|
||||||
|
moe_layer = model.layers[0].mlp
|
||||||
|
assert isinstance(moe_layer, DeepSeekMoE)
|
||||||
|
# routed experts use moe_intermediate_size
|
||||||
|
for expert in moe_layer.routed_experts:
|
||||||
|
assert expert.up.weight.shape[0] == 24
|
||||||
|
assert expert.gate.weight.shape[0] == 24
|
||||||
|
assert expert.down.weight.shape[1] == 24
|
||||||
|
# shared experts use shared_expert_intermediate_size
|
||||||
|
for expert in moe_layer.shared_experts:
|
||||||
|
assert expert.up.weight.shape[0] == 20
|
||||||
|
assert expert.gate.weight.shape[0] == 20
|
||||||
|
assert expert.down.weight.shape[1] == 20
|
||||||
|
|
||||||
|
|
||||||
|
def test_moe_defaults_preserve_normalized_routing():
|
||||||
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
|
|
||||||
|
config = AutoRegressiveLMConfig(
|
||||||
|
**TINY_CONFIG,
|
||||||
|
ffn_type="moe",
|
||||||
|
n_routed_experts=4,
|
||||||
|
n_shared_experts=1,
|
||||||
|
n_activated_experts=2,
|
||||||
|
topk_method="greedy",
|
||||||
|
)
|
||||||
|
model = AutoRegressiveLM(config)
|
||||||
|
|
||||||
|
assert config.norm_topk_prob is True
|
||||||
|
assert model.layers[0].mlp.norm_topk_prob is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_moe_aux_loss_only_emitted_during_training():
|
||||||
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
|
|
||||||
|
config = AutoRegressiveLMConfig(
|
||||||
|
**TINY_CONFIG,
|
||||||
|
ffn_type="moe",
|
||||||
|
n_routed_experts=4,
|
||||||
|
n_shared_experts=1,
|
||||||
|
n_activated_experts=2,
|
||||||
|
topk_method="greedy",
|
||||||
|
)
|
||||||
|
model = AutoRegressiveLM(config)
|
||||||
|
input_ids = torch.randint(0, config.vocab_size, (2, 8))
|
||||||
|
|
||||||
|
outputs = model(input_ids)
|
||||||
|
assert outputs["aux_loss"].ndim == 0
|
||||||
|
assert outputs["aux_loss"].requires_grad
|
||||||
|
assert torch.isfinite(outputs["aux_loss"])
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
outputs = model(input_ids)
|
||||||
|
assert "aux_loss" not in outputs
|
||||||
|
|
||||||
|
model.eval()
|
||||||
|
outputs = model(input_ids)
|
||||||
|
assert "aux_loss" not in outputs
|
||||||
|
|
||||||
|
|
||||||
|
def test_moe_component_forward_returns_ffn_output():
|
||||||
|
from astrai.model.components.mlp import DeepSeekMoE
|
||||||
|
|
||||||
|
moe = DeepSeekMoE(
|
||||||
|
dim=8,
|
||||||
|
dim_ffn=16,
|
||||||
|
n_routed_experts=4,
|
||||||
|
n_shared_experts=1,
|
||||||
|
n_activated_experts=2,
|
||||||
|
)
|
||||||
|
|
||||||
|
output = moe(torch.randn(2, 8, 8))
|
||||||
|
|
||||||
|
assert output["hidden_states"].shape == (2, 8, 8)
|
||||||
|
assert output["aux_loss"] is not None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("decoder_sparse_step", [0, -1])
|
||||||
|
def test_moe_rejects_invalid_decoder_sparse_step(decoder_sparse_step):
|
||||||
|
from pydantic import ValidationError
|
||||||
|
|
||||||
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
|
|
||||||
|
with pytest.raises(ValidationError, match="decoder_sparse_step must be at least 1"):
|
||||||
|
AutoRegressiveLMConfig(
|
||||||
|
**TINY_CONFIG,
|
||||||
|
ffn_type="moe",
|
||||||
|
n_routed_experts=4,
|
||||||
|
n_activated_experts=2,
|
||||||
|
decoder_sparse_step=decoder_sparse_step,
|
||||||
|
)
|
||||||
|
|||||||
@@ -98,18 +98,9 @@ def test_loralinear_merge():
|
|||||||
assert lora._merged
|
assert lora._merged
|
||||||
assert not hasattr(lora, "lora_A")
|
assert not hasattr(lora, "lora_A")
|
||||||
|
|
||||||
|
# merge is guarded by _merged — a second call is a no-op.
|
||||||
def test_loralinear_merge_is_idempotent():
|
|
||||||
base = Linear(4, 4)
|
|
||||||
with torch.no_grad():
|
|
||||||
base.weight.zero_()
|
|
||||||
|
|
||||||
lora = LoRALinear(base, r=2, alpha=2)
|
|
||||||
with torch.no_grad():
|
|
||||||
lora.lora_B.fill_(1.0)
|
|
||||||
|
|
||||||
lora.merge()
|
|
||||||
lora.merge()
|
lora.merge()
|
||||||
|
assert lora._merged
|
||||||
|
|
||||||
|
|
||||||
def test_inject_lora_default_target():
|
def test_inject_lora_default_target():
|
||||||
|
|||||||
@@ -71,23 +71,14 @@ def test_grpo_loss_backward(grpo_strategy):
|
|||||||
assert has_grad
|
assert has_grad
|
||||||
|
|
||||||
|
|
||||||
def test_grpo_ref_model_not_updated(grpo_strategy):
|
@pytest.mark.parametrize("model_name", ["ref_model", "old_model"])
|
||||||
"""Backward should not populate gradients on ref_model."""
|
def test_grpo_frozen_models_not_updated(grpo_strategy, model_name):
|
||||||
|
"""Backward should not populate gradients on ref_model or old_model."""
|
||||||
strategy, device = grpo_strategy
|
strategy, device = grpo_strategy
|
||||||
batch = _make_batch(device=device)
|
batch = _make_batch(device=device)
|
||||||
loss = strategy.compute_loss(batch)
|
loss = strategy.compute_loss(batch)
|
||||||
loss.backward()
|
loss.backward()
|
||||||
for p in strategy.ref_model.parameters():
|
for p in getattr(strategy, model_name).parameters():
|
||||||
assert p.grad is None
|
|
||||||
|
|
||||||
|
|
||||||
def test_grpo_old_model_not_updated(grpo_strategy):
|
|
||||||
"""Backward should not populate gradients on old_model."""
|
|
||||||
strategy, device = grpo_strategy
|
|
||||||
batch = _make_batch(device=device)
|
|
||||||
loss = strategy.compute_loss(batch)
|
|
||||||
loss.backward()
|
|
||||||
for p in strategy.old_model.parameters():
|
|
||||||
assert p.grad is None
|
assert p.grad is None
|
||||||
|
|
||||||
|
|
||||||
@@ -133,45 +124,3 @@ def test_grpo_sync_old_model(grpo_strategy):
|
|||||||
if k in old_sd_after
|
if k in old_sd_after
|
||||||
)
|
)
|
||||||
assert matches
|
assert matches
|
||||||
|
|
||||||
|
|
||||||
def test_grpo_partial_mask(grpo_strategy):
|
|
||||||
"""Only the first half of response tokens are valid."""
|
|
||||||
strategy, device = grpo_strategy
|
|
||||||
batch = _make_batch(device=device)
|
|
||||||
B, G, R = batch["masks"].shape
|
|
||||||
half = R // 2
|
|
||||||
batch["masks"][:, :, half:] = 0.0
|
|
||||||
loss = strategy.compute_loss(batch)
|
|
||||||
assert torch.isfinite(loss).item()
|
|
||||||
|
|
||||||
|
|
||||||
def test_grpo_clipping_effect(grpo_strategy):
|
|
||||||
"""After diverging policy from ref, ratio should be clipped to [1-eps, 1+eps]
|
|
||||||
on the surrogate. Verify loss is finite and non-zero for distinct rewards."""
|
|
||||||
strategy, device = grpo_strategy
|
|
||||||
with torch.no_grad():
|
|
||||||
for p in strategy.model.parameters():
|
|
||||||
p.add_(0.3)
|
|
||||||
batch = _make_batch(device=device)
|
|
||||||
loss = strategy.compute_loss(batch)
|
|
||||||
assert torch.isfinite(loss).item()
|
|
||||||
assert loss.abs().item() > 1e-4
|
|
||||||
|
|
||||||
|
|
||||||
def test_grpo_no_reduction_param():
|
|
||||||
"""GRPOStrategy.__init__ must not accept ``reduction`` (removed)."""
|
|
||||||
import inspect
|
|
||||||
|
|
||||||
sig = inspect.signature(GRPOStrategy.__init__)
|
|
||||||
assert "reduction" not in sig.parameters
|
|
||||||
|
|
||||||
|
|
||||||
def test_grpo_shapes_3d_batch(grpo_strategy):
|
|
||||||
"""Verify compute_loss handles non-square prompt/response lengths."""
|
|
||||||
strategy, device = grpo_strategy
|
|
||||||
batch = _make_batch(
|
|
||||||
batch_size=3, group_size=4, prompt_len=10, response_len=8, device=device
|
|
||||||
)
|
|
||||||
loss = strategy.compute_loss(batch)
|
|
||||||
assert torch.isfinite(loss).item()
|
|
||||||
|
|||||||
@@ -0,0 +1,105 @@
|
|||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.model.transformer import AutoRegressiveLM
|
||||||
|
from astrai.trainer.strategy import BaseStrategy, SEQStrategy
|
||||||
|
from astrai.trainer.train_callback import MetricCallback
|
||||||
|
from tests.helpers import make_tiny_config
|
||||||
|
|
||||||
|
|
||||||
|
def test_seq_strategy_combines_and_reports_moe_aux_loss(device):
|
||||||
|
config = make_tiny_config(
|
||||||
|
ffn_type="moe",
|
||||||
|
n_routed_experts=4,
|
||||||
|
n_shared_experts=1,
|
||||||
|
n_activated_experts=2,
|
||||||
|
topk_method="greedy",
|
||||||
|
)
|
||||||
|
model = AutoRegressiveLM(config).to(device=device)
|
||||||
|
strategy = SEQStrategy(model, device, moe_aux_loss_coef=0.25)
|
||||||
|
batch = {
|
||||||
|
"input_ids": torch.randint(0, config.vocab_size, (2, 8), device=device),
|
||||||
|
"target_ids": torch.randint(0, config.vocab_size, (2, 8), device=device),
|
||||||
|
}
|
||||||
|
|
||||||
|
output = strategy(batch)
|
||||||
|
legacy_loss = strategy.compute_loss(batch)
|
||||||
|
|
||||||
|
assert isinstance(legacy_loss, torch.Tensor)
|
||||||
|
assert set(output["metrics"]) == {
|
||||||
|
"loss",
|
||||||
|
"task_loss",
|
||||||
|
"moe_aux_loss",
|
||||||
|
"moe_aux_loss_weighted",
|
||||||
|
}
|
||||||
|
assert output["loss"].item() == pytest.approx(
|
||||||
|
output["metrics"]["task_loss"] + output["metrics"]["moe_aux_loss_weighted"],
|
||||||
|
)
|
||||||
|
assert output["metrics"]["moe_aux_loss_weighted"] == pytest.approx(
|
||||||
|
0.25 * output["metrics"]["moe_aux_loss"],
|
||||||
|
)
|
||||||
|
assert output["loss"].requires_grad
|
||||||
|
assert all(isinstance(metric, float) for metric in output["metrics"].values())
|
||||||
|
|
||||||
|
|
||||||
|
def test_metric_callback_includes_dynamic_strategy_metrics(tmp_path):
|
||||||
|
callback = MetricCallback(
|
||||||
|
ckpt_dir=tmp_path,
|
||||||
|
save_interval=1,
|
||||||
|
metrics=["loss", "lr"],
|
||||||
|
)
|
||||||
|
context = SimpleNamespace(
|
||||||
|
metrics={"task_loss": 2.0, "moe_aux_loss": 1.0},
|
||||||
|
loss=2.01,
|
||||||
|
optimizer=SimpleNamespace(param_groups=[{"lr": 1e-3}]),
|
||||||
|
val_loss=None,
|
||||||
|
grad_norm=None,
|
||||||
|
grad_snr_tracker=None,
|
||||||
|
world_size=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
metrics = callback._metrics(context, callback.metrics)
|
||||||
|
|
||||||
|
assert metrics == {
|
||||||
|
"loss": 2.01,
|
||||||
|
"lr": 1e-3,
|
||||||
|
"task_loss": 2.0,
|
||||||
|
"moe_aux_loss": 1.0,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_metric_callback_only_computes_requested_metrics(tmp_path):
|
||||||
|
def fail_metric(context):
|
||||||
|
_ = context
|
||||||
|
raise AssertionError("unrequested metric was computed")
|
||||||
|
|
||||||
|
callback = MetricCallback(
|
||||||
|
ckpt_dir=tmp_path,
|
||||||
|
save_interval=1,
|
||||||
|
metrics=["loss"],
|
||||||
|
)
|
||||||
|
callback._metric_funcs["grad_snr"] = fail_metric
|
||||||
|
context = SimpleNamespace(
|
||||||
|
metrics={},
|
||||||
|
loss=2.0,
|
||||||
|
world_size=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
metrics = callback._metrics(context, callback.metrics)
|
||||||
|
|
||||||
|
assert metrics == {"loss": 2.0}
|
||||||
|
|
||||||
|
|
||||||
|
def test_legacy_strategy_tensor_loss_is_normalized():
|
||||||
|
class LegacyStrategy(BaseStrategy):
|
||||||
|
def compute_loss(self, batch):
|
||||||
|
return torch.tensor(2.0, requires_grad=True)
|
||||||
|
|
||||||
|
strategy = LegacyStrategy(torch.nn.Linear(1, 1), "cpu")
|
||||||
|
|
||||||
|
output = strategy({})
|
||||||
|
|
||||||
|
assert output["loss"].item() == 2.0
|
||||||
|
assert output["metrics"]["loss"] == 2.0
|
||||||
@@ -96,12 +96,10 @@ def test_factory_registers_online_aliases():
|
|||||||
assert StrategyFactory.get_component_class("online_dpo") is DPOStrategy
|
assert StrategyFactory.get_component_class("online_dpo") is DPOStrategy
|
||||||
|
|
||||||
|
|
||||||
def test_grpo_supports_online(device):
|
@pytest.mark.parametrize("make_fn", ["_make_grpo", "_make_dpo"])
|
||||||
assert _make_grpo(device).supports_online() is True
|
def test_online_strategies_support_online(device, make_fn):
|
||||||
|
maker = {"_make_grpo": _make_grpo, "_make_dpo": _make_dpo}[make_fn]
|
||||||
|
assert maker(device).supports_online() is True
|
||||||
def test_dpo_supports_online(device):
|
|
||||||
assert _make_dpo(device).supports_online() is True
|
|
||||||
|
|
||||||
|
|
||||||
def test_base_strategy_prepare_from_rollout_raises_by_default(device):
|
def test_base_strategy_prepare_from_rollout_raises_by_default(device):
|
||||||
@@ -161,21 +159,21 @@ def test_call_without_runner_falls_back_to_compute_loss_grpo(device):
|
|||||||
"masks": torch.ones(2, 4, 6, device=device),
|
"masks": torch.ones(2, 4, 6, device=device),
|
||||||
"rewards": torch.randn(2, 4, device=device),
|
"rewards": torch.randn(2, 4, device=device),
|
||||||
}
|
}
|
||||||
loss = strat(batch)
|
loss = strat(batch)["loss"]
|
||||||
assert torch.isfinite(loss).item()
|
assert torch.isfinite(loss).item()
|
||||||
|
|
||||||
|
|
||||||
def test_call_with_runner_returns_finite_loss_grpo(device):
|
def test_call_with_runner_returns_finite_loss_grpo(device):
|
||||||
strat = _make_grpo(device)
|
strat = _make_grpo(device)
|
||||||
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))
|
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))
|
||||||
loss = strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
loss = strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})["loss"]
|
||||||
assert torch.isfinite(loss).item()
|
assert torch.isfinite(loss).item()
|
||||||
|
|
||||||
|
|
||||||
def test_call_with_runner_returns_finite_loss_dpo(device):
|
def test_call_with_runner_returns_finite_loss_dpo(device):
|
||||||
strat = _make_dpo(device)
|
strat = _make_dpo(device)
|
||||||
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))
|
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))
|
||||||
loss = strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
loss = strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})["loss"]
|
||||||
assert torch.isfinite(loss).item()
|
assert torch.isfinite(loss).item()
|
||||||
|
|
||||||
|
|
||||||
@@ -267,21 +265,10 @@ def test_step_called_when_sync_gradients_true(device):
|
|||||||
assert runner.step_calls == 1
|
assert runner.step_calls == 1
|
||||||
|
|
||||||
|
|
||||||
def test_loss_is_differentiable_grpo(device):
|
|
||||||
strat = _make_grpo(device)
|
|
||||||
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))
|
|
||||||
loss = strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
|
||||||
loss.backward()
|
|
||||||
has_grad = any(
|
|
||||||
p.grad is not None and p.grad.abs().sum() > 0 for p in strat.model.parameters()
|
|
||||||
)
|
|
||||||
assert has_grad
|
|
||||||
|
|
||||||
|
|
||||||
def test_loss_is_differentiable_dpo(device):
|
def test_loss_is_differentiable_dpo(device):
|
||||||
strat = _make_dpo(device)
|
strat = _make_dpo(device)
|
||||||
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))
|
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))
|
||||||
loss = strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
loss = strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})["loss"]
|
||||||
loss.backward()
|
loss.backward()
|
||||||
has_grad = any(
|
has_grad = any(
|
||||||
p.grad is not None and p.grad.abs().sum() > 0 for p in strat.model.parameters()
|
p.grad is not None and p.grad.abs().sum() > 0 for p in strat.model.parameters()
|
||||||
@@ -289,21 +276,10 @@ def test_loss_is_differentiable_dpo(device):
|
|||||||
assert has_grad
|
assert has_grad
|
||||||
|
|
||||||
|
|
||||||
def test_ref_and_old_model_not_updated_by_backward_grpo(device):
|
|
||||||
strat = _make_grpo(device)
|
|
||||||
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))
|
|
||||||
loss = strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
|
||||||
loss.backward()
|
|
||||||
for p in strat.ref_model.parameters():
|
|
||||||
assert p.grad is None
|
|
||||||
for p in strat.old_model.parameters():
|
|
||||||
assert p.grad is None
|
|
||||||
|
|
||||||
|
|
||||||
def test_ref_model_not_updated_by_backward_dpo(device):
|
def test_ref_model_not_updated_by_backward_dpo(device):
|
||||||
strat = _make_dpo(device)
|
strat = _make_dpo(device)
|
||||||
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))
|
strat.set_rollout_runner(_RecordingRunner(_make_rollout_result(device=device)))
|
||||||
loss = strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})
|
loss = strat({"input_ids": torch.randint(3, 200, (2, 4), device=device)})["loss"]
|
||||||
loss.backward()
|
loss.backward()
|
||||||
for p in strat.ref_model.parameters():
|
for p in strat.ref_model.parameters():
|
||||||
assert p.grad is None
|
assert p.grad is None
|
||||||
|
|||||||
@@ -81,6 +81,9 @@ def test_rollout_result_inherits_raw_rollout_fields():
|
|||||||
assert r.prompts.shape == (2, 4)
|
assert r.prompts.shape == (2, 4)
|
||||||
assert r.responses.shape == (2, 3, 5)
|
assert r.responses.shape == (2, 3, 5)
|
||||||
assert r.prompt_mask.shape == (2, 4)
|
assert r.prompt_mask.shape == (2, 4)
|
||||||
|
# RolloutResult must carry every RawRollout field.
|
||||||
|
raw_fields = {f for f in RawRollout.__dataclass_fields__}
|
||||||
|
assert raw_fields.issubset(set(RolloutResult.__dataclass_fields__))
|
||||||
|
|
||||||
|
|
||||||
def test_base_reward_model_is_abstract():
|
def test_base_reward_model_is_abstract():
|
||||||
|
|||||||
@@ -131,18 +131,10 @@ def test_register_signal_handlers():
|
|||||||
assert ctx.stop_requested
|
assert ctx.stop_requested
|
||||||
|
|
||||||
|
|
||||||
def test_sigterm_triggers_checkpoint_save(base_test_env):
|
|
||||||
exitcode = _spawn_train_and_signal(base_test_env["test_dir"], signal.SIGTERM)
|
|
||||||
assert exitcode == 0, f"Training process exited with code {exitcode} (expected 0)"
|
|
||||||
|
|
||||||
meta = load_checkpoint_meta(base_test_env["test_dir"])
|
|
||||||
assert "consumed_samples" in meta
|
|
||||||
assert meta["consumed_samples"] >= 0
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.slow
|
@pytest.mark.slow
|
||||||
def test_sigint_triggers_checkpoint_save(base_test_env):
|
@pytest.mark.parametrize("sig", [signal.SIGTERM, signal.SIGINT])
|
||||||
exitcode = _spawn_train_and_signal(base_test_env["test_dir"], signal.SIGINT)
|
def test_signal_triggers_checkpoint_save(base_test_env, sig):
|
||||||
|
exitcode = _spawn_train_and_signal(base_test_env["test_dir"], sig)
|
||||||
assert exitcode == 0, f"Training process exited with code {exitcode} (expected 0)"
|
assert exitcode == 0, f"Training process exited with code {exitcode} (expected 0)"
|
||||||
|
|
||||||
meta = load_checkpoint_meta(base_test_env["test_dir"])
|
meta = load_checkpoint_meta(base_test_env["test_dir"])
|
||||||
|
|||||||
@@ -1,13 +1,13 @@
|
|||||||
|
import pytest
|
||||||
|
|
||||||
from astrai.trainer import Trainer
|
from astrai.trainer import Trainer
|
||||||
|
|
||||||
# train_config_factory is injected via fixture
|
|
||||||
|
|
||||||
|
def test_training_runs_with_various_batch_sizes(
|
||||||
def test_different_batch_sizes(base_test_env, random_dataset, train_config_factory):
|
base_test_env, random_dataset, train_config_factory
|
||||||
"""Test training with different batch sizes"""
|
):
|
||||||
batch_sizes = [1, 2, 4, 8]
|
"""Training should complete for a range of batch sizes without error."""
|
||||||
|
for batch_per_device in [1, 2, 4]:
|
||||||
for batch_per_device in batch_sizes:
|
|
||||||
train_config = train_config_factory(
|
train_config = train_config_factory(
|
||||||
model_fn=lambda: base_test_env["model"],
|
model_fn=lambda: base_test_env["model"],
|
||||||
dataset=random_dataset,
|
dataset=random_dataset,
|
||||||
@@ -15,48 +15,22 @@ def test_different_batch_sizes(base_test_env, random_dataset, train_config_facto
|
|||||||
device=base_test_env["device"],
|
device=base_test_env["device"],
|
||||||
batch_per_device=batch_per_device,
|
batch_per_device=batch_per_device,
|
||||||
)
|
)
|
||||||
|
trainer = Trainer(train_config)
|
||||||
assert train_config.batch_per_device == batch_per_device
|
trainer.train()
|
||||||
|
|
||||||
|
|
||||||
def test_gradient_accumulation(base_test_env, random_dataset, train_config_factory):
|
@pytest.mark.slow
|
||||||
"""Test training with different gradient accumulation steps"""
|
def test_gradient_accumulation_runs(
|
||||||
grad_accum_steps_list = [1, 2, 4]
|
base_test_env, random_dataset, train_config_factory
|
||||||
|
):
|
||||||
for grad_accum_steps in grad_accum_steps_list:
|
"""Training with gradient accumulation should complete."""
|
||||||
train_config = train_config_factory(
|
train_config = train_config_factory(
|
||||||
model_fn=lambda: base_test_env["model"],
|
model_fn=lambda: base_test_env["model"],
|
||||||
dataset=random_dataset,
|
dataset=random_dataset,
|
||||||
test_dir=base_test_env["test_dir"],
|
test_dir=base_test_env["test_dir"],
|
||||||
device=base_test_env["device"],
|
device=base_test_env["device"],
|
||||||
batch_per_device=2,
|
batch_per_device=2,
|
||||||
grad_accum_steps=grad_accum_steps,
|
grad_accum_steps=4,
|
||||||
)
|
)
|
||||||
|
|
||||||
trainer = Trainer(train_config)
|
trainer = Trainer(train_config)
|
||||||
trainer.train()
|
trainer.train()
|
||||||
|
|
||||||
assert train_config.grad_accum_steps == grad_accum_steps
|
|
||||||
|
|
||||||
|
|
||||||
def test_memory_efficient_training(base_test_env, random_dataset, train_config_factory):
|
|
||||||
"""Test training with memory-efficient configurations"""
|
|
||||||
# Test with smaller batch sizes and gradient checkpointing
|
|
||||||
small_batch_configs = [
|
|
||||||
{"batch_per_device": 1, "grad_accum_steps": 8},
|
|
||||||
{"batch_per_device": 2, "grad_accum_steps": 4},
|
|
||||||
{"batch_per_device": 4, "grad_accum_steps": 2},
|
|
||||||
]
|
|
||||||
|
|
||||||
for config in small_batch_configs:
|
|
||||||
train_config = train_config_factory(
|
|
||||||
model_fn=lambda: base_test_env["model"],
|
|
||||||
dataset=random_dataset,
|
|
||||||
test_dir=base_test_env["test_dir"],
|
|
||||||
device=base_test_env["device"],
|
|
||||||
batch_per_device=config["batch_per_device"],
|
|
||||||
grad_accum_steps=config["grad_accum_steps"],
|
|
||||||
)
|
|
||||||
|
|
||||||
assert train_config.grad_accum_steps == config["grad_accum_steps"]
|
|
||||||
assert train_config.batch_per_device == config["batch_per_device"]
|
|
||||||
|
|||||||
Reference in New Issue
Block a user