Compare commits
27
Commits
3d3ea47d37
...
998b443aa3
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
998b443aa3 | ||
|
|
cebdd45d3a | ||
|
|
7da1439c9e | ||
|
|
29e5f571af | ||
|
|
74e694921c | ||
|
|
d5067af064 | ||
|
|
f6db546578 | ||
|
|
31ca357c61 | ||
|
|
34471252ab | ||
|
|
aa08479285 | ||
|
|
4b10d3ca37 | ||
|
|
2bc4d2b8a8 | ||
|
|
4244df2785 | ||
|
|
a29bdfae46 | ||
|
|
10fec8dca1 | ||
|
|
75304d084d | ||
|
|
16a55bb474 | ||
|
|
cb21af38ba | ||
|
|
dcc96de12a | ||
|
|
7d27f3e078 | ||
|
|
84753d3e08 | ||
|
|
53a7149577 | ||
|
|
c79d34eee1 | ||
|
|
398e8a3ea3 | ||
|
|
f252af495c | ||
|
|
00c2c80c8f | ||
|
|
a6c6a54ace |
+3
-2
@@ -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
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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}"
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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)
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
Vendored
-10
@@ -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:
|
||||
|
||||
Vendored
-2
@@ -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 ----
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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",
|
||||
)
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
@@ -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
@@ -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).
|
||||
|
||||
@@ -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
|
||||
+15
-10
@@ -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,
|
||||
+10
-10
@@ -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
|
||||
+14
-8
@@ -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
|
||||
@@ -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;
|
||||
|
||||
};
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)");
|
||||
}
|
||||
@@ -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");
|
||||
}
|
||||
@@ -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,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);
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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
@@ -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}
|
||||
|
||||
@@ -57,7 +57,7 @@ AstrAI 是一个覆盖模型构建、训练、评测与部署的端到端 Transf
|
||||
| **数据** | 声明式 JSON 预处理、可配置掩码与样本打包、二进制/JSONL 存储和流式数据集 |
|
||||
| **推理** | 连续批处理、分页 KV Cache、Radix 前缀缓存、流式生成,以及 Torch/CUDA/FlashAttention 后端 |
|
||||
| **服务** | 基于 FastAPI 的 OpenAI 与 Anthropic 聊天补全协议,支持 SSE 流式输出和工具调用 |
|
||||
| **评测** | Perplexity、MMLU、HumanEval、IFEval、IFD 和 ROUGE 评测工具 |
|
||||
| **评测** | Perplexity、MMLU、HumanEval、IFEval、IFD、ROUGE 和权重分析评测工具 |
|
||||
| **扩展** | 基于工厂与注册表扩展模型、数据集、训练策略、回调、内核和协议组件 |
|
||||
|
||||
### 快速上手
|
||||
@@ -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`) |
|
||||
|
||||
### 贡献
|
||||
|
||||
|
||||
@@ -1437,7 +1437,8 @@ classDiagram
|
||||
| **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
|
||||
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy–GRPOStrategy, StrategyFactory, BaseScheduler–WSDScheduler, SchedulerFactory, TrainCallback(Protocol)–MetricCallback, CallbackFactory, RawRollout, RolloutResult, BaseRewardModel, RolloutGenerator, RolloutRunner | Training workflow |
|
||||
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, InferenceWorkspace, PagePool, KVStorage, ReqToTokenPool, KVCache, Allocator, RadixCache, Task, TaskManager, TaskStatus, StreamDecoder, GenerateResult, BaseSamplingStrategy–SamplingPipeline, 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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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__" }
|
||||
|
||||
|
||||
@@ -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[@]}" "$@"
|
||||
|
||||
@@ -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 "$@"
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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
@@ -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,
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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)
|
||||
@@ -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 = {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"}
|
||||
|
||||
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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"
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user