27 Commits
Author SHA1 Message Date
ViperEkura 998b443aa3 refactor: dedupe fp8 meta state into per-operand rings
- collapse FP8TensorMeta's 12 slots + 6 copy-paste methods into three _ScaleRing objects (hist/idx/scale/initialized + update/seed)
- skip meta allocation entirely on the DynamicScaling path (zero rings, scales measured inline)
- drop write-only FP8State._last_device and unused E4M3_MAX alias
2026-08-24 21:19:45 +08:00
ViperEkura cebdd45d3a perf: batch crosswise stage loads in fp8 gemm
- load_operand_tile ColMajor path issued one LDG then immediately scattered 16 byte-granular shared stores, so every store waited on the preceding global load; the runs of one row group now batch into registers first (v[kPasses]) and scatter after, overlapping the LDG latencies
- hoist pass-invariant predicates: the alignment check folds to one uniform (base | ld) & 15 test since r0 is always a multiple of 16, and rows_full leaves the per-pass condition; the contract tail zero-fills without global traffic
- RowMajor path hoists the row bound and the (invariant) chunk-alignment check out of the per-chunk loop
- measured (cuda events, old/new interleaved): crosswise bwd gemms +5-9%, RowMajor and fwd NT within noise; model step unchanged (in the 563-617 ms band)
- file passed through clang-format with the new .clang-format config
2026-08-24 19:49:10 +08:00
ViperEkura 7da1439c9e feat: static fp8 weights and bias with fused epilogue
- linear_forward_fp8 accepts pre-quantized w8 (matching fmt) and skips the weight quantize; amax_w returns 0 on that path since no bf16 values are seen
- bias is now fused into the GEMM epilogue for both dtypes, replacing the separate torch-level add (one elementwise kernel per linear removed)
- FP8Params.bias becomes void* with a new bias_scale slot: null scale = raw bf16 bias, non-null = fp8 storage dequantized in the epilogue after the operand scaling and before any output quantization
- ops/fp8.py relaxes the w dtype check to bf16-or-fp8 and passes bias_scale through
- regression test covers w8/b8, w8/bf16-bias and the amax_w = 0 contract vs an explicit quantization reference
2026-08-24 19:25:23 +08:00
ViperEkura 29e5f571af fix: own fp8 linear backward via autograd Function
- backward used to read the global fp8 flag at loss.backward() time, so calling it outside fp8_autocast silently fell back to bf16 mm (953 ms cublas per step, 49.9% of the model step)
- _LinearFp8(torch.autograd.Function) now owns the fwd/bwd pair: forward captures fmt/recipe/meta on ctx inside the autocast region, backward reads only ctx (scales from the meta rings, masks from ctx.needs_input_grad), so backward is fp8 wherever it runs
- register the aten::linear impl on AutogradCUDA (replaces torch's generated linear formula that calls aten::linear_backward into the bf16 fallback) and keep the CUDA key for inference_mode
- drop the aten::linear_backward override and fp8_linear_backward (dead paths)
- regression test asserts the fp8 backward fires outside the autocast region and grads match the bf16 reference by direction/norm (E5M2 noise)
- model step (0.67B, CE loss, batch 4x1024): backward GEMMs 953 -> 618 ms (1.54x), full step ~1.2x
2026-08-24 19:03:01 +08:00
ViperEkura 74e694921c perf: speed up fp8 gemm tiles and scheduling
- K tile 32->64 (new default): fewer barriers, more MMA per stage; generalize tile_at swizzle and load_operand_tile accordingly
- 64x128 small-M CTA for m<=64 (2x at 64x4096x4096)
- L2 rasterization for crosswise-A layouts (+6..21%)
- micro-bench: NT 4096^3 +35%; linear fwd 1.24-1.76x, bwd 1.71-2.27x vs bf16
- add csrc/tests/fp8_test.cu (single MMA demo + GEMM layouts x K-tiles vs CPU reference)
2026-08-24 18:29:55 +08:00
ViperEkura d5067af064 refactor: harden param PODs and CUTLASS-style fp8 layout tags
- NSDMI null/-1 defaults for AttentionParams/FP8Params pointer+flag members: partially packed structs can no longer hold garbage non-null pointers that gate optional paths (root cause class of the paged test bug); still aggregates, still trivially copyable
- move per-lane ldmatrix wrappers (ldsm_x2/x4) from fp8/gemm.cuh to common/mma.cuh as ldmatrix_x2_lane/x4_lane, next to the single-address variants
- DEVICE_FORCEINLINE macro in common/mma.cuh (matches layout_policies.cuh, internal linkage)
- frag_addr now delegates to tile_at: the swizzle math has one source
- operand layouts as CUTLASS-style RowMajor/ColMajor tags threaded from launch_fp8_gemm through the kernel to load_operand_tile; B's operand view via transpose_layout_t; call sites read <Fmt, false, RowMajor, ColMajor> instead of <Fmt, false, false, true>
2026-08-24 15:27:52 +08:00
ViperEkura f6db546578 fix: zero-init AttentionParams in pure C tests
- paged decode test left new_k_ptr/new_v_ptr as stack garbage; PagedKV::decode_addr then took the new-KV path on wild pointers (illegal access or wrong last-token K/V)
- value-init the POD (= {}) at every construction site
2026-08-24 14:56:19 +08:00
ViperEkura 31ca357c61 refactor: namespace csrc kernels and extract common helpers
- attention family -> astrai::attention; fp8 family -> astrai::fp8
- new common/reduce.cuh (warp/group reductions, atomic_max_float)
- new common/cp_async.cuh (predicated cp_async_16, commit/wait group)
- move MAX_SPLITS into attention/common.h; delete warp_utils.cuh
- .cu bindings and pure C tests open family namespaces via using
2026-08-24 14:49:55 +08:00
ViperEkura 34471252ab perf: pipeline fp8 gemm fragment loads and pack bf16 epilogue
- software-pipeline A-fragment ldmatrix: row mt+1 loads hide behind row mt MMAs
- bf16 epilogue packs two columns into one bfloat162 store (half the stores)
- fp8 vs bf16 linear: fwd 1.15x@512, 1.5x@2048, 2.8x@4096; bwd up to 2.7x, peak 35 TFLOPS
2026-08-23 21:32:20 +08:00
ViperEkura aa08479285 perf: widen fp8 gemm tile and load fragments with ldmatrix
- 128x128 CTA of 8 warps x 64x32 warp tiles: 16 mma.sync per warp per K-segment (was 8)
- ldmatrix.x4/x2 with per-lane swizzled addresses replaces 36 scalar LDS per warp-tile step
- __launch_bounds__(256, 2) caps registers at 124 so two CTAs fit per SM
- fp8 linear vs bf16 cuBLAS: fwd 1.07x->2.65x, bwd 1.41x->2.65x by size, peak 31-33 TFLOPS
2026-08-23 21:13:09 +08:00
ViperEkura 4b10d3ca37 perf: vectorize fp8 quantize and swizzle gemm smem 2026-08-23 20:31:44 +08:00
ViperEkura 2bc4d2b8a8 refactor: unify kernel module loading and packaging
- loader.py: lazy/cached import; is_available defers the actual load; get_module raises on unavailable
- ops/{attention,rotary,fp8}: use get_module instead of touching private _modules or their own _mod() cache
- package-data: ship astrai.extension.lib *.so in built wheels (non-editable installs previously lost every kernel)
2026-08-23 15:57:04 +08:00
ViperEkura 4244df2785 perf: pure FP8 fwd/bwd and lean non-transposed GEMM
- drop the fused kernel; forward/backward are quantize + a pre-quantized GEMM
- rename module fp8_mm -> fp8_ops (mm.cu -> ops.cu)
- kernels/launchers fp8_gemm_kernel / launch_fp8_gemm; drop PqTraits/gather_trans/pack_fp8x4_vector
- remove the in-kernel transposed-operand branches (TransA/TransB)
- backward: quantize g once (amax_g here), explicit fp8 transposes, fast non-transposed GEMMs (dX = g@w^T, dW = g^T@x^T)
- each pass uses a single FP8 format (E4M3 fwd / E5M2 bwd)
2026-08-23 15:38:30 +08:00
ViperEkura a29bdfae46 refactor: rework attention backend resolution
- explicit attn_backend() context wins over ASTR_BACKEND env
- polymorphic available()/supports_call() replace isinstance dispatch
- cache singleton backend instances to avoid hot-path allocation
- training (fwd=None) resolves cuda > flash > torch by capability
- flash dense supports mask-free calls only; masked training falls back to torch
2026-08-23 14:47:02 +08:00
ViperEkura 10fec8dca1 docs: fix stale docs and align with code
- update cuda_kernels layout, arch flags, and add FP8 section
- fix install docs: kernels auto-build when nvcc + CUDA detected
- mark ignored OpenAI request params and complete KVCache fields
- add docker docs to indexes and astrai.optim to module overview
- refresh document update timestamps
2026-08-23 14:23:28 +08:00
ViperEkura 75304d084d refactor: unify CUDA skip guards in tests
- add skip_no_fp8 (CUDA + fp8_mm kernel + cc 8.9+) to tests/conftest.py
- use skip_no_cuda / skip_no_kernel / skip_no_fp8 directly in test modules
- drop _GPU alias and tests.extension.conftest re-exports
- remove unused imports (Union in hf_adapter, make_grpo_config in data conftest)
2026-08-22 21:08:06 +08:00
ViperEkura 16a55bb474 refactor: reorganize CUDA kernels into per-family directories
- move attention kernels to csrc/kernels/attention/ and rotary to rotary/
- add shared common/mma.cuh (mma_sync, ldmatrix) and device.cuh (sm checks)
- split fp8_mm into three-layer fp8/common.h, gemm.cuh, mm.cu
- fix fused FP8 GEMM ldmatrix lane indexing to fix OOB shared reads
- update extension ops, loader, and kernel tests
2026-08-22 20:40:31 +08:00
ViperEkura cb21af38ba feat: unify Docker serving configuration in YAML
- Add server.py --config serve.yaml; explicit CLI flags override YAML
- Add scripts/serve.sh and serve_runtime.py for the Compose lifecycle
- Template server/cpu ports and param mounts in docker-compose.yml
- Document schema in docs/developer/docker-serving.md and params guide
- Add tests for runtime parsing and server CLI merge logic
2026-08-21 23:16:45 +08:00
ViperEkura dcc96de12a test: refactor tests and fix inference edge cases
- Convert protocol and MoE test classes to plain functions
- Add real server/engine integration and generate_async tests
- Isolate test model per test and use pytest tmp_path
- Reset FastAPI engine state after inference tests
- Fix generate_async StopIteration handling on Python 3.12
- Fix HF adapter MoE dense/shared and Gemma qk_norm mapping
- Correct dev dependency httpx2 to httpx
2026-08-21 22:59:51 +08:00
ViperEkura 7d27f3e078 feat: load HuggingFace checkpoints via key/config conversion
- Add astrai.serialization.hf_adapter mapping LLaMA-style HF keys to AstrAI names (input_layernorm, gate_proj, MoE experts/shared_experts) with config aliases for dense and MoE (Mixtral/DeepSeek-V3) layouts; reject biased projections, mismatched head_dim and MLA
- Give AutoModel.from_pretrained weights_format=auto|astrai|hf with auto-detection; read sharded safetensors via model.safetensors.index.json
- Adapt preloaded weights/config in train_context and benchmark CLI
2026-08-20 11:34:59 +08:00
ViperEkura 84753d3e08 refactor: deduplicate and restructure test suite
- extract preprocessing config factories into tests/data/factories.py
- keep conftest.py fixtures-only; stop importing builders from it
- promote temp_dir fixture to root conftest for cross-directory reuse
- unify duplicate BPE tokenizer builders into build_test_tokenizer
- merge grpo/dpo online e2e tests into one parametrized integration test
- extract engine mock factory and shared model batch builders
- drop local tempfile usage in favor of shared fixtures

No behavior change: 519 tests pass.
2026-08-20 01:53:28 +08:00
ViperEkura 53a7149577 feat: add distributed rank to logs 2026-08-20 01:37:00 +08:00
ViperEkura c79d34eee1 refactor: simplify training and inference interfaces
- avoid constructing model_fn more than once when reading config
- keep inference package exports focused on public entry points
- rename extra strategy arguments to strategy_kwargs
2026-08-19 20:55:13 +08:00
ViperEkura 398e8a3ea3 refactor: deduplicate low-risk code paths 2026-08-19 16:17:40 +08:00
ViperEkura f252af495c refactor: remove dead code and deduplicate scheduler setup 2026-08-19 14:58:37 +08:00
ViperEkura 00c2c80c8f feat: unify Docker training configuration in YAML 2026-08-19 14:04:35 +08:00
ViperEkura a6c6a54ace feat: read host training vars from YAML infra section
- scripts/train.sh load_infra() parses the top-level infra: section of TRAIN_CONFIG_FILE
- exports TRAIN_JOB_NAME/DATA/MODEL/CHECKPOINT_DIR/TRAIN_GPU_COUNT/CUDA_VISIBLE_DEVICES
- infra overrides .env.train via compose interpolation precedence; keys absent fall back
- train.yaml is now the single per-job config: host mounts, GPU filter, and hyperparameters
- requires host python3 with PyYAML when TRAIN_CONFIG_FILE is set; errors fail fast
- docs: docker-training.md documents the infra overrides and precedence
2026-08-19 01:27:57 +08:00
109 changed files with 6138 additions and 3183 deletions
+3 -2
View File
@@ -39,11 +39,12 @@ ruff format . # re-format after fix
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 the system temp directory
> (`$TMPDIR` on Linux/macOS, `%TEMP%` on Windows). Clean them manually if needed.
### 4. (Optional) Full pre-commit check script
If you have Git Bash available:
If you have `bash` available (Git Bash on Windows works too):
```bash
bash scripts/pre_commit.sh
+9 -3
View File
@@ -51,7 +51,7 @@ AstrAI is an end-to-end Transformer framework for building, training, evaluating
| **Data** | Declarative JSON preprocessing, configurable masking and packing, binary/JSONL storage, and streaming datasets |
| **Inference** | Continuous batching, paged KV cache, radix prefix caching, streaming generation, and Torch/CUDA/FlashAttention backends |
| **Serving** | FastAPI server with OpenAI and Anthropic chat completion protocols, including SSE streaming and tool calls |
| **Evaluation** | Perplexity, MMLU, HumanEval, IFEval, IFD, and ROUGE evaluation tools |
| **Evaluation** | Perplexity, MMLU, HumanEval, IFEval, IFD, ROUGE, and weight-analysis evaluation tools |
| **Extensibility** | Factory and registry architecture for models, datasets, training strategies, callbacks, kernels, and protocol components |
### Getting Started
@@ -65,8 +65,9 @@ AstrAI requires Python 3.12+ and pins PyTorch exactly to `2.11.0`. Training, `sc
```bash
git clone https://github.com/ViperEkura/AstrAI.git
cd AstrAI
pip install -e . # pure PyTorch (no CUDA kernels)
# CSRC_KERNELS=true pip install -e . --no-build-isolation # optional: fused CUDA kernels
pip install -e . # kernels auto-build when nvcc + CUDA are detected
# CSRC_KERNELS=false pip install -e . # skip kernels (pure PyTorch)
# CSRC_KERNELS=true pip install -e . --no-build-isolation # force the fused CUDA kernel build
# pip install -e ".[dev]" # dev dependencies (pytest, ruff)
```
@@ -191,6 +192,9 @@ docker compose up -d
# Docker Compose CPU server profile (CUDA-only generation scripts/demos are unavailable)
docker compose --profile cpu up -d
# YAML-driven serving (see serve.yaml; up/run/down/logs/status...)
bash scripts/serve.sh up
```
> **Note**: `--gpus all` is required for CUDA support. Without it, `torch.cuda.is_available()` will return `False`.
@@ -236,6 +240,8 @@ See [Inference Guide](docs/guides/inference.md) for SSE streaming format, error
| [Data Flow](./docs/developer/dataflow.md) | Data pipeline, storage backends & dataset architecture |
| [Internals](./docs/developer/internals.md) | Training internals: loss formulas, callback lifecycle, KV cache |
| [CUDA Kernels](./docs/developer/cuda_kernels.md) | Custom CUDA attention kernels & benchmarks |
| [Docker Serving](./docs/developer/docker-serving.md) | YAML-driven containerized serving (`serve.yaml`, `serve.sh`) |
| [Docker Training](./docs/developer/docker-training.md) | YAML-driven containerized training (`train.yaml`, `train.sh`) |
### Contributing
+3 -8
View File
@@ -17,14 +17,9 @@ from astrai.dataset import (
StoreFactory,
)
from astrai.factory import BaseFactory
from astrai.inference import (
InferenceEngine,
ProtocolHandler,
SamplingPipeline,
get_app,
run_server,
sample,
)
from astrai.inference import InferenceEngine, get_app, run_server, sample
from astrai.inference.network import ProtocolHandler
from astrai.inference.runtime.sample import SamplingPipeline
from astrai.logging import setup_logging
from astrai.model import (
AutoModel,
+14 -14
View File
@@ -11,10 +11,10 @@ from torch.utils.data import Dataset
from astrai.config.base import BaseConfig
from astrai.model.components.lora import LoRAConfig
_TRAIN_TYPES = frozenset({"seq", "sft", "dpo", "grpo", "online_grpo", "online_dpo"})
_PARALLEL_MODES = frozenset({"none", "ddp", "fsdp"})
_BACKENDS = frozenset({"nccl", "gloo"})
_START_METHODS = frozenset({"spawn", "fork", "forkserver"})
TRAIN_TYPES = frozenset({"seq", "sft", "dpo", "grpo", "online_grpo", "online_dpo"})
PARALLEL_MODES = frozenset({"none", "ddp", "fsdp"})
BACKENDS = frozenset({"nccl", "gloo"})
START_METHODS = frozenset({"spawn", "fork", "forkserver"})
_COMPILE_MODES = frozenset({"default", "reduce-overhead", "max-autotune"})
@@ -70,7 +70,7 @@ class TrainConfig(BaseConfig):
rollout_max_tokens (int): Maximum generated tokens per response in rollout. Defaults to 1024.
reward_model_fn (Optional[Callable]): Factory for reward model, required for online RL strategies. Defaults to None.
executor_kwargs (Dict[str, Any]): Extra kwargs passed to ExecutorFactory.create(). Defaults to {}.
extra_kwargs (Dict[str, Any]): Other arguments. Defaults to {}.
strategy_kwargs (Dict[str, Any]): Extra strategy arguments. Defaults to {}.
"""
model_fn: Callable[[], nn.Module]
@@ -125,35 +125,35 @@ class TrainConfig(BaseConfig):
reward_model_fn: Optional[Callable] = None
executor_kwargs: Dict[str, Any] = field(default_factory=dict)
extra_kwargs: Dict[str, Any] = field(default_factory=dict)
strategy_kwargs: Dict[str, Any] = field(default_factory=dict)
@field_validator("strategy")
def _validate_strategy(cls, v: str) -> str:
if v not in _TRAIN_TYPES:
if v not in TRAIN_TYPES:
raise ValueError(
f"strategy must be one of {sorted(_TRAIN_TYPES)}, got {v!r}"
f"strategy must be one of {sorted(TRAIN_TYPES)}, got {v!r}"
)
return v
@field_validator("parallel_mode")
def _validate_parallel_mode(cls, v: str) -> str:
if v not in _PARALLEL_MODES:
if v not in PARALLEL_MODES:
raise ValueError(
f"parallel_mode must be one of {sorted(_PARALLEL_MODES)}, got {v!r}"
f"parallel_mode must be one of {sorted(PARALLEL_MODES)}, got {v!r}"
)
return v
@field_validator("backend")
def _validate_backend(cls, v: str) -> str:
if v not in _BACKENDS:
raise ValueError(f"backend must be one of {sorted(_BACKENDS)}, got {v!r}")
if v not in BACKENDS:
raise ValueError(f"backend must be one of {sorted(BACKENDS)}, got {v!r}")
return v
@field_validator("start_method")
def _validate_start_method(cls, v: str) -> str:
if v not in _START_METHODS:
if v not in START_METHODS:
raise ValueError(
f"start_method must be one of {sorted(_START_METHODS)}, got {v!r}"
f"start_method must be one of {sorted(START_METHODS)}, got {v!r}"
)
return v
+4 -7
View File
@@ -383,10 +383,10 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
transform = _build_jsonl_transform(load_path, tokenizer_path)
if transform is None:
raise FileNotFoundError(
f"JSONL dataset config not found. Expected "
f"dataset_config.json alongside *.jsonl files, pass "
f"tokenizer_path= for the built-in messages config, or "
f"use processor= for lazy on-the-fly tokenisation."
"JSONL dataset config not found. Expected "
"dataset_config.json alongside *.jsonl files, pass "
"tokenizer_path= for the built-in messages config, or "
"use processor= for lazy on-the-fly tokenisation."
)
store.load(load_path, transform=transform, **kwargs)
else:
@@ -492,9 +492,6 @@ class DPODataset(BaseDataset):
required_keys = ["chosen", "rejected", "chosen_mask", "rejected_mask"]
def make_processor(self, tokenizer, max_len: int):
return partial(dpo_processor, tokenizer=tokenizer, max_len=max_len)
def __getitem__(self, index: int) -> Dict[str, Tensor]:
return {
"chosen": self.store.fetch_record(index, "chosen").to(dtype=torch.long),
+1 -1
View File
@@ -217,7 +217,7 @@ class Store(ABC):
"""
if self._window_size <= 0:
raise RuntimeError("sample_window() requires window_size > 0 (stream mode)")
if self._window_size <= 0 or self._length <= self._window_size:
if self._length <= self._window_size:
raise IndexError(
f"Data too short for window: token_count={self._length}, "
f"window_size={self._window_size}"
+213 -97
View File
@@ -21,9 +21,20 @@ Usage — mirroring ``torch.nn.attention.sdpa_kernel``:
...
Thread-safe via ``contextvars`` — each scheduler thread gets its own
active backend. ``get_backend()`` returns the active one, falling back
to a process-wide default (cuda > flash > torch, overridable via
``ASTR_BACKEND``).
active backend. Backend resolution follows a strict precedence:
1. explicit ``attn_backend(...)`` context (wins over everything),
2. the process-wide ``ASTR_BACKEND`` environment override,
3. an implicit default picked from the available backends
(cuda > flash > torch).
Capability is polymorphic: every backend declares ``available()``
(machine-level) and ``supports_call(...)`` (per-call), so adding a new
backend requires no changes to the resolution logic. Training calls
(``fwd=None``, no KV cache) resolve through the same priority list: the
CUDA cache kernels cannot run without a cache, so they fall back to
flash (when it can handle the call — mask-free/causal only) and finally
to the reference ``TorchNativeBackend``.
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
(blhd). The backend returns ``[batch, seq_len, n_heads * head_dim]``.
@@ -32,11 +43,12 @@ Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
import contextvars
import enum
import functools
import logging
import os
import threading
from abc import ABC, abstractmethod
from contextlib import contextmanager
from typing import TYPE_CHECKING, Optional, Union
from typing import TYPE_CHECKING, Dict, Optional, Tuple, Union
import torch
import torch.nn.functional as F
@@ -57,8 +69,9 @@ except Exception:
if TYPE_CHECKING:
from astrai.inference.cache import KVCache
logger = logging.getLogger(__name__)
_default_backend: Optional["AttentionBackend"] = None
_default_backend_lock = threading.Lock()
_env_backend_name: Optional[str] = None
_env_backend: Optional["AttentionBackend"] = None
@@ -66,6 +79,10 @@ _current_backend: contextvars.ContextVar[Optional["AttentionBackend"]] = (
contextvars.ContextVar("attn_backend", default=None)
)
# Backends are stateless — one canonical instance per class, created lazily
# and reused everywhere (resolution, fallback, context managers).
_singletons: Dict[type, "AttentionBackend"] = {}
@functools.lru_cache(maxsize=1)
def flash_attn_available() -> bool:
@@ -102,58 +119,40 @@ class ATTN_BACKEND(enum.Enum):
FLASH = "flash"
def _priority_backends() -> list["AttentionBackend"]:
"""Available backends in priority order: cuda -> flash -> torch."""
backends: list[AttentionBackend] = []
if is_available("attn_paged_decode") and is_available("attn_paged_prefill"):
backends.append(CudaBackend())
if flash_attn_available():
backends.append(FlashAttnBackend())
backends.append(TorchNativeBackend())
return backends
def _instance(backend_cls: type) -> "AttentionBackend":
"""Return the canonical singleton instance for a backend class.
def _backend_supports(
backend: "AttentionBackend",
q: Tensor,
kv_cache: Optional["KVCache"],
attn_mask: Optional[Tensor],
is_causal: bool,
fwd: Optional[str],
) -> bool:
"""Whether ``backend`` can run this attention call.
The CUDA kernels are bf16-only, support head_dim in 32/64/128/256, and
need a KV cache (decode/prefill); everything else falls back to torch.
Backends hold no per-instance state, so a single cached instance is
safe and avoids per-call allocation on the attention hot path.
"""
if isinstance(backend, CudaBackend):
return (
fwd in ("prefill", "decode")
and kv_cache is not None
and q.ndim == 3
and q.dtype == torch.bfloat16
and q.size(-1) in (32, 64, 128, 256)
and is_available(f"attn_paged_{fwd}")
)
if isinstance(backend, FlashAttnBackend):
if not flash_attn_available():
return False
if q.dtype not in (torch.float16, torch.bfloat16):
return False
if fwd is not None:
return q.ndim == 3 and hasattr(_flash_attn, "flash_attn_varlen_func")
if attn_mask is None or is_causal:
return True
return attn_mask.dim() == 4
return True
backend = _singletons.get(backend_cls)
if backend is None:
backend = backend_cls()
_singletons[backend_cls] = backend
return backend
@functools.lru_cache(maxsize=1)
def _priority_backends() -> Tuple["AttentionBackend", ...]:
"""Available backends in priority order: cuda -> flash -> torch.
Computed once (machine availability cannot change at runtime) and
cached forever; the tuple always ends with ``TorchNativeBackend``,
which is unconditionally available.
"""
return tuple(
_instance(cls)
for cls in (CudaBackend, FlashAttnBackend, TorchNativeBackend)
if cls.available()
)
def _resolve_default_backend() -> "AttentionBackend":
"""Pick the highest-priority available backend (cuda -> flash -> torch).
Resolved lazily on first ``get_backend()`` and cached. Per-call
capability fallback happens in ``attention()``, so the default is
safe for training and fp32 models.
Resolved lazily on first use and cached via ``_priority_backends``.
Per-call capability fallback happens in ``attention()``, so the
default is safe for training and fp32 models.
"""
return _priority_backends()[0]
@@ -168,9 +167,14 @@ def _environment_backend() -> Optional["AttentionBackend"]:
with _default_backend_lock:
if name != _env_backend_name:
try:
_env_backend = AttentionBackendFactory.create(name)
_env_backend = _resolve_backend(name)
except (ValueError, RuntimeError):
_env_backend = None
logger.warning(
"ASTR_BACKEND=%r is not a registered attention backend; "
"falling back to default resolution",
name,
)
_env_backend_name = name
return _env_backend
@@ -178,43 +182,46 @@ def _environment_backend() -> Optional["AttentionBackend"]:
def _resolve_backend(
backend: Optional[Union[str, ATTN_BACKEND, "AttentionBackend", type]] = None,
) -> "AttentionBackend":
"""Resolve a backend configuration, defaulting to the process policy."""
"""Resolve a backend configuration to its canonical instance.
Accepts a registered name, ``ATTN_BACKEND`` enum value, backend class,
or instance. Names/classes resolve to the shared singleton; a caller
may still pass its own instance to opt out of sharing.
"""
if backend is not None:
if isinstance(backend, ATTN_BACKEND):
return AttentionBackendFactory.create(backend.value)
return _instance(AttentionBackendFactory.get_component_class(backend.value))
if isinstance(backend, str):
return AttentionBackendFactory.create(backend)
return _instance(AttentionBackendFactory.get_component_class(backend))
if isinstance(backend, type) and issubclass(backend, AttentionBackend):
return backend()
return _instance(backend)
if isinstance(backend, AttentionBackend):
return backend
raise TypeError(
f"expected a registered name, ATTN_BACKEND, AttentionBackend type, "
f"or instance, got {type(backend).__name__}"
)
global _default_backend
if _default_backend is None:
with _default_backend_lock:
if _default_backend is None:
_default_backend = _resolve_default_backend()
return _default_backend
return _resolve_default_backend()
def get_backend(
use_default: bool = True,
) -> Optional["AttentionBackend"]:
"""Return the context override, optionally falling back to the process default.
"""Resolve the active backend: explicit context > env > default.
``ASTR_BACKEND`` is a process-wide override and takes precedence over the
context value. Pass ``use_default=False`` at request submission to retain
only an environment override or the caller's :func:`attn_backend` value.
An ``attn_backend(...)`` context is the caller's explicit choice and
always wins. ``ASTR_BACKEND`` is a process-wide override consulted
only when no context is set. Pass ``use_default=False`` at request
submission to retain only an environment override or the caller's
:func:`attn_backend` value.
"""
return (
_environment_backend()
or _current_backend.get()
or (_resolve_backend() if use_default else None)
)
context_backend = _current_backend.get()
if context_backend is not None:
return context_backend
env_backend = _environment_backend()
if env_backend is not None:
return env_backend
return _resolve_default_backend() if use_default else None
@contextmanager
@@ -262,13 +269,23 @@ def attention(
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
fwd: Optional[str] = None,
backend: Optional[Union[str, ATTN_BACKEND, "AttentionBackend", type]] = None,
) -> Tensor:
"""Functional attention entry point — mirrors ``F.scaled_dot_product_attention``.
Delegates to the active backend (set via ``with attn_backend(...)``).
Delegates to the active backend. ``backend`` (optional) is an explicit
escape hatch; when omitted the backend is resolved as
explicit context > ``ASTR_BACKEND`` env > default (cuda > flash > torch).
Handles KV cache I/O, GQA head expansion, and causal masking so the
caller only needs to provide projected q/k/v.
Training calls (``fwd=None``, ``kv_cache=None``) resolve through the
same capability chain — the CUDA cache kernels cannot run without a
cache, so they fall back to flash (mask-free/causal calls only) and
finally to torch SDPA. An explicitly-selected backend that cannot
handle the call raises — an implicit one falls back down the priority
list to the first capable backend.
Args:
q: [batch, q_len, n_heads, head_dim] (blhd)
k: [batch, q_len, n_kv_heads, head_dim] (blhd)
@@ -277,30 +294,43 @@ def attention(
layer_id: transformer layer index for buffer access.
attn_mask: pre-built attention mask (SDPA-compatible).
is_causal: whether to apply causal masking.
fwd: "prefill" / "decode" for inference, None for training.
backend: optional explicit backend (name, enum, class, or instance).
Returns:
[batch, q_len, n_heads * head_dim]
"""
explicit = get_backend(use_default=False)
backend = get_backend()
if fwd is None and explicit is None:
backend = TorchNativeBackend()
if not _backend_supports(backend, q, kv_cache, attn_mask, is_causal, fwd):
if explicit is not None:
if backend is not None:
selected = _resolve_backend(backend)
explicit = True
else:
context_backend = _current_backend.get()
explicit = context_backend is not None
# Resolve through the same chain as inference: explicit context >
# ASTR_BACKEND env > default. Training calls (fwd=None, no cache)
# land on the CUDA backend and fall back by capability below —
# flash when it can handle the call, else torch SDPA.
selected = get_backend()
assert selected is not None
if not selected.supports_call(q, kv_cache, attn_mask, is_causal, fwd):
if explicit:
raise RuntimeError(
f"Explicitly-set backend {type(backend).__name__} cannot "
f"Explicitly-set backend {type(selected).__name__} cannot "
f"handle this attention call (shape={q.shape}, "
f"dtype={q.dtype}, kv_cache={'none' if kv_cache is None else 'present'}, "
f"attn_mask={'none' if attn_mask is None else 'present'}). "
f"Remove the attn_backend() context or switch to a compatible backend."
)
for candidate in _priority_backends():
if isinstance(candidate, type(backend)):
continue
if _backend_supports(candidate, q, kv_cache, attn_mask, is_causal, fwd):
backend = candidate
break
return backend.forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal, fwd)
selected = next(
(
candidate
for candidate in _priority_backends()
if candidate.supports_call(q, kv_cache, attn_mask, is_causal, fwd)
),
_instance(TorchNativeBackend),
)
return selected.forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal, fwd)
class AttentionBackend(ABC):
@@ -310,6 +340,17 @@ class AttentionBackend(ABC):
``fwd_prefill`` (q_len > 1, with or without cache). The public
``forward`` method dispatches based on q_len.
Capability contract — every backend declares:
* ``available()`` — machine-level: can this backend exist here
(kernel ``.so`` loaded, flash-attn present, GPU available)?
Used once to build the default priority list.
* ``supports_call(q, kv_cache, attn_mask, is_causal, fwd)`` — can this
backend run this *specific* call (shape/dtype/cache/mask)? Used by
``attention()`` for the per-call fallback. Resolution logic never
checks concrete backend types, so adding a backend requires no
changes outside its own class.
Three equivalent ways to activate a backend::
with attn_backend(ATTN_BACKEND.TORCH_NATIVE): # enum
@@ -327,6 +368,30 @@ class AttentionBackend(ABC):
def __exit__(self, *exc) -> None:
_current_backend.reset(self._token)
@classmethod
@abstractmethod
def available(cls) -> bool:
"""Return True if this backend can run on the current machine.
Checks static availability only (compiled kernels, optional
packages, GPU presence) — not call-specific constraints.
"""
@abstractmethod
def supports_call(
self,
q: Tensor,
kv_cache: Optional["KVCache"],
attn_mask: Optional[Tensor],
is_causal: bool,
fwd: Optional[str],
) -> bool:
"""Return True if this backend can run this specific attention call.
Called on the canonical singleton instance (or a caller-provided
one); must be side-effect free.
"""
def forward(
self,
q: Tensor,
@@ -412,8 +477,18 @@ class TorchNativeBackend(AttentionBackend):
runs SDPA directly on the projected q/k/v.
"""
@staticmethod
def supports(**kwargs) -> bool:
@classmethod
def available(cls) -> bool:
return True
def supports_call(
self,
q: Tensor,
kv_cache: Optional["KVCache"],
attn_mask: Optional[Tensor],
is_causal: bool,
fwd: Optional[str],
) -> bool:
return True
def fwd_decode(
@@ -516,16 +591,37 @@ class CudaBackend(AttentionBackend):
Raises ``RuntimeError`` if the required kernel is not available.
"""
@staticmethod
def supports(**kwargs) -> bool:
head_dim = kwargs.get("head_dim", -1)
# Head dims supported by the CUDA kernels (single source of truth).
HEAD_DIMS = (32, 64, 128, 256)
@classmethod
def available(cls) -> bool:
return (
torch.cuda.is_available()
and head_dim in (32, 64, 128, 256)
and is_available("attn_paged_decode")
and is_available("attn_paged_prefill")
)
def supports_call(
self,
q: Tensor,
kv_cache: Optional["KVCache"],
attn_mask: Optional[Tensor],
is_causal: bool,
fwd: Optional[str],
) -> bool:
# The CUDA kernels are bf16-only, support head_dim in
# HEAD_DIMS, and need a KV cache (decode/prefill); everything
# else falls back down the priority list to torch.
return (
fwd in ("prefill", "decode")
and kv_cache is not None
and q.ndim == 3
and q.dtype == torch.bfloat16
and q.size(-1) in self.HEAD_DIMS
and is_available(f"attn_paged_{fwd}")
)
@staticmethod
def supports_graph() -> bool:
return True
@@ -606,10 +702,30 @@ class FlashAttnBackend(AttentionBackend):
``flash_attn_func``.
"""
@staticmethod
def supports(**kwargs) -> bool:
@classmethod
def available(cls) -> bool:
return flash_attn_available()
def supports_call(
self,
q: Tensor,
kv_cache: Optional["KVCache"],
attn_mask: Optional[Tensor],
is_causal: bool,
fwd: Optional[str],
) -> bool:
if not self.available():
return False
if q.dtype not in (torch.float16, torch.bfloat16):
return False
if fwd is not None:
return q.ndim == 3 and hasattr(_flash_attn, "flash_attn_varlen_func")
# Dense (training) path: flash_attn_func cannot apply a custom
# mask, so only mask-free calls are supported — ``is_causal`` is
# a flag, not a mask. Masked training (SFT/DPO/GRPO) must fall
# back to TorchNativeBackend instead of silently ignoring the mask.
return attn_mask is None
def fwd_decode(
self,
q: Tensor,
@@ -649,9 +765,9 @@ class FlashAttnBackend(AttentionBackend):
k = repeat_kv(k, n_rep)
v = repeat_kv(v, n_rep)
if attn_mask is not None and not is_causal and attn_mask.dim() != 4:
if attn_mask is not None:
raise ValueError(
"FlashAttnBackend does not support a custom attention mask; "
"FlashAttnBackend cannot handle a custom attention mask; "
"use a causal mask or select TorchNativeBackend."
)
fa = _flash_attn
@@ -664,7 +780,7 @@ class FlashAttnBackend(AttentionBackend):
q.contiguous(),
k.contiguous(),
v.contiguous(),
causal=is_causal or (attn_mask is not None and attn_mask.dim() == 4),
causal=is_causal,
)
return out.contiguous()
+238 -215
View File
@@ -1,163 +1,175 @@
"""FP8 training: scaling state and aten::linear dispatch.
"""FP8 training: scaling recipes, per-tensor state, and aten::linear dispatch.
Layered (see also ``ops/fp8.py`` for the CUDA interface adapter):
1. Kernel interface: ``ops.fp8`` - the only module touching the pybind.
2. Training state (this module): per-tensor scales, amax history, delayed
scaling, and the ``fp8_autocast`` context (TE-style, like
``torch.autocast``).
1. Kernel interface: ``ops.fp8`` the only module touching the pybind.
2. Training state (this module): scaling *recipes* (TE-style delayed scaling
or dynamic current-amax scaling), per-tensor scales + amax history, and
the ``fp8_autocast`` context (like ``torch.autocast``).
3. aten::linear integration (this module): registers the CUDA impl and the
M/N alignment guard.
dtype guard.
Usage::
from astrai.extension.fp8 import fp8_autocast
with fp8_autocast(enabled=True):
with fp8_autocast(enabled=True, fp8_format="hybrid"):
logits = model(input_ids)
loss.backward()
loss.backward() # fp8 backward runs wherever it is called: the
# forward captures the fmt/recipe/meta on the autograd node
Importing this module registers the aten::linear CUDA implementation.
Importing this module registers the aten::linear CUDA and AutogradCUDA
implementations.
Format defaults follow the ecosystem consensus: E4M3 for the forward pass,
E5M2 for the backward (gradient) pass ("hybrid"); every operand's scale is a
quantization step derived from its amax history by the active recipe.
"""
from contextlib import contextmanager
from dataclasses import dataclass
from enum import Enum
from typing import Optional
import torch
from torch.library import Library
from astrai.extension.ops.fp8 import (
linear_backward_scaled,
linear_forward_scaled,
linear_backward_fp8,
linear_forward_fp8,
)
E4M3_MAX = 448.0
# Max representable value per FP8 format (E4M3: 448, E5M2: 57344).
FP8_MAX = {"e4m3": 448.0, "e5m2": 57344.0}
# ---------------------------------------------------------------------------
# Layer 2: training state (scales, amax history, delayed scaling, autocast)
# ---------------------------------------------------------------------------
class FP8Format(str, Enum):
"""Per-direction FP8 format. HYBRID = E4M3 forward / E5M2 backward."""
E4M3 = "e4m3"
E5M2 = "e5m2"
HYBRID = "hybrid"
def fwd(self) -> str:
return "e4m3" if self is FP8Format.HYBRID else self.value
def bwd(self) -> str:
return "e5m2" if self is FP8Format.HYBRID else self.value
class FP8Recipe:
"""Scale-from-amax policy; the scale computation is the injection point.
``scale_from_history`` receives the amax tensor for this operand (a ring
window for delayed scaling, the current amax for dynamic scaling) and
returns the quantization step: ``scale = (amax / FP8_MAX[fmt]) / 2^margin``.
"""
history_len: int = 16
margin: int = 0
def scale_from_history(self, amax: torch.Tensor, fmt: str) -> torch.Tensor:
raise NotImplementedError
@dataclass
class DelayedScaling(FP8Recipe):
"""TE-style delayed scaling: max over the amax history window.
The scale is computed from amax measured in *previous* steps (delayed one
step); the window length trades responsiveness against stability.
"""
history_len: int = 16
margin: int = 0
def scale_from_history(self, amax: torch.Tensor, fmt: str) -> torch.Tensor:
peak = amax.max()
return ((peak / FP8_MAX[fmt]) / (2**self.margin)).clamp_min(1e-12)
@dataclass
class DynamicScaling(FP8Recipe):
"""Current-amax scaling (torchao DYNAMIC): measure, then quantize.
No history — the scale is derived from the amax of the tensor being
quantized in the same step, at the cost of an extra reduction pass.
"""
history_len: int = 1
margin: int = 0
def scale_from_history(self, amax: torch.Tensor, fmt: str) -> torch.Tensor:
peak = amax.max()
return ((peak / FP8_MAX[fmt]) / (2**self.margin)).clamp_min(1e-12)
class _ScaleRing:
"""One operand's delayed-scaling state: amax history ring + derived scale.
The ring captures its recipe at construction; ``update`` records a fresh
amax and refreshes the scale for the *next* step (delayed one step).
"""
__slots__ = ("recipe", "hist", "idx", "scale", "initialized")
def __init__(self, device: torch.device, recipe: FP8Recipe):
self.recipe = recipe
n = recipe.history_len
self.hist = torch.ones(n, device=device, dtype=torch.float32)
self.idx = 0
self.scale = torch.ones(1, device=device, dtype=torch.float32)
self.initialized = False
def update(self, amax: torch.Tensor, fmt: str) -> None:
self.hist[self.idx] = amax.reshape(())
self.idx = (self.idx + 1) % self.hist.numel()
self.scale.copy_(self.recipe.scale_from_history(self.hist, fmt))
def seed(self, t: torch.Tensor, fmt: str) -> None:
amax = t.abs().amax().to(torch.float32).clamp_min(1e-12)
self.hist.fill_(amax)
self.scale.copy_(self.recipe.scale_from_history(self.hist, fmt))
self.initialized = True
class FP8TensorMeta:
"""Scales + amax state for one weight tensor and its paired activations.
"""Per-weight delayed-scaling state: one ring per operand role.
- weight: delayed scale from a 16-step amax history window (TE style)
- x/g: delayed one step, reuse the quantize kernel's free atomic amax
Holds the ``w`` / ``x`` / ``g`` rings; fused kernels record the amax
while quantizing, so the scale used at step N reflects amax from steps
< N. DynamicScaling never allocates a meta — it measures the current
amax inline (``_dynamic_scale``), so it needs no history storage.
"""
__slots__ = (
"scale",
"scale_inv",
"amax_history",
"idx",
"x_scale",
"x_scale_inv",
"x_history",
"x_idx",
"g_scale",
"g_scale_inv",
"g_history",
"g_idx",
"w_init",
"x_init",
"g_init",
)
__slots__ = ("w", "x", "g")
def __init__(self, device: torch.device, update_interval: int):
self.scale = torch.ones(1, device=device, dtype=torch.float32)
self.scale_inv = torch.ones(1, device=device, dtype=torch.float32)
self.amax_history = torch.ones(
update_interval, device=device, dtype=torch.float32
)
self.idx = 0
self.x_scale = torch.ones(1, device=device, dtype=torch.float32)
self.x_scale_inv = torch.ones(1, device=device, dtype=torch.float32)
self.x_history = torch.ones(update_interval, device=device, dtype=torch.float32)
self.x_idx = 0
self.g_scale = torch.ones(1, device=device, dtype=torch.float32)
self.g_scale_inv = torch.ones(1, device=device, dtype=torch.float32)
self.g_history = torch.ones(update_interval, device=device, dtype=torch.float32)
self.g_idx = 0
self.w_init = False
self.x_init = False
self.g_init = False
def init_scale(self, t: torch.Tensor) -> None:
"""Immediate scale from the current amax; used on the first call.
A scale of 1 would underflow small activations/gradients (e4m3 min
normal is 2^-6); initialize from the actual amax once, then delayed
updates take over.
"""
amax = t.abs().amax().to(torch.float32).clamp_min(1e-12)
self.scale.copy_(amax / E4M3_MAX)
self.scale_inv.copy_(E4M3_MAX / amax)
self.record(amax)
def push_x_scale(self, amax: torch.Tensor) -> None:
"""Window update for the activation scale (delayed, TE style)."""
self.x_history[self.x_idx] = amax.reshape(())
self.x_idx = (self.x_idx + 1) % self.x_history.numel()
m = self.x_history.max()
self.x_scale.copy_(m / E4M3_MAX)
self.x_scale_inv.copy_(E4M3_MAX / m)
def push_g_scale(self, amax: torch.Tensor) -> None:
"""Window update for the gradient scale (delayed, TE style)."""
self.g_history[self.g_idx] = amax.reshape(())
self.g_idx = (self.g_idx + 1) % self.g_history.numel()
m = self.g_history.max()
self.g_scale.copy_(m / E4M3_MAX)
self.g_scale_inv.copy_(E4M3_MAX / m)
def record(self, amax: torch.Tensor) -> None:
"""Push the latest amax into the ring buffer (device-side copy, no sync)."""
self.amax_history[self.idx] = amax.reshape(())
self.idx = (self.idx + 1) % self.amax_history.numel()
def refresh(self) -> None:
"""Recompute scale from the amax history window (delayed scaling)."""
amax = self.amax_history.max()
if amax > 0:
self.scale.copy_(amax / E4M3_MAX)
self.scale_inv.copy_(E4M3_MAX / amax)
def __init__(self, device: torch.device, recipe: FP8Recipe):
self.w = _ScaleRing(device, recipe)
self.x = _ScaleRing(device, recipe)
self.g = _ScaleRing(device, recipe)
class FP8State:
"""Global fp8 training state, TE-style."""
"""Global fp8 training state: active recipe + per-tensor metas."""
def __init__(self, update_interval: int = 16):
def __init__(self):
self.enabled = False
self.update_interval = update_interval
self.step_count = 0
self.recipe: FP8Recipe = DelayedScaling()
self.fp8_format: FP8Format = FP8Format.HYBRID
self._metas: dict[tuple, FP8TensorMeta] = {}
self._last_device: torch.device | None = None
def _get_device(self, t: torch.Tensor) -> torch.device:
if self._last_device is None:
self._last_device = t.device
return t.device
def get_weight_meta(self, w: torch.Tensor) -> FP8TensorMeta:
key = (w.data_ptr(), w.shape, w.dtype)
meta = self._metas.get(key)
if meta is None:
meta = FP8TensorMeta(self._get_device(w), self.update_interval)
meta = FP8TensorMeta(w.device, self.recipe)
self._metas[key] = meta
return meta
def step(self) -> None:
"""Advance the counter and refresh all weight scales every N steps."""
self.step_count += 1
if self.step_count % self.update_interval == 0:
for meta in self._metas.values():
meta.refresh()
def reset(self) -> None:
self.enabled = False
self.step_count = 0
self._metas.clear()
self._last_device = None
# Global singleton: autograd backward runs on the engine worker threads, so
@@ -171,100 +183,129 @@ def fp8_state() -> FP8State:
@contextmanager
def fp8_autocast(enabled: bool = True, update_interval: int = 16):
def fp8_autocast(
enabled: bool = True,
update_interval: int = 16,
recipe: Optional[FP8Recipe] = None,
fp8_format: str = "hybrid",
margin: int = 0,
):
"""Autocast-style context: fp8 linear dispatch on this thread.
Usage::
with fp8_autocast(enabled=True):
with fp8_autocast(enabled=True, fp8_format="hybrid"):
logits = model(input_ids) # aten::linear -> fp8 path
loss.backward()
loss.backward() # fp8 backward; state was captured at forward time
The scale-update counter advances once per ``enter`` (one training step),
refreshing weight scales from their amax history every ``update_interval``.
Args:
enabled: toggle fp8 dispatch for aten::linear.
update_interval: legacy alias for the delayed-scaling history window
(used only when ``recipe`` is not given).
recipe: scaling policy; defaults to ``DelayedScaling(update_interval)``.
fp8_format: ``"e4m3"`` / ``"e5m2"`` / ``"hybrid"`` (default) — hybrid
means E4M3 forward, E5M2 backward.
margin: scale headroom (``scale = (amax / FP8_MAX) / 2^margin``) used
with the default delayed recipe.
"""
state = fp8_state()
prev_enabled = state.enabled
prev_interval = state.update_interval
prev = (state.enabled, state.recipe, state.fp8_format)
if recipe is None:
recipe = DelayedScaling(history_len=update_interval, margin=margin)
state.enabled = enabled
state.update_interval = update_interval
state.recipe = recipe
state.fp8_format = FP8Format(fp8_format)
try:
if enabled:
state.step()
yield
finally:
state.enabled = prev_enabled
state.update_interval = prev_interval
state.enabled, state.recipe, state.fp8_format = prev
# ---------------------------------------------------------------------------
# Strategy-level forward / backward (called from the aten::linear impl)
# ---------------------------------------------------------------------------
def _dynamic_scale(t: torch.Tensor, recipe: FP8Recipe, fmt: str) -> torch.Tensor:
amax = t.abs().amax().to(torch.float32).clamp_min(1e-12)
return recipe.scale_from_history(amax, fmt)
def fp8_linear_forward(x: torch.Tensor, w: torch.Tensor, bias=None):
"""TE-style scaled fp8 linear forward (called from the aten::linear impl).
"""Scaled fp8 linear forward (called from the aten::linear impl).
x uses the delayed scale of its paired weight meta (amax from the previous
forward of this linear); the quantize kernel emits the current amax for the
next step. No extra abs/max reduce.
Pure FP8 path for both recipes: quantize x/w with the active scales, run
the pre-quantized GEMM, and feed the freshly measured amax back into the
delayed-scaling ring (dynamic scaling measures the current amax itself).
"""
if bias is None:
bias = torch.empty(0, device=x.device, dtype=x.dtype)
state = fp8_state()
meta = state.get_weight_meta(w)
if not meta.w_init:
meta.init_scale(w)
meta.w_init = True
if not meta.x_init:
amax = x.abs().amax().to(torch.float32).clamp_min(1e-12)
meta.x_history.fill_(amax)
meta.x_scale.copy_(amax / E4M3_MAX)
meta.x_scale_inv.copy_(E4M3_MAX / amax)
meta.x_init = True
amax_x = torch.empty(1, device=x.device, dtype=torch.float32)
amax_w = torch.empty(1, device=x.device, dtype=torch.float32)
out = linear_forward_scaled(
x,
w,
bias,
meta.x_scale,
meta.scale,
meta.x_scale_inv,
meta.scale_inv,
amax_x,
amax_w,
)
meta.record(amax_w)
meta.push_x_scale(amax_x)
fmt = state.fp8_format.fwd()
if isinstance(state.recipe, DynamicScaling):
meta = None
sx = _dynamic_scale(x.reshape(-1, w.size(1)), state.recipe, fmt)
sw = _dynamic_scale(w, state.recipe, fmt)
else:
meta = state.get_weight_meta(w)
if not meta.w.initialized:
meta.w.seed(w, fmt)
if not meta.x.initialized:
meta.x.seed(x, fmt)
sx, sw = meta.x.scale, meta.w.scale
out, amax_x, amax_w = linear_forward_fp8(x, w, bias, sx, sw, fmt)
if meta is not None:
meta.x.update(amax_x, fmt)
meta.w.update(amax_w, fmt)
return out
def fp8_linear_backward(g, x, w, masks):
"""TE-style scaled fp8 linear backward (called from aten::linear_backward)."""
state = fp8_state()
meta = state.get_weight_meta(w)
if not meta.g_init:
amax = g.abs().amax().to(torch.float32).clamp_min(1e-12)
meta.g_history.fill_(amax)
meta.g_scale.copy_(amax / E4M3_MAX)
meta.g_scale_inv.copy_(E4M3_MAX / amax)
meta.g_init = True
amax_g = torch.empty(1, device=g.device, dtype=torch.float32)
out = linear_backward_scaled(
g,
x,
w,
masks,
meta.g_scale,
meta.scale,
meta.x_scale,
meta.g_scale_inv,
meta.scale_inv,
meta.x_scale_inv,
amax_g,
)
meta.push_g_scale(amax_g)
return out
class _LinearFp8(torch.autograd.Function):
"""The fp8 linear forward/backward pair (standard Function style).
The forward runs inside the ``fp8_autocast`` region and captures the
active fmt/recipe/meta on ``ctx``; the backward reads only the captured
state, so ``loss.backward()`` may run after the context exits. The
gradient is quantized once (E5M2 in hybrid mode) and the dX / dW GEMMs
share that quantization; output masks come from ``needs_input_grad``.
"""
@staticmethod
def forward(ctx, x, w, bias):
out = fp8_linear_forward(x, w, bias)
state = fp8_state()
ctx.save_for_backward(x, w)
ctx.fmt_bwd = state.fp8_format.bwd()
ctx.recipe = state.recipe
ctx.is_dynamic = isinstance(state.recipe, DynamicScaling)
ctx.meta = None if ctx.is_dynamic else state.get_weight_meta(w)
return out
@staticmethod
@torch.autograd.function.once_differentiable
def backward(ctx, g):
x, w = ctx.saved_tensors
fmt = ctx.fmt_bwd
if ctx.is_dynamic:
sg = _dynamic_scale(g, ctx.recipe, fmt)
sw = _dynamic_scale(w, ctx.recipe, fmt)
sx = _dynamic_scale(x, ctx.recipe, fmt)
else:
meta = ctx.meta
if not meta.g.initialized:
meta.g.seed(g, fmt)
sg, sw, sx = meta.g.scale, meta.w.scale, meta.x.scale
masks = list(ctx.needs_input_grad)
grad_x, grad_w, grad_b, amax_g = linear_backward_fp8(
g, x, w, masks, sg, sw, sx, fmt
)
if not ctx.is_dynamic:
ctx.meta.g.update(amax_g, fmt)
return grad_x, grad_w, grad_b if masks[2] else None
# ---------------------------------------------------------------------------
# Layer 3: aten::linear integration
# aten::linear integration
# ---------------------------------------------------------------------------
@@ -279,9 +320,14 @@ def fp8_linear_enabled() -> bool:
def _fp8_supported(x: torch.Tensor, w: torch.Tensor) -> bool:
"""cuBLASLt fp8 requires M % 16 == 0 and N % 16 == 0 (K is padded)."""
m = x.numel() // x.size(-1)
return m % 16 == 0 and w.size(0) % 16 == 0
"""Shape guard for the fp8 linear path.
Unlike a strict 16-alignment requirement, the fp8 kernels handle unaligned
M/N via boundary checks (slower but correct) — so no whole-call bf16
fallback for small decode batches. Only the K-dimension contraction must
match, and the weight must be 2D.
"""
return x.dim() >= 2 and w.dim() == 2 and x.size(-1) == w.size(1)
def _linear_cuda_impl(x: torch.Tensor, w: torch.Tensor, bias=None):
@@ -291,7 +337,7 @@ def _linear_cuda_impl(x: torch.Tensor, w: torch.Tensor, bias=None):
and w.dtype == torch.bfloat16
and _fp8_supported(x, w)
):
return fp8_linear_forward(x, w, bias)
return _LinearFp8.apply(x, w, bias)
return torch.ops.aten.linear.default.redispatch(
torch._C.DispatchKeySet(torch._C.DispatchKey.CompositeImplicitAutograd),
x,
@@ -300,35 +346,12 @@ def _linear_cuda_impl(x: torch.Tensor, w: torch.Tensor, bias=None):
)
def _linear_backward_cuda_impl(input_tensor, grad_output, weight, output_mask):
if (
fp8_linear_enabled()
and weight.dtype == torch.bfloat16
and _fp8_supported(grad_output, weight)
):
return fp8_linear_backward(grad_output, input_tensor, weight, list(output_mask))
compute_dtype = weight.dtype
grad = grad_output.to(compute_dtype)
grad_2d = grad.reshape(-1, weight.size(0))
input_2d = input_tensor.reshape(-1, input_tensor.size(-1)).to(compute_dtype)
grad_input = (
torch.mm(grad_2d, weight)
if output_mask[0]
else torch.empty(0, device=input_tensor.device, dtype=input_tensor.dtype)
)
grad_weight = (
torch.mm(grad_2d.t(), input_2d)
if output_mask[1]
else torch.empty(0, device=input_tensor.device, dtype=input_tensor.dtype)
)
grad_bias = (
grad.sum(dim=0)
if output_mask[2]
else torch.empty(0, device=input_tensor.device, dtype=input_tensor.dtype)
)
return grad_input.reshape_as(input_tensor), grad_weight, grad_bias
_lib = Library("aten", "IMPL", "CUDA")
_lib.impl("linear", _linear_cuda_impl)
_lib.impl("linear_backward", _linear_backward_cuda_impl)
# Also replace torch's generated linear autograd formula (which would call
# aten::linear_backward after the fp8_autocast region exits). The fp8
# backward is owned by _LinearFp8 with its state captured at forward time,
# so loss.backward() works wherever it is called; the same CUDA registration
# still covers inference_mode, where autograd keys are skipped entirely.
_lib_autograd = Library("aten", "IMPL", "AutogradCUDA")
_lib_autograd.impl("linear", _linear_cuda_impl)
+63 -23
View File
@@ -1,43 +1,83 @@
"""Dynamic discovery and loading of compiled CUDA kernel modules.
Each kernel is registered in ``csrc/build.py`` and built into a ``.so`` placed
in this package directory. On import we try to load each one; kernels that
failed to build (or are running on a CPU-only machine) are marked unavailable
so the wrapper functions can fall back to ``torch`` SDPA.
Each kernel is built by the CMake build in ``csrc/CMakeLists.txt`` into a
``.so`` placed in ``astrai/extension/lib/`` — the module name equals the
``.so`` name equals the pybind name (e.g. ``attn_decode``, defined via
``TORCH_EXTENSION_NAME``). ``KERNEL_NAMES`` is discovered automatically from
the ``.so`` files present, so adding a kernel to the CMake ``KERNELS``
registry needs no change here.
Loading is **lazy and centralized**: module names are discovered eagerly
(cheap glob), but each ``.so`` is imported on first use via the single
``get_module`` accessor, then cached. The wrapper modules (``ops/*.py``) never
touch the internals or keep their own caches — they call ``get_module(name)``
(or ``is_available(name)`` when a torch fallback is acceptable). A kernel that
failed to build (or is running on a CPU-only machine) is ``None`` in the cache,
so ``is_available`` returns ``False`` and ``get_module`` raises a clear error.
"""
import glob
import importlib
import logging
import os
logger = logging.getLogger(__name__)
KERNEL_NAMES = [
"attn_decode",
"attn_prefill",
"attn_paged_decode",
"attn_paged_prefill",
"rotary_emb",
"fp8_mm",
]
_LIB_DIR = os.path.join(os.path.dirname(__file__), "lib")
def _discover_kernel_names() -> list[str]:
"""Return the module names of the compiled kernel ``.so`` files in lib/."""
names: list[str] = []
for path in glob.glob(os.path.join(_LIB_DIR, "*.so")):
# strip the "<soabi>.so" suffix, e.g. attn_decode.cpython-312-...so
names.append(os.path.basename(path).split(".", 1)[0])
return sorted(names)
KERNEL_NAMES = _discover_kernel_names()
_available: dict[str, bool] = {}
_modules: dict[str, object] = {}
for _name in KERNEL_NAMES:
try:
_mod = importlib.import_module(f".lib.{_name}", package=__package__)
_available[_name] = True
_modules[_name] = _mod
except ImportError:
_available[_name] = False
_modules[_name] = None
def _try_load(name: str) -> object:
"""Import and cache the ``name`` kernel module (lazy, one attempt).
Returns the module, or ``None`` if it is unavailable. Cached so each
``.so`` is imported at most once per process.
"""
if name not in _modules:
try:
_modules[name] = importlib.import_module(
f".lib.{name}", package=__package__
)
_available[name] = True
except ImportError:
logger.warning("kernel '%s' failed to import; marking unavailable", name)
_modules[name] = None
_available[name] = False
return _modules[name]
def is_available(name: str) -> bool:
"""Return ``True`` if the compiled kernel ``name`` was loaded."""
"""Return ``True`` if the compiled kernel ``name`` could be loaded."""
if name not in _available:
_try_load(name)
return _available.get(name, False)
def get_module(name: str) -> object:
"""Return the loaded kernel module for ``name``, or ``None`` if unavailable."""
return _modules.get(name)
"""Return the loaded kernel module for ``name``, importing it on first use.
Raises ``RuntimeError`` if the kernel is unavailable (not built, or failed
to import) — callers that can tolerate a torch fallback should check
``is_available(name)`` first instead.
"""
mod = _try_load(name)
if mod is None:
raise RuntimeError(
f"CUDA kernel '{name}' is not available. "
f"Build with CSRC_KERNELS=true (or use the torch-native fallback)."
)
return mod
+9 -17
View File
@@ -17,7 +17,7 @@ from typing import Optional
import torch
from astrai.extension.loader import _available, _modules
from astrai.extension.loader import get_module
class TensorLayout(enum.IntEnum):
@@ -30,14 +30,6 @@ class TensorLayout(enum.IntEnum):
BLHD = 1 # [batch, seq_len, n_heads, head_dim]
def _check_available(name: str):
if not _available.get(name):
raise RuntimeError(
f"CUDA kernel '{name}' is not available. "
f"Build with CSRC_KERNELS=true or use a torch-native backend."
)
def attn_decode(
q: torch.Tensor,
k: torch.Tensor,
@@ -57,9 +49,9 @@ def attn_decode(
Returns:
[batch, 1, n_heads, head_dim] (blhd, bf16)
"""
_check_available("attn_decode")
mod = get_module("attn_decode")
causal_offset = (k.size(1) - 1) if is_causal else -1
return _modules["attn_decode"].attn_decode(
return mod.attn_decode(
q, k, v, mask=mask, causal_offset=causal_offset, layout=TensorLayout.BLHD
)
@@ -83,9 +75,9 @@ def attn_prefill(
Returns:
[batch, q_len, n_heads, head_dim] (blhd, bf16)
"""
_check_available("attn_prefill")
mod = get_module("attn_prefill")
causal_offset = (k.size(1) - q.size(1)) if is_causal else -1
return _modules["attn_prefill"].attn_prefill(
return mod.attn_prefill(
q, k, v, mask=mask, causal_offset=causal_offset, layout=TensorLayout.BLHD
)
@@ -129,9 +121,9 @@ def attn_paged_decode(
Returns:
[batch, n_heads, head_dim] (bf16, 3D)
"""
_check_available("attn_paged_decode")
mod = get_module("attn_paged_decode")
causal_offset = 0 if is_causal else -1
return _modules["attn_paged_decode"].attn_paged_decode(
return mod.attn_paged_decode(
q,
k_cache,
v_cache,
@@ -183,9 +175,9 @@ def attn_paged_prefill(
Returns:
[total_q, n_heads, head_dim] (bf16, 3D)
"""
_check_available("attn_paged_prefill")
mod = get_module("attn_paged_prefill")
causal_offset = 0 if is_causal else -1
return _modules["attn_paged_prefill"].attn_paged_prefill(
return mod.attn_paged_prefill(
q,
k_cache,
v_cache,
+150 -98
View File
@@ -1,9 +1,16 @@
"""FP8 CUDA kernel interface adapter (the only module touching the pybind.
"""FP8 CUDA kernel interface adapter (the only module touching the pybind).
Isolates the ``fp8_mm`` CUDA extension behind stable Python functions:
- availability / dtype checks and clear errors
- torch.library ``custom::fp8_mm`` registration (meta + CPU fallback)
- quantize-in-GEMM primitives used by ``fp8.py`` training state
Isolates the ``fp8_ops`` CUDA extension behind stable Python primitives:
- ``quantize_bf16(x, scale, fmt) -> (x8, amax)`` — BF16 → FP8 with fused amax
- ``mm_fp8(a8, b8, sa, sb) -> out`` — pre-quantized FP8 GEMM (BF16 output)
- ``linear_forward_fp8(x, w, bias, sx, sw) -> (out, amax_x, amax_w)``
- ``linear_backward_fp8(g, x, w, masks, sg, sw, sx, fmt) -> (gx, gw, gb, amax_g)``
Scale semantics: scales are *quantization steps* — the value divided out when
quantizing (``x8 = x / scale``). Every primitive computes its own inverse
internally; callers never pass ``scale_inv``. ``amax`` values are *returned*,
never passed as output arguments. ``fmt`` is ``"e4m3"`` or ``"e5m2"``.
Policy (scales, amax history, delayed scaling, autocast) lives in ``fp8.py``;
this module is stateless.
@@ -12,108 +19,153 @@ this module is stateless.
import torch
from torch.library import custom_op
from astrai.extension.loader import get_module, is_available
from astrai.extension.loader import get_module
# fmt string -> kernel int (0 = E4M3, 1 = E5M2)
_FMT_TO_INT = {"e4m3": 0, "e5m2": 1}
def _mod():
if not is_available("fp8_mm"):
raise RuntimeError(
"CUDA kernel 'fp8_mm' is not available. Build with CSRC_KERNELS=true."
)
return get_module("fp8_mm")
def _fmt_int(fmt: str) -> int:
try:
return _FMT_TO_INT[fmt]
except KeyError:
raise ValueError(f"unsupported fp8 format {fmt!r} (expected 'e4m3' or 'e5m2')")
@custom_op("custom::fp8_mm", mutates_args=())
def fp8_mm(
a: torch.Tensor, b: torch.Tensor, sx: torch.Tensor, sw: torch.Tensor
) -> torch.Tensor:
"""BF16 inputs, fused FP8 GEMM with FP32 accumulation and BF16 output."""
def _fmt_dtype(fmt: str) -> torch.dtype:
return torch.float8_e5m2 if _fmt_int(fmt) else torch.float8_e4m3fn
@fp8_mm.register_fake
def _fp8_mm_fake(a, b, sx, sw):
return torch.empty((a.size(0), b.size(0)), device=a.device, dtype=torch.bfloat16)
@custom_op("custom::fp8_quantize", mutates_args=())
def fp8_quantize(
x: torch.Tensor, scale: torch.Tensor, fmt: int
) -> tuple[torch.Tensor, torch.Tensor]:
"""BF16 -> FP8 quantize with fused amax; returns ``(x8, amax)``."""
@fp8_mm.register_kernel("cuda")
def _fp8_mm_cuda(a, b, sx, sw):
if not (a.dtype == torch.bfloat16 and b.dtype == torch.bfloat16):
raise TypeError(f"bf16 GEMM requires bf16 inputs, got {a.dtype}/{b.dtype}")
return _mod().fp8_mm(a, b, sx, sw)
@fp8_mm.register_kernel("cpu")
def _fp8_mm_cpu(a, b, sx, sw):
return torch.mm(a.float(), b.float().t()).to(torch.bfloat16)
@custom_op("custom::fp8_mm_prequant", mutates_args=())
def fp8_mm_prequant(
a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor
) -> torch.Tensor:
"""Pre-quantized FP8 inputs, fused FP8 GEMM, FP32 accumulation, BF16 out."""
@fp8_mm_prequant.register_fake
def _fp8_mm_prequant_fake(a, b, scale):
return torch.empty((a.size(0), b.size(0)), device=a.device, dtype=torch.bfloat16)
@fp8_mm_prequant.register_kernel("cuda")
def _fp8_mm_prequant_cuda(a, b, scale):
if not (a.dtype == torch.float8_e4m3fn and b.dtype == torch.float8_e4m3fn):
raise TypeError(
f"pre-quantized FP8 GEMM requires fp8 inputs, got {a.dtype}/{b.dtype}"
)
return _mod().fp8_mm_prequant(a, b, scale)
@fp8_mm_prequant.register_kernel("cpu")
def _fp8_mm_prequant_cpu(a, b, scale):
return (a.float() @ b.float().t() * scale).to(torch.bfloat16)
@custom_op("custom::fp8_mm_prequant_fp8", mutates_args=())
def fp8_mm_prequant_fp8(
a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, out_scale: torch.Tensor
) -> torch.Tensor:
"""FP8 inputs and FP8 output: fused FP8 GEMM with FP32 accumulation."""
@fp8_mm_prequant_fp8.register_fake
def _fp8_mm_prequant_fp8_fake(a, b, scale, out_scale):
return torch.empty((a.size(0), b.size(0)), device=a.device, dtype=a.dtype)
@fp8_mm_prequant_fp8.register_kernel("cuda")
def _fp8_mm_prequant_fp8_cuda(a, b, scale, out_scale):
if not (a.dtype == torch.float8_e4m3fn and b.dtype == torch.float8_e4m3fn):
raise TypeError(
f"pre-quantized FP8 GEMM requires fp8 inputs, got {a.dtype}/{b.dtype}"
)
return _mod().fp8_mm_prequant_fp8(a, b, scale, out_scale)
@fp8_mm_prequant_fp8.register_kernel("cpu")
def _fp8_mm_prequant_fp8_cpu(a, b, scale, out_scale):
return (a.float() @ b.float().t() * scale * out_scale).to(torch.float8_e4m3fn)
def linear_forward_scaled(x, w, bias, sx, sw, sx_inv, sw_inv, amax_x, amax_w):
"""Quantize BF16 inputs to FP8, accumulate in FP32, and return BF16.
x/w: [..., K] / [N, K] bf16; sx/sw and their inverses control the fused
E4M3 conversion; amax_x/amax_w receive the input max-abs values.
"""
if not (x.dtype == torch.bfloat16 and w.dtype == torch.bfloat16):
raise TypeError(f"fp8 forward requires bf16 inputs, got {x.dtype}/{w.dtype}")
return _mod().fp8_linear_forward_scaled(
x, w, bias, sx, sw, sx_inv, sw_inv, amax_x, amax_w
@fp8_quantize.register_fake
def _fp8_quantize_fake(x, scale, fmt):
dtype = torch.float8_e5m2 if fmt else torch.float8_e4m3fn
return (
torch.empty(x.shape, device=x.device, dtype=dtype),
torch.empty(1, device=x.device, dtype=torch.float32),
)
def linear_backward_scaled(g, x, w, masks, sg, sw, sx, sg_inv, sw_inv, sx_inv, amax_g):
"""dX = g @ W, dW = g^T @ X, dB = sum(g) with per-tensor scales."""
@fp8_quantize.register_kernel("cuda")
def _fp8_quantize_cuda(x, scale, fmt):
if x.dtype != torch.bfloat16:
raise TypeError(f"fp8 quantize requires bf16 input, got {x.dtype}")
return get_module("fp8_ops").quantize_bf16(x, scale, int(fmt))
@fp8_quantize.register_kernel("cpu")
def _fp8_quantize_cpu(x, scale, fmt):
x8 = (x.float() / scale).to(_fmt_dtype("e5m2" if fmt else "e4m3"))
amax = x.abs().amax().float().reshape(1).clamp_min(1e-12)
return x8, amax
@custom_op("custom::fp8_gemm", mutates_args=())
def fp8_gemm(
a: torch.Tensor,
b: torch.Tensor,
sa: torch.Tensor,
sb: torch.Tensor,
out_dtype: int = 0,
out_scale: torch.Tensor | None = None,
) -> torch.Tensor:
"""FP8 GEMM: ``a @ b * (sa * sb)`` with FP32 accumulation.
``out_dtype``: 0 = BF16 (default), 1 = FP8 E4M3 (requires ``out_scale``,
the quantization step for the output — mirrors ``torch._scaled_mm``).
"""
@fp8_gemm.register_fake
def _fp8_gemm_fake(a, b, sa, sb, out_dtype=0, out_scale=None):
dtype = torch.float8_e4m3fn if out_dtype else torch.bfloat16
return torch.empty((a.size(0), b.size(1)), device=a.device, dtype=dtype)
@fp8_gemm.register_kernel("cuda")
def _fp8_gemm_cuda(a, b, sa, sb, out_dtype=0, out_scale=None):
if a.dtype != b.dtype or a.dtype not in (torch.float8_e4m3fn, torch.float8_e5m2):
raise TypeError(
f"fp8 GEMM requires matching fp8 inputs, got {a.dtype}/{b.dtype}"
)
return get_module("fp8_ops").mm_fp8(a, b, sa, sb, int(out_dtype), out_scale)
@fp8_gemm.register_kernel("cpu")
def _fp8_gemm_cpu(a, b, sa, sb, out_dtype=0, out_scale=None):
acc = a.float() @ b.float() * sa * sb
if out_dtype:
os_ = 1.0 if out_scale is None else out_scale
return (acc * os_).to(torch.float8_e4m3fn)
return acc.to(torch.bfloat16)
def quantize_bf16(x: torch.Tensor, scale: torch.Tensor, fmt: str = "e4m3"):
"""BF16 -> FP8 quantize with fused amax; returns ``(x8, amax)``.
``scale`` is the quantization step (device scalar); ``fmt`` selects
E4M3 or E5M2. ``amax`` is a fresh 1-element float32 tensor — the caller
never clears it.
"""
return fp8_quantize(x, scale, _fmt_int(fmt))
def mm_fp8(
a: torch.Tensor,
b: torch.Tensor,
sa: torch.Tensor,
sb: torch.Tensor,
out_dtype: str = "bf16",
out_scale: torch.Tensor | None = None,
) -> torch.Tensor:
"""Pre-quantized FP8 GEMM: ``a @ b * (sa * sb)``.
``a``/``b`` must be FP8 tensors of the same format (E4M3 or E5M2);
``sa``/``sb`` are their quantization steps. ``out_dtype`` is ``"bf16"``
(default) or ``"e4m3"`` — FP8 output for layer-to-layer pipelines, which
requires ``out_scale`` (the output quantization step).
"""
if out_dtype not in ("bf16", "e4m3"):
raise ValueError(
f"unsupported out_dtype {out_dtype!r} (expected 'bf16' or 'e4m3')"
)
return fp8_gemm(a, b, sa, sb, int(out_dtype == "e4m3"), out_scale)
def linear_forward_fp8(x, w, bias, sx, sw, fmt: str = "e4m3", bias_scale=None):
"""Pure FP8 linear forward: quantize x/w to ``fmt``, pre-quantized GEMM.
Returns ``(out, amax_x, amax_w)``. ``bias`` may be ``None``. For static
fp8 inference, ``w`` and ``bias`` may arrive pre-quantized to ``fmt``
(produced by :func:`quantize_bf16` with their scales as ``sw`` /
``bias_scale``); a pre-quantized ``bias`` requires ``bias_scale``, and
its ``amax_w`` comes back 0. The bias is fused into the GEMM epilogue.
"""
fmt8 = _fmt_dtype(fmt)
if x.dtype != torch.bfloat16 or w.dtype not in (torch.bfloat16, fmt8):
raise TypeError(
f"fp8 forward requires bf16 x and bf16-or-{fmt} w, got {x.dtype}/{w.dtype}"
)
if bias is None:
bias = torch.empty(0, device=x.device, dtype=x.dtype)
return get_module("fp8_ops").linear_forward_fp8(
x, w, bias, sx, sw, _fmt_int(fmt), bias_scale
)
def linear_backward_fp8(g, x, w, masks, sg, sw, sx, fmt: str = "e5m2"):
"""FP8 linear backward; returns ``(grad_input, grad_weight, grad_bias, amax_g)``.
The gradient (and the transposed w/x operands) are quantized to ``fmt``
(default E5M2 — larger dynamic range for gradients) and the two GEMMs run
as FP8 tensor-core products sharing a single gradient quantization.
"""
if not (
g.dtype == torch.bfloat16
and x.dtype == torch.bfloat16
@@ -122,6 +174,6 @@ def linear_backward_scaled(g, x, w, masks, sg, sw, sx, sg_inv, sw_inv, sx_inv, a
raise TypeError(
f"fp8 backward requires bf16 inputs, got {g.dtype}/{x.dtype}/{w.dtype}"
)
return _mod().fp8_linear_backward_scaled(
g, x, w, masks, sg, sw, sx, sg_inv, sw_inv, sx_inv, amax_g
return get_module("fp8_ops").linear_backward_fp8(
g, x, w, list(masks), sg, sw, sx, _fmt_int(fmt)
)
+3 -11
View File
@@ -10,15 +10,7 @@ Layout: x is packed [tokens, n_heads, head_dim] or dense
import torch
from astrai.extension.loader import _available, _modules
def _check_available():
if not _available.get("rotary_emb"):
raise RuntimeError(
"CUDA kernel 'rotary_emb' is not available. "
"Build with CSRC_KERNELS=true or use the torch fallback."
)
from astrai.extension.loader import get_module
def rotary_emb(x: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor:
@@ -31,9 +23,9 @@ def rotary_emb(x: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor:
Returns:
Tensor with the same shape as ``x``.
"""
_check_available()
mod = get_module("rotary_emb")
if not x.is_contiguous():
x = x.contiguous()
if not freqs_cis.is_contiguous():
freqs_cis = freqs_cis.contiguous()
return _modules["rotary_emb"].rotary_emb(x, freqs_cis)
return mod.rotary_emb(x, freqs_cis)
+2 -65
View File
@@ -12,45 +12,10 @@ Modules:
- engine.py: Facade (InferenceEngine)
"""
from astrai.inference.cache import (
Allocator,
KVCache,
KVStorage,
PagePool,
RadixCache,
ReqToTokenPool,
TaskCacheManager,
page_hash,
)
from astrai.inference.engine import InferenceEngine
from astrai.inference.network import (
AnthropicMessage,
BaseToolParser,
ChatCompletionRequest,
ChatMessage,
FunctionDef,
GenContext,
MessagesRequest,
ProtocolHandler,
SimpleJsonToolParser,
StopChecker,
ToolDef,
ToolParserFactory,
get_app,
run_server,
)
from astrai.inference.network.anthropic import AnthropicResponseBuilder
from astrai.inference.network.openai import OpenAIResponseBuilder
from astrai.inference.network import get_app, run_server
from astrai.inference.runtime.executor import Executor
from astrai.inference.runtime.sample import (
BaseSamplingStrategy,
FrequencyPenaltyStrategy,
SamplingPipeline,
TemperatureStrategy,
TopKStrategy,
TopPStrategy,
sample,
)
from astrai.inference.runtime.sample import sample
from astrai.inference.scheduler import InferenceScheduler
from astrai.inference.task import STOP, Task, TaskManager, TaskStatus
@@ -62,35 +27,7 @@ __all__ = [
"Task",
"TaskManager",
"TaskStatus",
"Allocator",
"KVCache",
"KVStorage",
"PagePool",
"RadixCache",
"ReqToTokenPool",
"TaskCacheManager",
"page_hash",
"sample",
"BaseSamplingStrategy",
"TemperatureStrategy",
"TopKStrategy",
"TopPStrategy",
"FrequencyPenaltyStrategy",
"SamplingPipeline",
"ProtocolHandler",
"StopChecker",
"GenContext",
"BaseToolParser",
"SimpleJsonToolParser",
"ToolParserFactory",
"OpenAIResponseBuilder",
"AnthropicResponseBuilder",
"ChatMessage",
"ChatCompletionRequest",
"FunctionDef",
"ToolDef",
"AnthropicMessage",
"MessagesRequest",
"get_app",
"run_server",
]
-10
View File
@@ -72,16 +72,6 @@ class KVStorage:
(n_layers, size, n_kv_heads, head_dim), device=device, dtype=dtype
)
def get_key_buffer(self, layer_id: int) -> Tensor:
return self.k_buffer[layer_id]
def get_value_buffer(self, layer_id: int) -> Tensor:
return self.v_buffer[layer_id]
def set_kv_buffer(self, layer_id: int, loc: Tensor, k: Tensor, v: Tensor) -> None:
self.k_buffer[layer_id, loc] = k
self.v_buffer[layer_id, loc] = v
@dataclass
class KVCache:
-2
View File
@@ -16,8 +16,6 @@ from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from typing import Callable, Dict, List, Optional, OrderedDict
import torch
from astrai.inference.cache.buffer import ReqToTokenPool
# ---- data contract: per-task slot state ----
+2 -3
View File
@@ -156,9 +156,8 @@ class InferenceEngine:
async def _agen():
loop = asyncio.get_event_loop()
while True:
try:
token = await loop.run_in_executor(None, next, sync_gen)
except StopIteration:
token = await loop.run_in_executor(None, next, sync_gen, None)
if token is None:
break
yield token
-22
View File
@@ -148,17 +148,6 @@ class MetricsCollector:
self._completed.append(timing)
self._accumulate(timing)
def clear(self):
"""Reset all state (e.g. on engine shutdown)."""
self._timings.clear()
self._completed.clear()
self._ttft_ms_sum = 0.0
self._ttft_ms_count = 0
self._decode_tps_sum = 0.0
self._decode_tps_count = 0
self._e2e_ms_sum = 0.0
self._e2e_ms_count = 0
# timing scopes
@contextmanager
@@ -180,17 +169,6 @@ class MetricsCollector:
t._decode_steps += 1
t._decode_total_s += dt
# access
def get_timing(self, task_id: str) -> Optional[TaskTiming]:
"""Return the timing record for *task_id* (active or completed)."""
if task_id in self._timings:
return self._timings[task_id]
for t in self._completed:
if t.task_id == task_id:
return t
return None
# aggregate stats
def get_stats(self) -> Dict[str, Any]:
+2 -2
View File
@@ -205,8 +205,8 @@ class Executor:
max_q_heads = config.num_attention_heads
head_dim = config.hidden_size // config.num_attention_heads
backend = get_backend()
self._graph_supported = backend.supports_graph() and CudaBackend.supports(
head_dim=head_dim
self._graph_supported = backend.supports_graph() and (
CudaBackend.available() and head_dim in CudaBackend.HEAD_DIMS
)
self._workspace = InferenceWorkspace(
max_batch_size=kv_cache.max_batch_size,
+22 -30
View File
@@ -79,29 +79,21 @@ class InferenceScheduler:
if backend is None:
self._backend = None
default_backend = get_backend()
self._backend_name = type(default_backend).__name__
with attn_backend(default_backend):
self._executor = Executor(
model=model,
kv_cache=self._cache,
task_cache=self._task_cache,
device=self.device,
dtype=self.dtype,
enable_cuda_graph=enable_cuda_graph,
)
active_backend = get_backend()
else:
with attn_backend(backend):
active_backend = backend
with attn_backend(active_backend):
if backend is not None:
self._backend = get_backend()
self._backend_name = type(self._backend).__name__
self._executor = Executor(
model=model,
kv_cache=self._cache,
task_cache=self._task_cache,
device=self.device,
dtype=self.dtype,
enable_cuda_graph=enable_cuda_graph,
)
self._backend_name = type(get_backend()).__name__
self._executor = Executor(
model=model,
kv_cache=self._cache,
task_cache=self._task_cache,
device=self.device,
dtype=self.dtype,
enable_cuda_graph=enable_cuda_graph,
)
self._stop_event = threading.Event()
self._loop_thread: Optional[threading.Thread] = None
@@ -283,12 +275,7 @@ class InferenceScheduler:
except Exception as e:
self._stop_event.set()
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
for task in self._task_mgr.get_active_tasks():
self._task_mgr.invoke_callback(task.task_id, STOP)
self._task_cache.task_free(task.task_id)
for task in self._task_mgr.get_waiting_tasks():
self._task_mgr.invoke_callback(task.task_id, STOP)
self._task_mgr.clear_queues()
self._abort_and_clear(free_waiting=False)
def start(self):
if self._loop_thread is not None and self._loop_thread.is_alive():
@@ -304,15 +291,20 @@ class InferenceScheduler:
if self._loop_thread is not None:
self._loop_thread.join(timeout=2.0)
self._loop_thread = None
self._abort_and_clear(free_waiting=True)
if torch.cuda.is_available():
torch.cuda.empty_cache()
def _abort_and_clear(self, free_waiting: bool):
"""Invoke STOP callbacks, release cache slots, and clear task queues."""
for task in self._task_mgr.get_active_tasks():
self._task_mgr.invoke_callback(task.task_id, STOP)
self._task_cache.task_free(task.task_id)
for task in self._task_mgr.get_waiting_tasks():
self._task_mgr.invoke_callback(task.task_id, STOP)
self._task_cache.task_free(task.task_id)
if free_waiting:
self._task_cache.task_free(task.task_id)
self._task_mgr.clear_queues()
if torch.cuda.is_available():
torch.cuda.empty_cache()
def run_batch(
self,
-7
View File
@@ -124,13 +124,6 @@ class InferenceWorkspace:
device=device,
)
def decode_buffers(self, batch: int, q_heads: int):
"""Return ``(o_part, ml_part)`` view sliced to live dimensions."""
return (
self.decode_o_part[:batch, :q_heads],
self.decode_ml_part[:batch, :q_heads],
)
def fill_input_ids(self, ids: "list[int]") -> Tensor:
"""Write ``ids`` into the device buffer and return ``[B]``.
+9 -1
View File
@@ -2,6 +2,13 @@ import logging
import os
class _DistributedContextFilter(logging.Filter):
def filter(self, record: logging.LogRecord) -> bool:
record.rank = os.environ.get("RANK", "0")
record.world_size = os.environ.get("WORLD_SIZE", "1")
return True
def setup_logging(level: str = "INFO"):
"""Attach a StreamHandler to the ``astrai`` logger (idempotent).
@@ -18,9 +25,10 @@ def setup_logging(level: str = "INFO"):
level_name = os.environ.get("ASTR_LOG_LEVEL", level).upper()
logger.setLevel(getattr(logging, level_name, logging.INFO))
handler = logging.StreamHandler()
handler.addFilter(_DistributedContextFilter())
handler.setFormatter(
logging.Formatter(
"%(asctime)s | %(levelname)-7s | %(name)s | %(message)s",
"%(asctime)s | %(levelname)-8s | rank=%(rank)2s/%(world_size)-2s | %(name)-32s | %(message)s",
datefmt="%Y-%m-%d %H:%M:%S",
)
)
+41 -7
View File
@@ -4,13 +4,21 @@ AutoModel base class for model loading and saving.
from contextlib import contextmanager
from pathlib import Path
from typing import Self, Union
from typing import Union
import torch.nn as nn
from astrai.config.model_config import BaseModelConfig, ConfigFactory
from astrai.factory import BaseFactory
from astrai.serialization import load_model_config, load_model_weights, save_model
from astrai.serialization import (
HF_MODEL_TYPES,
adapt_config,
convert_hf_weights,
load_model_config,
load_model_weights,
looks_like_hf_state_dict,
save_model,
)
@contextmanager
@@ -57,7 +65,25 @@ class AutoModel(nn.Module):
path: Union[str, Path],
disable_random_init: bool = True,
strict: bool = True,
weights_format: str = "auto",
) -> nn.Module:
"""Load a model directory.
Args:
path: Directory containing ``config.json`` and optionally
``model.safetensors``.
disable_random_init: Replace parameter initializers with no-ops
while building the model.
strict: Passed to ``load_state_dict``.
weights_format: ``"auto"`` detects HuggingFace checkpoints
(LLaMA-style keys and ``model_type``) and converts them;
``"astrai"`` skips conversion; ``"hf"`` forces it.
"""
if weights_format not in ("auto", "astrai", "hf"):
raise ValueError(
f"weights_format must be one of 'auto', 'astrai', 'hf', "
f"got {weights_format!r}"
)
model_path = Path(path)
@@ -66,6 +92,12 @@ class AutoModel(nn.Module):
raise FileNotFoundError(f"Config file not found: {config_path}")
raw = load_model_config(str(model_path))
is_hf_config = weights_format == "hf" or (
weights_format == "auto" and raw.get("model_type") in HF_MODEL_TYPES
)
if is_hf_config:
raw = adapt_config(raw)
config = ConfigFactory.load(raw)
model_type = config.model_type or "autoregressive_lm"
@@ -75,8 +107,14 @@ class AutoModel(nn.Module):
model = actual_cls(config)
weights_path = model_path / "model.safetensors"
if weights_path.exists():
index_path = model_path / "model.safetensors.index.json"
if weights_path.exists() or index_path.exists():
state_dict = load_model_weights(str(model_path))
is_hf_weights = is_hf_config or (
weights_format == "auto" and looks_like_hf_state_dict(state_dict)
)
if is_hf_weights:
state_dict = convert_hf_weights(state_dict, config)
model.load_state_dict(state_dict, strict=strict)
return model
@@ -90,7 +128,3 @@ class AutoModel(nn.Module):
state_dict=self.state_dict(),
save_directory=str(save_directory),
)
def to(self, *args, **kwargs) -> Self:
"""Move model to device/dtype."""
return super().to(*args, **kwargs)
+1 -7
View File
@@ -29,12 +29,6 @@ class FFNOutput(TypedDict):
router_stats: Optional[RouterStats]
class RoutedOutput(TypedDict):
hidden_states: Tensor
aux_loss: Optional[Tensor]
router_stats: Optional[RouterStats]
@FFNFactory.register("mlp")
class MLP(nn.Module):
def __init__(self, dim: int, dim_ffn: int, down_init_std: float = 0.02):
@@ -122,7 +116,7 @@ class DeepSeekMoE(nn.Module):
/ self.n_shared_experts
)
def _routed_forward(self, x: Tensor, include_aux_loss: bool) -> RoutedOutput:
def _routed_forward(self, x: Tensor, include_aux_loss: bool) -> FFNOutput:
N, D = x.shape
K = self.n_activated_experts
E = self.n_routed_experts
+7 -11
View File
@@ -247,18 +247,15 @@ class LocalStrategy(LaunchStrategy):
ctx.join()
def _detect_launcher() -> str:
"""Detect the distributed launcher from environment.
Returns one of: "torchelastic", "torchrun", "external", "local".
"""
def _is_external_launcher() -> bool:
"""Whether an external launcher (torchrun/elastic/manual env) started us."""
if dist.is_torchelastic_launched():
return "torchelastic"
return True
if "LOCAL_WORLD_SIZE" in os.environ:
return "torchrun"
return True
if "RANK" in os.environ and "WORLD_SIZE" in os.environ:
return "external"
return "local"
return True
return False
def spawn_parallel_fn(
@@ -273,8 +270,7 @@ def spawn_parallel_fn(
):
if master_port is None:
master_port = find_free_port()
launcher = _detect_launcher()
if launcher in ("torchelastic", "torchrun", "external"):
if _is_external_launcher():
strategy = TorchrunStrategy(
world_size, backend, master_addr, master_port, device_type, start_method
)
+12
View File
@@ -22,9 +22,21 @@ from astrai.serialization.dataset import (
load_bin_offsets,
save_bin,
)
from astrai.serialization.hf_adapter import (
HF_MODEL_TYPES,
adapt_config,
convert_hf_config,
convert_hf_weights,
looks_like_hf_state_dict,
)
__all__ = [
"Checkpoint",
"HF_MODEL_TYPES",
"adapt_config",
"convert_hf_config",
"convert_hf_weights",
"looks_like_hf_state_dict",
"load_json",
"load_model_config",
"load_model_weights",
+31 -23
View File
@@ -5,7 +5,7 @@ import json
import time
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Dict, Optional, Union
from typing import Any, Callable, Dict, Optional, Union
import safetensors.torch as st
import torch
@@ -22,39 +22,31 @@ def save_safetensors(state_dict: dict, path: Union[str, Path]):
st.save_file(state_dict, str(path))
def load_safetensors(path: Union[str, Path], broadcast: bool = False) -> dict:
def _broadcast_load(loader: Callable[[], dict], broadcast: bool) -> dict:
"""Load on rank 0 and broadcast the object to all ranks."""
if not broadcast or not dist.is_initialized():
return st.load_file(str(path))
return loader()
rank = get_rank()
if rank == 0:
state_dict = st.load_file(str(path))
data = loader()
else:
state_dict = {}
tmp = [state_dict]
data = {}
tmp = [data]
dist.broadcast_object_list(tmp, src=0)
return tmp[0]
def load_safetensors(path: Union[str, Path], broadcast: bool = False) -> dict:
return _broadcast_load(lambda: st.load_file(str(path)), broadcast)
def save_json(data: dict, path: Union[str, Path]):
with open(str(path), "w") as f:
json.dump(data, f, indent=2)
def load_json(path: Union[str, Path], broadcast: bool = False) -> dict:
if not broadcast or not dist.is_initialized():
with open(str(path), "r") as f:
return json.load(f)
rank = get_rank()
if rank == 0:
with open(str(path), "r") as f:
data = json.load(f)
else:
data = {}
tmp = [data]
dist.broadcast_object_list(tmp, src=0)
return tmp[0]
return _broadcast_load(lambda: json.loads(Path(path).read_text()), broadcast)
def save_torch(obj: Any, path: Union[str, Path]):
@@ -99,7 +91,21 @@ def load_model_config(save_directory: str) -> dict:
def load_model_weights(save_directory: str) -> dict:
return load_state_dict(Path(save_directory) / _WEIGHTS_FILE)
save_path = Path(save_directory)
weights_file = save_path / _WEIGHTS_FILE
if weights_file.exists():
return load_state_dict(weights_file)
index_path = save_path / "model.safetensors.index.json"
if index_path.exists():
index = load_json(index_path)
weight_map = index.get("weight_map", {})
state_dict = {}
for shard in sorted(set(weight_map.values())):
state_dict.update(load_state_dict(save_path / shard))
return state_dict
raise FileNotFoundError(f"No model weights found in {save_directory}")
def load_state_dict(path: Union[str, Path], broadcast: bool = False) -> dict:
@@ -190,8 +196,10 @@ class Checkpoint:
if meta_path.exists():
return cls.load(save_dir, broadcast=broadcast)
if weights_path.exists():
state_dict = load_state_dict(weights_path, broadcast=broadcast)
weights_path = save_path / _WEIGHTS_FILE
index_path = save_path / "model.safetensors.index.json"
if weights_path.exists() or index_path.exists():
state_dict = load_model_weights(save_dir)
config = {}
config_path = save_path / _CONFIG_FILE
if config_path.exists():
+271
View File
@@ -0,0 +1,271 @@
"""HuggingFace checkpoint adaptation for LLaMA-style decoder models.
AstrAI stores weights with its own key names (``layers.<i>.input_norm``,
``layers.<i>.mlp.gate``), while HuggingFace decoder-only checkpoints use
``model.layers.<i>.input_layernorm`` / ``model.layers.<i>.mlp.gate_proj``.
This module translates HF configs and state dicts so external checkpoints
can be loaded directly.
Supported families (LLaMA layout, dense and MoE):
- dense FFN: llama, mistral, qwen2, gemma, gemma2, phi3
- MoE FFN (Mixtral / Qwen2-MoE / DeepSeek-V3 layout): router
``mlp.gate``, routed experts ``mlp.experts.<j>``, shared experts
``mlp.shared_experts.<j>``
Not supported:
- MLA attention (DeepSeek-V2/V3 ``kv_a_proj_with_mqa``) uses a different
KV factorization and cannot be converted numerically.
- Attention/MLP bias (``attention_bias`` / ``mlp_bias``) AstrAI
projections are bias-free.
"""
import logging
import re
from typing import Any, Dict, Mapping
import torch
from astrai.config.base import BaseConfig
logger = logging.getLogger(__name__)
HF_MODEL_TYPES = frozenset(
{
"llama",
"mistral",
"mixtral",
"qwen2",
"qwen2_moe",
"gemma",
"gemma2",
"phi3",
}
)
_EMBED = re.compile(r"^model\.embed_tokens\.weight$")
_ATTN = re.compile(r"^model\.layers\.(\d+)\.self_attn\.(q|k|v|o)_proj\.(weight|bias)$")
_Q_NORM = re.compile(r"^model\.layers\.(\d+)\.self_attn\.q_norm\.weight$")
_K_NORM = re.compile(r"^model\.layers\.(\d+)\.self_attn\.k_norm\.weight$")
_INPUT_NORM = re.compile(r"^model\.layers\.(\d+)\.input_layernorm\.weight$")
_POST_NORM = re.compile(r"^model\.layers\.(\d+)\.post_attention_layernorm\.weight$")
_FINAL_NORM = re.compile(r"^model\.norm\.weight$")
_LM_HEAD = re.compile(r"^lm_head\.weight$")
_DENSE_MLP = re.compile(
r"^model\.layers\.(\d+)\.mlp\.(gate|up|down)_proj\.(weight|bias)$"
)
_MOE_ROUTER = re.compile(r"^model\.layers\.(\d+)\.mlp\.gate\.weight$")
_MOE_EXPERTS = re.compile(
r"^model\.layers\.(\d+)\.mlp\.experts\.(\d+)\.(gate|up|down)_proj\.(weight|bias)$"
)
_MOE_SHARED = re.compile(
r"^model\.layers\.(\d+)\.mlp\.shared_expert(?:s)?\.(\d+)\."
r"(gate|up|down)_proj\.(weight|bias)$"
)
_ASTR_PREFIXES = ("embed_tokens.", "layers.", "norm.", "lm_head.")
def looks_like_hf_state_dict(state_dict: Mapping[str, Any]) -> bool:
"""Return True if *state_dict* uses HuggingFace key names."""
return any(
key.startswith("model.")
or "self_attn." in key
or "input_layernorm" in key
or "mlp.experts." in key
for key in state_dict
)
def _is_dense_mlp_layer(config: BaseConfig, layer_id: int) -> bool:
"""Return whether a layer uses dense MLP instead of routed experts."""
if getattr(config, "ffn_type", "mlp") != "moe":
return True
mlp_only = getattr(config, "mlp_only_layers", None) or []
if layer_id in mlp_only:
return True
step = getattr(config, "decoder_sparse_step", 1) or 1
return step > 1 and (layer_id + 1) % step != 0
def adapt_config(raw: Dict[str, Any]) -> Dict[str, Any]:
"""Translate *raw* for AstrAI if it looks like an HF model config."""
if raw.get("model_type") in HF_MODEL_TYPES:
return convert_hf_config(raw)
return raw
def convert_hf_config(raw: Dict[str, Any]) -> Dict[str, Any]:
"""Convert an HF LLaMA-style config dict to AstrAI field names."""
if raw.get("attention_bias") or raw.get("mlp_bias"):
raise NotImplementedError(
"attention_bias / mlp_bias checkpoints are not supported; "
"AstrAI projections are bias-free"
)
cfg: Dict[str, Any] = {}
for key in (
"vocab_size",
"hidden_size",
"num_hidden_layers",
"intermediate_size",
"rms_norm_eps",
"tie_word_embeddings",
"max_position_embeddings",
"rope_theta",
"rope_scaling",
"num_attention_heads",
"num_key_value_heads",
"use_qk_norm",
"use_gated_attention",
"kv_lora_rank",
"qk_nope_head_dim",
"qk_rope_head_dim",
"moe_intermediate_size",
"shared_expert_intermediate_size",
"topk_method",
"norm_topk_prob",
"moe_aux_loss_coef",
"decoder_sparse_step",
"mlp_only_layers",
"neftune_alpha",
):
if key in raw:
cfg[key] = raw[key]
if "qk_norm" in raw and "use_qk_norm" not in cfg:
cfg["use_qk_norm"] = raw["qk_norm"]
if (
raw.get("model_type") in ("gemma", "gemma2")
and "use_qk_norm" not in cfg
and "qk_norm" not in raw
):
# Gemma/Gemma2 always apply RMSNorm to Q and K before attention.
cfg["use_qk_norm"] = True
n_heads = raw.get("num_attention_heads")
if cfg.get("num_key_value_heads") is None and n_heads is not None:
cfg["num_key_value_heads"] = n_heads
if raw.get("head_dim") is not None and n_heads and raw.get("hidden_size"):
expected = raw["hidden_size"] // n_heads
if raw["head_dim"] != expected:
raise NotImplementedError(
f"HF head_dim={raw['head_dim']} differs from the computed "
f"head dim {expected}; AstrAI derives head_dim from "
"hidden_size / num_attention_heads"
)
if "kv_lora_rank" in raw:
cfg["attn_type"] = "mla"
n_experts = raw.get("num_local_experts") or raw.get("n_routed_experts")
if n_experts:
cfg["ffn_type"] = "moe"
cfg["n_routed_experts"] = n_experts
if "num_experts_per_tok" in raw:
cfg["n_activated_experts"] = raw["num_experts_per_tok"]
if "n_activated_experts" in raw:
cfg["n_activated_experts"] = raw["n_activated_experts"]
if "n_shared_experts" in raw:
cfg["n_shared_experts"] = raw["n_shared_experts"]
else:
# Mixtral has no shared experts; AstrAI defaults to one.
cfg["n_shared_experts"] = 0
if cfg.get("moe_intermediate_size") is None and "intermediate_size" in raw:
# MoE configs store the per-expert FFN size in intermediate_size.
cfg["moe_intermediate_size"] = raw["intermediate_size"]
first_k_dense = raw.get("first_k_dense_replace")
if isinstance(first_k_dense, int) and first_k_dense > 0:
cfg["mlp_only_layers"] = list(range(first_k_dense))
cfg["decoder_sparse_step"] = 1
cfg["model_type"] = "autoregressive_lm"
return cfg
def convert_hf_weights(
state_dict: Mapping[str, Any],
config: BaseConfig,
) -> Dict[str, torch.Tensor]:
"""Rename HF state dict keys to AstrAI names.
Keys that are already AstrAI-style pass through unchanged; unmapped
HF keys are dropped with a warning. Use with ``strict=True`` to fail
loudly when the checkpoint does not match the config.
"""
if getattr(config, "attn_type", "gqa") == "mla":
if any("kv_a_proj_with_mqa" in key for key in state_dict):
raise NotImplementedError(
"MLA attention (DeepSeek-V2/V3 kv_a_proj_with_mqa) uses a "
"different KV factorization and cannot be converted"
)
ffn_type = getattr(config, "ffn_type", "mlp")
converted: Dict[str, torch.Tensor] = {}
skipped: list[str] = []
for key, tensor in state_dict.items():
if key.startswith(_ASTR_PREFIXES):
converted[key] = tensor
continue
new_key = None
if ffn_type == "moe":
m = _MOE_ROUTER.match(key)
if m:
new_key = f"layers.{m.group(1)}.mlp.router.weight"
else:
m = _MOE_EXPERTS.match(key)
if m:
new_key = (
f"layers.{m.group(1)}.mlp.routed_experts.{m.group(2)}."
f"{m.group(3)}.{m.group(4)}"
)
else:
m = _MOE_SHARED.match(key)
if m:
new_key = (
f"layers.{m.group(1)}.mlp.shared_experts.{m.group(2)}."
f"{m.group(3)}.{m.group(4)}"
)
if new_key is None:
m = _DENSE_MLP.match(key)
if m and _is_dense_mlp_layer(config, int(m.group(1))):
new_key = f"layers.{m.group(1)}.mlp.{m.group(2)}.{m.group(3)}"
else:
m = _DENSE_MLP.match(key)
if m:
new_key = f"layers.{m.group(1)}.mlp.{m.group(2)}.{m.group(3)}"
if new_key is None:
m = _ATTN.match(key)
if m:
new_key = (
f"layers.{m.group(1)}.attention.{m.group(2)}_proj.{m.group(3)}"
)
elif (m := _Q_NORM.match(key)) is not None:
new_key = f"layers.{m.group(1)}.attention.q_norm.weight"
elif (m := _K_NORM.match(key)) is not None:
new_key = f"layers.{m.group(1)}.attention.k_norm.weight"
elif (m := _INPUT_NORM.match(key)) is not None:
new_key = f"layers.{m.group(1)}.input_norm.weight"
elif (m := _POST_NORM.match(key)) is not None:
new_key = f"layers.{m.group(1)}.post_attention_norm.weight"
elif (m := _EMBED.match(key)) is not None:
new_key = "embed_tokens.weight"
elif (m := _FINAL_NORM.match(key)) is not None:
new_key = "norm.weight"
elif (m := _LM_HEAD.match(key)) is not None:
new_key = "lm_head.weight"
if new_key is None:
skipped.append(key)
else:
converted[new_key] = tensor
if skipped:
logger.warning(
"Dropped %d unmapped HuggingFace weight key(s): %s",
len(skipped),
", ".join(sorted(skipped)[:10]),
)
return converted
+2 -18
View File
@@ -94,21 +94,5 @@ def ctx_get_grad_snr(ctx):
return tracker.snr
def ctx_get_moe_aux_loss(ctx):
return ctx.strategy._moe_metrics.get("aux_loss")
def ctx_get_router_entropy(ctx):
return ctx.strategy._moe_metrics.get("router_entropy")
def ctx_get_dead_expert_fraction(ctx):
return ctx.strategy._moe_metrics.get("dead_expert_fraction")
def ctx_get_load_imbalance_mean(ctx):
return ctx.strategy._moe_metrics.get("load_imbalance_mean")
def ctx_get_load_imbalance_max(ctx):
return ctx.strategy._moe_metrics.get("load_imbalance_max")
def ctx_get_moe_metric(ctx, key):
return ctx.strategy._moe_metrics.get(key)
+1 -7
View File
@@ -2,7 +2,7 @@
import math
from abc import ABC, abstractmethod
from typing import Any, Dict, List
from typing import List
from torch.optim.lr_scheduler import LRScheduler
@@ -20,12 +20,6 @@ class BaseScheduler(LRScheduler, ABC):
"""Calculate the current learning rate."""
raise NotImplementedError
def state_dict(self) -> Dict[str, Any]:
return super().state_dict()
def load_state_dict(self, state_dict: Dict[str, Any]):
super().load_state_dict(state_dict)
class SchedulerFactory(BaseFactory["BaseScheduler"]):
"""Factory class for creating learning rate schedulers.
+3 -16
View File
@@ -1,6 +1,6 @@
"""Training strategy implementations with factory pattern."""
from abc import ABC, abstractmethod
from abc import ABC
from typing import Callable, Dict, List, Optional, TypedDict, Union
import torch
@@ -184,10 +184,9 @@ class BaseStrategy(ABC):
self.executor = kwargs.pop("executor", None)
self.moe_aux_loss_coef = kwargs.pop("moe_aux_loss_coef", 0.01)
self._moe_metrics: Dict[str, float] = {}
self.extra_kwargs = kwargs
self.strategy_kwargs = kwargs
self._rollout_runner = None
@abstractmethod
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
"""Compute loss for the given batch.
@@ -197,7 +196,7 @@ class BaseStrategy(ABC):
Returns:
Computed loss tensor
"""
raise NotImplementedError
return self.compute_loss_output(batch)["loss"]
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
return self._normalize_output(self.compute_loss(batch))
@@ -328,9 +327,6 @@ class SEQStrategy(BaseStrategy):
super().__init__(model, device, **kwargs)
self.label_smoothing = label_smoothing
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)
input_ids, target_ids = batch["input_ids"], batch["target_ids"]
@@ -369,9 +365,6 @@ class SFTStrategy(BaseStrategy):
super().__init__(model, device, **kwargs)
self.label_smoothing = label_smoothing
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)
input_ids, target_ids, position_ids, loss_mask = (
@@ -426,9 +419,6 @@ class DPOStrategy(BaseStrategy):
self.beta = beta
self.reduction = reduction
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)
chosen_ids, rejected_ids = batch["chosen"], batch["rejected"]
@@ -553,9 +543,6 @@ class GRPOStrategy(BaseStrategy):
if state_dict is not None:
self.old_model.load_state_dict(state_dict)
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)
prompts = batch["prompts"]
+11 -10
View File
@@ -3,6 +3,7 @@ import logging
import os
import sys
import time
from functools import partial
from pathlib import Path
from typing import IO, Callable, List, Optional, Protocol, runtime_checkable
@@ -17,15 +18,11 @@ from astrai.parallel import only_on_rank
from astrai.parallel.setup import get_current_device
from astrai.serialization import Checkpoint
from astrai.trainer.metric_util import (
ctx_get_dead_expert_fraction,
ctx_get_grad_norm,
ctx_get_grad_snr,
ctx_get_load_imbalance_max,
ctx_get_load_imbalance_mean,
ctx_get_loss,
ctx_get_lr,
ctx_get_moe_aux_loss,
ctx_get_router_entropy,
ctx_get_moe_metric,
ctx_get_val_loss,
)
from astrai.trainer.train_context import TrainContext
@@ -262,11 +259,15 @@ class MetricCallback(TrainCallback):
"val_loss": ctx_get_val_loss,
"grad_norm": ctx_get_grad_norm,
"grad_snr": ctx_get_grad_snr,
"moe_aux_loss": ctx_get_moe_aux_loss,
"router_entropy": ctx_get_router_entropy,
"dead_expert_fraction": ctx_get_dead_expert_fraction,
"load_imbalance_mean": ctx_get_load_imbalance_mean,
"load_imbalance_max": ctx_get_load_imbalance_max,
"moe_aux_loss": partial(ctx_get_moe_metric, key="aux_loss"),
"router_entropy": partial(ctx_get_moe_metric, key="router_entropy"),
"dead_expert_fraction": partial(
ctx_get_moe_metric, key="dead_expert_fraction"
),
"load_imbalance_mean": partial(
ctx_get_moe_metric, key="load_imbalance_mean"
),
"load_imbalance_max": partial(ctx_get_moe_metric, key="load_imbalance_max"),
}
def _metrics(self, context: TrainContext, names):
+23 -6
View File
@@ -8,6 +8,7 @@ import torch
import torch.nn as nn
from torch.utils.data import DataLoader, random_split
from astrai.config.model_config import ConfigFactory
from astrai.config.train_config import TrainConfig
from astrai.dataset import RDSampler
from astrai.inference.scheduler import InferenceScheduler
@@ -15,7 +16,13 @@ from astrai.model.components.lora import inject_lora
from astrai.parallel.executor import BaseExecutor, ExecutorFactory, create_ref_model
from astrai.parallel.setup import get_current_device, get_rank, get_world_size
from astrai.protocols import OptimizerProtocol, SchedulerProtocol
from astrai.serialization import Checkpoint, load_json
from astrai.serialization import (
Checkpoint,
adapt_config,
convert_hf_weights,
load_json,
looks_like_hf_state_dict,
)
from astrai.tokenize import AutoTokenizer
from astrai.trainer.metric_util import GradSNRTracker
from astrai.trainer.rollout import RolloutGenerator, RolloutRunner
@@ -126,9 +133,18 @@ class TrainContextBuilder:
if self._param_path:
config_path = Path(self._param_path) / "config.json"
if config_path.exists():
state.model_config = load_json(config_path)
state.model_config = adapt_config(load_json(config_path))
checkpoint = Checkpoint.load_any(self._param_path)
if checkpoint is not None:
if checkpoint.config:
checkpoint.config = adapt_config(checkpoint.config)
if checkpoint.state_dict and looks_like_hf_state_dict(
checkpoint.state_dict
):
checkpoint.state_dict = convert_hf_weights(
checkpoint.state_dict,
ConfigFactory.load(checkpoint.config or state.model_config),
)
state.state_dict = checkpoint.state_dict
state.model_config = checkpoint.config or state.model_config
if self._resume:
@@ -140,8 +156,10 @@ class TrainContextBuilder:
checkpoint.consumed_samples // per_step * per_step
)
state.checkpoint = checkpoint
if not state.model_config and hasattr(cfg.model_fn(), "config"):
state.model_config = cfg.model_fn().config.to_dict()
if not state.model_config:
model = cfg.model_fn()
if hasattr(model, "config"):
state.model_config = model.config.to_dict()
return state
def _create_context(
@@ -204,7 +222,6 @@ class TrainContextBuilder:
def _create_dataloaders(
self, context: TrainContext, train_dataset, val_dataset
) -> None:
cfg = self.config
sampler_offset = context.consumed_samples // context.world_size
if self._resume and sampler_offset > 0:
samples_per_replica = (
@@ -261,7 +278,7 @@ class TrainContextBuilder:
def _create_strategy(self, context: TrainContext, executor: BaseExecutor) -> dict:
cfg = self.config
kwargs = dict(cfg.extra_kwargs)
kwargs = dict(cfg.strategy_kwargs)
kwargs.setdefault("moe_aux_loss_coef", cfg.moe_aux_loss_coef)
if cfg.strategy in ("dpo", "grpo", "online_grpo", "online_dpo"):
kwargs["ref_model"] = create_ref_model(
+26 -3
View File
@@ -48,10 +48,33 @@ set(TORCH_LIBS
set(CMAKE_CUDA_ARCHITECTURES "${ASTRAI_CUDA_ARCH}")
set(KERNELS attn_decode attn_prefill attn_paged_decode attn_paged_prefill rotary_emb fp8_mm)
# Kernel registry parallel lists of module names (.so / pybind names,
# globally unique across families) and their per-family source paths under
# kernels/. `loader.py` auto-discovers the .so files in astrai/extension/lib/,
# so this CMake registry is the single place to register a new kernel.
set(KERNEL_NAMES
attn_decode
attn_prefill
attn_paged_decode
attn_paged_prefill
rotary_emb
fp8_ops
)
set(KERNEL_SRCS
attention/decode.cu
attention/prefill.cu
attention/paged_decode.cu
attention/paged_prefill.cu
rotary/rotary_emb.cu
fp8/ops.cu
)
foreach(name ${KERNELS})
add_library(${name} MODULE "${CMAKE_CURRENT_SOURCE_DIR}/kernels/${name}.cu")
list(LENGTH KERNEL_NAMES _kernel_count)
math(EXPR _kernel_last "${_kernel_count} - 1")
foreach(i RANGE ${_kernel_last})
list(GET KERNEL_NAMES ${i} name)
list(GET KERNEL_SRCS ${i} src)
add_library(${name} MODULE "${CMAKE_CURRENT_SOURCE_DIR}/kernels/${src}")
target_compile_definitions(${name} PRIVATE TORCH_EXTENSION_NAME=${name})
+1 -1
View File
@@ -1,2 +1,2 @@
# Source directory for CUDA kernels — build-time only.
# Compiled .so files live in astrAI/_ext/.
# Compiled .so files live in astrai/extension/lib/ (see csrc/CMakeLists.txt).
+95
View File
@@ -0,0 +1,95 @@
#pragma once
// Pure POD header
namespace astrai {
namespace attention {
// 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]
};
// Split-KV workspace cap: max decode splits per (batch, q_head).
constexpr int MAX_SPLITS = 32;
// Unified attention params covering BOTH addressing modes:
// - Contiguous K/V: dense [batch, kv_head, kv_len, head_dim] tensors (k/v).
// - Paged (SGLang-style): flat pool [size, kv_head, head_dim] + req_to_token.
// Each kernel selects the addressing via a KVSource policy (see
// layout_policies.cuh); a given call only touches the fields of one mode, so
// this is a POD shared by both paths rather than two parallel structs that
// drift out of sync.
//
// Pointer/flag members carry default member initializers: the pointers gate
// optional paths via null checks (new_k_ptr, mask, o_part, ...), so a stack
// `AttentionParams<T> p;` left partially packed must never see garbage
// non-null pointers or a garbage use_mask/causal_offset — that class of bug
// reads through wild addresses. NSDMI keeps the struct an aggregate (C++17)
// and trivially copyable, so `= {}`, memcpy-style packing and by-value kernel
// params all behave exactly as before.
template<typename T, typename AT = float>
struct AttentionParams {
// Shape
int batch;
int q_head;
int kv_head;
int head_dim;
int q_len; // Per-request in contiguous mode; total_q in paged mode.
int kv_len; // Contiguous mode; paged mode uses kv_indptr.
// Attention behavior
float scale;
// -1 = non-causal; >=0 = absolute position of first Q token
int causal_offset = -1;
int use_mask = 0;
// pointers
const T* __restrict__ q_ptr = nullptr;
const T* __restrict__ k_ptr = nullptr;
const T* __restrict__ v_ptr = nullptr;
const T* __restrict__ new_k_ptr = nullptr;
const T* __restrict__ new_v_ptr = nullptr;
T* __restrict__ o_ptr = nullptr;
const bool* __restrict__ mask = nullptr;
// strides
int q_b_stride;
int q_h_stride;
int q_l_stride;
int q_d_stride;
int kv_b_stride;
int kv_h_stride;
int kv_l_stride;
int kv_d_stride;
int new_kv_b_stride;
int new_kv_h_stride;
int mask_b_stride;
int mask_h_stride;
int mask_l_stride;
// Paged K/V addressing
const int* __restrict__ req_to_token = nullptr; // [num_reqs, max_context_len]
const int* __restrict__ req_pool_indices = nullptr; // [batch]
const int* __restrict__ kv_indptr = nullptr; // [batch + 1]
const int* __restrict__ qo_indptr = nullptr; // [batch + 1] or nullptr for decode
const int* __restrict__ q_tile_to_batch = nullptr; // [num_q_tiles], prefill only
const int* __restrict__ q_tile_to_index = nullptr; // [num_q_tiles], prefill only
int num_q_tiles;
int max_context_len; // req_to_token stride (dim 1)
// Decode split-KV workspace
int num_splits;
AT* __restrict__ o_part = nullptr;
AT* __restrict__ ml_part = nullptr;
};
} // namespace attention
} // namespace astrai
@@ -1,5 +1,7 @@
#include "attn_dispatchers.cuh"
#include "attn_entry_utils.cuh"
#include "dispatchers.cuh"
#include "entry_utils.cuh"
using namespace astrai::attention;
torch::Tensor attn_decode(
torch::Tensor q,
@@ -1,9 +1,13 @@
#pragma once
#include <cuda_bf16.h>
#include <float.h>
#include "attn_common.h"
#include "attn_layout_policies.cuh"
#include "attn_warp_utils.cuh"
#include "common.h"
#include "layout_policies.cuh"
#include "../common/reduce.cuh"
namespace astrai {
namespace attention {
constexpr int DC_CHUNK = 64;
// Scalar split-KV decode (fallback for sm < 80, no tensor cores), unified
@@ -142,3 +146,6 @@ __global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
int o_off = KV::q_decode_base(p, batch, q_head) + d * p.q_d_stride;
p.o_ptr[o_off] = __float2bfloat16(acc * inv);
}
} // namespace attention
} // namespace astrai
@@ -1,10 +1,12 @@
#pragma once
#include <cfloat>
#include <cuda_bf16.h>
#include "attn_common.h"
#include "attn_layout_policies.cuh"
#include "attn_mma_utils.cuh"
#include "attn_warp_utils.cuh"
#include "common.h"
#include "layout_policies.cuh"
#include "mma_utils.cuh"
namespace astrai {
namespace attention {
// Split-K (FlashDecoding) tensor-core decode via GQA head-packing, unified
// across contiguous and paged (SGLang flat-pool) K/V via the KV template
@@ -78,10 +80,10 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
KVAddr a = KV::template decode_addr<Traits::VEC>(
p, kctx, batch, kv_head, kc, d, valid, pass == 0);
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
cp_async_16_pred(&dK[off], a.k, a.valid);
cp_async_16_pred(&dV[off], a.v, a.valid);
astrai::cp_async_16(&dK[off], a.k, a.valid);
astrai::cp_async_16(&dV[off], a.v, a.valid);
}
cp_async_commit();
astrai::cp_async_commit_group();
};
// ---- Multi-stage cp.async pipeline ----
@@ -126,9 +128,9 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
for (int it = 0; it < ntiles; it++) {
if (it + 1 == ntiles)
cp_async_wait_group<0>();
astrai::cp_async_wait_group<0>();
else
cp_async_wait_group<STAGES - 1>();
astrai::cp_async_wait_group<STAGES - 1>();
__syncwarp();
process_tile(it, it & (STAGES - 1));
__syncwarp();
@@ -139,7 +141,7 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
// Fewer tiles than stages: load all, wait for all, process.
for (int i = 0; i < ntiles; i++)
load_tile(ti_begin + i, i);
cp_async_wait_group<0>();
astrai::cp_async_wait_all();
__syncwarp();
for (int it = 0; it < ntiles; it++)
process_tile(it, it);
@@ -181,3 +183,6 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
}
}
}
} // namespace attention
} // namespace astrai
@@ -3,22 +3,24 @@
// No torch dependency; pure CUDA.
//
// The paged and contiguous kernels are unified by the KVSource policy
// (ContigKV / PagedKV from attn_layout_policies.cuh), so each launcher struct
// (ContigKV / PagedKV from layout_policies.cuh), so each launcher struct
// below is templated on KV and the paged dispatch is just the same launcher
// instantiated with PagedKV. Only the grid/split math differs, and that is
// covered by KV::host_q_len / KV::host_kv_len.
#include <cuda_runtime.h>
#include <algorithm>
#include "attn_warp_utils.cuh"
#include "attn_layout_policies.cuh"
#include "attn_prefill_split_q.cuh"
#include "attn_decode_split_kv.cuh"
#include "layout_policies.cuh"
#include "prefill_split_q.cuh"
#include "decode_split_kv.cuh"
#ifndef ASTRAI_NO_MMA
#include "attn_prefill_split_q_mma.cuh"
#include "attn_decode_split_kv_mma.cuh"
#include "prefill_split_q_mma.cuh"
#include "decode_split_kv_mma.cuh"
#endif
namespace astrai {
namespace attention {
// 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,
// avoiding excessive loop/prologue overhead when tiles are small.
@@ -231,3 +233,6 @@ static inline void dispatch_paged_decode(AttentionParams<bf16>& p, cudaStream_t
attn_decode_combine_kernel<PagedKV><<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p);
}
} // namespace attention
} // namespace astrai
@@ -2,10 +2,7 @@
#include <float.h>
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include "attn_common.h"
#include "attn_warp_utils.cuh"
using bf16 = __nv_bfloat16;
#include "common.h"
// Dispatch head_dim: shared macro — avoids C++20 lambda template syntax.
// Usage: DISPATCH_HEAD_DIM(hd, fn, args...)
@@ -21,6 +18,11 @@ using bf16 = __nv_bfloat16;
" (supported: 32, 64, 128, 256)"); \
}
namespace astrai {
namespace attention {
using bf16 = __nv_bfloat16;
// The split kernel unconditionally writes every (batch, q_head, split) slot it
// owns — including empty split ranges, which store m = -FLT_MAX so the combine
// skips them. Allocators are therefore left uninitialized (torch::empty); the
@@ -356,3 +358,6 @@ inline void attn_pack_paged_prefill_params(
p.o_part = nullptr;
p.ml_part = nullptr;
}
} // namespace attention
} // namespace astrai
@@ -1,6 +1,6 @@
#pragma once
#include <cuda_bf16.h>
#include "attn_common.h"
#include "common.h"
// ============================================================================
// Attention layout policies keep Q scheduling independent from K/V storage.
@@ -26,6 +26,9 @@
#define DEVICE_FORCEINLINE static __device__ __forceinline__
#define HOST_DEV_FORCEINLINE static __host__ __device__ __forceinline__
namespace astrai {
namespace attention {
using bf16 = __nv_bfloat16;
// ============================================================================
@@ -253,3 +256,6 @@ struct PagedKV {
return kv_addr_from_token(p, c, token, d);
}
};
} // namespace attention
} // namespace astrai
@@ -3,12 +3,18 @@
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include "../common/cp_async.cuh"
#include "../common/mma.cuh"
// Predicated cp.async (4-operand form) requires CUDA 11.2+.
// bf16 mma.sync requires sm_80+ (guarded at build time by ASTRAI_NO_MMA).
#if CUDART_VERSION < 11020
#error "AstrAI CUDA kernels require CUDA 11.2 or later (CUDART_VERSION >= 11020)."
#endif
namespace astrai {
namespace attention {
// ============================================================================
// KernelTraits — FlashAttention-v2 style compile-time configuration bundle.
//
@@ -24,10 +30,10 @@ struct KernelTraits {
static constexpr int BR = 16; // Q rows per warp (mma M=16)
// Derived: mma.sync.m16n8k16 tile counts
static constexpr int KD = HEAD_DIM / 16; // Q/K k-slides
// Derived: mma tile counts from the shared mma_shape (m16n8k16 for bf16)
static constexpr int KD = HEAD_DIM / astrai::mma_shape<bf16>::k; // Q/K k-slides
static constexpr int NC8 = BC / 8; // S n-tiles (N=8)
static constexpr int KT2 = BC / 16; // P k-tiles (K=16)
static constexpr int KT2 = BC / astrai::mma_shape<bf16>::k; // P k-tiles (K=16)
static constexpr int DN8 = HEAD_DIM / 8; // O n-tiles (N=8)
static constexpr int LD = HEAD_DIM; // smem leading dim
@@ -43,16 +49,7 @@ struct KernelTraits {
// ---- PTX wrappers ----
using bf16 = __nv_bfloat16;
__device__ __forceinline__ void mma16816(float* d, const unsigned* a,
const unsigned* b, const float* c) {
asm volatile(
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%10,%11,%12,%13};"
: "=f"(d[0]), "=f"(d[1]), "=f"(d[2]), "=f"(d[3])
: "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]),
"f"(c[0]), "f"(c[1]), "f"(c[2]), "f"(c[3]));
}
// bf16 mma.sync lives in the shared astrai::mma_sync template (common/mma.cuh).
// read two adjacent bf16 from smem as one packed .b32 (elem0 low, elem1 high)
__device__ __forceinline__ unsigned ld2(const bf16* p) {
@@ -73,62 +70,18 @@ __device__ __forceinline__ unsigned pkb(bf16 a, bf16 b) {
return *reinterpret_cast<unsigned*>(&v);
}
// ldmatrix: cooperatively load mma fragments from smem (one instruction per
// 16x16 / 16x8 tile) with the exact register layout mma expects.
__device__ __forceinline__ void ldmatrix_x4(unsigned* r, const bf16* p) {
unsigned a = __cvta_generic_to_shared(p);
asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];"
: "=r"(r[0]), "=r"(r[1]), "=r"(r[2]), "=r"(r[3])
: "r"(a));
}
__device__ __forceinline__ void ldmatrix_x2(unsigned* r, const bf16* p) {
unsigned a = __cvta_generic_to_shared(p);
asm volatile("ldmatrix.sync.aligned.m8n8.x2.shared.b16 {%0,%1}, [%2];"
: "=r"(r[0]), "=r"(r[1])
: "r"(a));
}
__device__ __forceinline__ void ldmatrix_x2_trans(unsigned* r, const bf16* p) {
unsigned a = __cvta_generic_to_shared(p);
asm volatile("ldmatrix.sync.aligned.m8n8.x2.trans.shared.b16 {%0,%1}, [%2];"
: "=r"(r[0]), "=r"(r[1])
: "r"(a));
}
// ldmatrix lives in the shared template (common/mma.cuh):
// `astrai::ldmatrix_x2<bf16>` / `<bf16, /*Trans=*/true>` load the K/V
// fragments with the exact register layout mma expects.
// XOR swizzle for shared-memory column at 8-bf16 chunk granularity.
__device__ __forceinline__ int swiz_col(int d, int r, int mask = 7) {
return ((d >> 3) ^ (r & mask)) << 3 | (d & 7);
}
// Predicated cp.async: copy 16 bytes when `pred`, otherwise zero-fill.
// BypassL1 defaults to .cg (L2 only); false selects .ca (L1 + L2).
// src_size=0 means no bytes are read, so an out-of-bounds address is safe.
template <bool BypassL1 = true>
__device__ __forceinline__ void cp_async_16_pred(bf16* smem_ptr,
const void* gmem_ptr,
bool pred) {
unsigned smem_addr = __cvta_generic_to_shared(smem_ptr);
int src_size = pred ? 16 : 0;
if constexpr (BypassL1) {
asm volatile("cp.async.cg.shared.global [%0], [%1], 16, %2;"
:: "r"(smem_addr), "l"(gmem_ptr), "r"(src_size));
} else {
asm volatile("cp.async.ca.shared.global [%0], [%1], 16, %2;"
:: "r"(smem_addr), "l"(gmem_ptr), "r"(src_size));
}
}
__device__ __forceinline__ void cp_async_commit() {
asm volatile("cp.async.commit_group;");
}
__device__ __forceinline__ void cp_async_wait_all() {
asm volatile("cp.async.wait_all;");
}
template <int N>
__device__ __forceinline__ void cp_async_wait_group() {
asm volatile("cp.async.wait_group %0;" :: "n"(N));
}
// cp.async primitives live in the shared template (common/cp_async.cuh):
// `astrai::cp_async_16` (predicated), `astrai::cp_async_commit_group`,
// `astrai::cp_async_wait_group<N>` / `_wait_all` stage the K/V tiles.
// ---------------------------------------------------------------------------
// Q-load: load query rows directly from global memory into mma A-operand
@@ -180,9 +133,9 @@ __device__ inline void mma_compute_scores(
#pragma unroll
for (int kt = 0; kt < Traits::KD; kt++) {
unsigned b[2];
ldmatrix_x2(b, &sK[krow_l * Traits::LD
astrai::ldmatrix_x2<bf16>(b, &sK[krow_l * Traits::LD
+ swiz_col(kt * 16 + kcol_h, krow_l, Traits::SWIZ_MASK)]);
mma16816(Sacc[n8], Qa[kt], b, Sacc[n8]);
astrai::mma_sync<bf16>(Sacc[n8], Qa[kt], b, Sacc[n8]);
}
}
}
@@ -290,9 +243,12 @@ __device__ inline void mma_pv_accumulate(
#pragma unroll
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
unsigned b[2];
ldmatrix_x2_trans(b, &sV[vrow_l * Traits::LD
astrai::ldmatrix_x2<bf16, true>(b, &sV[vrow_l * Traits::LD
+ swiz_col(dn8 * 8, vrow_l, Traits::SWIZ_MASK)]);
mma16816(Oacc[dn8], Pa, b, Oacc[dn8]);
astrai::mma_sync<bf16>(Oacc[dn8], Pa, b, Oacc[dn8]);
}
}
}
} // namespace attention
} // namespace astrai
@@ -1,5 +1,7 @@
#include "attn_dispatchers.cuh"
#include "attn_entry_utils.cuh"
#include "dispatchers.cuh"
#include "entry_utils.cuh"
using namespace astrai::attention;
torch::Tensor attn_paged_decode(
torch::Tensor q,
@@ -1,5 +1,7 @@
#include "attn_dispatchers.cuh"
#include "attn_entry_utils.cuh"
#include "dispatchers.cuh"
#include "entry_utils.cuh"
using namespace astrai::attention;
torch::Tensor attn_paged_prefill(
torch::Tensor q,
@@ -1,5 +1,7 @@
#include "attn_dispatchers.cuh"
#include "attn_entry_utils.cuh"
#include "dispatchers.cuh"
#include "entry_utils.cuh"
using namespace astrai::attention;
torch::Tensor attn_prefill(
torch::Tensor q,
@@ -1,8 +1,12 @@
#pragma once
#include <cfloat>
#include <cuda_bf16.h>
#include "attn_common.h"
#include "attn_layout_policies.cuh"
#include "common.h"
#include "layout_policies.cuh"
#include "../common/reduce.cuh"
namespace astrai {
namespace attention {
using bf16 = __nv_bfloat16;
@@ -11,14 +15,7 @@ using bf16 = __nv_bfloat16;
// compile-time bools — the compiler eliminates dead branches.
// Unified across contiguous and paged (SGLang flat-pool) K/V via KV.
// Templated on <HEAD_DIM, KV, G, ROWS, P_BC, IsCausal, HasMask>.
template <int G>
__device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
#pragma unroll
for (int o = G / 2; o > 0; o >>= 1)
v += __shfl_xor_sync(mask, v, o);
return v;
}
// group_reduce_sum<G> lives in common/reduce.cuh (astrai::).
// load 8 contiguous bf16 from (16-byte aligned) smem as one float4
__device__ __forceinline__ void ld8(const bf16* p, float* o) {
@@ -155,3 +152,6 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
p.o_ptr[o_off + i * p.q_d_stride] = __float2bfloat16(acc[i] * rl);
}
}
} // namespace attention
} // namespace astrai
@@ -1,9 +1,12 @@
#pragma once
#include <cfloat>
#include <cuda_bf16.h>
#include "attn_common.h"
#include "attn_layout_policies.cuh"
#include "attn_mma_utils.cuh"
#include "common.h"
#include "layout_policies.cuh"
#include "mma_utils.cuh"
namespace astrai {
namespace attention {
// Tensor-core prefill flash attention (raw mma.sync PTX), unified across
// contiguous and paged (SGLang flat-pool) K/V via the KV template parameter.
@@ -85,10 +88,10 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
int token = KV::resolve_token(p, kctx, kc, valid);
KVAddr a = KV::kv_addr_from_token(p, kctx, token, d);
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
cp_async_16_pred(&dK[off], a.k, a.valid);
cp_async_16_pred(&dV[off], a.v, a.valid);
astrai::cp_async_16(&dK[off], a.k, a.valid);
astrai::cp_async_16(&dV[off], a.v, a.valid);
}
cp_async_commit();
astrai::cp_async_commit_group();
};
// ---- Prologue: issue first tile load ----
@@ -98,7 +101,7 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
int buf = ti & 1;
// Wait for current tile, then publish cross-warp + guard buffer reuse.
cp_async_wait_group<0>();
astrai::cp_async_wait_group<0>();
__syncthreads();
if (ti < t_end) load_tile(ti + 1, (ti + 1) & 1);
@@ -149,9 +152,12 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
}
if (qr1 < q_len) {
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][2] * rl1,
Oacc[dn8][3] * rl1);
Oacc[dn8][3] * rl1);
*reinterpret_cast<__nv_bfloat162*>(
&p.o_ptr[o_base + qr1 * p.q_l_stride + d * p.q_d_stride]) = v;
}
}
}
} // namespace attention
} // namespace astrai
-77
View File
@@ -1,77 +0,0 @@
#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]
};
// Unified attention params covering BOTH addressing modes:
// - Contiguous K/V: dense [batch, kv_head, kv_len, head_dim] tensors (k/v).
// - Paged (SGLang-style): flat pool [size, kv_head, head_dim] + req_to_token.
// Each kernel selects the addressing via a KVSource policy (see
// attn_layout_policies.cuh); a given call only touches the fields of one mode, so
// this is a POD shared by both paths rather than two parallel structs that
// drift out of sync.
template<typename T, typename AT = float>
struct AttentionParams {
// Shape
int batch;
int q_head;
int kv_head;
int head_dim;
int q_len; // Per-request in contiguous mode; total_q in paged mode.
int kv_len; // Contiguous mode; paged mode uses kv_indptr.
// Attention behavior
float scale;
// -1 = non-causal; >=0 = absolute position of first Q token
int causal_offset;
int use_mask;
// pointers
const T* __restrict__ q_ptr;
const T* __restrict__ k_ptr;
const T* __restrict__ v_ptr;
const T* __restrict__ new_k_ptr;
const T* __restrict__ new_v_ptr;
T* __restrict__ o_ptr;
const bool* __restrict__ mask;
// strides
int q_b_stride;
int q_h_stride;
int q_l_stride;
int q_d_stride;
int kv_b_stride;
int kv_h_stride;
int kv_l_stride;
int kv_d_stride;
int new_kv_b_stride;
int new_kv_h_stride;
int mask_b_stride;
int mask_h_stride;
int mask_l_stride;
// Paged K/V addressing
const int* __restrict__ req_to_token; // [num_reqs, max_context_len]
const int* __restrict__ req_pool_indices; // [batch]
const int* __restrict__ kv_indptr; // [batch + 1]
const int* __restrict__ qo_indptr; // [batch + 1] or nullptr for decode
const int* __restrict__ q_tile_to_batch; // [num_q_tiles], prefill only
const int* __restrict__ q_tile_to_index; // [num_q_tiles], prefill only
int num_q_tiles;
int max_context_len; // req_to_token stride (dim 1)
// Decode split-KV workspace
int num_splits;
AT* __restrict__ o_part;
AT* __restrict__ ml_part;
};
-13
View File
@@ -1,13 +0,0 @@
#pragma once
#include <cuda_bf16.h>
using bf16 = __nv_bfloat16;
static constexpr int MAX_SPLITS = 32;
__device__ inline float warp_reduce_sum(float val) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1)
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
return val;
}
+69
View File
@@ -0,0 +1,69 @@
// Shared cp.async primitives — pure CUDA, no torch.
//
// One header for the async-copy pipeline used by both the attention kernels
// (predicated 16-byte K/V tile staging) and the fp8 GEMM (predicated operand
// staging + wait_group dispatch). PTX requires wait_group's operand to be an
// immediate, hence the template forms.
#pragma once
#include <cuda_runtime.h>
namespace astrai {
// Predicated cp.async: copy 16 bytes when `pred`, otherwise zero-fill.
// src_size=0 means no bytes are read, so an out-of-bounds address is safe.
// BypassL1 defaults to .cg (L2 only); false selects .ca (L1 + L2).
// `T` is the smem element type; only the destination pointer's type matters.
template <typename T, bool BypassL1 = true>
__device__ __forceinline__ void cp_async_16(T* smem_ptr, const void* gmem_ptr,
bool pred) {
const unsigned smem_addr = __cvta_generic_to_shared(smem_ptr);
const int src_size = pred ? 16 : 0;
if constexpr (BypassL1) {
asm volatile("cp.async.cg.shared.global [%0], [%1], 16, %2;"
:: "r"(smem_addr), "l"(gmem_ptr), "r"(src_size));
} else {
asm volatile("cp.async.ca.shared.global [%0], [%1], 16, %2;"
:: "r"(smem_addr), "l"(gmem_ptr), "r"(src_size));
}
}
// Commit all outstanding cp.async ops of this thread as one group.
__device__ __forceinline__ void cp_async_commit_group() {
asm volatile("cp.async.commit_group;");
}
// Wait for every committed group (pipeline drain).
__device__ __forceinline__ void cp_async_wait_all() {
asm volatile("cp.async.wait_all;");
}
// Wait until at most KeepGroups committed groups are still in flight.
// PTX requires an immediate operand; keep it as a template argument so the
// stage policy stays compile-time configurable.
template <int KeepGroups>
__device__ __forceinline__ void cp_async_wait_group() {
static_assert(KeepGroups >= 0 && KeepGroups <= 7,
"cp.async.wait_group supports immediates in [0, 7]");
asm volatile("cp.async.wait_group %0;" :: "n"(KeepGroups));
}
// Runtime dispatch over cp_async_wait_group<N>: unrolls into a compare
// ladder over [0, MaxKeepGroups] so the immediate-only PTX constraint is
// hidden behind a runtime `keep_groups` (used by the fp8 GEMM pipeline,
// whose remaining-tile count is dynamic).
template <int MaxKeepGroups>
__device__ __forceinline__ void cp_async_wait_group_dispatch(int keep_groups) {
static_assert(MaxKeepGroups >= 0 && MaxKeepGroups <= 7,
"cp.async.wait_group supports immediates in [0, 7]");
if (keep_groups == MaxKeepGroups) {
cp_async_wait_group<MaxKeepGroups>();
} else if constexpr (MaxKeepGroups > 0) {
cp_async_wait_group_dispatch<MaxKeepGroups - 1>(keep_groups);
} else {
cp_async_wait_group<0>();
}
}
} // namespace astrai
+23
View File
@@ -0,0 +1,23 @@
// Pure-CUDA device helpers shared across kernel families (no torch).
//
// Family-local headers under kernels/<family>/ own their POD params and
// strategy traits; anything cross-cutting (compute-capability checks, device
// constants) lives here.
#pragma once
namespace astrai {
// Compute-capability comparison: is the device at least (major, minor)?
inline bool sm_at_least(int device_major, int device_minor, int major,
int minor) {
return device_major > major ||
(device_major == major && device_minor >= minor);
}
// FP8 tensor-core MMA (`mma.sync.aligned.m16n8k32` with fp8 inputs) exists on
// Ada (sm_89) and Hopper (sm_90+); sm_80 has no fp8 instructions.
inline constexpr int kMinSmForFp8Major = 8;
inline constexpr int kMinSmForFp8Minor = 9;
} // namespace astrai
+165
View File
@@ -0,0 +1,165 @@
// Shared mma.sync wrappers — pure CUDA, no torch.
//
// One template for every tensor-core MMA used by the kernel families. The
// instruction shape follows from the input element type:
// __nv_bfloat16 -> mma.sync.aligned.m16n8k16 (sm_80+), A = 4x b32, B = 2x b32
// __nv_fp8_e4m3/e5m2 -> mma.sync.aligned.m16n8k32 (sm_89+), A = 4x b32, B = 2x b32
// All variants accumulate into fp32: d = a*b + c, with the PTX mnemonic and
// the K dimension differing per type. `d` may alias `c` (in-place accumulate,
// as the FP8 GEMM does).
#pragma once
#include <cuda_bf16.h>
#include <cuda_fp8.h>
#include <cuda_runtime.h>
#include <type_traits>
#define DEVICE_FORCEINLINE static __device__ __forceinline__
namespace astrai {
// Compute capability of the current compilation pass: 0 in the host pass,
// the numeric CC (e.g. 890) in device passes where __CUDA_ARCH__ is defined.
// Defined() cannot appear in expressions, so this macro lets mma_sync use
// the arch in a static_assert instead of per-branch #if guards.
#ifndef __CUDA_ARCH__
#define ASTRAI_DEVICE_ARCH 0
#else
#define ASTRAI_DEVICE_ARCH __CUDA_ARCH__
#endif
// Compile-time shape of the MMA instruction for an input element type.
// `min_arch` is the numeric compute capability the instruction requires —
// the single place that encodes the hardware floor for each type.
template <typename InT>
struct mma_shape {
static constexpr int k = 16; // m16n8k16
static constexpr int a_regs = 4; // A fragment: 4x b32
static constexpr int b_regs = 2; // B fragment: 2x b32
static constexpr int min_arch = 800; // bf16 mma.sync, sm_80+
};
template <>
struct mma_shape<__nv_fp8_e4m3> {
static constexpr int k = 32; // m16n8k32
static constexpr int a_regs = 4;
static constexpr int b_regs = 2;
static constexpr int min_arch = 890; // fp8 mma.sync, sm_89+ (Ada/Hopper)
};
template <>
struct mma_shape<__nv_fp8_e5m2> {
static constexpr int k = 32;
static constexpr int a_regs = 4;
static constexpr int b_regs = 2;
static constexpr int min_arch = 890;
};
// d[4] = a[4] x b[2] + c[4], row-major A, col-major B, fp32 accumulator.
// The PTX mnemonic is selected from InT. Building for a compute capability
// below `mma_shape<InT>::min_arch` is a **compile error** — the instruction
// does not exist there, and a silent no-op would produce wrong results.
template <typename InT>
DEVICE_FORCEINLINE void mma_sync(float d[4], const unsigned a[4],
const unsigned b[2],
const float c[4]) {
static_assert(ASTRAI_DEVICE_ARCH == 0 ||
ASTRAI_DEVICE_ARCH >= mma_shape<InT>::min_arch,
"mma_sync: this MMA shape requires a newer compute "
"capability than the build target");
if constexpr (std::is_same_v<InT, __nv_bfloat16>) {
asm volatile(
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%10,%11,%12,%13};"
: "=f"(d[0]), "=f"(d[1]), "=f"(d[2]), "=f"(d[3])
: "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]),
"f"(c[0]), "f"(c[1]), "f"(c[2]), "f"(c[3]));
} else if constexpr (std::is_same_v<InT, __nv_fp8_e5m2>) {
asm volatile(
"mma.sync.aligned.m16n8k32.row.col.f32.e5m2.e5m2.f32 "
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%10,%11,%12,%13};"
: "=f"(d[0]), "=f"(d[1]), "=f"(d[2]), "=f"(d[3])
: "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]),
"f"(c[0]), "f"(c[1]), "f"(c[2]), "f"(c[3]));
} else {
asm volatile(
"mma.sync.aligned.m16n8k32.row.col.f32.e4m3.e4m3.f32 "
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%10,%11,%12,%13};"
: "=f"(d[0]), "=f"(d[1]), "=f"(d[2]), "=f"(d[3])
: "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]),
"f"(c[0]), "f"(c[1]), "f"(c[2]), "f"(c[3]));
}
}
#undef ASTRAI_DEVICE_ARCH
// ---------------------------------------------------------------------------
// ldmatrix — cooperatively load 8x8 b16 matrices from smem into registers.
//
// The instruction is identical for every 16-bit-storage element type: bf16
// maps 1:1 onto b16 slots; fp8 is stored packed two-per-slot (see
// fp8/gemm.cuh), so one b16 slot holds two fp8 values. `T` is the element
// type and only serves as a semantic tag.
//
// x2 (single address): matrix0 = p (8 rows), matrix1 = p + 8*16 bytes
// x4: four matrices at p, +128, +256, +384 bytes
// Trans: transpose variant (V fragments of attention)
//
// ldmatrix takes a *single* smem address per thread, but the addresses of
// the 32 lanes are *not* all the same: lane i supplies the start address of
// matrix-row i (modulo 8) for matrix (i/8) — lanes 0-7 feed matrix 0's rows,
// lanes 8-15 matrix 1's rows (x2/x4), lanes 16-23 / 24-31 matrix 2 / 3's rows
// (x4 only; their addresses are ignored by x2). Each matrix is 8 rows x 16
// bytes, and consecutive matrices of one instruction are contiguous at
// 128-byte strides. fp8 fragment layouts in fp8/gemm.cuh are arranged around
// this constraint.
// ---------------------------------------------------------------------------
template <typename T, bool Trans = false>
DEVICE_FORCEINLINE void ldmatrix_x2(unsigned r[2], const T* p) {
const unsigned a = __cvta_generic_to_shared(p);
if constexpr (Trans) {
asm volatile(
"ldmatrix.sync.aligned.m8n8.x2.trans.shared.b16 {%0,%1}, [%2];"
: "=r"(r[0]), "=r"(r[1])
: "r"(a));
} else {
asm volatile(
"ldmatrix.sync.aligned.m8n8.x2.shared.b16 {%0,%1}, [%2];"
: "=r"(r[0]), "=r"(r[1])
: "r"(a));
}
}
// Four matrices at p, p+128, p+256, p+384 bytes (16-byte row stride).
template <typename T>
DEVICE_FORCEINLINE void ldmatrix_x4(unsigned r[4], const T* p) {
const unsigned a = __cvta_generic_to_shared(p);
asm volatile(
"ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];"
: "=r"(r[0]), "=r"(r[1]), "=r"(r[2]), "=r"(r[3])
: "r"(a));
}
// Per-lane-address variants: the caller supplies a raw shared-memory address
// per lane instead of one common pointer. Use when the fragment tiles are
// XOR-swizzled per 16B chunk so each lane must compute its own row and chunk
// address (see fp8/gemm.cuh's frag_addr + lane selectors for the m16n8k32
// operand layouts).
DEVICE_FORCEINLINE void ldmatrix_x2_lane(unsigned r[2],
unsigned addr) {
asm volatile("ldmatrix.sync.aligned.m8n8.x2.shared.b16 {%0,%1}, [%2];"
: "=r"(r[0]), "=r"(r[1])
: "r"(addr));
}
DEVICE_FORCEINLINE void ldmatrix_x4_lane(unsigned r[4],
unsigned addr) {
asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];"
: "=r"(r[0]), "=r"(r[1]), "=r"(r[2]), "=r"(r[3])
: "r"(addr));
}
} // namespace astrai
+48
View File
@@ -0,0 +1,48 @@
// Shared warp/block reduction + atomic helpers — pure CUDA, no torch.
//
// Extracted from the attention and fp8 families so both share one
// implementation: warp_reduce_sum (decode scalar kernel), warp_reduce_max +
// atomic_max_float (fp8 quantize amax), group_reduce_sum<G> (prefill scalar
// kernel).
#pragma once
namespace astrai {
// Full-warp butterfly sum reduction (32 lanes).
__device__ __forceinline__ float warp_reduce_sum(float val) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1)
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
return val;
}
// Full-warp butterfly max reduction (32 lanes).
__device__ __forceinline__ float warp_reduce_max(float value) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1)
value = fmaxf(value, __shfl_xor_sync(0xffffffffu, value, offset));
return value;
}
// Sub-warp group reduction over G consecutive lanes (G a power of two).
// `mask` is the full participating-lane mask of the group (see the
// prefill scalar kernel's gmask computation).
template <int G>
__device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
#pragma unroll
for (int o = G / 2; o > 0; o >>= 1)
v += __shfl_xor_sync(mask, v, o);
return v;
}
// Unsigned-bit-pattern atomicMax for non-negative floats; a null
// destination disables the update (kernels with optional amax slots).
__device__ __forceinline__ void atomic_max_float(float* destination,
float value) {
if (destination)
atomicMax(reinterpret_cast<unsigned*>(destination),
__float_as_uint(value));
}
} // namespace astrai
+107
View File
@@ -0,0 +1,107 @@
#pragma once
#include <cuda_bf16.h>
#include <cuda_fp8.h>
#include <cuda_runtime.h>
#include <cstdint>
// Pure POD/traits header — no .cuh/CUDA-kernel includes; raw __nv_* type
// spellings only.
namespace astrai {
namespace fp8 {
// Compile-time FP8 format: E4M3 (forward / high precision, max 448) or
// E5M2 (gradient / large dynamic range, max 57344).
enum class FP8Format : int {
E4M3 = 0,
E5M2 = 1,
};
// Operand memory layouts as types (CUTLASS-style tags). The tag names the
// storage order of the raw buffer relative to the operand's canonical GEMM
// matrix — A is [M][K], B is [K][N]:
// A RowMajor = [M][K] storage (K-contiguous rows; the default)
// A ColMajor = [K][M] storage (M-contiguous; A^T)
// B RowMajor = [K][N] storage (N-contiguous; the plain a @ b operand)
// B ColMajor = [N][K] storage (K-contiguous; the nn.Linear weight layout)
// Empty tags: selection happens by type at compile time (see load_operand_tile).
struct RowMajor {};
struct ColMajor {};
// Transpose of a layout tag: the same buffer with the rows and contract dims
// swapped. B's tag is relative to the canonical [K][N] GEMM matrix, so the
// stage-load (which views any operand as [rows][contract]) sees the transposed
// tag — this trait makes that inversion explicit.
template <typename Layout>
struct transpose_layout;
template <>
struct transpose_layout<RowMajor> {
using type = ColMajor;
};
template <>
struct transpose_layout<ColMajor> {
using type = RowMajor;
};
template <typename Layout>
using transpose_layout_t = typename transpose_layout<Layout>::type;
// Compile-time tile configuration, mirroring KernelTraits<HEAD_DIM, BC,
// WARPS, STAGES> in the attention kernels. `Fmt` selects the FP8 conversion
// and the MMA PTX mnemonic; the remaining parameters shape the CTA tile and
// the cp.async pipeline depth.
template <FP8Format Fmt, int BlockM, int BlockN, int K, int Stages>
struct Fp8GemmTraits {
static constexpr FP8Format kFormat = Fmt;
static constexpr int kBlockM = BlockM;
static constexpr int kBlockN = BlockN;
static constexpr int kK = K;
static constexpr int kStages = Stages;
static constexpr bool kIsE5M2 = (Fmt == FP8Format::E5M2);
static constexpr __nv_fp8_interpretation_t kNvFormat =
kIsE5M2 ? __NV_E5M2 : __NV_E4M3;
static constexpr float kFp8Max = kIsE5M2 ? 57344.0f : 448.0f;
};
// Unified GEMM parameter POD, mirroring AttentionParams: one struct flows
// through quantize / fused / pre-quantized kernels. Each kernel touches only
// the fields it needs; buffers are raw pointers packed by the torch binding.
// Pointer members default to null (same NSDMI rationale as AttentionParams:
// bias / amax / out_scale gate optional paths via null checks, so a partially
// packed struct must never hold garbage non-null pointers). Still an
// aggregate, still trivially copyable.
struct FP8Params {
// Inputs: a/b are BF16 for the fused (quantize-in-GEMM) path, FP8 for
// the pre-quantized path. Scales are quantization steps (device scalars).
const void* __restrict__ a_ptr = nullptr;
const void* __restrict__ b_ptr = nullptr;
const void* __restrict__ bias = nullptr;
const float* __restrict__ scale_a = nullptr;
const float* __restrict__ scale_b = nullptr;
const float* __restrict__ bias_scale = nullptr;
// Output: BF16 or FP8 (E4M3). out_scale is the output quantization step
// (FP8 output only).
void* __restrict__ out_ptr = nullptr;
const float* __restrict__ out_scale = nullptr;
// Fused forward extras: bias (may be null) and amax slots (may be null).
float* __restrict__ amax_a = nullptr;
float* __restrict__ amax_b = nullptr;
// Shapes. total is only used by the elementwise quantize kernel. `int`
// covers every realistic LLM shape; the kernels promote to int64 for all
// pointer arithmetic.
int m, n, k;
// Physical leading dimensions (column count, i.e. row stride) of A and B.
// For a non-transposed operand the stride equals the contract dim; for a
// transposed operand it is the operand's own column count. The binding
// packs these so the kernel reads both buffers either naturally or
// transposed depending on the LayoutA/LayoutB tags (see gemm.cuh).
int a_ld, b_ld;
int total;
};
} // namespace fp8
} // namespace astrai
+575
View File
@@ -0,0 +1,575 @@
#pragma once
// FP8 GEMM device code — pure CUDA, no torch. Mirrors the attention kernel
// layout (attn_*_mma.cuh): kernels take the FP8Params POD, tile shape and
// FP8 format ride on compile-time template parameters, and launchers are
// plain functions usable from both the torch binding and pure C tests.
#include <cuda_bf16.h>
#include <cuda_fp8.h>
#include <cuda_runtime.h>
#include <type_traits>
#include "common.h"
#include "../common/cp_async.cuh"
#include "../common/mma.cuh"
#include "../common/reduce.cuh"
namespace astrai {
namespace fp8 {
// m16n8k32 (see astrai::mma_shape<fp8 type>::k in common/mma.cuh)
constexpr int kMmaK = 32;
constexpr int kWarps = 8; // 128x128 CTA = 8 warps
// log2 of a compile-time power of two (for tile_at's swizzle shift).
template <int N, int Acc = 0>
struct log2_const : log2_const<(N >> 1), Acc + 1> {};
template <int Acc>
struct log2_const<1, Acc> {
static constexpr int value = Acc;
};
// Map the FP8Format enum to the CUDA fp8 element type consumed by mma_sync.
template <FP8Format Fmt>
struct fp8_input {
using type = __nv_fp8_e4m3;
};
template <>
struct fp8_input<FP8Format::E5M2> {
using type = __nv_fp8_e5m2;
};
// ---------------------------------------------------------------------------
// Shared device helpers
// ---------------------------------------------------------------------------
// FP8 MMA lives in the shared astrai::mma_sync template (common/mma.cuh);
// instantiate it with fp8_input<Fmt>::type. Accumulates in-place: callers
// pass the same accumulator array as both `d` and `c`.
// warp_reduce_max / atomic_max_float (quantize amax) live in
// common/reduce.cuh; the cp.async pipeline primitives (predicated 16-byte
// copy, commit_group, wait_group + runtime dispatch) in common/cp_async.cuh.
// ---------------------------------------------------------------------------
// Quantize kernel: BF16 -> FP8 (E4M3 or E5M2), fused amax over raw values.
// ---------------------------------------------------------------------------
// Convert one packed bf16 pair to one packed fp8 pair. amax sees the *raw*
// (unscaled) values; the stored bytes see value * inv. Bit-identical to the
// scalar __nv_fp8_*(q) constructor path (round-nearest-even + satfinite).
template <FP8Format Fmt>
__device__ __forceinline__ unsigned quantize2(unsigned pair, float inv,
float& amax) {
const float lo = __bfloat162float(__ushort_as_bfloat16(pair & 0xffffu));
const float hi = __bfloat162float(__ushort_as_bfloat16(pair >> 16));
amax = fmaxf(amax, fmaxf(fabsf(lo), fabsf(hi)));
constexpr __nv_fp8_interpretation_t kFmt =
Fmt == FP8Format::E5M2 ? __NV_E5M2 : __NV_E4M3;
return static_cast<unsigned>(__nv_cvt_float2_to_fp8x2(
make_float2(lo * inv, hi * inv), __NV_SATFINITE, kFmt));
}
template <FP8Format Fmt>
__global__ void fp8_quantize_kernel(FP8Params p) {
const float inv = 1.0f / *p.scale_a;
const auto* x = reinterpret_cast<const __nv_bfloat16*>(p.a_ptr);
void* x8 = p.out_ptr;
float* amax = p.amax_a;
float local_amax = 0.0f;
const int64_t stride = (int64_t)blockDim.x * gridDim.x;
// Vectorized body: 8 bf16 (16B load) -> 8 fp8 (8B store) per step. Torch
// allocations are >=16B aligned and the binding passes freshly allocated
// contiguous buffers, so element 0 keeps the uint4/uint2 accesses
// natural; a misaligned base (contiguous view with an odd storage
// offset) falls back to the scalar loop below via total_vec = 0.
const bool aligned =
((reinterpret_cast<uintptr_t>(x) | reinterpret_cast<uintptr_t>(x8)) & 15) ==
0;
const int64_t total_vec = aligned ? p.total / 8 : 0;
const uint4* xv = reinterpret_cast<const uint4*>(x);
uint2* o8 = reinterpret_cast<uint2*>(x8);
for (int64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < total_vec;
i += stride) {
const uint4 v = xv[i];
const unsigned pair[4] = {v.x, v.y, v.z, v.w};
unsigned packed[2] = {0u, 0u};
#pragma unroll
for (int j = 0; j < 4; ++j)
packed[j >> 1] |= quantize2<Fmt>(pair[j], inv, local_amax)
<< (16 * (j & 1));
o8[i] = make_uint2(packed[0], packed[1]);
}
// Scalar tail (and full fallback for misaligned bases).
for (int64_t i = total_vec * 8 + blockIdx.x * blockDim.x + threadIdx.x;
i < p.total; i += stride) {
const float f = __bfloat162float(x[i]);
local_amax = fmaxf(local_amax, fabsf(f));
if constexpr (Fmt == FP8Format::E5M2) {
reinterpret_cast<__nv_fp8_e5m2*>(x8)[i] = __nv_fp8_e5m2(f * inv);
} else {
reinterpret_cast<__nv_fp8_e4m3*>(x8)[i] = __nv_fp8_e4m3(f * inv);
}
}
if (amax) {
local_amax = warp_reduce_max(local_amax);
__shared__ float slots[32];
if ((threadIdx.x & 31) == 0) slots[threadIdx.x >> 5] = local_amax;
__syncthreads();
if (threadIdx.x == 0) {
float v = 0.0f;
for (int w = 0; w < (blockDim.x >> 5); ++w) v = fmaxf(v, slots[w]);
atomic_max_float(amax, v);
}
}
}
// Swizzled address inside a flat [rows * K] staging tile: the 16-byte chunk
// index is XORed with a row-dependent slice so a warp's fragment load (8
// consecutive rows x 16B) hits all 32 banks exactly once. With kChunks
// power-of-two chunks per row, the XOR source is the top log2(kChunks) bits
// of the row index within each group of 8:
// kChunks=2 -> row bits [3] (K=32: rows r and r+4 diverge)
// kChunks=4 -> row bits [2:1] (K=64: rows diverge every 2)
// kChunks=8 -> row bits [2:0] (K=128: every row)
// (row word-stride is K/4 words = 4*kChunks, so unswizzled rows r and
// r + 8/kChunks collide mod 32 banks; the XOR spreads the 8 rows of one
// ldmatrix matrix across the 8 distinct 4-bank groups.) Chunks stay
// contiguous, so the cp.async 16B staging path is unaffected.
template <int K, typename T8>
__device__ __forceinline__ T8* tile_at(T8* tile, int row, int col) {
constexpr int kChunks = K / 16; // 16B chunks per row
static_assert(kChunks >= 1 && (kChunks & (kChunks - 1)) == 0,
"swizzle needs a power-of-two 16B-chunk count");
constexpr int kShift = 3 - log2_const<kChunks>::value;
return tile + row * K +
((((col >> 4) ^ ((row >> kShift) & (kChunks - 1))) << 4) + (col & 15));
}
// Stage-load one GEMM operand into the canonical flat [rows * K] shared tile
// (addressing via tile_at, so stores land in the swizzled layout). The
// transpose is folded into the staging step via a CUTLASS-style crosswise
// layout: RowMajor (stored [rows][contract]) copies 16-byte K-contiguous runs
// with cp.async, while ColMajor (stored [contract][rows]) reads 16-byte runs
// along the operand's contiguous non-contract dim and scatters them across
// the tile's rows. Crosswise runs cannot use cp.async (the 16 destination
// bytes land on 16 different rows), so their global loads are plain LDGs —
// issued as one batch per row group before the first scatter so their
// latencies overlap instead of serializing behind the shared stores.
// RowsTile is the tile's row capacity (kBlockM / kBlockN) and kThreads the
// CTA size; the runtime `rows` bound may be smaller (tail predication).
// `block_row` is this block's origin in the operand's row dim.
template <typename T8, int K, typename Layout, int RowsTile, int kThreads>
__device__ __forceinline__ void
load_operand_tile(T8* tile, const T8* __restrict__ operand, int64_t rows,
int64_t contract, int64_t ld, int tid, int64_t k_base,
int64_t block_row) {
constexpr int kChunks = K / 16;
static_assert(RowsTile * kChunks % kThreads == 0,
"tile chunks must divide evenly across threads");
constexpr int kCpt = RowsTile * kChunks / kThreads; // chunks per thread
if constexpr (std::is_same_v<Layout, ColMajor>) {
// Operand stored [contract][rows]: contiguous along the non-contract
// dim. Each thread scatters one 16-byte run per K/32 pass; when the
// tile has more 16-row groups than warps (RowsTile > kThreads/2),
// each thread covers several groups.
constexpr int kWarpsTile = kThreads / 32;
constexpr int kGroups = RowsTile / 16;
constexpr int kPasses = K / 32;
static_assert(kGroups % kWarpsTile == 0,
"row groups must divide evenly across warps");
const int kl = tid & 31; // byte column within a 32B pass
// r0 is always a multiple of 16 (block_row is a multiple of RowsTile
// and each group covers 16 rows), so every run shares the base+ld
// alignment: one uniform check instead of one per pass.
const bool run_aligned =
((reinterpret_cast<uintptr_t>(operand) | ld) & 15) == 0;
#pragma unroll
for (int g = 0; g < kGroups / kWarpsTile; ++g) {
const int rg = (tid >> 5) + g * kWarpsTile;
const int64_t r0 = block_row + rg * 16;
const bool rows_full = r0 + 15 < rows; // pass-invariant
// Batch every 16B run load of this row group before the first
// scatter: the LDGs are independent, and the byte-granular
// shared stores would otherwise serialize behind each one.
uint4 v[kPasses];
bool fast[kPasses];
#pragma unroll
for (int pass = 0; pass < kPasses; ++pass) {
const int64_t k_idx = k_base + kl + pass * 32;
fast[pass] = rows_full && run_aligned && k_idx < contract;
if (fast[pass])
v[pass] =
*reinterpret_cast<const uint4*>(operand + k_idx * ld + r0);
}
#pragma unroll
for (int pass = 0; pass < kPasses; ++pass) {
const int col = kl + pass * 32;
if (fast[pass]) {
const auto* bytes = reinterpret_cast<const T8*>(&v[pass]);
// Scatter 16 bytes along the tile rows through tile_at's
// swizzle. Rows sharing a physical chunk form groups of
// (8 / kChunks) consecutive rows (see tile_at), so each
// group is one tile_at address plus a K-byte row stride.
constexpr int kGrp = 8 / kChunks;
#pragma unroll
for (int j = 0; j < 16 / kGrp; ++j) {
T8* p = tile_at<K>(tile, rg * 16 + j * kGrp, col);
#pragma unroll
for (int i = 0; i < kGrp; ++i)
p[i * K] = bytes[j * kGrp + i];
}
} else if (k_base + col < contract) {
// Row-tail or misaligned run: byte-granular gather with
// per-row predication (the k column itself is in range).
#pragma unroll
for (int i = 0; i < 16; ++i) {
const int64_t r_idx = r0 + i;
*tile_at<K>(tile, rg * 16 + i, col) =
r_idx < rows ? operand[(k_base + col) * ld + r_idx]
: T8(0.0f);
}
} else {
// Contract tail: straight zero-fill, no global traffic.
#pragma unroll
for (int i = 0; i < 16; ++i)
*tile_at<K>(tile, rg * 16 + i, col) = T8(0.0f);
}
}
}
} else {
// Operand stored [rows][contract]: contiguous along the contract dim.
// Linear chunk mapping: thread covers kCpt consecutive 16B chunks of
// one row (K=64: a contiguous 32B pair; K=32: a single chunk).
constexpr int kCpr = kChunks / kCpt; // chunks per row slice
const int r = tid / kCpr;
const int c0 = (tid % kCpr) * kCpt * 16;
const int64_t row = block_row + r;
const bool row_ok = row < rows;
// k_base and every c are multiples of 16, so the per-chunk sources
// share the row base's alignment.
const auto* src = operand + row * ld + k_base;
const bool chunk_aligned = (reinterpret_cast<uintptr_t>(src) & 15) == 0;
#pragma unroll
for (int j = 0; j < kCpt; ++j) {
const int c = c0 + j * 16;
T8* dst = tile_at<K>(tile, r, c);
if (row_ok && chunk_aligned && k_base + c + 15 < contract) {
astrai::cp_async_16(dst, src + c, true);
} else {
// Tail chunk (or misaligned base): predicated scalar fill.
#pragma unroll
for (int i = 0; i < 16; ++i)
dst[i] =
row_ok && k_base + c + i < contract ? src[c + i] : T8(0.0f);
}
}
}
}
// ---------------------------------------------------------------------------
// Pre-quantized GEMM kernel: FP8 A/B read straight into shared memory, FP32
// accumulation, BF16 or FP8 output. The input format follows Traits; the
// tile is compact (row = kK bytes) so MMA fragments read directly — no
// in-kernel transpose of the operands (the binding handles transposes).
// ---------------------------------------------------------------------------
// Swizzled 16B-chunk address (tile_at's layout) as a raw shared-memory
// pointer for ldmatrix. Valid for kK in {32, 64} (the swizzle itself lives
// only in tile_at; this wrapper just converts the element address).
template <typename T8, int kK>
__device__ __forceinline__ unsigned frag_addr(const T8* tile, int row, int chunk) {
static_assert(kK == 32 || kK == 64,
"fragment swizzle offsets assume kK in {32, 64}");
return __cvta_generic_to_shared(tile_at<kK>(tile, row, chunk << 4));
}
// LayoutA / LayoutB tag the operands' storage (CUTLASS-style, see common.h):
// A RowMajor = [M][K] / ColMajor = [K][M]; B RowMajor = [K][N] /
// ColMajor = [N][K]. The kernel always computes
// out[m][n] = sum_p tileA[m][p] * tileB[n][p]
// with the tiles materialized in the canonical [M][kK] / [N][kK] layout, so the
// MMA fragments are read identically regardless of layout. The tags only
// change how the stage-load gathers the operand from global memory:
// A ColMajor: tileA[m][p] = a[p*a_ld + m]; A RowMajor: a[m*a_ld + p]
// B RowMajor: tileB[n][p] = b[p*b_ld + n]; B ColMajor: b[n*b_ld + p]
// BlockM x BlockN CTA as (BlockM/64) x (BlockN/32) warps of 64x32 warp tiles
// (mt x nt = 4x4 MMA each). The 64x128 variant runs 4 warps / 128 threads and
// exists for small-M calls: m <= 64 wastes half of every 128-row CTA, so the
// launcher dispatches to it there (see launch_fp8_gemm).
template <typename Traits, bool OutFp8 = false, typename LayoutA = RowMajor,
typename LayoutB = RowMajor>
__global__ void
__launch_bounds__((Traits::kBlockM / 64) * (Traits::kBlockN / 32) * 32, 2)
fp8_gemm_kernel(FP8Params p) {
using T8 = std::conditional_t<Traits::kIsE5M2, __nv_fp8_e5m2, __nv_fp8_e4m3>;
constexpr int kBlockM = Traits::kBlockM;
constexpr int kBlockN = Traits::kBlockN;
constexpr int kK = Traits::kK;
constexpr int kStages = Traits::kStages;
constexpr int kCtaThreads = (kBlockM / 64) * (kBlockN / 32) * 32;
static_assert(kStages >= 1 && kStages <= 8,
"FP8 GEMM stages must be in the range [1, 8]");
// Tiles are flat [rows * kK] with a 16B-chunk XOR swizzle (tile_at):
// ldmatrix reads whole 16B chunks through the same mapping the staging
// writes, and the swizzle removes the bank conflict the unswizzled
// 8-word row stride caused (see tile_at).
__shared__ __align__(16) T8 a_smem[kStages][kBlockM * kK];
__shared__ __align__(16) T8 b_smem[kStages][kBlockN * kK];
const auto* a = reinterpret_cast<const T8*>(p.a_ptr);
const auto* b = reinterpret_cast<const T8*>(p.b_ptr);
auto* out_bf16 = reinterpret_cast<__nv_bfloat16*>(p.out_ptr);
auto* out_fp8 = reinterpret_cast<__nv_fp8_e4m3*>(p.out_ptr);
const int64_t m = p.m, n = p.n, k = p.k;
const int64_t a_ld = p.a_ld, b_ld = p.b_ld;
const int tid = threadIdx.x;
const int warp = tid >> 5;
const int lane = tid & 31;
const int group = lane >> 2;
const int thread_in_group = lane & 3;
// L2-friendly rasterization (CUTLASS-style grouped launch order): remap
// the linear block id so consecutive CTAs cover a group of kGroupM M-tiles
// before advancing along N. All CTAs of one group share the same B column
// stripe, so B tiles stay hot in L2 across the wave (the default
// N-fastest order makes each wave touch every B tile instead).
// Measured win for the A-crosswise layouts (10-21% at K>=2048) and loss
// for A-congruous (-17..20%, A's cp.async stream prefers the N-fastest
// order) — so the branch follows LayoutA.
constexpr int kGroupM = 8;
int block_m, block_n;
if constexpr (std::is_same_v<LayoutA, ColMajor>) {
const int blocks_m = gridDim.y;
const int bid = blockIdx.y * gridDim.x + blockIdx.x;
const int group_first_m = (bid / (kGroupM * gridDim.x)) * kGroupM;
const int group_rows =
min(blocks_m - group_first_m, kGroupM); // M-tail group is short
block_m = group_first_m + bid % group_rows;
block_n = (bid % (kGroupM * gridDim.x)) / group_rows;
} else {
block_m = blockIdx.y;
block_n = blockIdx.x;
}
// 128x128 CTA = 8 warps as 2x4 warp tiles of 64x32 (mt x nt = 4x4 MMA).
constexpr int warps_n = kBlockN / 32;
const int warp_m = warp / warps_n;
const int warp_n = warp % warps_n;
const int64_t row_base = (int64_t)block_m * kBlockM + warp_m * 64 + group;
const int64_t output_col =
(int64_t)block_n * kBlockN + warp_n * 32 + thread_in_group * 2;
const int a_row0 = warp_m * 64; // + mt * 16 in the loop
const int b_row0 = warp_n * 32; // + nt * 8
const float sa = *p.scale_a;
const float sb = *p.scale_b;
float acc[4][4][4] = {}; // [nt][mt][acc]
// Both operands are staged into the canonical [M][kK] / [N][kK] shared
// tiles regardless of their global layout (see load_operand_tile), so the
// MMA fragment reads below stay unchanged across the four layout
// combinations. A's tag already names the operand view ([M][K] =
// [rows][contract]); B's tag is relative to the canonical [K][N], so the
// stage-load sees its transpose (transpose_layout_t, see common.h).
auto load_tile = [&](int stage, int64_t k_base) {
load_operand_tile<T8, kK, LayoutA, kBlockM, kCtaThreads>(
a_smem[stage], a, m, k, a_ld, tid, k_base, (int64_t)block_m * kBlockM);
load_operand_tile<T8, kK, transpose_layout_t<LayoutB>, kBlockN,
kCtaThreads>(b_smem[stage], b, n, k, b_ld, tid, k_base,
(int64_t)block_n * kBlockN);
};
const int64_t tile_count = (k + kK - 1) / kK;
// Per-lane ldmatrix row/chunk selectors for common/mma.cuh's
// ldmatrix_*_lane (the fragment tiles are XOR-swizzled per 16B chunk, so
// each lane computes its own row/chunk address). Layout contract for fp8
// m16n8k32 (values packed two-per-b16 slot, K-contiguous rows):
// x4 (A fragment): lane i points at tile row (i>>3 & 1)*8 + (i&7) of
// chunk (k_seg*2 + (i>>4)); reg j = matrix j = [row g][tig*4..+3] in
// the order (rows 0-7 c, rows 8-15 c, rows 0-7 c+1, rows 8-15 c+1) —
// exactly the mma.sync A operand layout.
// x2 (B fragment): lane i points at tile row (i&7) of chunk
// (k_seg*2 + ((i>>3) & 1)); reg j = [row(n) g][tig*4..+3] chunk c/c+1
// — exactly the mma.sync B operand layout (col operand, K-contiguous).
const int r7 = lane & 7; // row within the 8-row matrix
const int rh8 = (lane >> 3) & 1; // +8 rows (A: lanes 8-15, 24-31)
const int rh16 = lane >> 4; // +1 chunk (A: lanes 16-31; B uses rh8)
// Prime the pipeline. Each committed group occupies one circular shared
// memory stage; the loop also handles K dimensions smaller than kStages.
#pragma unroll
for (int stage = 0; stage < kStages; ++stage) {
if (stage < tile_count) {
load_tile(stage, static_cast<int64_t>(stage) * kK);
astrai::cp_async_commit_group();
}
}
for (int64_t tile_index = 0; tile_index < tile_count; ++tile_index) {
const int stage = static_cast<int>(tile_index % kStages);
const int64_t remaining = tile_count - tile_index - 1;
// Keep up to kStages - 1 younger groups in flight while making the
// oldest group (the current stage) ready for consumption.
const int keep_groups =
remaining < kStages - 1 ? static_cast<int>(remaining) : kStages - 1;
astrai::cp_async_wait_group_dispatch<kStages - 1>(keep_groups);
// Barrier 1: every thread's cp.async for this stage is complete
// before any thread reads tiles written by other threads.
__syncthreads();
// 4 ldmatrix.x2 (B) + 4 ldmatrix.x4 (A) feed 16 mma.sync per k_seg —
// 0.5 load instructions per MMA, versus 4.5 scalar LDS per MMA in
// the 128x64-tile version (the kernel was LSU-issue-bound there).
constexpr int kSegs = kK / kMmaK;
// B fragments double-buffered across k_segs: the next k_seg's B load
// is issued before the current k_seg's MMA sequence, so its LDS
// latency hides behind the A pipeline + tensor-pipe work (same trick
// as the A mt+1 prefetch below; costs kSegs x 8 registers).
unsigned b_frag[2][4][2];
#pragma unroll
for (int nt = 0; nt < 4; ++nt) {
const int row = b_row0 + nt * 8 + r7;
astrai::ldmatrix_x2_lane(b_frag[0][nt],
frag_addr<T8, kK>(b_smem[stage], row, rh8));
}
#pragma unroll
for (int k_seg = 0; k_seg < kSegs; ++k_seg) {
const int bcur = k_seg & 1, bnext = bcur ^ 1;
if (k_seg + 1 < kSegs) {
#pragma unroll
for (int nt = 0; nt < 4; ++nt) {
const int row = b_row0 + nt * 8 + r7;
astrai::ldmatrix_x2_lane(
b_frag[bnext][nt],
frag_addr<T8, kK>(b_smem[stage], row,
(k_seg + 1) * 2 + rh8));
}
}
// Software-pipelined A fragments: the ldmatrix.x4 for row mt+1
// is issued before the MMAs consuming row mt, so the LDS fixed
// latency hides behind tensor-pipe work (cuts the `wait` stall,
// ~2.3 cycles/issue before this). Costs 4 extra registers.
unsigned a_frag[5][4];
astrai::ldmatrix_x4_lane(
a_frag[0], frag_addr<T8, kK>(a_smem[stage], a_row0 + rh8 * 8 + r7,
k_seg * 2 + rh16));
#pragma unroll
for (int mt = 0; mt < 4; ++mt) {
if (mt < 3)
astrai::ldmatrix_x4_lane(
a_frag[mt + 1],
frag_addr<T8, kK>(a_smem[stage],
a_row0 + (mt + 1) * 16 + rh8 * 8 + r7,
k_seg * 2 + rh16));
#pragma unroll
for (int nt = 0; nt < 4; ++nt)
astrai::mma_sync<T8>(acc[nt][mt], a_frag[mt], b_frag[bcur][nt],
acc[nt][mt]);
}
}
// Barrier 2: every thread finished reading this stage's tiles before
// the prefetch for the (i+kStages)-th tile overwrites them.
__syncthreads();
if (tile_index + kStages < tile_count) {
load_tile(stage, (tile_index + kStages) * kK);
astrai::cp_async_commit_group();
}
}
const float output_scale = sa * sb;
// Fused bias: BF16 raw values, or FP8 storage dequantized by its own
// scale (bias_scale != null selects the FP8 path; the format follows the
// kernel's Traits). Added in real units after the operand dequantization
// and before any output quantization.
const auto* bias16 = static_cast<const __nv_bfloat16*>(p.bias);
const auto* bias8 = static_cast<const T8*>(p.bias);
auto bias_val = [&](int64_t col) -> float {
if (p.bias == nullptr || col >= n) return 0.0f;
if (p.bias_scale == nullptr) return __bfloat162float(bias16[col]);
return __half2float(__half(bias8[col])) * *p.bias_scale;
};
#pragma unroll
for (int nt = 0; nt < 4; ++nt) {
const int64_t col = output_col + nt * 8;
const float b0 = bias_val(col);
const float b1 = bias_val(col + 1);
// Per-row store: FP8 packs two adjacent columns into one 16-bit
// write, BF16 into one 32-bit __nv_bfloat162 (single cvt+pack
// instruction); boundary or unaligned columns fall back to scalar
// converts so a pack never crosses the row edge or misaligns.
auto store_out = [&](int64_t row, float v0, float v1) {
if (row >= m) return;
const float r0 = v0 * output_scale + b0;
const float r1 = v1 * output_scale + b1;
if constexpr (OutFp8) {
if (col + 1 < n) {
*reinterpret_cast<unsigned short*>(out_fp8 + row * n + col) =
static_cast<unsigned short>(__nv_cvt_float2_to_fp8x2(
make_float2(r0 * *p.out_scale, r1 * *p.out_scale),
__NV_SATFINITE, __NV_E4M3));
} else {
out_fp8[row * n + col] = __nv_fp8_e4m3(r0 * *p.out_scale);
}
} else {
auto* dst = out_bf16 + row * n + col;
if (col + 1 < n && (reinterpret_cast<uintptr_t>(dst) & 3) == 0) {
*reinterpret_cast<__nv_bfloat162*>(dst) =
__floats2bfloat162_rn(r0, r1);
} else {
dst[0] = __float2bfloat16(r0);
if (col + 1 < n) dst[1] = __float2bfloat16(r1);
}
}
};
#pragma unroll
for (int mt = 0; mt < 4; ++mt) {
const int64_t row0 = row_base + mt * 16;
float* tile_acc = acc[nt][mt];
if (col < n) {
store_out(row0, tile_acc[0], tile_acc[1]);
store_out(row0 + 8, tile_acc[2], tile_acc[3]);
}
}
}
}
// ---------------------------------------------------------------------------
// Launchers — pure CUDA (no torch), usable from the binding and pure C tests.
// ---------------------------------------------------------------------------
template <FP8Format Fmt>
void launch_fp8_quantize(const FP8Params& p, cudaStream_t stream) {
constexpr int kThreads = 256;
// One block per 256 vectors (8 elements each); at least one block so the
// scalar tail of a tiny / misaligned tensor is still covered.
int64_t blocks = (p.total / 8 + kThreads - 1) / kThreads;
if (blocks < 1) blocks = 1;
fp8_quantize_kernel<Fmt><<<blocks, kThreads, 0, stream>>>(p);
}
// Pre-quantized GEMM tile config: 128x128 CTA (8 warps x 64x32 warp tiles).
// kK selects the K tile (32 or 64; 64 halves the __syncthreads count per K
// and doubles the MMA work per stage, at 2x the smem per stage — measured
// 10-35% across shapes, so 64 is the default). Stages=2 with kK=64 keeps the
// pipeline at 32KB smem; deeper pipelines only win on K >= 4096 squares and
// lose elsewhere. LayoutA/LayoutB mirror the kernel template (defaults keep
// the NN layout: out = a @ b). m <= 64 dispatches to the 64x128 CTA — a
// 128-row CTA would waste half its MMA work on predicated-off rows.
template <FP8Format Fmt, bool OutFp8 = false, typename LayoutA = RowMajor,
typename LayoutB = RowMajor, int kK = 64, int Stages = 2>
void launch_fp8_gemm(const FP8Params& p, cudaStream_t stream) {
dim3 grid((p.n + 127) / 128, (p.m + 127) / 128);
if (p.m <= 64) {
using Traits = Fp8GemmTraits<Fmt, 64, 128, kK, Stages>;
fp8_gemm_kernel<Traits, OutFp8, LayoutA, LayoutB>
<<<grid, (64 / 64) * (128 / 32) * 32, 0, stream>>>(p);
} else {
using Traits = Fp8GemmTraits<Fmt, 128, 128, kK, Stages>;
fp8_gemm_kernel<Traits, OutFp8, LayoutA, LayoutB>
<<<grid, (128 / 64) * (128 / 32) * 32, 0, stream>>>(p);
}
}
} // namespace fp8
} // namespace astrai
+434
View File
@@ -0,0 +1,434 @@
// FP8 GEMM torch binding: tensor validation, FP8Params packing, template
// dispatch and pybind. Device code lives in gemm.cuh (pure CUDA) —
// mirroring the attn_*.cu / attn_*_mma.cuh split of the attention kernels.
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <cuda_bf16.h>
#include <cstdint>
#include <mutex>
#include <tuple>
#include <unordered_map>
#include "gemm.cuh"
#include "../common/device.cuh"
using namespace astrai::fp8;
namespace {
// FP8Format / FP8Params and the launchers live in astrai::fp8 (common.h /
// gemm.cuh); this TU opens the using-directive above so the binding reads
// them unqualified.
void check_fp8_device(const torch::Tensor& tensor) {
static std::mutex mutex;
static std::unordered_map<int, bool> supported;
const int device = tensor.device().index();
{
std::lock_guard<std::mutex> lock(mutex);
auto cached = supported.find(device);
if (cached != supported.end()) {
TORCH_CHECK(cached->second,
"fused FP8 MMA requires compute capability 8.9 or newer");
return;
}
}
const auto* properties = at::cuda::getDeviceProperties(device);
const bool is_supported =
astrai::sm_at_least(properties->major, properties->minor,
astrai::kMinSmForFp8Major,
astrai::kMinSmForFp8Minor);
{
std::lock_guard<std::mutex> lock(mutex);
supported.emplace(device, is_supported);
}
TORCH_CHECK(is_supported,
"fused FP8 MMA requires compute capability 8.9 or newer");
}
void check_scale(const torch::Tensor& scale, const torch::Tensor& input,
const char* name) {
TORCH_CHECK(scale.is_cuda() && scale.device() == input.device() &&
scale.scalar_type() == torch::kFloat32 && scale.numel() == 1,
name, " must be a CUDA float32 scalar on the input device");
}
// ---- FP8Params packing (mirrors attention/entry_utils.cuh pack_* helpers) ----
void pack_gemm_params(FP8Params& p, const void* a, const void* b, void* out,
const torch::Tensor& sa, const torch::Tensor& sb,
const torch::Tensor* out_scale, const void* bias,
const torch::Tensor* bias_scale, int64_t m, int64_t n,
int64_t k, int64_t a_ld, int64_t b_ld) {
p.a_ptr = a;
p.b_ptr = b;
p.out_ptr = out;
p.scale_a = sa.data_ptr<float>();
p.scale_b = sb.data_ptr<float>();
p.out_scale = out_scale ? out_scale->data_ptr<float>() : nullptr;
p.bias = bias;
p.bias_scale = bias_scale ? bias_scale->data_ptr<float>() : nullptr;
p.amax_a = nullptr;
p.amax_b = nullptr;
p.m = static_cast<int>(m);
p.n = static_cast<int>(n);
p.k = static_cast<int>(k);
p.a_ld = static_cast<int>(a_ld);
p.b_ld = static_cast<int>(b_ld);
p.total = 0;
}
void pack_quantize_params(FP8Params& p, const void* x, void* x8,
const torch::Tensor& scale, torch::Tensor* amax,
int64_t total) {
p.a_ptr = x;
p.b_ptr = nullptr;
p.out_ptr = x8;
p.scale_a = scale.data_ptr<float>();
p.scale_b = nullptr;
p.out_scale = nullptr;
p.bias = nullptr;
p.amax_a = amax ? amax->data_ptr<float>() : nullptr;
p.amax_b = nullptr;
p.m = p.n = p.k = 0;
p.a_ld = p.b_ld = 0;
p.total = static_cast<int>(total);
}
// ---- GEMM launch dispatch (runtime flags -> compile-time kernel variants) ----
template <FP8Format Fmt, int Variant>
void launch_gemm_variant(const FP8Params& p, cudaStream_t stream) {
static_assert(Variant >= 0 && Variant < 8,
"invalid FP8 GEMM dispatch variant");
constexpr bool out_fp8 = (Variant & 4) != 0;
// Variant bits 1/0 = trans_a/trans_b -> CUTLASS-style layout tags
// (trans_a ? A ColMajor : RowMajor, same for B; see common.h).
using LayoutA = std::conditional_t<(Variant & 2) != 0, ColMajor, RowMajor>;
using LayoutB = std::conditional_t<(Variant & 1) != 0, ColMajor, RowMajor>;
launch_fp8_gemm<Fmt, out_fp8, LayoutA, LayoutB>(p, stream);
}
template <FP8Format Fmt>
void dispatch_gemm(const FP8Params& p, cudaStream_t stream, bool out_fp8,
bool trans_a, bool trans_b) {
// Encode the runtime flags as [output FP8, transpose A, transpose B].
const int variant = (static_cast<int>(out_fp8) << 2) |
(static_cast<int>(trans_a) << 1) |
static_cast<int>(trans_b);
switch (variant) {
case 0: launch_gemm_variant<Fmt, 0>(p, stream); break;
case 1: launch_gemm_variant<Fmt, 1>(p, stream); break;
case 2: launch_gemm_variant<Fmt, 2>(p, stream); break;
case 3: launch_gemm_variant<Fmt, 3>(p, stream); break;
case 4: launch_gemm_variant<Fmt, 4>(p, stream); break;
case 5: launch_gemm_variant<Fmt, 5>(p, stream); break;
case 6: launch_gemm_variant<Fmt, 6>(p, stream); break;
case 7: launch_gemm_variant<Fmt, 7>(p, stream); break;
}
}
} // namespace
// ---------------------------------------------------------------------------
// Entry points
// ---------------------------------------------------------------------------
std::tuple<torch::Tensor, torch::Tensor> quantize_bf16(torch::Tensor x,
torch::Tensor scale,
int64_t fmt) {
// BF16 -> FP8 quantize with fused amax. fmt: 0 = E4M3, 1 = E5M2.
// Returns (x8, amax); the caller never clears amax (zero-initialized here).
TORCH_CHECK(x.is_cuda() && scale.is_cuda(), "CUDA tensors required");
TORCH_CHECK(x.scalar_type() == torch::kBFloat16, "x must be bf16");
check_scale(scale, x, "scale");
check_fp8_device(x);
const at::cuda::OptionalCUDAGuard guard(x.device());
auto stream = at::cuda::getCurrentCUDAStream();
auto x_c = x.contiguous();
auto x8 = torch::empty_like(
x_c, x_c.options().dtype(fmt ? torch::kFloat8_e5m2
: torch::kFloat8_e4m3fn));
auto amax = torch::zeros({1}, x_c.options().dtype(torch::kFloat32));
FP8Params p;
pack_quantize_params(p, x_c.data_ptr(), x8.data_ptr(), scale, &amax,
x_c.numel());
if (fmt) {
launch_fp8_quantize<FP8Format::E5M2>(p, stream.stream());
} else {
launch_fp8_quantize<FP8Format::E4M3>(p, stream.stream());
}
C10_CUDA_CHECK(cudaGetLastError());
return {x8, amax};
}
torch::Tensor mm_fp8(torch::Tensor a, torch::Tensor b, torch::Tensor sa,
torch::Tensor sb, int64_t out_dtype,
c10::optional<torch::Tensor> out_scale, int64_t trans_a,
int64_t trans_b) {
// Pre-quantized FP8 GEMM: out = op(a) @ op(b)^T * (sa * sb), FP32 accum.
// trans_a / trans_b select the operand layout (0 = stored [M,K]/[K,N],
// 1 = transposed [K,M]/[N,K]); the default (0/0) is the plain a @ b.
// out_dtype: 0 = BF16 (default), 1 = FP8 E4M3 (requires out_scale, the
// output quantization step). Both operands share one format.
TORCH_CHECK(a.is_cuda() && b.is_cuda(), "CUDA tensors required");
TORCH_CHECK(a.scalar_type() == torch::kFloat8_e4m3fn ||
a.scalar_type() == torch::kFloat8_e5m2,
"a and b must be fp8 (e4m3fn or e5m2)");
TORCH_CHECK(a.scalar_type() == b.scalar_type(),
"a and b must share the same fp8 format");
TORCH_CHECK(a.dim() == 2 && b.dim() == 2, "a and b must be 2D");
TORCH_CHECK(a.device() == b.device(), "a and b must be on the same device");
check_scale(sa, a, "sa");
check_scale(sb, a, "sb");
check_fp8_device(a);
const at::cuda::OptionalCUDAGuard guard(a.device());
auto stream = at::cuda::getCurrentCUDAStream();
auto a_c = a.contiguous();
auto b_c = b.contiguous();
const bool ta = (trans_a == 1), tb = (trans_b == 1);
// Physical leading dimension = column count of each contiguous buffer.
const int64_t a_ld = a_c.size(1);
const int64_t b_ld = b_c.size(1);
// Logical GEMM shape derived from the layout flags.
const int64_t m = ta ? a_c.size(1) : a_c.size(0);
const int64_t k = ta ? a_c.size(0) : a_c.size(1);
const int64_t n = tb ? b_c.size(0) : b_c.size(1);
const int64_t k2 = tb ? b_c.size(1) : b_c.size(0);
TORCH_CHECK(k == k2, "inner dim mismatch");
const bool out_fp8 = (out_dtype == 1);
TORCH_CHECK(out_dtype == 0 || out_fp8,
"out_dtype must be 0 (bf16) or 1 (fp8 e4m3)");
torch::Tensor os;
if (out_fp8) {
TORCH_CHECK(out_scale.has_value(), "fp8 output requires out_scale");
os = out_scale.value();
check_scale(os, a, "out_scale");
}
auto out = torch::empty(
{m, n},
out_fp8 ? a_c.options().dtype(torch::kFloat8_e4m3fn)
: a_c.options().dtype(torch::kBFloat16));
FP8Params p;
pack_gemm_params(p, a_c.data_ptr(), b_c.data_ptr(), out.data_ptr(), sa, sb,
out_fp8 ? &os : nullptr, nullptr, nullptr, m, n, k, a_ld,
b_ld);
if (a.scalar_type() == torch::kFloat8_e4m3fn)
dispatch_gemm<FP8Format::E4M3>(p, stream.stream(), out_fp8, ta, tb);
else
dispatch_gemm<FP8Format::E5M2>(p, stream.stream(), out_fp8, ta, tb);
C10_CUDA_CHECK(cudaGetLastError());
return out;
}
std::tuple<torch::Tensor, torch::Tensor, torch::Tensor> linear_forward_fp8(
torch::Tensor x, torch::Tensor w, torch::Tensor bias, torch::Tensor sx,
torch::Tensor sw, int64_t fmt,
c10::optional<torch::Tensor> bias_scale) {
// Pure FP8 forward: quantize x/w (fmt: 0 = E4M3, 1 = E5M2), then the
// pre-quantized GEMM; the dequantized BF16 output gets the bias added.
// amax_x / amax_w come from the quantize kernels (zero-initialized here;
// a pre-quantized w reports amax_w = 0 — nothing to feed a delayed ring).
// w may itself be pre-quantized fp8 storage matching fmt (static
// inference weights): the weight quantize is skipped, amax_w stays 0.
TORCH_CHECK(x.is_cuda() && w.is_cuda(), "CUDA tensors required");
const auto f8opt = fmt ? torch::kFloat8_e5m2 : torch::kFloat8_e4m3fn;
const bool w_prequant = w.scalar_type() == f8opt;
TORCH_CHECK(
x.scalar_type() == torch::kBFloat16 &&
(w.scalar_type() == torch::kBFloat16 || w_prequant),
"x must be bf16; w must be bf16 or pre-quantized fp8 matching fmt");
TORCH_CHECK(x.device() == w.device(), "x and w must be on the same device");
check_scale(sx, x, "sx");
check_scale(sw, x, "sw");
check_fp8_device(x);
const at::cuda::OptionalCUDAGuard guard(x.device());
auto stream = at::cuda::getCurrentCUDAStream();
auto x_c = x.reshape({-1, w.size(1)}).contiguous(); // [M, K]
auto w_c = w.contiguous(); // [N, K]
int64_t m = x_c.size(0), k = x_c.size(1), n = w_c.size(0);
TORCH_CHECK(w_c.dim() == 2 && w_c.size(1) == k, "inner dim mismatch");
const bool has_bias = bias.defined() && bias.numel() > 0;
const bool b_prequant = has_bias && bias.scalar_type() == f8opt;
if (has_bias) {
TORCH_CHECK(bias.is_cuda() && bias.device() == x.device() &&
bias.numel() == n &&
(bias.scalar_type() == torch::kBFloat16 || b_prequant),
"bias must be CUDA bf16 or pre-quantized fp8 matching fmt, "
"with shape [N]");
TORCH_CHECK(b_prequant == bias_scale.has_value(),
"fp8 bias requires bias_scale (and bf16 bias takes none)");
if (b_prequant) check_scale(*bias_scale, x, "bias_scale");
}
auto x8 = torch::empty({m, k}, x_c.options().dtype(f8opt));
auto amax_x = torch::zeros({1}, x.options().dtype(torch::kFloat32));
auto amax_w = torch::zeros({1}, x.options().dtype(torch::kFloat32));
auto out = torch::empty({m, n}, x_c.options());
auto quantize = [&](const torch::Tensor& src, torch::Tensor& dst,
const torch::Tensor& scale, torch::Tensor* amax) {
FP8Params qp;
pack_quantize_params(qp, src.data_ptr(), dst.data_ptr(), scale, amax,
src.numel());
if (fmt) {
launch_fp8_quantize<FP8Format::E5M2>(qp, stream.stream());
} else {
launch_fp8_quantize<FP8Format::E4M3>(qp, stream.stream());
}
};
quantize(x_c, x8, sx, &amax_x);
// Static inference weights arrive pre-quantized (w8 storage + its scale);
// only freshly-loaded bf16 weights quantize here.
torch::Tensor w8 = w_prequant
? w_c
: torch::empty({n, k}, x_c.options().dtype(f8opt));
if (!w_prequant) quantize(w_c, w8, sw, &amax_w);
FP8Params p;
// Forward is the NT layout: A = x8 [M,K] (a_ld = k), B = w8 [N,K]
// (b_ld = k), out = x @ w^T. The bias is fused into the epilogue (bf16
// raw, or fp8 + bias_scale on the static path).
auto bias_c = has_bias ? bias.contiguous() : bias;
pack_gemm_params(p, x8.data_ptr(), w8.data_ptr(), out.data_ptr(), sx, sw,
nullptr, has_bias ? bias_c.data_ptr() : nullptr,
b_prequant ? &*bias_scale : nullptr, m, n, k, k, k);
if (fmt) {
launch_fp8_gemm<FP8Format::E5M2, false, RowMajor, ColMajor>(
p, stream.stream());
} else {
launch_fp8_gemm<FP8Format::E4M3, false, RowMajor, ColMajor>(
p, stream.stream());
}
C10_CUDA_CHECK(cudaGetLastError());
std::vector<int64_t> shape(x.sizes().begin(), x.sizes().end() - 1);
shape.push_back(n);
return {out.reshape(shape), amax_x, amax_w};
}
std::tuple<torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor>
linear_backward_fp8(torch::Tensor g, torch::Tensor x, torch::Tensor w,
std::vector<int64_t> masks, torch::Tensor sg,
torch::Tensor sw, torch::Tensor sx, int64_t fmt) {
// Pre-quantized FP8 backward: grad is quantized once (E4M3 or E5M2 per
// `fmt`), then dX / dW run as FP8 tensor-core GEMMs sharing g8.
// Returns (grad_input, grad_weight, grad_bias, amax_g).
TORCH_CHECK(g.is_cuda() && x.is_cuda() && w.is_cuda(), "CUDA tensors required");
TORCH_CHECK(g.scalar_type() == torch::kBFloat16 &&
x.scalar_type() == torch::kBFloat16 &&
w.scalar_type() == torch::kBFloat16,
"g, x, and w must be bf16");
TORCH_CHECK(g.device() == x.device() && g.device() == w.device(),
"g, x, and w must be on the same device");
TORCH_CHECK(masks.size() == 3, "masks must contain three values");
check_fp8_device(g);
const at::cuda::OptionalCUDAGuard guard(g.device());
auto stream = at::cuda::getCurrentCUDAStream();
auto g_c = g.reshape({-1, w.size(0)}).contiguous(); // [M, N]
auto x_c = x.reshape({-1, x.size(-1)}).contiguous(); // [M, K]
auto w_c = w.contiguous(); // [N, K]
int64_t m = g_c.size(0), n = w_c.size(0), k = w_c.size(1);
TORCH_CHECK(x_c.size(0) == m && x_c.size(1) == k && g_c.size(1) == n,
"backward shape mismatch");
auto grad_input = torch::empty_like(x);
auto grad_weight = torch::empty_like(w);
auto grad_bias = torch::empty({0}, g.options());
auto amax_g = torch::zeros({1}, g.options().dtype(torch::kFloat32));
auto f8opt = fmt ? g.options().dtype(torch::kFloat8_e5m2)
: g.options().dtype(torch::kFloat8_e4m3fn);
auto quantize = [&](const torch::Tensor& src, torch::Tensor& dst,
const torch::Tensor& scale, torch::Tensor* amax) {
FP8Params qp;
pack_quantize_params(qp, src.data_ptr(), dst.data_ptr(), scale, amax,
src.numel());
if (fmt) {
launch_fp8_quantize<FP8Format::E5M2>(qp, stream.stream());
} else {
launch_fp8_quantize<FP8Format::E4M3>(qp, stream.stream());
}
};
// Four-layout backward: the gradient and activation tensors keep their
// natural row-major layout, and the kernel reads them transposed where the
// GEMM needs it (the ColMajor layout tags pick the crosswise stage-load).
// No torch-level `.transpose().contiguous()`
// copies are required — dX uses g8 [M,N] as A with w8 [N,K] read transposed
// as B; dW uses g8 transposed as A with x8 transposed as B.
// g is quantized once (amax_g measured here); both GEMMs share g8.
auto run_bwd_gemm = [&](const FP8Params& gp, bool trans_a, bool trans_b) {
if (fmt)
dispatch_gemm<FP8Format::E5M2>(gp, stream.stream(), false, trans_a,
trans_b);
else
dispatch_gemm<FP8Format::E4M3>(gp, stream.stream(), false, trans_a,
trans_b);
};
torch::Tensor g8;
if (masks[0] || masks[1]) {
g8 = torch::empty({m, n}, f8opt);
quantize(g_c, g8, sg, &amax_g);
}
// dX = g @ w: A = g8 [M,N] (contract over N), B = w8 [N,K] read transposed
// (b[p*b_ld + n] = w[p,n]); out = [M,K], a_ld = N, b_ld = K, contract = N.
if (masks[0]) {
auto w8 = torch::empty({n, k}, f8opt);
quantize(w_c, w8, sw, nullptr);
auto grad_input_2d = grad_input.reshape({m, k});
FP8Params gp;
pack_gemm_params(gp, g8.data_ptr(), w8.data_ptr(),
grad_input_2d.data_ptr(), sg, sw, nullptr, nullptr,
nullptr, m, k, n, n, k);
run_bwd_gemm(gp, false, false);
}
// dW = g^T @ x: A = g8 [M,N] read transposed (a[p*a_ld + m] = g[p,m]), B =
// x8 [M,K] read transposed (b[p*b_ld + n] = x[p,n]); out = [N,K], a_ld = N,
// b_ld = K, contract = M.
if (masks[1]) {
auto x8 = torch::empty({m, k}, f8opt);
quantize(x_c, x8, sx, nullptr);
FP8Params gp;
pack_gemm_params(gp, g8.data_ptr(), x8.data_ptr(),
grad_weight.data_ptr(), sg, sx, nullptr, nullptr,
nullptr, n, k, m, n, k);
run_bwd_gemm(gp, true, false);
}
if (!masks[0] && !masks[1]) {
amax_g.copy_(g_c.abs().amax().to(torch::kFloat32));
}
C10_CUDA_CHECK(cudaGetLastError());
if (masks[2]) grad_bias = g_c.sum(0).to(g.scalar_type());
return {grad_input, grad_weight, grad_bias, amax_g};
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("quantize_bf16", &quantize_bf16, py::arg("x"), py::arg("scale"),
py::arg("fmt"),
"BF16 to FP8 (E4M3/E5M2) quantize with fused amax; returns (x8, amax)");
m.def("mm_fp8", &mm_fp8, py::arg("a"), py::arg("b"), py::arg("sa"),
py::arg("sb"), py::arg("out_dtype") = 0,
py::arg("out_scale") = py::none(), py::arg("trans_a") = 0,
py::arg("trans_b") = 0,
"Pre-quantized FP8 GEMM: op(a) @ op(b)^T * (sa * sb); out_dtype "
"0=bf16, 1=fp8 e4m3 (requires out_scale); trans_a/trans_b select "
"the operand layout (default 0/0 = a@b)");
m.def("linear_forward_fp8", &linear_forward_fp8, py::arg("x"),
py::arg("w"), py::arg("bias"), py::arg("sx"), py::arg("sw"),
py::arg("fmt") = 0, py::arg("bias_scale") = py::none(),
"Pure FP8 linear forward: quantize x/w, pre-quantized GEMM with the "
"bias fused into the epilogue; w and bias may be pre-quantized fp8 "
"matching fmt (static inference path; fp8 bias requires bias_scale);"
" returns (out, amax_x, amax_w)");
m.def("linear_backward_fp8", &linear_backward_fp8, py::arg("g"),
py::arg("x"), py::arg("w"), py::arg("masks"), py::arg("sg"),
py::arg("sw"), py::arg("sx"), py::arg("fmt"),
"FP8 linear backward; returns (grad_input, grad_weight, grad_bias, amax_g)");
}
-817
View File
@@ -1,817 +0,0 @@
// Fused BF16 -> E4M3 MMA -> BF16 matrix multiplication
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <cuda_fp8.h>
#include <cuda_runtime.h>
#include <cstdint>
#include <mutex>
#include <unordered_map>
namespace {
constexpr int kMmaK = 32;
constexpr int kWarps = 8;
// Fast forward path: 128x64 CTA, 64x16 warp tile, 2-stage pipeline, dynamic
// shared memory. Mirrors the CUTLASS 58_ada_fp8_gemm threadblock geometry
// while keeping the fused BF16->FP8 quantize path. The FP8 tile overwrites
// the BF16 staging area in place. L20 opts in to only 101376 B shared per
// block; K=32 keeps the footprint at 24576 B so four CTAs/SM stay resident.
constexpr int kFastBlockM = 128;
constexpr int kFastBlockN = 64;
constexpr int kFastK = 32;
constexpr int kFastStages = 2;
constexpr int kFastSmemBytes =
kFastStages * (kFastBlockM * kFastK * 2 + kFastBlockN * kFastK * 2);
__device__ __forceinline__ unsigned pack_fp8x4_vector(float x0, float x1,
float x2, float x3) {
const auto low = __nv_cvt_float2_to_fp8x2(
make_float2(x0, x1), __NV_SATFINITE, __NV_E4M3);
const auto high = __nv_cvt_float2_to_fp8x2(
make_float2(x2, x3), __NV_SATFINITE, __NV_E4M3);
return static_cast<unsigned>(low) | (static_cast<unsigned>(high) << 16);
}
__device__ __forceinline__ void mma_fp8_16832(float d[4],
const unsigned a[4],
const unsigned b[2]) {
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 890
asm volatile(
"mma.sync.aligned.m16n8k32.row.col.f32.e4m3.e4m3.f32 "
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%0,%1,%2,%3};"
: "+f"(d[0]), "+f"(d[1]), "+f"(d[2]), "+f"(d[3])
: "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]),
"r"(b[0]), "r"(b[1]));
#endif
}
__device__ __forceinline__ void atomic_max_float(float* destination,
float value) {
if (destination)
atomicMax(reinterpret_cast<unsigned*>(destination), __float_as_uint(value));
}
__device__ __forceinline__ float warp_reduce_max(float value) {
#pragma unroll
for (int offset = 16; offset; offset >>= 1) {
value = fmaxf(value, __shfl_xor_sync(0xffffffffu, value, offset));
}
return value;
}
// Block-wide max reduction of a per-warp tracked value, then an atomic
// update of the global amax slot when `track` is set.
template <int NWarps>
__device__ __forceinline__ void block_reduce_amax(float& local, float* slots,
int warp, int lane,
bool track, float* global) {
local = warp_reduce_max(local);
if (lane == 0) slots[warp] = local;
__syncthreads();
if (warp == 0) {
float value = lane < NWarps ? slots[lane] : 0.0f;
value = warp_reduce_max(value);
if (lane == 0 && track && global) atomic_max_float(global, value);
}
}
// One thread moves eight BF16 values (16 bytes). The async copy is issued
// through a uint4-shaped pointer so the source and destination are both
// naturally 128-bit aligned for contiguous forward GEMMs.
__device__ __forceinline__ void cp_async_bf16_8(
__nv_bfloat16* destination, const __nv_bfloat16* source, bool valid) {
const unsigned shared_address = __cvta_generic_to_shared(destination);
const uint4* source_vec = reinterpret_cast<const uint4*>(source);
asm volatile("cp.async.cg.shared.global [%0], [%1], 16, %2;"
:: "r"(shared_address), "l"(source_vec),
"r"(valid ? 16 : 0));
}
template <bool TrackAmax = true>
__device__ __forceinline__ unsigned load_fp8x4_from_bf16(
const __nv_bfloat16* source, float scale_inv, float& amax,
bool track_amax = true) {
float x0 = __bfloat162float(source[0]);
float x1 = __bfloat162float(source[1]);
float x2 = __bfloat162float(source[2]);
float x3 = __bfloat162float(source[3]);
if constexpr (TrackAmax) {
if (track_amax) {
amax = fmaxf(amax, fmaxf(fabsf(x0), fmaxf(fabsf(x1),
fmaxf(fabsf(x2), fabsf(x3)))));
}
}
return pack_fp8x4_vector(x0 * scale_inv, x1 * scale_inv,
x2 * scale_inv, x3 * scale_inv);
}
template <bool AddBias, bool TrackAmax>
__global__ void fused_fp8_gemm_fast_kernel(
const __nv_bfloat16* __restrict__ a,
const __nv_bfloat16* __restrict__ b,
__nv_bfloat16* __restrict__ out,
const __nv_bfloat16* __restrict__ bias,
const float* __restrict__ scale_a,
const float* __restrict__ scale_b,
float* __restrict__ amax_a,
float* __restrict__ amax_b,
int64_t m, int64_t n, int64_t k) {
extern __shared__ char smem[];
constexpr int a_stride = kFastBlockM * kFastK;
constexpr int b_stride = kFastBlockN * kFastK;
constexpr int b_bf16_offset = kFastStages * a_stride;
auto* a_bf16 = reinterpret_cast<__nv_bfloat16*>(smem);
auto* b_bf16 =
reinterpret_cast<__nv_bfloat16*>(smem + b_bf16_offset * 2);
__shared__ float warp_amax_a[kWarps];
__shared__ float warp_amax_b[kWarps];
const int tid = threadIdx.x;
const int warp = tid >> 5;
const int lane = tid & 31;
const int group = lane >> 2;
const int thread_in_group = lane & 3;
constexpr int warps_n = kFastBlockN / 16;
const int warp_m = warp / warps_n;
const int warp_n = warp % warps_n;
const int64_t row_base =
blockIdx.y * kFastBlockM + warp_m * 64 + group;
const int64_t output_col =
blockIdx.x * kFastBlockN + warp_n * 16 + thread_in_group * 2;
const float sa = *scale_a;
const float sb = *scale_b;
const float inv_a = 1.0f / sa;
const float inv_b = 1.0f / sb;
float local_amax_a = 0.0f;
float local_amax_b = 0.0f;
float acc[4 * 4 * 2] = {};
const bool track_amax_a = TrackAmax && blockIdx.x == 0;
const bool track_amax_b = TrackAmax && blockIdx.y == 0;
// Each thread issues 8 A chunks and 4 B chunks of 8 BF16 (16B) per stage.
auto load_tile = [&](int stage, int64_t k_base) {
const int r0 = tid >> 2;
const int c0 = (tid & 3) * 8;
#pragma unroll
for (int j = 0; j < kFastK / 32; ++j) {
const int col = c0 + 32 * j;
const bool full_chunk = k_base + col + 7 < k;
const int64_t a_row = blockIdx.y * kFastBlockM + r0;
const int64_t b_row = blockIdx.x * kFastBlockN + r0;
auto* a_dst = &a_bf16[stage * a_stride + r0 * kFastK + col];
auto* b_dst = &b_bf16[stage * b_stride + r0 * kFastK + col];
const auto* a_ptr = a + a_row * k + k_base + col;
const auto* b_ptr = b + b_row * k + k_base + col;
const bool full_a = a_row < m && full_chunk;
const bool full_b = b_row < n && full_chunk;
const bool aligned_a =
(reinterpret_cast<uintptr_t>(a_ptr) & 15) == 0;
const bool aligned_b =
(reinterpret_cast<uintptr_t>(b_ptr) & 15) == 0;
if (full_a && aligned_a) {
cp_async_bf16_8(a_dst, a_ptr, true);
} else {
#pragma unroll
for (int i = 0; i < 8; ++i) {
a_dst[i] = a_row < m && k_base + col + i < k
? a_ptr[i]
: __float2bfloat16(0.0f);
}
}
if (full_b && aligned_b) {
cp_async_bf16_8(b_dst, b_ptr, true);
} else {
#pragma unroll
for (int i = 0; i < 8; ++i) {
b_dst[i] = b_row < n && k_base + col + i < k
? b_ptr[i]
: __float2bfloat16(0.0f);
}
}
if (r0 + 64 < kFastBlockM) {
const int64_t a_row_hi = blockIdx.y * kFastBlockM + r0 + 64;
auto* a_dst_hi =
&a_bf16[stage * a_stride + (r0 + 64) * kFastK + col];
const auto* a_ptr_hi = a + a_row_hi * k + k_base + col;
const bool full_a_hi = a_row_hi < m && full_chunk;
const bool aligned_a_hi =
(reinterpret_cast<uintptr_t>(a_ptr_hi) & 15) == 0;
if (full_a_hi && aligned_a_hi) {
cp_async_bf16_8(a_dst_hi, a_ptr_hi, true);
} else {
#pragma unroll
for (int i = 0; i < 8; ++i) {
a_dst_hi[i] = a_row_hi < m && k_base + col + i < k
? a_ptr_hi[i]
: __float2bfloat16(0.0f);
}
}
}
}
};
// Quantize must place each 4-BP8 group at the byte offset the MMA
// fragment reads: 8*(lane&3) + 64*k_seg for a row. With in-place storage
// (fp8 element k lives at byte 2k), the BF16 column of a group is
// 4*(tid&7) + 32*j, so partition by 4-element groups instead of the
// 8-element cp.async chunks.
auto quantize_tile = [&](int stage) {
const int r0 = tid >> 3;
const int c0 = (tid & 7) * 4;
#pragma unroll
for (int s = 0; s < 4; ++s) {
const int row = r0 + 32 * s;
auto* a_src = &a_bf16[stage * a_stride + row * kFastK + c0];
auto* a_dst = reinterpret_cast<unsigned*>(a_src);
#pragma unroll
for (int j = 0; j < kFastK / 32; ++j) {
a_dst[16 * j] = load_fp8x4_from_bf16<TrackAmax>(
a_src + 32 * j, inv_a, local_amax_a, track_amax_a);
}
}
#pragma unroll
for (int s = 0; s < 2; ++s) {
const int row = r0 + 32 * s;
auto* b_src = &b_bf16[stage * b_stride + row * kFastK + c0];
auto* b_dst = reinterpret_cast<unsigned*>(b_src);
#pragma unroll
for (int j = 0; j < kFastK / 32; ++j) {
b_dst[16 * j] = load_fp8x4_from_bf16<TrackAmax>(
b_src + 32 * j, inv_b, local_amax_b, track_amax_b);
}
}
};
const int64_t tile_count = (k + kFastK - 1) / kFastK;
load_tile(0, 0);
asm volatile("cp.async.commit_group;");
if (tile_count > 1) {
load_tile(1, kFastK);
asm volatile("cp.async.commit_group;");
}
for (int64_t tile_index = 0; tile_index < tile_count; ++tile_index) {
const int stage = static_cast<int>(tile_index % kFastStages);
if (tile_index + 1 == tile_count) {
asm volatile("cp.async.wait_group 0;");
} else {
asm volatile("cp.async.wait_group 1;");
}
// wait_group only waits for this thread's async copies. All threads
// must finish loading before the tile is read by the CTA.
__syncthreads();
quantize_tile(stage);
__syncthreads();
// Four m16n8k32 MMA segments per 128-K stage.
#pragma unroll
for (int k_seg = 0; k_seg < kFastK / kMmaK; ++k_seg) {
const int frag_col = thread_in_group * 4 + k_seg * 32;
#pragma unroll
for (int nt = 0; nt < 2; ++nt) {
const int b_row = warp_n * 16 + nt * 8 + group;
unsigned b_frag[2];
b_frag[0] = *reinterpret_cast<const unsigned*>(
&b_bf16[stage * b_stride + b_row * kFastK + frag_col]);
b_frag[1] = *reinterpret_cast<const unsigned*>(
&b_bf16[stage * b_stride + b_row * kFastK + frag_col + 16]);
#pragma unroll
for (int mt = 0; mt < 4; ++mt) {
const int a_row0 = warp_m * 64 + mt * 16 + group;
unsigned a_frag[4];
a_frag[0] = *reinterpret_cast<const unsigned*>(
&a_bf16[stage * a_stride + a_row0 * kFastK + frag_col]);
a_frag[1] = *reinterpret_cast<const unsigned*>(
&a_bf16[stage * a_stride + (a_row0 + 8) * kFastK + frag_col]);
a_frag[2] = *reinterpret_cast<const unsigned*>(
&a_bf16[stage * a_stride + a_row0 * kFastK + frag_col + 16]);
a_frag[3] = *reinterpret_cast<const unsigned*>(
&a_bf16[stage * a_stride + (a_row0 + 8) * kFastK + frag_col + 16]);
mma_fp8_16832(acc + (nt * 4 + mt) * 4, a_frag, b_frag);
}
}
}
__syncthreads();
if (tile_index + 2 < tile_count) {
load_tile(stage, (tile_index + 2) * kFastK);
asm volatile("cp.async.commit_group;");
}
}
if constexpr (TrackAmax) {
block_reduce_amax<kWarps>(local_amax_a, warp_amax_a, warp, lane,
track_amax_a, amax_a);
block_reduce_amax<kWarps>(local_amax_b, warp_amax_b, warp, lane,
track_amax_b, amax_b);
}
const float output_scale = sa * sb;
#pragma unroll
for (int nt = 0; nt < 2; ++nt) {
const int64_t col = output_col + nt * 8;
#pragma unroll
for (int mt = 0; mt < 4; ++mt) {
const int64_t row0 = row_base + mt * 16;
const int64_t row1 = row0 + 8;
float* tile_acc = acc + (nt * 4 + mt) * 4;
if (col < n) {
float bias0 = 0.0f;
float bias1 = 0.0f;
if constexpr (AddBias) {
bias0 = __bfloat162float(bias[col]);
if (col + 1 < n)
bias1 = __bfloat162float(bias[col + 1]);
}
if (row0 < m) {
out[row0 * n + col] =
__float2bfloat16(tile_acc[0] * output_scale + bias0);
if (col + 1 < n)
out[row0 * n + col + 1] = __float2bfloat16(
tile_acc[1] * output_scale + bias1);
}
if (row1 < m) {
out[row1 * n + col] =
__float2bfloat16(tile_acc[2] * output_scale + bias0);
if (col + 1 < n)
out[row1 * n + col + 1] = __float2bfloat16(
tile_acc[3] * output_scale + bias1);
}
}
}
}
}
// Pre-quantized FP8-in path: FP8 A/B read straight into shared memory (no
// BF16 staging, no inline quantization), FP32 accumulation, BF16 output.
// Same 128x64 CTA / 64x16 warp tile geometry as the fused kernel; the fp8
// tile is compact (row = kFastK bytes) so MMA fragments read directly.
constexpr int kPqBlockM = 128;
constexpr int kPqBlockN = 64;
constexpr int kPqK = 32;
constexpr int kPqStages = 3;
template <typename T>
__device__ __forceinline__ void cp_async_16b(T* destination,
const T* source, bool valid) {
const unsigned shared_address = __cvta_generic_to_shared(destination);
const uint4* source_vec = reinterpret_cast<const uint4*>(source);
asm volatile("cp.async.cg.shared.global [%0], [%1], 16, %2;"
:: "r"(shared_address), "l"(source_vec),
"r"(valid ? 16 : 0));
}
template <bool OutFp8>
__global__ void fp8_mm_pq_kernel(
const __nv_fp8_e4m3* __restrict__ a,
const __nv_fp8_e4m3* __restrict__ b,
__nv_bfloat16* __restrict__ out_bf16,
__nv_fp8_e4m3* __restrict__ out_fp8,
const float scale, const float out_scale,
int64_t m, int64_t n, int64_t k) {
__shared__ __align__(16) __nv_fp8_e4m3 a_tile[kPqStages][kPqBlockM][kPqK];
__shared__ __align__(16) __nv_fp8_e4m3 b_tile[kPqStages][kPqBlockN][kPqK];
const int tid = threadIdx.x;
const int warp = tid >> 5;
const int lane = tid & 31;
const int group = lane >> 2;
const int thread_in_group = lane & 3;
constexpr int warps_n = kPqBlockN / 16;
const int warp_m = warp / warps_n;
const int warp_n = warp % warps_n;
const int64_t row_base = blockIdx.y * kPqBlockM + warp_m * 64 + group;
const int64_t output_col =
blockIdx.x * kPqBlockN + warp_n * 16 + thread_in_group * 2;
float acc[4 * 4 * 2] = {};
// One A chunk (16 FP8) per thread covers the 128x32 tile; the first 128
// threads issue the 64x32 B chunks.
auto load_tile = [&](int stage, int64_t k_base) {
const int r0 = tid >> 1;
const int c0 = (tid & 1) * 16;
const bool full_chunk = k_base + c0 + 15 < k;
const int64_t a_row = blockIdx.y * kPqBlockM + r0;
auto* a_dst = &a_tile[stage][r0][c0];
const auto* a_ptr = a + a_row * k + k_base + c0;
const bool full_a = a_row < m && full_chunk;
const bool aligned_a =
(reinterpret_cast<uintptr_t>(a_ptr) & 15) == 0;
if (full_a && aligned_a) {
cp_async_16b(a_dst, a_ptr, true);
} else {
#pragma unroll
for (int i = 0; i < 16; ++i) {
a_dst[i] = a_row < m && k_base + c0 + i < k
? a_ptr[i]
: __nv_fp8_e4m3(0.0f);
}
}
if (tid < 128) {
const int64_t b_row = blockIdx.x * kPqBlockN + r0;
auto* b_dst = &b_tile[stage][r0][c0];
const auto* b_ptr = b + b_row * k + k_base + c0;
const bool full_b = b_row < n && full_chunk;
const bool aligned_b =
(reinterpret_cast<uintptr_t>(b_ptr) & 15) == 0;
if (full_b && aligned_b) {
cp_async_16b(b_dst, b_ptr, true);
} else {
#pragma unroll
for (int i = 0; i < 16; ++i) {
b_dst[i] = b_row < n && k_base + c0 + i < k
? b_ptr[i]
: __nv_fp8_e4m3(0.0f);
}
}
}
};
const int64_t tile_count = (k + kPqK - 1) / kPqK;
load_tile(0, 0);
asm volatile("cp.async.commit_group;");
if (tile_count > 1) {
load_tile(1, kPqK);
asm volatile("cp.async.commit_group;");
}
if (tile_count > 2) {
load_tile(2, 2 * kPqK);
asm volatile("cp.async.commit_group;");
}
for (int64_t tile_index = 0; tile_index < tile_count; ++tile_index) {
const int stage = static_cast<int>(tile_index % kPqStages);
const int64_t remaining = tile_count - tile_index - 1;
if (remaining >= 2) {
asm volatile("cp.async.wait_group 2;");
} else if (remaining == 1) {
asm volatile("cp.async.wait_group 1;");
} else {
asm volatile("cp.async.wait_group 0;");
}
// Barrier 1: every thread's cp.async for this stage is complete
// before any thread reads tiles written by other threads.
__syncthreads();
#pragma unroll
for (int k_seg = 0; k_seg < kPqK / kMmaK; ++k_seg) {
const int frag_col = thread_in_group * 4 + k_seg * 32;
#pragma unroll
for (int nt = 0; nt < 2; ++nt) {
const int b_row = warp_n * 16 + nt * 8 + group;
unsigned b_frag[2];
b_frag[0] = *reinterpret_cast<const unsigned*>(
&b_tile[stage][b_row][frag_col]);
b_frag[1] = *reinterpret_cast<const unsigned*>(
&b_tile[stage][b_row][frag_col + 16]);
#pragma unroll
for (int mt = 0; mt < 4; ++mt) {
const int a_row0 = warp_m * 64 + mt * 16 + group;
unsigned a_frag[4];
a_frag[0] = *reinterpret_cast<const unsigned*>(
&a_tile[stage][a_row0][frag_col]);
a_frag[1] = *reinterpret_cast<const unsigned*>(
&a_tile[stage][a_row0 + 8][frag_col]);
a_frag[2] = *reinterpret_cast<const unsigned*>(
&a_tile[stage][a_row0][frag_col + 16]);
a_frag[3] = *reinterpret_cast<const unsigned*>(
&a_tile[stage][a_row0 + 8][frag_col + 16]);
mma_fp8_16832(acc + (nt * 4 + mt) * 4, a_frag, b_frag);
}
}
}
// Barrier 2: every thread finished reading this stage's tiles before
// the prefetch for the (i+3)-th tile overwrites them.
__syncthreads();
if (tile_index + 3 < tile_count) {
load_tile(stage, (tile_index + 3) * kPqK);
asm volatile("cp.async.commit_group;");
}
}
const float output_scale = scale * out_scale;
#pragma unroll
for (int nt = 0; nt < 2; ++nt) {
const int64_t col = output_col + nt * 8;
// Per-row store: FP8 packs two adjacent columns into one 16-bit
// write; the BF16 path writes two scalars. Boundary columns fall
// back to a scalar convert so the pack never crosses the row edge.
auto store_out = [&](int64_t row, float v0, float v1) {
if (row >= m) return;
if constexpr (OutFp8) {
if (col + 1 < n) {
*reinterpret_cast<unsigned short*>(
out_fp8 + row * n + col) =
static_cast<unsigned short>(__nv_cvt_float2_to_fp8x2(
make_float2(v0 * output_scale, v1 * output_scale),
__NV_SATFINITE, __NV_E4M3));
} else {
out_fp8[row * n + col] = __nv_fp8_e4m3(v0 * output_scale);
}
} else {
out_bf16[row * n + col] = __float2bfloat16(v0 * scale);
if (col + 1 < n)
out_bf16[row * n + col + 1] = __float2bfloat16(v1 * scale);
}
};
#pragma unroll
for (int mt = 0; mt < 4; ++mt) {
const int64_t row0 = row_base + mt * 16;
float* tile_acc = acc + (nt * 4 + mt) * 4;
if (col < n) {
store_out(row0, tile_acc[0], tile_acc[1]);
store_out(row0 + 8, tile_acc[2], tile_acc[3]);
}
}
}
}
template <bool AddBias = false, bool TrackAmax = true>
void launch_fused_fp8_gemm_fast(
const torch::Tensor& a, const torch::Tensor& b, torch::Tensor& out,
const torch::Tensor& bias, const torch::Tensor& scale_a,
const torch::Tensor& scale_b, torch::Tensor* amax_a,
torch::Tensor* amax_b, int64_t m, int64_t n, int64_t k,
cudaStream_t stream) {
dim3 grid((n + kFastBlockN - 1) / kFastBlockN,
(m + kFastBlockM - 1) / kFastBlockM);
const auto* bias_ptr = AddBias
? reinterpret_cast<const __nv_bfloat16*>(bias.data_ptr())
: nullptr;
auto kernel = fused_fp8_gemm_fast_kernel<AddBias, TrackAmax>;
static bool attribute_set = false;
if (!attribute_set) {
C10_CUDA_CHECK(cudaFuncSetAttribute(
kernel, cudaFuncAttributeMaxDynamicSharedMemorySize,
kFastSmemBytes));
attribute_set = true;
}
kernel<<<grid, kWarps * 32, kFastSmemBytes, stream>>>(
reinterpret_cast<const __nv_bfloat16*>(a.data_ptr()),
reinterpret_cast<const __nv_bfloat16*>(b.data_ptr()),
reinterpret_cast<__nv_bfloat16*>(out.data_ptr()), bias_ptr,
scale_a.data_ptr<float>(), scale_b.data_ptr<float>(),
amax_a ? amax_a->data_ptr<float>() : nullptr,
amax_b ? amax_b->data_ptr<float>() : nullptr, m, n, k);
}
void check_fp8_device(const torch::Tensor& tensor) {
static std::mutex mutex;
static std::unordered_map<int, bool> supported;
const int device = tensor.device().index();
{
std::lock_guard<std::mutex> lock(mutex);
auto cached = supported.find(device);
if (cached != supported.end()) {
TORCH_CHECK(cached->second,
"fused FP8 MMA requires compute capability 8.9 or newer");
return;
}
}
const auto* properties = at::cuda::getDeviceProperties(device);
const bool is_supported = properties->major > 8 ||
(properties->major == 8 && properties->minor >= 9);
{
std::lock_guard<std::mutex> lock(mutex);
supported.emplace(device, is_supported);
}
TORCH_CHECK(is_supported,
"fused FP8 MMA requires compute capability 8.9 or newer");
}
void check_scale(const torch::Tensor& scale, const torch::Tensor& input,
const char* name) {
TORCH_CHECK(scale.is_cuda() && scale.device() == input.device() &&
scale.scalar_type() == torch::kFloat32 && scale.numel() == 1,
name, " must be a CUDA float32 scalar on the input device");
}
} // namespace
torch::Tensor fp8_mm(torch::Tensor a, torch::Tensor b, torch::Tensor sx,
torch::Tensor sw) {
TORCH_CHECK(a.is_cuda() && b.is_cuda(), "CUDA tensors required");
TORCH_CHECK(a.scalar_type() == torch::kBFloat16 &&
b.scalar_type() == torch::kBFloat16,
"a and b must be bf16");
TORCH_CHECK(a.dim() == 2 && b.dim() == 2, "a and b must be 2D");
TORCH_CHECK(a.device() == b.device(), "a and b must be on the same device");
TORCH_CHECK(a.size(1) == b.size(1), "inner dim mismatch");
check_scale(sx, a, "sx");
check_scale(sw, a, "sw");
check_fp8_device(a);
const at::cuda::OptionalCUDAGuard guard(a.device());
auto stream = at::cuda::getCurrentCUDAStream();
auto a_c = a.contiguous();
auto b_c = b.contiguous();
auto out = torch::empty({a_c.size(0), b_c.size(0)}, a_c.options());
torch::Tensor no_bias;
launch_fused_fp8_gemm_fast<false, false>(
a_c, b_c, out, no_bias, sx, sw, nullptr, nullptr,
a_c.size(0), b_c.size(0), a_c.size(1), stream.stream());
C10_CUDA_CHECK(cudaGetLastError());
return out;
}
torch::Tensor fp8_linear_forward_scaled(
torch::Tensor x, torch::Tensor w, torch::Tensor bias, torch::Tensor sx,
torch::Tensor sw, torch::Tensor sx_inv, torch::Tensor sw_inv,
torch::Tensor amax_x, torch::Tensor amax_w) {
TORCH_CHECK(x.is_cuda() && w.is_cuda(), "CUDA tensors required");
TORCH_CHECK(x.scalar_type() == torch::kBFloat16 &&
w.scalar_type() == torch::kBFloat16,
"x and w must be bf16");
TORCH_CHECK(x.device() == w.device(), "x and w must be on the same device");
check_scale(sx, x, "sx");
check_scale(sw, x, "sw");
check_fp8_device(x);
const at::cuda::OptionalCUDAGuard guard(x.device());
auto stream = at::cuda::getCurrentCUDAStream();
auto x_c = x.reshape({-1, w.size(1)}).contiguous();
auto w_c = w.contiguous();
int64_t m = x_c.size(0), k = x_c.size(1), n = w_c.size(0);
TORCH_CHECK(w_c.dim() == 2 && w_c.size(1) == k, "inner dim mismatch");
C10_CUDA_CHECK(cudaMemsetAsync(amax_x.data_ptr<float>(), 0, sizeof(float),
stream.stream()));
C10_CUDA_CHECK(cudaMemsetAsync(amax_w.data_ptr<float>(), 0, sizeof(float),
stream.stream()));
auto out = torch::empty({m, n}, x_c.options());
if (bias.defined() && bias.numel() > 0) {
TORCH_CHECK(bias.is_cuda() && bias.device() == x.device() &&
bias.scalar_type() == torch::kBFloat16 &&
bias.numel() == n,
"bias must be CUDA bf16 with shape [N]");
launch_fused_fp8_gemm_fast<true, true>(
x_c, w_c, out, bias, sx, sw, &amax_x, &amax_w,
m, n, k, stream.stream());
} else {
launch_fused_fp8_gemm_fast<false, true>(
x_c, w_c, out, bias, sx, sw, &amax_x, &amax_w,
m, n, k, stream.stream());
}
C10_CUDA_CHECK(cudaGetLastError());
(void)sx_inv;
(void)sw_inv;
std::vector<int64_t> shape(x.sizes().begin(), x.sizes().end() - 1);
shape.push_back(n);
return out.reshape(shape);
}
std::tuple<torch::Tensor, torch::Tensor, torch::Tensor> fp8_linear_backward_scaled(
torch::Tensor g, torch::Tensor x, torch::Tensor w,
std::vector<int64_t> masks, torch::Tensor sg, torch::Tensor sw,
torch::Tensor sx, torch::Tensor sg_inv, torch::Tensor sw_inv,
torch::Tensor sx_inv, torch::Tensor amax_g) {
TORCH_CHECK(g.is_cuda() && x.is_cuda() && w.is_cuda(), "CUDA tensors required");
TORCH_CHECK(g.scalar_type() == torch::kBFloat16 &&
x.scalar_type() == torch::kBFloat16 &&
w.scalar_type() == torch::kBFloat16,
"g, x, and w must be bf16");
TORCH_CHECK(g.device() == x.device() && g.device() == w.device(),
"g, x, and w must be on the same device");
TORCH_CHECK(masks.size() == 3, "masks must contain three values");
check_fp8_device(g);
const at::cuda::OptionalCUDAGuard guard(g.device());
auto stream = at::cuda::getCurrentCUDAStream();
auto g_c = g.reshape({-1, w.size(0)}).contiguous();
auto x_c = x.reshape({-1, x.size(-1)}).contiguous();
auto w_c = w.contiguous();
int64_t m = g_c.size(0), n = w_c.size(0), k = w_c.size(1);
TORCH_CHECK(x_c.size(0) == m && x_c.size(1) == k && g_c.size(1) == n,
"backward shape mismatch");
auto grad_input = torch::empty_like(x);
auto grad_weight = torch::empty_like(w);
auto grad_bias = torch::empty({0}, g.options());
C10_CUDA_CHECK(cudaMemsetAsync(amax_g.data_ptr<float>(), 0, sizeof(float),
stream.stream()));
torch::Tensor no_bias;
bool recorded_amax = false;
if (masks[0]) {
auto grad_input_2d = grad_input.reshape({m, k});
// The fast kernel computes A @ B^T. A contiguous W^T makes dX use
// the same coalesced forward tile path instead of scalar fragments.
auto w_t = w_c.transpose(0, 1).contiguous();
launch_fused_fp8_gemm_fast<false, true>(
g_c, w_t, grad_input_2d, no_bias, sg, sw, &amax_g, nullptr,
m, k, n, stream.stream());
recorded_amax = true;
}
if (masks[1]) {
// dW = G^T @ X, expressed as (G^T) @ (X^T)^T for the same kernel.
auto g_t = g_c.transpose(0, 1).contiguous();
auto x_t = x_c.transpose(0, 1).contiguous();
if (recorded_amax) {
launch_fused_fp8_gemm_fast<false, false>(
g_t, x_t, grad_weight, no_bias, sg, sx, nullptr, nullptr,
n, k, m, stream.stream());
} else {
launch_fused_fp8_gemm_fast<false, true>(
g_t, x_t, grad_weight, no_bias, sg, sx, &amax_g, nullptr,
n, k, m, stream.stream());
}
recorded_amax = true;
}
if (!recorded_amax) {
amax_g.copy_(g_c.abs().amax().to(torch::kFloat32));
}
C10_CUDA_CHECK(cudaGetLastError());
if (masks[2]) grad_bias = g_c.sum(0).to(g.scalar_type());
(void)sg_inv;
(void)sw_inv;
(void)sx_inv;
return {grad_input, grad_weight, grad_bias};
}
torch::Tensor fp8_mm_prequant(torch::Tensor a, torch::Tensor b,
torch::Tensor scale) {
TORCH_CHECK(a.is_cuda() && b.is_cuda(), "CUDA tensors required");
TORCH_CHECK(a.scalar_type() == torch::kFloat8_e4m3fn &&
b.scalar_type() == torch::kFloat8_e4m3fn,
"a and b must be fp8_e4m3fn");
TORCH_CHECK(a.dim() == 2 && b.dim() == 2, "a and b must be 2D");
TORCH_CHECK(a.device() == b.device(), "a and b must be on the same device");
TORCH_CHECK(a.size(1) == b.size(1), "inner dim mismatch");
check_scale(scale, a, "scale");
check_fp8_device(a);
const at::cuda::OptionalCUDAGuard guard(a.device());
auto stream = at::cuda::getCurrentCUDAStream();
auto a_c = a.contiguous();
auto b_c = b.contiguous();
int64_t m = a_c.size(0), k = a_c.size(1), n = b_c.size(0);
auto out = torch::empty({m, n},
a_c.options().dtype(torch::kBFloat16));
const float scale_value = scale.item<float>();
dim3 grid((n + kPqBlockN - 1) / kPqBlockN,
(m + kPqBlockM - 1) / kPqBlockM);
fp8_mm_pq_kernel<false><<<grid, kWarps * 32, 0, stream>>>(
reinterpret_cast<const __nv_fp8_e4m3*>(a_c.data_ptr()),
reinterpret_cast<const __nv_fp8_e4m3*>(b_c.data_ptr()),
reinterpret_cast<__nv_bfloat16*>(out.data_ptr()), nullptr,
scale_value, 1.0f, m, n, k);
C10_CUDA_CHECK(cudaGetLastError());
return out;
}
torch::Tensor fp8_mm_prequant_fp8(torch::Tensor a, torch::Tensor b,
torch::Tensor scale,
torch::Tensor out_scale) {
TORCH_CHECK(a.is_cuda() && b.is_cuda(), "CUDA tensors required");
TORCH_CHECK(a.scalar_type() == torch::kFloat8_e4m3fn &&
b.scalar_type() == torch::kFloat8_e4m3fn,
"a and b must be fp8_e4m3fn");
TORCH_CHECK(a.dim() == 2 && b.dim() == 2, "a and b must be 2D");
TORCH_CHECK(a.device() == b.device(), "a and b must be on the same device");
TORCH_CHECK(a.size(1) == b.size(1), "inner dim mismatch");
check_scale(scale, a, "scale");
check_scale(out_scale, a, "out_scale");
check_fp8_device(a);
const at::cuda::OptionalCUDAGuard guard(a.device());
auto stream = at::cuda::getCurrentCUDAStream();
auto a_c = a.contiguous();
auto b_c = b.contiguous();
int64_t m = a_c.size(0), k = a_c.size(1), n = b_c.size(0);
auto out = torch::empty({m, n}, a_c.options());
const float scale_value = scale.item<float>();
const float out_scale_value = out_scale.item<float>();
dim3 grid((n + kPqBlockN - 1) / kPqBlockN,
(m + kPqBlockM - 1) / kPqBlockM);
fp8_mm_pq_kernel<true><<<grid, kWarps * 32, 0, stream>>>(
reinterpret_cast<const __nv_fp8_e4m3*>(a_c.data_ptr()),
reinterpret_cast<const __nv_fp8_e4m3*>(b_c.data_ptr()), nullptr,
reinterpret_cast<__nv_fp8_e4m3*>(out.data_ptr()),
scale_value, out_scale_value, m, n, k);
C10_CUDA_CHECK(cudaGetLastError());
return out;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("fp8_mm", &fp8_mm, py::arg("a"), py::arg("b"), py::arg("sx"),
py::arg("sw"),
"Fused BF16 input, E4M3 MMA, FP32 accumulation, BF16 output GEMM");
m.def("fp8_mm_prequant", &fp8_mm_prequant, py::arg("a"), py::arg("b"),
py::arg("scale"),
"Pre-quantized FP8 GEMM with FP32 accumulation and BF16 output");
m.def("fp8_mm_prequant_fp8", &fp8_mm_prequant_fp8, py::arg("a"),
py::arg("b"), py::arg("scale"), py::arg("out_scale"),
"Pre-quantized FP8 GEMM with FP32 accumulation and FP8 output");
m.def("fp8_linear_forward_scaled", &fp8_linear_forward_scaled,
py::arg("x"), py::arg("w"), py::arg("bias"), py::arg("sx"),
py::arg("sw"), py::arg("sx_inv"), py::arg("sw_inv"),
py::arg("amax_x"), py::arg("amax_w"),
"Fused BF16-to-FP8 linear forward with FP32 accumulation");
m.def("fp8_linear_backward_scaled", &fp8_linear_backward_scaled,
py::arg("g"), py::arg("x"), py::arg("w"), py::arg("masks"),
py::arg("sg"), py::arg("sw"), py::arg("sx"), py::arg("sg_inv"),
py::arg("sw_inv"), py::arg("sx_inv"), py::arg("amax_g"),
"Fused BF16-to-FP8 linear backward with FP32 accumulation");
}
+9 -7
View File
@@ -7,7 +7,9 @@
#include <cstring>
#include <vector>
#include "test_utils.cuh"
#include "../kernels/attn_dispatchers.cuh"
#include "../kernels/attention/dispatchers.cuh"
using namespace astrai::attention;
struct PagedDecodeDispatch { AttentionParams<bf16>& p; template<int H> void operator()() { dispatch_paged_decode<H>(p, 0); } };
struct PagedPrefillDispatch { AttentionParams<bf16>& p; template<int H> void operator()() { dispatch_paged_prefill<H>(p, 0); } };
@@ -236,7 +238,7 @@ static int run_decode_test(int B, int Hq, int Hkv, int max_seq,
B, Hq, Hkv, HEAD_DIM, max_ctx, h_o_ref);
// Kernel launch
AttentionParams<bf16> p;
AttentionParams<bf16> p = {};
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
@@ -371,7 +373,7 @@ static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq,
h_mask, max_sl,
B, Hq, Hkv, HEAD_DIM, max_ctx, h_o_ref);
AttentionParams<bf16> p;
AttentionParams<bf16> p = {};
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
@@ -507,7 +509,7 @@ static int run_prefill_test(int B, int Hq, int Hkv,
int num_q_tiles = make_q_tile_mapping(q_lens, &d_qtb, &d_qti);
// Kernel launch
AttentionParams<bf16> p;
AttentionParams<bf16> p = {};
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
@@ -649,7 +651,7 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) {
int *d_qtb, *d_qti;
int num_q_tiles = make_q_tile_mapping(q_lens, &d_qtb, &d_qti);
AttentionParams<bf16> p;
AttentionParams<bf16> p = {};
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
@@ -742,7 +744,7 @@ static void bench_decode(int B, int Hq, int Hkv, int seq_len) {
for (int b = 0; b < B; b++) h_kvi[b + 1] = h_kvi[b] + seq_len;
cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice);
AttentionParams<bf16> p;
AttentionParams<bf16> p = {};
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
@@ -823,7 +825,7 @@ static void bench_prefill(int B, int Hq, int Hkv, int q_len, int kv_len, int cau
int *d_qtb, *d_qti;
int num_q_tiles = make_q_tile_mapping(q_lens, &d_qtb, &d_qti);
AttentionParams<bf16> p;
AttentionParams<bf16> p = {};
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
p.head_dim = HEAD_DIM;
p.q_l_stride = Hq * HEAD_DIM; p.q_h_stride = HEAD_DIM; p.q_d_stride = 1;
+7 -5
View File
@@ -7,7 +7,9 @@ nvcc -I csrc -arch=sm_89 -O3 \
*/
#include "test_utils.cuh"
#include "../kernels/attn_dispatchers.cuh"
#include "../kernels/attention/dispatchers.cuh"
using namespace astrai::attention;
struct DecodeDispatch { AttentionParams<bf16>& p; template<int H> void operator()() { dispatch_decode<H>(p, 0); } };
struct PrefillDispatch { AttentionParams<bf16>& p; template<int H> void operator()() { dispatch_prefill<H>(p, 0); } };
@@ -56,7 +58,7 @@ static int run_decode_test(int B, int Hq, int Hk, int sl, int D, int causal) {
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
cudaMemcpy(dMask,hMask,B*sl,cudaMemcpyHostToDevice);
AttentionParams<bf16> p;
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);
@@ -137,7 +139,7 @@ static void bench_decode() {
cudaMemcpy(dV, tmp, nKV*2, cudaMemcpyHostToDevice);
delete[] tmp;
AttentionParams<bf16> p;
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);
@@ -184,7 +186,7 @@ static int run_prefill_test(int B, int Hq, int Hk, int ql, int kl, int D, int ca
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
AttentionParams<bf16> p;
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);
@@ -264,7 +266,7 @@ static void bench_prefill() {
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(randf());
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
AttentionParams<bf16> p;
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);
-155
View File
@@ -1,155 +0,0 @@
/*
Single-kernel BF16 -> FP8 MMA -> BF16 demo for Ada (sm_89).
nvcc -I csrc -arch=sm_89 -std=c++17 -O3 --use_fast_math \
--ptxas-options=-O3,-v csrc/tests/fp8_mma_test.cu -o fp8_mma_test \
&& ./fp8_mma_test
*/
#include "test_utils.cuh"
#include <cuda_fp8.h>
#include <algorithm>
#include <vector>
constexpr int M = 16;
constexpr int N = 8;
constexpr int K = 32;
__device__ __forceinline__ unsigned pack_fp8x4(float x0, float x1, float x2,
float x3) {
__nv_fp8_e4m3 q0(x0);
__nv_fp8_e4m3 q1(x1);
__nv_fp8_e4m3 q2(x2);
__nv_fp8_e4m3 q3(x3);
return static_cast<unsigned>(q0.__x) |
(static_cast<unsigned>(q1.__x) << 8) |
(static_cast<unsigned>(q2.__x) << 16) |
(static_cast<unsigned>(q3.__x) << 24);
}
__device__ __forceinline__ unsigned load_quantize_fp8x4(
const bf16* src, float scale_inv) {
return pack_fp8x4(__bfloat162float(src[0]) * scale_inv,
__bfloat162float(src[1]) * scale_inv,
__bfloat162float(src[2]) * scale_inv,
__bfloat162float(src[3]) * scale_inv);
}
__device__ __forceinline__ void mma_fp8_16832(float d[4],
const unsigned a[4],
const unsigned b[2]) {
asm volatile(
"mma.sync.aligned.m16n8k32.row.col.f32.e4m3.e4m3.f32 "
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%0,%1,%2,%3};"
: "+f"(d[0]), "+f"(d[1]), "+f"(d[2]), "+f"(d[3])
: "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]),
"r"(b[0]), "r"(b[1]));
}
__global__ void fused_bf16_fp8_mma_kernel(
const bf16* __restrict__ a, const bf16* __restrict__ b,
bf16* __restrict__ out, float scale_a, float scale_b) {
const int lane = threadIdx.x;
const int group = lane >> 2;
const int thread_in_group = lane & 3;
const int k0 = thread_in_group * 4;
// PTX m16n8k32 A fragment: two rows, two 16-column K partitions.
unsigned a_frag[4];
a_frag[0] = load_quantize_fp8x4(&a[group * K + k0], 1.0f / scale_a);
a_frag[1] = load_quantize_fp8x4(&a[(group + 8) * K + k0], 1.0f / scale_a);
a_frag[2] = load_quantize_fp8x4(&a[group * K + k0 + 16], 1.0f / scale_a);
a_frag[3] = load_quantize_fp8x4(&a[(group + 8) * K + k0 + 16],
1.0f / scale_a);
// B is supplied as row-major [N,K], equivalent to the col-major [K,N]
// operand required by the MMA instruction.
unsigned b_frag[2];
b_frag[0] = load_quantize_fp8x4(&b[group * K + k0], 1.0f / scale_b);
b_frag[1] = load_quantize_fp8x4(&b[group * K + k0 + 16], 1.0f / scale_b);
float acc[4] = {0.0f, 0.0f, 0.0f, 0.0f};
mma_fp8_16832(acc, a_frag, b_frag);
const int col = thread_in_group * 2;
const float output_scale = scale_a * scale_b;
*reinterpret_cast<__nv_bfloat162*>(&out[group * N + col]) =
__floats2bfloat162_rn(acc[0] * output_scale,
acc[1] * output_scale);
*reinterpret_cast<__nv_bfloat162*>(&out[(group + 8) * N + col]) =
__floats2bfloat162_rn(acc[2] * output_scale,
acc[3] * output_scale);
}
static float quantize_e4m3(float value) {
return static_cast<float>(__nv_fp8_e4m3(value));
}
int main() {
srand(0);
std::vector<float> a(M * K), b(N * K), reference(M * N, 0.0f);
std::vector<bf16> a_bf16(M * K), b_bf16(N * K), output(M * N);
for (float& value : a) value = randf() * 4.0f;
for (float& value : b) value = randf() * 4.0f;
for (int i = 0; i < M * K; ++i) {
a_bf16[i] = f2bf(a[i]);
a[i] = bf2f(a_bf16[i]);
}
for (int i = 0; i < N * K; ++i) {
b_bf16[i] = f2bf(b[i]);
b[i] = bf2f(b_bf16[i]);
}
const float amax = *std::max_element(
a.begin(), a.end(), [](float x, float y) { return fabsf(x) < fabsf(y); });
const float bmax = *std::max_element(
b.begin(), b.end(), [](float x, float y) { return fabsf(x) < fabsf(y); });
const float scale_a = fabsf(amax) / 448.0f;
const float scale_b = fabsf(bmax) / 448.0f;
for (int row = 0; row < M; ++row) {
for (int col = 0; col < N; ++col) {
float sum = 0.0f;
for (int k = 0; k < K; ++k) {
float qa = quantize_e4m3(a[row * K + k] / scale_a);
float qb = quantize_e4m3(b[col * K + k] / scale_b);
sum = fmaf(qa, qb, sum);
}
reference[row * N + col] = sum * scale_a * scale_b;
}
}
bf16 *d_a, *d_b, *d_out;
CUDA_CHECK(cudaMalloc(&d_a, a_bf16.size() * sizeof(bf16)));
CUDA_CHECK(cudaMalloc(&d_b, b_bf16.size() * sizeof(bf16)));
CUDA_CHECK(cudaMalloc(&d_out, output.size() * sizeof(bf16)));
CUDA_CHECK(cudaMemcpy(d_a, a_bf16.data(), a_bf16.size() * sizeof(bf16),
cudaMemcpyHostToDevice));
CUDA_CHECK(cudaMemcpy(d_b, b_bf16.data(), b_bf16.size() * sizeof(bf16),
cudaMemcpyHostToDevice));
fused_bf16_fp8_mma_kernel<<<1, 32>>>(d_a, d_b, d_out, scale_a, scale_b);
CUDA_CHECK(cudaDeviceSynchronize());
CUDA_CHECK(cudaMemcpy(output.data(), d_out, output.size() * sizeof(bf16),
cudaMemcpyDeviceToHost));
float max_abs_error = 0.0f;
float max_rel_error = 0.0f;
for (int i = 0; i < M * N; ++i) {
float error = fabsf(bf2f(output[i]) - reference[i]);
max_abs_error = fmaxf(max_abs_error, error);
max_rel_error = fmaxf(max_rel_error,
error / fmaxf(fabsf(reference[i]), 1e-4f));
}
const bool pass = max_abs_error < 0.05f;
print_test_header();
print_test_row("M=16 N=8 K=32 fused BF16->E4M3 MMA", max_abs_error,
max_rel_error, pass);
cudaFree(d_a);
cudaFree(d_b);
cudaFree(d_out);
return pass ? 0 : 1;
}
+312
View File
@@ -0,0 +1,312 @@
/*
FP8 family tests: single-warp MMA demo + full GEMM correctness.
Part 1 exercises one bf16 -> fp8 -> mma.sync m16n8k32 instruction pair
(sanity for astrai::mma_sync + the fragment layout contract).
Part 2 checks launch_fp8_gemm across all four operand layouts, both K
tiles, and ragged shapes against an fp32 CPU reference.
nvcc -I csrc -arch=sm_89 -std=c++17 -O3 csrc/tests/fp8_test.cu -o /tmp/fp8_test \
&& /tmp/fp8_test
*/
#include "test_utils.cuh"
#include <cuda_fp8.h>
#include <algorithm>
#include <cmath>
#include <cstdio>
#include <cstdlib>
#include <cuda_runtime.h>
#include <type_traits>
#include <vector>
#include "../kernels/common/mma.cuh"
#include "../kernels/fp8/gemm.cuh"
using namespace astrai::fp8;
// ---------------------------------------------------------------------------
// Part 1: single-kernel BF16 -> FP8 MMA -> BF16 demo (m16n8k32)
// ---------------------------------------------------------------------------
namespace {
constexpr int kMmaM = 16;
constexpr int kMmaN = 8;
constexpr int kMmaK = 32;
__device__ __forceinline__ unsigned pack_fp8x4(float x0, float x1, float x2,
float x3) {
__nv_fp8_e4m3 q0(x0);
__nv_fp8_e4m3 q1(x1);
__nv_fp8_e4m3 q2(x2);
__nv_fp8_e4m3 q3(x3);
return static_cast<unsigned>(q0.__x) |
(static_cast<unsigned>(q1.__x) << 8) |
(static_cast<unsigned>(q2.__x) << 16) |
(static_cast<unsigned>(q3.__x) << 24);
}
__device__ __forceinline__ unsigned load_quantize_fp8x4(
const bf16* src, float scale_inv) {
return pack_fp8x4(__bfloat162float(src[0]) * scale_inv,
__bfloat162float(src[1]) * scale_inv,
__bfloat162float(src[2]) * scale_inv,
__bfloat162float(src[3]) * scale_inv);
}
__global__ void fused_bf16_fp8_mma_kernel(
const bf16* __restrict__ a, const bf16* __restrict__ b,
bf16* __restrict__ out, float scale_a, float scale_b) {
const int lane = threadIdx.x;
const int group = lane >> 2;
const int thread_in_group = lane & 3;
const int k0 = thread_in_group * 4;
// PTX m16n8k32 A fragment: two rows, two 16-column K partitions.
unsigned a_frag[4];
a_frag[0] = load_quantize_fp8x4(&a[group * kMmaK + k0], 1.0f / scale_a);
a_frag[1] =
load_quantize_fp8x4(&a[(group + 8) * kMmaK + k0], 1.0f / scale_a);
a_frag[2] =
load_quantize_fp8x4(&a[group * kMmaK + k0 + 16], 1.0f / scale_a);
a_frag[3] = load_quantize_fp8x4(&a[(group + 8) * kMmaK + k0 + 16],
1.0f / scale_a);
// B is supplied as row-major [N,K], equivalent to the col-major [K,N]
// operand required by the MMA instruction.
unsigned b_frag[2];
b_frag[0] = load_quantize_fp8x4(&b[group * kMmaK + k0], 1.0f / scale_b);
b_frag[1] =
load_quantize_fp8x4(&b[group * kMmaK + k0 + 16], 1.0f / scale_b);
float acc[4] = {0.0f, 0.0f, 0.0f, 0.0f};
astrai::mma_sync<__nv_fp8_e4m3>(acc, a_frag, b_frag, acc);
const int col = thread_in_group * 2;
const float output_scale = scale_a * scale_b;
*reinterpret_cast<__nv_bfloat162*>(&out[group * kMmaN + col]) =
__floats2bfloat162_rn(acc[0] * output_scale, acc[1] * output_scale);
*reinterpret_cast<__nv_bfloat162*>(&out[(group + 8) * kMmaN + col]) =
__floats2bfloat162_rn(acc[2] * output_scale, acc[3] * output_scale);
}
static float quantize_e4m3(float value) {
return static_cast<float>(__nv_fp8_e4m3(value));
}
static bool test_single_mma() {
srand(0);
std::vector<float> a(kMmaM * kMmaK), b(kMmaN * kMmaK),
reference(kMmaM * kMmaN, 0.0f);
std::vector<bf16> a_bf16(kMmaM * kMmaK), b_bf16(kMmaN * kMmaK),
output(kMmaM * kMmaN);
for (float& value : a) value = randf() * 4.0f;
for (float& value : b) value = randf() * 4.0f;
for (int i = 0; i < kMmaM * kMmaK; ++i) {
a_bf16[i] = f2bf(a[i]);
a[i] = bf2f(a_bf16[i]);
}
for (int i = 0; i < kMmaN * kMmaK; ++i) {
b_bf16[i] = f2bf(b[i]);
b[i] = bf2f(b_bf16[i]);
}
const float amax = *std::max_element(
a.begin(), a.end(),
[](float x, float y) { return fabsf(x) < fabsf(y); });
const float bmax = *std::max_element(
b.begin(), b.end(),
[](float x, float y) { return fabsf(x) < fabsf(y); });
const float scale_a = fabsf(amax) / 448.0f;
const float scale_b = fabsf(bmax) / 448.0f;
for (int row = 0; row < kMmaM; ++row) {
for (int col = 0; col < kMmaN; ++col) {
float sum = 0.0f;
for (int k = 0; k < kMmaK; ++k) {
float qa = quantize_e4m3(a[row * kMmaK + k] / scale_a);
float qb = quantize_e4m3(b[col * kMmaK + k] / scale_b);
sum = fmaf(qa, qb, sum);
}
reference[row * kMmaN + col] = sum * scale_a * scale_b;
}
}
bf16 *d_a, *d_b, *d_out;
CUDA_CHECK(cudaMalloc(&d_a, a_bf16.size() * sizeof(bf16)));
CUDA_CHECK(cudaMalloc(&d_b, b_bf16.size() * sizeof(bf16)));
CUDA_CHECK(cudaMalloc(&d_out, output.size() * sizeof(bf16)));
CUDA_CHECK(cudaMemcpy(d_a, a_bf16.data(), a_bf16.size() * sizeof(bf16),
cudaMemcpyHostToDevice));
CUDA_CHECK(cudaMemcpy(d_b, b_bf16.data(), b_bf16.size() * sizeof(bf16),
cudaMemcpyHostToDevice));
fused_bf16_fp8_mma_kernel<<<1, 32>>>(d_a, d_b, d_out, scale_a, scale_b);
CUDA_CHECK(cudaDeviceSynchronize());
CUDA_CHECK(cudaMemcpy(output.data(), d_out, output.size() * sizeof(bf16),
cudaMemcpyDeviceToHost));
float max_abs_error = 0.0f;
float max_rel_error = 0.0f;
for (int i = 0; i < kMmaM * kMmaN; ++i) {
float error = fabsf(bf2f(output[i]) - reference[i]);
max_abs_error = fmaxf(max_abs_error, error);
max_rel_error = fmaxf(
max_rel_error, error / fmaxf(fabsf(reference[i]), 1e-4f));
}
const bool pass = max_abs_error < 0.05f;
print_test_row("M=16 N=8 K=32 fused BF16->E4M3 MMA", max_abs_error,
max_rel_error, pass);
cudaFree(d_a);
cudaFree(d_b);
cudaFree(d_out);
return pass;
}
// ---------------------------------------------------------------------------
// Part 2: GEMM correctness — layouts x K-tiles vs fp32 CPU reference
// ---------------------------------------------------------------------------
template <typename LA, typename LB, int kK, int Stages>
static bool run_gemm_case(const float* ha, const float* hb, int m, int n,
int k, int a_ld, int b_ld) {
__nv_fp8_e4m3 *da, *db;
__nv_bfloat16* dout;
float *dsa, *dsb;
cudaMalloc(&da, (size_t)m * k);
cudaMalloc(&db, (size_t)n * k);
cudaMalloc(&dout, (size_t)m * n * 2);
cudaMalloc(&dsa, 4);
cudaMalloc(&dsb, 4);
float one = 1.0f;
cudaMemcpy(dsa, &one, 4, cudaMemcpyHostToDevice);
cudaMemcpy(dsb, &one, 4, cudaMemcpyHostToDevice);
// quantize inputs to e4m3 on host and upload byte-by-byte
std::vector<unsigned char> qa(m * k), qb(n * k);
for (int i = 0; i < m * k; ++i) {
__nv_fp8_e4m3 q(ha[i]);
qa[i] = *(unsigned char*)&q;
}
for (int i = 0; i < n * k; ++i) {
__nv_fp8_e4m3 q(hb[i]);
qb[i] = *(unsigned char*)&q;
}
cudaMemcpy(da, qa.data(), qa.size(), cudaMemcpyHostToDevice);
cudaMemcpy(db, qb.data(), qb.size(), cudaMemcpyHostToDevice);
FP8Params p = {};
p.a_ptr = da;
p.b_ptr = db;
p.out_ptr = dout;
p.scale_a = dsa;
p.scale_b = dsb;
p.m = m;
p.n = n;
p.k = k;
p.a_ld = a_ld;
p.b_ld = b_ld;
launch_fp8_gemm<FP8Format::E4M3, false, LA, LB, kK, Stages>(p, 0);
cudaError_t e = cudaDeviceSynchronize();
if (e != cudaSuccess) {
printf(" CUDA err: %s\n", cudaGetErrorString(e));
return false;
}
std::vector<unsigned short> hb16(m * n);
cudaMemcpy(hb16.data(), dout, (size_t)m * n * 2, cudaMemcpyDeviceToHost);
const float tol = 0.06f;
double max_rel = 0;
bool ok = true;
for (int i = 0; i < m && ok; ++i) {
for (int j = 0; j < n && ok; ++j) {
float ref = 0;
for (int kk = 0; kk < k; ++kk) {
// A reference reads the actual uploaded buffer: LA ColMajor
// means the buffer is [K][M] (ha_t), else [M][K].
float av = std::is_same_v<LA, ColMajor>
? (float)__nv_fp8_e4m3(ha[kk * m + i])
: (float)__nv_fp8_e4m3(ha[i * k + kk]);
float bv;
if (std::is_same_v<LB, ColMajor>)
bv = (float)__nv_fp8_e4m3(hb[j * k + kk]);
else
bv = (float)__nv_fp8_e4m3(hb[kk * n + j]);
ref += av * bv;
}
float got =
__bfloat162float(__ushort_as_bfloat16(hb16[i * n + j]));
float err = fabsf(got - ref);
float rel = err / fmaxf(fabsf(ref), 0.5f);
if (rel > max_rel) max_rel = rel;
if (err > tol * fmaxf(fabsf(ref), 1.0f)) ok = false;
}
}
printf(" max_rel=%.4f %s\n", max_rel, ok ? "PASS" : "FAIL");
cudaFree(da);
cudaFree(db);
cudaFree(dout);
cudaFree(dsa);
cudaFree(dsb);
return ok;
}
static bool test_gemm() {
struct {
int m, n, k;
} cfgs[] = {
{128, 128, 128}, {256, 128, 256}, {128, 256, 64},
{100, 130, 96}, {64, 64, 160}, {300, 200, 320},
};
bool all = true;
for (auto& c : cfgs) {
float* ha = new float[c.m * c.k];
float* hb_rowmajor = new float[c.k * c.n]; // [K][N] for B RowMajor
float* hb_colmajor = new float[c.n * c.k]; // [N][K] for B ColMajor
for (int i = 0; i < c.m * c.k; ++i) ha[i] = randf();
for (int i = 0; i < c.k * c.n; ++i) hb_rowmajor[i] = randf();
for (int i = 0; i < c.k * c.n; ++i)
hb_colmajor[i / c.k * c.k + i % c.k] = hb_rowmajor[i];
float* ha_t = new float[c.k * c.m]; // [K][M] for A ColMajor
for (int i = 0; i < c.m; ++i)
for (int p = 0; p < c.k; ++p) ha_t[p * c.m + i] = ha[i * c.k + p];
printf("%dx%dx%d:\n", c.m, c.n, c.k);
printf(" NT K32:");
all &= run_gemm_case<RowMajor, ColMajor, 32, 3>(ha, hb_colmajor, c.m,
c.n, c.k, c.k, c.k);
printf(" NT K64:");
all &= run_gemm_case<RowMajor, ColMajor, 64, 2>(ha, hb_colmajor, c.m,
c.n, c.k, c.k, c.k);
printf(" NN K32:");
all &= run_gemm_case<RowMajor, RowMajor, 32, 3>(ha, hb_rowmajor, c.m,
c.n, c.k, c.k, c.n);
printf(" NN K64:");
all &= run_gemm_case<RowMajor, RowMajor, 64, 2>(ha, hb_rowmajor, c.m,
c.n, c.k, c.k, c.n);
printf(" TN K32:");
all &= run_gemm_case<ColMajor, ColMajor, 32, 3>(ha_t, hb_colmajor, c.m,
c.n, c.k, c.m, c.k);
printf(" TN K64:");
all &= run_gemm_case<ColMajor, ColMajor, 64, 2>(ha_t, hb_colmajor, c.m,
c.n, c.k, c.m, c.k);
printf(" TT K64:");
all &= run_gemm_case<ColMajor, RowMajor, 64, 2>(ha_t, hb_rowmajor, c.m,
c.n, c.k, c.m, c.n);
delete[] ha;
delete[] hb_rowmajor;
delete[] hb_colmajor;
delete[] ha_t;
}
return all;
}
} // namespace
int main() {
print_test_header();
bool ok = test_single_mma();
ok &= test_gemm();
printf(ok ? "All PASS\n" : "FAILURES\n");
return ok ? 0 : 1;
}
+7 -6
View File
@@ -9,9 +9,11 @@ services:
USER_GID: ${ASTRAI_GID:-1000}
user: "${ASTRAI_UID:-1000}:${ASTRAI_GID:-1000}"
ports:
- "8000:8000"
- "${SERVE_PORT:-8000}:${SERVE_CONTAINER_PORT:-8000}"
volumes:
- ./params:/app/params:ro
- ${SERVE_PARAM_DIR:-./params}:/app/params:ro
environment:
- CUDA_VISIBLE_DEVICES
command: python -m scripts.tools.server --port 8000 --device cuda
deploy:
resources:
@@ -39,9 +41,9 @@ services:
USER_GID: ${ASTRAI_GID:-1000}
user: "${ASTRAI_UID:-1000}:${ASTRAI_GID:-1000}"
ports:
- "8000:8000"
- "${SERVE_PORT:-8000}:${SERVE_CONTAINER_PORT:-8000}"
volumes:
- ./params:/app/params:ro
- ${SERVE_PARAM_DIR:-./params}:/app/params:ro
command: python -m scripts.tools.server --port 8000 --device cpu
healthcheck:
test: ["CMD", "curl", "-f", "http://localhost:8000/health"]
@@ -72,9 +74,8 @@ services:
- BASE_MODEL=${BASE_MODEL:-/models/base}
- CHECKPOINT_ROOT=/checkpoints
- TRAIN_GPU_COUNT=${TRAIN_GPU_COUNT:-all}
- TRAIN_PARALLEL_MODE=${TRAIN_PARALLEL_MODE:-auto}
- CUDA_VISIBLE_DEVICES
- NCCL_P2P_DISABLE
- NCCL_NET_GDR_LEVEL
entrypoint: ["bash", "/app/scripts/docker/train-entrypoint.sh"]
ipc: ${TRAIN_IPC_MODE:-host}
stop_grace_period: ${TRAIN_STOP_GRACE_PERIOD:-10m}
+6 -3
View File
@@ -57,7 +57,7 @@ AstrAI 是一个覆盖模型构建、训练、评测与部署的端到端 Transf
| **数据** | 声明式 JSON 预处理、可配置掩码与样本打包、二进制/JSONL 存储和流式数据集 |
| **推理** | 连续批处理、分页 KV Cache、Radix 前缀缓存、流式生成,以及 Torch/CUDA/FlashAttention 后端 |
| **服务** | 基于 FastAPI 的 OpenAI 与 Anthropic 聊天补全协议,支持 SSE 流式输出和工具调用 |
| **评测** | Perplexity、MMLU、HumanEval、IFEval、IFDROUGE 评测工具 |
| **评测** | Perplexity、MMLU、HumanEval、IFEval、IFDROUGE 和权重分析评测工具 |
| **扩展** | 基于工厂与注册表扩展模型、数据集、训练策略、回调、内核和协议组件 |
### 快速上手
@@ -71,8 +71,9 @@ AstrAI 需要 Python 3.12+,并精确固定 PyTorch 版本为 `2.11.0`。训练
```bash
git clone https://github.com/ViperEkura/AstrAI.git
cd AstrAI
pip install -e . # 纯 PyTorch(不含 CUDA 内核
# CSRC_KERNELS=true pip install -e . --no-build-isolation # 可选:融合 CUDA 内核加速
pip install -e . # 检测到 nvcc + CUDA 时自动构建内核
# CSRC_KERNELS=false pip install -e . # 跳过内核(纯 PyTorch
# CSRC_KERNELS=true pip install -e . --no-build-isolation # 强制构建融合 CUDA 内核
# pip install -e ".[dev]" # 可选:开发依赖(pytest, ruff
```
@@ -242,6 +243,8 @@ SSE 流式格式、错误码和统计端点详见[推理文档](guides/inference
| [数据流程](./developer/dataflow.md) | 数据管道、存储后端与数据集架构 |
| [内部实现](./developer/internals.md) | 训练原理:损失公式、回调生命周期、KV Cache |
| [CUDA 内核](./developer/cuda_kernels.md) | 自定义 CUDA 注意力内核与基准测试 |
| [Docker 服务部署](./developer/docker-serving.md) | YAML 驱动的容器化服务(`serve.yaml``serve.sh` |
| [Docker 训练部署](./developer/docker-training.md) | YAML 驱动的容器化训练(`train.yaml``train.sh` |
### 贡献
+4 -2
View File
@@ -1437,7 +1437,8 @@ classDiagram
| **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategyGRPOStrategy, StrategyFactory, BaseSchedulerWSDScheduler, SchedulerFactory, TrainCallback(Protocol)MetricCallback, CallbackFactory, RawRollout, RolloutResult, BaseRewardModel, RolloutGenerator, RolloutRunner | Training workflow |
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, InferenceWorkspace, PagePool, KVStorage, ReqToTokenPool, KVCache, Allocator, RadixCache, Task, TaskManager, TaskStatus, StreamDecoder, GenerateResult, BaseSamplingStrategySamplingPipeline, FrequencyPenaltyStrategy, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, StopInfo, ChatMessage, FunctionDef, ToolDef, ChatCompletionRequest, AnthropicMessage, MessagesRequest, BaseToolParser, ToolParserFactory, SimpleJsonToolParser | Inference service |
| **astrai.extension** | `backend` policy package, `ops` kernel-wrapper package, AttentionBackend, TorchNativeBackend, CudaBackend, FlashAttnBackend, attention, attn_backend, ATTN_BACKEND, apply_rotary_emb, is_available | Stable API over attention/rotary execution policy and optional CUDA kernels |
| **astrai.extension** | `backend` policy package, `ops` kernel-wrapper package, `fp8.py` FP8 strategy layer, AttentionBackend, TorchNativeBackend, CudaBackend, FlashAttnBackend, attention, attn_backend, ATTN_BACKEND, apply_rotary_emb, is_available | Stable API over attention/rotary/FP8 execution policy and optional CUDA kernels |
| **astrai.optim** | OptimizerFactory, MuonAdamW, NoraNadamW, ManoAdamW, composite_step/composite_zero_grad/composite_state_dict, partition_optimizer_parameters | Built-in optimizers (`muon_adamw` / `nora_nadamw` / `mano_adamw`) with shared composite-optimizer helpers |
| **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank/get_world_size/get_current_device, only_on_rank, LaunchStrategy, TorchrunStrategy, LocalStrategy, BaseExecutor, ExecutorFactory, NoneExecutor, DDPExecutor, FSDPExecutor, GradientState, AccumOptimizer, AccumScheduler | Distributed parallel & gradient accumulation |
| **astrai.factory** | BaseFactory | Component registration |
| **astrai.protocols** | OptimizerProtocol, SchedulerProtocol | Structural subtyping for optimizer/scheduler wrappers |
@@ -1461,6 +1462,7 @@ classDiagram
| **Storage** | `Store`, `MmapStore`, `JsonlStore` | Format-agnostic data access with multi-segment support |
| **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching |
| **Model Registry** | `ModelFactory`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading |
| **Optimizer Routing** | `OptimizerFactory`, `MuonAdamW`, `NoraNadamW`, `ManoAdamW` | Route parameter groups (matrices vs. embeddings/heads/norms) through different optimizers |
## Core Relationships
@@ -1476,4 +1478,4 @@ classDiagram
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
> Document Update Time: 2026-08-16
> Document Update Time: 2026-08-22
+87 -37
View File
@@ -1,23 +1,24 @@
# CUDA Kernels
AstrAI includes optional custom CUDA kernels for attention and rotary embedding. These are built when `nvcc` is available and CUDA is detected, and are dispatched via the `CudaBackend` attention backend or auto-dispatched for rotary.
AstrAI includes optional custom CUDA kernels for attention, rotary embedding, and FP8 GEMM. These are built when `nvcc` is available and CUDA is detected, and are dispatched via the `CudaBackend` attention backend, auto-dispatched for rotary, or invoked through the FP8 linear primitives.
## Overview
| Kernel | File | Description |
|--------|------|-------------|
| `attn_decode` | `attn_decode.cu` | GQA decode attention (split-KV) |
| `attn_prefill` | `attn_prefill.cu` | GQA prefill attention (split-Q) |
| `attn_paged_decode` | `attn_paged_decode.cu` | Paged KV cache decode attention |
| `attn_paged_prefill` | `attn_paged_prefill.cu` | Paged KV cache prefill attention (ragged batch) |
| `rotary_emb` | `rotary_emb.cu` | Fused rotary embedding (cos/sin lookup + rotation) |
| `attn_decode` | `attention/decode.cu` | GQA decode attention (split-KV) |
| `attn_prefill` | `attention/prefill.cu` | GQA prefill attention (split-Q) |
| `attn_paged_decode` | `attention/paged_decode.cu` | Paged KV cache decode attention |
| `attn_paged_prefill` | `attention/paged_prefill.cu` | Paged KV cache prefill attention (ragged batch) |
| `rotary_emb` | `rotary/rotary_emb.cu` | Fused rotary embedding (cos/sin lookup + rotation) |
| `fp8_ops` | `fp8/ops.cu` | FP8 quantization + tensor-core GEMM (sm_89+) |
Additionally, optimized `.cuh` variants with tensor-core MMA (Matrix Multiply-Accumulate) exist:
| Variant | File | Optimization |
|---------|------|--------------|
| Split-KV MMA decode | `attn_decode_split_kv_mma.cuh` | Split KV across warps + MMA (sm_80+) |
| Split-Q MMA prefill | `attn_prefill_split_q_mma.cuh` | Split Q across warps + MMA (sm_80+) |
| Split-KV MMA decode | `attention/decode_split_kv_mma.cuh` | Split KV across warps + MMA (sm_80+) |
| Split-Q MMA prefill | `attention/prefill_split_q_mma.cuh` | Split Q across warps + MMA (sm_80+) |
> The paged and non-paged paths share one kernel body. Prefill is templated on
> an independent Q schedule (`DenseQSchedule` / `PackedQSchedule`) and KV
@@ -26,7 +27,7 @@ Additionally, optimized `.cuh` variants with tensor-core MMA (Matrix Multiply-Ac
### Rotary Embedding Kernel
The `rotary_emb` kernel (`csrc/kernels/rotary_emb.cu`) fuses cos/sin lookup and rotation into a single kernel:
The `rotary_emb` kernel (`csrc/kernels/rotary/rotary_emb.cu`) fuses cos/sin lookup and rotation into a single kernel:
- One thread per (head, dim-pair), vectorized `__nv_bfloat162` load/store
- f32 cos/sin input, bf16 compute and output
@@ -36,6 +37,30 @@ The `rotary_emb` kernel (`csrc/kernels/rotary_emb.cu`) fuses cos/sin lookup and
Standalone benchmark vs torch complex-multiply (48 calls = 24 layers × q+k): 6-9x faster, max diff 0 (decode) to 3e-2 (large prefill, bf16).
### FP8 GEMM / Linear Kernel
The `fp8_ops` family (`csrc/kernels/fp8/`) accelerates bf16 linear layers by
quantizing to FP8 and running tensor-core GEMMs (**requires sm_89+**; fp8
`mma.sync.m16n8k32` only exists on Ada/Hopper). It follows the same three-layer
style as attention, but split into **three** files:
| File | Role |
|------|------|
| `fp8/common.h` | `FP8Format` enum (E4M3/E5M2), `Fp8GemmTraits<Fmt, BlockM, BlockN, K, Stages>`, `FP8Params` POD — no torch |
| `fp8/gemm.cuh` | pure-CUDA device code: `fp8_quantize_kernel` (BF16→FP8 + amax), `fp8_gemm_kernel` (pre-quantized GEMM, 128×64 CTA / 64×16 warp / 3-stage cp.async) — no torch |
| `fp8/ops.cu` | binding only: `check_fp8_device` (sm_89+), param packing, launch dispatch, pybind → module `fp8_ops` |
Scale semantics follow `torch._scaled_mm` (quantization step size: divide by
`scale`; the kernel computes the reciprocal internally — the interface never
takes `*_inv`). `amax` is always returned in the original bf16 domain.
Python layer (two levels): `astrai/extension/ops/fp8.py` provides stateless
primitives (`quantize_bf16` / `mm_fp8` / `linear_forward_fp8` /
`linear_backward_fp8`) via `torch.library.custom_op`, and
`astrai/extension/fp8.py` is the strategy layer (`fp8_autocast`, delayed /
dynamic scaling recipes, `fp8_linear_forward/backward` wiring `aten::linear`
on CUDA). See the FP8 section in `AGENTS.md` for full detail.
## Build System
### Auto-detection
@@ -66,10 +91,17 @@ cmake --build build/cmake -j 16
### Architecture flags
`setup.py` passes the GPU compute capability to CMake via `ASTRAI_CUDA_ARCH` (default `89`, i.e. sm_89 / L20):
`setup.py` passes the GPU compute capability to CMake via `ASTRAI_CUDA_ARCH`. When
unset, `setup.py` auto-detects the real GPU capability through
`torch.cuda.get_device_capability()`; the CMake fallback default is `80` (sm_80):
- **sm_80+** (Ampere and later): enables tensor-core MMA path (`mma.sync.m16n8k16.bf16`)
- **Below sm_80**: adds `-DASTRAI_NO_MMA` to disable the MMA path at compile time
- **sm_80+** (Ampere and later): enables the tensor-core MMA path
(`mma.sync.m16n8k16.bf16` for bf16 attention, `mma.sync.m16n8k32` for FP8).
- **sm_89+**: required for the FP8 family (`fp8_ops`) — FP8 tensor-core
instructions only exist on Ada/Hopper and newer.
- **`-DASTRAI_NO_MMA`** is a manual escape hatch only — the build never defines
it automatically. To disable the MMA path, add it to `NVCC_FLAGS` yourself;
all supported build targets are sm_80+.
### Build configuration
@@ -80,7 +112,7 @@ NVCC_FLAGS = -O3 --expt-relaxed-constexpr --use_fast_math
--ptxas-options=-O3,-v --extra-device-vectorization --threads=16
```
Each kernel in `astrai/extension/lib` is compiled as an independent pybind11 module (one `.so` per kernel, named `<kernel>.cpython-*-x86_64-linux-gnu.so`). CMake builds all five kernel targets in parallel via `cmake --build -j N`.
Each kernel in `astrai/extension/lib` is compiled as an independent pybind11 module (one `.so` per kernel, named `<kernel>.cpython-*-x86_64-linux-gnu.so`). CMake builds all six kernel targets in parallel via `cmake --build -j N`. The target list is the **single source of truth**: `KERNEL_NAMES` and the parallel `KERNEL_SRCS` list in `csrc/CMakeLists.txt`; `astrai/extension/loader.py` auto-discovers the compiled `.so` files.
## Python Extension Architecture
@@ -93,7 +125,9 @@ astrai/extension/
├── loader.py # Optional compiled-module discovery and loading
├── ops/
│ ├── attention.py # Stateless attention kernel wrappers
── rotary.py # Stateless rotary kernel wrapper
── rotary.py # Stateless rotary kernel wrapper
│ └── fp8.py # Stateless FP8 primitives (custom_op)
├── fp8.py # FP8 strategy layer (fp8_autocast, recipes)
└── backend/
├── attention.py # Backend selection, KV cache I/O, and fallback
└── rotary.py # Per-call CUDA/torch rotary dispatch
@@ -212,9 +246,14 @@ with attn_backend(ATTN_BACKEND.CUDA):
The `attention(...)` policy entry point falls back to `FlashAttnBackend` (when
flash-attn is installed and supports the call) or `TorchNativeBackend` when the
automatically selected CUDA backend cannot handle an input. An explicit
`ASTR_BACKEND` or `attn_backend(...)` selection is strict and raises instead of
silently switching implementations.
automatically selected CUDA backend cannot handle an input. Resolution
precedence is: explicit `attn_backend(...)` context > `ASTR_BACKEND` env >
default. An explicit `attn_backend(...)` selection is strict and raises instead
of silently switching implementations; the env override (and the implicit
default) fall back to the first compatible backend when incapable. Training
calls (`fwd=None`, no KV cache) resolve by capability: the CUDA cache kernels
cannot run without a cache, so they fall back to flash (mask-free/causal calls
only) and finally to torch SDPA.
### Rotary Backend
@@ -293,6 +332,7 @@ nvcc -I csrc -arch=sm_89 -O3 --use_fast_math \
Test files:
- `attn_test.cu` — decode + prefill kernels (correctness tables + benchmarks)
- `attn_paged_test.cu` — paged decode/prefill kernels
- `fp8_mma_test.cu` — BF16→FP8→BF16 MMA demo (sm_89)
## Benchmarks
@@ -315,29 +355,39 @@ nvcc -I csrc -arch=sm_89 -O3 --use_fast_math \
```
csrc/
├── CMakeLists.txt # CMake build: 5 kernel targets, torch/pybind11 linking
├── CMakeLists.txt # CMake build: kernel registry (KERNEL_NAMES / KERNEL_SRCS), torch/pybind11 linking
├── kernels/
│ ├── attn_common.h # Unified attention params (contig + paged modes)
│ ├── attn_decode.cu # Basic decode kernel (registered)
├── attn_prefill.cu # Basic prefill kernel (registered)
│ ├── attn_paged_decode.cu # Paged decode kernel (registered)
│ ├── attn_paged_prefill.cu # Paged prefill kernel (registered)
│ ├── rotary_emb.cu # Fused rotary embedding kernel (registered)
│ ├── attn_decode_split_kv.cuh # Split-KV variant (contig + paged via KVSource)
│ ├── attn_decode_split_kv_mma.cuh # Split-KV + MMA variant (contig + paged)
│ ├── attn_prefill_split_q.cuh # Split-Q variant (contig + paged via KVSource)
│ ├── attn_prefill_split_q_mma.cuh # Split-Q + MMA variant (contig + paged)
│ ├── attn_layout_policies.cuh # Q schedules and KVSource policies
│ ├── attn_dispatchers.cuh # Kernel dispatch macros + KV-templated launchers
│ ├── attn_entry_utils.cuh # Entry point helpers
│ ├── attn_mma_utils.cuh # MMA utilities
└── attn_warp_utils.cuh # Warp-level utilities
│ ├── common/ # cross-family pure-CUDA helpers (no torch)
│ ├── device.cuh # sm_at_least(), kMinSmForFp8* constants
│ └── mma.cuh # shared mma_sync<InT> + mma_shape<InT> (bf16 m16n8k16 / fp8 m16n8k32) + ldmatrix_x2/x4<T>
│ ├── attention/ # attention family (module names keep the attn_* prefix)
│ ├── common.h # AttentionParams POD, TensorLayout enum (BHLD/BLHD)
│ ├── warp_utils.cuh # warp reduction helpers
│ ├── layout_policies.cuh # KV addressing policies: DenseQSchedule/PackedQSchedule, ContigKV/PagedKV
│ ├── mma_utils.cuh # ldmatrix/pack helpers + online-softmax (bf16 mma via common/mma.cuh)
│ ├── entry_utils.cuh # torch binding helpers: DISPATCH_HEAD_DIM, pack_*_params
│ ├── dispatchers.cuh # pure-CUDA launchers: dispatch_decode/prefill (+paged), split-K math
│ ├── decode_split_kv.cuh # decode kernel, scalar (split-KV)
│ ├── decode_split_kv_mma.cuh # decode kernel, MMA + split-K
│ ├── prefill_split_q.cuh # prefill kernel, scalar (split-Q)
│ ├── prefill_split_q_mma.cuh # prefill kernel, MMA (split-Q, packed/ragged Q schedule)
│ ├── decode.cu # → module attn_decode
│ │ ├── prefill.cu # → module attn_prefill
│ │ ├── paged_decode.cu # → module attn_paged_decode
│ │ └── paged_prefill.cu # → module attn_paged_prefill
│ ├── rotary/
│ │ └── rotary_emb.cu # rotary embedding (kernel + binding in one file) → module rotary_emb
│ └── fp8/ # FP8 family (module name fp8_ops)
│ ├── common.h # FP8Format enum, Fp8GemmTraits, FP8Params POD (no torch)
│ ├── gemm.cuh # FP8 device code: quantize + pre-quantized GEMM kernels (no torch)
│ └── mm.cu # binding only: validation, param packing, launch dispatch, pybind
└── tests/
├── test_utils.cuh # Shared test utilities
├── attn_test.cu # Decode + prefill kernels
── attn_paged_test.cu # Paged decode/prefill kernels
├── test_utils.cuh # Shared test utilities (now_ms, f2bf, bf2f, randf)
├── attn_test.cu # Decode + prefill kernels
── attn_paged_test.cu # Paged decode/prefill kernels
└── fp8_mma_test.cu # BF16→FP8→BF16 MMA demo
```
Compiled `.so` files are placed in `astrai/extension/lib/`, separate from Python source files.
> Document Update Time: 2026-08-16
> Document Update Time: 2026-08-22
+110
View File
@@ -0,0 +1,110 @@
# Containerized Serving Deployment
AstrAI uses one serving YAML as the declaration for both host-side container
runtime settings and in-container server settings. `scripts/serve.sh` wraps the
Compose commands so preflight validation and container lifecycle stay
consistent with the trainer.
## Architecture
```text
serve.yaml
├── runtime parsed on the host before Docker starts
└── server parsed by server.py inside the container
scripts/serve.sh preflight, Compose wrapper, lifecycle
└── docker-compose.yml GPU passthrough, mounts, image, port mapping
└── server.py --config /run/astrai/serve.yaml
```
`scripts/tools/serve_runtime.py` reads `runtime:` plus the two container-side
values Compose needs (`server.port` for the port mapping, `server.device` for
the preflight GPU check). `scripts/tools/server.py --config` reads `server:`.
Explicit CLI arguments to `server.py` override `server:` YAML values.
## Runtime Schema
```yaml
runtime:
job_name: serve
port: 8000
paths:
param: ./params
gpu:
enabled: true # false → cpu profile (server-cpu service)
devices: all # all | [0]
container:
cuda_tag: cu128
# environment:
# TOKENIZERS_PARALLELISM: "false"
server:
host: 0.0.0.0
port: 8000
device: cuda # cuda | cpu
dtype: bfloat16 # bfloat16 | float16 | float32
max_batch_size: 16
max_seq_len: null # falls back to model config
```
- Relative paths resolve from the YAML file's directory, not the current shell.
- `runtime.port` is the host publish port; `server.port` is the port the
container listens on. The Compose mapping is
`${SERVE_PORT}:${SERVE_CONTAINER_PORT}`.
- `runtime.gpu.enabled: true` (default) selects the `server` service with an
NVIDIA device reservation; `false` selects `server-cpu` (no GPU passthrough).
When disabled, `server.device` must be `cpu`.
- `runtime.gpu.devices` is either `all` or a single-device list such as `[0]`;
the list becomes `CUDA_VISIBLE_DEVICES`. Compose passes `count: 1`.
- `environment` values are explicitly passed to the serving container. Keep
host-specific settings here; they are not universal defaults.
- `server.device` must agree with `runtime.gpu.enabled`; `preflight` enforces it.
## Fixed Container Paths
| Runtime path | Container path | Access |
|---|---|---|
| `runtime.paths.param` | `/app/params` | read-only |
| the selected YAML | `/run/astrai/serve.yaml` | read-only |
`server.param_path` is optional: the server default is
`project_root/params`, which is exactly `/app/params` inside the container
(the working directory is `/app`). Set it explicitly only when serving from a
different location; in Docker it must be a container path.
## Operations
The config argument defaults to `./serve.yaml`:
```bash
bash scripts/serve.sh init [CONFIG]
bash scripts/serve.sh preflight [CONFIG]
bash scripts/serve.sh up [CONFIG]
bash scripts/serve.sh run [CONFIG]
bash scripts/serve.sh down [CONFIG]
bash scripts/serve.sh restart [CONFIG]
bash scripts/serve.sh logs [CONFIG]
bash scripts/serve.sh status [CONFIG]
```
`preflight` validates Docker, the model directory
(`config.json` + `model.safetensors`), GPU/device consistency, and the
rendered Compose configuration. `up` starts the container detached and
rebuilds the image when the code changed (`--build`); `run` keeps it in the
foreground. The wrapper manages a fixed container name
(`astrai-server` or `astrai-server-<job_name>`); the plain
`docker compose up -d` / `docker compose --profile cpu up -d` path keeps
working with defaults (port 8000, `./params`).
## Hard Rules
1. Keep Docker settings in `runtime` and server settings in `server`.
2. Filter GPUs once: the `server` service reserves one device; a `devices`
list becomes `CUDA_VISIBLE_DEVICES`.
3. `runtime.gpu.enabled: false` requires `server.device: cpu`.
4. In Docker, `server.port` must match the published container port (default
`8000`); change `runtime.port` to publish on a different host port.
5. The image user is built with the host UID/GID so the mounted model
directory stays readable.
> Document Update Time: 2026-08-22
+108 -38
View File
@@ -1,57 +1,127 @@
# Containerized Training Deployment
Rules for running AstrAI distributed training in containers, distilled from real deployment failures. Read before touching `Dockerfile`, `docker-compose.yml`, `scripts/train.sh`, `train-entrypoint.sh`. AGENTS.md mirrors this locally; this file is the committed version.
AstrAI uses one training YAML as the declaration for both host-side container
runtime settings and in-container training settings. Do not invoke the trainer
with raw `docker compose up`; use `scripts/train.sh` so preflight validation,
checkpoint recovery, and graceful shutdown remain active.
## Architecture
```
scripts/train.sh host-side CLI: env loading, preflight, compose wrapper, lifecycle
── docker-compose.yml GPU passthrough, mounts, in-container env vars, entrypoint
└── train-entrypoint.sh GPU-count resolution, parallel-mode selection, auto-resume
```text
train.yaml
── runtime parsed on the host before Docker starts
└── model/data/... parsed by train.py inside the container
scripts/train.sh preflight, Compose wrapper, lifecycle, timer
└── docker-compose.yml GPU passthrough, mounts, image, container limits
└── scripts/docker/train-entrypoint.sh process count, parallel mode, auto-resume
└── train.py --config /run/astrai/train.yaml
```
| Layer | Responsible for | NOT responsible for |
|-------|-----------------|---------------------|
| `train.sh` | host paths, `.env.train`, preflight, lifecycle | training args, GPU selection, parallel mode |
| compose | GPU passthrough, mounts, in-container env (NCCL) | training args (beyond `TRAIN_*` forwarding) |
| entrypoint | `--ckpt_dir/--nprocs/--parallel_mode/--param_path`, resume | hyperparameters (YAML/CLI) |
| `train.yaml` | hyperparameters (`_merge_yaml_into_kwargs`, CLI wins) | container paths, process count |
The two parsers deliberately own different sections. `scripts/tools/train_runtime.py`
reads only `runtime`; `scripts/tools/train.py` reads only
`model/data/parallel/training/ckpt/log`. Explicit trainer arguments after `--`
override training YAML values.
## Path Conventions
## Runtime Schema
| Host var | Container | Perm | Purpose |
|---|---|---|---|
| `TRAIN_DATA_DIR` | `/data` | ro | dataset (`data_root_path` must be `/data`) |
| `TRAIN_MODEL_DIR` | `/models/base` | ro | base model (`config.json` + `model.safetensors`) |
| `TRAIN_CHECKPOINT_DIR` | `/checkpoints` | rw | checkpoint root, per-`TRAIN_JOB_NAME` subdirs |
| `TRAIN_CONFIG_FILE` | `/run/astrai/train.yaml` | ro | training YAML (mounted only on `start`) |
| code | `/app` | image | **not a mount**; rebuild image for code changes |
```yaml
runtime:
job_name: astrai-train
paths:
data: ./data
model: ./params
checkpoints: ./checkpoints
gpu:
devices: all
parallel_mode: auto # one GPU: none; multiple GPUs: ddp
container:
cuda_tag: cu128
ipc: host
stop_grace_period: 10m
stop_timeout_seconds: 600
checkpoint_keep_last: 5
# max_duration_hours: 12
# Add host-specific workarounds only when required:
# environment:
# NCCL_P2P_DISABLE: "1"
# NCCL_NET_GDR_LEVEL: "0"
```
## Hard Rules
- Relative paths resolve from the YAML file's directory, not the current shell.
- `devices` is either `all` or a non-empty physical GPU index list. Compose
passes all GPUs once; `CUDA_VISIBLE_DEVICES` performs the only filtering.
- The process count is derived from `devices`. With `all`, the entrypoint uses
`torch.cuda.device_count()` after Docker starts.
- `parallel_mode: auto` selects `none` for one GPU and `ddp` for multiple GPUs.
Use `fsdp` explicitly when model sharding is required.
- To select specific physical GPUs, replace `all` with a list such as
`devices: [0, 1]`.
- `environment` values are explicitly passed to the training container. Keep
host-specific NCCL workarounds here; they are not universal defaults.
- `max_duration_hours` starts a detached host timer that calls the same graceful
`stop` command. A manual stop cancels the timer.
1. **Filter GPUs once**: compose passes the full physical set (`count: all`); `CUDA_VISIBLE_DEVICES` filters inside by physical index. Never `count: N` + physical indices (double filter leaves 1 card → `device_id out of range`).
2. **In-container UID = host UID**: Dockerfile builds the user via `USER_UID/USER_GID` args; `train.sh` injects `ASTRAI_UID/GID` (bash `UID` is readonly). compose `user:` alone does not create the /etc/passwd entry — torch's `getpass.getuser()` then dies with `uid not found`.
3. **In-container env vars are explicit**: `.env.train` (`--env-file`) is only compose's interpolation dictionary — never reaches the container. A var arrives only via a value-less `environment` entry (`- VAR`, read from the calling process env).
4. **NCCL hang workaround** (this host): `NCCL_P2P_DISABLE=1` + `NCCL_NET_GDR_LEVEL=0` must be in-container.
5. **Checkpoint complete =** `meta.json + config.json + model.safetensors + optimizer.pt + scheduler.pt`; `start` auto-resumes the latest complete one.
6. **tqdm is silent without a TTY**: add `disable=False` in `astrai/trainer/train_callback.py`; `metric.jsonl` (per step) works as progress evidence regardless.
## Fixed Container Paths
| Runtime path | Container path | Access |
|---|---|---|
| `runtime.paths.data` | `/data` | read-only |
| `runtime.paths.model` | `/models/base` | read-only |
| `runtime.paths.checkpoints` | `/checkpoints` | read-write |
| the selected YAML | `/run/astrai/train.yaml` | read-only |
Training configuration must therefore use `data_root_path: /data`. The source
code is baked into `/app`; `start` uses `--build`, so code changes rebuild the
image when necessary.
## Operations
The config argument defaults to `./train.yaml`:
```bash
bash scripts/train.sh init # first run: dirs + .env.train (edit per machine)
bash scripts/train.sh preflight # validate Docker/paths/GPU/model/YAML/compose
bash scripts/train.sh start # build + start in background (auto-resume)
bash scripts/train.sh start --foreground -- --dry-run # print plan only
bash scripts/train.sh logs | status | stop | restart
bash scripts/train.sh clean --keep 5 # prune old checkpoints (--force to delete)
bash scripts/train.sh init [CONFIG]
bash scripts/train.sh preflight [CONFIG]
bash scripts/train.sh start [CONFIG]
bash scripts/train.sh start [CONFIG] --foreground -- --dry-run
bash scripts/train.sh logs [CONFIG]
bash scripts/train.sh status [CONFIG]
bash scripts/train.sh stop [CONFIG]
bash scripts/train.sh restart [CONFIG]
bash scripts/train.sh clean [CONFIG] --keep 5
bash scripts/train.sh clean [CONFIG] --keep 5 --force
```
## Files
`init` creates the declared runtime directories but does not generate or mutate
the YAML. `preflight` validates Docker, paths, base model files, checkpoint
writability, GPU configuration, and rendered Compose configuration.
- `docker-compose.yml` — trainer service: `count: all`, `ASTRAI_UID/GID` build args + `user:`, env whitelist, mounts
- `Dockerfile` — production stage builds user from `USER_UID/USER_GID`; `ENV HOME=/home/astrai`; `USER astrai`
- `scripts/train.sh``load_env` filters `UID=` lines (readonly var); `compose()` injects `ASTRAI_UID/GID`
- `scripts/docker/train-entrypoint.sh` — GPU-count resolution, parallel mode, resume
- `.env.train`, `train.yaml` — host-specific; templates from `scripts/train.sh init`; scientific-notation floats (`2e-5`) parse correctly since train.py uses the YAML 1.2 float schema
## Checkpoint Recovery
Checkpoints are stored below
`runtime.paths.checkpoints/<job_name>/epoch_<N>_step_<N>`. A checkpoint is
complete only when it contains:
```text
meta.json
config.json
model.safetensors
optimizer.pt
scheduler.pt
```
`start` resumes the latest complete checkpoint and ignores partial writes. If no
complete checkpoint exists, `/models/base/config.json` and
`/models/base/model.safetensors` are required. `stop` sends `SIGTERM`; the
trainer finishes at a batch boundary and saves an emergency checkpoint before
the Docker timeout expires.
## Hard Rules
1. Keep Docker settings in `runtime` and trainer settings in the remaining YAML sections.
2. Filter GPUs once: Compose passes `count: all`; `devices` becomes `CUDA_VISIBLE_DEVICES`.
3. Do not force DDP for a model that requires FSDP; declare the mode explicitly.
4. Do not use `kill -9` for routine shutdown; use `scripts/train.sh stop CONFIG`.
5. The image user is built with the host UID/GID so mounted checkpoints retain usable ownership.
> Document Update Time: 2026-08-22
+1 -1
View File
@@ -188,7 +188,7 @@ Attention computation is decoupled from the model via `AttentionBackend` ABC (`a
- **`FlashAttnBackend`**: optional flash-attn dispatch with `flash_attn_with_kvcache` fast path for contiguous cache; falls back to KV gather + `flash_attn_func`.
- **`TorchNativeBackend`** (always-available fallback): writes K/V to cache, gathers via `req_to_token` indirect indexing, calls `F.scaled_dot_product_attention`.
- The `attention(...)` entry point uses cuda > flash > torch priority and chooses another compatible backend when an automatically selected backend cannot handle a call.
- `ASTR_BACKEND=cuda|torch_native|flash` and `attn_backend(...)` are explicit selections; incompatible calls raise instead of silently changing backend.
- Resolution precedence is: explicit `attn_backend(...)` context > `ASTR_BACKEND` env > default. An explicit `attn_backend(...)` selection is strict (incompatible calls raise); `ASTR_BACKEND` is a default-level override that falls back to a compatible backend when incapable. Training calls (`fwd=None`, no KV cache) resolve by capability: the CUDA cache kernels cannot run without a cache, so they fall back to flash (mask-free/causal calls only) and finally to torch SDPA.
Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/backend/rotary.py`, which auto-dispatches to the fused CUDA kernel (`rotary_emb.cu`) during inference or torch complex multiply during training (for autograd compatibility). Both attention backends share the same rotary dispatch.
+21 -4
View File
@@ -26,17 +26,24 @@ This guide walks you through installing AstrAI, downloading a model, running inf
git clone https://github.com/ViperEkura/AstrAI.git
cd AstrAI
# Basic install (pure PyTorch, no custom CUDA kernels)
# Kernels auto-build when nvcc + CUDA are detected; skip with CSRC_KERNELS=false
pip install -e .
# With CUDA kernels (optional, for fused attention and rotary embedding)
# Force the CUDA kernel build (fused attention, rotary embedding, FP8 GEMM)
# CSRC_KERNELS=true pip install -e . --no-build-isolation
# With dev dependencies (pytest, ruff)
# pip install -e ".[dev]"
```
> **CUDA kernels** are opt-in at build time (`CSRC_KERNELS=true`). Once built, `CudaBackend` is the default attention backend on GPU (cuda > flash > torch priority). Override via `ASTR_BACKEND` env var or `attn_backend()` context manager. Fused rotary embedding kernel is auto-dispatched when available. Skip for CPU-only usage.
> **CUDA kernels** build automatically when `nvcc` is on `PATH` and
> `torch.cuda.is_available()` returns `True`; set `CSRC_KERNELS=false` to skip
> them, or `CSRC_KERNELS=true` to force them (required when building in an
> isolated environment with `--no-build-isolation`). Once built, `CudaBackend`
> is the default attention backend on GPU (cuda > flash > torch priority).
> Override via `ASTR_BACKEND` env var or `attn_backend()` context manager.
> Fused rotary embedding kernel is auto-dispatched when available. Skip for
> CPU-only usage.
## 2. Download Model Weights
@@ -58,6 +65,14 @@ The model directory contains:
- `model.safetensors` — model weights
- `tokenizer.json` + `tokenizer_config.json` — tokenizer files (including chat template)
External HuggingFace checkpoints of the LLaMA layout (e.g. `meta-llama/...`,
`mistralai/...`, `Qwen/Qwen2-...`) can be loaded directly: `AutoModel.from_pretrained`
auto-detects HF `model_type` / key names (`input_layernorm`, `gate_proj`, MoE
`experts.<j>` ...) and converts config and weights in place. Dense and MoE
(Mixtral / DeepSeek-V3 layout) FFNs are supported; MLA attention
(DeepSeek-V2/V3) and biased projections (`attention_bias`) are not. Pass
`weights_format="astrai"` to skip conversion, or `"hf"` to force it.
## 3. Run Inference
### Interactive Chat (Simplest)
@@ -250,5 +265,7 @@ docker compose up -d
| Multi-GPU DDP / FSDP | [Distributed Guide](guides/distributed.md) |
| System architecture | [Architecture](developer/architecture.md) |
| Data pipeline internals | [Data Flow](developer/dataflow.md) |
| YAML-driven containerized serving | [Docker Serving](developer/docker-serving.md) |
| YAML-driven containerized training | [Docker Training](developer/docker-training.md) |
> Document Update Time: 2026-07-31
> Document Update Time: 2026-08-22
+44 -8
View File
@@ -54,7 +54,12 @@ KVCache
├── out_cache_loc [batch, seq_len] — write indices for this forward
├── max_len int — max(seq_lens), avoids GPU sync in decode
├── kv_indptr [batch + 1] int32 — prefix sum of seq_lens, precomputed once per step
── qo_indptr [batch + 1] int32 — prefix sum of per-request q_lens (prefill), precomputed once per step
── qo_indptr [batch + 1] int32 — prefix sum of per-request q_lens (prefill), precomputed once per step
├── q_tile_to_batch [num_q_tiles] int32 — prefill: Q tile → request (precomputed once per step)
├── q_tile_to_index [num_q_tiles] int32 — prefill: Q tile → request-local tile index
├── decode_o_part [batch, n_heads, head_dim] — decode split-K partial output buffer
├── decode_ml_part [batch, n_heads] — decode split-K partial max/logsum buffer
└── decode_out [batch, n_heads, head_dim] — decode output accumulator
```
Attention layers do raw buffer indexing: `k_buffer[layer_id, out_cache_loc] = k` to write, `k_buffer[layer_id, indices]` to gather.
@@ -80,7 +85,9 @@ AttentionBackend (ABC)
Default priority is cuda > flash > torch. Automatic selection may choose a
compatible fallback for a particular call. Set
`ASTR_BACKEND=cuda|torch_native|flash` to require one backend process-wide.
`ASTR_BACKEND=cuda|torch_native|flash` to override the default process-wide;
an explicit `attn_backend(...)` context still takes precedence over the env
override.
Select via context manager (mirrors `torch.nn.attention.sdpa_kernel`):
@@ -94,7 +101,7 @@ with attn_backend(ATTN_BACKEND.CUDA):
Environment and context selections are strict: if the selected backend cannot
handle the call, inference raises an error rather than silently switching.
`CudaBackend` decode path: writes K/V to cache, then calls `attn_paged_decode` with `page_size=1` — the `req_to_token` table serves directly as the page table, each token slot is a single-token "page". No explicit K/V gather needed.
`CudaBackend` decode path: writes K/V via `new_k`/`new_v` while calling `attn_paged_decode` — the `req_to_token` table serves directly as the page table (conceptually a single-token "page" per slot, i.e. `page_size=1`; the op itself takes no `page_size` argument). No explicit K/V gather needed.
`CudaBackend` prefill path: writes K/V, then calls `attn_paged_prefill` — a ragged-batch (paged) prefill kernel that reads K/V directly from the flat pool via `req_to_token`, addressing each request's `q_len`/`kv_len` through `qo_indptr` and `kv_indptr`. No explicit K/V gather needed.
@@ -171,6 +178,31 @@ InferenceEngine
`GenerateResult` uses `Condition` for non-streaming (`wait_completion()`) and `Event` for streaming (`wait()`). Stream callback is `cb(token)`.
## Launching the Server
`scripts/tools/server.py` accepts every option as a CLI flag or from a YAML
config file (`--config serve.yaml`); explicit CLI flags override YAML values.
The YAML `server:` section mirrors the flags:
```yaml
server:
host: 0.0.0.0
port: 8000
device: cuda
dtype: bfloat16
max_batch_size: 16
max_seq_len: null
```
```bash
python scripts/tools/server.py --config serve.yaml
python scripts/tools/server.py --config serve.yaml --port 9000 # CLI wins
```
In Docker, `scripts/serve.sh` drives the same YAML (a `runtime:` section
controls ports/GPU/mounts); see
[Docker Serving](../developer/docker-serving.md).
## HTTP API
```
@@ -228,14 +260,18 @@ The HTTP protocols and direct engine API have distinct request models and defaul
| `max_tokens` | Optional[int] | 2048 | Max generation length |
| `stream` | Optional[bool] | False | Stream output |
| `stop` | Optional[Union[str, List[str]]] | None | Stop sequences |
| `n` | Optional[int] | 1 | Number of choices requested |
| `presence_penalty` | Optional[float] | 0.0 | Presence penalty (-2.0 to 2.0) |
| `n` | Optional[int] | 1 | Accepted for API compatibility, **ignored** (always returns a single choice) |
| `presence_penalty` | Optional[float] | 0.0 | Accepted for API compatibility, **ignored** |
| `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 |
| `logit_bias` | Optional[Dict[int, float]] | None | Accepted for API compatibility, **ignored** |
| `user` | Optional[str] | None | Accepted for API compatibility, **ignored** |
| `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 |
> `n`, `presence_penalty`, `logit_bias`, and `user` are validated by the request
> model but ignored by the server (a warning is logged when a non-default value
> is supplied).
**Anthropic** (`MessagesRequest`):
| Param | Type | Default | Description |
@@ -346,4 +382,4 @@ async for token in engine.generate_async("Hello", ...): # -> AsyncGenerator[s
print(token)
```
> Document Update Time: 2026-08-16
> Document Update Time: 2026-08-22
+19 -2
View File
@@ -28,7 +28,7 @@
|-----------|-------------|---------|
| `--warmup_ratio` | Fraction of total steps used for LR warmup | 0.05 |
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
| `--max_grad_norm` | Maximum gradient norm for clipping; the current CLI requires a positive number | 1.0 |
| `--max_grad_norm` | Maximum gradient norm for clipping; `TrainConfig` validates it as positive (or `None`) | 1.0 |
### Optimizer
@@ -203,6 +203,7 @@ nohup python scripts/tools/train.py \
| Parameter | Type | Default | Description |
|-----------|------|---------|-------------|
| `--config`, `-c` | path | `None` | Serving YAML config. CLI flags override YAML values |
| `--host` | str | `0.0.0.0` | Host address |
| `--port` | int | `8000` | Port number |
| `--param_path` | path | `project_root/params` | Path to model parameters |
@@ -217,6 +218,22 @@ Usage:
python scripts/tools/server.py --param_path ./params --device cuda --dtype bfloat16
```
YAML config (a `server:` section; explicit CLI flags override YAML values):
```bash
python scripts/tools/server.py --config serve.yaml
```
```yaml
server:
host: 0.0.0.0
port: 8000
device: cuda
dtype: bfloat16
max_batch_size: 16
max_seq_len: null
```
`serve.yaml` also carries a `runtime:` section for the Docker wrapper; see
[Docker Serving](../developer/docker-serving.md).
See [Inference Guide](inference.md) for HTTP API documentation.
## Generate (`generate.py`)
@@ -264,4 +281,4 @@ See [Preprocessing Guide](preprocessing.md) for config file format and examples.
---
> Document Update Time: 2026-07-20
> Document Update Time: 2026-08-22
+4 -1
View File
@@ -31,12 +31,15 @@ classifiers = [
urls = { Homepage = "https://github.com/ViperEkura/AstrAI" }
[project.optional-dependencies]
dev = ["pytest==9.0.2", "ruff", "httpx2"]
dev = ["pytest==9.0.2", "ruff", "httpx"]
flash = ["flash-attn>=2.6"]
[tool.setuptools.packages.find]
where = ["."]
[tool.setuptools.package-data]
"astrai.extension.lib" = ["*.so"]
[tool.setuptools.dynamic]
version = { attr = "astrai.__version__" }
+17 -6
View File
@@ -10,12 +10,27 @@ CHECKPOINT_DIR="${CHECKPOINT_ROOT}/${TRAIN_JOB_NAME}"
BASE_MODEL="${BASE_MODEL:-/models/base}"
TRAIN_CONFIG="${TRAIN_CONFIG:-}"
TRAIN_GPU_COUNT="${TRAIN_GPU_COUNT:-all}"
TRAIN_PARALLEL_MODE="${TRAIN_PARALLEL_MODE:-auto}"
validate_job_name "${TRAIN_JOB_NAME}"
if [[ "${TRAIN_GPU_COUNT}" == "all" ]]; then
TRAIN_GPU_COUNT="$(python -c 'import torch; print(torch.cuda.device_count())')"
fi
[[ "${TRAIN_GPU_COUNT}" =~ ^[1-9][0-9]*$ ]] || die "No visible GPU found"
if [[ "${TRAIN_PARALLEL_MODE}" == "auto" ]]; then
if (( TRAIN_GPU_COUNT > 1 )); then
TRAIN_PARALLEL_MODE=ddp
else
TRAIN_PARALLEL_MODE=none
fi
fi
[[ "${TRAIN_PARALLEL_MODE}" =~ ^(none|ddp|fsdp)$ ]] || die "Invalid parallel mode: ${TRAIN_PARALLEL_MODE}"
if [[ "${TRAIN_PARALLEL_MODE}" == "none" ]] && (( TRAIN_GPU_COUNT != 1 )); then
die "Parallel mode none requires exactly one GPU"
fi
if [[ "${TRAIN_PARALLEL_MODE}" != "none" ]] && (( TRAIN_GPU_COUNT < 2 )); then
die "Parallel mode ${TRAIN_PARALLEL_MODE} requires at least two GPUs"
fi
if [[ -n "${TRAIN_CONFIG}" ]]; then
[[ -f "${TRAIN_CONFIG}" ]] || die "Training config not found: ${TRAIN_CONFIG}"
fi
@@ -36,11 +51,7 @@ if [[ -n "${TRAIN_CONFIG}" ]]; then
train_args+=(--config "${TRAIN_CONFIG}")
fi
if (( TRAIN_GPU_COUNT > 1 )); then
train_args+=(--parallel_mode ddp)
else
train_args+=(--parallel_mode none)
fi
train_args+=(--parallel_mode "${TRAIN_PARALLEL_MODE}")
if [[ -n "${latest_checkpoint}" ]]; then
log_info "Resuming ${TRAIN_JOB_NAME} from ${latest_checkpoint}"
@@ -52,7 +63,7 @@ else
train_args+=(--param_path "${BASE_MODEL}")
fi
log_info "GPUs=${TRAIN_GPU_COUNT}, checkpoints=${CHECKPOINT_DIR}"
log_info "GPUs=${TRAIN_GPU_COUNT}, parallel=${TRAIN_PARALLEL_MODE}, checkpoints=${CHECKPOINT_DIR}"
# Replace the shell so the container init forwards SIGTERM to the trainer.
exec "${train_args[@]}" "$@"
+189
View File
@@ -0,0 +1,189 @@
#!/usr/bin/env bash
set -euo pipefail
ROOT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")/.." && pwd)"
source "${ROOT_DIR}/scripts/docker/lib/train-common.sh"
COMPOSE_BASE=(
docker compose
--project-directory "${ROOT_DIR}"
--file "${ROOT_DIR}/docker-compose.yml"
)
usage() {
cat <<'EOF'
Usage: scripts/serve.sh <command> [CONFIG] [options]
CONFIG defaults to ./serve.yaml. The same file declares host runtime settings
under `runtime:` and server settings under `server:`.
Commands:
init [CONFIG] Create the model directory
preflight [CONFIG] Validate Docker, paths, GPU, and Compose
build [CONFIG] Build the serving image
up [CONFIG] Start the server container (detached)
run [CONFIG] Start the server container (foreground)
down [CONFIG] Stop and remove the server container
restart [CONFIG] Down, then up
logs [CONFIG] Follow server logs
status [CONFIG] Show container status
EOF
}
resolve_path() {
if [[ "$1" = /* ]]; then
printf '%s\n' "$1"
else
printf '%s/%s\n' "${ROOT_DIR}" "${1#./}"
fi
}
load_config() {
CONFIG_FILE="$(resolve_path "$1")"
[[ -f "${CONFIG_FILE}" ]] || die "Serving config not found: ${CONFIG_FILE}"
require_command python3
python3 -c 'import yaml' >/dev/null 2>&1 ||
die "PyYAML is required on the host (install python3-yaml)"
local exports
exports="$(python3 "${ROOT_DIR}/scripts/tools/serve_runtime.py" exports "${CONFIG_FILE}")" ||
die "Failed to load runtime configuration"
eval "${exports}"
if [[ -n "${SERVE_JOB_NAME}" ]]; then
validate_job_name "${SERVE_JOB_NAME}"
fi
}
compose() {
ASTRAI_UID="$(id -u)" ASTRAI_GID="$(id -g)" "${COMPOSE_BASE[@]}" "$@"
}
container_name() {
if [[ -n "${SERVE_JOB_NAME}" ]]; then
printf 'astrai-server-%s\n' "${SERVE_JOB_NAME}"
else
printf 'astrai-server\n'
fi
}
service_name() {
if [[ "${SERVE_GPU_ENABLED:-true}" == "false" ]]; then
printf 'server-cpu\n'
else
printf 'server\n'
fi
}
set_profile_args() {
PROFILE_ARGS=()
if [[ "${SERVE_GPU_ENABLED:-true}" == "false" ]]; then
PROFILE_ARGS=(--profile cpu)
fi
}
init_environment() {
mkdir -p "${SERVE_PARAM_DIR}"
log_info "Model: ${SERVE_PARAM_DIR}"
}
preflight() {
require_command docker
docker info >/dev/null 2>&1 || die "Docker daemon is unavailable"
[[ -d "${SERVE_PARAM_DIR}" ]] || die "Model directory not found: ${SERVE_PARAM_DIR}"
[[ -s "${SERVE_PARAM_DIR}/config.json" ]] ||
die "Model config not found: ${SERVE_PARAM_DIR}/config.json"
[[ -s "${SERVE_PARAM_DIR}/model.safetensors" ]] ||
die "Model weights not found: ${SERVE_PARAM_DIR}/model.safetensors"
if [[ "${SERVE_GPU_ENABLED}" == "false" ]] && [[ "${SERVE_DEVICE}" != "cpu" ]]; then
die "runtime.gpu.enabled is false but server.device is '${SERVE_DEVICE}'; use server.device: cpu"
fi
compose config --quiet
log_info "Preflight passed (service: $(service_name), device: ${SERVE_DEVICE})"
}
runtime_environment_args() {
RUNTIME_ENV_ARGS=()
local pair
while IFS= read -r -d '' pair; do
RUNTIME_ENV_ARGS+=(--env "${pair}")
done < <(python3 "${ROOT_DIR}/scripts/tools/serve_runtime.py" environment "${CONFIG_FILE}")
}
start_server() {
local foreground="$1"
shift
local container running
local -a run_options
preflight
runtime_environment_args
set_profile_args
container="$(container_name)"
running="$(docker inspect --format '{{.State.Running}}' "${container}" 2>/dev/null || true)"
[[ "${running}" != "true" ]] || die "Server is already running: ${container}"
docker rm "${container}" >/dev/null 2>&1 || true
run_options=(
--volume "${CONFIG_FILE}:/run/astrai/serve.yaml:ro"
"${RUNTIME_ENV_ARGS[@]}"
)
if [[ "${foreground}" == "true" ]]; then
compose "${PROFILE_ARGS[@]}" run --build --rm --service-ports \
"${run_options[@]}" "$(service_name)" \
python -m scripts.tools.server --config /run/astrai/serve.yaml "$@"
else
compose "${PROFILE_ARGS[@]}" run -d --build --service-ports \
--name "${container}" "${run_options[@]}" "$(service_name)" \
python -m scripts.tools.server --config /run/astrai/serve.yaml "$@"
log_info "Server started; run scripts/serve.sh logs ${CONFIG_FILE} to follow it"
fi
}
stop_server() {
local container
container="$(container_name)"
docker stop --timeout 30 "${container}" >/dev/null 2>&1 ||
log_warn "Server container is not running"
docker rm "${container}" >/dev/null 2>&1 || true
}
show_status() {
docker ps -a --filter "name=^/$(container_name)$"
}
main() {
local command="${1:-}" config="${SERVE_CONFIG_FILE:-${ROOT_DIR}/serve.yaml}"
[[ -n "${command}" ]] || { usage; exit 1; }
shift || true
if [[ "${command}" =~ ^(help|-h|--help)$ ]]; then
usage
return
fi
if [[ $# -gt 0 && "$1" != --* ]]; then
config="$1"
shift
fi
load_config "${config}"
case "${command}" in
init) init_environment ;;
preflight) preflight ;;
build)
set_profile_args
preflight
compose "${PROFILE_ARGS[@]}" build "$(service_name)"
;;
up) start_server false "$@" ;;
run) start_server true "$@" ;;
down) stop_server ;;
restart) stop_server; start_server false ;;
logs) docker logs -f --tail "${SERVE_LOG_TAIL:-200}" "$(container_name)" ;;
status) show_status ;;
*) die "Unknown command: ${command}" ;;
esac
}
main "$@"
+4 -1
View File
@@ -13,6 +13,7 @@ from astrai.inference.engine import InferenceEngine
from astrai.inference.runtime.graph import CudaGraphContext
from astrai.inference.workspace import InferenceWorkspace
from astrai.model import AutoModel, AutoRegressiveLM
from astrai.serialization import adapt_config
from astrai.tokenize import AutoTokenizer
_DTYPES = ["bfloat16", "float16", "float32"]
@@ -478,7 +479,9 @@ def benchmark_command(
if ckpt is not None:
click.echo(f"Loading model from {ckpt} ...")
config = ConfigFactory.load(
json.loads((Path(ckpt) / "config.json").read_text(encoding="utf-8-sig"))
adapt_config(
json.loads((Path(ckpt) / "config.json").read_text(encoding="utf-8-sig"))
)
)
model = AutoModel.from_pretrained(ckpt)
else:
+154
View File
@@ -0,0 +1,154 @@
"""Parse the host-side runtime section of a serving configuration.
The Compose wrapper needs a few container-side values on the host as well:
``server.port`` (the port the container listens on) and ``server.device``
(used by the preflight GPU consistency check). Everything else under
``server:`` is owned by ``scripts/tools/server.py --config`` inside the
container.
"""
import argparse
import re
import shlex
from pathlib import Path
import yaml
ENV_NAME = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$")
def _mapping(value, name: str) -> dict:
if value is None:
return {}
if not isinstance(value, dict):
raise ValueError(f"{name} must be a mapping")
return value
def _path(value, name: str, config_dir: Path) -> str:
if not isinstance(value, str) or not value.strip():
raise ValueError(f"runtime.paths.{name} is required")
path = Path(value).expanduser()
if not path.is_absolute():
path = config_dir / path
return str(path.resolve())
def _port(value, name: str) -> int:
if isinstance(value, bool) or not isinstance(value, int):
raise ValueError(f"{name} must be an integer")
if not 1 <= value <= 65535:
raise ValueError(f"{name} must be between 1 and 65535")
return value
def load_runtime(config_path: str) -> dict[str, str]:
path = Path(config_path).resolve()
with path.open(encoding="utf-8") as file:
config = yaml.safe_load(file) or {}
if not isinstance(config, dict):
raise ValueError("serving configuration must be a mapping")
runtime = _mapping(config.get("runtime"), "runtime")
if not runtime:
raise ValueError("top-level runtime section is required")
paths = _mapping(runtime.get("paths"), "paths")
gpu = _mapping(runtime.get("gpu"), "gpu")
container = _mapping(runtime.get("container"), "container")
environment = _mapping(runtime.get("environment"), "environment")
server = _mapping(config.get("server"), "server")
job_name = runtime.get("job_name", "")
if job_name and not isinstance(job_name, str):
raise ValueError("runtime.job_name must be a string")
if job_name and not re.fullmatch(r"[A-Za-z0-9][A-Za-z0-9._-]*", job_name):
raise ValueError(
"runtime.job_name must use letters, numbers, dot, underscore, or dash"
)
port = _port(runtime.get("port", 8000), "runtime.port")
container_port = _port(server.get("port", 8000), "server.port")
device = server.get("device", "cuda")
if not isinstance(device, str) or not device.strip():
raise ValueError("server.device must be a string")
gpu_enabled = gpu.get("enabled", True)
if not isinstance(gpu_enabled, bool):
raise ValueError("runtime.gpu.enabled must be a boolean")
devices = gpu.get("devices", "all")
if gpu_enabled:
if devices == "all":
visible_devices = ""
elif isinstance(devices, list) and len(devices) == 1:
text = str(devices[0])
if not text.isdigit():
raise ValueError(
"runtime.gpu.devices entries must be non-negative integers"
)
visible_devices = text
else:
raise ValueError(
"runtime.gpu.devices must be 'all' or a single-device list such as [0]"
)
else:
visible_devices = ""
if devices != "all":
raise ValueError(
"runtime.gpu.devices is ignored when runtime.gpu.enabled is false"
)
if device != "cpu":
raise ValueError(
"server.device must be 'cpu' when runtime.gpu.enabled is false"
)
values = {
"SERVE_JOB_NAME": job_name,
"SERVE_PORT": str(port),
"SERVE_CONTAINER_PORT": str(container_port),
"SERVE_PARAM_DIR": _path(paths.get("param", "./params"), "param", path.parent),
"SERVE_GPU_ENABLED": "true" if gpu_enabled else "false",
"SERVE_DEVICE": device,
"CUDA_VISIBLE_DEVICES": visible_devices,
"CUDA_TAG": str(container.get("cuda_tag", "cu128")),
}
for name, value in environment.items():
if not isinstance(name, str) or not ENV_NAME.fullmatch(name):
raise ValueError(f"invalid runtime.environment name: {name!r}")
if value is not None and not isinstance(value, (str, int, float, bool)):
raise ValueError(f"runtime.environment.{name} must be a scalar")
values["environment"] = environment
return values
def shell_exports(runtime: dict[str, str]) -> str:
return "\n".join(
f"export {name}={shlex.quote(value)}"
for name, value in runtime.items()
if name != "environment"
)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("command", choices=("exports", "environment"))
parser.add_argument("config")
args = parser.parse_args()
try:
runtime = load_runtime(args.config)
except (OSError, ValueError, yaml.YAMLError) as exc:
parser.error(str(exc))
if args.command == "exports":
print(shell_exports(runtime))
return
for name, value in runtime["environment"].items():
rendered = "" if value is None else str(value)
print(f"{name}={rendered}", end="\0")
if __name__ == "__main__":
main()
+128 -2
View File
@@ -2,16 +2,105 @@ from pathlib import Path
import click
import torch
import yaml
from click.core import ParameterSource
from astrai.inference import run_server
_DTYPES = ["bfloat16", "float16", "float32"]
_SERVER_KEYS = (
"host",
"port",
"reload",
"param_path",
"device",
"dtype",
"max_batch_size",
"max_seq_len",
)
def _merge_yaml_into_kwargs(
config_path: str,
passed_kwargs: dict,
explicit_keys: set[str] | None = None,
) -> dict:
"""Merge Click defaults, YAML server values, then explicit CLI values."""
with open(config_path, encoding="utf-8") as file:
config = yaml.safe_load(file) or {}
if not isinstance(config, dict):
raise click.UsageError(f"Serving config must be a mapping: {config_path}")
server = config.get("server") or {}
if not isinstance(server, dict):
raise click.UsageError("top-level server section must be a mapping")
unknown = sorted(set(server) - set(_SERVER_KEYS))
if unknown:
click.echo(
f"Warning: ignoring unknown server config keys: {', '.join(unknown)}",
err=True,
)
merged = dict(passed_kwargs)
merged.update({key: server[key] for key in _SERVER_KEYS if key in server})
if explicit_keys is None:
explicit_keys = set(passed_kwargs)
for key in explicit_keys:
if key in passed_kwargs:
merged[key] = passed_kwargs[key]
return merged
def _as_int(value, name: str) -> int | None:
if value is None:
return None
if isinstance(value, bool):
raise click.UsageError(f"{name} must be an integer")
try:
return int(value)
except (TypeError, ValueError):
raise click.UsageError(f"{name} must be an integer, got {value!r}") from None
def _resolve_server_config(
config_path: str,
passed_kwargs: dict,
explicit_keys: set[str] | None = None,
) -> dict:
"""Merge YAML values, then coerce and validate the resolved settings.
``explicit_keys`` are CLI flags that win over YAML; when None, YAML values
win over Click defaults.
"""
merged = _merge_yaml_into_kwargs(config_path, passed_kwargs, explicit_keys or set())
resolved = dict(merged)
resolved["port"] = _as_int(resolved["port"], "server.port") or 8000
resolved["max_batch_size"] = (
_as_int(resolved["max_batch_size"], "server.max_batch_size") or 16
)
resolved["max_seq_len"] = _as_int(resolved["max_seq_len"], "server.max_seq_len")
resolved["reload"] = bool(resolved["reload"])
if resolved["dtype"] not in _DTYPES:
raise click.UsageError(
f"server.dtype must be one of {', '.join(_DTYPES)}, got {resolved['dtype']!r}"
)
return resolved
@click.command(name="serve", help="Launch inference server (OpenAI-compatible API).")
@click.option(
"--config",
"-c",
"config_path",
type=click.Path(exists=True, dir_okay=False),
default=None,
help="Serving YAML config. CLI flags override YAML values.",
)
@click.option("--host", default="0.0.0.0", help="Host address.")
@click.option("--port", type=int, default=8000, help="Port number.")
@click.option("--reload", is_flag=True, help="Enable auto-reload for development.")
@click.option(
"--reload", is_flag=True, default=False, help="Enable auto-reload for development."
)
@click.option(
"--param_path",
type=click.Path(exists=True),
@@ -37,10 +126,47 @@ _DTYPES = ["bfloat16", "float16", "float32"]
default=None,
help="Maximum sequence length (KV cache size + prompt truncation). Uses model config if not set.",
)
@click.pass_context
def server_command(
host, port, reload, param_path, device, dtype, max_batch_size, max_seq_len
ctx,
config_path,
host,
port,
reload,
param_path,
device,
dtype,
max_batch_size,
max_seq_len,
):
"""Launch inference server (OpenAI-compatible API)."""
if config_path:
passed_kwargs = {
"host": host,
"port": port,
"reload": reload,
"param_path": param_path,
"device": device,
"dtype": dtype,
"max_batch_size": max_batch_size,
"max_seq_len": max_seq_len,
}
explicit_keys = {
key
for key in passed_kwargs
if ctx.get_parameter_source(key) is ParameterSource.COMMANDLINE
}
resolved = _resolve_server_config(config_path, passed_kwargs, explicit_keys)
host = resolved["host"]
port = resolved["port"]
reload = resolved["reload"]
param_path = resolved["param_path"]
device = resolved["device"]
dtype = resolved["dtype"]
max_batch_size = resolved["max_batch_size"]
max_seq_len = resolved["max_seq_len"]
click.echo(f"Config: {config_path}")
dtype_map = {
"bfloat16": torch.bfloat16,
"float16": torch.float16,
+13 -14
View File
@@ -11,6 +11,12 @@ from click.core import ParameterSource
from torch import optim
from astrai.config import AutoRegressiveLMConfig, TrainConfig
from astrai.config.train_config import (
BACKENDS,
PARALLEL_MODES,
START_METHODS,
TRAIN_TYPES,
)
from astrai.dataset import DatasetFactory, dpo_collate_fn, grpo_collate_fn
from astrai.model import AutoRegressiveLM
from astrai.model.components.decoder_block import DecoderBlock
@@ -92,12 +98,12 @@ def _merge_yaml_into_kwargs(
return merged
_TRAIN_TYPE = ["seq", "sft", "dpo", "grpo", "online_grpo", "online_dpo"]
_PARALLEL = ["none", "ddp", "fsdp"]
_TRAIN_TYPE = sorted(TRAIN_TYPES)
_PARALLEL = sorted(PARALLEL_MODES)
_SCHEDULES = ["cosine", "sgdr", "wsd"]
_OPTIMIZERS = OptimizerFactory.list_registered()
_BACKENDS = ["nccl", "gloo"]
_START_METHODS = ["spawn", "fork", "forkserver"]
_BACKENDS = sorted(BACKENDS)
_START_METHODS = sorted(START_METHODS)
@click.command(
@@ -651,17 +657,10 @@ def train(
decay_steps: int,
**kwargs,
):
if train_type not in [
"seq",
"sft",
"dpo",
"grpo",
"online_grpo",
"online_dpo",
]:
if train_type not in _TRAIN_TYPE:
raise ValueError(
f"Invalid train_type '{train_type}'. "
f"Must be one of: seq, sft, dpo, grpo, online_grpo, online_dpo"
f"Must be one of: {', '.join(_TRAIN_TYPE)}"
)
if not os.path.exists(param_path):
raise FileNotFoundError(f"Model directory not found: {param_path}")
@@ -837,7 +836,7 @@ def train(
gradient_checkpointing_modules=grad_ckpt_modules,
compile_mode=compile_mode,
executor_kwargs=executor_kwargs,
extra_kwargs=strategy_kwargs,
strategy_kwargs=strategy_kwargs,
neftune_alpha=neftune_alpha,
collate_fn=collate_fn,
rollout_interval=rollout_interval,
+152
View File
@@ -0,0 +1,152 @@
"""Parse the host-side runtime section of a training configuration."""
import argparse
import math
import re
import shlex
from pathlib import Path
import yaml
ENV_NAME = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$")
PARALLEL_MODES = {"auto", "none", "ddp", "fsdp"}
def _mapping(value, name: str) -> dict:
if value is None:
return {}
if not isinstance(value, dict):
raise ValueError(f"runtime.{name} must be a mapping")
return value
def _path(value, name: str, config_dir: Path) -> str:
if not isinstance(value, str) or not value.strip():
raise ValueError(f"runtime.paths.{name} is required")
path = Path(value).expanduser()
if not path.is_absolute():
path = config_dir / path
return str(path.resolve())
def load_runtime(config_path: str) -> dict[str, str]:
path = Path(config_path).resolve()
with path.open(encoding="utf-8") as file:
config = yaml.safe_load(file) or {}
if not isinstance(config, dict):
raise ValueError("training configuration must be a mapping")
runtime = _mapping(config.get("runtime"), "runtime")
if not runtime:
raise ValueError("top-level runtime section is required")
paths = _mapping(runtime.get("paths"), "paths")
gpu = _mapping(runtime.get("gpu"), "gpu")
container = _mapping(runtime.get("container"), "container")
environment = _mapping(runtime.get("environment"), "environment")
job_name = runtime.get("job_name")
if not isinstance(job_name, str) or not re.fullmatch(
r"[A-Za-z0-9][A-Za-z0-9._-]*", job_name
):
raise ValueError(
"runtime.job_name must use letters, numbers, dot, underscore, or dash"
)
devices = gpu.get("devices", "all")
if devices == "all":
gpu_count = "all"
visible_devices = ""
elif isinstance(devices, list) and devices:
normalized = []
for device in devices:
text = str(device)
if not text.isdigit():
raise ValueError(
"runtime.gpu.devices entries must be non-negative integers"
)
normalized.append(text)
if len(set(normalized)) != len(normalized):
raise ValueError("runtime.gpu.devices must not contain duplicates")
gpu_count = str(len(normalized))
visible_devices = ",".join(normalized)
else:
raise ValueError("runtime.gpu.devices must be 'all' or a non-empty list")
parallel_mode = str(gpu.get("parallel_mode", "auto"))
if parallel_mode not in PARALLEL_MODES:
raise ValueError("runtime.gpu.parallel_mode must be auto, none, ddp, or fsdp")
if gpu_count != "all":
count = int(gpu_count)
if parallel_mode == "none" and count != 1:
raise ValueError("parallel_mode none requires exactly one GPU")
if parallel_mode in {"ddp", "fsdp"} and count < 2:
raise ValueError(
f"parallel_mode {parallel_mode} requires at least two GPUs"
)
max_hours = container.get("max_duration_hours", 0)
try:
max_seconds = math.ceil(float(max_hours) * 3600) if max_hours else 0
except (TypeError, ValueError) as exc:
raise ValueError(
"runtime.container.max_duration_hours must be a number"
) from exc
if max_seconds < 0:
raise ValueError("runtime.container.max_duration_hours must not be negative")
values = {
"TRAIN_JOB_NAME": job_name,
"TRAIN_DATA_DIR": _path(paths.get("data"), "data", path.parent),
"TRAIN_MODEL_DIR": _path(paths.get("model"), "model", path.parent),
"TRAIN_CHECKPOINT_DIR": _path(
paths.get("checkpoints"), "checkpoints", path.parent
),
"TRAIN_GPU_COUNT": gpu_count,
"CUDA_VISIBLE_DEVICES": visible_devices,
"TRAIN_PARALLEL_MODE": parallel_mode,
"CUDA_TAG": str(container.get("cuda_tag", "cu128")),
"TRAIN_IPC_MODE": str(container.get("ipc", "host")),
"TRAIN_STOP_GRACE_PERIOD": str(container.get("stop_grace_period", "10m")),
"TRAIN_STOP_TIMEOUT": str(container.get("stop_timeout_seconds", 600)),
"CHECKPOINT_KEEP_LAST": str(container.get("checkpoint_keep_last", 5)),
"TRAIN_MAX_DURATION_SECONDS": str(max_seconds),
}
for name, value in environment.items():
if not isinstance(name, str) or not ENV_NAME.fullmatch(name):
raise ValueError(f"invalid runtime.environment name: {name!r}")
if value is not None and not isinstance(value, (str, int, float, bool)):
raise ValueError(f"runtime.environment.{name} must be a scalar")
values["environment"] = environment
return values
def shell_exports(runtime: dict[str, str]) -> str:
return "\n".join(
f"export {name}={shlex.quote(value)}"
for name, value in runtime.items()
if name != "environment"
)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("command", choices=("exports", "environment"))
parser.add_argument("config")
args = parser.parse_args()
try:
runtime = load_runtime(args.config)
except (OSError, ValueError, yaml.YAMLError) as exc:
parser.error(str(exc))
if args.command == "exports":
print(shell_exports(runtime))
return
for name, value in runtime["environment"].items():
rendered = "" if value is None else str(value)
print(f"{name}={rendered}", end="\0")
if __name__ == "__main__":
main()
+141 -183
View File
@@ -4,7 +4,6 @@ set -euo pipefail
ROOT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")/.." && pwd)"
source "${ROOT_DIR}/scripts/docker/lib/train-common.sh"
ENV_FILE="${TRAIN_ENV_FILE:-${ROOT_DIR}/.env.train}"
COMPOSE_BASE=(
docker compose
--project-directory "${ROOT_DIR}"
@@ -14,50 +13,28 @@ COMPOSE_BASE=(
usage() {
cat <<'EOF'
Usage: scripts/train.sh <command> [options]
Usage: scripts/train.sh <command> [CONFIG] [options]
CONFIG defaults to ./train.yaml. The same file declares host runtime settings
under `runtime:` and trainer settings under model/data/parallel/training/ckpt/log.
Commands:
init Create local directories and .env.train
preflight Validate Docker, paths, GPU settings, and Compose
build Build the trainer image
start [--foreground] [-- ARGS...] Start or resume training
stop Gracefully stop and checkpoint training
restart Stop, then start training
logs Follow trainer logs
status Show container and latest checkpoint status
latest Print the latest complete checkpoint path
list List all complete checkpoints
clean [--keep N] Preview old checkpoint removal
clean --force Remove old checkpoints after previewing
Environment:
TRAIN_ENV_FILE Env file path (default: .env.train)
TRAIN_CONFIG_FILE Optional host YAML mounted only when the job starts
Training arguments come from an externally mounted TRAIN_CONFIG or ARGS passed
after --. The image does not contain experiment configuration.
init [CONFIG] Create runtime directories
preflight [CONFIG] Validate Docker, paths, GPUs, and Compose
build [CONFIG] Build the trainer image
start [CONFIG] [--foreground] [-- ARGS...]
Start or resume training
stop [CONFIG] Gracefully stop and checkpoint training
restart [CONFIG] Stop, then start training
logs [CONFIG] Follow trainer logs
status [CONFIG] Show container and checkpoint status
latest [CONFIG] Print the latest complete checkpoint
list [CONFIG] List complete checkpoints
clean [CONFIG] [--keep N] [--force]
Preview or remove old checkpoints
EOF
}
load_env() {
if [[ -f "${ENV_FILE}" ]]; then
set -a
# UID/GID are readonly in bash; compose gets them via ASTRAI_UID/GID in compose()
# shellcheck disable=SC1090
source <(grep -v -E '^[[:space:]]*(UID|GID)=' "${ENV_FILE}")
set +a
fi
TRAIN_JOB_NAME="${TRAIN_JOB_NAME:-astrai-train}"
TRAIN_DATA_DIR="${TRAIN_DATA_DIR:-./data}"
TRAIN_MODEL_DIR="${TRAIN_MODEL_DIR:-./params}"
TRAIN_CHECKPOINT_DIR="${TRAIN_CHECKPOINT_DIR:-./checkpoints}"
TRAIN_GPU_COUNT="${TRAIN_GPU_COUNT:-all}"
TRAIN_STOP_TIMEOUT="${TRAIN_STOP_TIMEOUT:-600}"
validate_job_name "${TRAIN_JOB_NAME}"
}
resolve_path() {
if [[ "$1" = /* ]]; then
printf '%s\n' "$1"
@@ -66,142 +43,147 @@ resolve_path() {
fi
}
checkpoint_dir() {
printf '%s/%s\n' "$(resolve_path "${TRAIN_CHECKPOINT_DIR}")" "${TRAIN_JOB_NAME}"
load_config() {
CONFIG_FILE="$(resolve_path "$1")"
[[ -f "${CONFIG_FILE}" ]] || die "Training config not found: ${CONFIG_FILE}"
require_command python3
python3 -c 'import yaml' >/dev/null 2>&1 ||
die "PyYAML is required on the host (install python3-yaml)"
local exports
exports="$(python3 "${ROOT_DIR}/scripts/tools/train_runtime.py" exports "${CONFIG_FILE}")" ||
die "Failed to load runtime configuration"
eval "${exports}"
validate_job_name "${TRAIN_JOB_NAME}"
}
compose() {
local -a command=("${COMPOSE_BASE[@]}")
ASTRAI_UID="$(id -u)" ASTRAI_GID="$(id -g)" "${COMPOSE_BASE[@]}" "$@"
}
if [[ -f "${ENV_FILE}" ]]; then
command+=(--env-file "${ENV_FILE}")
checkpoint_dir() {
printf '%s/%s\n' "${TRAIN_CHECKPOINT_DIR}" "${TRAIN_JOB_NAME}"
}
container_name() {
printf 'astrai-trainer-%s\n' "${TRAIN_JOB_NAME}"
}
timer_pid_file() {
printf '/tmp/astrai-timer-%s.pid\n' "${TRAIN_JOB_NAME}"
}
timer_log_file() {
printf '/tmp/astrai-timer-%s.log\n' "${TRAIN_JOB_NAME}"
}
cancel_timer() {
local pid_file pid
pid_file="$(timer_pid_file)"
[[ -f "${pid_file}" ]] || return 0
pid="$(<"${pid_file}")"
if [[ "${pid}" =~ ^[1-9][0-9]*$ ]] && kill -0 "${pid}" 2>/dev/null; then
kill "${pid}" 2>/dev/null || true
fi
rm -f -- "${pid_file}"
}
# Inject the host user into compose so container processes share the
# checkpoint directory ownership (bash UID/GID are readonly).
ASTRAI_UID="$(id -u)" ASTRAI_GID="$(id -g)" "${command[@]}" "$@"
schedule_timer() {
(( TRAIN_MAX_DURATION_SECONDS > 0 )) || return 0
cancel_timer
local pid_file log_file
pid_file="$(timer_pid_file)"
log_file="$(timer_log_file)"
(
sleep "${TRAIN_MAX_DURATION_SECONDS}"
"${ROOT_DIR}/scripts/train.sh" stop "${CONFIG_FILE}" --from-timer
) >"${log_file}" 2>&1 &
printf '%s\n' "$!" >"${pid_file}"
log_info "Automatic stop scheduled in ${TRAIN_MAX_DURATION_SECONDS}s"
}
init_environment() {
local data_dir model_dir checkpoints_dir
data_dir="$(resolve_path "${TRAIN_DATA_DIR}")"
model_dir="$(resolve_path "${TRAIN_MODEL_DIR}")"
checkpoints_dir="$(resolve_path "${TRAIN_CHECKPOINT_DIR}")"
mkdir -p "${data_dir}" "${model_dir}" "${checkpoints_dir}"
if [[ ! -f "${ENV_FILE}" ]]; then
cat >"${ENV_FILE}" <<'EOF'
TRAIN_JOB_NAME=astrai-train
TRAIN_DATA_DIR=./data
TRAIN_MODEL_DIR=./params
TRAIN_CHECKPOINT_DIR=./checkpoints
TRAIN_CONFIG_FILE=
TRAIN_GPU_COUNT=all
# CUDA_VISIBLE_DEVICES=0,1
CUDA_TAG=cu128
TRAIN_IPC_MODE=host
TRAIN_STOP_GRACE_PERIOD=10m
TRAIN_STOP_TIMEOUT=600
CHECKPOINT_KEEP_LAST=5
EOF
log_info "Created ${ENV_FILE}"
else
log_info "Keeping existing ${ENV_FILE}"
fi
log_info "Data: ${data_dir}"
log_info "Model: ${model_dir}"
log_info "Checkpoints: ${checkpoints_dir}"
mkdir -p "${TRAIN_DATA_DIR}" "${TRAIN_MODEL_DIR}" "${TRAIN_CHECKPOINT_DIR}"
log_info "Data: ${TRAIN_DATA_DIR}"
log_info "Model: ${TRAIN_MODEL_DIR}"
log_info "Checkpoints: ${TRAIN_CHECKPOINT_DIR}"
}
preflight() {
local data_dir model_dir checkpoints_dir config_file latest visible_count
local latest visible_count
require_command docker
docker info >/dev/null 2>&1 || die "Docker daemon is unavailable"
[[ "${TRAIN_GPU_COUNT}" == "all" || "${TRAIN_GPU_COUNT}" =~ ^[1-9][0-9]*$ ]] ||
die "TRAIN_GPU_COUNT must be 'all' or a positive integer"
[[ -d "${TRAIN_DATA_DIR}" ]] || die "Training data directory not found: ${TRAIN_DATA_DIR}"
mkdir -p "$(checkpoint_dir)"
[[ -w "$(checkpoint_dir)" ]] || die "Checkpoint directory is not writable: $(checkpoint_dir)"
data_dir="$(resolve_path "${TRAIN_DATA_DIR}")"
model_dir="$(resolve_path "${TRAIN_MODEL_DIR}")"
checkpoints_dir="$(resolve_path "${TRAIN_CHECKPOINT_DIR}")"
[[ -d "${data_dir}" ]] || die "Training data directory not found: ${data_dir}"
mkdir -p "${checkpoints_dir}/${TRAIN_JOB_NAME}"
[[ -w "${checkpoints_dir}/${TRAIN_JOB_NAME}" ]] || die "Checkpoint directory is not writable"
if [[ -n "${TRAIN_CONFIG_FILE:-}" ]]; then
config_file="$(resolve_path "${TRAIN_CONFIG_FILE}")"
[[ -f "${config_file}" ]] || die "Training config not found: ${config_file}"
fi
latest="$(find_latest_checkpoint "${checkpoints_dir}/${TRAIN_JOB_NAME}" || true)"
latest="$(find_latest_checkpoint "$(checkpoint_dir)" || true)"
if [[ -z "${latest}" ]]; then
[[ -s "${model_dir}/config.json" ]] || die "Model config not found: ${model_dir}/config.json"
[[ -s "${model_dir}/model.safetensors" ]] || die "Model weights not found: ${model_dir}/model.safetensors"
[[ -s "${TRAIN_MODEL_DIR}/config.json" ]] ||
die "Model config not found: ${TRAIN_MODEL_DIR}/config.json"
[[ -s "${TRAIN_MODEL_DIR}/model.safetensors" ]] ||
die "Model weights not found: ${TRAIN_MODEL_DIR}/model.safetensors"
else
log_info "Resume candidate: ${latest}"
fi
if [[ -n "${CUDA_VISIBLE_DEVICES:-}" && "${TRAIN_GPU_COUNT}" != "all" ]]; then
if [[ -n "${CUDA_VISIBLE_DEVICES}" ]]; then
IFS=',' read -r -a visible_gpus <<<"${CUDA_VISIBLE_DEVICES}"
visible_count="${#visible_gpus[@]}"
(( visible_count == TRAIN_GPU_COUNT )) ||
die "TRAIN_GPU_COUNT=${TRAIN_GPU_COUNT}, but CUDA_VISIBLE_DEVICES exposes ${visible_count} GPU(s)"
die "Configured GPU count and visible device list disagree"
fi
compose config --quiet
log_info "Preflight passed for ${TRAIN_JOB_NAME} (GPU request: ${TRAIN_GPU_COUNT})"
log_info "Preflight passed for ${TRAIN_JOB_NAME} (GPU request: ${TRAIN_GPU_COUNT}, parallel: ${TRAIN_PARALLEL_MODE})"
}
runtime_environment_args() {
RUNTIME_ENV_ARGS=()
local pair
while IFS= read -r -d '' pair; do
RUNTIME_ENV_ARGS+=(--env "${pair}")
done < <(python3 "${ROOT_DIR}/scripts/tools/train_runtime.py" environment "${CONFIG_FILE}")
}
start_training() {
local foreground="$1"
local config_file container running
local -a run_options=()
shift
local container running
local -a run_options
preflight
if [[ -n "${TRAIN_CONFIG_FILE:-}" ]]; then
config_file="$(resolve_path "${TRAIN_CONFIG_FILE}")"
run_options+=(
--volume "${config_file}:/run/astrai/train.yaml:ro"
--env TRAIN_CONFIG=/run/astrai/train.yaml
)
elif [[ -z "${TRAIN_CONFIG:-}" && $# -eq 0 ]]; then
die "Set TRAIN_CONFIG_FILE or pass complete trainer arguments after --"
fi
container="astrai-trainer-${TRAIN_JOB_NAME}"
runtime_environment_args
container="$(container_name)"
running="$(docker inspect --format '{{.State.Running}}' "${container}" 2>/dev/null || true)"
[[ "${running}" != "true" ]] || die "Trainer is already running: ${container}"
docker rm "${container}" >/dev/null 2>&1 || true
run_options=(
--volume "${CONFIG_FILE}:/run/astrai/train.yaml:ro"
--env TRAIN_CONFIG=/run/astrai/train.yaml
"${RUNTIME_ENV_ARGS[@]}"
)
if [[ "${foreground}" == "true" ]]; then
compose run --build --rm "${run_options[@]}" trainer "$@"
else
compose run -d --build --name "${container}" \
"${run_options[@]}" trainer "$@"
log_info "Training started; run scripts/train.sh logs to follow it"
compose run -d --build --name "${container}" "${run_options[@]}" trainer "$@"
schedule_timer
log_info "Training started; run scripts/train.sh logs ${CONFIG_FILE} to follow it"
fi
}
stop_training() {
local from_timer="$1"
[[ "${from_timer}" == "true" ]] || cancel_timer
log_info "Stopping trainer with ${TRAIN_STOP_TIMEOUT}s grace period"
docker stop --timeout "${TRAIN_STOP_TIMEOUT}" "astrai-trainer-${TRAIN_JOB_NAME}" >/dev/null 2>&1 ||
docker stop --timeout "${TRAIN_STOP_TIMEOUT}" "$(container_name)" >/dev/null 2>&1 ||
log_warn "Trainer container is not running"
}
restart_training() {
local container="astrai-trainer-${TRAIN_JOB_NAME}"
docker inspect "${container}" >/dev/null 2>&1 ||
die "Trainer container not found; use start with a config or CLI arguments first"
log_info "Restarting trainer with ${TRAIN_STOP_TIMEOUT}s grace period"
docker restart --timeout "${TRAIN_STOP_TIMEOUT}" "${container}" >/dev/null
[[ "${from_timer}" != "true" ]] || rm -f -- "$(timer_pid_file)"
}
show_status() {
local latest
docker ps -a --filter "name=^/astrai-trainer-${TRAIN_JOB_NAME}$"
docker ps -a --filter "name=^/$(container_name)$"
latest="$(find_latest_checkpoint "$(checkpoint_dir)" || true)"
if [[ -n "${latest}" ]]; then
log_info "Latest checkpoint: ${latest}"
@@ -213,7 +195,6 @@ show_status() {
clean_checkpoints() {
local keep="$1" force="$2" dir count remove_count index path
local -a checkpoints=()
[[ "${keep}" =~ ^[1-9][0-9]*$ ]] || die "--keep must be a positive integer"
dir="$(checkpoint_dir)"
while IFS= read -r line; do
@@ -226,7 +207,6 @@ clean_checkpoints() {
log_info "Nothing to clean; ${count} complete checkpoint(s), keeping ${keep}"
return
fi
for ((index = 0; index < remove_count; index++)); do
path="${checkpoints[index]}"
if [[ "${force}" == "true" ]]; then
@@ -240,57 +220,49 @@ clean_checkpoints() {
}
main() {
local command="${1:-}" foreground=false keep="${CHECKPOINT_KEEP_LAST:-5}" force=false
local command="${1:-}" config="${TRAIN_CONFIG_FILE:-${ROOT_DIR}/train.yaml}"
local foreground=false keep force=false from_timer=false
local -a train_args=()
[[ -n "${command}" ]] || { usage; exit 1; }
shift || true
load_env
if [[ "${command}" =~ ^(help|-h|--help)$ ]]; then
usage
return
fi
if [[ $# -gt 0 && "$1" != --* ]]; then
config="$1"
shift
fi
load_config "${config}"
keep="${CHECKPOINT_KEEP_LAST}"
case "${command}" in
init)
init_environment
;;
preflight)
preflight
;;
build)
preflight
compose build trainer
;;
init) init_environment ;;
preflight) preflight ;;
build) preflight; compose build trainer ;;
start)
while [[ $# -gt 0 ]]; do
case "$1" in
--foreground)
foreground=true
shift
;;
--)
shift
train_args=("$@")
break
;;
*)
die "Unknown start option: $1 (put trainer arguments after --)"
;;
--foreground) foreground=true; shift ;;
--) shift; train_args=("$@"); break ;;
*) die "Unknown start option: $1 (put trainer arguments after --)" ;;
esac
done
start_training "${foreground}" "${train_args[@]}"
;;
stop)
stop_training
[[ "${1:-}" != "--from-timer" ]] || from_timer=true
stop_training "${from_timer}"
;;
restart)
restart_training
;;
logs)
docker logs -f --tail "${TRAIN_LOG_TAIL:-200}" "astrai-trainer-${TRAIN_JOB_NAME}"
;;
status)
show_status
;;
latest)
find_latest_checkpoint "$(checkpoint_dir)" || die "No complete checkpoint found"
stop_training false
start_training false
;;
logs) docker logs -f --tail "${TRAIN_LOG_TAIL:-200}" "$(container_name)" ;;
status) show_status ;;
latest) find_latest_checkpoint "$(checkpoint_dir)" || die "No complete checkpoint found" ;;
list)
list_complete_checkpoints "$(checkpoint_dir)" | while read -r _epoch _step path; do
printf '%s\n' "${path}"
@@ -299,28 +271,14 @@ main() {
clean)
while [[ $# -gt 0 ]]; do
case "$1" in
--keep)
[[ $# -ge 2 ]] || die "--keep requires a value"
keep="$2"
shift 2
;;
--force)
force=true
shift
;;
*)
die "Unknown clean option: $1"
;;
--keep) [[ $# -ge 2 ]] || die "--keep requires a value"; keep="$2"; shift 2 ;;
--force) force=true; shift ;;
*) die "Unknown clean option: $1" ;;
esac
done
clean_checkpoints "${keep}" "${force}"
;;
help|-h|--help)
usage
;;
*)
die "Unknown command: ${command}"
;;
*) die "Unknown command: ${command}" ;;
esac
}
+1 -1
View File
@@ -78,7 +78,7 @@ class _CMakeBuildExt(_build_ext):
if cmake is None:
raise RuntimeError("cmake not found on PATH; install it to build kernels")
parallel = os.environ.get("BUILD_PARALLEL", "16")
parallel = os.environ.get("BUILD_PARALLEL", "4")
cfg = [
cmake,
"-S",
+29 -25
View File
@@ -1,21 +1,31 @@
import json
import os
import shutil
import tempfile
import pytest
import torch
from tokenizers import Tokenizer, models, pre_tokenizers, trainers
from astrai.extension import KERNEL_NAMES, is_available
from astrai.model.transformer import AutoRegressiveLM
from astrai.tokenize import AutoTokenizer
from tests.helpers import TINY_CONFIG, RandomTokenDataset, make_tiny_config
from tests.helpers import (
TINY_CONFIG,
RandomTokenDataset,
build_test_tokenizer,
make_tiny_config,
)
CUDA_AVAIL = torch.cuda.is_available()
KERNEL_AVAIL = CUDA_AVAIL and all(is_available(k) for k in KERNEL_NAMES)
FP8_AVAIL = (
CUDA_AVAIL
and is_available("fp8_ops")
and torch.cuda.get_device_capability() >= (8, 9)
)
skip_no_cuda = pytest.mark.skipif(not CUDA_AVAIL, reason="CUDA not available")
skip_no_kernel = pytest.mark.skipif(not KERNEL_AVAIL, reason="CUDA kernels not built")
skip_no_fp8 = pytest.mark.skipif(
not FP8_AVAIL,
reason="fused FP8 MMA requires a built kernel and compute capability 8.9+",
)
def pytest_configure(config):
@@ -30,18 +40,9 @@ def device():
return "cuda" if torch.cuda.is_available() else "cpu"
def create_test_tokenizer(vocab_size: int = 1000) -> AutoTokenizer:
def create_test_tokenizer(vocab_size: int = 1000):
"""Create a simple tokenizer for testing purposes."""
tokenizer = Tokenizer(models.BPE())
tokenizer.pre_tokenizer = pre_tokenizers.ByteLevel()
trainer = trainers.BpeTrainer(
vocab_size=vocab_size, min_frequency=1, special_tokens=["<unk>", "<pad>"]
)
tokenizer.train_from_iterator([chr(i) for i in range(256)], trainer)
auto_tokenizer = AutoTokenizer()
auto_tokenizer._tokenizer = tokenizer
auto_tokenizer._special_token_map = {"unk_token": "<unk>", "pad_token": "<pad>"}
return auto_tokenizer
return build_test_tokenizer(vocab_size)
@pytest.fixture(scope="session")
@@ -50,33 +51,36 @@ def test_tokenizer():
return create_test_tokenizer()
@pytest.fixture(scope="session")
@pytest.fixture
def test_model(device):
"""Session-scoped small AutoRegressiveLM model, created once."""
"""Function-scoped small AutoRegressiveLM model, isolated per test."""
config = make_tiny_config()
model = AutoRegressiveLM(config).to(device=device)
return {"model": model, "device": device, "config": config}
@pytest.fixture
def base_test_env(test_model, test_tokenizer):
def temp_dir(tmp_path):
"""Function-scoped temporary directory, cleaned up by pytest."""
return str(tmp_path)
@pytest.fixture
def base_test_env(test_model, test_tokenizer, temp_dir):
"""Function-scoped test environment with isolated temp directory."""
test_dir = tempfile.mkdtemp()
config_path = os.path.join(test_dir, "config.json")
config_path = os.path.join(temp_dir, "config.json")
with open(config_path, "w") as f:
json.dump(TINY_CONFIG, f)
yield {
return {
"device": test_model["device"],
"test_dir": str(test_dir),
"test_dir": temp_dir,
"config_path": config_path,
"transformer_config": test_model["config"],
"model": test_model["model"],
"tokenizer": test_tokenizer,
}
shutil.rmtree(test_dir)
@pytest.fixture
def random_dataset():
+43 -169
View File
@@ -1,21 +1,14 @@
import json
import os
import tempfile
import pytest
from tokenizers import Tokenizer, models, pre_tokenizers, trainers
from astrai.config.preprocess_config import (
InputConfig,
PipelineConfig,
ProcessingConfig,
)
from astrai.preprocessing.builder import (
MultiOutputMaskBuilder,
SectionedMaskBuilder,
SingleOutputMaskBuilder,
)
from astrai.tokenize import AutoTokenizer
from tests.helpers import build_test_tokenizer
_SPECIAL_TOKENS_CONFIG = {
"bos_token": "<|begin_of_sentence|>",
@@ -41,55 +34,42 @@ _CHAT_TEMPLATE = (
"{% if add_generation_prompt %}<|im_start|>assistant\n{% endif %}"
)
_CHAT_SECTIONS = [{"field": "messages", "action": "$role", "template": True}]
_INSTRUCTION_SECTIONS = [
{"field": "prompt", "action": "mask", "add_special_tokens": True},
{"field": "response", "action": "train"},
_CHAT_TOKENIZER_DATA = [
"hello world",
"Hi there!",
"You are helpful.",
"What is 2+2?",
"Tell me a story about dragons and knights.",
"Sure, here is a tale.",
"Translate to French: Hello",
"Bonjour",
"Artificial Intelligence is a field of computer science.",
"system",
"user",
"assistant",
"<|im_start|>",
"<|im_end|>",
*[chr(i) for i in range(32, 127)],
]
_TEXT_SECTIONS = [{"field": "text", "action": "train"}]
_GRPO_RESPONSE_SECTIONS = [{"field": "responses", "action": "train"}]
_CHAT_TOKENIZER_MAP = {
"bos_token": "<|begin_of_sentence|>",
"eos_token": "<|end_of_sentence|>",
"pad_token": "<|_pad_|>",
"unk_token": "<|_unk_|>",
}
def _build_chat_tokenizer():
tok = Tokenizer(models.BPE())
tok.pre_tokenizer = pre_tokenizers.ByteLevel(add_prefix_space=False)
tr = trainers.BpeTrainer(
return build_test_tokenizer(
vocab_size=512,
min_frequency=1,
special_tokens=_SPECIAL_TOKENS,
special_token_map=_CHAT_TOKENIZER_MAP,
add_prefix_space=False,
train_data=_CHAT_TOKENIZER_DATA,
chat_template=_CHAT_TEMPLATE,
)
train_data = [
"hello world",
"Hi there!",
"You are helpful.",
"What is 2+2?",
"Tell me a story about dragons and knights.",
"Sure, here is a tale.",
"Translate to French: Hello",
"Bonjour",
"Artificial Intelligence is a field of computer science.",
"system",
"user",
"assistant",
"<|im_start|>",
"<|im_end|>",
*[chr(i) for i in range(32, 127)],
]
tok.train_from_iterator(train_data, tr)
auto_tok = AutoTokenizer()
auto_tok._tokenizer = tok
auto_tok._special_token_map = {
"bos_token": "<|begin_of_sentence|>",
"eos_token": "<|end_of_sentence|>",
"pad_token": "<|_pad_|>",
"unk_token": "<|_unk_|>",
}
auto_tok.set_chat_template(_CHAT_TEMPLATE)
return auto_tok
@pytest.fixture(scope="session")
@@ -97,116 +77,11 @@ def chat_tokenizer():
return _build_chat_tokenizer()
@pytest.fixture
def temp_dir():
d = tempfile.mkdtemp()
yield d
import shutil
shutil.rmtree(d, ignore_errors=True)
def make_chat_config():
return PipelineConfig(
input=InputConfig(sections=_CHAT_SECTIONS),
mask={"system": "mask", "user": "mask", "assistant": "train"},
mask_default="mask",
preprocessing=ProcessingConfig(max_seq_len=2048),
)
def make_instruction_config():
return PipelineConfig(
input=InputConfig(sections=_INSTRUCTION_SECTIONS),
mask={"prompt": "mask", "response": "train"},
mask_default="mask",
preprocessing=ProcessingConfig(max_seq_len=2048),
)
def make_text_config():
return PipelineConfig(
input=InputConfig(sections=_TEXT_SECTIONS),
preprocessing=ProcessingConfig(
max_seq_len=2048, min_chars=1, max_chars=2_000_000
),
)
def make_dpo_chat_config():
return PipelineConfig(
input=InputConfig(
sources={
"chosen": {
"sections": [
{"field": "chosen", "action": "$role", "template": True}
]
},
"rejected": {
"sections": [
{"field": "rejected", "action": "$role", "template": True}
]
},
}
),
mask={"user": "mask", "assistant": "train"},
mask_default="mask",
preprocessing=ProcessingConfig(max_seq_len=2048),
)
def make_grpo_config():
return PipelineConfig(
input=InputConfig(
sources={
"prompts": {
"sections": [
{"field": "prompt", "action": "mask", "template": True}
]
},
"responses": {
"sections": _GRPO_RESPONSE_SECTIONS,
"list_field": True,
"mask_key": "masks",
},
"rewards": {
"sections": [{"field": "rewards", "action": "value"}],
},
}
),
mask={"user": "mask", "assistant": "train"},
mask_default="mask",
preprocessing=ProcessingConfig(max_seq_len=2048),
)
def make_grpo_no_template_config():
return PipelineConfig(
input=InputConfig(
sources={
"prompts": {
"sections": [
{
"field": "prompt",
"action": "mask",
"add_special_tokens": True,
}
]
},
"responses": {
"sections": _GRPO_RESPONSE_SECTIONS,
"list_field": True,
"mask_key": "masks",
},
"rewards": {
"sections": [{"field": "rewards", "action": "value"}],
},
}
),
mask={"user": "mask", "assistant": "train"},
mask_default="mask",
preprocessing=ProcessingConfig(max_seq_len=2048),
)
def _write_tokenizer_dir(dir_path, tokenizer, tokenizer_config):
"""Persist a tokenizer plus ``tokenizer_config.json`` into *dir_path*."""
tokenizer._tokenizer.save(os.path.join(dir_path, "tokenizer.json"))
with open(os.path.join(dir_path, "tokenizer_config.json"), "w") as f:
json.dump(tokenizer_config, f)
@pytest.fixture
@@ -228,11 +103,11 @@ def multi_builder():
def tokenizer_dir(temp_dir, test_tokenizer):
d = os.path.join(temp_dir, "tok")
os.makedirs(d, exist_ok=True)
test_tokenizer._tokenizer.save(os.path.join(d, "tokenizer.json"))
with open(os.path.join(d, "tokenizer_config.json"), "w") as f:
json.dump(
{"special_tokens": {"pad_token": "<|_pad_|>", "unk_token": "<|_unk_|>"}}, f
)
_write_tokenizer_dir(
d,
test_tokenizer,
{"special_tokens": {"pad_token": "<|_pad_|>", "unk_token": "<|_unk_|>"}},
)
return d
@@ -240,10 +115,9 @@ def tokenizer_dir(temp_dir, test_tokenizer):
def chat_tokenizer_dir(temp_dir, chat_tokenizer):
d = os.path.join(temp_dir, "tok")
os.makedirs(d, exist_ok=True)
chat_tokenizer._tokenizer.save(os.path.join(d, "tokenizer.json"))
with open(os.path.join(d, "tokenizer_config.json"), "w") as f:
json.dump(
{"special_tokens": _SPECIAL_TOKENS_CONFIG, "chat_template": _CHAT_TEMPLATE},
f,
)
_write_tokenizer_dir(
d,
chat_tokenizer,
{"special_tokens": _SPECIAL_TOKENS_CONFIG, "chat_template": _CHAT_TEMPLATE},
)
return d
+86
View File
@@ -0,0 +1,86 @@
"""Test data builders for preprocessing and dataset scenarios."""
from astrai.config.preprocess_config import (
InputConfig,
PipelineConfig,
ProcessingConfig,
)
CHAT_SECTIONS = [{"field": "messages", "action": "$role", "template": True}]
INSTRUCTION_SECTIONS = [
{"field": "prompt", "action": "mask", "add_special_tokens": True},
{"field": "response", "action": "train"},
]
TEXT_SECTIONS = [{"field": "text", "action": "train"}]
GRPO_RESPONSE_SECTIONS = [{"field": "responses", "action": "train"}]
def make_pipeline_config(sections, *, mask=None, preprocessing=None, sources=None):
"""Build a pipeline config with the common test defaults."""
return PipelineConfig(
input=InputConfig(sections=sections, sources=sources),
mask={} if mask is None else mask,
mask_default="mask",
preprocessing=preprocessing or ProcessingConfig(max_seq_len=2048),
)
def make_chat_config():
return make_pipeline_config(
CHAT_SECTIONS,
mask={"system": "mask", "user": "mask", "assistant": "train"},
)
def make_instruction_config():
return make_pipeline_config(
INSTRUCTION_SECTIONS,
mask={"prompt": "mask", "response": "train"},
)
def make_text_config():
return make_pipeline_config(
TEXT_SECTIONS,
preprocessing=ProcessingConfig(
max_seq_len=2048, min_chars=1, max_chars=2_000_000
),
)
def make_dpo_chat_config():
sources = {
name: {"sections": [{"field": name, "action": "$role", "template": True}]}
for name in ("chosen", "rejected")
}
return make_pipeline_config(
None,
mask={"user": "mask", "assistant": "train"},
sources=sources,
)
def make_grpo_config(*, template=True):
prompt_section = {"field": "prompt", "action": "mask"}
if template:
prompt_section["template"] = True
else:
prompt_section["add_special_tokens"] = True
sources = {
"prompts": {"sections": [prompt_section]},
"responses": {
"sections": GRPO_RESPONSE_SECTIONS,
"list_field": True,
"mask_key": "masks",
},
"rewards": {"sections": [{"field": "rewards", "action": "value"}]},
}
return make_pipeline_config(
None,
mask={"user": "mask", "assistant": "train"},
sources=sources,
)
def make_grpo_no_template_config():
return make_grpo_config(template=False)
+2 -2
View File
@@ -24,7 +24,7 @@ from astrai.serialization import (
load_bin,
save_bin,
)
from tests.data.conftest import make_grpo_no_template_config
from tests.data.factories import make_grpo_config
def _rand_seq(length, vocab=1000):
@@ -797,7 +797,7 @@ def test_grpo_builder_preserves_response_boundaries(base_test_env):
_save_test_tokenizer(base_test_env["test_dir"], tokenizer)
builder = SectionedMaskBuilder()
config = make_grpo_no_template_config()
config = make_grpo_config(template=False)
config.preprocessing.max_seq_len = 128
item = {
+13 -13
View File
@@ -12,10 +12,10 @@ from astrai.preprocessing.builder import (
SectionedMaskBuilder,
SingleOutputMaskBuilder,
)
from tests.data.conftest import (
_CHAT_SECTIONS,
_INSTRUCTION_SECTIONS,
_TEXT_SECTIONS,
from tests.data.factories import (
CHAT_SECTIONS,
INSTRUCTION_SECTIONS,
TEXT_SECTIONS,
make_chat_config,
make_dpo_chat_config,
make_grpo_config,
@@ -101,7 +101,7 @@ def test_chat_uniform_masking(
mask_rules, mask_default, expect_nonzero, chat_tokenizer, builder
):
config = PipelineConfig(
input=InputConfig(sections=_CHAT_SECTIONS),
input=InputConfig(sections=CHAT_SECTIONS),
mask=mask_rules,
mask_default=mask_default,
preprocessing=ProcessingConfig(max_seq_len=2048),
@@ -128,7 +128,7 @@ def test_chat_empty_messages(chat_tokenizer, builder):
def test_chat_domain_extraction(chat_tokenizer, builder):
config = PipelineConfig(
input=InputConfig(sections=_CHAT_SECTIONS),
input=InputConfig(sections=CHAT_SECTIONS),
mask={"assistant": "train"},
mask_default="mask",
preprocessing=ProcessingConfig(max_seq_len=2048),
@@ -147,7 +147,7 @@ def test_chat_domain_extraction(chat_tokenizer, builder):
def test_chat_truncation(chat_tokenizer, builder):
config = PipelineConfig(
input=InputConfig(sections=_CHAT_SECTIONS),
input=InputConfig(sections=CHAT_SECTIONS),
mask={"assistant": "train"},
mask_default="mask",
preprocessing=ProcessingConfig(max_seq_len=10),
@@ -237,7 +237,7 @@ def test_text_empty(test_tokenizer, builder):
def test_text_too_short(test_tokenizer, builder):
config = PipelineConfig(
input=InputConfig(sections=_TEXT_SECTIONS),
input=InputConfig(sections=TEXT_SECTIONS),
preprocessing=ProcessingConfig(min_chars=100),
)
assert builder.build({"text": "short"}, config, test_tokenizer) is None
@@ -245,7 +245,7 @@ def test_text_too_short(test_tokenizer, builder):
def test_text_truncation(test_tokenizer, builder):
config = PipelineConfig(
input=InputConfig(sections=_TEXT_SECTIONS),
input=InputConfig(sections=TEXT_SECTIONS),
preprocessing=ProcessingConfig(max_seq_len=3, min_chars=1),
)
item = {"text": "This is a very long text that should be truncated"}
@@ -255,7 +255,7 @@ def test_text_truncation(test_tokenizer, builder):
def test_sectioned_chat(chat_tokenizer, builder):
config = PipelineConfig(
input=InputConfig(sections=_CHAT_SECTIONS),
input=InputConfig(sections=CHAT_SECTIONS),
mask={"system": "mask", "user": "mask", "assistant": "train"},
mask_default="mask",
preprocessing=ProcessingConfig(max_seq_len=2048),
@@ -275,7 +275,7 @@ def test_sectioned_chat(chat_tokenizer, builder):
def test_sectioned_instruction(test_tokenizer, builder):
config = PipelineConfig(
input=InputConfig(sections=_INSTRUCTION_SECTIONS),
input=InputConfig(sections=INSTRUCTION_SECTIONS),
preprocessing=ProcessingConfig(max_seq_len=2048, min_chars=0),
)
item = {"prompt": "Q: Why?", "response": "A: Because."}
@@ -288,7 +288,7 @@ def test_sectioned_instruction(test_tokenizer, builder):
def test_sectioned_text(test_tokenizer, builder):
config = PipelineConfig(
input=InputConfig(sections=_TEXT_SECTIONS),
input=InputConfig(sections=TEXT_SECTIONS),
preprocessing=ProcessingConfig(max_seq_len=2048, min_chars=1),
)
item = {"text": "Hello world, this is a test."}
@@ -299,7 +299,7 @@ def test_sectioned_text(test_tokenizer, builder):
def test_sectioned_text_too_short(test_tokenizer, builder):
config = PipelineConfig(
input=InputConfig(sections=_TEXT_SECTIONS),
input=InputConfig(sections=TEXT_SECTIONS),
preprocessing=ProcessingConfig(max_seq_len=2048, min_chars=100),
)
assert builder.build({"text": "short"}, config, test_tokenizer) is None
+7 -7
View File
@@ -4,9 +4,9 @@ from astrai.config.preprocess_config import (
InputConfig,
PipelineConfig,
)
from tests.data.conftest import (
_INSTRUCTION_SECTIONS,
_TEXT_SECTIONS,
from tests.data.factories import (
INSTRUCTION_SECTIONS,
TEXT_SECTIONS,
make_dpo_chat_config,
)
@@ -43,26 +43,26 @@ def test_from_dict_flat():
def test_to_dict_roundtrip():
config = PipelineConfig(
input=InputConfig(sections=_INSTRUCTION_SECTIONS),
input=InputConfig(sections=INSTRUCTION_SECTIONS),
mask={"prompt": "mask", "response": "train"},
mask_default="mask",
)
d = config.to_dict()
config2 = PipelineConfig.from_dict(d)
assert config2.input.sections == _INSTRUCTION_SECTIONS
assert config2.input.sections == INSTRUCTION_SECTIONS
assert config2.mask == {"prompt": "mask", "response": "train"}
def test_to_file_from_file(temp_dir):
config = PipelineConfig(
input=InputConfig(sections=_TEXT_SECTIONS),
input=InputConfig(sections=TEXT_SECTIONS),
mask={"text": "train"},
mask_default="mask",
)
path = os.path.join(temp_dir, "config.json")
config.to_file(path)
loaded = PipelineConfig.from_file(path)
assert loaded.input.sections == _TEXT_SECTIONS
assert loaded.input.sections == TEXT_SECTIONS
assert loaded.mask == {"text": "train"}
+8 -8
View File
@@ -9,10 +9,10 @@ from astrai.config.preprocess_config import (
)
from astrai.preprocessing.packing import PackingStrategyFactory
from astrai.preprocessing.pipeline import Pipeline, filter_by_length
from tests.data.conftest import (
_CHAT_SECTIONS,
_INSTRUCTION_SECTIONS,
_TEXT_SECTIONS,
from tests.data.factories import (
CHAT_SECTIONS,
INSTRUCTION_SECTIONS,
TEXT_SECTIONS,
make_dpo_chat_config,
make_grpo_no_template_config,
)
@@ -54,7 +54,7 @@ def test_full_chat_pipeline(temp_dir, chat_tokenizer_dir):
)
config = PipelineConfig(
input=InputConfig(sections=_CHAT_SECTIONS),
input=InputConfig(sections=CHAT_SECTIONS),
mask={"system": "mask", "user": "mask", "assistant": "train"},
mask_default="mask",
preprocessing=ProcessingConfig(max_seq_len=2048),
@@ -97,7 +97,7 @@ def test_full_text_pipeline(temp_dir, tokenizer_dir):
)
config = PipelineConfig(
input=InputConfig(sections=_TEXT_SECTIONS),
input=InputConfig(sections=TEXT_SECTIONS),
preprocessing=ProcessingConfig(max_seq_len=2048, min_chars=10),
output=OutputConfig(storage_format="bin"),
)
@@ -138,7 +138,7 @@ def test_full_instruction_pipeline(temp_dir, tokenizer_dir):
)
config = PipelineConfig(
input=InputConfig(sections=_INSTRUCTION_SECTIONS),
input=InputConfig(sections=INSTRUCTION_SECTIONS),
mask={"prompt": "mask", "response": "train"},
mask_default="mask",
preprocessing=ProcessingConfig(max_seq_len=2048),
@@ -164,7 +164,7 @@ def test_dtype_override(temp_dir, tokenizer_dir):
f.write(json.dumps({"prompt": "Q", "response": "A"}) + "\n")
config = PipelineConfig(
input=InputConfig(sections=_INSTRUCTION_SECTIONS),
input=InputConfig(sections=INSTRUCTION_SECTIONS),
mask={"prompt": "mask", "response": "train"},
mask_default="mask",
preprocessing=ProcessingConfig(max_seq_len=2048),
-1
View File
@@ -5,7 +5,6 @@ import torch
from astrai.config.model_config import AutoRegressiveLMConfig
from astrai.model.transformer import AutoRegressiveLM
from tests.conftest import skip_no_kernel # noqa: F401 re-export for test modules
D = 64
CFG = dict(
+129 -3
View File
@@ -2,20 +2,33 @@
These tests do not require CUDA they only check that the active
backend is correctly set and restored.
Resolution precedence under test: explicit ``attn_backend(...)``
context > ``ASTR_BACKEND`` env override > implicit default. Training
calls (``fwd=None``, no KV cache) resolve by capability: the CUDA cache
kernels cannot run without a cache, so they fall back to flash (mask-free
calls only) and finally to torch SDPA.
"""
import importlib
import pytest
import torch
from astrai.extension import (
ATTN_BACKEND,
AttentionBackend,
AttentionBackendFactory,
CudaBackend,
FlashAttnBackend,
TorchNativeBackend,
attention,
attn_backend,
get_backend,
)
_attn_module = importlib.import_module("astrai.extension.backend.attention")
def test_default_backend_resolves_to_available():
"""Default backend is the first available in cuda > flash > torch order."""
@@ -23,6 +36,10 @@ def test_default_backend_resolves_to_available():
assert isinstance(backend, (CudaBackend, FlashAttnBackend, TorchNativeBackend))
def test_default_backend_is_cached_singleton():
assert get_backend() is get_backend()
def test_attn_backend_context_with_enum():
default = get_backend()
with attn_backend(ATTN_BACKEND.CUDA):
@@ -44,11 +61,120 @@ def test_backend_can_read_only_context_selection():
assert get_backend(use_default=False) is None
def test_environment_backend_overrides_context(monkeypatch):
def test_context_beats_environment_backend(monkeypatch):
"""An explicit attn_backend() context wins over ASTR_BACKEND."""
monkeypatch.setenv("ASTR_BACKEND", "torch_native")
with attn_backend("cuda"):
assert type(get_backend()).__name__ == "TorchNativeBackend"
assert type(get_backend(use_default=False)).__name__ == "TorchNativeBackend"
assert isinstance(get_backend(), CudaBackend)
assert isinstance(get_backend(use_default=False), CudaBackend)
def test_environment_backend_used_without_context(monkeypatch):
monkeypatch.setenv("ASTR_BACKEND", "torch_native")
assert isinstance(get_backend(), TorchNativeBackend)
assert isinstance(get_backend(use_default=False), TorchNativeBackend)
def test_explicit_backend_mismatch_raises(monkeypatch):
monkeypatch.delenv("ASTR_BACKEND", raising=False)
q = torch.zeros(1, 2, 4, 8, dtype=torch.bfloat16)
with pytest.raises(RuntimeError, match="Explicitly-set backend"):
with attn_backend("cuda"):
attention(q, q, q) # cuda + no KV cache -> cannot handle
def test_implicit_backend_falls_back_when_incapable(monkeypatch):
"""An implicit (env) backend that cannot run the call falls back."""
monkeypatch.setenv("ASTR_BACKEND", "cuda")
q = torch.zeros(1, 2, 4, 8, dtype=torch.float32) # fp32: cuda kernels can't
out = attention(q, q, q, fwd="prefill", is_causal=True)
assert out.shape == q.shape
def _flash_available(monkeypatch) -> None:
"""Pretend flash-attn is usable and rebuild the priority list."""
monkeypatch.setattr(_attn_module, "flash_attn_available", lambda: True)
_attn_module._priority_backends.cache_clear()
def test_training_falls_back_to_flash_before_torch_when_capable(monkeypatch):
"""Training (no cache) prefers flash over torch when flash can run the call."""
_flash_available(monkeypatch)
try:
prio = _attn_module._priority_backends()
names = [type(b).__name__ for b in prio]
assert "FlashAttnBackend" in names
assert names.index("FlashAttnBackend") < names.index("TorchNativeBackend")
q = torch.zeros(1, 2, 4, 8, dtype=torch.bfloat16)
# Mask-free training call resolves to flash, not torch.
resolved = next(b for b in prio if b.supports_call(q, None, None, False, None))
assert isinstance(resolved, FlashAttnBackend)
finally:
_attn_module._priority_backends.cache_clear()
def test_flash_dense_supports_only_mask_free_calls(monkeypatch):
"""FlashAttnBackend cannot apply custom masks in the dense path."""
_flash_available(monkeypatch)
flash = _attn_module._instance(_attn_module.FlashAttnBackend)
q = torch.zeros(1, 2, 4, 8, dtype=torch.bfloat16)
mask_4d = torch.zeros(1, 1, 2, 2, dtype=torch.bool)
assert flash.supports_call(q, None, None, False, None) is True
assert flash.supports_call(q, None, None, True, None) is True
assert flash.supports_call(q, None, mask_4d, False, None) is False
def test_flash_dense_rejects_custom_mask(monkeypatch):
"""A masked dense call must fail loudly, never silently ignore the mask."""
_flash_available(monkeypatch)
flash = _attn_module._instance(_attn_module.FlashAttnBackend)
q = torch.zeros(1, 2, 4, 8, dtype=torch.bfloat16)
mask_4d = torch.zeros(1, 1, 2, 2, dtype=torch.bool)
with pytest.raises(ValueError, match="custom attention mask"):
flash._forward_dense(q, q, q, attn_mask=mask_4d, is_causal=False)
def test_backend_resolution_returns_shared_singletons():
with attn_backend("cuda") as first:
pass
with attn_backend("cuda") as second:
assert first is second
class _DummyBackend(AttentionBackend):
"""Minimal backend used only to prove capability is polymorphic."""
@classmethod
def available(cls) -> bool:
return True
def supports_call(self, q, kv_cache, attn_mask, is_causal, fwd) -> bool:
return True
def fwd_decode(
self, q, k, v, kv_cache=None, layer_id=0, attn_mask=None, is_causal=False
):
return q
def fwd_prefill(
self, q, k, v, kv_cache=None, layer_id=0, attn_mask=None, is_causal=False
):
return q
def test_custom_backend_usable_without_touching_resolution():
"""A third-party backend plugs in via context or explicit param."""
custom = _DummyBackend()
q = torch.zeros(1, 2, 4, 8)
with attn_backend(custom):
assert get_backend() is custom
out = attention(q, q, q, backend=custom)
assert out is q
def test_attention_backend_factory_lists_builtin_backends():
+14 -7
View File
@@ -12,7 +12,8 @@ from astrai.inference.cache import PagePool, TaskCacheManager
from astrai.inference.runtime.graph import CudaGraphContext
from astrai.inference.scheduler import InferenceScheduler
from astrai.inference.workspace import InferenceWorkspace
from tests.extension.conftest import D, skip_no_kernel
from tests.conftest import skip_no_kernel
from tests.extension.conftest import D
from tests.helpers import FakeTokenizer
@@ -33,10 +34,13 @@ def _ws(pool: PagePool) -> InferenceWorkspace:
@skip_no_kernel
def test_training_forward_matches_torch(cuda_model):
"""Training forward (kv_cache=None) uses torch-native SDPA.
"""Training forward (kv_cache=None) resolves to a capable dense backend.
CudaBackend does not support training (requires kv_cache).
Torch-native backend must match default (which falls back to torch).
CudaBackend cannot run training (requires a KV cache), so the default
falls back by capability flash when it can handle the call
(mask-free/causal), otherwise torch SDPA. The default path must not
raise and must produce finite logits; explicitly selected torch SDPA
must be deterministic across runs.
"""
model, _ = cuda_model
@@ -47,12 +51,15 @@ def test_training_forward_matches_torch(cuda_model):
with attn_backend(ATTN_BACKEND.TORCH_NATIVE):
with torch.no_grad():
out_torch = model(input_ids)
out_torch_a = model(input_ids)
with torch.no_grad():
out_torch_b = model(input_ids)
assert out_default["logits"].shape == out_torch_a["logits"].shape
assert torch.isfinite(out_default["logits"]).all()
torch.testing.assert_close(
out_torch["logits"], out_default["logits"], atol=1e-6, rtol=1e-6
out_torch_a["logits"], out_torch_b["logits"], atol=0, rtol=0
)
assert out_default["logits"].shape[0] == 2
@skip_no_kernel
+301 -82
View File
@@ -1,16 +1,29 @@
"""Fused BF16-boundary FP8 MMA kernel tests."""
"""FP8 primitives: kernel-level (CUDA) and policy-level (CPU-verifiable) tests.
The kernel-level tests exercise the pure FP8 path (quantize_bf16 + mm_fp8 for
the forward GEMM, quantize + pre-quantized GEMMs for the backward); the
policy-level tests (recipes, autocast context, per-tensor meta, CPU fallbacks
of the custom ops) run without a GPU.
"""
import pytest
import torch
from astrai.extension.loader import get_module, is_available
pytestmark = pytest.mark.skipif(
not torch.cuda.is_available()
or torch.cuda.get_device_capability() < (8, 9)
or not is_available("fp8_mm"),
reason="fused FP8 MMA requires a built kernel and compute capability 8.9+",
from astrai.extension.fp8 import (
DelayedScaling,
DynamicScaling,
FP8Format,
FP8TensorMeta,
fp8_autocast,
fp8_state,
)
from astrai.extension.ops.fp8 import (
linear_backward_fp8,
linear_forward_fp8,
mm_fp8,
quantize_bf16,
)
from tests.conftest import skip_no_fp8
def _scale(tensor):
@@ -21,28 +34,60 @@ def _quantize(tensor, scale):
return (tensor.float() / scale).to(torch.float8_e4m3fn).float()
# --------------------------------------------------------------------------
# Kernel-level (CUDA)
# --------------------------------------------------------------------------
@skip_no_fp8
@pytest.mark.parametrize(
("m", "n", "k"),
[(16, 8, 32), (17, 9, 33), (31, 15, 64), (32, 48, 96)],
)
def test_fused_fp8_mma_matches_explicit_quantization(m, n, k):
def test_fp8_mm_matches_explicit_quantization(m, n, k):
torch.manual_seed(m + n + k)
a = torch.randn(m, k, device="cuda", dtype=torch.bfloat16)
b = torch.randn(n, k, device="cuda", dtype=torch.bfloat16)
b = torch.randn(k, n, device="cuda", dtype=torch.bfloat16)
scale_a = _scale(a)
scale_b = _scale(b)
out = get_module("fp8_mm").fp8_mm(a, b, scale_a, scale_b)
expected = (
_quantize(a, scale_a) @ _quantize(b, scale_b).t() * scale_a * scale_b
).to(torch.bfloat16)
a8, _ = quantize_bf16(a, scale_a, "e4m3")
b8, _ = quantize_bf16(b, scale_b, "e4m3")
out = mm_fp8(a8, b8, scale_a, scale_b)
expected = (_quantize(a, scale_a) @ _quantize(b, scale_b) * scale_a * scale_b).to(
torch.bfloat16
)
assert out.dtype == torch.bfloat16
assert out.shape == (m, n)
torch.testing.assert_close(out, expected, atol=0.125, rtol=0.01)
def test_fused_fp8_linear_forward_and_backward():
@skip_no_fp8
def test_quantize_bf16_returns_amax():
"""quantize_bf16 returns (x8, amax); amax tracks the *raw* values and the
caller never clears it (zero-initialized inside the kernel entry)."""
torch.manual_seed(3)
x = torch.randn(64, 128, device="cuda", dtype=torch.bfloat16)
scale = torch.tensor([0.5], device="cuda")
x8, amax = quantize_bf16(x, scale, "e4m3")
assert x8.dtype == torch.float8_e4m3fn
assert x8.shape == x.shape
assert amax.shape == (1,)
torch.testing.assert_close(amax, x.abs().amax().float().reshape(1))
ref = (x.float() / 0.5).to(torch.float8_e4m3fn)
assert torch.equal(x8, ref)
@skip_no_fp8
def test_quantize_bf16_e5m2_format():
x = torch.randn(32, 64, device="cuda", dtype=torch.bfloat16)
x8, amax = quantize_bf16(x, torch.tensor([0.1], device="cuda"), "e5m2")
assert x8.dtype == torch.float8_e5m2
torch.testing.assert_close(amax, x.abs().amax().float().reshape(1))
@skip_no_fp8
def test_fp8_linear_forward_and_backward():
torch.manual_seed(7)
m, n, k = 19, 13, 37
x = torch.randn(m, k, device="cuda", dtype=torch.bfloat16)
@@ -50,34 +95,10 @@ def test_fused_fp8_linear_forward_and_backward():
grad = torch.randn(m, n, device="cuda", dtype=torch.bfloat16)
bias = torch.randn(n, device="cuda", dtype=torch.bfloat16)
scale_x, scale_w, scale_g = _scale(x), _scale(weight), _scale(grad)
amax_x = torch.empty(1, device="cuda", dtype=torch.float32)
amax_w = torch.empty(1, device="cuda", dtype=torch.float32)
amax_g = torch.empty(1, device="cuda", dtype=torch.float32)
module = get_module("fp8_mm")
out = module.fp8_linear_forward_scaled(
x,
weight,
bias,
scale_x,
scale_w,
scale_x.reciprocal(),
scale_w.reciprocal(),
amax_x,
amax_w,
)
grad_x, grad_w, grad_b = module.fp8_linear_backward_scaled(
grad,
x,
weight,
[1, 1, 1],
scale_g,
scale_w,
scale_x,
scale_g.reciprocal(),
scale_w.reciprocal(),
scale_x.reciprocal(),
amax_g,
out, amax_x, amax_w = linear_forward_fp8(x, weight, bias, scale_x, scale_w)
grad_x, grad_w, grad_b, amax_g = linear_backward_fp8(
grad, x, weight, [1, 1, 1], scale_g, scale_w, scale_x, "e4m3"
)
qx = _quantize(x, scale_x)
@@ -96,61 +117,259 @@ def test_fused_fp8_linear_forward_and_backward():
torch.testing.assert_close(amax_g, grad.abs().amax().float().reshape(1))
def test_fp8_mm_prequant_matches_scaled_mm():
@skip_no_fp8
def test_linear_backward_e5m2_gradients():
"""Hybrid backward: gradient GEMMs run in E5M2 (larger dynamic range)."""
torch.manual_seed(5)
m, n, k = 32, 16, 64
x = torch.randn(m, k, device="cuda", dtype=torch.bfloat16) * 3.0
weight = torch.randn(n, k, device="cuda", dtype=torch.bfloat16)
grad = torch.randn(m, n, device="cuda", dtype=torch.bfloat16) * 10.0
sg = _scale(grad) * 0.5
sw = _scale(weight)
sx = _scale(x)
grad_x, grad_w, grad_b, amax_g = linear_backward_fp8(
grad, x, weight, [1, 1, 1], sg, sw, sx, "e5m2"
)
def q5(t, s):
return (t.float() / s).to(torch.float8_e5m2).float()
qg = q5(grad, sg)
qw = q5(weight, sw)
qx = q5(x, sx)
expected_grad_x = (qg @ qw * sg * sw).to(torch.bfloat16)
expected_grad_w = (qg.t() @ qx * sg * sx).to(torch.bfloat16)
torch.testing.assert_close(grad_x, expected_grad_x, atol=0.5, rtol=0.05)
torch.testing.assert_close(grad_w, expected_grad_w, atol=0.5, rtol=0.05)
torch.testing.assert_close(amax_g, grad.abs().amax().float().reshape(1))
@skip_no_fp8
def test_fp8_linear_static_fp8_weight_and_bias():
"""Static fp8 inference: pre-quantized w8/b8 + their scales take the GEMM
directly (no weight quantize, amax_w = 0); the bias is fused in the
epilogue (bf16 and fp8 bias share the fused path)."""
torch.manual_seed(9)
m, n, k = 67, 45, 129
x = torch.randn(m, k, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(n, k, device="cuda", dtype=torch.bfloat16) * 0.5
bias = torch.randn(n, device="cuda", dtype=torch.bfloat16) * 0.5
sx, sw, sb = _scale(x), _scale(weight), _scale(bias)
w8, _ = quantize_bf16(weight, sw, "e4m3")
b8, _ = quantize_bf16(bias, sb, "e4m3")
out, amax_x, amax_w = linear_forward_fp8(x, w8, b8, sx, sw, "e4m3", sb)
qx = _quantize(x, sx)
qw = _quantize(weight, sw)
qb = _quantize(bias, sb)
expected = (qx @ qw.t() * sx * sw + qb * sb).to(torch.bfloat16)
torch.testing.assert_close(out, expected, atol=0.125, rtol=0.01)
torch.testing.assert_close(amax_x, x.abs().amax().float().reshape(1))
assert amax_w.item() == 0.0 # nothing measured on the static path
# bf16 bias stays bf16 on the same fused-epilogue path
out_bf16bias, _, _ = linear_forward_fp8(x, w8, bias, sx, sw, "e4m3")
expected_b = (qx @ qw.t() * sx * sw + bias.float()).to(torch.bfloat16)
torch.testing.assert_close(out_bf16bias, expected_b, atol=0.125, rtol=0.01)
@skip_no_fp8
def test_fp8_linear_backward_outside_autocast():
"""aten::linear records an fp8 autograd node inside fp8_autocast; the
backward runs fp8 kernels even after the context exits (loss.backward()
placement is free), instead of falling back to bf16 mm."""
import torch.nn.functional as F
import astrai.extension.fp8 as f8mod
torch.manual_seed(5)
x = torch.randn(64, 128, device="cuda", dtype=torch.bfloat16, requires_grad=True)
weight = torch.randn(
96, 128, device="cuda", dtype=torch.bfloat16, requires_grad=True
)
bias = torch.randn(96, device="cuda", dtype=torch.bfloat16, requires_grad=True)
xr, wr, br = (t.detach().clone().requires_grad_() for t in (x, weight, bias))
calls = {"bwd": 0}
orig = f8mod.linear_backward_fp8
def spy(g, xx, ww, masks, sg, sw, sx, fmt="e5m2"):
calls["bwd"] += 1
return orig(g, xx, ww, masks, sg, sw, sx, fmt)
f8mod.linear_backward_fp8 = spy
try:
with fp8_autocast(enabled=True):
out = F.linear(x, weight, bias)
assert type(out.grad_fn).__name__ == "_LinearFp8Backward"
out.float().pow(2).sum().backward() # outside the autocast region
finally:
f8mod.linear_backward_fp8 = orig
f8mod.fp8_state().reset()
assert calls["bwd"] == 1 # fp8 kernels, not the bf16 fallback
ref = F.linear(xr, wr, br)
ref.float().pow(2).sum().backward()
# E5M2 backward quantization noise: compare directions/norms (the
# torchao/TE style) rather than elementwise against the bf16 reference.
def _direction(a, b):
cos = torch.nn.functional.cosine_similarity(
a.float().flatten(), b.float().flatten(), dim=0
)
return cos > 0.99 and 0.9 < a.float().norm() / b.float().norm() < 1.1
assert _direction(x.grad, xr.grad)
assert _direction(weight.grad, wr.grad)
assert _direction(bias.grad, br.grad)
@skip_no_fp8
def test_mm_fp8_matches_scaled_mm():
torch.manual_seed(11)
m, n, k = 512, 4096, 4096
a = torch.randn(m, k, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(n, k, device="cuda", dtype=torch.bfloat16)
a8 = a.to(torch.float8_e4m3fn)
w8 = weight.to(torch.float8_e4m3fn)
scale = torch.tensor([2.5], device="cuda")
out = get_module("fp8_mm").fp8_mm_prequant(a8, w8, scale)
# Reference via fp64: FP8 quantization error is dominated by the 3-bit
# mantissa, so the tolerance must track the input quantization scale.
ref = (a8.float().double() @ w8.float().double().t() * 2.5).to(torch.bfloat16)
b = torch.randn(k, n, device="cuda", dtype=torch.bfloat16)
sa = torch.tensor([2.5], device="cuda")
sb = torch.tensor([1.5], device="cuda")
a8, _ = quantize_bf16(a, sa, "e4m3")
b8, _ = quantize_bf16(b, sb, "e4m3")
out = mm_fp8(a8, b8, sa, sb)
assert out.dtype == torch.bfloat16
assert out.shape == (m, n)
ref = (a8.float().double() @ b8.float().double() * 2.5 * 1.5).to(torch.bfloat16)
torch.testing.assert_close(out, ref, atol=6.0, rtol=0.05)
# Cross-check against torch's native FP8 GEMM on identical inputs.
try:
torch._scaled_mm(
a8,
w8.t(),
torch.full((m, 1), 2.5, device="cuda"),
torch.ones((1, n), device="cuda"),
out_dtype=torch.bfloat16,
)
torch._scaled_mm(a8, b8, sa, sb, out_dtype=torch.bfloat16)
except (RuntimeError, NotImplementedError):
return
torch.testing.assert_close(
out,
torch._scaled_mm(
a8,
w8.t(),
torch.full((m, 1), 2.5, device="cuda"),
torch.ones((1, n), device="cuda"),
out_dtype=torch.bfloat16,
),
torch._scaled_mm(a8, b8, sa, sb, out_dtype=torch.bfloat16),
atol=2.0,
rtol=0.01,
)
def test_fp8_mm_prequant_fp8_output():
torch.manual_seed(13)
m, n, k = 512, 4096, 4096
@skip_no_fp8
def test_mm_fp8_fp8_output():
"""mm_fp8 with out_dtype='e4m3' produces an FP8 output (layer-to-layer)."""
torch.manual_seed(12)
m, n, k = 256, 128, 64
a = torch.randn(m, k, device="cuda", dtype=torch.bfloat16)
weight = torch.randn(n, k, device="cuda", dtype=torch.bfloat16)
a8 = a.to(torch.float8_e4m3fn)
w8 = weight.to(torch.float8_e4m3fn)
scale = torch.tensor([2.5], device="cuda")
out_scale = torch.tensor([0.1], device="cuda")
b = torch.randn(k, n, device="cuda", dtype=torch.bfloat16)
sa = torch.tensor([2.0], device="cuda")
sb = torch.tensor([1.0], device="cuda")
os_ = torch.tensor([0.5], device="cuda")
a8, _ = quantize_bf16(a, sa, "e4m3")
b8, _ = quantize_bf16(b, sb, "e4m3")
out8 = mm_fp8(a8, b8, sa, sb, out_dtype="e4m3", out_scale=os_)
assert out8.dtype == torch.float8_e4m3fn
assert out8.shape == (m, n)
out = get_module("fp8_mm").fp8_mm_prequant_fp8(a8, w8, scale, out_scale)
assert out.dtype == torch.float8_e4m3fn
assert out.shape == (m, n)
ref = (a8.float().double() @ w8.float().double().t() * 2.5 * 0.1).to(torch.bfloat16)
torch.testing.assert_close(out.float().to(torch.bfloat16), ref, atol=1.0, rtol=0.05)
ref = (a8.float().double() @ b8.float().double() * 2.0 * 1.0 * 0.5).to(
torch.bfloat16
)
torch.testing.assert_close(
out8.float().to(torch.bfloat16), ref, atol=6.0, rtol=0.05
)
# --------------------------------------------------------------------------
# Policy-level (CPU-verifiable)
# --------------------------------------------------------------------------
def test_recipe_scale_from_history():
"""Delayed: max over the window + margin; dynamic: current amax."""
hist = torch.tensor([1.0, 2.0, 0.5])
d = DelayedScaling(history_len=3, margin=0)
assert torch.allclose(d.scale_from_history(hist, "e4m3"), torch.tensor(2.0 / 448.0))
d_m = DelayedScaling(history_len=3, margin=2)
assert torch.allclose(
d_m.scale_from_history(hist, "e4m3"), torch.tensor(2.0 / 448.0 / 4.0)
)
dyn = DynamicScaling()
amax = torch.tensor([0.25])
assert torch.allclose(
dyn.scale_from_history(amax, "e4m3"), torch.tensor(0.25 / 448.0)
)
assert torch.allclose(
dyn.scale_from_history(amax, "e5m2"), torch.tensor(0.25 / 57344.0)
)
def test_fp8_format_enum():
assert FP8Format.HYBRID.fwd() == "e4m3"
assert FP8Format.HYBRID.bwd() == "e5m2"
assert FP8Format.E4M3.fwd() == FP8Format.E4M3.bwd() == "e4m3"
assert FP8Format.E5M2.fwd() == FP8Format.E5M2.bwd() == "e5m2"
def test_fp8_autocast_context():
"""fp8_autocast sets and restores recipe + format on the global state."""
state = fp8_state()
prev = (state.enabled, state.recipe, state.fp8_format)
try:
with fp8_autocast(enabled=True, fp8_format="hybrid", update_interval=8):
assert state.enabled
assert isinstance(state.recipe, DelayedScaling)
assert state.recipe.history_len == 8
assert state.fp8_format is FP8Format.HYBRID
with fp8_autocast(enabled=True, recipe=DynamicScaling(), fp8_format="e4m3"):
assert isinstance(state.recipe, DynamicScaling)
assert state.fp8_format is FP8Format.E4M3
assert state.fp8_format is FP8Format.HYBRID # restored on exit
assert not state.enabled
finally:
state.enabled, state.recipe, state.fp8_format = prev
def test_fp8_tensor_meta_delayed_update():
"""Meta seeds from data and refreshes the scale from the amax ring."""
meta = FP8TensorMeta(torch.device("cpu"), DelayedScaling(history_len=4, margin=0))
w = torch.randn(8, 8)
meta.w.seed(w, "e4m3")
assert meta.w.initialized
torch.testing.assert_close(meta.w.scale, (w.abs().amax() / 448.0).reshape(1))
meta.w.update(torch.tensor([4.0]), "e4m3")
torch.testing.assert_close(meta.w.scale, torch.tensor(4.0 / 448.0).reshape(1))
def test_quantize_bf16_cpu_fallback():
"""CPU fallback of the quantize primitive (scale semantics + amax)."""
x = torch.randn(16, 32, dtype=torch.bfloat16)
scale = torch.tensor([0.5])
x8, amax = quantize_bf16(x, scale, "e4m3")
assert x8.dtype == torch.float8_e4m3fn
ref = (x.float() / 0.5).to(torch.float8_e4m3fn)
assert torch.equal(x8, ref)
torch.testing.assert_close(amax, x.abs().amax().float().reshape(1))
def test_mm_fp8_cpu_fallback():
a8 = torch.tensor([[1.0, 2.0]], dtype=torch.float8_e4m3fn)
b8 = torch.tensor([[3.0], [4.0]], dtype=torch.float8_e4m3fn)
sa = torch.tensor([2.0])
sb = torch.tensor([0.5])
out = mm_fp8(a8, b8, sa, sb)
ref = (a8.float() @ b8.float() * 2.0 * 0.5).to(torch.bfloat16)
torch.testing.assert_close(out, ref)
def test_mm_fp8_fp8_output_cpu():
"""CPU fallback with an FP8 output (out_dtype='e4m3' + out_scale)."""
a8 = torch.tensor([[1.0, 2.0]], dtype=torch.float8_e4m3fn)
b8 = torch.tensor([[3.0], [4.0]], dtype=torch.float8_e4m3fn)
sa = torch.tensor([2.0])
sb = torch.tensor([0.5])
os_ = torch.tensor([0.25])
out8 = mm_fp8(a8, b8, sa, sb, out_dtype="e4m3", out_scale=os_)
assert out8.dtype == torch.float8_e4m3fn
ref = (a8.float() @ b8.float() * 2.0 * 0.5 * 0.25).to(torch.float8_e4m3fn)
assert torch.equal(out8, ref)
+2 -1
View File
@@ -3,7 +3,8 @@
import torch
from astrai.extension.ops.attention import attn_prefill
from tests.extension.conftest import D, skip_no_kernel
from tests.conftest import skip_no_kernel
from tests.extension.conftest import D
@skip_no_kernel
+40
View File
@@ -4,10 +4,12 @@ import json
import os
import torch
from tokenizers import Tokenizer, models, pre_tokenizers, trainers
from torch.utils.data import Dataset
from astrai.config.model_config import AutoRegressiveLMConfig
from astrai.model.transformer import AutoRegressiveLM
from astrai.tokenize import AutoTokenizer
TINY_CONFIG = dict(
vocab_size=1000,
@@ -57,6 +59,44 @@ def make_model(device, **cfg_overrides):
return model, cfg
def build_test_tokenizer(
vocab_size: int = 1000,
*,
special_tokens=("<unk>", "<pad>"),
special_token_map=None,
add_prefix_space: bool = True,
train_data=None,
chat_template: str | None = None,
) -> AutoTokenizer:
"""Build a lightweight BPE ``AutoTokenizer`` for tests.
``special_token_map`` defaults to ``{"unk_token", "pad_token"}``
pointing at the first two special tokens.
"""
tokenizer = Tokenizer(models.BPE())
tokenizer.pre_tokenizer = pre_tokenizers.ByteLevel(
add_prefix_space=add_prefix_space
)
trainer = trainers.BpeTrainer(
vocab_size=vocab_size,
min_frequency=1,
special_tokens=list(special_tokens),
)
tokenizer.train_from_iterator(
train_data if train_data is not None else [chr(i) for i in range(256)],
trainer,
)
auto_tokenizer = AutoTokenizer()
auto_tokenizer._tokenizer = tokenizer
auto_tokenizer._special_token_map = special_token_map or {
"unk_token": special_tokens[0],
"pad_token": special_tokens[1],
}
if chat_template is not None:
auto_tokenizer.set_chat_template(chat_template)
return auto_tokenizer
def make_frozen(model, device):
"""Create a frozen, eval-mode copy of *model* with identical weights."""
cfg = make_rollout_config()
+7
View File
@@ -8,6 +8,13 @@ from fastapi.testclient import TestClient
from astrai.inference import get_app
@pytest.fixture(autouse=True)
def _cleanup_app_engine():
"""Reset the lazy FastAPI singleton engine after each inference test."""
yield
get_app().state.engine = None
@pytest.fixture
def client():
"""Provide a test client for the FastAPI app."""
+1 -18
View File
@@ -2,7 +2,7 @@
import torch
from astrai.inference import (
from astrai.inference.cache import (
Allocator,
KVStorage,
PagePool,
@@ -201,23 +201,6 @@ def test_req_to_token_pool_write():
# ---- KVStorage ----
def test_kv_storage_set_and_get():
storage = KVStorage(
size=16,
n_layers=2,
n_kv_heads=4,
head_dim=8,
device=torch.device("cpu"),
dtype=torch.float32,
)
loc = torch.tensor([[0, 1]], dtype=torch.long)
k = torch.randn(1, 2, 4, 8)
v = torch.randn(1, 2, 4, 8)
storage.set_kv_buffer(0, loc, k, v)
assert torch.allclose(storage.get_key_buffer(0)[loc], k)
assert torch.allclose(storage.get_value_buffer(0)[loc], v)
def test_kv_storage_buffer_shape():
storage = KVStorage(
size=32,
+48 -25
View File
@@ -1,5 +1,6 @@
"""Unit tests for GenerateResult accumulator and InferenceEngine.generate()."""
import asyncio
import threading
from unittest.mock import MagicMock, patch
@@ -8,6 +9,17 @@ from astrai.inference import STOP
from astrai.inference.engine import GenerateResult, InferenceEngine
def _make_engine_mocks(decode=None):
"""Build the standard mock model/tokenizer pair used by engine tests."""
mock_model = MagicMock()
mock_tokenizer = MagicMock()
mock_tokenizer.encode.return_value = [1, 2, 3]
mock_tokenizer.stop_ids = [0]
if decode is not None:
mock_tokenizer.decode.return_value = decode
return mock_model, mock_tokenizer
def test_result_append_single():
r = GenerateResult(count=1)
r.append("hello", 0)
@@ -102,11 +114,7 @@ def test_result_get_results():
def test_engine_generate_non_streaming_single():
mock_model = MagicMock()
mock_tokenizer = MagicMock()
mock_tokenizer.encode.return_value = [1, 2, 3]
mock_tokenizer.decode.return_value = "response"
mock_tokenizer.stop_ids = [0]
mock_model, mock_tokenizer = _make_engine_mocks(decode="response")
with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
instance = MockSched.return_value
@@ -125,11 +133,7 @@ def test_engine_generate_non_streaming_single():
def test_engine_generate_streaming_yields_tokens():
mock_model = MagicMock()
mock_tokenizer = MagicMock()
mock_tokenizer.encode.return_value = [1, 2, 3]
mock_tokenizer.decode.return_value = "tok"
mock_tokenizer.stop_ids = [0]
mock_model, mock_tokenizer = _make_engine_mocks(decode="tok")
callbacks_saved = []
@@ -153,12 +157,37 @@ def test_engine_generate_streaming_yields_tokens():
assert tokens == ["t1", "t2"]
def test_engine_generate_async_yields_tokens_until_stop():
mock_model, mock_tokenizer = _make_engine_mocks(decode="tok")
callbacks_saved = []
def capture_cb(prompt, **kw):
callbacks_saved.append(kw.get("stream_callback"))
with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
instance = MockSched.return_value
instance.add_task.side_effect = capture_cb
instance.remove_task.return_value = []
eng = InferenceEngine(mock_model, mock_tokenizer, max_batch_size=1)
agen = eng.generate_async("hello")
async def collect():
out = []
async for token in agen:
out.append(token)
return out
cb = callbacks_saved[0]
cb("t1")
cb("t2")
cb(STOP)
assert asyncio.run(collect()) == ["t1", "t2"]
def test_engine_generate_non_streaming_batch():
mock_model = MagicMock()
mock_tokenizer = MagicMock()
mock_tokenizer.encode.return_value = [1, 2, 3]
mock_tokenizer.decode.return_value = "r"
mock_tokenizer.stop_ids = [0]
mock_model, mock_tokenizer = _make_engine_mocks(decode="r")
with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
instance = MockSched.return_value
@@ -177,10 +206,7 @@ def test_engine_generate_non_streaming_batch():
def test_engine_generate_zero_max_tokens_returns_empty():
mock_model = MagicMock()
mock_tokenizer = MagicMock()
mock_tokenizer.encode.return_value = [1, 2, 3]
mock_tokenizer.stop_ids = [0]
mock_model, mock_tokenizer = _make_engine_mocks()
with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
instance = MockSched.return_value
@@ -192,8 +218,7 @@ def test_engine_generate_zero_max_tokens_returns_empty():
def test_engine_generate_zero_max_tokens_stream_is_empty():
mock_model = MagicMock()
mock_tokenizer = MagicMock()
mock_model, mock_tokenizer = _make_engine_mocks()
with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
instance = MockSched.return_value
@@ -203,8 +228,7 @@ def test_engine_generate_zero_max_tokens_stream_is_empty():
def test_engine_passes_backend_to_scheduler():
mock_model = MagicMock()
mock_tokenizer = MagicMock()
mock_model, mock_tokenizer = _make_engine_mocks()
with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
InferenceEngine(
@@ -218,8 +242,7 @@ def test_engine_passes_backend_to_scheduler():
def test_generate_captures_calling_backend_context():
mock_model = MagicMock()
mock_tokenizer = MagicMock()
mock_model, mock_tokenizer = _make_engine_mocks()
captured = []
with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
+227 -205
View File
@@ -3,8 +3,6 @@
import json
from unittest.mock import MagicMock
import pytest
from astrai.inference.network.anthropic import AnthropicResponseBuilder
from astrai.inference.network.openai import OpenAIResponseBuilder
from astrai.inference.network.protocol import GenContext, StopChecker, StopInfo
@@ -34,223 +32,247 @@ def _sse_payloads(events):
return payloads
class TestStopChecker:
def test_check_finds_match(self):
sc = StopChecker(["stop", "end"])
assert sc.check("hello stop world") == "stop"
def test_check_returns_none_when_no_match(self):
sc = StopChecker(["stop"])
assert sc.check("hello world") is None
def test_check_empty_sequences(self):
sc = StopChecker([])
assert sc.check("hello") is None
def _make_openai_builder():
builder = OpenAIResponseBuilder()
req = MagicMock()
req.messages = [MagicMock(role="user", content="Hello")]
req.stop = None
req.model = "astrai"
engine = MagicMock()
engine.tokenizer.apply_chat_template.return_value = "Hello"
builder.prepare(req, engine)
return builder
class TestGenContext:
def test_defaults(self):
ctx = GenContext(resp_id="a", created=1, model="m", prompt_tokens=10)
assert ctx.completion_tokens == 0
def test_fields_mutable(self):
ctx = GenContext(resp_id="a", created=1, model="m", prompt_tokens=10)
ctx.completion_tokens = 42
assert ctx.completion_tokens == 42
def _make_anthropic_builder():
builder = AnthropicResponseBuilder()
req = MagicMock()
req.messages = [MagicMock(role="user", content="Hello")]
req.model = "claude"
req.system = None
engine = MagicMock()
engine.tokenizer.apply_chat_template.return_value = "Hello"
builder.prepare(req, engine)
return builder
class TestStopInfo:
def test_defaults(self):
s = StopInfo()
assert s.matched is None
assert s.body == ""
assert s.yielded == ""
def test_with_values(self):
s = StopInfo(matched="stop", body="hello stop", yielded="hello ")
assert s.matched == "stop"
assert s.body == "hello stop"
assert s.yielded == "hello "
def test_check_finds_match():
sc = StopChecker(["stop", "end"])
assert sc.check("hello stop world") == "stop"
class TestOpenAIResponseBuilder:
@pytest.fixture
def builder(self):
builder = OpenAIResponseBuilder()
req = MagicMock()
req.messages = [MagicMock(role="user", content="Hello")]
req.stop = None
req.model = "astrai"
engine = MagicMock()
engine.tokenizer.apply_chat_template.return_value = "Hello"
builder.prepare(req, engine)
return builder
def test_prepare_returns_prompt_ctx_stops(self, builder):
req = MagicMock()
req.messages = [MagicMock(role="user", content="Hi")]
req.stop = ["END"]
req.model = "gpt"
engine = MagicMock()
engine.tokenizer.apply_chat_template.return_value = "Hi"
prompt, ctx, stops = builder.prepare(req, engine)
assert prompt == "Hi"
assert ctx.model == "gpt"
assert ctx.prompt_tokens == 0
assert stops == ["END"]
def test_prepare_no_stop_returns_empty_list(self, builder):
req = MagicMock()
req.messages = []
req.stop = None
req.model = "x"
engine = MagicMock()
engine.tokenizer.apply_chat_template.return_value = ""
_, _, stops = builder.prepare(req, engine)
assert stops == []
def test_format_stream_start(self, builder):
ctx = _make_ctx()
events = builder.format_stream_start(ctx)
payloads = _sse_payloads(events)
assert len(payloads) == 1
p = payloads[0]
assert p["object"] == "chat.completion.chunk"
assert p["choices"][0]["delta"]["role"] == "assistant"
assert p["choices"][0]["finish_reason"] is None
def test_format_chunk(self, builder):
events = builder.format_chunk("hello", body="hello")
payload = json.loads(events[0].split("data: ", 1)[1])
assert payload["choices"][0]["delta"]["content"] == "hello"
assert payload["choices"][0]["finish_reason"] is None
def test_format_stream_end(self, builder):
ctx = _make_ctx(completion_tokens=5)
stop = StopInfo(matched="stop")
events = builder.format_stream_end(ctx, stop)
payloads = _sse_payloads(events)
finish = payloads[0]
assert finish["choices"][0]["finish_reason"] == "stop"
usage = payloads[1]
assert usage["completion_tokens"] == 5
assert usage["total_tokens"] == 15
def test_format_response(self, builder):
ctx = _make_ctx()
stop = StopInfo()
resp = builder.format_response(ctx, "hello", stop)
assert resp["object"] == "chat.completion"
assert resp["choices"][0]["message"]["content"] == "hello"
assert resp["usage"]["prompt_tokens"] == 10
def test_check_returns_none_when_no_match():
sc = StopChecker(["stop"])
assert sc.check("hello world") is None
class TestAnthropicResponseBuilder:
@pytest.fixture
def builder(self):
builder = AnthropicResponseBuilder()
req = MagicMock()
req.messages = [MagicMock(role="user", content="Hello")]
req.model = "claude"
engine = MagicMock()
engine.tokenizer.apply_chat_template.return_value = "Hello"
req.system = None
builder.prepare(req, engine)
return builder
def test_check_empty_sequences():
sc = StopChecker([])
assert sc.check("hello") is None
def test_prepare_messages(self, builder):
req = MagicMock()
req.messages = [MagicMock(role="user", content="Hi")]
req.model = "claude"
req.system = None
req.stop_sequences = None
engine = MagicMock()
engine.tokenizer.apply_chat_template.return_value = "Hi"
prompt, ctx, stops = builder.prepare(req, engine)
assert prompt == "Hi"
assert stops == []
def test_prepare_with_stop_sequences(self, builder):
req = MagicMock()
req.messages = []
req.model = "x"
req.stop_sequences = ["stop", "end"]
req.system = None
engine = MagicMock()
engine.tokenizer.apply_chat_template.return_value = ""
_, _, stops = builder.prepare(req, engine)
assert stops == ["stop", "end"]
def test_gen_context_defaults():
ctx = GenContext(resp_id="a", created=1, model="m", prompt_tokens=10)
assert ctx.completion_tokens == 0
def test_format_stream_start(self, builder):
ctx = _make_ctx(prompt_tokens=3)
events = builder.format_stream_start(ctx)
payloads = _sse_payloads(events)
assert len(payloads) == 2
assert payloads[0]["type"] == "message_start"
assert payloads[0]["message"]["usage"]["input_tokens"] == 3
assert payloads[1]["type"] == "content_block_start"
def test_format_chunk(self, builder):
events = builder.format_chunk("tok", body="tok")
payload = json.loads(events[0].split("data: ", 1)[1])
assert payload["type"] == "content_block_delta"
assert payload["delta"]["text"] == "tok"
def test_gen_context_fields_mutable():
ctx = GenContext(resp_id="a", created=1, model="m", prompt_tokens=10)
ctx.completion_tokens = 42
assert ctx.completion_tokens == 42
def test_format_stream_end_no_stop(self, builder):
ctx = _make_ctx(completion_tokens=3)
stop = StopInfo()
events = builder.format_stream_end(ctx, stop)
payloads = _sse_payloads(events)
# content_block_stop, message_delta, message_stop
types = [p["type"] for p in payloads]
assert types == ["content_block_stop", "message_delta", "message_stop"]
assert payloads[1]["delta"]["stop_reason"] == "end_turn"
def test_format_stream_end_with_stop_trims_and_emits_remaining(self, builder):
ctx = _make_ctx(completion_tokens=7)
stop = StopInfo(
matched="END",
body="Hello world END extra",
yielded="Hello ",
)
events = builder.format_stream_end(ctx, stop)
payloads = _sse_payloads(events)
# unyielded delta, content_block_stop, message_delta, message_stop
types = [p["type"] for p in payloads]
assert types == [
"content_block_delta",
"content_block_stop",
"message_delta",
"message_stop",
]
assert payloads[0]["delta"]["text"] == "world "
assert payloads[2]["delta"]["stop_reason"] == "stop_sequence"
assert payloads[2]["delta"]["stop_sequence"] == "END"
def test_stop_info_defaults():
s = StopInfo()
assert s.matched is None
assert s.body == ""
assert s.yielded == ""
def test_format_stream_end_stop_trimmed_already_yielded(self, builder):
ctx = _make_ctx()
stop = StopInfo(
matched="END",
body="Hello END",
yielded="Hello ",
)
events = builder.format_stream_end(ctx, stop)
payloads = _sse_payloads(events)
# No unyielded delta (everything already sent)
types = [p["type"] for p in payloads]
assert types == ["content_block_stop", "message_delta", "message_stop"]
def test_format_response_with_stop_trims_content(self, builder):
ctx = _make_ctx()
stop = StopInfo(matched="STOP", body="text STOP extra", yielded="text ")
resp = builder.format_response(ctx, "text STOP extra", stop)
assert resp["content"][0]["text"] == "text "
assert resp["stop_reason"] == "stop_sequence"
assert resp["stop_sequence"] == "STOP"
def test_stop_info_with_values():
s = StopInfo(matched="stop", body="hello stop", yielded="hello ")
assert s.matched == "stop"
assert s.body == "hello stop"
assert s.yielded == "hello "
def test_format_response_no_stop(self, builder):
ctx = _make_ctx()
stop = StopInfo()
resp = builder.format_response(ctx, "full text", stop)
assert resp["content"][0]["text"] == "full text"
assert resp["stop_reason"] == "end_turn"
def test_openai_prepare_returns_prompt_ctx_stops():
builder = _make_openai_builder()
req = MagicMock()
req.messages = [MagicMock(role="user", content="Hi")]
req.stop = ["END"]
req.model = "gpt"
engine = MagicMock()
engine.tokenizer.apply_chat_template.return_value = "Hi"
prompt, ctx, stops = builder.prepare(req, engine)
assert prompt == "Hi"
assert ctx.model == "gpt"
assert ctx.prompt_tokens == 0
assert stops == ["END"]
def test_openai_prepare_no_stop_returns_empty_list():
builder = _make_openai_builder()
req = MagicMock()
req.messages = []
req.stop = None
req.model = "x"
engine = MagicMock()
engine.tokenizer.apply_chat_template.return_value = ""
_, _, stops = builder.prepare(req, engine)
assert stops == []
def test_openai_format_stream_start():
builder = _make_openai_builder()
ctx = _make_ctx()
events = builder.format_stream_start(ctx)
payloads = _sse_payloads(events)
assert len(payloads) == 1
p = payloads[0]
assert p["object"] == "chat.completion.chunk"
assert p["choices"][0]["delta"]["role"] == "assistant"
assert p["choices"][0]["finish_reason"] is None
def test_openai_format_chunk():
builder = _make_openai_builder()
events = builder.format_chunk("hello", body="hello")
payload = json.loads(events[0].split("data: ", 1)[1])
assert payload["choices"][0]["delta"]["content"] == "hello"
assert payload["choices"][0]["finish_reason"] is None
def test_openai_format_stream_end():
builder = _make_openai_builder()
ctx = _make_ctx(completion_tokens=5)
stop = StopInfo(matched="stop")
events = builder.format_stream_end(ctx, stop)
payloads = _sse_payloads(events)
finish = payloads[0]
assert finish["choices"][0]["finish_reason"] == "stop"
usage = payloads[1]
assert usage["completion_tokens"] == 5
assert usage["total_tokens"] == 15
def test_openai_format_response():
builder = _make_openai_builder()
ctx = _make_ctx()
stop = StopInfo()
resp = builder.format_response(ctx, "hello", stop)
assert resp["object"] == "chat.completion"
assert resp["choices"][0]["message"]["content"] == "hello"
assert resp["usage"]["prompt_tokens"] == 10
def test_anthropic_prepare_messages():
builder = _make_anthropic_builder()
req = MagicMock()
req.messages = [MagicMock(role="user", content="Hi")]
req.model = "claude"
req.system = None
req.stop_sequences = None
engine = MagicMock()
engine.tokenizer.apply_chat_template.return_value = "Hi"
prompt, ctx, stops = builder.prepare(req, engine)
assert prompt == "Hi"
assert stops == []
def test_anthropic_prepare_with_stop_sequences():
builder = _make_anthropic_builder()
req = MagicMock()
req.messages = []
req.model = "x"
req.stop_sequences = ["stop", "end"]
req.system = None
engine = MagicMock()
engine.tokenizer.apply_chat_template.return_value = ""
_, _, stops = builder.prepare(req, engine)
assert stops == ["stop", "end"]
def test_anthropic_format_stream_start():
builder = _make_anthropic_builder()
ctx = _make_ctx(prompt_tokens=3)
events = builder.format_stream_start(ctx)
payloads = _sse_payloads(events)
assert len(payloads) == 2
assert payloads[0]["type"] == "message_start"
assert payloads[0]["message"]["usage"]["input_tokens"] == 3
assert payloads[1]["type"] == "content_block_start"
def test_anthropic_format_chunk():
builder = _make_anthropic_builder()
events = builder.format_chunk("tok", body="tok")
payload = json.loads(events[0].split("data: ", 1)[1])
assert payload["type"] == "content_block_delta"
assert payload["delta"]["text"] == "tok"
def test_anthropic_format_stream_end_no_stop():
builder = _make_anthropic_builder()
ctx = _make_ctx(completion_tokens=3)
stop = StopInfo()
events = builder.format_stream_end(ctx, stop)
payloads = _sse_payloads(events)
types = [p["type"] for p in payloads]
assert types == ["content_block_stop", "message_delta", "message_stop"]
assert payloads[1]["delta"]["stop_reason"] == "end_turn"
def test_anthropic_format_stream_end_with_stop_trims_and_emits_remaining():
builder = _make_anthropic_builder()
ctx = _make_ctx(completion_tokens=7)
stop = StopInfo(
matched="END",
body="Hello world END extra",
yielded="Hello ",
)
events = builder.format_stream_end(ctx, stop)
payloads = _sse_payloads(events)
types = [p["type"] for p in payloads]
assert types == [
"content_block_delta",
"content_block_stop",
"message_delta",
"message_stop",
]
assert payloads[0]["delta"]["text"] == "world "
assert payloads[2]["delta"]["stop_reason"] == "stop_sequence"
assert payloads[2]["delta"]["stop_sequence"] == "END"
def test_anthropic_format_stream_end_stop_trimmed_already_yielded():
builder = _make_anthropic_builder()
ctx = _make_ctx()
stop = StopInfo(
matched="END",
body="Hello END",
yielded="Hello ",
)
events = builder.format_stream_end(ctx, stop)
payloads = _sse_payloads(events)
types = [p["type"] for p in payloads]
assert types == ["content_block_stop", "message_delta", "message_stop"]
def test_anthropic_format_response_with_stop_trims_content():
builder = _make_anthropic_builder()
ctx = _make_ctx()
stop = StopInfo(matched="STOP", body="text STOP extra", yielded="text ")
resp = builder.format_response(ctx, "text STOP extra", stop)
assert resp["content"][0]["text"] == "text "
assert resp["stop_reason"] == "stop_sequence"
assert resp["stop_sequence"] == "STOP"
def test_anthropic_format_response_no_stop():
builder = _make_anthropic_builder()
ctx = _make_ctx()
stop = StopInfo()
resp = builder.format_response(ctx, "full text", stop)
assert resp["content"][0]["text"] == "full text"
assert resp["stop_reason"] == "end_turn"
+44
View File
@@ -1,8 +1,15 @@
"""Unit tests for the inference HTTP server."""
from pathlib import Path
import pytest
import torch
from astrai.inference import get_app
from astrai.inference.network.app import _create_engine
from astrai.model.transformer import AutoRegressiveLM
from astrai.serialization import save_model
from tests.helpers import CHAT_TEMPLATE, build_test_tokenizer, make_tiny_config
def test_health_no_model(client):
@@ -212,5 +219,42 @@ def test_chat_completions_stop_sequence_stream(client, loaded_model):
)
def test_chat_completions_real_engine(tmp_path, client):
"""POST /v1/chat/completions with a real tiny model and tokenizer."""
cfg = make_tiny_config(vocab_size=256)
model = AutoRegressiveLM(cfg).eval()
save_model(cfg.to_dict(), model.state_dict(), str(tmp_path))
tokenizer = build_test_tokenizer(vocab_size=256, chat_template=CHAT_TEMPLATE)
tokenizer.save_pretrained(str(tmp_path))
engine = _create_engine(
Path(tmp_path),
device="cpu",
dtype=torch.float32,
max_batch_size=1,
max_seq_len=64,
)
try:
get_app().state.engine = engine
response = client.post(
"/v1/chat/completions",
json={
"messages": [{"role": "user", "content": "Hello"}],
"max_tokens": 4,
"temperature": 0.0,
"stream": False,
},
)
assert response.status_code == 200
data = response.json()
content = data["choices"][0]["message"]["content"]
assert isinstance(content, str)
assert data["usage"]["completion_tokens"] > 0
finally:
engine.shutdown()
get_app().state.engine = None
if __name__ == "__main__":
pytest.main([__file__, "-v"])

Some files were not shown because too many files have changed in this diff Show More