Compare commits
55
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3639b50b4a | ||
|
|
d855c09cf3 | ||
|
|
d6bfb09863 | ||
|
|
6db276f37a | ||
|
|
6c76c16480 | ||
|
|
11073bd1d2 | ||
|
|
25c9e81b2b | ||
|
|
ffbd9b57c9 | ||
|
|
04899a2b15 | ||
|
|
530d280e33 | ||
|
|
21ddead238 | ||
|
|
7aa5ed09d9 | ||
|
|
75411ce0cc | ||
|
|
9f83d982ec | ||
|
|
3e67b4f88d | ||
|
|
50cfd0d555 | ||
|
|
5756054d38 | ||
|
|
738cb8f128 | ||
|
|
28d1bd07cf | ||
|
|
02625739fe | ||
|
|
f688cd9c5a | ||
|
|
8055027df7 | ||
|
|
3067a8e1a6 | ||
|
|
97114b95a4 | ||
|
|
32fd03a025 | ||
|
|
21bf37dd83 | ||
|
|
5b67d5865a | ||
|
|
df979b4469 | ||
|
|
deb2d7e127 | ||
|
|
fc47319240 | ||
|
|
22cf798d81 | ||
|
|
164be9708b | ||
|
|
6a97524db4 | ||
|
|
c8b1e40f71 | ||
|
|
bcaa2d1ae0 | ||
|
|
8206afefd9 | ||
|
|
646b1b0f46 | ||
|
|
8150ab6c32 | ||
|
|
0b0693a0a2 | ||
|
|
115192c67c | ||
|
|
c2b04d8458 | ||
|
|
db487ab48b | ||
|
|
a95794d3db | ||
|
|
39f84f3b4c | ||
|
|
9f7cf50c56 | ||
|
|
d9a0c72149 | ||
|
|
5ab18bec48 | ||
|
|
2e29ed45d3 | ||
|
|
5ba21f4eb3 | ||
|
|
c26a47b0df | ||
|
|
b1a87b22bb | ||
|
|
07625057f2 | ||
|
|
53c804e233 | ||
|
|
05c7432964 | ||
|
|
4de42d83c2 |
+3
-1
@@ -4,6 +4,8 @@
|
|||||||
# Allow necessary files
|
# Allow necessary files
|
||||||
!astrai/
|
!astrai/
|
||||||
!scripts/
|
!scripts/
|
||||||
!assets/
|
!docs/
|
||||||
|
!csrc/
|
||||||
|
!setup.py
|
||||||
!pyproject.toml
|
!pyproject.toml
|
||||||
!README.md
|
!README.md
|
||||||
|
|||||||
@@ -26,22 +26,30 @@ jobs:
|
|||||||
if-no-files-found: error
|
if-no-files-found: error
|
||||||
|
|
||||||
build-cuda-linux:
|
build-cuda-linux:
|
||||||
name: Build CUDA wheel (Linux)
|
name: Build CUDA wheel (Linux, ${{ matrix.cuda_tag }})
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
|
strategy:
|
||||||
|
fail-fast: false
|
||||||
|
matrix:
|
||||||
|
include:
|
||||||
|
- cuda_tag: "cu128"
|
||||||
|
cuda_ver: "12.8.0"
|
||||||
|
- cuda_tag: "cu130"
|
||||||
|
cuda_ver: "13.0.0"
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v4
|
- uses: actions/checkout@v4
|
||||||
- uses: actions/setup-python@v5
|
- uses: actions/setup-python@v5
|
||||||
with:
|
with:
|
||||||
python-version: "3.12"
|
python-version: "3.12"
|
||||||
|
|
||||||
- name: Install torch (CUDA 12.8)
|
- name: Install torch (${{ matrix.cuda_tag }})
|
||||||
run: |
|
run: |
|
||||||
pip install torch --index-url https://download.pytorch.org/whl/cu128
|
pip install torch --index-url https://download.pytorch.org/whl/${{ matrix.cuda_tag }}
|
||||||
|
|
||||||
- name: Setup CUDA
|
- name: Setup CUDA (${{ matrix.cuda_ver }})
|
||||||
uses: Jimver/cuda-toolkit@v0.2.35
|
uses: Jimver/cuda-toolkit@v0.2.35
|
||||||
with:
|
with:
|
||||||
cuda: "12.8.0"
|
cuda: "${{ matrix.cuda_ver }}"
|
||||||
|
|
||||||
- name: Build wheel (with CUDA kernels)
|
- name: Build wheel (with CUDA kernels)
|
||||||
run: |
|
run: |
|
||||||
@@ -49,7 +57,7 @@ jobs:
|
|||||||
|
|
||||||
- uses: actions/upload-artifact@v4
|
- uses: actions/upload-artifact@v4
|
||||||
with:
|
with:
|
||||||
name: cuda-wheel-linux
|
name: cuda-wheel-linux-${{ matrix.cuda_tag }}
|
||||||
path: dist/*.whl
|
path: dist/*.whl
|
||||||
if-no-files-found: error
|
if-no-files-found: error
|
||||||
|
|
||||||
@@ -66,10 +74,11 @@ jobs:
|
|||||||
name: pure-wheel
|
name: pure-wheel
|
||||||
path: release-assets/pure
|
path: release-assets/pure
|
||||||
|
|
||||||
- name: Download CUDA wheel
|
- name: Download CUDA wheels (all variants)
|
||||||
uses: actions/download-artifact@v4
|
uses: actions/download-artifact@v4
|
||||||
with:
|
with:
|
||||||
name: cuda-wheel-linux
|
pattern: cuda-wheel-linux-*
|
||||||
|
merge-multiple: true
|
||||||
path: release-assets/cuda
|
path: release-assets/cuda
|
||||||
|
|
||||||
- name: Verify release assets
|
- name: Verify release assets
|
||||||
@@ -79,8 +88,7 @@ jobs:
|
|||||||
pure_wheels=(release-assets/pure/*.whl)
|
pure_wheels=(release-assets/pure/*.whl)
|
||||||
cuda_wheels=(release-assets/cuda/*.whl)
|
cuda_wheels=(release-assets/cuda/*.whl)
|
||||||
test "${#pure_wheels[@]}" -eq 1
|
test "${#pure_wheels[@]}" -eq 1
|
||||||
test "${#cuda_wheels[@]}" -eq 1
|
test "${#cuda_wheels[@]}" -ge 1
|
||||||
test "$(basename "${pure_wheels[0]}")" != "$(basename "${cuda_wheels[0]}")"
|
|
||||||
|
|
||||||
- name: Create release & upload assets
|
- name: Create release & upload assets
|
||||||
uses: softprops/action-gh-release@v2
|
uses: softprops/action-gh-release@v2
|
||||||
|
|||||||
+1
-1
@@ -24,7 +24,7 @@
|
|||||||
!/.dockerignore
|
!/.dockerignore
|
||||||
!/Dockerfile
|
!/Dockerfile
|
||||||
!/docker-compose.yml
|
!/docker-compose.yml
|
||||||
!/assets/**
|
!/docs/**
|
||||||
!/CONTRIBUTING.md
|
!/CONTRIBUTING.md
|
||||||
!/LICENSE
|
!/LICENSE
|
||||||
!/pyproject.toml
|
!/pyproject.toml
|
||||||
|
|||||||
+12
-2
@@ -1,8 +1,16 @@
|
|||||||
# AstrAI Dockerfile - Multi-stage Build (Optimized)
|
# AstrAI Dockerfile - Multi-stage Build (Optimized)
|
||||||
|
#
|
||||||
|
# CUDA version selection:
|
||||||
|
# docker build -t astrai .
|
||||||
|
# docker build -t astrai --build-arg CUDA_TAG=cu128 .
|
||||||
|
# docker build -t astrai --build-arg CUDA_TAG=cu130 .
|
||||||
|
# Default: cu128
|
||||||
|
|
||||||
# Build stage - use base image with minimal build tools
|
# Build stage - use base image with minimal build tools
|
||||||
FROM ubuntu:24.04 AS builder
|
FROM ubuntu:24.04 AS builder
|
||||||
|
|
||||||
|
ARG CUDA_TAG=cu128
|
||||||
|
|
||||||
WORKDIR /app
|
WORKDIR /app
|
||||||
|
|
||||||
# Install Python 3.12 and minimal build dependencies
|
# Install Python 3.12 and minimal build dependencies
|
||||||
@@ -20,10 +28,12 @@ ENV PATH="/opt/venv/bin:$PATH"
|
|||||||
|
|
||||||
# Copy source code and install (deps read from pyproject.toml)
|
# Copy source code and install (deps read from pyproject.toml)
|
||||||
COPY astrai/ ./astrai/
|
COPY astrai/ ./astrai/
|
||||||
|
COPY csrc/ ./csrc/
|
||||||
|
COPY setup.py .
|
||||||
COPY pyproject.toml .
|
COPY pyproject.toml .
|
||||||
RUN pip install --no-cache-dir --upgrade pip \
|
RUN pip install --no-cache-dir --upgrade pip \
|
||||||
&& pip install --no-cache-dir . \
|
&& pip install --no-cache-dir . \
|
||||||
--extra-index-url https://download.pytorch.org/whl/cu128
|
--extra-index-url "https://download.pytorch.org/whl/${CUDA_TAG}"
|
||||||
|
|
||||||
# Production stage
|
# Production stage
|
||||||
FROM ubuntu:24.04 AS production
|
FROM ubuntu:24.04 AS production
|
||||||
@@ -43,7 +53,7 @@ ENV PATH="/opt/venv/bin:$PATH"
|
|||||||
# Copy application code
|
# Copy application code
|
||||||
COPY astrai/ ./astrai/
|
COPY astrai/ ./astrai/
|
||||||
COPY scripts/ ./scripts/
|
COPY scripts/ ./scripts/
|
||||||
COPY assets/ ./assets/
|
COPY docs/ ./docs/
|
||||||
COPY pyproject.toml .
|
COPY pyproject.toml .
|
||||||
COPY README.md .
|
COPY README.md .
|
||||||
|
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
<div align="center">
|
<div align="center">
|
||||||
|
|
||||||
<img src="assets/images/logo.png" width="auto" alt="Logo">
|
<img src="docs/images/logo.png" width="auto" alt="Logo">
|
||||||
<p>
|
<p>
|
||||||
<strong>A lightweight Transformer training & inference framework</strong>
|
<strong>A lightweight Transformer training & inference framework</strong>
|
||||||
</p>
|
</p>
|
||||||
@@ -17,7 +17,7 @@
|
|||||||
|
|
||||||
<div align="center">
|
<div align="center">
|
||||||
<a href="#english">English</a> •
|
<a href="#english">English</a> •
|
||||||
<a href="assets/docs/README-zh-CN.md">中文</a> •
|
<a href="docs/README-zh-CN.md">中文</a> •
|
||||||
<a href="https://github.com/ViperEkura/AstrAI/issues">Issue Tracker</a> •
|
<a href="https://github.com/ViperEkura/AstrAI/issues">Issue Tracker</a> •
|
||||||
<a href="https://github.com/ViperEkura/AstrAI/discussions">Discussions</a> •
|
<a href="https://github.com/ViperEkura/AstrAI/discussions">Discussions</a> •
|
||||||
<a href="https://huggingface.co/ViperEkura">HuggingFace</a>
|
<a href="https://huggingface.co/ViperEkura">HuggingFace</a>
|
||||||
@@ -213,18 +213,23 @@ curl -X POST http://localhost:8000/v1/messages \
|
|||||||
curl http://localhost:8000/health
|
curl http://localhost:8000/health
|
||||||
```
|
```
|
||||||
|
|
||||||
See [Inference Guide](assets/docs/inference.md) for SSE streaming format, error codes, and stats endpoint.
|
See [Inference Guide](docs/guides/inference.md) for SSE streaming format, error codes, and stats endpoint.
|
||||||
|
|
||||||
### Documentation
|
### Documentation
|
||||||
|
|
||||||
| Document | Description |
|
| Document | Description |
|
||||||
|----------|-------------|
|
|----------|-------------|
|
||||||
| [CLI Reference](./assets/docs/params.md) | Parameters for all CLI tools (train, server, generate, preprocess) |
|
| [Get Started](./docs/get-started.md) | Installation and quickstart |
|
||||||
| [Architecture](./assets/docs/architecture.md) | System architecture, class diagram & design patterns |
|
| [CLI Reference](./docs/guides/params.md) | Parameters for all CLI tools (train, server, generate, preprocess) |
|
||||||
| [Training](./assets/docs/training.md) | Training loop, strategies & formulas |
|
| [Preprocessing](./docs/guides/preprocessing.md) | Declarative JSON-driven data preprocessing |
|
||||||
| [Inference](./assets/docs/inference.md) | KVCache, continuous batching, sampling & HTTP API |
|
| [Training](./docs/guides/training.md) | Training loop, strategies & formulas |
|
||||||
| [Data Flow](./assets/docs/dataflow.md) | Data pipeline, storage backends & dataset architecture |
|
| [Inference](./docs/guides/inference.md) | KVCache, continuous batching, sampling & HTTP API |
|
||||||
| [Preprocessing](./assets/docs/preprocessing.md) | Declarative JSON-driven data preprocessing |
|
| [Evaluation](./docs/guides/evaluation.md) | HumanEval, MMLU, PPL, ROUGE, IFD, IFEval |
|
||||||
|
| [Distributed](./docs/guides/distributed.md) | Multi-GPU DDP / FSDP training |
|
||||||
|
| [Architecture](./docs/developer/architecture.md) | System architecture, class diagram & design patterns |
|
||||||
|
| [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 |
|
||||||
|
|
||||||
### Contributing
|
### Contributing
|
||||||
|
|
||||||
|
|||||||
+29
-1
@@ -1,6 +1,9 @@
|
|||||||
__version__ = "1.3.11"
|
__version__ = "1.3.12"
|
||||||
__author__ = "ViperEkura"
|
__author__ = "ViperEkura"
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
|
||||||
from astrai.config import (
|
from astrai.config import (
|
||||||
AutoRegressiveLMConfig,
|
AutoRegressiveLMConfig,
|
||||||
BaseModelConfig,
|
BaseModelConfig,
|
||||||
@@ -53,6 +56,30 @@ from astrai.trainer import (
|
|||||||
Trainer,
|
Trainer,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def setup_logging(level: str = "INFO"):
|
||||||
|
"""Attach a handler to the ``astrai`` logger (only, not root).
|
||||||
|
|
||||||
|
Call once per process, e.g. at the top of CLI scripts.
|
||||||
|
Set ``ASTR_LOG_LEVEL`` to override the default ``INFO``.
|
||||||
|
"""
|
||||||
|
_logger = logging.getLogger("astrai")
|
||||||
|
if _logger.handlers:
|
||||||
|
return
|
||||||
|
_level = getattr(
|
||||||
|
logging, os.environ.get("ASTR_LOG_LEVEL", level).upper(), logging.INFO
|
||||||
|
)
|
||||||
|
_logger.setLevel(_level)
|
||||||
|
_handler = logging.StreamHandler()
|
||||||
|
_handler.setFormatter(
|
||||||
|
logging.Formatter(
|
||||||
|
"%(asctime)s | %(levelname)-7s | %(name)s | %(message)s",
|
||||||
|
datefmt="%Y-%m-%d %H:%M:%S",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
_logger.addHandler(_handler)
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"AutoRegressiveLM",
|
"AutoRegressiveLM",
|
||||||
"AutoRegressiveLMConfig",
|
"AutoRegressiveLMConfig",
|
||||||
@@ -94,5 +121,6 @@ __all__ = [
|
|||||||
"only_on_rank",
|
"only_on_rank",
|
||||||
"run_server",
|
"run_server",
|
||||||
"sample",
|
"sample",
|
||||||
|
"setup_logging",
|
||||||
"spawn_parallel_fn",
|
"spawn_parallel_fn",
|
||||||
]
|
]
|
||||||
|
|||||||
+20
-80
@@ -1,92 +1,32 @@
|
|||||||
import json
|
import json
|
||||||
from dataclasses import MISSING, dataclass, fields
|
from dataclasses import asdict
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Dict, Optional, Self, Union, get_type_hints
|
from typing import Any, Dict, Self, Union
|
||||||
|
|
||||||
|
from pydantic import ConfigDict
|
||||||
|
from pydantic.dataclasses import dataclass
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass(config=ConfigDict(use_attribute_docstrings=True))
|
||||||
class BaseConfig:
|
class BaseConfig:
|
||||||
def to_dict(self) -> Dict[str, Any]:
|
def to_dict(self) -> Dict[str, Any]:
|
||||||
d = {}
|
result = {}
|
||||||
for fld in fields(self):
|
for k, v in asdict(self).items():
|
||||||
v = getattr(self, fld.name)
|
if isinstance(v, tuple):
|
||||||
if isinstance(v, (str, int, float, bool)):
|
v = list(v)
|
||||||
d[fld.name] = v
|
try:
|
||||||
elif v is None:
|
json.dumps(v)
|
||||||
d[fld.name] = None
|
result[k] = v
|
||||||
elif isinstance(v, (dict, list, tuple)):
|
except (TypeError, ValueError):
|
||||||
try:
|
# Skip non-serializable runtime objects (e.g. model_fn, dataset).
|
||||||
val = list(v) if isinstance(v, tuple) else v
|
# TrainConfig mixes hyperparams with callables/datasets; only the
|
||||||
json.dumps(val)
|
# JSON-serializable subset is written to checkpoint meta.
|
||||||
d[fld.name] = val
|
pass
|
||||||
except (TypeError, ValueError):
|
return result
|
||||||
pass
|
|
||||||
elif isinstance(v, BaseConfig):
|
|
||||||
d[fld.name] = v.to_dict()
|
|
||||||
elif hasattr(v, "__dataclass_fields__"):
|
|
||||||
sub = {}
|
|
||||||
for f in fields(v):
|
|
||||||
a = getattr(v, f.name)
|
|
||||||
sub[f.name] = list(a) if isinstance(a, tuple) else a
|
|
||||||
d[fld.name] = sub
|
|
||||||
return d
|
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_dict(cls, d: Dict[str, Any]) -> Self:
|
def from_dict(cls, d: Dict[str, Any]) -> Self:
|
||||||
hints = get_type_hints(cls)
|
return cls(**d)
|
||||||
inst = cls.__new__(cls)
|
|
||||||
for fld in fields(cls):
|
|
||||||
if fld.name in d:
|
|
||||||
v = d[fld.name]
|
|
||||||
target = cls._unwrap_optional(hints.get(fld.name))
|
|
||||||
if target is not None:
|
|
||||||
try:
|
|
||||||
v = cls._coerce(v, target)
|
|
||||||
except (TypeError, ValueError):
|
|
||||||
pass
|
|
||||||
object.__setattr__(inst, fld.name, v)
|
|
||||||
elif fld.default is not MISSING:
|
|
||||||
object.__setattr__(inst, fld.name, fld.default)
|
|
||||||
elif fld.default_factory is not MISSING:
|
|
||||||
object.__setattr__(inst, fld.name, fld.default_factory())
|
|
||||||
else:
|
|
||||||
object.__setattr__(inst, fld.name, None)
|
|
||||||
return inst
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _unwrap_optional(tp) -> Optional[type]:
|
|
||||||
if tp is None:
|
|
||||||
return None
|
|
||||||
origin = getattr(tp, "__origin__", None)
|
|
||||||
if origin is not None:
|
|
||||||
args = getattr(tp, "__args__", ())
|
|
||||||
non_none = [a for a in args if a is not type(None)]
|
|
||||||
return non_none[0] if non_none else None
|
|
||||||
return tp
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _coerce(value: Any, target_type: type) -> Any:
|
|
||||||
if target_type is bool and isinstance(value, bool):
|
|
||||||
return value
|
|
||||||
if (
|
|
||||||
target_type is int
|
|
||||||
and isinstance(value, (int, float))
|
|
||||||
and not isinstance(value, bool)
|
|
||||||
):
|
|
||||||
return int(value)
|
|
||||||
if (
|
|
||||||
target_type is float
|
|
||||||
and isinstance(value, (int, float))
|
|
||||||
and not isinstance(value, bool)
|
|
||||||
):
|
|
||||||
return float(value)
|
|
||||||
if target_type is str and isinstance(value, str):
|
|
||||||
return value
|
|
||||||
if isinstance(value, target_type):
|
|
||||||
return value
|
|
||||||
if isinstance(value, dict) and issubclass(target_type, BaseConfig):
|
|
||||||
return target_type.from_dict(value)
|
|
||||||
raise TypeError
|
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_file(cls, path: Union[str, Path]) -> Self:
|
def from_file(cls, path: Union[str, Path]) -> Self:
|
||||||
|
|||||||
@@ -1,9 +1,14 @@
|
|||||||
from dataclasses import dataclass
|
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any, Dict, Optional
|
||||||
|
|
||||||
|
from pydantic import field_validator
|
||||||
|
from pydantic.dataclasses import dataclass
|
||||||
|
|
||||||
from astrai.config.base import BaseConfig
|
from astrai.config.base import BaseConfig
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
|
|
||||||
|
_ATTN_TYPES = frozenset({"gqa", "mla"})
|
||||||
|
_FFN_TYPES = frozenset({"mlp", "moe"})
|
||||||
|
|
||||||
|
|
||||||
class ConfigFactory(BaseFactory[BaseConfig]):
|
class ConfigFactory(BaseFactory[BaseConfig]):
|
||||||
"""Factory that dispatches config classes by ``model_type``."""
|
"""Factory that dispatches config classes by ``model_type``."""
|
||||||
@@ -17,7 +22,12 @@ class ConfigFactory(BaseFactory[BaseConfig]):
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class BaseModelConfig(BaseConfig):
|
class BaseModelConfig(BaseConfig):
|
||||||
"""Base config with ``model_type`` dispatch and file I/O."""
|
"""Base config with ``model_type`` dispatch and file I/O.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model_type (Optional[str]): Model type identifier for AutoModel dispatch. Defaults to None.
|
||||||
|
neftune_alpha (float): NEFTune noise alpha, 0=disabled, typical: 5.0. Defaults to 0.0.
|
||||||
|
"""
|
||||||
|
|
||||||
model_type: Optional[str] = None
|
model_type: Optional[str] = None
|
||||||
neftune_alpha: float = 0.0
|
neftune_alpha: float = 0.0
|
||||||
@@ -26,7 +36,34 @@ class BaseModelConfig(BaseConfig):
|
|||||||
@dataclass
|
@dataclass
|
||||||
@ConfigFactory.register("autoregressive_lm")
|
@ConfigFactory.register("autoregressive_lm")
|
||||||
class AutoRegressiveLMConfig(BaseModelConfig):
|
class AutoRegressiveLMConfig(BaseModelConfig):
|
||||||
"""Configuration for autoregressive language model."""
|
"""Configuration for autoregressive language model.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model_type (Optional[str]): Model type identifier for AutoModel dispatch. Defaults to None.
|
||||||
|
neftune_alpha (float): NEFTune noise alpha, 0=disabled, typical: 5.0. Defaults to 0.0.
|
||||||
|
vocab_size (Optional[int]): Vocabulary size. Defaults to None.
|
||||||
|
hidden_size (Optional[int]): Hidden dimension size. Defaults to None.
|
||||||
|
num_hidden_layers (Optional[int]): Number of transformer layers. Defaults to None.
|
||||||
|
rms_norm_eps (Optional[float]): Epsilon for RMSNorm. Defaults to None.
|
||||||
|
intermediate_size (Optional[int]): Intermediate size in FFN. Defaults to None.
|
||||||
|
tie_word_embeddings (Optional[bool]): Whether to tie embedding and lm_head weights. Defaults to None.
|
||||||
|
max_position_embeddings (Optional[int]): Maximum sequence length the model was trained with. Defaults to None.
|
||||||
|
rope_theta (Optional[float]): Base frequency for RoPE. Defaults to None.
|
||||||
|
rope_scaling (Optional[dict]): RoPE scaling config, e.g. {"type": "linear", "factor": 4.0}. Defaults to None.
|
||||||
|
attn_type (str): Attention type: 'gqa' or 'mla'. Defaults to "gqa".
|
||||||
|
num_attention_heads (Optional[int]): Number of query attention heads. Defaults to None.
|
||||||
|
num_key_value_heads (Optional[int]): Number of key/value heads for GQA. Defaults to None.
|
||||||
|
use_qk_norm (Optional[bool]): Whether to apply RMSNorm to Q/K. Defaults to None.
|
||||||
|
use_gated_attention (Optional[bool]): Whether to use gated attention. Defaults to None.
|
||||||
|
kv_lora_rank (Optional[int]): KV compression rank, MLA only. Defaults to None.
|
||||||
|
qk_nope_head_dim (Optional[int]): Non-RoPE head dimension, MLA only. Defaults to None.
|
||||||
|
qk_rope_head_dim (Optional[int]): RoPE head dimension, MLA only. Defaults to None.
|
||||||
|
ffn_type (str): FFN type: 'mlp' or 'moe'. Defaults to "mlp".
|
||||||
|
n_routed_experts (Optional[int]): Number of routed experts, MoE only. Defaults to None.
|
||||||
|
n_shared_experts (Optional[int]): Number of shared experts, MoE only. Defaults to None.
|
||||||
|
n_activated_experts (Optional[int]): Number of activated experts per token, MoE only. Defaults to None.
|
||||||
|
topk_method (Optional[str]): Top-k routing method, MoE only. Defaults to None.
|
||||||
|
"""
|
||||||
|
|
||||||
vocab_size: Optional[int] = None
|
vocab_size: Optional[int] = None
|
||||||
hidden_size: Optional[int] = None
|
hidden_size: Optional[int] = None
|
||||||
@@ -34,49 +71,91 @@ class AutoRegressiveLMConfig(BaseModelConfig):
|
|||||||
rms_norm_eps: Optional[float] = None
|
rms_norm_eps: Optional[float] = None
|
||||||
intermediate_size: Optional[int] = None
|
intermediate_size: Optional[int] = None
|
||||||
tie_word_embeddings: Optional[bool] = None
|
tie_word_embeddings: Optional[bool] = None
|
||||||
|
|
||||||
max_position_embeddings: Optional[int] = None
|
max_position_embeddings: Optional[int] = None
|
||||||
rope_theta: Optional[float] = None
|
rope_theta: Optional[float] = None
|
||||||
rope_scaling: Optional[dict] = None
|
rope_scaling: Optional[dict] = None
|
||||||
|
|
||||||
attn_type: str = "gqa"
|
attn_type: str = "gqa"
|
||||||
num_attention_heads: Optional[int] = None
|
num_attention_heads: Optional[int] = None
|
||||||
num_key_value_heads: Optional[int] = None
|
num_key_value_heads: Optional[int] = None
|
||||||
use_qk_norm: Optional[bool] = None
|
use_qk_norm: Optional[bool] = None
|
||||||
use_gated_attention: Optional[bool] = None
|
use_gated_attention: Optional[bool] = None
|
||||||
|
|
||||||
kv_lora_rank: Optional[int] = None
|
kv_lora_rank: Optional[int] = None
|
||||||
qk_nope_head_dim: Optional[int] = None
|
qk_nope_head_dim: Optional[int] = None
|
||||||
qk_rope_head_dim: Optional[int] = None
|
qk_rope_head_dim: Optional[int] = None
|
||||||
|
|
||||||
ffn_type: str = "mlp"
|
ffn_type: str = "mlp"
|
||||||
n_routed_experts: Optional[int] = None
|
n_routed_experts: Optional[int] = None
|
||||||
n_shared_experts: Optional[int] = None
|
n_shared_experts: Optional[int] = None
|
||||||
n_activated_experts: Optional[int] = None
|
n_activated_experts: Optional[int] = None
|
||||||
topk_method: Optional[str] = None
|
topk_method: Optional[str] = None
|
||||||
|
|
||||||
|
@field_validator("attn_type")
|
||||||
|
def _validate_attn_type(cls, v: str) -> str:
|
||||||
|
if v not in _ATTN_TYPES:
|
||||||
|
raise ValueError(
|
||||||
|
f"attn_type must be one of {sorted(_ATTN_TYPES)}, got {v!r}"
|
||||||
|
)
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("ffn_type")
|
||||||
|
def _validate_ffn_type(cls, v: str) -> str:
|
||||||
|
if v not in _FFN_TYPES:
|
||||||
|
raise ValueError(f"ffn_type must be one of {sorted(_FFN_TYPES)}, got {v!r}")
|
||||||
|
return v
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@ConfigFactory.register("embedding")
|
@ConfigFactory.register("embedding")
|
||||||
class EncoderConfig(BaseModelConfig):
|
class EncoderConfig(BaseModelConfig):
|
||||||
"""Configuration for embedding encoder model."""
|
"""Configuration for embedding encoder model.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model_type (Optional[str]): Model type identifier for AutoModel dispatch. Defaults to None.
|
||||||
|
neftune_alpha (float): NEFTune noise alpha, 0=disabled, typical: 5.0. Defaults to 0.0.
|
||||||
|
vocab_size (Optional[int]): Vocabulary size. Defaults to None.
|
||||||
|
hidden_size (Optional[int]): Hidden dimension size. Defaults to None.
|
||||||
|
num_hidden_layers (Optional[int]): Number of transformer layers. Defaults to None.
|
||||||
|
rms_norm_eps (Optional[float]): Epsilon for RMSNorm. Defaults to None.
|
||||||
|
intermediate_size (Optional[int]): Intermediate size in FFN. Defaults to None.
|
||||||
|
max_position_embeddings (Optional[int]): Maximum sequence length the model was trained with. Defaults to None.
|
||||||
|
rope_theta (Optional[float]): Base frequency for RoPE. Defaults to None.
|
||||||
|
rope_scaling (Optional[dict]): RoPE scaling config, e.g. {"type": "linear", "factor": 4.0}. Defaults to None.
|
||||||
|
attn_type (str): Attention type: 'gqa' or 'mla'. Defaults to "gqa".
|
||||||
|
num_attention_heads (Optional[int]): Number of query attention heads. Defaults to None.
|
||||||
|
num_key_value_heads (Optional[int]): Number of key/value heads for GQA. Defaults to None.
|
||||||
|
use_qk_norm (Optional[bool]): Whether to apply RMSNorm to Q/K. Defaults to None.
|
||||||
|
use_gated_attention (Optional[bool]): Whether to use gated attention. Defaults to None.
|
||||||
|
ffn_type (str): FFN type: 'mlp' or 'moe'. Defaults to "mlp".
|
||||||
|
pooling_type (Optional[str]): Pooling strategy for embedding, e.g. 'mean', 'cls'. Defaults to None.
|
||||||
|
normalize_embeddings (Optional[bool]): Whether to L2-normalize output embeddings. Defaults to None.
|
||||||
|
"""
|
||||||
|
|
||||||
vocab_size: Optional[int] = None
|
vocab_size: Optional[int] = None
|
||||||
hidden_size: Optional[int] = None
|
hidden_size: Optional[int] = None
|
||||||
num_hidden_layers: Optional[int] = None
|
num_hidden_layers: Optional[int] = None
|
||||||
rms_norm_eps: Optional[float] = None
|
rms_norm_eps: Optional[float] = None
|
||||||
intermediate_size: Optional[int] = None
|
intermediate_size: Optional[int] = None
|
||||||
|
|
||||||
max_position_embeddings: Optional[int] = None
|
max_position_embeddings: Optional[int] = None
|
||||||
rope_theta: Optional[float] = None
|
rope_theta: Optional[float] = None
|
||||||
rope_scaling: Optional[dict] = None
|
rope_scaling: Optional[dict] = None
|
||||||
|
|
||||||
attn_type: str = "gqa"
|
attn_type: str = "gqa"
|
||||||
num_attention_heads: Optional[int] = None
|
num_attention_heads: Optional[int] = None
|
||||||
num_key_value_heads: Optional[int] = None
|
num_key_value_heads: Optional[int] = None
|
||||||
use_qk_norm: Optional[bool] = None
|
use_qk_norm: Optional[bool] = None
|
||||||
use_gated_attention: Optional[bool] = None
|
use_gated_attention: Optional[bool] = None
|
||||||
|
|
||||||
ffn_type: str = "mlp"
|
ffn_type: str = "mlp"
|
||||||
pooling_type: Optional[str] = None
|
pooling_type: Optional[str] = None
|
||||||
normalize_embeddings: Optional[bool] = None
|
normalize_embeddings: Optional[bool] = None
|
||||||
|
|
||||||
|
@field_validator("attn_type")
|
||||||
|
def _validate_attn_type(cls, v: str) -> str:
|
||||||
|
if v not in _ATTN_TYPES:
|
||||||
|
raise ValueError(
|
||||||
|
f"attn_type must be one of {sorted(_ATTN_TYPES)}, got {v!r}"
|
||||||
|
)
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("ffn_type")
|
||||||
|
def _validate_ffn_type(cls, v: str) -> str:
|
||||||
|
if v not in _FFN_TYPES:
|
||||||
|
raise ValueError(f"ffn_type must be one of {sorted(_FFN_TYPES)}, got {v!r}")
|
||||||
|
return v
|
||||||
|
|||||||
@@ -5,11 +5,19 @@ modes, both driven declaratively through ``input.sections`` or
|
|||||||
``input.sources``.
|
``input.sources``.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import field
|
||||||
from typing import Dict, List, Optional
|
from typing import Dict, List, Optional
|
||||||
|
|
||||||
|
from pydantic import field_validator
|
||||||
|
from pydantic.dataclasses import dataclass
|
||||||
|
|
||||||
from astrai.config.base import BaseConfig
|
from astrai.config.base import BaseConfig
|
||||||
|
|
||||||
|
_PACKING_STRATEGIES = frozenset({"simple", "bfd", "bfd_split"})
|
||||||
|
_TRUNCATION_MODES = frozenset({"keep_start", "keep_end"})
|
||||||
|
_STORAGE_FORMATS = frozenset({"bin", "jsonl"})
|
||||||
|
_POSITION_IDS_MODES = frozenset({"none", "doc_reset", "continuous"})
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class InputConfig(BaseConfig):
|
class InputConfig(BaseConfig):
|
||||||
@@ -25,6 +33,10 @@ class InputConfig(BaseConfig):
|
|||||||
"chosen": {"sections": [{"field": "chosen", ...}]},
|
"chosen": {"sections": [{"field": "chosen", ...}]},
|
||||||
"rejected": {"sections": [{"field": "rejected", ...}]},
|
"rejected": {"sections": [{"field": "rejected", ...}]},
|
||||||
}}}
|
}}}
|
||||||
|
|
||||||
|
Args:
|
||||||
|
sections (Optional[List[Dict]]): Section list for single-output mode. Defaults to None.
|
||||||
|
sources (Optional[Dict[str, Dict]]): Source map for multi-output mode, DPO/GRPO. Defaults to None.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
sections: Optional[List[Dict]] = None
|
sections: Optional[List[Dict]] = None
|
||||||
@@ -33,34 +45,17 @@ class InputConfig(BaseConfig):
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ProcessingConfig(BaseConfig):
|
class ProcessingConfig(BaseConfig):
|
||||||
"""Processing configuration.
|
"""Processing configuration for tokenization and packing.
|
||||||
|
|
||||||
Parameters
|
Args:
|
||||||
----------
|
max_seq_len (int): Maximum sequence length. Defaults to 2048.
|
||||||
max_seq_len : int
|
min_chars (int): Minimum number of characters to keep. Defaults to 50.
|
||||||
Maximum sequence length (default: 2048).
|
max_chars (int): Maximum number of characters to keep. Defaults to 2_000_000.
|
||||||
min_chars : int
|
max_items (Optional[int]): Maximum number of items to process, None=unlimited. Defaults to None.
|
||||||
Minimum number of characters to keep (default: 50).
|
batch_size (int): Number of records tokenized together. Defaults to 256.
|
||||||
max_chars : int
|
packing_strategy (str): How to pack sequences: 'simple', 'bfd', or 'bfd_split'. Defaults to "simple".
|
||||||
Maximum number of characters to keep (default: 2_000_000).
|
max_packed_len (int): Maximum length of a packed bin. Defaults to 8192.
|
||||||
max_items : Optional[int]
|
truncation_mode (str): How to truncate over-length sequences: 'keep_start' or 'keep_end'. Defaults to "keep_start".
|
||||||
Maximum number of items to process (default: None, unlimited).
|
|
||||||
batch_size : int
|
|
||||||
Number of records tokenized together (default: 256).
|
|
||||||
packing_strategy : str
|
|
||||||
How to pack sequences into a contiguous stream.
|
|
||||||
|
|
||||||
- ``"simple"``: sequential concatenation (default, backward compatible).
|
|
||||||
- ``"bfd"``: best-fit decreasing bin packing, minimises wasted tokens.
|
|
||||||
- ``"bfd_split"``: BFD with over-length sequences split into chunks.
|
|
||||||
max_packed_len : int
|
|
||||||
Maximum length of a packed bin. Sequences longer than this are
|
|
||||||
truncated or split depending on ``packing_strategy`` (default: 8192).
|
|
||||||
truncation_mode : str
|
|
||||||
How to truncate sequences longer than ``max_packed_len``.
|
|
||||||
|
|
||||||
- ``"keep_start"``: keep the first ``max_packed_len`` tokens (default).
|
|
||||||
- ``"keep_end"``: keep the last ``max_packed_len`` tokens.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
max_seq_len: int = 2048
|
max_seq_len: int = 2048
|
||||||
@@ -72,27 +67,45 @@ class ProcessingConfig(BaseConfig):
|
|||||||
max_packed_len: int = 8192
|
max_packed_len: int = 8192
|
||||||
truncation_mode: str = "keep_start"
|
truncation_mode: str = "keep_start"
|
||||||
|
|
||||||
|
@field_validator("packing_strategy")
|
||||||
|
def _validate_packing_strategy(cls, v: str) -> str:
|
||||||
|
if v not in _PACKING_STRATEGIES:
|
||||||
|
raise ValueError(
|
||||||
|
f"packing_strategy must be one of {sorted(_PACKING_STRATEGIES)}, got {v!r}"
|
||||||
|
)
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("truncation_mode")
|
||||||
|
def _validate_truncation_mode(cls, v: str) -> str:
|
||||||
|
if v not in _TRUNCATION_MODES:
|
||||||
|
raise ValueError(
|
||||||
|
f"truncation_mode must be one of {sorted(_TRUNCATION_MODES)}, got {v!r}"
|
||||||
|
)
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("max_seq_len", "batch_size", "max_packed_len")
|
||||||
|
def _validate_positive_int(cls, v: int) -> int:
|
||||||
|
if v <= 0:
|
||||||
|
raise ValueError(f"must be positive, got {v}")
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("min_chars")
|
||||||
|
def _validate_non_negative(cls, v: int) -> int:
|
||||||
|
if v < 0:
|
||||||
|
raise ValueError(f"min_chars must be non-negative, got {v}")
|
||||||
|
return v
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class OutputConfig(BaseConfig):
|
class OutputConfig(BaseConfig):
|
||||||
"""Output configuration.
|
"""Output configuration for storage.
|
||||||
|
|
||||||
Parameters
|
Args:
|
||||||
----------
|
domain_key (Optional[str]): Domain key for the output store. Defaults to None.
|
||||||
domain_key : Optional[str]
|
storage_format (str): Storage format: 'bin' or 'jsonl'. Defaults to "bin".
|
||||||
Domain key for the output store (default: None).
|
max_tokens_per_shard (int): Maximum tokens per shard before splitting. Defaults to 100_000_000.
|
||||||
storage_format : str
|
dtype (Dict[str, str]): Per-key dtype overrides, e.g. {"input_ids": "int32"}. Defaults to {}.
|
||||||
Storage format, one of ``"bin"``, ``"jsonl"`` (default: ``"bin"``).
|
position_ids_mode (str): Position ids mode: 'none', 'doc_reset', or 'continuous'. Defaults to "doc_reset".
|
||||||
max_tokens_per_shard : int
|
|
||||||
Maximum tokens per shard before splitting (default: 100_000_000).
|
|
||||||
dtype : Dict[str, str]
|
|
||||||
Per-key dtype overrides, e.g. ``{"input_ids": "int32"}`` (default: {}).
|
|
||||||
position_ids_mode : Optional[str]
|
|
||||||
How to compute position_ids in packed sequences.
|
|
||||||
|
|
||||||
- ``"none"``: do not generate (default).
|
|
||||||
- ``"doc_reset"``: reset to 0 at each document boundary.
|
|
||||||
- ``"continuous"``: sequential 0, 1, 2, ... (pretrain, single doc).
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
domain_key: Optional[str] = None
|
domain_key: Optional[str] = None
|
||||||
@@ -101,9 +114,36 @@ class OutputConfig(BaseConfig):
|
|||||||
dtype: Dict[str, str] = field(default_factory=dict)
|
dtype: Dict[str, str] = field(default_factory=dict)
|
||||||
position_ids_mode: str = "doc_reset"
|
position_ids_mode: str = "doc_reset"
|
||||||
|
|
||||||
|
@field_validator("storage_format")
|
||||||
|
def _validate_storage_format(cls, v: str) -> str:
|
||||||
|
if v not in _STORAGE_FORMATS:
|
||||||
|
raise ValueError(
|
||||||
|
f"storage_format must be one of {sorted(_STORAGE_FORMATS)}, got {v!r}"
|
||||||
|
)
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("position_ids_mode")
|
||||||
|
def _validate_position_ids_mode(cls, v: str) -> str:
|
||||||
|
if v not in _POSITION_IDS_MODES:
|
||||||
|
raise ValueError(
|
||||||
|
f"position_ids_mode must be one of {sorted(_POSITION_IDS_MODES)}, got {v!r}"
|
||||||
|
)
|
||||||
|
return v
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class PipelineConfig(BaseConfig):
|
class PipelineConfig(BaseConfig):
|
||||||
|
"""Top-level preprocessing pipeline config.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
version (int): Config schema version. Defaults to 1.
|
||||||
|
input (InputConfig): Input mapping config.
|
||||||
|
mask (Dict[str, str]): Per-field mask labels, e.g. {"system": "mask", "assistant": "train"}. Defaults to {}.
|
||||||
|
mask_default (str): Default mask label for unlisted fields. Defaults to "mask".
|
||||||
|
preprocessing (ProcessingConfig): Processing config.
|
||||||
|
output (OutputConfig): Output config.
|
||||||
|
"""
|
||||||
|
|
||||||
version: int = 1
|
version: int = 1
|
||||||
input: InputConfig = field(default_factory=InputConfig)
|
input: InputConfig = field(default_factory=InputConfig)
|
||||||
mask: Dict[str, str] = field(default_factory=dict)
|
mask: Dict[str, str] = field(default_factory=dict)
|
||||||
|
|||||||
+191
-158
@@ -1,7 +1,9 @@
|
|||||||
from dataclasses import dataclass, field, fields
|
from dataclasses import field
|
||||||
from typing import Any, Callable, Dict, List, Optional
|
from typing import Any, Callable, Dict, List, Optional
|
||||||
|
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
from pydantic import ConfigDict, field_validator, model_validator
|
||||||
|
from pydantic.dataclasses import dataclass
|
||||||
from torch.optim import Optimizer
|
from torch.optim import Optimizer
|
||||||
from torch.optim.lr_scheduler import LRScheduler
|
from torch.optim.lr_scheduler import LRScheduler
|
||||||
from torch.utils.data import Dataset
|
from torch.utils.data import Dataset
|
||||||
@@ -9,173 +11,204 @@ from torch.utils.data import Dataset
|
|||||||
from astrai.config.base import BaseConfig
|
from astrai.config.base import BaseConfig
|
||||||
from astrai.model.components.lora import LoRAConfig
|
from astrai.model.components.lora import LoRAConfig
|
||||||
|
|
||||||
|
_TRAIN_TYPES = frozenset({"seq", "sft", "dpo", "grpo", "online_grpo", "online_dpo"})
|
||||||
def required(**kw):
|
_PARALLEL_MODES = frozenset({"none", "ddp", "fsdp"})
|
||||||
return {"required": True, **kw}
|
_BACKENDS = frozenset({"nccl", "gloo"})
|
||||||
|
_START_METHODS = frozenset({"spawn", "fork", "forkserver"})
|
||||||
|
_COMPILE_MODES = frozenset({"default", "reduce-overhead", "max-autotune"})
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass(config=ConfigDict(arbitrary_types_allowed=True))
|
||||||
class TrainConfig(BaseConfig):
|
class TrainConfig(BaseConfig):
|
||||||
# basic setting
|
"""Training configuration.
|
||||||
model_fn: Callable[[], nn.Module] = field(
|
|
||||||
default=None, metadata=required(help="Model factory for training.")
|
|
||||||
)
|
|
||||||
strategy: str = field(default=None, metadata=required(help="Training strategy."))
|
|
||||||
dataset: Dataset = field(
|
|
||||||
default=None, metadata=required(help="Dataset for training.")
|
|
||||||
)
|
|
||||||
optimizer_fn: Callable[[nn.Module], Optimizer] = field(
|
|
||||||
default=None, metadata=required(help="Optimizer factory for training.")
|
|
||||||
)
|
|
||||||
scheduler_fn: Callable[[Optimizer], LRScheduler] = field(
|
|
||||||
default=None, metadata=required(help="Scheduler factory for training.")
|
|
||||||
)
|
|
||||||
n_epoch: int = field(default=1, metadata={"help": "Number of epochs for training."})
|
|
||||||
batch_per_device: int = field(
|
|
||||||
default=4, metadata={"help": "Batch size per device."}
|
|
||||||
)
|
|
||||||
grad_accum_steps: int = field(
|
|
||||||
default=1, metadata={"help": "Number of iterations between steps."}
|
|
||||||
)
|
|
||||||
max_grad_norm: Optional[float] = field(
|
|
||||||
default=1.0,
|
|
||||||
metadata={"help": "Maximum gradient norm. None disables clipping."},
|
|
||||||
)
|
|
||||||
gradient_checkpointing_modules: List[str] = field(
|
|
||||||
default_factory=list,
|
|
||||||
metadata={"help": "Module types to enable activation checkpointing for."},
|
|
||||||
)
|
|
||||||
|
|
||||||
# checkpoint setting
|
Combines hyperparameters with runtime objects (model_fn, dataset, etc.).
|
||||||
start_epoch: int = field(default=0, metadata={"help": "Start epoch for training."})
|
Only JSON-serializable fields are written to checkpoint meta via to_dict().
|
||||||
start_samples: int = field(
|
|
||||||
default=0,
|
|
||||||
metadata={
|
|
||||||
"help": "Start samples count (per rank). Superseded by checkpoint consumed_samples."
|
|
||||||
},
|
|
||||||
)
|
|
||||||
ckpt_dir: str = field(
|
|
||||||
default="./checkpoint", metadata={"help": "Checkpoint directory."}
|
|
||||||
)
|
|
||||||
ckpt_interval: int = field(
|
|
||||||
default=5000,
|
|
||||||
metadata={"help": "Number of optimizer steps between checkpoints."},
|
|
||||||
)
|
|
||||||
|
|
||||||
# lora setting
|
Args:
|
||||||
lora: Optional[LoRAConfig] = field(
|
model_fn (Callable[[], nn.Module]): Model factory for training.
|
||||||
default=None,
|
strategy (str): Training strategy (seq, sft, dpo, grpo, online_*).
|
||||||
metadata={"help": "LoRA config. None means full fine-tuning."},
|
dataset (Dataset): Dataset for training.
|
||||||
)
|
optimizer_fn (Callable[[nn.Module], Optimizer]): Optimizer factory for training.
|
||||||
|
optimizer_name (Optional[str]): Serializable built-in optimizer identifier. Defaults to None.
|
||||||
|
optimizer_hyperparameters (Dict[str, Any]): Serializable optimizer settings. Defaults to {}.
|
||||||
|
scheduler_fn (Callable[[Optimizer], LRScheduler]): Scheduler factory for training.
|
||||||
|
n_epoch (int): Number of epochs for training. Defaults to 1.
|
||||||
|
batch_per_device (int): Batch size per device. Defaults to 4.
|
||||||
|
grad_accum_steps (int): Number of iterations between optimizer steps. Defaults to 1.
|
||||||
|
max_grad_norm (Optional[float]): Maximum gradient norm. None disables clipping. Defaults to 1.0.
|
||||||
|
gradient_checkpointing_modules (List[type]): Module types to enable activation checkpointing for. Defaults to [].
|
||||||
|
compile_mode (Optional[str]): torch.compile mode: 'default', 'reduce-overhead', 'max-autotune', or None. Defaults to None.
|
||||||
|
start_epoch (int): Start epoch for training. Defaults to 0.
|
||||||
|
start_samples (int): Start samples count (per rank). Superseded by checkpoint consumed_samples. Defaults to 0.
|
||||||
|
ckpt_dir (str): Checkpoint directory. Defaults to "./checkpoint".
|
||||||
|
ckpt_interval (int): Number of optimizer steps between checkpoints. Defaults to 5000.
|
||||||
|
lora (Optional[LoRAConfig]): LoRA config. None means full fine-tuning. Defaults to None.
|
||||||
|
metrics (List[str]): Metrics to record during training. Defaults to ["loss", "lr", "grad_norm"].
|
||||||
|
random_seed (int): Random seed. Defaults to 3407.
|
||||||
|
num_workers (int): Number of workers for dataloader. Defaults to 0.
|
||||||
|
prefetch_factor (Optional[int]): Prefetch factor for dataloader. Defaults to None.
|
||||||
|
pin_memory (bool): Pin memory for dataloader. Defaults to False.
|
||||||
|
collate_fn (Optional[Callable[[List[Any]], Any]]): Collate function for dataloader (e.g. dpo_collate_fn). Defaults to None.
|
||||||
|
nprocs (int): Number of processes for distributed training. Defaults to 1.
|
||||||
|
backend (str): Distributed training backend. Defaults to "nccl".
|
||||||
|
master_addr (str): Master address for distributed training. Defaults to "localhost".
|
||||||
|
master_port (str): Master port for distributed training. Defaults to "29500".
|
||||||
|
parallel_mode (str): Parallel strategy: none, ddp, fsdp. Defaults to "none".
|
||||||
|
start_method (str): Multiprocessing start method: spawn/fork/forkserver. Defaults to "spawn".
|
||||||
|
device_type (str): Device type for distributed training. Defaults to "cuda".
|
||||||
|
val_dataset (Optional[Dataset]): Dataset for validation. Defaults to None.
|
||||||
|
val_split (Optional[float]): Ratio to split from training dataset for validation, e.g. 0.05. Defaults to None.
|
||||||
|
val_step (int): Number of optimizer steps between validation runs. Defaults to 1000.
|
||||||
|
neftune_alpha (float): NEFTune noise alpha, 0=disabled, typical: 5.0. Defaults to 0.0.
|
||||||
|
rollout_interval (int): Number of optimizer steps between online rollouts. Defaults to 512.
|
||||||
|
rollout_temperature (float): Sampling temperature for online rollout. Defaults to 0.7.
|
||||||
|
rollout_top_k (int): Top-k filtering for online rollout, 0=disable. Defaults to 0.
|
||||||
|
rollout_top_p (float): Top-p (nucleus) filtering for online rollout. Defaults to 0.9.
|
||||||
|
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 {}.
|
||||||
|
"""
|
||||||
|
|
||||||
# metric setting
|
model_fn: Callable[[], nn.Module]
|
||||||
log_dir: str = field(
|
strategy: str
|
||||||
default="./checkpoint/logs", metadata={"help": "Directory for metric logs."}
|
dataset: Dataset
|
||||||
)
|
optimizer_fn: Callable[[nn.Module], Optimizer]
|
||||||
metrics: List[str] = field(
|
scheduler_fn: Callable[[Optimizer], LRScheduler]
|
||||||
default_factory=lambda: ["loss", "lr", "grad_norm"],
|
optimizer_name: Optional[str] = None
|
||||||
metadata={"help": "Metrics to record during training."},
|
optimizer_hyperparameters: Dict[str, Any] = field(default_factory=dict)
|
||||||
)
|
n_epoch: int = 1
|
||||||
|
batch_per_device: int = 4
|
||||||
|
grad_accum_steps: int = 1
|
||||||
|
max_grad_norm: Optional[float] = 1.0
|
||||||
|
gradient_checkpointing_modules: List[type] = field(default_factory=list)
|
||||||
|
compile_mode: Optional[str] = None
|
||||||
|
|
||||||
# dataloader setting
|
start_epoch: int = 0
|
||||||
random_seed: int = field(default=3407, metadata={"help": "Random seed."})
|
start_samples: int = 0
|
||||||
num_workers: int = field(
|
ckpt_dir: str = "./checkpoint"
|
||||||
default=0, metadata={"help": "Number of workers for dataloader."}
|
ckpt_interval: int = 5000
|
||||||
)
|
|
||||||
prefetch_factor: Optional[int] = field(
|
|
||||||
default=None, metadata={"help": "Prefetch factor for dataloader."}
|
|
||||||
)
|
|
||||||
pin_memory: bool = field(
|
|
||||||
default=False, metadata={"help": "Pin memory for dataloader."}
|
|
||||||
)
|
|
||||||
collate_fn: Optional[Callable[[List[Any]], Any]] = field(
|
|
||||||
default=None,
|
|
||||||
metadata={"help": "Collate function for dataloader (e.g. dpo_collate_fn)."},
|
|
||||||
)
|
|
||||||
|
|
||||||
# distributed training
|
lora: Optional[LoRAConfig] = None
|
||||||
nprocs: int = field(
|
|
||||||
default=1, metadata={"help": "Number of processes for distributed training."}
|
|
||||||
)
|
|
||||||
backend: str = field(
|
|
||||||
default="nccl", metadata={"help": "Distributed training backend."}
|
|
||||||
)
|
|
||||||
master_addr: str = field(
|
|
||||||
default="localhost",
|
|
||||||
metadata={"help": "Master address for distributed training."},
|
|
||||||
)
|
|
||||||
master_port: str = field(
|
|
||||||
default="29500", metadata={"help": "Master port for distributed training."}
|
|
||||||
)
|
|
||||||
parallel_mode: str = field(
|
|
||||||
default="none",
|
|
||||||
metadata={"help": "Parallel strategy: none, ddp, fsdp."},
|
|
||||||
)
|
|
||||||
start_method: str = field(
|
|
||||||
default="spawn",
|
|
||||||
metadata={"help": "Multiprocessing start method (spawn/fork/forkserver)."},
|
|
||||||
)
|
|
||||||
|
|
||||||
# others
|
metrics: List[str] = field(default_factory=lambda: ["loss", "lr", "grad_norm"])
|
||||||
device_type: str = field(
|
|
||||||
default="cuda", metadata={"help": "Device type for distributed training."}
|
|
||||||
)
|
|
||||||
val_dataset: Optional[Dataset] = field(
|
|
||||||
default=None, metadata={"help": "Dataset for validation."}
|
|
||||||
)
|
|
||||||
val_split: Optional[float] = field(
|
|
||||||
default=None,
|
|
||||||
metadata={
|
|
||||||
"help": "Ratio to split from training dataset for validation (e.g. 0.05). Ignored if val_dataset is set."
|
|
||||||
},
|
|
||||||
)
|
|
||||||
val_step: int = field(
|
|
||||||
default=1000,
|
|
||||||
metadata={"help": "Number of optimizer steps between validation runs."},
|
|
||||||
)
|
|
||||||
neftune_alpha: float = field(
|
|
||||||
default=0.0,
|
|
||||||
metadata={"help": "NEFTune noise alpha (0=disabled, typical: 5.0)."},
|
|
||||||
)
|
|
||||||
|
|
||||||
# online rollout
|
random_seed: int = 3407
|
||||||
rollout_interval: int = field(
|
num_workers: int = 0
|
||||||
default=512,
|
prefetch_factor: Optional[int] = None
|
||||||
metadata={"help": "Number of optimizer steps between online rollouts."},
|
pin_memory: bool = False
|
||||||
)
|
collate_fn: Optional[Callable[[List[Any]], Any]] = None
|
||||||
rollout_temperature: float = field(
|
|
||||||
default=0.7, metadata={"help": "Sampling temperature for online rollout."}
|
|
||||||
)
|
|
||||||
rollout_top_k: int = field(
|
|
||||||
default=0, metadata={"help": "Top-k filtering for online rollout (0=disable)."}
|
|
||||||
)
|
|
||||||
rollout_top_p: float = field(
|
|
||||||
default=0.9,
|
|
||||||
metadata={"help": "Top-p (nucleus) filtering for online rollout."},
|
|
||||||
)
|
|
||||||
rollout_max_tokens: int = field(
|
|
||||||
default=1024,
|
|
||||||
metadata={"help": "Maximum generated tokens per response in rollout."},
|
|
||||||
)
|
|
||||||
reward_model_fn: Optional[Callable] = field(
|
|
||||||
default=None,
|
|
||||||
metadata={
|
|
||||||
"help": "Factory for reward model (required for online RL strategies)."
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
executor_kwargs: Dict[str, Any] = field(
|
nprocs: int = 1
|
||||||
default_factory=dict,
|
backend: str = "nccl"
|
||||||
metadata={"help": "Extra kwargs passed to ExecutorFactory.create()."},
|
master_addr: str = "localhost"
|
||||||
)
|
master_port: str = "29500"
|
||||||
extra_kwargs: Dict[str, Any] = field(
|
parallel_mode: str = "none"
|
||||||
default_factory=dict, metadata={"help": "Other arguments."}
|
start_method: str = "spawn"
|
||||||
)
|
|
||||||
|
|
||||||
def __post_init__(self):
|
device_type: str = "cuda"
|
||||||
self.validate()
|
val_dataset: Optional[Dataset] = None
|
||||||
|
val_split: Optional[float] = None
|
||||||
|
val_step: int = 1000
|
||||||
|
neftune_alpha: float = 0.0
|
||||||
|
|
||||||
def validate(self):
|
rollout_interval: int = 512
|
||||||
for fld in fields(self):
|
rollout_temperature: float = 0.7
|
||||||
if fld.metadata.get("required") and getattr(self, fld.name) is None:
|
rollout_top_k: int = 0
|
||||||
raise ValueError(f"TrainConfig.{fld.name} is required but got None.")
|
rollout_top_p: float = 0.9
|
||||||
|
rollout_max_tokens: int = 1024
|
||||||
|
reward_model_fn: Optional[Callable] = None
|
||||||
|
|
||||||
|
executor_kwargs: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
extra_kwargs: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
@field_validator("strategy")
|
||||||
|
def _validate_strategy(cls, v: str) -> str:
|
||||||
|
if v not in _TRAIN_TYPES:
|
||||||
|
raise ValueError(
|
||||||
|
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:
|
||||||
|
raise ValueError(
|
||||||
|
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}")
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("start_method")
|
||||||
|
def _validate_start_method(cls, v: str) -> str:
|
||||||
|
if v not in _START_METHODS:
|
||||||
|
raise ValueError(
|
||||||
|
f"start_method must be one of {sorted(_START_METHODS)}, got {v!r}"
|
||||||
|
)
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("compile_mode")
|
||||||
|
def _validate_compile_mode(cls, v: Optional[str]) -> Optional[str]:
|
||||||
|
if v is not None and v not in _COMPILE_MODES:
|
||||||
|
raise ValueError(
|
||||||
|
f"compile_mode must be one of {sorted(_COMPILE_MODES)} or None, got {v!r}"
|
||||||
|
)
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator(
|
||||||
|
"n_epoch",
|
||||||
|
"batch_per_device",
|
||||||
|
"grad_accum_steps",
|
||||||
|
"ckpt_interval",
|
||||||
|
"val_step",
|
||||||
|
"rollout_interval",
|
||||||
|
"rollout_max_tokens",
|
||||||
|
)
|
||||||
|
def _validate_positive_int(cls, v: int) -> int:
|
||||||
|
if v <= 0:
|
||||||
|
raise ValueError(f"must be positive, got {v}")
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("rollout_temperature")
|
||||||
|
def _validate_positive_float(cls, v: float) -> float:
|
||||||
|
if v <= 0:
|
||||||
|
raise ValueError(f"must be positive, got {v}")
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("rollout_top_p")
|
||||||
|
def _validate_top_p(cls, v: float) -> float:
|
||||||
|
if not 0 < v <= 1:
|
||||||
|
raise ValueError(f"rollout_top_p must be in (0, 1], got {v}")
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("rollout_top_k", "num_workers", "neftune_alpha")
|
||||||
|
def _validate_non_negative(cls, v):
|
||||||
|
if v < 0:
|
||||||
|
raise ValueError(f"must be non-negative, got {v}")
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("max_grad_norm")
|
||||||
|
def _validate_max_grad_norm(cls, v: Optional[float]) -> Optional[float]:
|
||||||
|
if v is not None and v <= 0:
|
||||||
|
raise ValueError(f"max_grad_norm must be positive or None, got {v}")
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("val_split")
|
||||||
|
def _validate_val_split(cls, v: Optional[float]) -> Optional[float]:
|
||||||
|
if v is not None and not 0 < v < 1:
|
||||||
|
raise ValueError(f"val_split must be in (0, 1) or None, got {v}")
|
||||||
|
return v
|
||||||
|
|
||||||
|
@model_validator(mode="after")
|
||||||
|
def _validate_online_strategy(self) -> "TrainConfig":
|
||||||
|
if self.strategy.startswith("online_") and self.reward_model_fn is None:
|
||||||
|
raise ValueError(
|
||||||
|
f"reward_model_fn is required for online RL strategy {self.strategy!r}"
|
||||||
|
)
|
||||||
|
return self
|
||||||
|
|||||||
@@ -6,7 +6,6 @@ from astrai.dataset.dataset import (
|
|||||||
)
|
)
|
||||||
from astrai.dataset.sampler import RDSampler
|
from astrai.dataset.sampler import RDSampler
|
||||||
from astrai.dataset.storage import (
|
from astrai.dataset.storage import (
|
||||||
H5Store,
|
|
||||||
JsonlStore,
|
JsonlStore,
|
||||||
MmapStore,
|
MmapStore,
|
||||||
Recordable,
|
Recordable,
|
||||||
@@ -17,9 +16,7 @@ from astrai.dataset.storage import (
|
|||||||
)
|
)
|
||||||
from astrai.serialization import (
|
from astrai.serialization import (
|
||||||
load_bin,
|
load_bin,
|
||||||
load_h5,
|
|
||||||
save_bin,
|
save_bin,
|
||||||
save_h5,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
@@ -31,12 +28,9 @@ __all__ = [
|
|||||||
"Streamable",
|
"Streamable",
|
||||||
"Recordable",
|
"Recordable",
|
||||||
"StoreFactory",
|
"StoreFactory",
|
||||||
"H5Store",
|
|
||||||
"MmapStore",
|
"MmapStore",
|
||||||
"JsonlStore",
|
"JsonlStore",
|
||||||
"detect_format",
|
"detect_format",
|
||||||
"save_h5",
|
|
||||||
"load_h5",
|
|
||||||
"save_bin",
|
"save_bin",
|
||||||
"load_bin",
|
"load_bin",
|
||||||
"RDSampler",
|
"RDSampler",
|
||||||
|
|||||||
@@ -314,7 +314,7 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
|
|||||||
stream datasets (SEQ/SFT). Record datasets ignore it.
|
stream datasets (SEQ/SFT). Record datasets ignore it.
|
||||||
stride: Stride between consecutive stream samples
|
stride: Stride between consecutive stream samples
|
||||||
(default: same as *window_size*).
|
(default: same as *window_size*).
|
||||||
storage_type: Storage backend ("h5", "bin", "jsonl") or
|
storage_type: Storage backend ("bin", "jsonl") or
|
||||||
None for auto-detection.
|
None for auto-detection.
|
||||||
tokenizer_path: Path to tokenizer for lazy JSONL
|
tokenizer_path: Path to tokenizer for lazy JSONL
|
||||||
tokenisation (record datasets only).
|
tokenisation (record datasets only).
|
||||||
@@ -384,7 +384,7 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
|
|||||||
"""Build an on-the-fly tokenisation processor if applicable.
|
"""Build an on-the-fly tokenisation processor if applicable.
|
||||||
|
|
||||||
Only raw JSONL + record datasets (DPO/GRPO) need a processor;
|
Only raw JSONL + record datasets (DPO/GRPO) need a processor;
|
||||||
pre-tokenised backends (H5/bin) and stream datasets (SEQ/SFT)
|
pre-tokenised backends (bin) and stream datasets (SEQ/SFT)
|
||||||
return ``None`` so no tokenizer is loaded.
|
return ``None`` so no tokenizer is loaded.
|
||||||
"""
|
"""
|
||||||
if tokenizer_path is None or storage_type != "jsonl":
|
if tokenizer_path is None or storage_type != "jsonl":
|
||||||
@@ -451,7 +451,7 @@ class DPODataset(BaseDataset):
|
|||||||
|
|
||||||
Two loading paths (handled by :class:`DatasetFactory`):
|
Two loading paths (handled by :class:`DatasetFactory`):
|
||||||
|
|
||||||
- **Pre-tokenized** (H5/bin): ``store.load(path)`` reads per-record
|
- **Pre-tokenized** (bin): ``store.load(path)`` reads per-record
|
||||||
tensors; ``__getitem__`` returns them directly.
|
tensors; ``__getitem__`` returns them directly.
|
||||||
- **Raw JSONL** (``tokenizer_path=...``): builds a lazy processor
|
- **Raw JSONL** (``tokenizer_path=...``): builds a lazy processor
|
||||||
via :func:`dpo_processor` that tokenises on the fly — no packing,
|
via :func:`dpo_processor` that tokenises on the fly — no packing,
|
||||||
|
|||||||
@@ -10,7 +10,6 @@ Architecture (composition over inheritance):
|
|||||||
Streamable (mixin) — raw token slice fetch(begin, end, keys)
|
Streamable (mixin) — raw token slice fetch(begin, end, keys)
|
||||||
Recordable (mixin) — raw record slice fetch_record(idx, keys)
|
Recordable (mixin) — raw record slice fetch_record(idx, keys)
|
||||||
|
|
||||||
H5Store(Store, Streamable, Recordable)
|
|
||||||
MmapStore(Store, Streamable, Recordable)
|
MmapStore(Store, Streamable, Recordable)
|
||||||
JsonlStore(Store, Streamable, Recordable)
|
JsonlStore(Store, Streamable, Recordable)
|
||||||
|
|
||||||
@@ -36,9 +35,9 @@ control. ``store.token_count`` is the total stream token count (what
|
|||||||
``len(store)`` used to mean in the legacy stream-only API).
|
``len(store)`` used to mean in the legacy stream-only API).
|
||||||
|
|
||||||
``segments_are_records`` (class attribute on each Store subclass)
|
``segments_are_records`` (class attribute on each Store subclass)
|
||||||
tells ``_normalize`` whether segments are inherently per-record (H5/
|
tells ``_normalize`` whether segments are inherently per-record (JSONL)
|
||||||
JSONL) or opaque shards (bin). Record access for bin relies on
|
or opaque shards (bin). Record access for bin relies on ``_offsets``
|
||||||
``_offsets`` instead.
|
instead.
|
||||||
|
|
||||||
:class:`JsonlStore` supports a lazy mode (``processor=fn``) that keeps
|
:class:`JsonlStore` supports a lazy mode (``processor=fn``) that keeps
|
||||||
raw records and defers tokenisation to ``fetch_record`` — used by DPO
|
raw records and defers tokenisation to ``fetch_record`` — used by DPO
|
||||||
@@ -62,7 +61,6 @@ from astrai.preprocessing.transform import TokenizeTransform
|
|||||||
from astrai.serialization import (
|
from astrai.serialization import (
|
||||||
load_bin,
|
load_bin,
|
||||||
load_bin_offsets,
|
load_bin_offsets,
|
||||||
load_h5,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -83,19 +81,10 @@ def detect_format(load_path: str) -> str:
|
|||||||
root = Path(load_path)
|
root = Path(load_path)
|
||||||
if root.is_file():
|
if root.is_file():
|
||||||
suffix = root.suffix.lower()
|
suffix = root.suffix.lower()
|
||||||
if suffix in (".h5", ".hdf5"):
|
|
||||||
return "h5"
|
|
||||||
if suffix == ".jsonl":
|
if suffix == ".jsonl":
|
||||||
return "jsonl"
|
return "jsonl"
|
||||||
raise ValueError(f"Unsupported file format: {suffix}")
|
raise ValueError(f"Unsupported file format: {suffix}")
|
||||||
|
|
||||||
h5_files = [
|
|
||||||
Path(p)
|
|
||||||
for pattern in ("*.h5", "*.hdf5")
|
|
||||||
for p in glob.glob(str(root / "**" / pattern), recursive=True)
|
|
||||||
]
|
|
||||||
if h5_files:
|
|
||||||
return "h5"
|
|
||||||
bin_files = [Path(p) for p in glob.glob(str(root / "**" / "*.bin"), recursive=True)]
|
bin_files = [Path(p) for p in glob.glob(str(root / "**" / "*.bin"), recursive=True)]
|
||||||
if bin_files:
|
if bin_files:
|
||||||
has_meta = (root / "meta.json").exists() or len(
|
has_meta = (root / "meta.json").exists() or len(
|
||||||
@@ -185,7 +174,7 @@ class Store(ABC):
|
|||||||
"""Number of records available via :meth:`fetch_record`.
|
"""Number of records available via :meth:`fetch_record`.
|
||||||
|
|
||||||
Non-zero only when the backing layout provides per-record
|
Non-zero only when the backing layout provides per-record
|
||||||
indexing (H5/JSONL segments or bin ``_offsets``).
|
indexing (JSONL segments or bin ``_offsets``).
|
||||||
"""
|
"""
|
||||||
return self._num_records
|
return self._num_records
|
||||||
|
|
||||||
@@ -269,7 +258,7 @@ class Store(ABC):
|
|||||||
Record mode: if *offsets* is provided (bin layout),
|
Record mode: if *offsets* is provided (bin layout),
|
||||||
``_offsets[key]`` stores cumulative per-record offsets into the
|
``_offsets[key]`` stores cumulative per-record offsets into the
|
||||||
single concatenated segment. Otherwise, when
|
single concatenated segment. Otherwise, when
|
||||||
``segments_are_records`` is True (H5/JSONL), ``_data[key]`` is
|
``segments_are_records`` is True (JSONL), ``_data[key]`` is
|
||||||
a per-record list and ``fetch_record`` indexes it directly.
|
a per-record list and ``fetch_record`` indexes it directly.
|
||||||
|
|
||||||
Nested keys (GRPO ``responses``/``masks`` as
|
Nested keys (GRPO ``responses``/``masks`` as
|
||||||
@@ -305,7 +294,7 @@ class Store(ABC):
|
|||||||
logger.warning(
|
logger.warning(
|
||||||
"Key '%s' has %d segments with offsets — record mode "
|
"Key '%s' has %d segments with offsets — record mode "
|
||||||
"disabled for this key (multi-shard bin+offsets not "
|
"disabled for this key (multi-shard bin+offsets not "
|
||||||
"supported). Merge shards or use H5/JSONL.",
|
"supported). Merge shards or use JSONL.",
|
||||||
key,
|
key,
|
||||||
len(segs),
|
len(segs),
|
||||||
)
|
)
|
||||||
@@ -330,7 +319,7 @@ class Streamable:
|
|||||||
Stateless trait relying on ``self._data``, ``self._cum``,
|
Stateless trait relying on ``self._data``, ``self._cum``,
|
||||||
``self._length`` maintained by :class:`Store`. Stream mode is
|
``self._length`` maintained by :class:`Store`. Stream mode is
|
||||||
active when the owning store has ``window_size > 0``; for stores
|
active when the owning store has ``window_size > 0``; for stores
|
||||||
that can also serve record access (H5/JSONL/bin+offsets), the
|
that can also serve record access (JSONL/bin+offsets), the
|
||||||
``fetch_record`` API from :class:`Recordable` is used instead.
|
``fetch_record`` API from :class:`Recordable` is used instead.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@@ -415,33 +404,6 @@ class StoreFactory(BaseFactory["Store"]):
|
|||||||
"""Factory for creating Store instances by type name."""
|
"""Factory for creating Store instances by type name."""
|
||||||
|
|
||||||
|
|
||||||
@StoreFactory.register("h5")
|
|
||||||
class H5Store(Store, Streamable, Recordable):
|
|
||||||
"""HDF5-based storage backend (pre-tokenized data).
|
|
||||||
|
|
||||||
Each key is stored as a group of per-record datasets (``data_0``,
|
|
||||||
``data_1``, …). Supports both access modes:
|
|
||||||
|
|
||||||
- **Stream**: ``fetch(begin, end, key)`` and ``store[i]`` slice
|
|
||||||
across concatenated records via ``_cum`` — used by SEQ/SFT.
|
|
||||||
- **Record**: ``fetch_record(i, key)`` and ``store[i]`` (when
|
|
||||||
``window_size == 0``) index ``_data[key]`` directly — used by
|
|
||||||
DPO/GRPO.
|
|
||||||
"""
|
|
||||||
|
|
||||||
segments_are_records = True
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
window_size: int = 0,
|
|
||||||
stride: Optional[int] = None,
|
|
||||||
):
|
|
||||||
super().__init__(window_size=window_size, stride=stride)
|
|
||||||
|
|
||||||
def load(self, path: str, **kwargs):
|
|
||||||
self._normalize(load_h5(path))
|
|
||||||
|
|
||||||
|
|
||||||
@StoreFactory.register("bin")
|
@StoreFactory.register("bin")
|
||||||
class MmapStore(Store, Streamable, Recordable):
|
class MmapStore(Store, Streamable, Recordable):
|
||||||
"""Memory-mapped binary storage backend.
|
"""Memory-mapped binary storage backend.
|
||||||
|
|||||||
@@ -4,27 +4,46 @@ Public API:
|
|||||||
- ``attn_decode`` — single-query decode attention
|
- ``attn_decode`` — single-query decode attention
|
||||||
- ``attn_prefill`` — multi-query prefill attention
|
- ``attn_prefill`` — multi-query prefill attention
|
||||||
- ``attn_paged_decode`` — paged decode attention (direct page-table access)
|
- ``attn_paged_decode`` — paged decode attention (direct page-table access)
|
||||||
|
- ``AttentionBackend`` — ABC for attention computation strategies
|
||||||
|
- ``TorchNativeBackend`` — default SDPA backend with KV cache I/O
|
||||||
|
- ``CudaBackend`` — CUDA kernel backend with paged decode + prefill
|
||||||
|
|
||||||
Interface (shared by all wrappers):
|
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
|
||||||
causal_offset: -1 = non-causal; >=0 = absolute position of first Q token
|
(blhd). Scale is always ``1/sqrt(head_dim)``.
|
||||||
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True = keep)
|
|
||||||
scale: 0.0 = auto (1/sqrt(head_dim)); >0 = explicit
|
|
||||||
layout: "bhld" (default) or "blhd"
|
|
||||||
|
|
||||||
Causal and mask can coexist — both are applied simultaneously.
|
Each wrapper calls its compiled CUDA kernel directly. Fallback to torch
|
||||||
|
SDPA is handled by the attention backend, not the wrapper functions.
|
||||||
Each wrapper dispatches to its compiled CUDA kernel (``astrai.extension.attn_*``)
|
|
||||||
when available, otherwise falls back to ``torch.nn.functional.scaled_dot_product_attention``.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from astrai.extension.attention_backend import (
|
||||||
|
ATTN_BACKEND,
|
||||||
|
AttentionBackend,
|
||||||
|
CudaBackend,
|
||||||
|
TorchNativeBackend,
|
||||||
|
attention,
|
||||||
|
attn_backend,
|
||||||
|
get_backend,
|
||||||
|
)
|
||||||
|
from astrai.extension.attention_ops import (
|
||||||
|
attn_decode,
|
||||||
|
attn_paged_decode,
|
||||||
|
attn_prefill,
|
||||||
|
)
|
||||||
from astrai.extension.loader import KERNEL_NAMES, is_available
|
from astrai.extension.loader import KERNEL_NAMES, is_available
|
||||||
from astrai.extension.ops import attention, attn_decode, attn_paged_decode, attn_prefill
|
from astrai.extension.rotary_backend import apply_rotary_emb
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
|
"ATTN_BACKEND",
|
||||||
|
"AttentionBackend",
|
||||||
|
"CudaBackend",
|
||||||
|
"TorchNativeBackend",
|
||||||
|
"attention",
|
||||||
|
"attn_backend",
|
||||||
|
"get_backend",
|
||||||
"attn_decode",
|
"attn_decode",
|
||||||
"attn_paged_decode",
|
"attn_paged_decode",
|
||||||
"attn_prefill",
|
"attn_prefill",
|
||||||
"attention",
|
|
||||||
"is_available",
|
"is_available",
|
||||||
"KERNEL_NAMES",
|
"KERNEL_NAMES",
|
||||||
|
"apply_rotary_emb",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -0,0 +1,422 @@
|
|||||||
|
"""Attention backend abstraction with context-manager switching.
|
||||||
|
|
||||||
|
The backend encapsulates KV cache I/O and attention computation. The
|
||||||
|
attention module (GQA/MLA) keeps projections, rotary, QK-norm, gating,
|
||||||
|
and output projection; the backend handles everything from "write K/V
|
||||||
|
to cache" through "SDPA output".
|
||||||
|
|
||||||
|
Usage — mirroring ``torch.nn.attention.sdpa_kernel``:
|
||||||
|
|
||||||
|
from astrai.extension import attn_backend, ATTN_BACKEND
|
||||||
|
|
||||||
|
with attn_backend(ATTN_BACKEND.TORCH_NATIVE):
|
||||||
|
engine.generate("hello")
|
||||||
|
|
||||||
|
# or with an instance:
|
||||||
|
with attn_backend(TorchNativeBackend()):
|
||||||
|
...
|
||||||
|
|
||||||
|
# or the shorthand (instance is itself a context manager):
|
||||||
|
with TorchNativeBackend():
|
||||||
|
...
|
||||||
|
|
||||||
|
Thread-safe via ``contextvars`` — each scheduler thread gets its own
|
||||||
|
active backend. ``get_backend()`` returns the active one, falling back
|
||||||
|
to a process-wide ``TorchNativeBackend`` singleton.
|
||||||
|
|
||||||
|
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]``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import contextvars
|
||||||
|
import enum
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from typing import Optional, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.extension.attention_ops import attn_paged_decode, attn_prefill
|
||||||
|
from astrai.extension.loader import is_available
|
||||||
|
from astrai.inference.core.cache import KVCache
|
||||||
|
|
||||||
|
_current_backend: contextvars.ContextVar["AttentionBackend"] = contextvars.ContextVar(
|
||||||
|
"attn_backend"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class ATTN_BACKEND(enum.Enum):
|
||||||
|
"""Backend selector enum, mirroring ``torch.nn.attention.SDPBackend``."""
|
||||||
|
|
||||||
|
TORCH_NATIVE = "torch_native"
|
||||||
|
CUDA = "cuda"
|
||||||
|
|
||||||
|
|
||||||
|
def get_backend() -> "AttentionBackend":
|
||||||
|
"""Return the active backend for the current thread/context.
|
||||||
|
|
||||||
|
Falls back to a ``TorchNativeBackend`` singleton when no backend
|
||||||
|
has been activated via ``with``.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
return _current_backend.get()
|
||||||
|
except LookupError:
|
||||||
|
return _default_backend
|
||||||
|
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def attn_backend(backend: Union[ATTN_BACKEND, "AttentionBackend", type]):
|
||||||
|
"""Context manager to select an attention backend.
|
||||||
|
|
||||||
|
Mirrors ``torch.nn.attention.sdpa_kernel``. Accepts an
|
||||||
|
``ATTN_BACKEND`` enum value, a backend class, or a backend instance.
|
||||||
|
|
||||||
|
Examples::
|
||||||
|
|
||||||
|
with attn_backend(ATTN_BACKEND.TORCH_NATIVE):
|
||||||
|
...
|
||||||
|
with attn_backend(TorchNativeBackend):
|
||||||
|
...
|
||||||
|
with attn_backend(TorchNativeBackend()):
|
||||||
|
...
|
||||||
|
"""
|
||||||
|
if isinstance(backend, ATTN_BACKEND):
|
||||||
|
instance = _BACKEND_REGISTRY[backend]()
|
||||||
|
elif isinstance(backend, type) and issubclass(backend, AttentionBackend):
|
||||||
|
instance = backend()
|
||||||
|
elif isinstance(backend, AttentionBackend):
|
||||||
|
instance = backend
|
||||||
|
else:
|
||||||
|
raise TypeError(
|
||||||
|
f"expected ATTN_BACKEND, AttentionBackend type, or instance, "
|
||||||
|
f"got {type(backend).__name__}"
|
||||||
|
)
|
||||||
|
token = _current_backend.set(instance)
|
||||||
|
try:
|
||||||
|
yield instance
|
||||||
|
finally:
|
||||||
|
_current_backend.reset(token)
|
||||||
|
|
||||||
|
|
||||||
|
def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
|
||||||
|
"""Expand KV heads to match Q heads for GQA."""
|
||||||
|
bs, slen, n_heads, head_dim = x.shape
|
||||||
|
if n_rep == 1:
|
||||||
|
return x
|
||||||
|
return (
|
||||||
|
x[:, :, :, None, :]
|
||||||
|
.expand(bs, slen, n_heads, n_rep, head_dim)
|
||||||
|
.reshape(bs, slen, n_heads * n_rep, head_dim)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def attention(
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional[KVCache] = None,
|
||||||
|
layer_id: int = 0,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Functional attention entry point — mirrors ``F.scaled_dot_product_attention``.
|
||||||
|
|
||||||
|
Delegates to the active backend (set via ``with attn_backend(...)``).
|
||||||
|
Handles KV cache I/O, GQA head expansion, and causal masking so the
|
||||||
|
caller only needs to provide projected q/k/v.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
q: [batch, q_len, n_heads, head_dim] (blhd)
|
||||||
|
k: [batch, q_len, n_kv_heads, head_dim] (blhd)
|
||||||
|
v: [batch, q_len, n_kv_heads, head_dim] (blhd)
|
||||||
|
kv_cache: cache dataclass, or None for training (no cache).
|
||||||
|
layer_id: transformer layer index for buffer access.
|
||||||
|
attn_mask: pre-built attention mask (SDPA-compatible).
|
||||||
|
is_causal: whether to apply causal masking.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[batch, q_len, n_heads * head_dim]
|
||||||
|
"""
|
||||||
|
backend = get_backend()
|
||||||
|
return backend.forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||||
|
|
||||||
|
|
||||||
|
class AttentionBackend(ABC):
|
||||||
|
"""Abstract base for attention computation strategies.
|
||||||
|
|
||||||
|
Subclasses implement ``fwd_decode`` (q_len == 1, with cache) and
|
||||||
|
``fwd_prefill`` (q_len > 1, with or without cache). The public
|
||||||
|
``forward`` method dispatches based on q_len.
|
||||||
|
|
||||||
|
Three equivalent ways to activate a backend::
|
||||||
|
|
||||||
|
with attn_backend(ATTN_BACKEND.TORCH_NATIVE): # enum
|
||||||
|
...
|
||||||
|
with attn_backend(TorchNativeBackend): # class
|
||||||
|
...
|
||||||
|
with TorchNativeBackend(): # instance
|
||||||
|
...
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __enter__(self) -> "AttentionBackend":
|
||||||
|
self._token = _current_backend.set(self)
|
||||||
|
return self
|
||||||
|
|
||||||
|
def __exit__(self, *exc) -> None:
|
||||||
|
_current_backend.reset(self._token)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional[KVCache],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Dispatch to decode or extend based on q_len.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
q: [batch, q_len, n_heads, head_dim]
|
||||||
|
k: [batch, q_len, n_kv_heads, head_dim]
|
||||||
|
v: [batch, q_len, n_kv_heads, head_dim]
|
||||||
|
kv_cache: cache dataclass, or None for training (no cache).
|
||||||
|
layer_id: transformer layer index for buffer access.
|
||||||
|
attn_mask: pre-built attention mask compatible with SDPA.
|
||||||
|
is_causal: whether to apply causal masking.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[batch, q_len, n_heads * head_dim]
|
||||||
|
"""
|
||||||
|
if kv_cache is not None and q.size(1) == 1:
|
||||||
|
return self.fwd_decode(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||||
|
return self.fwd_prefill(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def fwd_decode(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional[KVCache],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Single-token decode with KV cache."""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def fwd_prefill(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional[KVCache],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Multi-token prefill or training forward."""
|
||||||
|
|
||||||
|
|
||||||
|
class TorchNativeBackend(AttentionBackend):
|
||||||
|
"""Reference backend using torch SDPA with indirect KV cache indexing.
|
||||||
|
|
||||||
|
Writes new K/V into the cache buffers, gathers the full sequence K/V
|
||||||
|
via ``req_to_token`` indirect indexing, then calls
|
||||||
|
``F.scaled_dot_product_attention``.
|
||||||
|
|
||||||
|
For training (``kv_cache is None``), skips cache I/O entirely and
|
||||||
|
runs SDPA directly on the projected q/k/v.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def fwd_decode(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional[KVCache],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||||
|
|
||||||
|
def fwd_prefill(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional[KVCache],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||||
|
|
||||||
|
def _forward(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional[KVCache],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
if kv_cache is not None:
|
||||||
|
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
||||||
|
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
||||||
|
|
||||||
|
max_len = kv_cache.max_len
|
||||||
|
if kv_cache.page_table is not None:
|
||||||
|
indices = kv_cache.page_table
|
||||||
|
else:
|
||||||
|
indices = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len]
|
||||||
|
if kv_cache.decode_mask is not None:
|
||||||
|
pos_mask = kv_cache.decode_mask
|
||||||
|
else:
|
||||||
|
pos_mask = (
|
||||||
|
torch.arange(max_len, device=q.device)[None, :]
|
||||||
|
< kv_cache.seq_lens[:, None]
|
||||||
|
)
|
||||||
|
indices = torch.where(pos_mask, indices, torch.zeros_like(indices))
|
||||||
|
k = kv_cache.k_buffer[layer_id, indices]
|
||||||
|
v = kv_cache.v_buffer[layer_id, indices]
|
||||||
|
|
||||||
|
n_rep = q.size(2) // k.size(2)
|
||||||
|
if n_rep > 1:
|
||||||
|
k = repeat_kv(k, n_rep)
|
||||||
|
v = repeat_kv(v, n_rep)
|
||||||
|
|
||||||
|
q = q.permute(0, 2, 1, 3)
|
||||||
|
k = k.permute(0, 2, 1, 3)
|
||||||
|
v = v.permute(0, 2, 1, 3)
|
||||||
|
|
||||||
|
out = F.scaled_dot_product_attention(q, k, v, attn_mask, is_causal=is_causal)
|
||||||
|
out = out.permute(0, 2, 1, 3).contiguous().flatten(2)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
_default_backend = TorchNativeBackend()
|
||||||
|
|
||||||
|
|
||||||
|
class CudaBackend(AttentionBackend):
|
||||||
|
"""CUDA kernel backend with direct KV cache access.
|
||||||
|
|
||||||
|
Decode path: writes K/V to cache, then calls ``attn_paged_decode``
|
||||||
|
with ``page_size=1`` (each token slot is a single-token "page").
|
||||||
|
The ``req_to_token`` table serves directly as the page table.
|
||||||
|
|
||||||
|
Prefill path: writes K/V to cache, gathers full-sequence K/V via
|
||||||
|
indirect indexing (same as TorchNativeBackend), then calls
|
||||||
|
``attn_prefill``.
|
||||||
|
|
||||||
|
Training path (``kv_cache is None``): calls ``attn_prefill`` directly
|
||||||
|
on the projected q/k/v.
|
||||||
|
|
||||||
|
Falls back to ``TorchNativeBackend`` for any path where the
|
||||||
|
corresponding CUDA kernel is not available.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self._fallback = TorchNativeBackend()
|
||||||
|
|
||||||
|
def fwd_decode(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional[KVCache],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
if kv_cache is None or not is_available("attn_paged_decode"):
|
||||||
|
return self._fallback.fwd_decode(
|
||||||
|
q, k, v, kv_cache, layer_id, attn_mask, is_causal
|
||||||
|
)
|
||||||
|
|
||||||
|
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
||||||
|
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
||||||
|
|
||||||
|
max_len = kv_cache.max_len
|
||||||
|
|
||||||
|
if kv_cache.page_table is not None:
|
||||||
|
page_table = kv_cache.page_table
|
||||||
|
else:
|
||||||
|
page_table = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len]
|
||||||
|
|
||||||
|
k_cache = kv_cache.k_buffer[layer_id].unsqueeze(1)
|
||||||
|
v_cache = kv_cache.v_buffer[layer_id].unsqueeze(1)
|
||||||
|
|
||||||
|
if q.size(0) == 1:
|
||||||
|
mask = None
|
||||||
|
elif kv_cache.decode_mask is not None:
|
||||||
|
mask = kv_cache.decode_mask
|
||||||
|
else:
|
||||||
|
mask = (
|
||||||
|
torch.arange(max_len, device=q.device)[None, :]
|
||||||
|
< kv_cache.seq_lens[:, None]
|
||||||
|
)
|
||||||
|
|
||||||
|
out = attn_paged_decode(
|
||||||
|
q,
|
||||||
|
page_table,
|
||||||
|
k_cache,
|
||||||
|
v_cache,
|
||||||
|
page_size=1,
|
||||||
|
kv_len=max_len,
|
||||||
|
mask=mask,
|
||||||
|
is_causal=is_causal,
|
||||||
|
)
|
||||||
|
|
||||||
|
out = out.flatten(2)
|
||||||
|
return out
|
||||||
|
|
||||||
|
def fwd_prefill(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional[KVCache],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
if kv_cache is None:
|
||||||
|
if is_available("attn_prefill"):
|
||||||
|
out = attn_prefill(q, k, v, mask=attn_mask, is_causal=is_causal)
|
||||||
|
return out.flatten(2)
|
||||||
|
return self._fallback.fwd_prefill(
|
||||||
|
q, k, v, kv_cache, layer_id, attn_mask, is_causal
|
||||||
|
)
|
||||||
|
|
||||||
|
if not is_available("attn_prefill"):
|
||||||
|
return self._fallback.fwd_prefill(
|
||||||
|
q, k, v, kv_cache, layer_id, attn_mask, is_causal
|
||||||
|
)
|
||||||
|
|
||||||
|
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
||||||
|
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
||||||
|
|
||||||
|
max_len = kv_cache.max_len
|
||||||
|
indices = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len]
|
||||||
|
pos_mask = (
|
||||||
|
torch.arange(max_len, device=q.device)[None, :] < kv_cache.seq_lens[:, None]
|
||||||
|
)
|
||||||
|
indices = torch.where(pos_mask, indices, torch.zeros_like(indices))
|
||||||
|
k_full = kv_cache.k_buffer[layer_id, indices]
|
||||||
|
v_full = kv_cache.v_buffer[layer_id, indices]
|
||||||
|
|
||||||
|
out = attn_prefill(q, k_full, v_full, mask=attn_mask, is_causal=is_causal)
|
||||||
|
return out.flatten(2)
|
||||||
|
|
||||||
|
|
||||||
|
_BACKEND_REGISTRY: dict[ATTN_BACKEND, type[AttentionBackend]] = {
|
||||||
|
ATTN_BACKEND.TORCH_NATIVE: TorchNativeBackend,
|
||||||
|
ATTN_BACKEND.CUDA: CudaBackend,
|
||||||
|
}
|
||||||
@@ -0,0 +1,117 @@
|
|||||||
|
"""Attention kernel wrapper functions — one entry point per compiled kernel.
|
||||||
|
|
||||||
|
Each wrapper calls its CUDA kernel directly. If the kernel is not
|
||||||
|
available, raises ``RuntimeError``. Fallback to torch SDPA is the
|
||||||
|
responsibility of the attention backend, not this module.
|
||||||
|
|
||||||
|
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
|
||||||
|
(blhd). Scale is always ``1/sqrt(head_dim)``.
|
||||||
|
|
||||||
|
Interface (all functions):
|
||||||
|
is_causal: True = causal mask; False = non-causal
|
||||||
|
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True=keep)
|
||||||
|
"""
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.extension.loader import _available, _modules
|
||||||
|
|
||||||
|
|
||||||
|
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,
|
||||||
|
v: torch.Tensor,
|
||||||
|
mask: torch.Tensor | None = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""GQA decode attention (q_len == 1).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
q: [batch, 1, n_heads, head_dim] (blhd, bf16)
|
||||||
|
k: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
||||||
|
v: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
||||||
|
mask: 2D [batch, kv_len] or 3D [batch, 1, kv_len] (bool, True=keep)
|
||||||
|
is_causal: apply causal mask
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[batch, 1, n_heads, head_dim] (blhd, bf16)
|
||||||
|
"""
|
||||||
|
_check_available("attn_decode")
|
||||||
|
causal_offset = (k.size(1) - 1) if is_causal else -1
|
||||||
|
return _modules["attn_decode"].attn_decode(
|
||||||
|
q, k, v, mask=mask, causal_offset=causal_offset, layout=1
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def attn_prefill(
|
||||||
|
q: torch.Tensor,
|
||||||
|
k: torch.Tensor,
|
||||||
|
v: torch.Tensor,
|
||||||
|
mask: torch.Tensor | None = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""GQA prefill attention (q_len > 1).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
q: [batch, q_len, n_heads, head_dim] (blhd, bf16)
|
||||||
|
k: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
||||||
|
v: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
||||||
|
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True=keep)
|
||||||
|
is_causal: apply causal mask
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[batch, q_len, n_heads, head_dim] (blhd, bf16)
|
||||||
|
"""
|
||||||
|
_check_available("attn_prefill")
|
||||||
|
causal_offset = (k.size(1) - q.size(1)) if is_causal else -1
|
||||||
|
return _modules["attn_prefill"].attn_prefill(
|
||||||
|
q, k, v, mask=mask, causal_offset=causal_offset, layout=1
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def attn_paged_decode(
|
||||||
|
q: torch.Tensor,
|
||||||
|
page_table: torch.Tensor,
|
||||||
|
k_cache: torch.Tensor,
|
||||||
|
v_cache: torch.Tensor,
|
||||||
|
page_size: int,
|
||||||
|
kv_len: int,
|
||||||
|
mask: torch.Tensor | None = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Paged GQA decode attention (q_len == 1, direct page-table access).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
q: [batch, 1, n_heads, head_dim] (blhd, bf16)
|
||||||
|
page_table: [batch, max_pages] (int64)
|
||||||
|
k_cache: [n_pages, page_size, n_kv_heads, head_dim] (bf16)
|
||||||
|
v_cache: same as k_cache
|
||||||
|
page_size: tokens per page
|
||||||
|
kv_len: actual sequence length per request
|
||||||
|
mask: 2D [batch, kv_len] or 3D [batch, 1, kv_len] (bool, True=keep)
|
||||||
|
is_causal: apply causal mask
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[batch, 1, n_heads, head_dim] (blhd, bf16)
|
||||||
|
"""
|
||||||
|
_check_available("attn_paged_decode")
|
||||||
|
causal_offset = (kv_len - 1) if is_causal else -1
|
||||||
|
return _modules["attn_paged_decode"].attn_paged_decode(
|
||||||
|
q,
|
||||||
|
page_table,
|
||||||
|
k_cache,
|
||||||
|
v_cache,
|
||||||
|
page_size,
|
||||||
|
kv_len,
|
||||||
|
mask=mask,
|
||||||
|
causal_offset=causal_offset,
|
||||||
|
layout=1,
|
||||||
|
)
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
"""Compiled CUDA kernel modules (``*.so``) live here, kept separate from Python source."""
|
||||||
@@ -11,14 +11,14 @@ import logging
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
KERNEL_NAMES = ["attn_decode", "attn_prefill", "attn_paged_decode"]
|
KERNEL_NAMES = ["attn_decode", "attn_prefill", "attn_paged_decode", "rotary_emb"]
|
||||||
|
|
||||||
_available: dict[str, bool] = {}
|
_available: dict[str, bool] = {}
|
||||||
_modules: dict[str, object] = {}
|
_modules: dict[str, object] = {}
|
||||||
|
|
||||||
for _name in KERNEL_NAMES:
|
for _name in KERNEL_NAMES:
|
||||||
try:
|
try:
|
||||||
_mod = importlib.import_module(f".{_name}", package=__package__)
|
_mod = importlib.import_module(f".lib.{_name}", package=__package__)
|
||||||
_available[_name] = True
|
_available[_name] = True
|
||||||
_modules[_name] = _mod
|
_modules[_name] = _mod
|
||||||
except ImportError:
|
except ImportError:
|
||||||
|
|||||||
@@ -1,298 +0,0 @@
|
|||||||
"""GQA attention wrapper functions — one entry point per compiled kernel.
|
|
||||||
|
|
||||||
Each wrapper dispatches to its CUDA kernel (loaded in ``loader.py``) when
|
|
||||||
available, otherwise falls back to ``torch`` SDPA.
|
|
||||||
|
|
||||||
Interface (all functions):
|
|
||||||
causal_offset: -1 = non-causal; >=0 = absolute position of first Q token
|
|
||||||
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool)
|
|
||||||
scale: 0.0 = auto (1/sqrt(head_dim)); >0 = explicit
|
|
||||||
layout: "bhld" (default) or "blhd"
|
|
||||||
|
|
||||||
Add new kernel wrappers here; split into per-variant files only if this file
|
|
||||||
grows large.
|
|
||||||
"""
|
|
||||||
|
|
||||||
import math
|
|
||||||
|
|
||||||
import torch
|
|
||||||
import torch.nn.functional as F
|
|
||||||
|
|
||||||
from astrai.extension.loader import _available, _modules
|
|
||||||
|
|
||||||
_LAYOUT_CODES: dict[str, int] = {"bhld": 0, "blhd": 1}
|
|
||||||
|
|
||||||
|
|
||||||
def _parse_layout(layout: str | int) -> int:
|
|
||||||
if isinstance(layout, int):
|
|
||||||
return layout
|
|
||||||
code = _LAYOUT_CODES.get(layout.lower())
|
|
||||||
if code is None:
|
|
||||||
raise ValueError(
|
|
||||||
f"unknown layout '{layout}', expected one of {list(_LAYOUT_CODES)}"
|
|
||||||
)
|
|
||||||
return code
|
|
||||||
|
|
||||||
|
|
||||||
def _to_bhld(t: torch.Tensor, layout: int) -> torch.Tensor:
|
|
||||||
"""Normalize to b h l d view. Zero-copy transpose if layout==1 (b l h d)."""
|
|
||||||
if layout == 1:
|
|
||||||
return t.transpose(1, 2)
|
|
||||||
return t
|
|
||||||
|
|
||||||
|
|
||||||
def _expand_kv_heads(
|
|
||||||
k: torch.Tensor, v: torch.Tensor, q_head: int
|
|
||||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
|
||||||
"""Expand K/V heads to match Q heads for GQA fallback."""
|
|
||||||
kv_head = k.size(1)
|
|
||||||
if kv_head == q_head:
|
|
||||||
return k, v
|
|
||||||
group = q_head // kv_head
|
|
||||||
k = k.repeat_interleave(group, dim=1)
|
|
||||||
v = v.repeat_interleave(group, dim=1)
|
|
||||||
return k, v
|
|
||||||
|
|
||||||
|
|
||||||
def _build_attn_mask(
|
|
||||||
q: torch.Tensor,
|
|
||||||
k: torch.Tensor,
|
|
||||||
mask: torch.Tensor | None,
|
|
||||||
causal_offset: int,
|
|
||||||
scale: float,
|
|
||||||
) -> tuple[torch.Tensor | None, float]:
|
|
||||||
"""Build SDPA-compatible attn_mask + resolved scale.
|
|
||||||
|
|
||||||
q and k must already be in b h l d layout.
|
|
||||||
Causal and mask can coexist: causal sets -inf above the diagonal, mask
|
|
||||||
sets -inf for padded positions. Both are OR'd into a single bool mask.
|
|
||||||
"""
|
|
||||||
q_len = q.size(2)
|
|
||||||
kv_len = k.size(2)
|
|
||||||
head_dim = q.size(3)
|
|
||||||
resolved_scale = scale if scale and scale > 0 else 1.0 / math.sqrt(head_dim)
|
|
||||||
|
|
||||||
attn_mask = None
|
|
||||||
|
|
||||||
if mask is not None:
|
|
||||||
if mask.dim() == 2:
|
|
||||||
# [batch, kv_len] → [batch, 1, 1, kv_len]
|
|
||||||
attn_mask = mask[:, None, None, :]
|
|
||||||
elif mask.dim() == 3:
|
|
||||||
# [batch, q_len, kv_len] → [batch, 1, q_len, kv_len]
|
|
||||||
attn_mask = mask[:, None, :, :]
|
|
||||||
else:
|
|
||||||
raise ValueError(f"mask must be 2D or 3D, got {mask.dim()}D")
|
|
||||||
|
|
||||||
if causal_offset >= 0:
|
|
||||||
batch = q.size(0)
|
|
||||||
# q row i attends to kv cols 0..(causal_offset + i)
|
|
||||||
q_idx = torch.arange(q_len, device=q.device).unsqueeze(1) # [q_len, 1]
|
|
||||||
kv_idx = torch.arange(kv_len, device=q.device).unsqueeze(0) # [1, kv_len]
|
|
||||||
causal_bool = kv_idx > (causal_offset + q_idx) # True = masked out
|
|
||||||
causal_mask = causal_bool.unsqueeze(0).expand(
|
|
||||||
batch, -1, -1
|
|
||||||
) # [batch, q_len, kv_len]
|
|
||||||
causal_mask = causal_mask[:, None, :, :] # [batch, 1, q_len, kv_len]
|
|
||||||
|
|
||||||
if attn_mask is not None:
|
|
||||||
attn_mask = attn_mask | causal_mask
|
|
||||||
else:
|
|
||||||
attn_mask = causal_mask
|
|
||||||
|
|
||||||
return attn_mask, resolved_scale
|
|
||||||
|
|
||||||
|
|
||||||
def _torch_fallback(
|
|
||||||
q: torch.Tensor,
|
|
||||||
k: torch.Tensor,
|
|
||||||
v: torch.Tensor,
|
|
||||||
mask: torch.Tensor | None,
|
|
||||||
causal_offset: int,
|
|
||||||
scale: float,
|
|
||||||
q_layout: int,
|
|
||||||
kv_layout: int | None = None,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
"""Reference attention via ``scaled_dot_product_attention``.
|
|
||||||
|
|
||||||
q_layout / kv_layout: 0 = b h l d, 1 = b l h d.
|
|
||||||
If kv_layout is None, uses q_layout (Q and K/V share the same layout).
|
|
||||||
"""
|
|
||||||
if kv_layout is None:
|
|
||||||
kv_layout = q_layout
|
|
||||||
q = _to_bhld(q, q_layout)
|
|
||||||
k = _to_bhld(k, kv_layout)
|
|
||||||
v = _to_bhld(v, kv_layout)
|
|
||||||
k, v = _expand_kv_heads(k, v, q.size(1))
|
|
||||||
attn_mask, resolved_scale = _build_attn_mask(q, k, mask, causal_offset, scale)
|
|
||||||
out = F.scaled_dot_product_attention(
|
|
||||||
q, k, v, attn_mask=attn_mask, is_causal=False, scale=resolved_scale
|
|
||||||
)
|
|
||||||
# Restore Q's original layout
|
|
||||||
if q_layout == 1:
|
|
||||||
out = out.transpose(1, 2)
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
def _gather_kv_from_pages(
|
|
||||||
page_table: torch.Tensor,
|
|
||||||
k_cache: torch.Tensor,
|
|
||||||
v_cache: torch.Tensor,
|
|
||||||
page_size: int,
|
|
||||||
kv_len: int,
|
|
||||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
|
||||||
"""Gather contiguous K/V from paged cache for torch SDPA fallback.
|
|
||||||
|
|
||||||
Shapes:
|
|
||||||
page_table : [batch, max_pages] (int64)
|
|
||||||
k_cache : [n_pages, page_size, n_kv_heads, head_dim]
|
|
||||||
v_cache : same as k_cache
|
|
||||||
Returns:
|
|
||||||
k, v : [batch, kv_len, n_kv_heads, head_dim] (b l h d)
|
|
||||||
"""
|
|
||||||
batch, max_pages = page_table.shape
|
|
||||||
_, ps, n_kv_heads, head_dim = k_cache.shape
|
|
||||||
if ps != page_size:
|
|
||||||
raise ValueError(f"k_cache page_size mismatch: {ps} vs {page_size}")
|
|
||||||
|
|
||||||
# Vectorized gather: build physical page + offset indices, then advanced-index
|
|
||||||
positions = torch.arange(kv_len, device=page_table.device)
|
|
||||||
logical_pages = positions // page_size # [kv_len]
|
|
||||||
page_offsets = positions % page_size # [kv_len]
|
|
||||||
|
|
||||||
phys_pages = page_table[:, logical_pages] # [batch, kv_len]
|
|
||||||
# k_cache[phys_pages, page_offsets] → [batch, kv_len, n_kv_heads, head_dim] (b l h d)
|
|
||||||
k = k_cache[phys_pages, page_offsets]
|
|
||||||
v = v_cache[phys_pages, page_offsets]
|
|
||||||
return k, v
|
|
||||||
|
|
||||||
|
|
||||||
def attn_decode(
|
|
||||||
q: torch.Tensor,
|
|
||||||
k: torch.Tensor,
|
|
||||||
v: torch.Tensor,
|
|
||||||
mask: torch.Tensor | None = None,
|
|
||||||
causal_offset: int = -1,
|
|
||||||
scale: float = 0.0,
|
|
||||||
layout: str = "bhld",
|
|
||||||
) -> torch.Tensor:
|
|
||||||
li = _parse_layout(layout)
|
|
||||||
if _available["attn_decode"]:
|
|
||||||
return _modules["attn_decode"].attn_decode(
|
|
||||||
q,
|
|
||||||
k,
|
|
||||||
v,
|
|
||||||
mask=mask,
|
|
||||||
causal_offset=causal_offset,
|
|
||||||
scale=scale,
|
|
||||||
layout=li,
|
|
||||||
)
|
|
||||||
return _torch_fallback(q, k, v, mask, causal_offset, scale, q_layout=li)
|
|
||||||
|
|
||||||
|
|
||||||
def attn_prefill(
|
|
||||||
q: torch.Tensor,
|
|
||||||
k: torch.Tensor,
|
|
||||||
v: torch.Tensor,
|
|
||||||
mask: torch.Tensor | None = None,
|
|
||||||
causal_offset: int = -1,
|
|
||||||
scale: float = 0.0,
|
|
||||||
layout: str = "bhld",
|
|
||||||
) -> torch.Tensor:
|
|
||||||
li = _parse_layout(layout)
|
|
||||||
if _available["attn_prefill"]:
|
|
||||||
return _modules["attn_prefill"].attn_prefill(
|
|
||||||
q,
|
|
||||||
k,
|
|
||||||
v,
|
|
||||||
mask=mask,
|
|
||||||
causal_offset=causal_offset,
|
|
||||||
scale=scale,
|
|
||||||
layout=li,
|
|
||||||
)
|
|
||||||
return _torch_fallback(q, k, v, mask, causal_offset, scale, q_layout=li)
|
|
||||||
|
|
||||||
|
|
||||||
def attn_paged_decode(
|
|
||||||
q: torch.Tensor,
|
|
||||||
page_table: torch.Tensor,
|
|
||||||
k_cache: torch.Tensor,
|
|
||||||
v_cache: torch.Tensor,
|
|
||||||
page_size: int,
|
|
||||||
kv_len: int,
|
|
||||||
mask: torch.Tensor | None = None,
|
|
||||||
causal_offset: int = -1,
|
|
||||||
scale: float = 0.0,
|
|
||||||
layout: str = "bhld",
|
|
||||||
) -> torch.Tensor:
|
|
||||||
li = _parse_layout(layout)
|
|
||||||
if _available["attn_paged_decode"]:
|
|
||||||
return _modules["attn_paged_decode"].attn_paged_decode(
|
|
||||||
q,
|
|
||||||
page_table,
|
|
||||||
k_cache,
|
|
||||||
v_cache,
|
|
||||||
page_size,
|
|
||||||
kv_len,
|
|
||||||
mask=mask,
|
|
||||||
causal_offset=causal_offset,
|
|
||||||
scale=scale,
|
|
||||||
layout=li,
|
|
||||||
)
|
|
||||||
# Gathered K/V are always b l h d
|
|
||||||
k, v = _gather_kv_from_pages(page_table, k_cache, v_cache, page_size, kv_len)
|
|
||||||
return _torch_fallback(
|
|
||||||
q, k, v, mask, causal_offset, scale, q_layout=li, kv_layout=1
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def attention(
|
|
||||||
q: torch.Tensor,
|
|
||||||
k: torch.Tensor,
|
|
||||||
v: torch.Tensor,
|
|
||||||
mask: torch.Tensor | None = None,
|
|
||||||
causal_offset: int = -1,
|
|
||||||
scale: float = 0.0,
|
|
||||||
layout: str = "bhld",
|
|
||||||
) -> torch.Tensor:
|
|
||||||
"""Dispatch to decode or prefill attention based on the query length.
|
|
||||||
|
|
||||||
A query length of one is the decode case; longer queries use prefill.
|
|
||||||
The paged-cache decode path cannot be selected here because its page-table
|
|
||||||
arguments are not part of this interface.
|
|
||||||
"""
|
|
||||||
li = _parse_layout(layout)
|
|
||||||
|
|
||||||
if q.ndim not in (2, 3, 4) or k.ndim != q.ndim or v.ndim != q.ndim:
|
|
||||||
raise ValueError(
|
|
||||||
"q, k, and v must all have the same rank in {2, 3, 4}, "
|
|
||||||
f"got {q.ndim}D, {k.ndim}D, {v.ndim}D"
|
|
||||||
)
|
|
||||||
if k.shape != v.shape:
|
|
||||||
raise ValueError(
|
|
||||||
f"k and v must have the same shape, got {k.shape} and {v.shape}"
|
|
||||||
)
|
|
||||||
|
|
||||||
original_ndim = q.ndim
|
|
||||||
if original_ndim == 2:
|
|
||||||
# [L, D] -> [1, 1, L, D] or [1, L, 1, D]
|
|
||||||
q = q.unsqueeze(0).unsqueeze(1 if li == 0 else 2)
|
|
||||||
k = k.unsqueeze(0).unsqueeze(1 if li == 0 else 2)
|
|
||||||
v = v.unsqueeze(0).unsqueeze(1 if li == 0 else 2)
|
|
||||||
elif original_ndim == 3:
|
|
||||||
# [B, L, D] -> single-head 4D input.
|
|
||||||
q = q.unsqueeze(1 if li == 0 else 2)
|
|
||||||
k = k.unsqueeze(1 if li == 0 else 2)
|
|
||||||
v = v.unsqueeze(1 if li == 0 else 2)
|
|
||||||
|
|
||||||
q_len = q.size(2 if li == 0 else 1)
|
|
||||||
if q_len == 1:
|
|
||||||
out = attn_decode(q, k, v, mask, causal_offset, scale, layout)
|
|
||||||
else:
|
|
||||||
out = attn_prefill(q, k, v, mask, causal_offset, scale, layout)
|
|
||||||
|
|
||||||
if original_ndim == 2:
|
|
||||||
return out.squeeze(0).squeeze(0 if li == 0 else 1)
|
|
||||||
if original_ndim == 3:
|
|
||||||
return out.squeeze(1 if li == 0 else 2)
|
|
||||||
return out
|
|
||||||
@@ -0,0 +1,54 @@
|
|||||||
|
"""Rotary embedding with auto-dispatch to CUDA kernel.
|
||||||
|
|
||||||
|
Single entry point ``apply_rotary_emb(x, freqs_cis)`` — uses the fused
|
||||||
|
CUDA kernel when available, falls back to torch complex multiply otherwise.
|
||||||
|
|
||||||
|
Layout: x is [batch, seq_len, n_heads, head_dim] (bf16).
|
||||||
|
freqs_cis is [batch, seq_len, dim/2, 2] (f32) — [cos, sin] pairs.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.extension.loader import is_available
|
||||||
|
|
||||||
|
_cache = {"available": None}
|
||||||
|
|
||||||
|
|
||||||
|
def _cuda_available() -> bool:
|
||||||
|
if _cache["available"] is None:
|
||||||
|
_cache["available"] = is_available("rotary_emb")
|
||||||
|
return _cache["available"]
|
||||||
|
|
||||||
|
|
||||||
|
def _torch_apply(x: Tensor, freqs_cis: Tensor) -> Tensor:
|
||||||
|
cos, sin = freqs_cis[..., 0], freqs_cis[..., 1]
|
||||||
|
dtype = x.dtype
|
||||||
|
x_ = x.float().reshape(*x.shape[:-1], -1, 2)
|
||||||
|
x_complex = torch.view_as_complex(x_)
|
||||||
|
freqs_cis_complex = torch.complex(cos, sin).unsqueeze(2)
|
||||||
|
x_rotated = x_complex * freqs_cis_complex
|
||||||
|
x_out = torch.view_as_real(x_rotated).flatten(-2)
|
||||||
|
return x_out.to(dtype)
|
||||||
|
|
||||||
|
|
||||||
|
def apply_rotary_emb(x: Tensor, freqs_cis: Tensor) -> Tensor:
|
||||||
|
"""Apply rotary embedding to x.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
x: [batch, seq_len, n_heads, head_dim] (bf16)
|
||||||
|
freqs_cis: [batch, seq_len, dim/2, 2] (f32) — [cos, sin] pairs
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[batch, seq_len, n_heads, head_dim] (bf16)
|
||||||
|
"""
|
||||||
|
if (
|
||||||
|
_cuda_available()
|
||||||
|
and not torch.is_grad_enabled()
|
||||||
|
and x.is_cuda
|
||||||
|
and x.dtype == torch.bfloat16
|
||||||
|
):
|
||||||
|
from astrai.extension.rotary_ops import rotary_emb as _cuda_rotary
|
||||||
|
|
||||||
|
return _cuda_rotary(x, freqs_cis)
|
||||||
|
return _torch_apply(x, freqs_cis)
|
||||||
@@ -0,0 +1,39 @@
|
|||||||
|
"""Rotary embedding CUDA kernel wrapper.
|
||||||
|
|
||||||
|
Calls the compiled CUDA kernel directly. If the kernel is not available,
|
||||||
|
raises ``RuntimeError``. Fallback to torch complex multiply is the
|
||||||
|
responsibility of ``astrai.extension.rotary_backend.apply_rotary_emb``.
|
||||||
|
|
||||||
|
Layout: x is [batch, seq_len, n_heads, head_dim] (bf16, contiguous).
|
||||||
|
freqs_cis is [batch, seq_len, head_dim/2, 2] (f32, contiguous) — [cos, sin] pairs.
|
||||||
|
"""
|
||||||
|
|
||||||
|
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."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def rotary_emb(x: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""Fused rotary embedding kernel.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
x: [batch, seq_len, n_heads, head_dim] (bf16, contiguous)
|
||||||
|
freqs_cis: [batch, seq_len, head_dim/2, 2] (f32, contiguous) — [cos, sin] pairs
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[batch, seq_len, n_heads, head_dim] (bf16)
|
||||||
|
"""
|
||||||
|
_check_available()
|
||||||
|
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)
|
||||||
+41
-35
@@ -13,41 +13,63 @@ from typing import (
|
|||||||
Type,
|
Type,
|
||||||
TypeVar,
|
TypeVar,
|
||||||
Union,
|
Union,
|
||||||
|
get_args,
|
||||||
|
get_origin,
|
||||||
)
|
)
|
||||||
from typing import get_args as _get_args
|
|
||||||
from typing import get_origin as _get_origin
|
|
||||||
|
|
||||||
T = TypeVar("T")
|
T = TypeVar("T")
|
||||||
|
|
||||||
|
|
||||||
def _resolve_type(
|
def _resolve_base_type(
|
||||||
arg: Union[Type, str, ForwardRef], factory_cls: type
|
arg: Union[Type, str, ForwardRef], factory_cls: type
|
||||||
) -> Optional[Type]:
|
) -> Optional[Type]:
|
||||||
"""Resolve a generic type-arg (str forward-ref, ForwardRef, or class)."""
|
"""Resolve the generic type-arg T to a concrete class.
|
||||||
if not isinstance(arg, (str, ForwardRef)):
|
|
||||||
|
- Concrete class (``BaseFactory[MyBase]``): returned directly.
|
||||||
|
- Forward reference (``BaseFactory["MyBase"]``): ``Base["X"]``
|
||||||
|
produces a ``ForwardRef("X")`` at class-creation time. We
|
||||||
|
extract the name and evaluate it in the factory module's
|
||||||
|
global namespace — the same mechanism ``typing.get_type_hints``
|
||||||
|
uses internally.
|
||||||
|
"""
|
||||||
|
if isinstance(arg, type):
|
||||||
return arg
|
return arg
|
||||||
|
|
||||||
name = arg if isinstance(arg, str) else arg.__forward_arg__
|
if isinstance(arg, str):
|
||||||
if name == factory_cls.__name__:
|
name = arg
|
||||||
return factory_cls
|
elif isinstance(arg, ForwardRef):
|
||||||
|
name = arg.__forward_arg__
|
||||||
|
else:
|
||||||
|
return None
|
||||||
|
|
||||||
mod = sys.modules.get(factory_cls.__module__)
|
mod = sys.modules.get(factory_cls.__module__)
|
||||||
if mod is None:
|
if mod is None:
|
||||||
return None
|
return None
|
||||||
ns = vars(mod)
|
try:
|
||||||
|
return eval(name, vars(mod)) # noqa: S307
|
||||||
|
except NameError:
|
||||||
|
return None
|
||||||
|
|
||||||
if isinstance(arg, ForwardRef):
|
|
||||||
return arg._evaluate(ns, None, recursive_guard=frozenset())
|
|
||||||
|
|
||||||
return ns.get(name)
|
def _validate_component(component_cls: Type, base: Optional[Type]) -> None:
|
||||||
|
"""Validate that *component_cls* inherits from *base*.
|
||||||
|
|
||||||
|
No-op when *base* is ``None`` (e.g. forward-ref resolution failed).
|
||||||
|
"""
|
||||||
|
if base is not None and not issubclass(component_cls, base):
|
||||||
|
raise TypeError(f"{component_cls.__name__} must inherit from {base.__name__}")
|
||||||
|
|
||||||
|
|
||||||
class BaseFactory(ABC, Generic[T]):
|
class BaseFactory(ABC, Generic[T]):
|
||||||
"""Generic factory with decorator-based component registration.
|
"""Generic factory with decorator-based registration.
|
||||||
|
|
||||||
|
Create a factory by subclassing with the desired base type::
|
||||||
|
|
||||||
class MyFactory(BaseFactory[MyBase]):
|
class MyFactory(BaseFactory[MyBase]):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
Register components with the ``register`` decorator::
|
||||||
|
|
||||||
@MyFactory.register("custom")
|
@MyFactory.register("custom")
|
||||||
class CustomComponent(MyBase):
|
class CustomComponent(MyBase):
|
||||||
...
|
...
|
||||||
@@ -64,13 +86,10 @@ class BaseFactory(ABC, Generic[T]):
|
|||||||
def __init_subclass__(cls, **kwargs):
|
def __init_subclass__(cls, **kwargs):
|
||||||
super().__init_subclass__(**kwargs)
|
super().__init_subclass__(**kwargs)
|
||||||
for orig_base in getattr(cls, "__orig_bases__", ()):
|
for orig_base in getattr(cls, "__orig_bases__", ()):
|
||||||
if _get_origin(orig_base) is BaseFactory:
|
if get_origin(orig_base) is BaseFactory:
|
||||||
(arg,) = _get_args(orig_base)
|
(arg,) = get_args(orig_base)
|
||||||
cls._entries = {}
|
cls._entries = {}
|
||||||
try:
|
cls._component_base = _resolve_base_type(arg, cls)
|
||||||
cls._component_base = _resolve_type(arg, cls)
|
|
||||||
except Exception:
|
|
||||||
cls._component_base = None
|
|
||||||
return
|
return
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -82,7 +101,7 @@ class BaseFactory(ABC, Generic[T]):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
def decorator(component_cls: Type[T]) -> Type[T]:
|
def decorator(component_cls: Type[T]) -> Type[T]:
|
||||||
cls._validate_component(component_cls)
|
_validate_component(component_cls, cls._component_base)
|
||||||
if name in cls._entries:
|
if name in cls._entries:
|
||||||
raise ValueError(f"Component '{name}' is already registered")
|
raise ValueError(f"Component '{name}' is already registered")
|
||||||
cls._entries[name] = component_cls
|
cls._entries[name] = component_cls
|
||||||
@@ -95,12 +114,11 @@ class BaseFactory(ABC, Generic[T]):
|
|||||||
"""Create a component instance by name, filtering kwargs to match
|
"""Create a component instance by name, filtering kwargs to match
|
||||||
the component's ``__init__`` signature.
|
the component's ``__init__`` signature.
|
||||||
"""
|
"""
|
||||||
entry = cls._entries.get(name)
|
component_cls = cls._entries.get(name)
|
||||||
if entry is None:
|
if component_cls is None:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Unknown component: '{name}'. Supported types: {sorted(cls._entries)}"
|
f"Unknown component: '{name}'. Supported types: {sorted(cls._entries)}"
|
||||||
)
|
)
|
||||||
component_cls = entry
|
|
||||||
sig = inspect.signature(component_cls.__init__)
|
sig = inspect.signature(component_cls.__init__)
|
||||||
has_var_kwargs = any(
|
has_var_kwargs = any(
|
||||||
p.kind == inspect.Parameter.VAR_KEYWORD for p in sig.parameters.values()
|
p.kind == inspect.Parameter.VAR_KEYWORD for p in sig.parameters.values()
|
||||||
@@ -114,18 +132,6 @@ class BaseFactory(ABC, Generic[T]):
|
|||||||
kwargs = {k: v for k, v in kwargs.items() if k in valid}
|
kwargs = {k: v for k, v in kwargs.items() if k in valid}
|
||||||
return component_cls(*args, **kwargs)
|
return component_cls(*args, **kwargs)
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def _validate_component(cls, component_cls: Type[T]):
|
|
||||||
"""Validate the decorated class inherits from the factory's base type.
|
|
||||||
|
|
||||||
Override for custom validation beyond ``issubclass``.
|
|
||||||
"""
|
|
||||||
base = cls._component_base
|
|
||||||
if base is not None and not issubclass(component_cls, base):
|
|
||||||
raise TypeError(
|
|
||||||
f"{component_cls.__name__} must inherit from {base.__name__}"
|
|
||||||
)
|
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_component_class(cls, name: str) -> Type[T]:
|
def get_component_class(cls, name: str) -> Type[T]:
|
||||||
"""Get the registered component class without instantiating it."""
|
"""Get the registered component class without instantiating it."""
|
||||||
|
|||||||
@@ -30,21 +30,16 @@ from astrai.inference.api.openai import OpenAIResponseBuilder
|
|||||||
from astrai.inference.core import (
|
from astrai.inference.core import (
|
||||||
STOP,
|
STOP,
|
||||||
Allocator,
|
Allocator,
|
||||||
CacheView,
|
|
||||||
ContiguousCache,
|
|
||||||
ContiguousCacheView,
|
|
||||||
Executor,
|
Executor,
|
||||||
InferenceScheduler,
|
InferenceScheduler,
|
||||||
KVCache,
|
KVCache,
|
||||||
PageCache,
|
KVStorage,
|
||||||
PageCacheView,
|
|
||||||
PagePool,
|
PagePool,
|
||||||
PrefixCache,
|
PrefixCache,
|
||||||
Storage,
|
ReqToTokenPool,
|
||||||
Task,
|
Task,
|
||||||
TaskManager,
|
TaskManager,
|
||||||
TaskStatus,
|
TaskStatus,
|
||||||
TaskTable,
|
|
||||||
page_hash,
|
page_hash,
|
||||||
)
|
)
|
||||||
from astrai.inference.engine import GenerationRequest, InferenceEngine
|
from astrai.inference.engine import GenerationRequest, InferenceEngine
|
||||||
@@ -68,16 +63,11 @@ __all__ = [
|
|||||||
"TaskManager",
|
"TaskManager",
|
||||||
"TaskStatus",
|
"TaskStatus",
|
||||||
"Allocator",
|
"Allocator",
|
||||||
"CacheView",
|
|
||||||
"KVCache",
|
"KVCache",
|
||||||
"ContiguousCache",
|
"KVStorage",
|
||||||
"ContiguousCacheView",
|
|
||||||
"PageCache",
|
|
||||||
"PageCacheView",
|
|
||||||
"PagePool",
|
"PagePool",
|
||||||
"PrefixCache",
|
"PrefixCache",
|
||||||
"Storage",
|
"ReqToTokenPool",
|
||||||
"TaskTable",
|
|
||||||
"page_hash",
|
"page_hash",
|
||||||
"sample",
|
"sample",
|
||||||
"BaseSamplingStrategy",
|
"BaseSamplingStrategy",
|
||||||
|
|||||||
@@ -110,6 +110,7 @@ def _create_engine(
|
|||||||
device: str = "cuda",
|
device: str = "cuda",
|
||||||
dtype: torch.dtype = torch.bfloat16,
|
dtype: torch.dtype = torch.bfloat16,
|
||||||
max_batch_size: int = 16,
|
max_batch_size: int = 16,
|
||||||
|
max_seq_len: Optional[int] = None,
|
||||||
) -> InferenceEngine:
|
) -> InferenceEngine:
|
||||||
if not param_path.exists():
|
if not param_path.exists():
|
||||||
raise FileNotFoundError(f"Parameter directory not found: {param_path}")
|
raise FileNotFoundError(f"Parameter directory not found: {param_path}")
|
||||||
@@ -123,6 +124,7 @@ def _create_engine(
|
|||||||
model=model,
|
model=model,
|
||||||
tokenizer=tokenizer,
|
tokenizer=tokenizer,
|
||||||
max_batch_size=max_batch_size,
|
max_batch_size=max_batch_size,
|
||||||
|
max_seq_len=max_seq_len,
|
||||||
)
|
)
|
||||||
logger.info(f"Inference engine initialized with max_batch_size={max_batch_size}")
|
logger.info(f"Inference engine initialized with max_batch_size={max_batch_size}")
|
||||||
return engine
|
return engine
|
||||||
@@ -186,6 +188,7 @@ def run_server(
|
|||||||
device: str = "cuda",
|
device: str = "cuda",
|
||||||
dtype: torch.dtype = torch.bfloat16,
|
dtype: torch.dtype = torch.bfloat16,
|
||||||
max_batch_size: int = 16,
|
max_batch_size: int = 16,
|
||||||
|
max_seq_len: Optional[int] = None,
|
||||||
):
|
):
|
||||||
app = get_app()
|
app = get_app()
|
||||||
app.state.server_config = {
|
app.state.server_config = {
|
||||||
@@ -193,6 +196,7 @@ def run_server(
|
|||||||
"dtype": dtype,
|
"dtype": dtype,
|
||||||
"param_path": param_path,
|
"param_path": param_path,
|
||||||
"max_batch_size": max_batch_size,
|
"max_batch_size": max_batch_size,
|
||||||
|
"max_seq_len": max_seq_len,
|
||||||
}
|
}
|
||||||
uvicorn.run(
|
uvicorn.run(
|
||||||
app,
|
app,
|
||||||
|
|||||||
@@ -22,13 +22,10 @@ class BaseToolParser(ABC):
|
|||||||
Maintains streaming state internally so that each call to :meth:`feed`
|
Maintains streaming state internally so that each call to :meth:`feed`
|
||||||
can diff against previously emitted content.
|
can diff against previously emitted content.
|
||||||
|
|
||||||
Parameters
|
Args:
|
||||||
----------
|
tools (list of dict, optional): Tool definitions from the request.
|
||||||
tools : list of dict, optional
|
tool_choice (str): ``"auto"`` / ``"required"`` / ``"none"`` or a named
|
||||||
Tool definitions from the request.
|
tool choice dict.
|
||||||
tool_choice : str
|
|
||||||
``"auto"`` / ``"required"`` / ``"none"`` or a named tool choice
|
|
||||||
dict.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, tools: Optional[List[Dict]] = None, tool_choice: str = "auto"):
|
def __init__(self, tools: Optional[List[Dict]] = None, tool_choice: str = "auto"):
|
||||||
@@ -51,14 +48,12 @@ class BaseToolParser(ABC):
|
|||||||
|
|
||||||
Returns an empty list when nothing new should be emitted.
|
Returns an empty list when nothing new should be emitted.
|
||||||
|
|
||||||
Parameters
|
Args:
|
||||||
----------
|
body (str): The complete accumulated generated text so far.
|
||||||
body : str
|
current_token_ids (list of int, optional): All token IDs decoded
|
||||||
The complete accumulated generated text so far.
|
into *body* (cumulative).
|
||||||
current_token_ids : list of int, optional
|
delta_token_ids (list of int, optional): Only the token IDs for
|
||||||
All token IDs decoded into *body* (cumulative).
|
this chunk.
|
||||||
delta_token_ids : list of int, optional
|
|
||||||
Only the token IDs for this chunk.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
|
|||||||
@@ -2,16 +2,11 @@
|
|||||||
|
|
||||||
from astrai.inference.core.cache import (
|
from astrai.inference.core.cache import (
|
||||||
Allocator,
|
Allocator,
|
||||||
CacheView,
|
|
||||||
ContiguousCache,
|
|
||||||
ContiguousCacheView,
|
|
||||||
KVCache,
|
KVCache,
|
||||||
PageCache,
|
KVStorage,
|
||||||
PageCacheView,
|
|
||||||
PagePool,
|
PagePool,
|
||||||
PrefixCache,
|
PrefixCache,
|
||||||
Storage,
|
ReqToTokenPool,
|
||||||
TaskTable,
|
|
||||||
page_hash,
|
page_hash,
|
||||||
)
|
)
|
||||||
from astrai.inference.core.executor import Executor
|
from astrai.inference.core.executor import Executor
|
||||||
@@ -20,16 +15,11 @@ from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
|
|||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"Allocator",
|
"Allocator",
|
||||||
"CacheView",
|
|
||||||
"KVCache",
|
"KVCache",
|
||||||
"ContiguousCache",
|
"KVStorage",
|
||||||
"ContiguousCacheView",
|
|
||||||
"PageCache",
|
|
||||||
"PageCacheView",
|
|
||||||
"PagePool",
|
"PagePool",
|
||||||
"PrefixCache",
|
"PrefixCache",
|
||||||
"Storage",
|
"ReqToTokenPool",
|
||||||
"TaskTable",
|
|
||||||
"page_hash",
|
"page_hash",
|
||||||
"Executor",
|
"Executor",
|
||||||
"InferenceScheduler",
|
"InferenceScheduler",
|
||||||
|
|||||||
+333
-357
@@ -1,7 +1,21 @@
|
|||||||
|
"""KV cache architecture: three-layer separation (SGLang-inspired).
|
||||||
|
|
||||||
|
Layer 1 — KVStorage: flat token-level K/V buffers [n_layers, size, H, D]
|
||||||
|
Layer 2 — ReqToTokenPool: index table [req_idx, pos] → physical token slot
|
||||||
|
Layer 3 — Allocator: slot/page allocation with ref-counting and LRU
|
||||||
|
|
||||||
|
PagePool orchestrates all three plus PrefixCache (content addressing).
|
||||||
|
KVCache is a pure dataclass passed to the model for direct buffer access.
|
||||||
|
|
||||||
|
Two modes:
|
||||||
|
- contiguous (default): pre-allocated per-request blocks, no dynamic alloc
|
||||||
|
- paged: shared pool with on-demand allocation, prefix caching support
|
||||||
|
"""
|
||||||
|
|
||||||
import threading
|
import threading
|
||||||
from abc import ABC, abstractmethod
|
|
||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
from typing import Callable, Dict, List, Optional, Tuple
|
from dataclasses import dataclass
|
||||||
|
from typing import Callable, Dict, List, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
@@ -108,418 +122,380 @@ class PrefixCache:
|
|||||||
self._hash_to_page[h] = page_idx
|
self._hash_to_page[h] = page_idx
|
||||||
|
|
||||||
|
|
||||||
class PagePool:
|
class ReqToTokenPool:
|
||||||
"""Orchestrates allocator (page management) and PrefixCache (content addressing)."""
|
"""Maps [req_idx, pos] -> physical token slot in KV storage.
|
||||||
|
|
||||||
def __init__(self, allocator: Allocator, prefix: PrefixCache):
|
Each row is one request; each column is a sequence position. The value
|
||||||
self._alloc = allocator
|
at [req_idx, pos] is the flat index into the KV storage buffers.
|
||||||
self._prefix = prefix
|
"""
|
||||||
self._alloc.on_evict = prefix.evict
|
|
||||||
|
|
||||||
@property
|
def __init__(self, size: int, max_context_len: int, device: torch.device):
|
||||||
def allocator(self) -> Allocator:
|
self.size = size
|
||||||
return self._alloc
|
self.max_context_len = max_context_len
|
||||||
|
self.req_to_token = torch.zeros(
|
||||||
@property
|
(size, max_context_len), dtype=torch.long, device=device
|
||||||
def prefix(self) -> PrefixCache:
|
)
|
||||||
return self._prefix
|
self.free_slots = list(range(size))
|
||||||
|
|
||||||
def alloc(self) -> int:
|
|
||||||
return self._alloc.alloc()
|
|
||||||
|
|
||||||
def free(self, idx: int):
|
|
||||||
keep = self._prefix.has_page(idx)
|
|
||||||
self._alloc.free(idx, keep_cached=keep)
|
|
||||||
if not keep:
|
|
||||||
self._prefix.evict(idx)
|
|
||||||
|
|
||||||
def inc_ref(self, idx: int):
|
|
||||||
self._alloc.inc_ref(idx)
|
|
||||||
|
|
||||||
def lookup(self, token_ids: List[int]) -> List[int]:
|
|
||||||
hits = self._prefix.lookup(token_ids)
|
|
||||||
for p in hits:
|
|
||||||
self._alloc.touch(p)
|
|
||||||
return hits
|
|
||||||
|
|
||||||
def record(self, page_idx: int, token_ids: List[int], logical_page_idx: int):
|
|
||||||
self._prefix.record(page_idx, token_ids, logical_page_idx)
|
|
||||||
|
|
||||||
|
|
||||||
class TaskTable:
|
|
||||||
"""Maps task_ids to page tables and cached token counts."""
|
|
||||||
|
|
||||||
def __init__(self, page_size: int):
|
|
||||||
self._page_size = page_size
|
|
||||||
self._pages: Dict[str, List[int]] = {}
|
|
||||||
self._cached: Dict[str, int] = {}
|
|
||||||
self._lock = threading.Lock()
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
def set(self, task_id: str, page_table: List[int], cached: int):
|
def alloc(self, num_reqs: int) -> Optional[List[int]]:
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self._pages[task_id] = page_table
|
if num_reqs > len(self.free_slots):
|
||||||
self._cached[task_id] = cached
|
return None
|
||||||
|
slots = self.free_slots[:num_reqs]
|
||||||
|
self.free_slots = self.free_slots[num_reqs:]
|
||||||
|
return slots
|
||||||
|
|
||||||
def get(self, task_id: str) -> List[int]:
|
def free(self, req_indices: List[int]):
|
||||||
with self._lock:
|
with self._lock:
|
||||||
return self._pages.get(task_id, [])
|
self.free_slots.extend(req_indices)
|
||||||
|
|
||||||
def get_cached(self, task_id: str) -> int:
|
def write(self, indices, values):
|
||||||
with self._lock:
|
self.req_to_token[indices] = values
|
||||||
return self._cached.get(task_id, 0)
|
|
||||||
|
|
||||||
def pop(self, task_id: str) -> Tuple[List[int], int]:
|
|
||||||
with self._lock:
|
|
||||||
pages = self._pages.pop(task_id, [])
|
|
||||||
cached = self._cached.pop(task_id, 0)
|
|
||||||
return pages, cached
|
|
||||||
|
|
||||||
def get_ref(self, task_id: str) -> List[int]:
|
|
||||||
with self._lock:
|
|
||||||
return self._pages.setdefault(task_id, [])
|
|
||||||
|
|
||||||
def table_tensor(self, task_ids: List[str], device: torch.device) -> Tensor:
|
|
||||||
with self._lock:
|
|
||||||
states = [self._pages.get(tid, []) for tid in task_ids]
|
|
||||||
max_pages = max((len(s) for s in states), default=0)
|
|
||||||
rows = [s + [-1] * (max_pages - len(s)) for s in states]
|
|
||||||
return torch.tensor(rows, dtype=torch.long, device=device)
|
|
||||||
|
|
||||||
|
|
||||||
class Storage:
|
class KVStorage:
|
||||||
"""KV-cache tensor storage with paged write/gather."""
|
"""Token-level KV cache storage.
|
||||||
|
|
||||||
|
Buffers: [n_layers, size, n_kv_heads, head_dim]. Each token occupies
|
||||||
|
one slot indexed by ReqToTokenPool.
|
||||||
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
|
size: int,
|
||||||
n_layers: int,
|
n_layers: int,
|
||||||
n_pages: int,
|
|
||||||
page_size: int,
|
|
||||||
n_kv_heads: int,
|
n_kv_heads: int,
|
||||||
head_dim: int,
|
head_dim: int,
|
||||||
device: torch.device,
|
device: torch.device,
|
||||||
dtype: torch.dtype,
|
dtype: torch.dtype,
|
||||||
):
|
):
|
||||||
self.page_size = page_size
|
self.size = size
|
||||||
self.k_cache = torch.empty(
|
self.k_buffer = torch.empty(
|
||||||
(n_layers, n_pages, page_size, n_kv_heads, head_dim),
|
(n_layers, size, n_kv_heads, head_dim), device=device, dtype=dtype
|
||||||
device=device,
|
|
||||||
dtype=dtype,
|
|
||||||
)
|
)
|
||||||
self.v_cache = torch.empty(
|
self.v_buffer = torch.empty(
|
||||||
(n_layers, n_pages, page_size, n_kv_heads, head_dim),
|
(n_layers, size, n_kv_heads, head_dim), device=device, dtype=dtype
|
||||||
device=device,
|
|
||||||
dtype=dtype,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def write(
|
def get_key_buffer(self, layer_id: int) -> Tensor:
|
||||||
self,
|
return self.k_buffer[layer_id]
|
||||||
layer_id: int,
|
|
||||||
page_table: Tensor,
|
|
||||||
start_pos: int,
|
|
||||||
k: Tensor,
|
|
||||||
v: Tensor,
|
|
||||||
):
|
|
||||||
seq_len = k.size(1)
|
|
||||||
if seq_len == 0:
|
|
||||||
return
|
|
||||||
page_size = self.page_size
|
|
||||||
written = 0
|
|
||||||
first_page = start_pos // page_size
|
|
||||||
last_page = (start_pos + seq_len - 1) // page_size
|
|
||||||
for pi in range(first_page, last_page + 1):
|
|
||||||
phys_pages = page_table[:, pi]
|
|
||||||
page_start = pi * page_size
|
|
||||||
write_start = max(page_start, start_pos)
|
|
||||||
write_end = min(page_start + page_size, start_pos + seq_len)
|
|
||||||
offset = write_start - page_start
|
|
||||||
chunk = write_end - write_start
|
|
||||||
valid = phys_pages >= 0
|
|
||||||
if not valid.all():
|
|
||||||
if valid.any():
|
|
||||||
valid_pages = phys_pages[valid]
|
|
||||||
self.k_cache[layer_id, valid_pages, offset : offset + chunk] = k[
|
|
||||||
valid, written : written + chunk
|
|
||||||
]
|
|
||||||
self.v_cache[layer_id, valid_pages, offset : offset + chunk] = v[
|
|
||||||
valid, written : written + chunk
|
|
||||||
]
|
|
||||||
written += chunk
|
|
||||||
continue
|
|
||||||
self.k_cache[layer_id, phys_pages, offset : offset + chunk] = k[
|
|
||||||
:, written : written + chunk
|
|
||||||
]
|
|
||||||
self.v_cache[layer_id, phys_pages, offset : offset + chunk] = v[
|
|
||||||
:, written : written + chunk
|
|
||||||
]
|
|
||||||
written += chunk
|
|
||||||
|
|
||||||
def gather(
|
def get_value_buffer(self, layer_id: int) -> Tensor:
|
||||||
self, layer_id: int, page_table: Tensor, total_len: int
|
return self.v_buffer[layer_id]
|
||||||
) -> Tuple[Tensor, Tensor]:
|
|
||||||
safe = page_table.clamp(min=0)
|
def set_kv_buffer(self, layer_id: int, loc: Tensor, k: Tensor, v: Tensor) -> None:
|
||||||
k = self.k_cache[layer_id, safe]
|
self.k_buffer[layer_id, loc] = k
|
||||||
v = self.v_cache[layer_id, safe]
|
self.v_buffer[layer_id, loc] = v
|
||||||
k = k.flatten(1, 2)
|
|
||||||
v = v.flatten(1, 2)
|
|
||||||
if (page_table < 0).any():
|
|
||||||
invalid = (
|
|
||||||
(page_table < 0)
|
|
||||||
.unsqueeze(-1)
|
|
||||||
.expand(-1, -1, self.page_size)
|
|
||||||
.flatten(1, 2)
|
|
||||||
)
|
|
||||||
invalid = invalid[:, :, None, None].expand_as(k)
|
|
||||||
k = k.masked_fill(invalid, 0.0)
|
|
||||||
v = v.masked_fill(invalid, 0.0)
|
|
||||||
k = k[:, :total_len]
|
|
||||||
v = v[:, :total_len]
|
|
||||||
return k, v
|
|
||||||
|
|
||||||
|
|
||||||
class CacheView(ABC):
|
@dataclass
|
||||||
"""Abstract view passed to attention layers for KV-cache I/O."""
|
class KVCache:
|
||||||
|
"""Pure data struct passed to model for KV cache I/O.
|
||||||
|
|
||||||
@abstractmethod
|
The attention layer does raw buffer indexing — no methods, no abstraction.
|
||||||
def write(self, layer_id: int, k: Tensor, v: Tensor): ...
|
|
||||||
|
|
||||||
@abstractmethod
|
Attributes:
|
||||||
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]: ...
|
k_buffer: [n_layers, size, n_kv_heads, head_dim]
|
||||||
|
v_buffer: [n_layers, size, n_kv_heads, head_dim]
|
||||||
|
req_to_token: [num_reqs, max_ctx_len] — index table
|
||||||
|
req_pool_indices: [batch_size] — row indices into req_to_token
|
||||||
|
seq_lens: [batch_size] — per-request total sequence lengths
|
||||||
|
out_cache_loc: [batch, new_seq_len] or [batch, 1] — write indices
|
||||||
|
max_len: max(seq_lens) as Python int — avoids GPU sync in decode
|
||||||
|
page_table: [batch, max_len] — precomputed gather indices for decode;
|
||||||
|
None for prefill or when not yet computed.
|
||||||
|
decode_mask: [batch, max_len] bool — precomputed position validity
|
||||||
|
mask for decode; None for prefill or single-batch decode.
|
||||||
|
"""
|
||||||
|
|
||||||
|
k_buffer: Tensor
|
||||||
|
v_buffer: Tensor
|
||||||
|
req_to_token: Tensor
|
||||||
|
req_pool_indices: Tensor
|
||||||
|
seq_lens: Tensor
|
||||||
|
out_cache_loc: Tensor
|
||||||
|
max_len: int = 0
|
||||||
|
page_table: Optional[Tensor] = None
|
||||||
|
decode_mask: Optional[Tensor] = None
|
||||||
|
|
||||||
|
|
||||||
class KVCache(ABC):
|
class PagePool:
|
||||||
"""Abstract KV-cache facade for scheduler/executor."""
|
"""Top-level KV cache manager.
|
||||||
|
|
||||||
@abstractmethod
|
Combines KVStorage + ReqToTokenPool + Allocator + PrefixCache.
|
||||||
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool: ...
|
|
||||||
|
|
||||||
@abstractmethod
|
Args:
|
||||||
def task_free(self, task_id: str): ...
|
n_layers: Number of transformer layers.
|
||||||
|
n_kv_heads: Number of KV attention heads.
|
||||||
@abstractmethod
|
head_dim: Dimension per head.
|
||||||
def task_extend(self, task_id: str, pos: int) -> bool: ...
|
max_batch_size: Maximum concurrent requests.
|
||||||
|
max_seq_len: Maximum sequence length per request.
|
||||||
@abstractmethod
|
device, dtype: Tensor device and dtype.
|
||||||
def bind_tasks(
|
page_size: Page size for paged mode (1 = token-level).
|
||||||
self,
|
n_tokens: Total token slots for paged mode. None = contiguous mode
|
||||||
task_ids: List[str],
|
(pre-allocates max_batch_size * max_seq_len).
|
||||||
total_len: int,
|
"""
|
||||||
device: torch.device,
|
|
||||||
write_positions: Optional[Tensor] = None,
|
|
||||||
) -> CacheView: ...
|
|
||||||
|
|
||||||
def task_cached(self, task_id: str) -> int:
|
|
||||||
return 0
|
|
||||||
|
|
||||||
def task_record_hashes(
|
|
||||||
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
|
|
||||||
): ...
|
|
||||||
|
|
||||||
|
|
||||||
class PageCacheView(CacheView):
|
|
||||||
"""Bundles Storage + page_table + total_len for attention layers."""
|
|
||||||
|
|
||||||
def __init__(self, storage: Storage, page_table: Tensor, total_len: int = 0):
|
|
||||||
self._storage = storage
|
|
||||||
self._page_table = page_table
|
|
||||||
self._total_len = total_len
|
|
||||||
|
|
||||||
def write(self, layer_id: int, k: Tensor, v: Tensor):
|
|
||||||
start_pos = self._total_len - k.size(1)
|
|
||||||
self._storage.write(layer_id, self._page_table, start_pos, k, v)
|
|
||||||
|
|
||||||
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]:
|
|
||||||
return self._storage.gather(layer_id, self._page_table, self._total_len)
|
|
||||||
|
|
||||||
|
|
||||||
class PageCache(KVCache):
|
|
||||||
"""Paged KV-cache with prefix sharing."""
|
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
n_layers: int,
|
n_layers: int,
|
||||||
n_pages: int,
|
|
||||||
page_size: int,
|
|
||||||
n_kv_heads: int,
|
n_kv_heads: int,
|
||||||
head_dim: int,
|
head_dim: int,
|
||||||
device: torch.device,
|
|
||||||
dtype: torch.dtype,
|
|
||||||
):
|
|
||||||
self.page_size = page_size
|
|
||||||
self._pool = PagePool(Allocator(n_pages), PrefixCache(page_size))
|
|
||||||
self._table = TaskTable(page_size)
|
|
||||||
self._storage = Storage(
|
|
||||||
n_layers, n_pages, page_size, n_kv_heads, head_dim, device, dtype
|
|
||||||
)
|
|
||||||
|
|
||||||
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
|
|
||||||
hits = self._pool.lookup(prompt_ids)
|
|
||||||
cached = len(hits) * self.page_size
|
|
||||||
for p in hits:
|
|
||||||
self._pool.inc_ref(p)
|
|
||||||
|
|
||||||
remaining = len(prompt_ids) - cached
|
|
||||||
n_new = (
|
|
||||||
(remaining + self.page_size - 1) // self.page_size if remaining > 0 else 0
|
|
||||||
)
|
|
||||||
new_pages: List[int] = []
|
|
||||||
if n_new > 0:
|
|
||||||
for _ in range(n_new):
|
|
||||||
p = self._pool.alloc()
|
|
||||||
if p < 0:
|
|
||||||
for hp in hits:
|
|
||||||
self._pool.free(hp)
|
|
||||||
for np in new_pages:
|
|
||||||
self._pool.free(np)
|
|
||||||
return False
|
|
||||||
new_pages.append(p)
|
|
||||||
|
|
||||||
self._table.set(task_id, hits + new_pages, cached)
|
|
||||||
return True
|
|
||||||
|
|
||||||
def task_free(self, task_id: str):
|
|
||||||
page_table, _ = self._table.pop(task_id)
|
|
||||||
for idx in page_table:
|
|
||||||
self._pool.free(idx)
|
|
||||||
|
|
||||||
def task_extend(self, task_id: str, pos: int) -> bool:
|
|
||||||
page_table = self._table.get(task_id)
|
|
||||||
needed = (pos + 1 + self.page_size - 1) // self.page_size
|
|
||||||
while len(page_table) < needed:
|
|
||||||
p = self._pool.alloc()
|
|
||||||
if p < 0:
|
|
||||||
return False
|
|
||||||
page_table.append(p)
|
|
||||||
return True
|
|
||||||
|
|
||||||
def task_cached(self, task_id: str) -> int:
|
|
||||||
return self._table.get_cached(task_id)
|
|
||||||
|
|
||||||
def task_record_hashes(
|
|
||||||
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
|
|
||||||
):
|
|
||||||
page_table = self._table.get(task_id)
|
|
||||||
full_pages = len(prompt_ids) // self.page_size
|
|
||||||
for i in range(start_logical_page, full_pages):
|
|
||||||
self._pool.record(page_table[i], prompt_ids, i)
|
|
||||||
|
|
||||||
def bind_tasks(
|
|
||||||
self,
|
|
||||||
task_ids: List[str],
|
|
||||||
total_len: int,
|
|
||||||
device: torch.device,
|
|
||||||
write_positions: Optional[Tensor] = None,
|
|
||||||
) -> PageCacheView:
|
|
||||||
page_table = self._table.table_tensor(task_ids, device)
|
|
||||||
return PageCacheView(self._storage, page_table, total_len)
|
|
||||||
|
|
||||||
|
|
||||||
class ContiguousCacheView(CacheView):
|
|
||||||
"""Contiguous KV-cache view for attention layers."""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
cache: "ContiguousCache",
|
|
||||||
batch_indices: Tensor,
|
|
||||||
total_len: int = 0,
|
|
||||||
write_positions: Optional[Tensor] = None,
|
|
||||||
):
|
|
||||||
self._cache = cache
|
|
||||||
self._batch_indices = batch_indices
|
|
||||||
self._total_len = total_len
|
|
||||||
self._write_positions = write_positions
|
|
||||||
|
|
||||||
def write(self, layer_id: int, k: Tensor, v: Tensor):
|
|
||||||
seq_len = k.size(1)
|
|
||||||
indices = self._batch_indices
|
|
||||||
if self._write_positions is not None and seq_len == 1:
|
|
||||||
pos = self._write_positions
|
|
||||||
self._cache.k[layer_id, indices, pos] = k.squeeze(1)
|
|
||||||
self._cache.v[layer_id, indices, pos] = v.squeeze(1)
|
|
||||||
else:
|
|
||||||
start_pos = self._total_len - seq_len
|
|
||||||
self._cache.k[layer_id, indices, start_pos : start_pos + seq_len] = k
|
|
||||||
self._cache.v[layer_id, indices, start_pos : start_pos + seq_len] = v
|
|
||||||
|
|
||||||
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]:
|
|
||||||
max_len = self._total_len
|
|
||||||
indices = self._batch_indices
|
|
||||||
k = self._cache.k[layer_id, indices, :max_len]
|
|
||||||
v = self._cache.v[layer_id, indices, :max_len]
|
|
||||||
return k, v
|
|
||||||
|
|
||||||
|
|
||||||
class ContiguousCache(KVCache):
|
|
||||||
"""Contiguous per-slot KV cache (default implementation)."""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
n_layers: int,
|
|
||||||
max_batch_size: int,
|
max_batch_size: int,
|
||||||
max_seq_len: int,
|
max_seq_len: int,
|
||||||
n_kv_heads: int,
|
|
||||||
head_dim: int,
|
|
||||||
device: torch.device,
|
device: torch.device,
|
||||||
dtype: torch.dtype,
|
dtype: torch.dtype,
|
||||||
|
page_size: int = 1,
|
||||||
|
n_tokens: Optional[int] = None,
|
||||||
):
|
):
|
||||||
|
self.page_size = page_size
|
||||||
|
self.max_batch_size = max_batch_size
|
||||||
self.max_seq_len = max_seq_len
|
self.max_seq_len = max_seq_len
|
||||||
self.k = torch.zeros(
|
self.device = device
|
||||||
n_layers,
|
self.dtype = dtype
|
||||||
max_batch_size,
|
self.n_layers = n_layers
|
||||||
max_seq_len,
|
self.n_kv_heads = n_kv_heads
|
||||||
n_kv_heads,
|
self.head_dim = head_dim
|
||||||
head_dim,
|
|
||||||
device=device,
|
self.contiguous = n_tokens is None
|
||||||
dtype=dtype,
|
if self.contiguous:
|
||||||
|
self.n_tokens = max_batch_size * max_seq_len
|
||||||
|
else:
|
||||||
|
self.n_tokens = n_tokens
|
||||||
|
|
||||||
|
self._storage = KVStorage(
|
||||||
|
self.n_tokens, n_layers, n_kv_heads, head_dim, device, dtype
|
||||||
)
|
)
|
||||||
self.v = torch.zeros(
|
self._req_pool = ReqToTokenPool(max_batch_size, max_seq_len, device)
|
||||||
n_layers,
|
|
||||||
max_batch_size,
|
if self.contiguous:
|
||||||
max_seq_len,
|
for i in range(max_batch_size):
|
||||||
n_kv_heads,
|
self._req_pool.req_to_token[i] = torch.arange(
|
||||||
head_dim,
|
i * max_seq_len, (i + 1) * max_seq_len, device=device
|
||||||
device=device,
|
)
|
||||||
dtype=dtype,
|
self._alloc: Optional[Allocator] = None
|
||||||
)
|
self._prefix: Optional[PrefixCache] = None
|
||||||
self._slot_len: Dict[int, int] = {}
|
else:
|
||||||
self._task_slot: Dict[str, int] = {}
|
n_pages = self.n_tokens // page_size
|
||||||
self._free_slots = list(range(max_batch_size))
|
self._alloc = Allocator(n_pages)
|
||||||
self._device = device
|
self._prefix = PrefixCache(page_size) if page_size > 1 else None
|
||||||
|
if self._prefix is not None:
|
||||||
|
self._alloc.on_evict = self._prefix.evict
|
||||||
|
|
||||||
|
self._task_req: Dict[str, int] = {}
|
||||||
|
self._task_len: Dict[int, int] = {}
|
||||||
|
self._task_cached: Dict[str, int] = {}
|
||||||
|
self._task_slots: Dict[str, List[int]] = {}
|
||||||
|
self._task_pages: Dict[str, List[int]] = {}
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
|
# ---- task lifecycle ----
|
||||||
|
|
||||||
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
|
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
|
||||||
if not self._free_slots:
|
req_slots = self._req_pool.alloc(1)
|
||||||
|
if req_slots is None:
|
||||||
return False
|
return False
|
||||||
slot = self._free_slots.pop(0)
|
req_idx = req_slots[0]
|
||||||
self._task_slot[task_id] = slot
|
self._task_req[task_id] = req_idx
|
||||||
self._slot_len[slot] = 0
|
|
||||||
|
if self.contiguous:
|
||||||
|
self._task_len[req_idx] = len(prompt_ids)
|
||||||
|
self._task_cached[task_id] = 0
|
||||||
|
return True
|
||||||
|
|
||||||
|
n_tokens_needed = len(prompt_ids)
|
||||||
|
cached = 0
|
||||||
|
|
||||||
|
if self._prefix is not None:
|
||||||
|
hits = self._prefix.lookup(prompt_ids)
|
||||||
|
cached = len(hits) * self.page_size
|
||||||
|
for p in hits:
|
||||||
|
self._alloc.inc_ref(p)
|
||||||
|
self._task_pages[task_id] = list(hits)
|
||||||
|
self._task_slots[task_id] = []
|
||||||
|
else:
|
||||||
|
self._task_pages[task_id] = []
|
||||||
|
self._task_slots[task_id] = []
|
||||||
|
|
||||||
|
remaining = n_tokens_needed - cached
|
||||||
|
if remaining > 0:
|
||||||
|
if self.page_size == 1:
|
||||||
|
slots = self._alloc_tokens(remaining)
|
||||||
|
if slots is None:
|
||||||
|
for p in self._task_pages[task_id]:
|
||||||
|
self._alloc.free(p)
|
||||||
|
self._req_pool.free([req_idx])
|
||||||
|
del self._task_req[task_id]
|
||||||
|
return False
|
||||||
|
self._task_slots[task_id] = slots
|
||||||
|
else:
|
||||||
|
n_new_pages = (remaining + self.page_size - 1) // self.page_size
|
||||||
|
new_pages = []
|
||||||
|
for _ in range(n_new_pages):
|
||||||
|
p = self._alloc.alloc()
|
||||||
|
if p < 0:
|
||||||
|
for hp in self._task_pages[task_id]:
|
||||||
|
self._alloc.free(hp)
|
||||||
|
for np_ in new_pages:
|
||||||
|
self._alloc.free(np_)
|
||||||
|
self._req_pool.free([req_idx])
|
||||||
|
del self._task_req[task_id]
|
||||||
|
return False
|
||||||
|
new_pages.append(p)
|
||||||
|
self._task_pages[task_id].extend(new_pages)
|
||||||
|
|
||||||
|
self._write_req_to_token(task_id, prompt_ids, cached)
|
||||||
|
self._task_len[req_idx] = len(prompt_ids)
|
||||||
|
self._task_cached[task_id] = cached
|
||||||
return True
|
return True
|
||||||
|
|
||||||
def task_free(self, task_id: str):
|
def task_free(self, task_id: str):
|
||||||
slot = self._task_slot.pop(task_id, None)
|
req_idx = self._task_req.pop(task_id, None)
|
||||||
if slot is not None:
|
if req_idx is None:
|
||||||
self._slot_len.pop(slot, None)
|
return
|
||||||
self._free_slots.append(slot)
|
self._task_len.pop(req_idx, None)
|
||||||
|
self._task_cached.pop(task_id, None)
|
||||||
|
|
||||||
|
if not self.contiguous:
|
||||||
|
if self._prefix is not None:
|
||||||
|
for p in self._task_pages.get(task_id, []):
|
||||||
|
keep = self._prefix.has_page(p)
|
||||||
|
self._alloc.free(p, keep_cached=keep)
|
||||||
|
if not keep:
|
||||||
|
self._prefix.evict(p)
|
||||||
|
else:
|
||||||
|
for p in self._task_pages.get(task_id, []):
|
||||||
|
self._alloc.free(p)
|
||||||
|
self._task_pages.pop(task_id, None)
|
||||||
|
self._task_slots.pop(task_id, None)
|
||||||
|
|
||||||
|
self._req_pool.free([req_idx])
|
||||||
|
|
||||||
def task_extend(self, task_id: str, pos: int) -> bool:
|
def task_extend(self, task_id: str, pos: int) -> bool:
|
||||||
return pos < self.max_seq_len
|
req_idx = self._task_req.get(task_id)
|
||||||
|
if req_idx is None:
|
||||||
|
return False
|
||||||
|
|
||||||
|
if self.contiguous:
|
||||||
|
return pos < self.max_seq_len
|
||||||
|
|
||||||
|
if self.page_size == 1:
|
||||||
|
slots = self._alloc_tokens(1)
|
||||||
|
if slots is None:
|
||||||
|
return False
|
||||||
|
self._task_slots.setdefault(task_id, []).extend(slots)
|
||||||
|
self._req_pool.req_to_token[req_idx, pos] = slots[0]
|
||||||
|
else:
|
||||||
|
page_idx = pos // self.page_size
|
||||||
|
existing = self._task_pages.get(task_id, [])
|
||||||
|
if page_idx >= len(existing):
|
||||||
|
p = self._alloc.alloc()
|
||||||
|
if p < 0:
|
||||||
|
return False
|
||||||
|
existing.append(p)
|
||||||
|
self._task_pages[task_id] = existing
|
||||||
|
page_offset = pos % self.page_size
|
||||||
|
page = existing[page_idx]
|
||||||
|
token_slot = page * self.page_size + page_offset
|
||||||
|
self._req_pool.req_to_token[req_idx, pos] = token_slot
|
||||||
|
|
||||||
|
self._task_len[req_idx] = pos + 1
|
||||||
|
return True
|
||||||
|
|
||||||
def task_cached(self, task_id: str) -> int:
|
def task_cached(self, task_id: str) -> int:
|
||||||
slot = self._task_slot.get(task_id)
|
return self._task_cached.get(task_id, 0)
|
||||||
if slot is None:
|
|
||||||
return 0
|
def task_record_hashes(
|
||||||
return self._slot_len.get(slot, 0)
|
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
|
||||||
|
):
|
||||||
|
if self._prefix is None or self.contiguous:
|
||||||
|
return
|
||||||
|
pages = self._task_pages.get(task_id, [])
|
||||||
|
full_pages = len(prompt_ids) // self.page_size
|
||||||
|
for i in range(start_logical_page, min(full_pages, len(pages))):
|
||||||
|
self._prefix.record(pages[i], prompt_ids, i)
|
||||||
|
|
||||||
|
# ---- bind for forward ----
|
||||||
|
|
||||||
def bind_tasks(
|
def bind_tasks(
|
||||||
self,
|
self,
|
||||||
task_ids: List[str],
|
task_ids: List[str],
|
||||||
total_len: int,
|
seq_lens: List[int],
|
||||||
device: torch.device,
|
device: torch.device,
|
||||||
write_positions: Optional[Tensor] = None,
|
start_pos: Optional[int] = None,
|
||||||
) -> ContiguousCacheView:
|
) -> KVCache:
|
||||||
slots = [self._task_slot[tid] for tid in task_ids]
|
req_indices = [self._task_req[tid] for tid in task_ids]
|
||||||
batch_indices = torch.tensor(slots, dtype=torch.long, device=device)
|
req_pool_indices = torch.tensor(req_indices, dtype=torch.long, device=device)
|
||||||
for slot in slots:
|
seq_lens_t = torch.tensor(seq_lens, dtype=torch.long, device=device)
|
||||||
if total_len > self._slot_len.get(slot, 0):
|
|
||||||
self._slot_len[slot] = total_len
|
if start_pos is not None:
|
||||||
return ContiguousCacheView(
|
seq_len = seq_lens[0]
|
||||||
self, batch_indices, total_len, write_positions=write_positions
|
out_cache_loc = self._req_pool.req_to_token[
|
||||||
|
req_pool_indices, start_pos:seq_len
|
||||||
|
]
|
||||||
|
page_table = None
|
||||||
|
decode_mask = None
|
||||||
|
else:
|
||||||
|
write_pos = seq_lens_t - 1
|
||||||
|
out_cache_loc = self._req_pool.req_to_token[
|
||||||
|
req_pool_indices, write_pos
|
||||||
|
].unsqueeze(-1)
|
||||||
|
ml = max(seq_lens)
|
||||||
|
page_table = self._req_pool.req_to_token[req_pool_indices, :ml]
|
||||||
|
if len(task_ids) > 1:
|
||||||
|
decode_mask = (
|
||||||
|
torch.arange(ml, device=device)[None, :] < seq_lens_t[:, None]
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
decode_mask = None
|
||||||
|
|
||||||
|
return KVCache(
|
||||||
|
k_buffer=self._storage.k_buffer,
|
||||||
|
v_buffer=self._storage.v_buffer,
|
||||||
|
req_to_token=self._req_pool.req_to_token,
|
||||||
|
req_pool_indices=req_pool_indices,
|
||||||
|
seq_lens=seq_lens_t,
|
||||||
|
out_cache_loc=out_cache_loc,
|
||||||
|
max_len=max(seq_lens),
|
||||||
|
page_table=page_table,
|
||||||
|
decode_mask=decode_mask,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# ---- internals ----
|
||||||
|
|
||||||
|
def _alloc_tokens(self, n: int) -> Optional[List[int]]:
|
||||||
|
if self.page_size != 1:
|
||||||
|
raise RuntimeError("_alloc_tokens is for page_size=1 only")
|
||||||
|
slots = []
|
||||||
|
for _ in range(n):
|
||||||
|
p = self._alloc.alloc()
|
||||||
|
if p < 0:
|
||||||
|
for s in slots:
|
||||||
|
self._alloc.free(s)
|
||||||
|
return None
|
||||||
|
slots.append(p)
|
||||||
|
return slots
|
||||||
|
|
||||||
|
def _write_req_to_token(self, task_id: str, prompt_ids: List[int], cached: int):
|
||||||
|
req_idx = self._task_req[task_id]
|
||||||
|
total = len(prompt_ids)
|
||||||
|
|
||||||
|
if self.contiguous:
|
||||||
|
return
|
||||||
|
|
||||||
|
if self.page_size == 1:
|
||||||
|
slots = self._task_slots.get(task_id, [])
|
||||||
|
all_slots = slots[: total - cached]
|
||||||
|
if all_slots:
|
||||||
|
self._req_pool.req_to_token[req_idx, cached:total] = torch.tensor(
|
||||||
|
all_slots, dtype=torch.long, device=self.device
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
pages = self._task_pages.get(task_id, [])
|
||||||
|
for pos in range(cached, total):
|
||||||
|
page_idx = pos // self.page_size
|
||||||
|
page_offset = pos % self.page_size
|
||||||
|
if page_idx < len(pages):
|
||||||
|
token_slot = pages[page_idx] * self.page_size + page_offset
|
||||||
|
self._req_pool.req_to_token[req_idx, pos] = token_slot
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ from typing import List, Optional
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from astrai.inference.core.cache import KVCache
|
from astrai.inference.core.cache import PagePool
|
||||||
from astrai.inference.core.task import Task
|
from astrai.inference.core.task import Task
|
||||||
from astrai.inference.sample import sample
|
from astrai.inference.sample import sample
|
||||||
from astrai.model.automodel import AutoModel
|
from astrai.model.automodel import AutoModel
|
||||||
@@ -19,7 +19,7 @@ class Executor:
|
|||||||
self,
|
self,
|
||||||
model: AutoModel,
|
model: AutoModel,
|
||||||
tokenizer: AutoTokenizer,
|
tokenizer: AutoTokenizer,
|
||||||
kv_cache: KVCache,
|
kv_cache: PagePool,
|
||||||
device: Optional[str] = None,
|
device: Optional[str] = None,
|
||||||
dtype: Optional[torch.dtype] = None,
|
dtype: Optional[torch.dtype] = None,
|
||||||
):
|
):
|
||||||
@@ -57,7 +57,9 @@ class Executor:
|
|||||||
input_ids,
|
input_ids,
|
||||||
input_mask=input_mask,
|
input_mask=input_mask,
|
||||||
position_ids=position_ids,
|
position_ids=position_ids,
|
||||||
paged_cache=self.kv_cache.bind_tasks(task_ids, prompt_len, self.device),
|
kv_cache=self.kv_cache.bind_tasks(
|
||||||
|
task_ids, [prompt_len] * batch_sz, self.device, start_pos=start_pos
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
def execute_decode(
|
def execute_decode(
|
||||||
@@ -103,36 +105,42 @@ class Executor:
|
|||||||
[t.frequency_penalty for t in tasks], device=self.device
|
[t.frequency_penalty for t in tasks], device=self.device
|
||||||
)
|
)
|
||||||
|
|
||||||
history_lists = []
|
has_freq = bool((freq_penalties != 0).any())
|
||||||
history_lens = []
|
if has_freq:
|
||||||
for t in tasks:
|
history_lists = []
|
||||||
window = t.rep_window
|
history_lens = []
|
||||||
prompt_part = t.prompt_ids[-window:]
|
for t in tasks:
|
||||||
ids = prompt_part + t.output_ids
|
window = t.rep_window
|
||||||
history_lists.append(ids)
|
prompt_part = t.prompt_ids[-window:]
|
||||||
history_lens.append(len(ids))
|
ids = prompt_part + t.output_ids
|
||||||
|
history_lists.append(ids)
|
||||||
|
history_lens.append(len(ids))
|
||||||
|
|
||||||
max_len = max(history_lens) if history_lens else 0
|
max_len = max(history_lens) if history_lens else 0
|
||||||
padded_ids = torch.zeros(
|
padded_ids = torch.zeros(
|
||||||
len(tasks), max_len, dtype=torch.long, device=self.device
|
len(tasks), max_len, dtype=torch.long, device=self.device
|
||||||
)
|
)
|
||||||
padded_mask = torch.zeros(
|
padded_mask = torch.zeros(
|
||||||
len(tasks), max_len, dtype=torch.bool, device=self.device
|
len(tasks), max_len, dtype=torch.bool, device=self.device
|
||||||
)
|
)
|
||||||
for i, h in enumerate(history_lists):
|
for i, h in enumerate(history_lists):
|
||||||
L = history_lens[i]
|
L = history_lens[i]
|
||||||
padded_ids[i, :L] = torch.as_tensor(h, dtype=torch.long, device=self.device)
|
padded_ids[i, :L] = torch.as_tensor(
|
||||||
padded_mask[i, :L] = True
|
h, dtype=torch.long, device=self.device
|
||||||
|
)
|
||||||
|
padded_mask[i, :L] = True
|
||||||
|
else:
|
||||||
|
padded_ids = None
|
||||||
|
padded_mask = None
|
||||||
|
|
||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
outputs = self.model(
|
outputs = self.model(
|
||||||
input_ids.unsqueeze(1),
|
input_ids.unsqueeze(1),
|
||||||
input_mask=input_mask,
|
input_mask=input_mask,
|
||||||
paged_cache=self.kv_cache.bind_tasks(
|
kv_cache=self.kv_cache.bind_tasks(
|
||||||
task_ids,
|
task_ids,
|
||||||
total_len,
|
[t.next_pos + 1 for t in tasks],
|
||||||
self.device,
|
self.device,
|
||||||
write_positions=position_ids,
|
|
||||||
),
|
),
|
||||||
position_ids=position_ids.unsqueeze(1),
|
position_ids=position_ids.unsqueeze(1),
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ from typing import Any, Dict, List, Optional, Tuple
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from astrai.inference.core.cache import ContiguousCache, KVCache
|
from astrai.inference.core.cache import PagePool
|
||||||
from astrai.inference.core.executor import Executor
|
from astrai.inference.core.executor import Executor
|
||||||
from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
|
from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
|
||||||
from astrai.model.automodel import AutoModel
|
from astrai.model.automodel import AutoModel
|
||||||
@@ -23,10 +23,9 @@ class InferenceScheduler:
|
|||||||
tokenizer: AutoTokenizer,
|
tokenizer: AutoTokenizer,
|
||||||
max_batch_size: int = 16,
|
max_batch_size: int = 16,
|
||||||
max_seq_len: Optional[int] = None,
|
max_seq_len: Optional[int] = None,
|
||||||
max_prompt_len: int = 2048,
|
|
||||||
device: Optional[str] = None,
|
device: Optional[str] = None,
|
||||||
dtype: Optional[torch.dtype] = None,
|
dtype: Optional[torch.dtype] = None,
|
||||||
cache: Optional[KVCache] = None,
|
cache: Optional[PagePool] = None,
|
||||||
):
|
):
|
||||||
config = model.config
|
config = model.config
|
||||||
|
|
||||||
@@ -47,21 +46,20 @@ class InferenceScheduler:
|
|||||||
if cache is not None:
|
if cache is not None:
|
||||||
self._cache = cache
|
self._cache = cache
|
||||||
else:
|
else:
|
||||||
self._cache = ContiguousCache(
|
self._cache = PagePool(
|
||||||
config.num_hidden_layers,
|
n_layers=config.num_hidden_layers,
|
||||||
max_batch_size,
|
n_kv_heads=config.num_key_value_heads,
|
||||||
self.max_seq_len,
|
head_dim=head_dim,
|
||||||
config.num_key_value_heads,
|
max_batch_size=max_batch_size,
|
||||||
head_dim,
|
max_seq_len=self.max_seq_len,
|
||||||
self.device,
|
device=self.device,
|
||||||
self.dtype,
|
dtype=self.dtype,
|
||||||
)
|
)
|
||||||
|
|
||||||
self._task_mgr = TaskManager(
|
self._task_mgr = TaskManager(
|
||||||
tokenizer=tokenizer,
|
tokenizer=tokenizer,
|
||||||
max_batch_size=max_batch_size,
|
max_batch_size=max_batch_size,
|
||||||
max_seq_len=self.max_seq_len,
|
max_seq_len=self.max_seq_len,
|
||||||
max_prompt_len=max_prompt_len,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
self._executor = Executor(
|
self._executor = Executor(
|
||||||
@@ -111,9 +109,11 @@ class InferenceScheduler:
|
|||||||
self._task_mgr.wait_for_tasks(timeout=1.0)
|
self._task_mgr.wait_for_tasks(timeout=1.0)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
active = self._task_mgr.get_active_tasks()
|
||||||
|
|
||||||
to_prefill = [
|
to_prefill = [
|
||||||
t
|
t
|
||||||
for t in self._task_mgr.get_active_tasks()
|
for t in active
|
||||||
if t.output_tokens == 0
|
if t.output_tokens == 0
|
||||||
and cache.task_cached(t.task_id) < len(t.prompt_ids)
|
and cache.task_cached(t.task_id) < len(t.prompt_ids)
|
||||||
]
|
]
|
||||||
@@ -139,10 +139,10 @@ class InferenceScheduler:
|
|||||||
t.task_id, t.prompt_ids, start_logical_page
|
t.task_id, t.prompt_ids, start_logical_page
|
||||||
)
|
)
|
||||||
|
|
||||||
decode_tasks = self._task_mgr.get_active_tasks()
|
decode_tasks = active
|
||||||
|
|
||||||
valid: List[Task] = []
|
valid: List[Task] = []
|
||||||
for t in sorted(decode_tasks, key=lambda t: t.task_id):
|
for t in decode_tasks:
|
||||||
if cache.task_extend(t.task_id, t.next_pos):
|
if cache.task_extend(t.task_id, t.next_pos):
|
||||||
valid.append(t)
|
valid.append(t)
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -6,6 +6,8 @@ from collections import deque
|
|||||||
from enum import Enum
|
from enum import Enum
|
||||||
from typing import Any, Callable, Deque, Dict, List, Optional
|
from typing import Any, Callable, Deque, Dict, List, Optional
|
||||||
|
|
||||||
|
from tokenizers.decoders import DecodeStream
|
||||||
|
|
||||||
from astrai.tokenize.tokenizer import AutoTokenizer
|
from astrai.tokenize.tokenizer import AutoTokenizer
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -14,37 +16,30 @@ STOP = object()
|
|||||||
|
|
||||||
|
|
||||||
class StreamDecoder:
|
class StreamDecoder:
|
||||||
"""Incremental decoder for byte-level BPE streaming.
|
"""Incremental decoder backed by the tokenizers library's DecodeStream.
|
||||||
|
|
||||||
Byte-level BPE may split a single Unicode character (e.g. em-dash,
|
Delegates to the Rust-native streaming decoder which maintains an
|
||||||
smart quotes) across multiple tokens. Decoding such a token in
|
O(1) bounded token buffer internally (via prefix drain), avoiding
|
||||||
isolation produces U+FFFD (replacement char). This decoder
|
the O(n²) cost of re-decoding the full history on each step.
|
||||||
accumulates token IDs and only emits text once the trailing
|
|
||||||
characters are complete, buffering incomplete multi-byte sequences
|
Multi-byte UTF-8 sequences split across token boundaries are
|
||||||
until the next token arrives.
|
buffered until complete; ``push`` returns "" while the trailing
|
||||||
|
sequence is still incomplete.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
__slots__ = ("_tokenizer", "_ids", "_emitted")
|
__slots__ = ("_stream", "_tok")
|
||||||
|
|
||||||
def __init__(self, tokenizer: AutoTokenizer):
|
def __init__(self, tokenizer: AutoTokenizer):
|
||||||
self._tokenizer = tokenizer
|
self._tok = tokenizer._tokenizer
|
||||||
self._ids: List[int] = []
|
self._stream = DecodeStream(skip_special_tokens=True)
|
||||||
self._emitted: str = ""
|
|
||||||
|
|
||||||
def push(self, token_id: int) -> str:
|
def push(self, token_id: int) -> str:
|
||||||
"""Append a token ID and return newly completed text.
|
"""Append a token ID and return newly completed text.
|
||||||
|
|
||||||
Returns "" while a multi-byte character is still incomplete.
|
Returns "" while a multi-byte character is still incomplete.
|
||||||
"""
|
"""
|
||||||
self._ids.append(token_id)
|
chunk = self._stream.step(self._tok, token_id)
|
||||||
full = self._tokenizer.decode(self._ids, skip_special_tokens=True)
|
return chunk or ""
|
||||||
if full.endswith("\ufffd"):
|
|
||||||
return ""
|
|
||||||
if len(full) > len(self._emitted):
|
|
||||||
diff = full[len(self._emitted) :]
|
|
||||||
self._emitted = full
|
|
||||||
return diff
|
|
||||||
return ""
|
|
||||||
|
|
||||||
|
|
||||||
class TaskStatus(Enum):
|
class TaskStatus(Enum):
|
||||||
@@ -101,18 +96,11 @@ class Task:
|
|||||||
def flush_remaining(self, tokenizer: AutoTokenizer) -> str:
|
def flush_remaining(self, tokenizer: AutoTokenizer) -> str:
|
||||||
"""Emit any text still buffered in the decoder.
|
"""Emit any text still buffered in the decoder.
|
||||||
|
|
||||||
Called when generation terminates (max_tokens reached, stop
|
With the Rust-native DecodeStream, the stream is always in a
|
||||||
sequence, or external removal) to avoid dropping a final
|
correct state — any completed text was already emitted by the
|
||||||
incomplete-looking fragment that is actually complete when
|
last ``push``. A trailing incomplete multi-byte sequence has no
|
||||||
adjacent to the stop token.
|
valid text to emit, so this is a no-op.
|
||||||
"""
|
"""
|
||||||
if self._decoder is None or not self.output_ids:
|
|
||||||
return ""
|
|
||||||
full = tokenizer.decode(self.output_ids, skip_special_tokens=True)
|
|
||||||
if len(full) > len(self._decoder._emitted):
|
|
||||||
diff = full[len(self._decoder._emitted) :]
|
|
||||||
self._decoder._emitted = full
|
|
||||||
return diff
|
|
||||||
return ""
|
return ""
|
||||||
|
|
||||||
@property
|
@property
|
||||||
@@ -135,12 +123,10 @@ class TaskManager:
|
|||||||
tokenizer: AutoTokenizer,
|
tokenizer: AutoTokenizer,
|
||||||
max_batch_size: int = 16,
|
max_batch_size: int = 16,
|
||||||
max_seq_len: int = 8192,
|
max_seq_len: int = 8192,
|
||||||
max_prompt_len: int = 512,
|
|
||||||
):
|
):
|
||||||
self.tokenizer = tokenizer
|
self.tokenizer = tokenizer
|
||||||
self.max_batch_size = max_batch_size
|
self.max_batch_size = max_batch_size
|
||||||
self.max_seq_len = max_seq_len
|
self.max_seq_len = max_seq_len
|
||||||
self.max_prompt_len = max_prompt_len
|
|
||||||
|
|
||||||
self.waiting_queue: Deque[Task] = deque()
|
self.waiting_queue: Deque[Task] = deque()
|
||||||
self.active_tasks: List[Task] = []
|
self.active_tasks: List[Task] = []
|
||||||
@@ -165,10 +151,10 @@ class TaskManager:
|
|||||||
) -> str:
|
) -> str:
|
||||||
task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}"
|
task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}"
|
||||||
prompt_ids = self.tokenizer.encode(prompt)
|
prompt_ids = self.tokenizer.encode(prompt)
|
||||||
if len(prompt_ids) > self.max_prompt_len:
|
if len(prompt_ids) > self.max_seq_len:
|
||||||
prompt_ids = prompt_ids[-self.max_prompt_len :]
|
prompt_ids = prompt_ids[-self.max_seq_len :]
|
||||||
|
|
||||||
if len(prompt_ids) >= self.max_seq_len:
|
if len(prompt_ids) > self.max_seq_len:
|
||||||
if stream_callback:
|
if stream_callback:
|
||||||
stream_callback(STOP)
|
stream_callback(STOP)
|
||||||
return task_id
|
return task_id
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ from typing import Any, AsyncGenerator, Dict, Generator, List, Optional, Tuple,
|
|||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
|
||||||
from astrai.inference.core.cache import KVCache
|
from astrai.inference.core.cache import PagePool
|
||||||
from astrai.inference.core.scheduler import InferenceScheduler
|
from astrai.inference.core.scheduler import InferenceScheduler
|
||||||
from astrai.inference.core.task import STOP
|
from astrai.inference.core.task import STOP
|
||||||
from astrai.tokenize import AutoTokenizer
|
from astrai.tokenize import AutoTokenizer
|
||||||
@@ -111,9 +111,7 @@ class InferenceEngine:
|
|||||||
tokenizer: AutoTokenizer,
|
tokenizer: AutoTokenizer,
|
||||||
max_batch_size: int = 1,
|
max_batch_size: int = 1,
|
||||||
max_seq_len: Optional[int] = None,
|
max_seq_len: Optional[int] = None,
|
||||||
max_prompt_len: int = 2048,
|
cache: Optional[PagePool] = None,
|
||||||
page_size: int = 128,
|
|
||||||
cache: Optional[KVCache] = None,
|
|
||||||
):
|
):
|
||||||
self.model = model
|
self.model = model
|
||||||
self.tokenizer = tokenizer
|
self.tokenizer = tokenizer
|
||||||
@@ -122,7 +120,6 @@ class InferenceEngine:
|
|||||||
tokenizer=self.tokenizer,
|
tokenizer=self.tokenizer,
|
||||||
max_batch_size=max_batch_size,
|
max_batch_size=max_batch_size,
|
||||||
max_seq_len=max_seq_len,
|
max_seq_len=max_seq_len,
|
||||||
max_prompt_len=max_prompt_len,
|
|
||||||
cache=cache,
|
cache=cache,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -343,6 +343,10 @@ def sample(
|
|||||||
When **temperature** is exactly 0 (scalar or single-element tensor)
|
When **temperature** is exactly 0 (scalar or single-element tensor)
|
||||||
the function short-circuits to ``argmax`` for deterministic decode.
|
the function short-circuits to ``argmax`` for deterministic decode.
|
||||||
|
|
||||||
|
When **frequency_penalty** is 0 (the common decode case), the entire
|
||||||
|
frequency penalty computation — including the O(batch * vocab) count
|
||||||
|
tensor allocation — is skipped.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
logits: Raw logits ``[batch, vocab_size]``.
|
logits: Raw logits ``[batch, vocab_size]``.
|
||||||
frequency_penalty: Penalty per occurrence for repeated tokens
|
frequency_penalty: Penalty per occurrence for repeated tokens
|
||||||
@@ -359,14 +363,39 @@ def sample(
|
|||||||
``True`` — a ``(token_ids, chosen_logprobs)`` tuple where
|
``True`` — a ``(token_ids, chosen_logprobs)`` tuple where
|
||||||
``chosen_logprobs`` has shape ``[batch]``.
|
``chosen_logprobs`` has shape ``[batch]``.
|
||||||
"""
|
"""
|
||||||
return SamplingPipeline(
|
greedy = (
|
||||||
[
|
(
|
||||||
TemperatureStrategy(temperature),
|
isinstance(temperature, Tensor)
|
||||||
TopKStrategy(top_k),
|
and temperature.numel() == 1
|
||||||
TopPStrategy(top_p),
|
and temperature.item() == 0
|
||||||
FrequencyPenaltyStrategy(frequency_penalty),
|
)
|
||||||
]
|
if isinstance(temperature, Tensor)
|
||||||
).sample(
|
else temperature == 0
|
||||||
|
)
|
||||||
|
|
||||||
|
if greedy:
|
||||||
|
tokens = logits.argmax(dim=-1)
|
||||||
|
if not return_logprobs:
|
||||||
|
return tokens
|
||||||
|
log_probs = torch.log_softmax(logits.float(), dim=-1)
|
||||||
|
chosen = torch.gather(log_probs, -1, tokens.unsqueeze(-1)).squeeze(-1)
|
||||||
|
return tokens, chosen
|
||||||
|
|
||||||
|
has_freq = (
|
||||||
|
(isinstance(frequency_penalty, Tensor) and (frequency_penalty != 0).any())
|
||||||
|
if isinstance(frequency_penalty, Tensor)
|
||||||
|
else frequency_penalty != 0
|
||||||
|
)
|
||||||
|
|
||||||
|
strategies: List[BaseSamplingStrategy] = [
|
||||||
|
TemperatureStrategy(temperature),
|
||||||
|
TopKStrategy(top_k),
|
||||||
|
TopPStrategy(top_p),
|
||||||
|
]
|
||||||
|
if has_freq:
|
||||||
|
strategies.append(FrequencyPenaltyStrategy(frequency_penalty))
|
||||||
|
|
||||||
|
return SamplingPipeline(strategies).sample(
|
||||||
logits,
|
logits,
|
||||||
filter_value=filter_value,
|
filter_value=filter_value,
|
||||||
input_ids=input_ids,
|
input_ids=input_ids,
|
||||||
|
|||||||
@@ -40,11 +40,12 @@ def _disable_random_init(enable: bool = True):
|
|||||||
setattr(nn.init, n, fn)
|
setattr(nn.init, n, fn)
|
||||||
|
|
||||||
|
|
||||||
class AutoModel(BaseFactory["AutoModel"], nn.Module):
|
class ModelFactory(BaseFactory[nn.Module]):
|
||||||
"""
|
"""Pure factory for model dispatch, separated from nn.Module state."""
|
||||||
Autoregressive language model base class.
|
|
||||||
Provides model loading/saving, registration, and generation.
|
|
||||||
"""
|
class AutoModel(nn.Module):
|
||||||
|
"""Model base class with loading/saving and generation."""
|
||||||
|
|
||||||
def __init__(self, config: BaseModelConfig):
|
def __init__(self, config: BaseModelConfig):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
@@ -68,7 +69,7 @@ class AutoModel(BaseFactory["AutoModel"], nn.Module):
|
|||||||
config = ConfigFactory.load(raw)
|
config = ConfigFactory.load(raw)
|
||||||
model_type = config.model_type or "autoregressive_lm"
|
model_type = config.model_type or "autoregressive_lm"
|
||||||
|
|
||||||
actual_cls = AutoModel.get_component_class(model_type)
|
actual_cls = ModelFactory.get_component_class(model_type)
|
||||||
|
|
||||||
with _disable_random_init(enable=disable_random_init):
|
with _disable_random_init(enable=disable_random_init):
|
||||||
model = actual_cls(config)
|
model = actual_cls(config)
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
from astrai.model.components.attention import GQA, MLA, repeat_kv
|
from astrai.extension.rotary_backend import apply_rotary_emb
|
||||||
|
from astrai.model.components.attention import GQA, MLA
|
||||||
from astrai.model.components.decoder_block import DecoderBlock
|
from astrai.model.components.decoder_block import DecoderBlock
|
||||||
from astrai.model.components.embedding import Embedding
|
from astrai.model.components.embedding import Embedding
|
||||||
from astrai.model.components.linear import Linear
|
from astrai.model.components.linear import Linear
|
||||||
@@ -6,7 +7,6 @@ from astrai.model.components.mlp import MLP
|
|||||||
from astrai.model.components.norm import RMSNorm
|
from astrai.model.components.norm import RMSNorm
|
||||||
from astrai.model.components.rope import (
|
from astrai.model.components.rope import (
|
||||||
RotaryEmbedding,
|
RotaryEmbedding,
|
||||||
apply_rotary_emb,
|
|
||||||
get_rotary_emb,
|
get_rotary_emb,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -21,5 +21,4 @@ __all__ = [
|
|||||||
"RotaryEmbedding",
|
"RotaryEmbedding",
|
||||||
"apply_rotary_emb",
|
"apply_rotary_emb",
|
||||||
"get_rotary_emb",
|
"get_rotary_emb",
|
||||||
"repeat_kv",
|
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -5,22 +5,12 @@ import torch.nn as nn
|
|||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.extension import attention
|
||||||
|
from astrai.extension.rotary_backend import apply_rotary_emb
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
from astrai.inference.core.cache import CacheView
|
from astrai.inference.core.cache import KVCache
|
||||||
from astrai.model.components.linear import Linear
|
from astrai.model.components.linear import Linear
|
||||||
from astrai.model.components.norm import RMSNorm
|
from astrai.model.components.norm import RMSNorm
|
||||||
from astrai.model.components.rope import apply_rotary_emb
|
|
||||||
|
|
||||||
|
|
||||||
def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
|
|
||||||
bs, slen, n_heads, head_dim = x.shape
|
|
||||||
if n_rep == 1:
|
|
||||||
return x
|
|
||||||
return (
|
|
||||||
x[:, :, :, None, :]
|
|
||||||
.expand(bs, slen, n_heads, n_rep, head_dim)
|
|
||||||
.reshape(bs, slen, n_heads * n_rep, head_dim)
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class AttnFactory(BaseFactory[nn.Module]):
|
class AttnFactory(BaseFactory[nn.Module]):
|
||||||
@@ -75,7 +65,7 @@ class GQA(nn.Module):
|
|||||||
x: Tensor,
|
x: Tensor,
|
||||||
rotary_emb: Tensor,
|
rotary_emb: Tensor,
|
||||||
attn_mask: Tensor = None,
|
attn_mask: Tensor = None,
|
||||||
paged_cache: Optional[CacheView] = None,
|
kv_cache: Optional[KVCache] = None,
|
||||||
is_causal: bool = False,
|
is_causal: bool = False,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
q = self._split_heads(self.q_proj(x), self.n_heads)
|
q = self._split_heads(self.q_proj(x), self.n_heads)
|
||||||
@@ -86,19 +76,7 @@ class GQA(nn.Module):
|
|||||||
if self.use_qk_norm:
|
if self.use_qk_norm:
|
||||||
q, k = self.q_norm(q), self.k_norm(k)
|
q, k = self.q_norm(q), self.k_norm(k)
|
||||||
|
|
||||||
if paged_cache is not None:
|
sdqa_out = attention(q, k, v, kv_cache, self.layer_id, attn_mask, is_causal)
|
||||||
paged_cache.write(self.layer_id, k, v)
|
|
||||||
k, v = paged_cache.gather(self.layer_id)
|
|
||||||
|
|
||||||
k, v = repeat_kv(k, self.n_rep), repeat_kv(v, self.n_rep)
|
|
||||||
|
|
||||||
q, k, v = q.permute(0, 2, 1, 3), k.permute(0, 2, 1, 3), v.permute(0, 2, 1, 3)
|
|
||||||
sdqa_out = (
|
|
||||||
F.scaled_dot_product_attention(q, k, v, attn_mask, is_causal=is_causal)
|
|
||||||
.permute(0, 2, 1, 3)
|
|
||||||
.contiguous()
|
|
||||||
.flatten(2)
|
|
||||||
)
|
|
||||||
|
|
||||||
if self.use_gated_attention:
|
if self.use_gated_attention:
|
||||||
sdqa_out = sdqa_out * F.sigmoid(self.gate(x))
|
sdqa_out = sdqa_out * F.sigmoid(self.gate(x))
|
||||||
@@ -161,7 +139,7 @@ class MLA(nn.Module):
|
|||||||
x: Tensor,
|
x: Tensor,
|
||||||
rotary_emb: Tensor,
|
rotary_emb: Tensor,
|
||||||
attn_mask: Tensor = None,
|
attn_mask: Tensor = None,
|
||||||
paged_cache: Optional[CacheView] = None,
|
kv_cache: Optional[KVCache] = None,
|
||||||
is_causal: bool = False,
|
is_causal: bool = False,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
bsz, seq_len, _ = x.size()
|
bsz, seq_len, _ = x.size()
|
||||||
@@ -193,18 +171,7 @@ class MLA(nn.Module):
|
|||||||
q = self.q_norm(q)
|
q = self.q_norm(q)
|
||||||
k = self.k_norm(k)
|
k = self.k_norm(k)
|
||||||
|
|
||||||
if paged_cache is not None:
|
attn_out = attention(q, k, v, kv_cache, self.layer_id, attn_mask, is_causal)
|
||||||
paged_cache.write(self.layer_id, k, v)
|
|
||||||
k, v = paged_cache.gather(self.layer_id)
|
|
||||||
|
|
||||||
q = q.permute(0, 2, 1, 3)
|
|
||||||
k = k.permute(0, 2, 1, 3)
|
|
||||||
v = v.permute(0, 2, 1, 3)
|
|
||||||
|
|
||||||
attn_out = F.scaled_dot_product_attention(
|
|
||||||
q, k, v, attn_mask, is_causal=is_causal
|
|
||||||
)
|
|
||||||
attn_out = attn_out.permute(0, 2, 1, 3).contiguous().flatten(2)
|
|
||||||
|
|
||||||
if self.use_gated_attention:
|
if self.use_gated_attention:
|
||||||
attn_out = attn_out * F.sigmoid(self.gate(x))
|
attn_out = attn_out * F.sigmoid(self.gate(x))
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ from typing import Optional
|
|||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from astrai.inference.core.cache import CacheView
|
from astrai.inference.core.cache import KVCache
|
||||||
from astrai.model.components.attention import AttnFactory
|
from astrai.model.components.attention import AttnFactory
|
||||||
from astrai.model.components.mlp import FFNFactory
|
from astrai.model.components.mlp import FFNFactory
|
||||||
from astrai.model.components.norm import RMSNorm
|
from astrai.model.components.norm import RMSNorm
|
||||||
@@ -33,14 +33,14 @@ class DecoderBlock(nn.Module):
|
|||||||
x: Tensor,
|
x: Tensor,
|
||||||
rotary_emb: Tensor,
|
rotary_emb: Tensor,
|
||||||
attention_mask: Optional[Tensor] = None,
|
attention_mask: Optional[Tensor] = None,
|
||||||
paged_cache: Optional[CacheView] = None,
|
kv_cache: Optional[KVCache] = None,
|
||||||
is_causal: bool = False,
|
is_causal: bool = False,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
attn_output = self.attention(
|
attn_output = self.attention(
|
||||||
self.input_norm(x),
|
self.input_norm(x),
|
||||||
rotary_emb,
|
rotary_emb,
|
||||||
attention_mask,
|
attention_mask,
|
||||||
paged_cache,
|
kv_cache,
|
||||||
is_causal,
|
is_causal,
|
||||||
)
|
)
|
||||||
x = attn_output + x
|
x = attn_output + x
|
||||||
|
|||||||
@@ -1,11 +1,12 @@
|
|||||||
import logging
|
import logging
|
||||||
from dataclasses import asdict, dataclass
|
from dataclasses import asdict
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Optional, Set
|
from typing import Optional, Set
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
|
from pydantic.dataclasses import dataclass
|
||||||
|
|
||||||
from astrai.model.components.linear import Linear
|
from astrai.model.components.linear import Linear
|
||||||
from astrai.serialization import (
|
from astrai.serialization import (
|
||||||
|
|||||||
@@ -11,28 +11,23 @@ def get_rotary_emb(
|
|||||||
base: float = 10000,
|
base: float = 10000,
|
||||||
device: Optional[torch.device] = None,
|
device: Optional[torch.device] = None,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
|
"""Precompute cos/sin tables for rotary embedding.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[max_len, dim/2, 2] (f32) — [cos, sin] pairs.
|
||||||
|
"""
|
||||||
theta = base ** (-torch.arange(0, dim, 2, dtype=torch.float64, device=device) / dim)
|
theta = base ** (-torch.arange(0, dim, 2, dtype=torch.float64, device=device) / dim)
|
||||||
t = torch.arange(0, max_len, dtype=torch.float64, device=device)
|
t = torch.arange(0, max_len, dtype=torch.float64, device=device)
|
||||||
freqs = torch.outer(t, theta).float()
|
freqs = torch.outer(t, theta).float()
|
||||||
cos = torch.cos(freqs)
|
cos = torch.cos(freqs)
|
||||||
sin = torch.sin(freqs)
|
sin = torch.sin(freqs)
|
||||||
return torch.complex(cos, sin)
|
return torch.stack([cos, sin], dim=-1)
|
||||||
|
|
||||||
|
|
||||||
def ntk_base(base: float, dim: int, factor: float) -> float:
|
def ntk_base(base: float, dim: int, factor: float) -> float:
|
||||||
return base * (factor ** (dim / (dim - 2)))
|
return base * (factor ** (dim / (dim - 2)))
|
||||||
|
|
||||||
|
|
||||||
def apply_rotary_emb(x: torch.Tensor, freqs_cis: Tensor) -> Tensor:
|
|
||||||
dtype = x.dtype
|
|
||||||
x_ = x.float().reshape(*x.shape[:-1], -1, 2)
|
|
||||||
x_complex = torch.view_as_complex(x_)
|
|
||||||
freqs_cis = freqs_cis.unsqueeze(2)
|
|
||||||
x_rotated = x_complex * freqs_cis
|
|
||||||
x_out = torch.view_as_real(x_rotated).flatten(-2)
|
|
||||||
return x_out.to(dtype)
|
|
||||||
|
|
||||||
|
|
||||||
class RotaryEmbedding(nn.Module):
|
class RotaryEmbedding(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -56,16 +51,23 @@ class RotaryEmbedding(nn.Module):
|
|||||||
self._set_rotary_buffer(self.max_len)
|
self._set_rotary_buffer(self.max_len)
|
||||||
|
|
||||||
def _set_rotary_buffer(self, max_len: int):
|
def _set_rotary_buffer(self, max_len: int):
|
||||||
rotary_emb = get_rotary_emb(self.dim, max_len, self.base)
|
freqs_cis = get_rotary_emb(self.dim, max_len, self.base)
|
||||||
freqs_cis = torch.view_as_real(rotary_emb)
|
|
||||||
self.register_buffer("freqs_cis", freqs_cis, persistent=False)
|
self.register_buffer("freqs_cis", freqs_cis, persistent=False)
|
||||||
|
|
||||||
def forward(self, x: Tensor, position_ids: Optional[Tensor] = None) -> Tensor:
|
def forward(self, x: Tensor, position_ids: Optional[Tensor] = None) -> Tensor:
|
||||||
|
"""Lookup cos/sin for the given positions.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
x: [batch, seq_len, ...] — only batch and seq_len are used.
|
||||||
|
position_ids: [batch, seq_len] optional position indices.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[batch, seq_len, dim/2, 2] (f32) — [cos, sin] pairs.
|
||||||
|
"""
|
||||||
if position_ids is None:
|
if position_ids is None:
|
||||||
position_ids = (
|
position_ids = (
|
||||||
torch.arange(x.size(1), device=x.device)
|
torch.arange(x.size(1), device=x.device)
|
||||||
.unsqueeze(0)
|
.unsqueeze(0)
|
||||||
.expand(x.size(0), -1)
|
.expand(x.size(0), -1)
|
||||||
)
|
)
|
||||||
position_freq_cis = self.freqs_cis[position_ids].float()
|
return self.freqs_cis[position_ids].float()
|
||||||
return torch.view_as_complex(position_freq_cis)
|
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ import torch.nn as nn
|
|||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from astrai.config.model_config import EncoderConfig
|
from astrai.config.model_config import EncoderConfig
|
||||||
from astrai.model.automodel import AutoModel
|
from astrai.model.automodel import AutoModel, ModelFactory
|
||||||
from astrai.model.components.decoder_block import DecoderBlock
|
from astrai.model.components.decoder_block import DecoderBlock
|
||||||
from astrai.model.components.embedding import Embedding
|
from astrai.model.components.embedding import Embedding
|
||||||
from astrai.model.components.norm import RMSNorm
|
from astrai.model.components.norm import RMSNorm
|
||||||
@@ -13,7 +13,7 @@ from astrai.model.components.rope import RotaryEmbedding
|
|||||||
from astrai.model.transformer import process_attention_mask
|
from astrai.model.transformer import process_attention_mask
|
||||||
|
|
||||||
|
|
||||||
@AutoModel.register("embedding")
|
@ModelFactory.register("embedding")
|
||||||
class EmbeddingEncoder(AutoModel):
|
class EmbeddingEncoder(AutoModel):
|
||||||
def __init__(self, config: EncoderConfig):
|
def __init__(self, config: EncoderConfig):
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
|
|||||||
@@ -5,8 +5,8 @@ import torch.nn as nn
|
|||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
from astrai.inference.core.cache import CacheView
|
from astrai.inference.core.cache import KVCache
|
||||||
from astrai.model.automodel import AutoModel
|
from astrai.model.automodel import AutoModel, ModelFactory
|
||||||
from astrai.model.components.decoder_block import DecoderBlock
|
from astrai.model.components.decoder_block import DecoderBlock
|
||||||
from astrai.model.components.embedding import Embedding
|
from astrai.model.components.embedding import Embedding
|
||||||
from astrai.model.components.linear import Linear
|
from astrai.model.components.linear import Linear
|
||||||
@@ -26,7 +26,7 @@ def process_attention_mask(
|
|||||||
return input_mask
|
return input_mask
|
||||||
|
|
||||||
|
|
||||||
@AutoModel.register("autoregressive_lm")
|
@ModelFactory.register("autoregressive_lm")
|
||||||
class AutoRegressiveLM(AutoModel):
|
class AutoRegressiveLM(AutoModel):
|
||||||
"""Autoregressive language model with paged KV cache."""
|
"""Autoregressive language model with paged KV cache."""
|
||||||
|
|
||||||
@@ -103,7 +103,7 @@ class AutoRegressiveLM(AutoModel):
|
|||||||
self,
|
self,
|
||||||
input_ids: Tensor,
|
input_ids: Tensor,
|
||||||
input_mask: Optional[Tensor] = None,
|
input_mask: Optional[Tensor] = None,
|
||||||
paged_cache: Optional[CacheView] = None,
|
kv_cache: Optional[KVCache] = None,
|
||||||
position_ids: Optional[Tensor] = None,
|
position_ids: Optional[Tensor] = None,
|
||||||
) -> Dict[str, Tensor]:
|
) -> Dict[str, Tensor]:
|
||||||
assert input_ids.ndim == 2
|
assert input_ids.ndim == 2
|
||||||
@@ -114,7 +114,7 @@ class AutoRegressiveLM(AutoModel):
|
|||||||
use_sdpa_causal_mask = attn_mask is None
|
use_sdpa_causal_mask = attn_mask is None
|
||||||
|
|
||||||
for layer in self.layers:
|
for layer in self.layers:
|
||||||
x = layer(x, rotary_emb, attn_mask, paged_cache, use_sdpa_causal_mask)
|
x = layer(x, rotary_emb, attn_mask, kv_cache, use_sdpa_causal_mask)
|
||||||
|
|
||||||
hidden_states = self.norm(x)
|
hidden_states = self.norm(x)
|
||||||
logits = self.lm_head(hidden_states)
|
logits = self.lm_head(hidden_states)
|
||||||
|
|||||||
@@ -0,0 +1,38 @@
|
|||||||
|
"""Optimizer implementations and factory registration."""
|
||||||
|
|
||||||
|
from astrai.optim.composite import (
|
||||||
|
OptimizerFactory,
|
||||||
|
composite_state_dict,
|
||||||
|
composite_step,
|
||||||
|
composite_zero_grad,
|
||||||
|
refresh_param_groups,
|
||||||
|
)
|
||||||
|
from astrai.optim.mano_adamw import Mano, ManoAdamW
|
||||||
|
from astrai.optim.muon_adamw import MuonAdamW
|
||||||
|
from astrai.optim.nora_nadamw import (
|
||||||
|
NAdamW,
|
||||||
|
Nora,
|
||||||
|
NoraNAdamW,
|
||||||
|
OptimizerParameterGroups,
|
||||||
|
nora_direction,
|
||||||
|
nora_lr_scale,
|
||||||
|
partition_optimizer_parameters,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"Mano",
|
||||||
|
"ManoAdamW",
|
||||||
|
"MuonAdamW",
|
||||||
|
"NAdamW",
|
||||||
|
"Nora",
|
||||||
|
"NoraNAdamW",
|
||||||
|
"OptimizerFactory",
|
||||||
|
"OptimizerParameterGroups",
|
||||||
|
"composite_state_dict",
|
||||||
|
"composite_step",
|
||||||
|
"composite_zero_grad",
|
||||||
|
"nora_direction",
|
||||||
|
"nora_lr_scale",
|
||||||
|
"partition_optimizer_parameters",
|
||||||
|
"refresh_param_groups",
|
||||||
|
]
|
||||||
@@ -0,0 +1,71 @@
|
|||||||
|
"""Shared infrastructure for the optim package.
|
||||||
|
|
||||||
|
This module hosts two things:
|
||||||
|
|
||||||
|
* ``OptimizerFactory`` — the registry for built-in optimizers. Defining it
|
||||||
|
here (rather than in ``__init__.py``) lets each optimizer module import it
|
||||||
|
and register itself with a decorator, avoiding circular imports.
|
||||||
|
* Composite-optimizer helpers — ``step``/``zero_grad``/``state_dict``/
|
||||||
|
``param_groups`` delegation shared by every optimizer that routes different
|
||||||
|
parameter groups through distinct sub-optimizers.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch.optim import Optimizer
|
||||||
|
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
|
||||||
|
|
||||||
|
class OptimizerFactory(BaseFactory[Optimizer]):
|
||||||
|
"""Factory for built-in training optimizers."""
|
||||||
|
|
||||||
|
|
||||||
|
def composite_step(
|
||||||
|
sub_optimizers: list[Optimizer],
|
||||||
|
closure=None,
|
||||||
|
) -> torch.Tensor | None:
|
||||||
|
"""Run ``step`` on every sub-optimizer, invoking the closure once.
|
||||||
|
|
||||||
|
The closure (if given) is executed inside ``torch.enable_grad`` exactly
|
||||||
|
once before any sub-optimizer steps, matching the contract of a single
|
||||||
|
``Optimizer.step``. Sub-optimizers receive ``None`` so they do not
|
||||||
|
re-execute it.
|
||||||
|
"""
|
||||||
|
loss = None
|
||||||
|
if closure is not None:
|
||||||
|
with torch.enable_grad():
|
||||||
|
loss = closure()
|
||||||
|
for sub in sub_optimizers:
|
||||||
|
sub.step()
|
||||||
|
return loss
|
||||||
|
|
||||||
|
|
||||||
|
def composite_zero_grad(
|
||||||
|
sub_optimizers: list[Optimizer],
|
||||||
|
set_to_none: bool = True,
|
||||||
|
) -> None:
|
||||||
|
for sub in sub_optimizers:
|
||||||
|
sub.zero_grad(set_to_none=set_to_none)
|
||||||
|
|
||||||
|
|
||||||
|
def composite_state_dict(
|
||||||
|
named_sub_optimizers: dict[str, Optimizer | None],
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Serialize sub-optimizers, preserving ``None`` slots."""
|
||||||
|
return {
|
||||||
|
name: sub.state_dict() if sub is not None else None
|
||||||
|
for name, sub in named_sub_optimizers.items()
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def refresh_param_groups(
|
||||||
|
sub_optimizers: list[Optimizer],
|
||||||
|
) -> list[dict]:
|
||||||
|
"""Concatenate param_groups from every non-None sub-optimizer."""
|
||||||
|
groups: list[dict] = []
|
||||||
|
for sub in sub_optimizers:
|
||||||
|
if sub is not None:
|
||||||
|
groups.extend(sub.param_groups)
|
||||||
|
return groups
|
||||||
@@ -0,0 +1,214 @@
|
|||||||
|
"""Mano manifold optimizer combined with AdamW.
|
||||||
|
|
||||||
|
Mano projects the momentum onto the tangent space of the Oblique manifold
|
||||||
|
(axis-wise tangent projection) and normalizes it, replacing the expensive
|
||||||
|
Newton-Schulz iteration in Muon with a cheaper manifold normalization.
|
||||||
|
|
||||||
|
Reference: https://arxiv.org/abs/2601.23000
|
||||||
|
"""
|
||||||
|
|
||||||
|
import math
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import nn, optim
|
||||||
|
from torch.optim import Optimizer
|
||||||
|
|
||||||
|
from astrai.optim.composite import (
|
||||||
|
OptimizerFactory,
|
||||||
|
composite_state_dict,
|
||||||
|
composite_step,
|
||||||
|
composite_zero_grad,
|
||||||
|
refresh_param_groups,
|
||||||
|
)
|
||||||
|
from astrai.optim.nora_nadamw import partition_optimizer_parameters
|
||||||
|
|
||||||
|
|
||||||
|
class Mano(Optimizer):
|
||||||
|
"""Manifold Normalized Optimizer for two-dimensional matrices.
|
||||||
|
|
||||||
|
Each step alternates the projection axis (dim 0 / dim 1) to restrike the
|
||||||
|
manifold along both rows and columns. The tangent momentum is computed
|
||||||
|
without normalizing the parameter itself (v2 simplification) and the
|
||||||
|
epsilon is added (not clamped) to the norm denominator.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
params,
|
||||||
|
lr: float = 1e-3,
|
||||||
|
weight_decay: float = 0.1,
|
||||||
|
momentum: float = 0.95,
|
||||||
|
nesterov: bool = True,
|
||||||
|
eps: float = 1e-8,
|
||||||
|
):
|
||||||
|
if lr < 0:
|
||||||
|
raise ValueError(f"Invalid learning rate: {lr}")
|
||||||
|
if weight_decay < 0:
|
||||||
|
raise ValueError(f"Invalid weight decay: {weight_decay}")
|
||||||
|
if not 0 <= momentum <= 1:
|
||||||
|
raise ValueError(f"Invalid momentum: {momentum}")
|
||||||
|
if eps <= 0:
|
||||||
|
raise ValueError(f"Invalid epsilon: {eps}")
|
||||||
|
|
||||||
|
defaults = {
|
||||||
|
"lr": lr,
|
||||||
|
"weight_decay": weight_decay,
|
||||||
|
"momentum": momentum,
|
||||||
|
"nesterov": nesterov,
|
||||||
|
"eps": eps,
|
||||||
|
"steps": 0,
|
||||||
|
}
|
||||||
|
super().__init__(params, defaults)
|
||||||
|
for group in self.param_groups:
|
||||||
|
for param in group["params"]:
|
||||||
|
if param.ndim != 2:
|
||||||
|
raise ValueError(
|
||||||
|
f"Mano only supports 2D matrices, got shape {tuple(param.shape)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def step(self, closure=None):
|
||||||
|
loss = None
|
||||||
|
if closure is not None:
|
||||||
|
with torch.enable_grad():
|
||||||
|
loss = closure()
|
||||||
|
|
||||||
|
for group in self.param_groups:
|
||||||
|
lr = group["lr"]
|
||||||
|
weight_decay = group["weight_decay"]
|
||||||
|
momentum = group["momentum"]
|
||||||
|
nesterov = group["nesterov"]
|
||||||
|
eps = group["eps"]
|
||||||
|
dim = int(group["steps"] % 2)
|
||||||
|
|
||||||
|
for param in group["params"]:
|
||||||
|
if param.grad is None:
|
||||||
|
continue
|
||||||
|
if param.grad.is_sparse:
|
||||||
|
raise RuntimeError("Mano does not support sparse gradients")
|
||||||
|
|
||||||
|
grad = param.grad
|
||||||
|
state = self.state[param]
|
||||||
|
momentum_buffer = state.get("momentum_buffer")
|
||||||
|
if momentum_buffer is None:
|
||||||
|
momentum_buffer = torch.zeros_like(grad)
|
||||||
|
momentum_buffer.mul_(momentum).add_(grad)
|
||||||
|
update = (
|
||||||
|
grad.add(momentum_buffer, alpha=momentum)
|
||||||
|
if nesterov
|
||||||
|
else momentum_buffer
|
||||||
|
)
|
||||||
|
|
||||||
|
tangent = update - (
|
||||||
|
torch.sum(update * param.data, dim=dim, keepdim=True) * param.data
|
||||||
|
)
|
||||||
|
direction = tangent / (
|
||||||
|
torch.norm(tangent, p=2, dim=dim, keepdim=True) + eps
|
||||||
|
)
|
||||||
|
|
||||||
|
if weight_decay != 0:
|
||||||
|
param.mul_(1 - lr * weight_decay)
|
||||||
|
adjusted_lr = lr * 0.2 * math.sqrt(direction.shape[dim])
|
||||||
|
param.add_(direction, alpha=-adjusted_lr)
|
||||||
|
state["momentum_buffer"] = momentum_buffer
|
||||||
|
|
||||||
|
group["steps"] += 1
|
||||||
|
|
||||||
|
return loss
|
||||||
|
|
||||||
|
|
||||||
|
@OptimizerFactory.register("mano_adamw")
|
||||||
|
class ManoAdamW(Optimizer):
|
||||||
|
"""Mano for internal linear weights and AdamW for remaining parameters."""
|
||||||
|
|
||||||
|
optimizer_name = "mano_adamw"
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model: nn.Module,
|
||||||
|
lr: float = 3e-4,
|
||||||
|
weight_decay: float = 0.1,
|
||||||
|
momentum: float = 0.95,
|
||||||
|
nesterov: bool = True,
|
||||||
|
):
|
||||||
|
groups = partition_optimizer_parameters(model)
|
||||||
|
all_params = [
|
||||||
|
*groups.nora,
|
||||||
|
*groups.nadamw_decay,
|
||||||
|
*groups.nadamw_no_decay,
|
||||||
|
]
|
||||||
|
if not all_params:
|
||||||
|
raise ValueError(
|
||||||
|
"Cannot build an optimizer for a model with no trainable parameters"
|
||||||
|
)
|
||||||
|
super().__init__(all_params, {})
|
||||||
|
|
||||||
|
self.mano = (
|
||||||
|
Mano(
|
||||||
|
groups.nora,
|
||||||
|
lr=lr,
|
||||||
|
weight_decay=weight_decay,
|
||||||
|
momentum=momentum,
|
||||||
|
nesterov=nesterov,
|
||||||
|
)
|
||||||
|
if groups.nora
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
|
||||||
|
adamw_groups = []
|
||||||
|
if groups.nadamw_decay:
|
||||||
|
adamw_groups.append(
|
||||||
|
{"params": groups.nadamw_decay, "weight_decay": weight_decay}
|
||||||
|
)
|
||||||
|
if groups.nadamw_no_decay:
|
||||||
|
adamw_groups.append({"params": groups.nadamw_no_decay, "weight_decay": 0.0})
|
||||||
|
self.adamw = (
|
||||||
|
optim.AdamW(
|
||||||
|
adamw_groups,
|
||||||
|
lr=lr,
|
||||||
|
betas=(0.9, 0.95),
|
||||||
|
fused=True,
|
||||||
|
)
|
||||||
|
if adamw_groups
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
self.param_groups = refresh_param_groups([self.mano, self.adamw])
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def step(self, closure=None):
|
||||||
|
return composite_step(
|
||||||
|
[opt for opt in (self.mano, self.adamw) if opt is not None],
|
||||||
|
closure,
|
||||||
|
)
|
||||||
|
|
||||||
|
def zero_grad(self, set_to_none: bool = True):
|
||||||
|
composite_zero_grad(
|
||||||
|
[opt for opt in (self.mano, self.adamw) if opt is not None],
|
||||||
|
set_to_none,
|
||||||
|
)
|
||||||
|
|
||||||
|
def state_dict(self) -> dict:
|
||||||
|
return composite_state_dict({"mano": self.mano, "adamw": self.adamw})
|
||||||
|
|
||||||
|
def load_state_dict(self, state_dict: dict):
|
||||||
|
if "muon" in state_dict or "nora" in state_dict:
|
||||||
|
raise ValueError(
|
||||||
|
"Checkpoint uses a different optimizer; select the matching "
|
||||||
|
"--optimizer to resume it"
|
||||||
|
)
|
||||||
|
if "mano" not in state_dict or "adamw" not in state_dict:
|
||||||
|
raise ValueError(
|
||||||
|
"Checkpoint optimizer state is not compatible with mano_adamw"
|
||||||
|
)
|
||||||
|
|
||||||
|
saved_mano = state_dict["mano"]
|
||||||
|
saved_adamw = state_dict["adamw"]
|
||||||
|
if (self.mano is None) != (saved_mano is None):
|
||||||
|
raise ValueError("Checkpoint Mano parameter groups do not match the model")
|
||||||
|
if (self.adamw is None) != (saved_adamw is None):
|
||||||
|
raise ValueError("Checkpoint AdamW parameter groups do not match the model")
|
||||||
|
if self.mano is not None:
|
||||||
|
self.mano.load_state_dict(saved_mano)
|
||||||
|
if self.adamw is not None:
|
||||||
|
self.adamw.load_state_dict(saved_adamw)
|
||||||
|
self.param_groups = refresh_param_groups([self.mano, self.adamw])
|
||||||
@@ -0,0 +1,95 @@
|
|||||||
|
"""Legacy Muon + AdamW combined optimizer."""
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import Tensor, nn, optim
|
||||||
|
|
||||||
|
from astrai.optim.composite import (
|
||||||
|
OptimizerFactory,
|
||||||
|
composite_state_dict,
|
||||||
|
composite_step,
|
||||||
|
composite_zero_grad,
|
||||||
|
refresh_param_groups,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@OptimizerFactory.register("muon_adamw")
|
||||||
|
class MuonAdamW(optim.Optimizer):
|
||||||
|
"""Combined Muon (matrix) + AdamW (non-matrix) optimizer."""
|
||||||
|
|
||||||
|
optimizer_name = "muon_adamw"
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model: nn.Module,
|
||||||
|
lr: float = 3e-4,
|
||||||
|
weight_decay: float = 0.1,
|
||||||
|
momentum: float = 0.95,
|
||||||
|
nesterov: bool = True,
|
||||||
|
ns_steps: int = 5,
|
||||||
|
adjust_lr_fn: str = "match_rms_adamw",
|
||||||
|
):
|
||||||
|
defaults = {
|
||||||
|
"lr": lr,
|
||||||
|
"weight_decay": weight_decay,
|
||||||
|
"momentum": momentum,
|
||||||
|
"nesterov": nesterov,
|
||||||
|
"ns_steps": ns_steps,
|
||||||
|
"adjust_lr_fn": adjust_lr_fn,
|
||||||
|
}
|
||||||
|
params = [param for param in model.parameters() if param.requires_grad]
|
||||||
|
super().__init__(params, defaults)
|
||||||
|
|
||||||
|
matrix_params: list[Tensor] = []
|
||||||
|
other_params: list[Tensor] = []
|
||||||
|
for name, param in model.named_parameters():
|
||||||
|
if not param.requires_grad:
|
||||||
|
continue
|
||||||
|
if (
|
||||||
|
param.dim() >= 2
|
||||||
|
and "norm" not in name
|
||||||
|
and "bias" not in name
|
||||||
|
and "embed" not in name
|
||||||
|
and "lm_head" not in name
|
||||||
|
):
|
||||||
|
matrix_params.append(param)
|
||||||
|
else:
|
||||||
|
other_params.append(param)
|
||||||
|
|
||||||
|
self.muon = optim.Muon(
|
||||||
|
matrix_params,
|
||||||
|
lr=lr,
|
||||||
|
weight_decay=weight_decay,
|
||||||
|
momentum=momentum,
|
||||||
|
nesterov=nesterov,
|
||||||
|
ns_steps=ns_steps,
|
||||||
|
adjust_lr_fn=adjust_lr_fn,
|
||||||
|
)
|
||||||
|
self.adamw = optim.AdamW(
|
||||||
|
[{"params": other_params, "weight_decay": 0.0}],
|
||||||
|
lr=lr,
|
||||||
|
betas=(0.9, 0.95),
|
||||||
|
fused=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.param_groups = refresh_param_groups([self.muon, self.adamw])
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def step(self, closure=None):
|
||||||
|
return composite_step([self.muon, self.adamw], closure)
|
||||||
|
|
||||||
|
def zero_grad(self, set_to_none: bool = True):
|
||||||
|
composite_zero_grad([self.muon, self.adamw], set_to_none)
|
||||||
|
|
||||||
|
def state_dict(self) -> dict[str, Any]:
|
||||||
|
return composite_state_dict({"muon": self.muon, "adamw": self.adamw})
|
||||||
|
|
||||||
|
def load_state_dict(self, state_dict: dict[str, Any]):
|
||||||
|
if "muon" not in state_dict or "adamw" not in state_dict:
|
||||||
|
raise ValueError(
|
||||||
|
"Checkpoint optimizer state is not compatible with muon_adamw"
|
||||||
|
)
|
||||||
|
self.muon.load_state_dict(state_dict["muon"])
|
||||||
|
self.adamw.load_state_dict(state_dict["adamw"])
|
||||||
|
self.param_groups = refresh_param_groups([self.muon, self.adamw])
|
||||||
@@ -0,0 +1,372 @@
|
|||||||
|
"""Nora matrix optimizer combined with Nesterov AdamW."""
|
||||||
|
|
||||||
|
import math
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import Tensor, nn
|
||||||
|
from torch.distributed.tensor import DTensor, Shard
|
||||||
|
from torch.optim import Optimizer
|
||||||
|
|
||||||
|
from astrai.model.components.embedding import Embedding
|
||||||
|
from astrai.model.components.linear import Linear
|
||||||
|
from astrai.model.components.lora import LoRALinear
|
||||||
|
from astrai.model.components.norm import RMSNorm
|
||||||
|
from astrai.optim.composite import (
|
||||||
|
OptimizerFactory,
|
||||||
|
composite_state_dict,
|
||||||
|
composite_step,
|
||||||
|
composite_zero_grad,
|
||||||
|
refresh_param_groups,
|
||||||
|
)
|
||||||
|
|
||||||
|
NORA_EPS = 1e-10
|
||||||
|
|
||||||
|
|
||||||
|
def _row_normalize(tensor: Tensor, eps: float) -> Tensor:
|
||||||
|
return tensor / tensor.norm(dim=-1, keepdim=True).clamp(min=eps)
|
||||||
|
|
||||||
|
|
||||||
|
def nora_direction(update: Tensor, param: Tensor, eps: float = NORA_EPS) -> Tensor:
|
||||||
|
"""Project an update onto each parameter row's tangent space and normalize."""
|
||||||
|
theta_hat = _row_normalize(param.to(torch.float32), eps)
|
||||||
|
update_fp32 = update.to(torch.float32)
|
||||||
|
radial = (update_fp32 * theta_hat).sum(dim=-1, keepdim=True) * theta_hat
|
||||||
|
direction = _row_normalize(update_fp32 - radial, eps)
|
||||||
|
return direction.to(update.dtype)
|
||||||
|
|
||||||
|
|
||||||
|
def nora_lr_scale(lr: float, shape: torch.Size) -> float:
|
||||||
|
"""Scale Nora's LR for tall ``[d_out, d_in]`` linear weights."""
|
||||||
|
return lr * math.sqrt(max(1.0, shape[-2] / shape[-1]))
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_complete_rows(param: Tensor) -> None:
|
||||||
|
if not isinstance(param, DTensor):
|
||||||
|
return
|
||||||
|
last_dim = param.ndim - 1
|
||||||
|
for placement in param.placements:
|
||||||
|
if isinstance(placement, Shard) and placement.dim % param.ndim == last_dim:
|
||||||
|
raise ValueError(
|
||||||
|
"Nora requires complete parameter rows, but this DTensor is sharded "
|
||||||
|
"along its last dimension"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class Nora(Optimizer):
|
||||||
|
"""Normalized Orthogonal Row Alignment for two-dimensional matrices."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
params,
|
||||||
|
lr: float = 5e-3,
|
||||||
|
weight_decay: float = 0.0,
|
||||||
|
momentum: float = 0.95,
|
||||||
|
beta: float = 0.95,
|
||||||
|
nesterov: bool = True,
|
||||||
|
eps: float = NORA_EPS,
|
||||||
|
):
|
||||||
|
if lr < 0:
|
||||||
|
raise ValueError(f"Invalid learning rate: {lr}")
|
||||||
|
if weight_decay < 0:
|
||||||
|
raise ValueError(f"Invalid weight decay: {weight_decay}")
|
||||||
|
if not 0 <= momentum <= 1:
|
||||||
|
raise ValueError(f"Invalid momentum: {momentum}")
|
||||||
|
if not 0 <= beta < 1:
|
||||||
|
raise ValueError(f"Invalid beta: {beta}")
|
||||||
|
if eps <= 0:
|
||||||
|
raise ValueError(f"Invalid epsilon: {eps}")
|
||||||
|
|
||||||
|
defaults = {
|
||||||
|
"lr": lr,
|
||||||
|
"weight_decay": weight_decay,
|
||||||
|
"momentum": momentum,
|
||||||
|
"beta": beta,
|
||||||
|
"nesterov": nesterov,
|
||||||
|
"eps": eps,
|
||||||
|
}
|
||||||
|
super().__init__(params, defaults)
|
||||||
|
for group in self.param_groups:
|
||||||
|
for param in group["params"]:
|
||||||
|
if param.ndim != 2:
|
||||||
|
raise ValueError(
|
||||||
|
f"Nora only supports 2D matrices, got shape {tuple(param.shape)}"
|
||||||
|
)
|
||||||
|
_validate_complete_rows(param)
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def step(self, closure=None):
|
||||||
|
loss = None
|
||||||
|
if closure is not None:
|
||||||
|
with torch.enable_grad():
|
||||||
|
loss = closure()
|
||||||
|
|
||||||
|
for group in self.param_groups:
|
||||||
|
lr = group["lr"]
|
||||||
|
weight_decay = group["weight_decay"]
|
||||||
|
momentum = group["momentum"]
|
||||||
|
beta = group["beta"]
|
||||||
|
nesterov = group["nesterov"]
|
||||||
|
eps = group["eps"]
|
||||||
|
for param in group["params"]:
|
||||||
|
if param.grad is None:
|
||||||
|
continue
|
||||||
|
if param.grad.is_sparse:
|
||||||
|
raise RuntimeError("Nora does not support sparse gradients")
|
||||||
|
|
||||||
|
grad = param.grad
|
||||||
|
state = self.state[param]
|
||||||
|
momentum_buffer = state.get("momentum_buffer")
|
||||||
|
if momentum_buffer is None:
|
||||||
|
momentum_buffer = torch.zeros_like(grad)
|
||||||
|
momentum_buffer.lerp_(grad, 1 - beta)
|
||||||
|
update = (
|
||||||
|
grad.lerp(momentum_buffer, momentum)
|
||||||
|
if nesterov
|
||||||
|
else momentum_buffer
|
||||||
|
)
|
||||||
|
direction = nora_direction(update, param, eps)
|
||||||
|
|
||||||
|
if weight_decay != 0:
|
||||||
|
param.mul_(1 - lr * weight_decay)
|
||||||
|
param.add_(direction, alpha=-nora_lr_scale(lr, param.shape))
|
||||||
|
state["momentum_buffer"] = momentum_buffer
|
||||||
|
|
||||||
|
return loss
|
||||||
|
|
||||||
|
|
||||||
|
class NAdamW(Optimizer):
|
||||||
|
"""AdamW using the reference Nesterov first-moment update."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
params,
|
||||||
|
lr: float = 3e-4,
|
||||||
|
betas: tuple[float, float] = (0.9, 0.999),
|
||||||
|
eps: float = 1e-8,
|
||||||
|
weight_decay: float = 0.1,
|
||||||
|
):
|
||||||
|
beta1, beta2 = betas
|
||||||
|
if lr < 0:
|
||||||
|
raise ValueError(f"Invalid learning rate: {lr}")
|
||||||
|
if not 0 <= beta1 < 1 or not 0 <= beta2 < 1:
|
||||||
|
raise ValueError(f"Invalid betas: {betas}")
|
||||||
|
if eps <= 0:
|
||||||
|
raise ValueError(f"Invalid epsilon: {eps}")
|
||||||
|
if weight_decay < 0:
|
||||||
|
raise ValueError(f"Invalid weight decay: {weight_decay}")
|
||||||
|
defaults = {
|
||||||
|
"lr": lr,
|
||||||
|
"betas": betas,
|
||||||
|
"eps": eps,
|
||||||
|
"weight_decay": weight_decay,
|
||||||
|
}
|
||||||
|
super().__init__(params, defaults)
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def step(self, closure=None):
|
||||||
|
loss = None
|
||||||
|
if closure is not None:
|
||||||
|
with torch.enable_grad():
|
||||||
|
loss = closure()
|
||||||
|
|
||||||
|
for group in self.param_groups:
|
||||||
|
beta1, beta2 = group["betas"]
|
||||||
|
eps = group["eps"]
|
||||||
|
lr = group["lr"]
|
||||||
|
weight_decay = group["weight_decay"]
|
||||||
|
for param in group["params"]:
|
||||||
|
if param.grad is None:
|
||||||
|
continue
|
||||||
|
if param.grad.is_sparse:
|
||||||
|
raise RuntimeError("NAdamW does not support sparse gradients")
|
||||||
|
|
||||||
|
grad = param.grad
|
||||||
|
state = self.state[param]
|
||||||
|
if not state:
|
||||||
|
state["step"] = 0
|
||||||
|
state["m"] = torch.zeros_like(param)
|
||||||
|
state["v"] = torch.zeros_like(param)
|
||||||
|
|
||||||
|
state["step"] += 1
|
||||||
|
first_moment = state["m"]
|
||||||
|
second_moment = state["v"]
|
||||||
|
first_moment.mul_(beta1).add_(grad, alpha=1 - beta1)
|
||||||
|
second_moment.mul_(beta2).addcmul_(grad, grad, value=1 - beta2)
|
||||||
|
|
||||||
|
bias_correction1 = 1 - beta1 ** state["step"]
|
||||||
|
bias_correction2 = 1 - beta2 ** state["step"]
|
||||||
|
nesterov_moment = (
|
||||||
|
beta1 * first_moment + (1 - beta1) * grad
|
||||||
|
) / bias_correction1
|
||||||
|
corrected_second_moment = second_moment / bias_correction2
|
||||||
|
|
||||||
|
if weight_decay != 0:
|
||||||
|
param.mul_(1 - lr * weight_decay)
|
||||||
|
param.addcdiv_(
|
||||||
|
nesterov_moment,
|
||||||
|
corrected_second_moment.sqrt().add_(eps),
|
||||||
|
value=-lr,
|
||||||
|
)
|
||||||
|
|
||||||
|
return loss
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class OptimizerParameterGroups:
|
||||||
|
nora: list[Tensor]
|
||||||
|
nadamw_decay: list[Tensor]
|
||||||
|
nadamw_no_decay: list[Tensor]
|
||||||
|
|
||||||
|
|
||||||
|
def partition_optimizer_parameters(model: nn.Module) -> OptimizerParameterGroups:
|
||||||
|
"""Partition trainable parameters by module role and parameter identity."""
|
||||||
|
nora_ids: set[int] = set()
|
||||||
|
no_decay_ids: set[int] = set()
|
||||||
|
|
||||||
|
for module_name, module in model.named_modules():
|
||||||
|
if isinstance(module, LoRALinear):
|
||||||
|
for param in module.parameters(recurse=False):
|
||||||
|
if param.requires_grad:
|
||||||
|
no_decay_ids.add(id(param))
|
||||||
|
continue
|
||||||
|
|
||||||
|
if isinstance(module, (Embedding, RMSNorm)):
|
||||||
|
for param in module.parameters(recurse=False):
|
||||||
|
if param.requires_grad:
|
||||||
|
no_decay_ids.add(id(param))
|
||||||
|
continue
|
||||||
|
|
||||||
|
if not isinstance(module, Linear):
|
||||||
|
continue
|
||||||
|
|
||||||
|
if module.bias is not None and module.bias.requires_grad:
|
||||||
|
no_decay_ids.add(id(module.bias))
|
||||||
|
if not module.weight.requires_grad:
|
||||||
|
continue
|
||||||
|
if module_name.rsplit(".", 1)[-1] == "lm_head":
|
||||||
|
no_decay_ids.add(id(module.weight))
|
||||||
|
elif module.weight.ndim == 2:
|
||||||
|
nora_ids.add(id(module.weight))
|
||||||
|
|
||||||
|
nora: list[Tensor] = []
|
||||||
|
nadamw_decay: list[Tensor] = []
|
||||||
|
nadamw_no_decay: list[Tensor] = []
|
||||||
|
seen: set[int] = set()
|
||||||
|
for param in model.parameters():
|
||||||
|
param_id = id(param)
|
||||||
|
if not param.requires_grad or param_id in seen:
|
||||||
|
continue
|
||||||
|
seen.add(param_id)
|
||||||
|
if param_id in no_decay_ids or param.ndim <= 1:
|
||||||
|
nadamw_no_decay.append(param)
|
||||||
|
elif param_id in nora_ids:
|
||||||
|
nora.append(param)
|
||||||
|
else:
|
||||||
|
nadamw_decay.append(param)
|
||||||
|
|
||||||
|
trainable_ids = {id(param) for param in model.parameters() if param.requires_grad}
|
||||||
|
grouped_ids = {id(param) for param in [*nora, *nadamw_decay, *nadamw_no_decay]}
|
||||||
|
if grouped_ids != trainable_ids:
|
||||||
|
missing = len(trainable_ids - grouped_ids)
|
||||||
|
extra = len(grouped_ids - trainable_ids)
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Optimizer parameter partition is incomplete: missing={missing}, extra={extra}"
|
||||||
|
)
|
||||||
|
|
||||||
|
return OptimizerParameterGroups(nora, nadamw_decay, nadamw_no_decay)
|
||||||
|
|
||||||
|
|
||||||
|
@OptimizerFactory.register("nora_nadamw")
|
||||||
|
class NoraNAdamW(Optimizer):
|
||||||
|
"""Nora for internal linear weights and NAdamW for remaining parameters."""
|
||||||
|
|
||||||
|
optimizer_name = "nora_nadamw"
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model: nn.Module,
|
||||||
|
lr: float = 3e-4,
|
||||||
|
weight_decay: float = 0.1,
|
||||||
|
nora_lr: float = 5e-3,
|
||||||
|
nora_weight_decay: float = 0.0,
|
||||||
|
nora_beta: float = 0.95,
|
||||||
|
nora_momentum: float = 0.95,
|
||||||
|
):
|
||||||
|
groups = partition_optimizer_parameters(model)
|
||||||
|
all_params = [
|
||||||
|
*groups.nora,
|
||||||
|
*groups.nadamw_decay,
|
||||||
|
*groups.nadamw_no_decay,
|
||||||
|
]
|
||||||
|
if not all_params:
|
||||||
|
raise ValueError(
|
||||||
|
"Cannot build an optimizer for a model with no trainable parameters"
|
||||||
|
)
|
||||||
|
super().__init__(all_params, {})
|
||||||
|
|
||||||
|
self.nora = (
|
||||||
|
Nora(
|
||||||
|
groups.nora,
|
||||||
|
lr=nora_lr,
|
||||||
|
weight_decay=nora_weight_decay,
|
||||||
|
momentum=nora_momentum,
|
||||||
|
beta=nora_beta,
|
||||||
|
)
|
||||||
|
if groups.nora
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
|
||||||
|
nadamw_groups = []
|
||||||
|
if groups.nadamw_decay:
|
||||||
|
nadamw_groups.append(
|
||||||
|
{"params": groups.nadamw_decay, "weight_decay": weight_decay}
|
||||||
|
)
|
||||||
|
if groups.nadamw_no_decay:
|
||||||
|
nadamw_groups.append(
|
||||||
|
{"params": groups.nadamw_no_decay, "weight_decay": 0.0}
|
||||||
|
)
|
||||||
|
self.nadamw = NAdamW(nadamw_groups, lr=lr) if nadamw_groups else None
|
||||||
|
self.param_groups = refresh_param_groups([self.nora, self.nadamw])
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def step(self, closure=None):
|
||||||
|
return composite_step(
|
||||||
|
[opt for opt in (self.nora, self.nadamw) if opt is not None],
|
||||||
|
closure,
|
||||||
|
)
|
||||||
|
|
||||||
|
def zero_grad(self, set_to_none: bool = True):
|
||||||
|
composite_zero_grad(
|
||||||
|
[opt for opt in (self.nora, self.nadamw) if opt is not None],
|
||||||
|
set_to_none,
|
||||||
|
)
|
||||||
|
|
||||||
|
def state_dict(self) -> dict[str, Any]:
|
||||||
|
return composite_state_dict({"nora": self.nora, "nadamw": self.nadamw})
|
||||||
|
|
||||||
|
def load_state_dict(self, state_dict: dict[str, Any]):
|
||||||
|
if "muon" in state_dict or "adamw" in state_dict:
|
||||||
|
raise ValueError(
|
||||||
|
"Checkpoint uses muon_adamw state; select optimizer='muon_adamw' "
|
||||||
|
"to resume it"
|
||||||
|
)
|
||||||
|
if "nora" not in state_dict or "nadamw" not in state_dict:
|
||||||
|
raise ValueError(
|
||||||
|
"Checkpoint optimizer state is not compatible with nora_nadamw"
|
||||||
|
)
|
||||||
|
|
||||||
|
saved_nora = state_dict["nora"]
|
||||||
|
saved_nadamw = state_dict["nadamw"]
|
||||||
|
if (self.nora is None) != (saved_nora is None):
|
||||||
|
raise ValueError("Checkpoint Nora parameter groups do not match the model")
|
||||||
|
if (self.nadamw is None) != (saved_nadamw is None):
|
||||||
|
raise ValueError(
|
||||||
|
"Checkpoint NAdamW parameter groups do not match the model"
|
||||||
|
)
|
||||||
|
if self.nora is not None:
|
||||||
|
self.nora.load_state_dict(saved_nora)
|
||||||
|
if self.nadamw is not None:
|
||||||
|
self.nadamw.load_state_dict(saved_nadamw)
|
||||||
|
self.param_groups = refresh_param_groups([self.nora, self.nadamw])
|
||||||
@@ -4,12 +4,12 @@ from astrai.parallel.executor import (
|
|||||||
BaseExecutor,
|
BaseExecutor,
|
||||||
DDPExecutor,
|
DDPExecutor,
|
||||||
ExecutorFactory,
|
ExecutorFactory,
|
||||||
FSDP2Executor,
|
|
||||||
FSDPExecutor,
|
FSDPExecutor,
|
||||||
GradientState,
|
GradientState,
|
||||||
NoneExecutor,
|
NoneExecutor,
|
||||||
|
broadcast_state_dict,
|
||||||
|
create_ref_model,
|
||||||
)
|
)
|
||||||
from astrai.parallel.module import ColumnParallelLinear, RowParallelLinear
|
|
||||||
from astrai.parallel.setup import (
|
from astrai.parallel.setup import (
|
||||||
get_current_device,
|
get_current_device,
|
||||||
get_rank,
|
get_rank,
|
||||||
@@ -26,8 +26,6 @@ __all__ = [
|
|||||||
"only_on_rank",
|
"only_on_rank",
|
||||||
"setup_parallel",
|
"setup_parallel",
|
||||||
"spawn_parallel_fn",
|
"spawn_parallel_fn",
|
||||||
"RowParallelLinear",
|
|
||||||
"ColumnParallelLinear",
|
|
||||||
"ExecutorFactory",
|
"ExecutorFactory",
|
||||||
"BaseExecutor",
|
"BaseExecutor",
|
||||||
"GradientState",
|
"GradientState",
|
||||||
@@ -36,5 +34,6 @@ __all__ = [
|
|||||||
"NoneExecutor",
|
"NoneExecutor",
|
||||||
"DDPExecutor",
|
"DDPExecutor",
|
||||||
"FSDPExecutor",
|
"FSDPExecutor",
|
||||||
"FSDP2Executor",
|
"create_ref_model",
|
||||||
|
"broadcast_state_dict",
|
||||||
]
|
]
|
||||||
|
|||||||
+120
-99
@@ -4,18 +4,15 @@ import contextlib
|
|||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from typing import Any, Callable, Optional, Tuple
|
from typing import Any, Callable, Dict, Optional, Tuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.distributed as dist
|
import torch.distributed as dist
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from torch.distributed.fsdp import (
|
from torch.distributed.fsdp import (
|
||||||
FSDPModule,
|
FSDPModule,
|
||||||
FullStateDictConfig,
|
|
||||||
StateDictType,
|
|
||||||
fully_shard,
|
fully_shard,
|
||||||
)
|
)
|
||||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
|
||||||
from torch.distributed.tensor import DTensor
|
from torch.distributed.tensor import DTensor
|
||||||
from torch.nn.parallel import DistributedDataParallel as DDP
|
from torch.nn.parallel import DistributedDataParallel as DDP
|
||||||
from torch.optim import Optimizer
|
from torch.optim import Optimizer
|
||||||
@@ -27,6 +24,82 @@ from astrai.parallel.setup import get_rank, get_world_size
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def broadcast_state_dict(
|
||||||
|
state_dict: Optional[Dict[str, torch.Tensor]],
|
||||||
|
src: int = 0,
|
||||||
|
) -> Optional[Dict[str, torch.Tensor]]:
|
||||||
|
"""Broadcast a state_dict from *src* rank to all ranks.
|
||||||
|
|
||||||
|
Tensors stay on their original device (GPU) for the broadcast.
|
||||||
|
All ranks must call this collectively.
|
||||||
|
|
||||||
|
On non-distributed runs, returns *state_dict* unchanged.
|
||||||
|
"""
|
||||||
|
if not dist.is_initialized() or dist.get_world_size() == 1:
|
||||||
|
return state_dict
|
||||||
|
|
||||||
|
rank = dist.get_rank()
|
||||||
|
|
||||||
|
# Broadcast metadata (keys, shapes, dtypes, device) so non-src ranks
|
||||||
|
# can allocate matching empty tensors on the correct device.
|
||||||
|
if rank == src:
|
||||||
|
device = next(iter(state_dict.values())).device
|
||||||
|
metadata = [
|
||||||
|
(k, tuple(v.shape), v.dtype, str(device)) for k, v in state_dict.items()
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
metadata = None
|
||||||
|
metadata_list = [metadata]
|
||||||
|
dist.broadcast_object_list(metadata_list, src=src)
|
||||||
|
metadata = metadata_list[0]
|
||||||
|
|
||||||
|
# Non-src ranks allocate empty tensors with the broadcasted metadata.
|
||||||
|
if rank != src:
|
||||||
|
state_dict = {
|
||||||
|
k: torch.empty(s, dtype=d, device=torch.device(dev))
|
||||||
|
for k, s, d, dev in metadata
|
||||||
|
}
|
||||||
|
|
||||||
|
# Broadcast each tensor in-place.
|
||||||
|
for tensor in state_dict.values():
|
||||||
|
dist.broadcast(tensor, src=src)
|
||||||
|
|
||||||
|
return state_dict
|
||||||
|
|
||||||
|
|
||||||
|
def create_ref_model(
|
||||||
|
model_fn: Callable[[], nn.Module],
|
||||||
|
executor: Optional["BaseExecutor"] = None,
|
||||||
|
model: Optional[nn.Module] = None,
|
||||||
|
state_dict: Optional[Dict[str, torch.Tensor]] = None,
|
||||||
|
device: Optional[str] = None,
|
||||||
|
) -> Optional[nn.Module]:
|
||||||
|
"""Create a frozen reference model from executor or state dict.
|
||||||
|
|
||||||
|
In distributed mode (FSDP), ``unwrap_model`` returns ``None`` on
|
||||||
|
non-rank-0. The state_dict is broadcast from rank-0 to all ranks
|
||||||
|
so every rank gets a complete copy.
|
||||||
|
"""
|
||||||
|
if state_dict is None and executor is not None and model is not None:
|
||||||
|
state_dict = executor.unwrap_model(model)
|
||||||
|
|
||||||
|
# FSDP's unwrap_model returns None on non-rank-0. Broadcast from
|
||||||
|
# rank-0 so every rank receives a complete state_dict.
|
||||||
|
if executor is not None and executor.use_distributed:
|
||||||
|
state_dict = broadcast_state_dict(state_dict)
|
||||||
|
|
||||||
|
if state_dict is None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
ref_model = model_fn()
|
||||||
|
ref_model.load_state_dict(state_dict)
|
||||||
|
ref_model.requires_grad_(False)
|
||||||
|
ref_model.eval()
|
||||||
|
if device is not None:
|
||||||
|
ref_model = ref_model.to(device=device)
|
||||||
|
return ref_model
|
||||||
|
|
||||||
|
|
||||||
class GradientState:
|
class GradientState:
|
||||||
def __init__(self, grad_accum_steps: int = 1):
|
def __init__(self, grad_accum_steps: int = 1):
|
||||||
self.num_steps = max(grad_accum_steps, 1)
|
self.num_steps = max(grad_accum_steps, 1)
|
||||||
@@ -95,11 +168,14 @@ class BaseExecutor:
|
|||||||
optimizer_fn: Optional[Callable[[nn.Module], Optimizer]] = None,
|
optimizer_fn: Optional[Callable[[nn.Module], Optimizer]] = None,
|
||||||
scheduler_fn: Optional[Callable[[Optimizer], LRScheduler]] = None,
|
scheduler_fn: Optional[Callable[[Optimizer], LRScheduler]] = None,
|
||||||
before_wrap: Optional[Callable[[nn.Module], nn.Module]] = None,
|
before_wrap: Optional[Callable[[nn.Module], nn.Module]] = None,
|
||||||
|
after_wrap: Optional[Callable[[nn.Module], nn.Module]] = None,
|
||||||
) -> Tuple[nn.Module, Optional[Optimizer], Optional[LRScheduler]]:
|
) -> Tuple[nn.Module, Optional[Optimizer], Optional[LRScheduler]]:
|
||||||
model = model_fn()
|
model = model_fn()
|
||||||
if before_wrap is not None:
|
if before_wrap is not None:
|
||||||
model = before_wrap(model)
|
model = before_wrap(model)
|
||||||
model = self._prepare_model(model)
|
model = self._prepare_model(model)
|
||||||
|
if after_wrap is not None:
|
||||||
|
model = after_wrap(model)
|
||||||
optimizer = None
|
optimizer = None
|
||||||
scheduler = None
|
scheduler = None
|
||||||
if optimizer_fn is not None:
|
if optimizer_fn is not None:
|
||||||
@@ -238,88 +314,11 @@ class DDPExecutor(BaseExecutor):
|
|||||||
|
|
||||||
@ExecutorFactory.register("fsdp")
|
@ExecutorFactory.register("fsdp")
|
||||||
class FSDPExecutor(BaseExecutor):
|
class FSDPExecutor(BaseExecutor):
|
||||||
def __init__(
|
"""FSDP executor using `torch.distributed.fsdp.fully_shard` (per-module API).
|
||||||
self,
|
|
||||||
grad_accum_steps: int = 1,
|
|
||||||
process_group=None,
|
|
||||||
sharding_strategy=None,
|
|
||||||
cpu_offload=None,
|
|
||||||
auto_wrap_policy=None,
|
|
||||||
backward_prefetch=None,
|
|
||||||
mixed_precision=None,
|
|
||||||
ignored_modules=None,
|
|
||||||
param_init_fn=None,
|
|
||||||
sync_module_states: bool = False,
|
|
||||||
forward_prefetch: bool = False,
|
|
||||||
limit_all_gathers: bool = True,
|
|
||||||
ignored_states=None,
|
|
||||||
device_mesh=None,
|
|
||||||
):
|
|
||||||
super().__init__(grad_accum_steps=grad_accum_steps)
|
|
||||||
self._fsdp_kwargs = {
|
|
||||||
k: v
|
|
||||||
for k, v in dict(
|
|
||||||
process_group=process_group,
|
|
||||||
sharding_strategy=sharding_strategy,
|
|
||||||
cpu_offload=cpu_offload,
|
|
||||||
auto_wrap_policy=auto_wrap_policy,
|
|
||||||
backward_prefetch=backward_prefetch,
|
|
||||||
mixed_precision=mixed_precision,
|
|
||||||
ignored_modules=ignored_modules,
|
|
||||||
param_init_fn=param_init_fn,
|
|
||||||
sync_module_states=sync_module_states,
|
|
||||||
forward_prefetch=forward_prefetch,
|
|
||||||
limit_all_gathers=limit_all_gathers,
|
|
||||||
use_orig_params=True,
|
|
||||||
ignored_states=ignored_states,
|
|
||||||
device_mesh=device_mesh,
|
|
||||||
).items()
|
|
||||||
if v is not None
|
|
||||||
}
|
|
||||||
self._original_model: Optional[nn.Module] = None
|
|
||||||
|
|
||||||
def _prepare_model(self, model: nn.Module) -> nn.Module:
|
|
||||||
if not self.use_distributed:
|
|
||||||
logger.warning("FSDP backend selected but world_size=1, model not wrapped")
|
|
||||||
return model
|
|
||||||
self._original_model = model
|
|
||||||
device_id = torch.device("cuda", get_rank())
|
|
||||||
model = FSDP(model, device_id=device_id, **self._fsdp_kwargs)
|
|
||||||
logger.info("Model wrapped with FSDP (world_size=%d)", get_world_size())
|
|
||||||
return model
|
|
||||||
|
|
||||||
def _no_sync(self, model: nn.Module):
|
|
||||||
if isinstance(model, FSDP):
|
|
||||||
return model.no_sync()
|
|
||||||
return contextlib.nullcontext()
|
|
||||||
|
|
||||||
def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float:
|
|
||||||
if isinstance(model, FSDP) and self.use_distributed:
|
|
||||||
total_norm = model.clip_grad_norm_(max_norm)
|
|
||||||
if isinstance(total_norm, torch.Tensor):
|
|
||||||
return total_norm.item()
|
|
||||||
return total_norm
|
|
||||||
return super().clip_grad_norm(model, max_norm)
|
|
||||||
|
|
||||||
def unwrap_model(self, model: nn.Module):
|
|
||||||
if isinstance(model, FSDP) and self.use_distributed:
|
|
||||||
with FSDP.state_dict_type(
|
|
||||||
model,
|
|
||||||
StateDictType.FULL_STATE_DICT,
|
|
||||||
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
|
|
||||||
):
|
|
||||||
return model.state_dict()
|
|
||||||
|
|
||||||
return model.state_dict()
|
|
||||||
|
|
||||||
|
|
||||||
@ExecutorFactory.register("fsdp2")
|
|
||||||
class FSDP2Executor(BaseExecutor):
|
|
||||||
"""FSDP2 executor using `torch.distributed.fsdp.fully_shard` (per-module API).
|
|
||||||
|
|
||||||
Wraps each child module individually via ``fully_shard``.
|
Wraps each child module individually via ``fully_shard``.
|
||||||
Skips the root model because ``ABC + Generic[T]`` in the MRO makes
|
Skips the root model because ``ABC + Generic[T]`` in the MRO makes
|
||||||
FSDP2's dynamic ``__class__`` assignment fail at the CPython level.
|
``fully_shard``'s dynamic ``__class__`` assignment fail at the CPython level.
|
||||||
Original ``Parameter`` objects are preserved (as DTensors) — no
|
Original ``Parameter`` objects are preserved (as DTensors) — no
|
||||||
``FlatParameter``, no ``use_orig_params=True`` hack.
|
``FlatParameter``, no ``use_orig_params=True`` hack.
|
||||||
"""
|
"""
|
||||||
@@ -329,7 +328,7 @@ class FSDP2Executor(BaseExecutor):
|
|||||||
grad_accum_steps: int = 1,
|
grad_accum_steps: int = 1,
|
||||||
mesh: Optional[Any] = None,
|
mesh: Optional[Any] = None,
|
||||||
mp_policy: Optional[Any] = None,
|
mp_policy: Optional[Any] = None,
|
||||||
reshard_after_forward: bool = True,
|
reshard_after_forward: bool = False,
|
||||||
):
|
):
|
||||||
super().__init__(grad_accum_steps=grad_accum_steps)
|
super().__init__(grad_accum_steps=grad_accum_steps)
|
||||||
self._mesh = mesh
|
self._mesh = mesh
|
||||||
@@ -338,7 +337,7 @@ class FSDP2Executor(BaseExecutor):
|
|||||||
|
|
||||||
def _prepare_model(self, model: nn.Module) -> nn.Module:
|
def _prepare_model(self, model: nn.Module) -> nn.Module:
|
||||||
if not self.use_distributed:
|
if not self.use_distributed:
|
||||||
logger.warning("FSDP2 backend selected but world_size=1, model not wrapped")
|
logger.warning("FSDP backend selected but world_size=1, model not wrapped")
|
||||||
return model
|
return model
|
||||||
|
|
||||||
kwargs = dict(
|
kwargs = dict(
|
||||||
@@ -356,7 +355,7 @@ class FSDP2Executor(BaseExecutor):
|
|||||||
fully_shard(child, **kwargs)
|
fully_shard(child, **kwargs)
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
"FSDP2 wrapping applied to %d direct children (root skipped for ABC compat)",
|
"FSDP wrapping applied to %d direct children (root skipped for ABC compat)",
|
||||||
len(list(model.children())),
|
len(list(model.children())),
|
||||||
)
|
)
|
||||||
return model
|
return model
|
||||||
@@ -376,32 +375,54 @@ class FSDP2Executor(BaseExecutor):
|
|||||||
yield
|
yield
|
||||||
|
|
||||||
def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float:
|
def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float:
|
||||||
if self.use_distributed:
|
if not self.use_distributed:
|
||||||
total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm)
|
return super().clip_grad_norm(model, max_norm)
|
||||||
if isinstance(total_norm, torch.Tensor):
|
|
||||||
return total_norm.item()
|
# FSDP params are DTensors (sharded across ranks).
|
||||||
return total_norm
|
# torch.nn.utils.clip_grad_norm_ computes LOCAL norm per rank,
|
||||||
return super().clip_grad_norm(model, max_norm)
|
# so we must all-reduce to get the global norm before clipping.
|
||||||
|
local_norm = torch.nn.utils.get_total_norm(
|
||||||
|
[p.grad for p in model.parameters() if p.grad is not None],
|
||||||
|
)
|
||||||
|
if isinstance(local_norm, DTensor):
|
||||||
|
local_norm = local_norm.to_local()
|
||||||
|
total_norm_sq = local_norm**2
|
||||||
|
dist.all_reduce(total_norm_sq, op=dist.ReduceOp.SUM)
|
||||||
|
total_norm = total_norm_sq.sqrt()
|
||||||
|
|
||||||
|
clip_coef = max_norm / (total_norm + 1e-6)
|
||||||
|
clip_coef_clamped = torch.clamp(clip_coef, max=1.0)
|
||||||
|
for p in model.parameters():
|
||||||
|
if p.grad is not None:
|
||||||
|
p.grad.mul_(clip_coef_clamped)
|
||||||
|
|
||||||
|
return total_norm.item()
|
||||||
|
|
||||||
def unwrap_model(self, model: nn.Module):
|
def unwrap_model(self, model: nn.Module):
|
||||||
if not self.use_distributed:
|
if not self.use_distributed:
|
||||||
return model.state_dict()
|
return model.state_dict()
|
||||||
|
|
||||||
if get_rank() != 0:
|
# unshard() and full_tensor() are collective ops — all ranks must
|
||||||
return None
|
# participate. Non-rank-0 ranks still call them but discard results.
|
||||||
|
|
||||||
for module in model.modules():
|
for module in model.modules():
|
||||||
if isinstance(module, FSDPModule):
|
if isinstance(module, FSDPModule):
|
||||||
module.unshard()
|
module.unshard()
|
||||||
|
|
||||||
state_dict = model.state_dict()
|
state_dict = model.state_dict()
|
||||||
result = {
|
result = {}
|
||||||
k: (v.full_tensor() if isinstance(v, DTensor) else v)
|
for k, v in state_dict.items():
|
||||||
for k, v in state_dict.items()
|
if isinstance(v, DTensor):
|
||||||
}
|
full = v.full_tensor()
|
||||||
|
if get_rank() == 0:
|
||||||
|
result[k] = full
|
||||||
|
elif get_rank() == 0:
|
||||||
|
result[k] = v
|
||||||
|
|
||||||
for module in model.modules():
|
for module in model.modules():
|
||||||
if isinstance(module, FSDPModule):
|
if isinstance(module, FSDPModule):
|
||||||
module.reshard()
|
module.reshard()
|
||||||
|
|
||||||
|
if get_rank() != 0:
|
||||||
|
return None
|
||||||
|
|
||||||
return result
|
return result
|
||||||
|
|||||||
@@ -1,115 +0,0 @@
|
|||||||
from typing import Dict
|
|
||||||
|
|
||||||
import torch
|
|
||||||
import torch.distributed as dist
|
|
||||||
import torch.nn as nn
|
|
||||||
import torch.nn.functional as F
|
|
||||||
from torch import Tensor
|
|
||||||
|
|
||||||
|
|
||||||
class ParallelModel(nn.Module):
|
|
||||||
def __init__(self, process_group: dist.ProcessGroup):
|
|
||||||
super().__init__()
|
|
||||||
self.process_group = process_group
|
|
||||||
self.rank = dist.get_rank(self.process_group)
|
|
||||||
self.world_size = dist.get_world_size(self.process_group)
|
|
||||||
|
|
||||||
|
|
||||||
class RowParallelLinear(ParallelModel):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
process_group: dist.ProcessGroup,
|
|
||||||
in_features: int,
|
|
||||||
out_features: int,
|
|
||||||
bias: bool = True,
|
|
||||||
reduce_results: bool = True,
|
|
||||||
):
|
|
||||||
super().__init__(process_group)
|
|
||||||
|
|
||||||
self.in_features = in_features
|
|
||||||
self.out_features = out_features
|
|
||||||
self.in_features_per_rank = in_features // self.world_size
|
|
||||||
self.reduce_results = reduce_results
|
|
||||||
|
|
||||||
if in_features % self.world_size != 0:
|
|
||||||
raise ValueError(
|
|
||||||
f"in_features must be divisible by world_size. Got {in_features} and {self.world_size}"
|
|
||||||
)
|
|
||||||
|
|
||||||
self.weight = nn.Parameter(torch.empty(out_features, self.in_features_per_rank))
|
|
||||||
self.bias = nn.Parameter(torch.zeros(out_features)) if bias else None
|
|
||||||
|
|
||||||
def forward(self, input: Tensor) -> Tensor:
|
|
||||||
output = F.linear(input, self.weight)
|
|
||||||
|
|
||||||
if self.reduce_results:
|
|
||||||
dist.all_reduce(output, op=dist.ReduceOp.SUM, group=self.process_group)
|
|
||||||
|
|
||||||
if self.bias is not None:
|
|
||||||
output += self.bias
|
|
||||||
|
|
||||||
return output
|
|
||||||
|
|
||||||
def load_state_dict(self, state_dict: Dict[str, Tensor]):
|
|
||||||
full_weight = state_dict.get("weight")
|
|
||||||
full_bias = state_dict.get("bias")
|
|
||||||
|
|
||||||
start_idx = self.rank * self.in_features_per_rank
|
|
||||||
end_idx = start_idx + self.in_features_per_rank
|
|
||||||
weight_slice = full_weight[:, start_idx:end_idx]
|
|
||||||
self.weight.data.copy_(weight_slice)
|
|
||||||
|
|
||||||
if self.bias is not None:
|
|
||||||
self.bias.data.copy_(full_bias)
|
|
||||||
|
|
||||||
|
|
||||||
class ColumnParallelLinear(ParallelModel):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
process_group: dist.ProcessGroup,
|
|
||||||
in_features: int,
|
|
||||||
out_features: int,
|
|
||||||
bias: bool = True,
|
|
||||||
gather_results: bool = True,
|
|
||||||
):
|
|
||||||
super().__init__(process_group)
|
|
||||||
|
|
||||||
self.in_features = in_features
|
|
||||||
self.out_features = out_features
|
|
||||||
self.out_features_per_rank = out_features // self.world_size
|
|
||||||
self.gather_results = gather_results
|
|
||||||
|
|
||||||
if out_features % self.world_size != 0:
|
|
||||||
raise ValueError(
|
|
||||||
f"out_features must be divisible by world_size. Got {out_features} and {self.world_size}"
|
|
||||||
)
|
|
||||||
|
|
||||||
self.weight = nn.Parameter(
|
|
||||||
torch.empty(self.out_features_per_rank, self.in_features)
|
|
||||||
)
|
|
||||||
self.bias = (
|
|
||||||
nn.Parameter(torch.zeros(self.out_features_per_rank)) if bias else None
|
|
||||||
)
|
|
||||||
|
|
||||||
def forward(self, input: Tensor) -> Tensor:
|
|
||||||
output = F.linear(input, self.weight, self.bias)
|
|
||||||
|
|
||||||
if self.gather_results:
|
|
||||||
output_list = [torch.empty_like(output) for _ in range(self.world_size)]
|
|
||||||
dist.all_gather(output_list, output, group=self.process_group)
|
|
||||||
output = torch.cat(output_list, dim=-1)
|
|
||||||
|
|
||||||
return output
|
|
||||||
|
|
||||||
def load_state_dict(self, state_dict: Dict[str, Tensor]):
|
|
||||||
full_weight = state_dict.get("weight")
|
|
||||||
full_bias = state_dict.get("bias")
|
|
||||||
|
|
||||||
start_idx = self.rank * self.out_features_per_rank
|
|
||||||
end_idx = start_idx + self.out_features_per_rank
|
|
||||||
weight_slice = full_weight[start_idx:end_idx, :]
|
|
||||||
self.weight.data.copy_(weight_slice)
|
|
||||||
|
|
||||||
if self.bias is not None:
|
|
||||||
bias_slice = full_bias[start_idx:end_idx]
|
|
||||||
self.bias.data.copy_(bias_slice)
|
|
||||||
@@ -12,7 +12,7 @@ import torch
|
|||||||
import torch.distributed as dist
|
import torch.distributed as dist
|
||||||
import torch.multiprocessing as mp
|
import torch.multiprocessing as mp
|
||||||
|
|
||||||
from astrai.parallel.signal_handler import install_early_signal_handlers
|
from astrai.signal_handler import install_early_signal_handlers
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
"""Config-driven JSONL preprocessing pipeline.
|
"""Config-driven JSONL preprocessing pipeline.
|
||||||
|
|
||||||
Composes a :class:`BaseMaskBuilder` (selected by ``input.type``) with
|
Composes a :class:`BaseMaskBuilder` (selected by ``input.type``) with
|
||||||
sharding and flush to ``.h5`` / ``.bin`` storage. Packing, position-id
|
sharding and flush to ``.bin`` storage. Packing, position-id
|
||||||
generation and storage writing are each delegated to pluggable strategies,
|
generation and storage writing are each delegated to pluggable strategies,
|
||||||
dispatched by configuration keys.
|
dispatched by configuration keys.
|
||||||
|
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
"""Storage writer strategies for pipeline output.
|
"""Storage writer strategies for pipeline output.
|
||||||
|
|
||||||
The :class:`StoreWriter` abstraction decouples the pipeline from the
|
The :class:`StoreWriter` abstraction decouples the pipeline from the
|
||||||
concrete storage format (bin / h5). The pipeline builds a ``{key:
|
concrete storage format (bin). The pipeline builds a ``{key:
|
||||||
List[Tensor]}`` dict and delegates the write to the writer selected
|
List[Tensor]}`` dict and delegates the write to the writer selected
|
||||||
by ``output.storage_format``.
|
by ``output.storage_format``.
|
||||||
"""
|
"""
|
||||||
@@ -15,7 +15,7 @@ from typing import Dict, List
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
from astrai.serialization import save_bin, save_h5
|
from astrai.serialization import save_bin
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -54,22 +54,3 @@ class BinWriter(StoreWriter):
|
|||||||
exc_info=True,
|
exc_info=True,
|
||||||
)
|
)
|
||||||
raise
|
raise
|
||||||
|
|
||||||
|
|
||||||
@StoreWriterFactory.register("h5")
|
|
||||||
class H5Writer(StoreWriter):
|
|
||||||
def save(self, output_dir, domain, shard_idx, tensors):
|
|
||||||
chunk_dir = os.path.join(output_dir, domain)
|
|
||||||
file_path = os.path.join(chunk_dir, f"data_{shard_idx:04d}.h5")
|
|
||||||
try:
|
|
||||||
save_h5(chunk_dir, f"data_{shard_idx:04d}", tensors)
|
|
||||||
except Exception:
|
|
||||||
if os.path.exists(file_path):
|
|
||||||
os.remove(file_path)
|
|
||||||
logger.error(
|
|
||||||
"Failed to write shard %s/data_%04d.h5, cleaned up partial output",
|
|
||||||
domain,
|
|
||||||
shard_idx,
|
|
||||||
exc_info=True,
|
|
||||||
)
|
|
||||||
raise
|
|
||||||
|
|||||||
@@ -20,9 +20,7 @@ from astrai.serialization.checkpoint import (
|
|||||||
from astrai.serialization.dataset import (
|
from astrai.serialization.dataset import (
|
||||||
load_bin,
|
load_bin,
|
||||||
load_bin_offsets,
|
load_bin_offsets,
|
||||||
load_h5,
|
|
||||||
save_bin,
|
save_bin,
|
||||||
save_h5,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
@@ -39,7 +37,5 @@ __all__ = [
|
|||||||
"save_torch",
|
"save_torch",
|
||||||
"load_bin",
|
"load_bin",
|
||||||
"load_bin_offsets",
|
"load_bin_offsets",
|
||||||
"load_h5",
|
|
||||||
"save_bin",
|
"save_bin",
|
||||||
"save_h5",
|
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -1,55 +1,14 @@
|
|||||||
"""Dataset storage serialization helpers (HDF5 / memory-mapped binary)."""
|
"""Dataset storage serialization helpers (memory-mapped binary)."""
|
||||||
|
|
||||||
import json
|
import json
|
||||||
import os
|
import os
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any, Dict, List, Optional
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
import h5py
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
|
|
||||||
def save_h5(file_path: str, file_name: str, tensor_group: Dict[str, List[Tensor]]):
|
|
||||||
os.makedirs(file_path, exist_ok=True)
|
|
||||||
full_file_path = os.path.join(file_path, f"{file_name}.h5")
|
|
||||||
with h5py.File(full_file_path, "w") as f:
|
|
||||||
for key, tensors in tensor_group.items():
|
|
||||||
grp = f.create_group(key)
|
|
||||||
for idx, tensor in enumerate(tensors):
|
|
||||||
arr = tensor.cpu().numpy()
|
|
||||||
grp.create_dataset(f"data_{idx}", data=arr)
|
|
||||||
|
|
||||||
|
|
||||||
def load_h5(file_path: str, share_memory=True) -> Dict[str, List[Tensor]]:
|
|
||||||
tensor_group: Dict[str, List[Tensor]] = {}
|
|
||||||
|
|
||||||
root_path = Path(file_path)
|
|
||||||
if root_path.is_file() and root_path.suffix in (".h5", ".hdf5"):
|
|
||||||
h5_files = [root_path]
|
|
||||||
else:
|
|
||||||
h5_files = list(root_path.rglob("*.h5")) + list(root_path.rglob("*.hdf5"))
|
|
||||||
|
|
||||||
for h5_file in h5_files:
|
|
||||||
with h5py.File(h5_file, "r") as f:
|
|
||||||
for key in f.keys():
|
|
||||||
grp = f[key]
|
|
||||||
dsets = []
|
|
||||||
for dset_name in grp.keys():
|
|
||||||
dset = grp[dset_name]
|
|
||||||
tensor = torch.from_numpy(dset[:])
|
|
||||||
if share_memory:
|
|
||||||
tensor = tensor.share_memory_()
|
|
||||||
dsets.append(tensor)
|
|
||||||
|
|
||||||
if tensor_group.get(key) is None:
|
|
||||||
tensor_group[key] = []
|
|
||||||
tensor_group[key].extend(dsets)
|
|
||||||
|
|
||||||
return tensor_group
|
|
||||||
|
|
||||||
|
|
||||||
def save_bin(
|
def save_bin(
|
||||||
file_path: str,
|
file_path: str,
|
||||||
tensor_group: Dict[str, List[Tensor]],
|
tensor_group: Dict[str, List[Tensor]],
|
||||||
@@ -65,7 +24,7 @@ def save_bin(
|
|||||||
offsets, preserving backward compatibility.
|
offsets, preserving backward compatibility.
|
||||||
|
|
||||||
Nested keys (``List[List[Tensor]]`` such as GRPO ``responses``) are
|
Nested keys (``List[List[Tensor]]`` such as GRPO ``responses``) are
|
||||||
not supported in bin format — use H5 for those.
|
not supported in bin format — use JSONL for those.
|
||||||
"""
|
"""
|
||||||
os.makedirs(file_path, exist_ok=True)
|
os.makedirs(file_path, exist_ok=True)
|
||||||
record_keys = set(record_keys or [])
|
record_keys = set(record_keys or [])
|
||||||
@@ -74,7 +33,7 @@ def save_bin(
|
|||||||
if tensors and isinstance(tensors[0], list):
|
if tensors and isinstance(tensors[0], list):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Nested key '{key}' (List[List[Tensor]]) is not supported "
|
f"Nested key '{key}' (List[List[Tensor]]) is not supported "
|
||||||
f"in bin format. Use H5 or JSONL storage instead."
|
f"in bin format. Use JSONL storage instead."
|
||||||
)
|
)
|
||||||
cat = torch.cat(tensors, dim=0)
|
cat = torch.cat(tensors, dim=0)
|
||||||
entry: Dict[str, Any] = {
|
entry: Dict[str, Any] = {
|
||||||
@@ -112,7 +71,7 @@ def load_bin_offsets(file_path: str) -> Dict[str, List[int]]:
|
|||||||
|
|
||||||
Returns an empty dict when no key has offsets (legacy bin files),
|
Returns an empty dict when no key has offsets (legacy bin files),
|
||||||
in which case record-mode access falls back to per-record segment
|
in which case record-mode access falls back to per-record segment
|
||||||
indexing (H5/JSONL layout).
|
indexing (JSONL layout).
|
||||||
"""
|
"""
|
||||||
with open(os.path.join(file_path, "meta.json"), "r") as f:
|
with open(os.path.join(file_path, "meta.json"), "r") as f:
|
||||||
meta = json.load(f)
|
meta = json.load(f)
|
||||||
|
|||||||
@@ -38,12 +38,27 @@ class ChatTemplate:
|
|||||||
The compiled :class:`~jinja2.Template` holds a dynamically-generated
|
The compiled :class:`~jinja2.Template` holds a dynamically-generated
|
||||||
``root`` render function whose ``__module__`` is ``None``; under
|
``root`` render function whose ``__module__`` is ``None``; under
|
||||||
``pickle`` it falls back to ``__main__`` and breaks ``spawn``-based
|
``pickle`` it falls back to ``__main__`` and breaks ``spawn``-based
|
||||||
multiprocessing. By deferring compilation to first access, the
|
multiprocessing. :meth:`__getstate__` drops the cached template so
|
||||||
default pickle protocol serialises only ``template_str``; each
|
that pickle serialises only ``template_str``; each worker rebuilds
|
||||||
worker rebuilds the cache on first render.
|
the cache on first render.
|
||||||
"""
|
"""
|
||||||
return Template(self.template_str)
|
return Template(self.template_str)
|
||||||
|
|
||||||
|
def __getstate__(self) -> Dict[str, Any]:
|
||||||
|
"""Exclude the cached Jinja2 template from pickling.
|
||||||
|
|
||||||
|
``Template.root_render_func`` is a dynamically generated closure
|
||||||
|
that cannot be pickled by reference. Dropping ``_compiled`` here
|
||||||
|
lets :class:`cached_property` rebuild it on first access after
|
||||||
|
unpickle.
|
||||||
|
"""
|
||||||
|
state = self.__dict__.copy()
|
||||||
|
state.pop("_compiled", None)
|
||||||
|
return state
|
||||||
|
|
||||||
|
def __setstate__(self, state: Dict[str, Any]) -> None:
|
||||||
|
self.__dict__.update(state)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_string(
|
def from_string(
|
||||||
cls,
|
cls,
|
||||||
|
|||||||
@@ -20,8 +20,6 @@ Messages = List[Message]
|
|||||||
class AutoTokenizer:
|
class AutoTokenizer:
|
||||||
"""Base tokenizer class with automatic loading support"""
|
"""Base tokenizer class with automatic loading support"""
|
||||||
|
|
||||||
TOKENIZER_CLASSES = {} # Registry for auto-loading
|
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
path: Optional[Union[str, Path]] = None,
|
path: Optional[Union[str, Path]] = None,
|
||||||
@@ -108,17 +106,6 @@ class AutoTokenizer:
|
|||||||
with open(save_path / "tokenizer_config.json", "w", encoding="utf-8") as f:
|
with open(save_path / "tokenizer_config.json", "w", encoding="utf-8") as f:
|
||||||
json.dump(config, f, ensure_ascii=False, indent=2)
|
json.dump(config, f, ensure_ascii=False, indent=2)
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def register_tokenizer(cls, name: str, tokenizer_class: type):
|
|
||||||
"""
|
|
||||||
Register a new tokenizer class.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
name: Name to register the tokenizer class under
|
|
||||||
tokenizer_class: The tokenizer class to register
|
|
||||||
"""
|
|
||||||
cls.TOKENIZER_CLASSES[name] = tokenizer_class
|
|
||||||
|
|
||||||
def encode(
|
def encode(
|
||||||
self,
|
self,
|
||||||
tokens: Union[str, List[str]],
|
tokens: Union[str, List[str]],
|
||||||
|
|||||||
@@ -22,6 +22,51 @@ def grad_norm(model: nn.Module, per_param: bool = False) -> float | Dict[str, fl
|
|||||||
return total_sq.sqrt().item()
|
return total_sq.sqrt().item()
|
||||||
|
|
||||||
|
|
||||||
|
class GradSNRTracker:
|
||||||
|
"""Track gradient signal-to-noise ratio via EMA of first/second moments.
|
||||||
|
|
||||||
|
SNR = E[g]^2 / Var(g) = E[g]^2 / (E[g^2] - E[g]^2)
|
||||||
|
|
||||||
|
The tracker accumulates per-parameter EMA moments across optimizer steps.
|
||||||
|
Call ``update`` after backward (before ``optimizer.step``) and read
|
||||||
|
``snr`` to get the aggregate SNR across all parameters.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, beta: float = 0.999, eps: float = 1e-8):
|
||||||
|
self.beta = beta
|
||||||
|
self.eps = eps
|
||||||
|
self._first: Dict[int, torch.Tensor] = {}
|
||||||
|
self._second: Dict[int, torch.Tensor] = {}
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def update(self, model: nn.Module) -> None:
|
||||||
|
beta = self.beta
|
||||||
|
for param in model.parameters():
|
||||||
|
if param.grad is None:
|
||||||
|
continue
|
||||||
|
pid = id(param)
|
||||||
|
g = param.grad.detach()
|
||||||
|
if pid not in self._first:
|
||||||
|
self._first[pid] = g.clone()
|
||||||
|
self._second[pid] = g.pow(2).clone()
|
||||||
|
else:
|
||||||
|
self._first[pid].mul_(beta).add_(g, alpha=1 - beta)
|
||||||
|
self._second[pid].mul_(beta).addcmul_(g, g, value=1 - beta)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def snr(self) -> float:
|
||||||
|
if not self._first:
|
||||||
|
return 0.0
|
||||||
|
total_signal = 0.0
|
||||||
|
total_noise = 0.0
|
||||||
|
for m, v in zip(self._first.values(), self._second.values()):
|
||||||
|
signal = m.pow(2).sum().item()
|
||||||
|
noise = (v - m.pow(2)).clamp(min=0).sum().item()
|
||||||
|
total_signal += signal
|
||||||
|
total_noise += noise
|
||||||
|
return total_signal / (total_noise + self.eps)
|
||||||
|
|
||||||
|
|
||||||
def ctx_get_loss(ctx):
|
def ctx_get_loss(ctx):
|
||||||
return ctx.loss
|
return ctx.loss
|
||||||
|
|
||||||
@@ -36,3 +81,10 @@ def ctx_get_val_loss(ctx):
|
|||||||
|
|
||||||
def ctx_get_grad_norm(ctx):
|
def ctx_get_grad_norm(ctx):
|
||||||
return ctx.grad_norm
|
return ctx.grad_norm
|
||||||
|
|
||||||
|
|
||||||
|
def ctx_get_grad_snr(ctx):
|
||||||
|
tracker = getattr(ctx, "grad_snr_tracker", None)
|
||||||
|
if tracker is None:
|
||||||
|
return None
|
||||||
|
return tracker.snr
|
||||||
|
|||||||
@@ -9,20 +9,10 @@ import torch.nn.functional as F
|
|||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.parallel.executor import broadcast_state_dict
|
||||||
from astrai.trainer.rollout import RolloutResult
|
from astrai.trainer.rollout import RolloutResult
|
||||||
|
|
||||||
|
|
||||||
def create_ref_model(
|
|
||||||
model_fn: Callable[[], nn.Module], state_dict: Dict[str, Tensor]
|
|
||||||
) -> nn.Module:
|
|
||||||
"""Create a frozen reference model from model_fn + full state dict."""
|
|
||||||
ref_model = model_fn()
|
|
||||||
ref_model.load_state_dict(state_dict)
|
|
||||||
ref_model.requires_grad_(False)
|
|
||||||
ref_model.eval()
|
|
||||||
return ref_model
|
|
||||||
|
|
||||||
|
|
||||||
def move_to_device(batch: Dict[str, Tensor], device: str) -> Dict[str, Tensor]:
|
def move_to_device(batch: Dict[str, Tensor], device: str) -> Dict[str, Tensor]:
|
||||||
"""Move batch tensors to specified device with non-blocking transfer."""
|
"""Move batch tensors to specified device with non-blocking transfer."""
|
||||||
return {key: value.to(device, non_blocking=True) for key, value in batch.items()}
|
return {key: value.to(device, non_blocking=True) for key, value in batch.items()}
|
||||||
@@ -401,7 +391,11 @@ class GRPOStrategy(BaseStrategy):
|
|||||||
|
|
||||||
def sync_old_model(self):
|
def sync_old_model(self):
|
||||||
"""Copy current policy weights to old model."""
|
"""Copy current policy weights to old model."""
|
||||||
self.old_model.load_state_dict(self.executor.unwrap_model(self.model))
|
state_dict = self.executor.unwrap_model(self.model)
|
||||||
|
if self.executor.use_distributed:
|
||||||
|
state_dict = broadcast_state_dict(state_dict)
|
||||||
|
if state_dict is not None:
|
||||||
|
self.old_model.load_state_dict(state_dict)
|
||||||
|
|
||||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||||
batch = move_to_device(batch, self.device)
|
batch = move_to_device(batch, self.device)
|
||||||
@@ -510,5 +504,5 @@ class GRPOStrategy(BaseStrategy):
|
|||||||
# Factory aliases: online variants use the same strategy class; the
|
# Factory aliases: online variants use the same strategy class; the
|
||||||
# ``RolloutRunner`` is injected by ``TrainContextBuilder`` to enable
|
# ``RolloutRunner`` is injected by ``TrainContextBuilder`` to enable
|
||||||
# online mode, so no separate subclass is needed.
|
# online mode, so no separate subclass is needed.
|
||||||
StrategyFactory._entries["online_grpo"] = GRPOStrategy
|
StrategyFactory.register("online_grpo")(GRPOStrategy)
|
||||||
StrategyFactory._entries["online_dpo"] = DPOStrategy
|
StrategyFactory.register("online_dpo")(DPOStrategy)
|
||||||
|
|||||||
@@ -18,6 +18,7 @@ from astrai.parallel.setup import get_current_device
|
|||||||
from astrai.serialization import Checkpoint
|
from astrai.serialization import Checkpoint
|
||||||
from astrai.trainer.metric_util import (
|
from astrai.trainer.metric_util import (
|
||||||
ctx_get_grad_norm,
|
ctx_get_grad_norm,
|
||||||
|
ctx_get_grad_snr,
|
||||||
ctx_get_loss,
|
ctx_get_loss,
|
||||||
ctx_get_lr,
|
ctx_get_lr,
|
||||||
ctx_get_val_loss,
|
ctx_get_val_loss,
|
||||||
@@ -235,7 +236,7 @@ class ProgressBarCallback(TrainCallback):
|
|||||||
class MetricCallback(TrainCallback):
|
class MetricCallback(TrainCallback):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
log_dir: str,
|
ckpt_dir: str,
|
||||||
save_interval: int,
|
save_interval: int,
|
||||||
metrics: List[str] = None,
|
metrics: List[str] = None,
|
||||||
val_step: int = 0,
|
val_step: int = 0,
|
||||||
@@ -246,8 +247,7 @@ class MetricCallback(TrainCallback):
|
|||||||
self.val_step = val_step
|
self.val_step = val_step
|
||||||
self._next_val_step = 0
|
self._next_val_step = 0
|
||||||
|
|
||||||
self.log_dir = Path(log_dir) if log_dir else Path.cwd() / "logs"
|
self.ckpt_dir = Path(ckpt_dir) if ckpt_dir else Path.cwd() / "checkpoint"
|
||||||
self.log_dir.mkdir(parents=True, exist_ok=True)
|
|
||||||
|
|
||||||
self.log_cache = []
|
self.log_cache = []
|
||||||
|
|
||||||
@@ -256,6 +256,7 @@ class MetricCallback(TrainCallback):
|
|||||||
"lr": ctx_get_lr,
|
"lr": ctx_get_lr,
|
||||||
"val_loss": ctx_get_val_loss,
|
"val_loss": ctx_get_val_loss,
|
||||||
"grad_norm": ctx_get_grad_norm,
|
"grad_norm": ctx_get_grad_norm,
|
||||||
|
"grad_snr": ctx_get_grad_snr,
|
||||||
}
|
}
|
||||||
|
|
||||||
def _metrics(self, context: TrainContext, names):
|
def _metrics(self, context: TrainContext, names):
|
||||||
@@ -306,13 +307,15 @@ class MetricCallback(TrainCallback):
|
|||||||
|
|
||||||
@only_on_rank(0)
|
@only_on_rank(0)
|
||||||
def _flush(self, epoch, step):
|
def _flush(self, epoch, step):
|
||||||
log_file = self.log_dir / f"epoch_{epoch}_step_{step}_metric.jsonl"
|
log_file = self.ckpt_dir / f"epoch_{epoch}_step_{step}" / "metric.jsonl"
|
||||||
log_file.parent.mkdir(parents=True, exist_ok=True)
|
log_file.parent.mkdir(parents=True, exist_ok=True)
|
||||||
with open(log_file, "w") as f:
|
with open(log_file, "w") as f:
|
||||||
for log in self.log_cache:
|
for log in self.log_cache:
|
||||||
f.write(json.dumps(log) + "\n")
|
f.write(json.dumps(log) + "\n")
|
||||||
|
|
||||||
def on_optimizer_step(self, context):
|
def on_optimizer_step(self, context):
|
||||||
|
context.grad_snr_tracker.update(context.model)
|
||||||
|
|
||||||
if (
|
if (
|
||||||
context.val_dataloader is not None
|
context.val_dataloader is not None
|
||||||
and self.val_step > 0
|
and self.val_step > 0
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
import logging
|
||||||
import threading
|
import threading
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
@@ -11,13 +12,16 @@ from astrai.config.train_config import TrainConfig
|
|||||||
from astrai.dataset import RDSampler
|
from astrai.dataset import RDSampler
|
||||||
from astrai.inference.core.scheduler import InferenceScheduler
|
from astrai.inference.core.scheduler import InferenceScheduler
|
||||||
from astrai.model.components.lora import inject_lora
|
from astrai.model.components.lora import inject_lora
|
||||||
from astrai.parallel.executor import BaseExecutor, ExecutorFactory
|
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.parallel.setup import get_current_device, get_rank, get_world_size
|
||||||
from astrai.protocols import OptimizerProtocol, SchedulerProtocol
|
from astrai.protocols import OptimizerProtocol, SchedulerProtocol
|
||||||
from astrai.serialization import Checkpoint, load_json
|
from astrai.serialization import Checkpoint, load_json
|
||||||
from astrai.tokenize import AutoTokenizer
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
from astrai.trainer.metric_util import GradSNRTracker
|
||||||
from astrai.trainer.rollout import RolloutGenerator, RolloutRunner
|
from astrai.trainer.rollout import RolloutGenerator, RolloutRunner
|
||||||
from astrai.trainer.strategy import BaseStrategy, StrategyFactory, create_ref_model
|
from astrai.trainer.strategy import BaseStrategy, StrategyFactory
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -35,6 +39,7 @@ class TrainContext:
|
|||||||
consumed_samples: int = field(default=0)
|
consumed_samples: int = field(default=0)
|
||||||
loss: float = field(default=0.0)
|
loss: float = field(default=0.0)
|
||||||
grad_norm: Optional[float] = field(default=None)
|
grad_norm: Optional[float] = field(default=None)
|
||||||
|
grad_snr_tracker: GradSNRTracker = field(default_factory=GradSNRTracker)
|
||||||
val_dataloader: Optional[DataLoader] = field(default=None)
|
val_dataloader: Optional[DataLoader] = field(default=None)
|
||||||
val_loss: Optional[float] = field(default=None)
|
val_loss: Optional[float] = field(default=None)
|
||||||
|
|
||||||
@@ -101,18 +106,13 @@ class TrainContextBuilder:
|
|||||||
if checkpoint.config:
|
if checkpoint.config:
|
||||||
model_config = checkpoint.config
|
model_config = checkpoint.config
|
||||||
if self._resume:
|
if self._resume:
|
||||||
preloaded_epoch = checkpoint.epoch or cfg.start_epoch
|
preloaded_epoch = checkpoint.epoch
|
||||||
if checkpoint.consumed_samples > 0:
|
per_step = (
|
||||||
per_step = (
|
cfg.batch_per_device * get_world_size() * cfg.grad_accum_steps
|
||||||
cfg.batch_per_device
|
)
|
||||||
* get_world_size()
|
preloaded_consumed = (
|
||||||
* cfg.grad_accum_steps
|
checkpoint.consumed_samples // per_step
|
||||||
)
|
) * per_step
|
||||||
preloaded_consumed = (
|
|
||||||
checkpoint.consumed_samples // per_step
|
|
||||||
) * per_step
|
|
||||||
else:
|
|
||||||
preloaded_consumed = cfg.start_samples * get_world_size()
|
|
||||||
preloaded_checkpoint = checkpoint
|
preloaded_checkpoint = checkpoint
|
||||||
|
|
||||||
if not model_config and hasattr(cfg.model_fn(), "config"):
|
if not model_config and hasattr(cfg.model_fn(), "config"):
|
||||||
@@ -131,6 +131,12 @@ class TrainContextBuilder:
|
|||||||
m.load_state_dict(preloaded_state_dict, strict=False)
|
m.load_state_dict(preloaded_state_dict, strict=False)
|
||||||
return m
|
return m
|
||||||
|
|
||||||
|
def _after_wrap(m):
|
||||||
|
if cfg.compile_mode is not None:
|
||||||
|
logger.info("torch.compile enabled (mode=%s)", cfg.compile_mode)
|
||||||
|
m = torch.compile(m, mode=cfg.compile_mode)
|
||||||
|
return m
|
||||||
|
|
||||||
context = TrainContext(
|
context = TrainContext(
|
||||||
world_size=get_world_size(),
|
world_size=get_world_size(),
|
||||||
rank=get_rank(),
|
rank=get_rank(),
|
||||||
@@ -147,6 +153,7 @@ class TrainContextBuilder:
|
|||||||
cfg.optimizer_fn,
|
cfg.optimizer_fn,
|
||||||
cfg.scheduler_fn,
|
cfg.scheduler_fn,
|
||||||
before_wrap=_before_wrap,
|
before_wrap=_before_wrap,
|
||||||
|
after_wrap=_after_wrap,
|
||||||
)
|
)
|
||||||
|
|
||||||
train_dataset = cfg.dataset
|
train_dataset = cfg.dataset
|
||||||
@@ -162,6 +169,15 @@ class TrainContextBuilder:
|
|||||||
)
|
)
|
||||||
|
|
||||||
sampler_offset = context.consumed_samples // context.world_size
|
sampler_offset = context.consumed_samples // context.world_size
|
||||||
|
|
||||||
|
if self._resume and sampler_offset > 0:
|
||||||
|
offset = context.world_size - 1
|
||||||
|
num_samples_per_replica = (
|
||||||
|
len(train_dataset) + offset
|
||||||
|
) // context.world_size
|
||||||
|
if num_samples_per_replica > 0:
|
||||||
|
context.epoch = sampler_offset // num_samples_per_replica
|
||||||
|
|
||||||
sampler = RDSampler(
|
sampler = RDSampler(
|
||||||
data_source=train_dataset,
|
data_source=train_dataset,
|
||||||
start_epoch=context.epoch,
|
start_epoch=context.epoch,
|
||||||
@@ -215,17 +231,14 @@ class TrainContextBuilder:
|
|||||||
needs_old = cfg.strategy in ("grpo", "online_grpo")
|
needs_old = cfg.strategy in ("grpo", "online_grpo")
|
||||||
|
|
||||||
if needs_ref:
|
if needs_ref:
|
||||||
ref_model = create_ref_model(
|
strategy_kwargs["ref_model"] = create_ref_model(
|
||||||
cfg.model_fn, executor.unwrap_model(context.model)
|
cfg.model_fn, executor=executor, model=context.model, device=device
|
||||||
).to(device=device)
|
)
|
||||||
strategy_kwargs["ref_model"] = ref_model
|
|
||||||
|
|
||||||
old_model = None
|
|
||||||
if needs_old:
|
if needs_old:
|
||||||
old_model = create_ref_model(
|
strategy_kwargs["old_model"] = create_ref_model(
|
||||||
cfg.model_fn, executor.unwrap_model(context.model)
|
cfg.model_fn, executor=executor, model=context.model, device=device
|
||||||
).to(device=device)
|
)
|
||||||
strategy_kwargs["old_model"] = old_model
|
|
||||||
|
|
||||||
context.strategy = StrategyFactory.create(
|
context.strategy = StrategyFactory.create(
|
||||||
cfg.strategy,
|
cfg.strategy,
|
||||||
@@ -257,7 +270,6 @@ class TrainContextBuilder:
|
|||||||
tokenizer=tokenizer,
|
tokenizer=tokenizer,
|
||||||
max_batch_size=rollout_batch_size,
|
max_batch_size=rollout_batch_size,
|
||||||
max_seq_len=max_seq_len,
|
max_seq_len=max_seq_len,
|
||||||
max_prompt_len=max_seq_len or 4096,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
generator = RolloutGenerator(
|
generator = RolloutGenerator(
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ import torch.distributed as dist
|
|||||||
|
|
||||||
from astrai.config import TrainConfig
|
from astrai.config import TrainConfig
|
||||||
from astrai.parallel.setup import spawn_parallel_fn
|
from astrai.parallel.setup import spawn_parallel_fn
|
||||||
from astrai.parallel.signal_handler import (
|
from astrai.signal_handler import (
|
||||||
register_signal_handlers,
|
register_signal_handlers,
|
||||||
unregister_signal_handlers,
|
unregister_signal_handlers,
|
||||||
)
|
)
|
||||||
@@ -42,7 +42,7 @@ class Trainer:
|
|||||||
),
|
),
|
||||||
CallbackFactory.create(
|
CallbackFactory.create(
|
||||||
"metric",
|
"metric",
|
||||||
log_dir=cfg.log_dir,
|
ckpt_dir=cfg.ckpt_dir,
|
||||||
save_interval=cfg.ckpt_interval,
|
save_interval=cfg.ckpt_interval,
|
||||||
metrics=cfg.metrics,
|
metrics=cfg.metrics,
|
||||||
val_step=cfg.val_step,
|
val_step=cfg.val_step,
|
||||||
|
|||||||
+28
-1
@@ -1,6 +1,32 @@
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
|
def cuda_toolkit_version() -> tuple[int, int] | None:
|
||||||
|
"""Return ``(major, minor)`` of the nvcc on PATH, or ``None``.
|
||||||
|
|
||||||
|
Used by ``setup.py`` to detect nvcc/torch CUDA version mismatches
|
||||||
|
(e.g. nvcc 13.0 with a cu128 torch wheel) which cause cryptic ABI errors.
|
||||||
|
"""
|
||||||
|
import shutil
|
||||||
|
import subprocess
|
||||||
|
|
||||||
|
nvcc = shutil.which("nvcc")
|
||||||
|
if nvcc is None:
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
out = subprocess.check_output(
|
||||||
|
[nvcc, "--version"], stderr=subprocess.STDOUT, text=True
|
||||||
|
)
|
||||||
|
for line in out.splitlines():
|
||||||
|
if "release" in line:
|
||||||
|
ver = line.split("release")[1].split(",")[0].strip()
|
||||||
|
major, minor = ver.split(".")
|
||||||
|
return (int(major), int(minor))
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
def _arch_flags() -> list[str]:
|
def _arch_flags() -> list[str]:
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -27,7 +53,7 @@ NVCC_FLAGS = [
|
|||||||
"--use_fast_math",
|
"--use_fast_math",
|
||||||
"--ptxas-options=-O3,-v",
|
"--ptxas-options=-O3,-v",
|
||||||
"--extra-device-vectorization",
|
"--extra-device-vectorization",
|
||||||
"--threads=8",
|
"--threads=16",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
@@ -46,3 +72,4 @@ def register(name: str, sources: list[str] | None = None, **kwargs):
|
|||||||
register("attn_decode")
|
register("attn_decode")
|
||||||
register("attn_prefill")
|
register("attn_prefill")
|
||||||
register("attn_paged_decode")
|
register("attn_paged_decode")
|
||||||
|
register("rotary_emb")
|
||||||
|
|||||||
@@ -19,9 +19,11 @@ struct AttentionParams {
|
|||||||
// KV strides (K and V share the same layout — only base pointers differ)
|
// KV strides (K and V share the same layout — only base pointers differ)
|
||||||
int kv_stride_b, kv_stride_h, kv_stride_l, kv_stride_d;
|
int kv_stride_b, kv_stride_h, kv_stride_l, kv_stride_d;
|
||||||
|
|
||||||
// Mask: 2D [batch, kv_len] (mask_q_stride=0) or 3D [batch, q_len, kv_len]
|
// Mask: 2D [batch, kv_len], 3D [batch, q_len, kv_len],
|
||||||
int mask_b_stride; // = kv_len (both 2D and 3D)
|
// or 4D [batch, n_heads, q_len, kv_len] (head dim broadcasts when stride=0)
|
||||||
int mask_q_stride; // 2D: 0 (all q rows share); 3D: kv_len
|
int mask_b_stride; // batch stride
|
||||||
|
int mask_h_stride; // head stride (0 = broadcast across heads)
|
||||||
|
int mask_q_stride; // q stride (0 = all q rows share)
|
||||||
|
|
||||||
const T* __restrict__ q;
|
const T* __restrict__ q;
|
||||||
const T* __restrict__ k;
|
const T* __restrict__ k;
|
||||||
@@ -52,8 +54,9 @@ struct PagedAttentionParams {
|
|||||||
// Q strides (layout-agnostic)
|
// Q strides (layout-agnostic)
|
||||||
int q_stride_b, q_stride_h, q_stride_l, q_stride_d;
|
int q_stride_b, q_stride_h, q_stride_l, q_stride_d;
|
||||||
|
|
||||||
// Mask strides (2D or 3D)
|
// Mask strides (2D, 3D, or 4D)
|
||||||
int mask_b_stride;
|
int mask_b_stride;
|
||||||
|
int mask_h_stride;
|
||||||
int mask_q_stride;
|
int mask_q_stride;
|
||||||
|
|
||||||
const T* __restrict__ q;
|
const T* __restrict__ q;
|
||||||
|
|||||||
@@ -24,7 +24,7 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
|
|||||||
|
|
||||||
// KV: [batch, kv_head, kv_len, head_dim] — stride-based base
|
// KV: [batch, kv_head, kv_len, head_dim] — stride-based base
|
||||||
int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
||||||
int mask_base = batch * p.mask_b_stride;
|
int mask_base = batch * p.mask_b_stride + q_head * p.mask_h_stride;
|
||||||
|
|
||||||
float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f};
|
float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f};
|
||||||
|
|
||||||
@@ -70,8 +70,8 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
float new_m = fmaxf(m, partial);
|
float new_m = fmaxf(m, partial);
|
||||||
float alpha = expf(m - new_m);
|
float alpha = __expf(m - new_m);
|
||||||
float beta = expf(partial - new_m);
|
float beta = __expf(partial - new_m);
|
||||||
d = d * alpha + beta;
|
d = d * alpha + beta;
|
||||||
|
|
||||||
int v_off = kv_base + kv_idx * p.kv_stride_l
|
int v_off = kv_base + kv_idx * p.kv_stride_l
|
||||||
@@ -116,8 +116,8 @@ __global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
|
|||||||
if (mi <= -FLT_MAX) continue;
|
if (mi <= -FLT_MAX) continue;
|
||||||
float li = mlp[s * 2 + 1];
|
float li = mlp[s * 2 + 1];
|
||||||
float nm = fmaxf(m, mi);
|
float nm = fmaxf(m, mi);
|
||||||
float corr = expf(m - nm);
|
float corr = __expf(m - nm);
|
||||||
float e = expf(mi - nm);
|
float e = __expf(mi - nm);
|
||||||
acc = fmaf(acc, corr, op[s * p.head_dim + d] * e);
|
acc = fmaf(acc, corr, op[s * p.head_dim + d] * e);
|
||||||
l = fmaf(l, corr, li * e);
|
l = fmaf(l, corr, li * e);
|
||||||
m = nm;
|
m = nm;
|
||||||
|
|||||||
@@ -76,26 +76,17 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
|||||||
cp_async_commit();
|
cp_async_commit();
|
||||||
};
|
};
|
||||||
|
|
||||||
constexpr int BUF_MASK = (Traits::STAGES > 1) ? (Traits::STAGES - 1) : 0;
|
// ---- Multi-stage cp.async pipeline ----
|
||||||
|
// Prologue loads STAGES tiles; each loop iteration waits only for the
|
||||||
// Prologue
|
// oldest outstanding group (wait_group<STAGES-1>) so the STAGES-1 newer
|
||||||
if (ti_begin < ti_end) {
|
// tile loads stay in flight and overlap with the current tile's compute.
|
||||||
load_tile(ti_begin, 0);
|
constexpr int STAGES = Traits::STAGES;
|
||||||
}
|
const int ntiles = ti_end - ti_begin;
|
||||||
|
|
||||||
for (int ti = ti_begin; ti < ti_end; ti++) {
|
|
||||||
int buf = (ti - ti_begin) & BUF_MASK;
|
|
||||||
|
|
||||||
cp_async_wait_group<0>();
|
|
||||||
__syncwarp();
|
|
||||||
if constexpr (Traits::STAGES > 1) {
|
|
||||||
if (ti + 1 < ti_end)
|
|
||||||
load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK);
|
|
||||||
}
|
|
||||||
|
|
||||||
|
auto process_tile = [&](int it, int buf) {
|
||||||
const bf16* bK = sK + buf * Traits::BC * Traits::LD;
|
const bf16* bK = sK + buf * Traits::BC * Traits::LD;
|
||||||
const bf16* bV = sV + buf * Traits::BC * Traits::LD;
|
const bf16* bV = sV + buf * Traits::BC * Traits::LD;
|
||||||
int kv0 = ti * Traits::BC;
|
int kv0 = (ti_begin + it) * Traits::BC;
|
||||||
|
|
||||||
float Sacc[Traits::NC8][4];
|
float Sacc[Traits::NC8][4];
|
||||||
mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
|
mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
|
||||||
@@ -109,18 +100,35 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
|||||||
int maxc = IsCausal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
|
int maxc = IsCausal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
|
||||||
mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc,
|
mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc,
|
||||||
0, 0,
|
0, 0,
|
||||||
p.mask_b_stride, 0,
|
p.mask_b_stride, 0, 0,
|
||||||
batch,
|
batch, 0,
|
||||||
p.mask,
|
p.mask,
|
||||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||||
|
|
||||||
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
|
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
|
||||||
__syncwarp();
|
};
|
||||||
|
|
||||||
if constexpr (Traits::STAGES == 1) {
|
if (ntiles >= STAGES) {
|
||||||
if (ti + 1 < ti_end)
|
#pragma unroll
|
||||||
load_tile(ti + 1, 0);
|
for (int i = 0; i < STAGES; i++)
|
||||||
|
load_tile(ti_begin + i, i);
|
||||||
|
|
||||||
|
for (int it = 0; it < ntiles; it++) {
|
||||||
|
cp_async_wait_group<STAGES - 1>();
|
||||||
|
__syncwarp();
|
||||||
|
process_tile(it, it & (STAGES - 1));
|
||||||
|
__syncwarp();
|
||||||
|
if (it + STAGES < ntiles)
|
||||||
|
load_tile(ti_begin + it + STAGES, (it + STAGES) & (STAGES - 1));
|
||||||
}
|
}
|
||||||
|
} else {
|
||||||
|
// 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>();
|
||||||
|
__syncwarp();
|
||||||
|
for (int it = 0; it < ntiles; it++)
|
||||||
|
process_tile(it, it);
|
||||||
}
|
}
|
||||||
|
|
||||||
// ---- write UN-normalised partials for this split ----
|
// ---- write UN-normalised partials for this split ----
|
||||||
|
|||||||
@@ -15,11 +15,15 @@
|
|||||||
#endif
|
#endif
|
||||||
|
|
||||||
// Split-KV: compute number of splits to fill all SMs for small-batch decode.
|
// Split-KV: compute number of splits to fill all SMs for small-batch decode.
|
||||||
inline int compute_num_splits(int base_blocks, int tiles_total) {
|
// Caps splits so each split processes at least `min_tiles_per_split` tiles,
|
||||||
|
// avoiding excessive loop/prologue overhead when tiles are small.
|
||||||
|
inline int compute_num_splits(int base_blocks, int tiles_total,
|
||||||
|
int min_tiles_per_split = 1) {
|
||||||
int sm_count = 0;
|
int sm_count = 0;
|
||||||
cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0);
|
cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0);
|
||||||
int n = (2 * sm_count + base_blocks - 1) / base_blocks;
|
int n = (2 * sm_count + base_blocks - 1) / base_blocks;
|
||||||
return std::max(1, std::min(n, std::min(tiles_total, MAX_SPLITS)));
|
int max_by_work = tiles_total / min_tiles_per_split;
|
||||||
|
return std::max(1, std::min(n, std::min(max_by_work, MAX_SPLITS)));
|
||||||
}
|
}
|
||||||
|
|
||||||
// ======================================================================
|
// ======================================================================
|
||||||
@@ -75,15 +79,20 @@ static inline void dispatch_prefill(AttentionParams<bf16>& p) {
|
|||||||
// ======================================================================
|
// ======================================================================
|
||||||
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
// BC=16: halves smem (16KB vs 32KB) → doubles occupancy (6 vs 3 blocks/SM).
|
||||||
|
// For D=256, BC=16 also reduces register pressure (fewer Sacc/PV frags),
|
||||||
|
// enabling STAGES=2 (double-buffer) within the 32KB smem budget — eliminates
|
||||||
|
// the 176-byte spill that STAGES=1+BC=32 suffered.
|
||||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
static inline void launch_decode_mma(AttentionParams<bf16>& p, int group_size) {
|
static inline void launch_decode_mma(AttentionParams<bf16>& p, int group_size) {
|
||||||
int G = p.q_head / p.kv_head;
|
int G = p.q_head / p.kv_head;
|
||||||
constexpr int MAX_G = 16;
|
constexpr int MAX_G = 16;
|
||||||
int num_passes = (G + MAX_G - 1) / MAX_G;
|
int num_passes = (G + MAX_G - 1) / MAX_G;
|
||||||
int tiles_total = (p.kv_len + 32 - 1) / 32;
|
constexpr int BC = 16;
|
||||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
|
int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||||
constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1;
|
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total, 2);
|
||||||
using Traits = KernelTraits<HEAD_DIM, 32, 1, STAGES>;
|
constexpr int STAGES = 2;
|
||||||
|
using Traits = KernelTraits<HEAD_DIM, BC, 1, STAGES>;
|
||||||
dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
|
dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
|
||||||
attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask><<<grid, 32>>>(p);
|
attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask><<<grid, 32>>>(p);
|
||||||
}
|
}
|
||||||
@@ -136,23 +145,14 @@ template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
|||||||
static inline void launch_paged_decode_mma(PagedAttentionParams<bf16>& p, int group_size) {
|
static inline void launch_paged_decode_mma(PagedAttentionParams<bf16>& p, int group_size) {
|
||||||
int G = p.q_head / p.kv_head;
|
int G = p.q_head / p.kv_head;
|
||||||
constexpr int MAX_G = 16;
|
constexpr int MAX_G = 16;
|
||||||
bool page_ok = (p.page_size >= 32);
|
constexpr int BC = 16;
|
||||||
if (G >= 1 && page_ok) {
|
int num_passes = (G + MAX_G - 1) / MAX_G;
|
||||||
int num_passes = (G + MAX_G - 1) / MAX_G;
|
int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||||
int tiles_total = (p.kv_len + 32 - 1) / 32;
|
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total, 2);
|
||||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
|
constexpr int STAGES = 2;
|
||||||
constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1;
|
using Traits = KernelTraits<HEAD_DIM, BC, 1, STAGES>;
|
||||||
using Traits = KernelTraits<HEAD_DIM, 32, 1, STAGES>;
|
dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
|
||||||
dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
|
paged_attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask> <<<grid, 32>>>(p);
|
||||||
paged_attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask> <<<grid, 32>>>(p);
|
|
||||||
} else {
|
|
||||||
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
|
|
||||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
|
||||||
size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16);
|
|
||||||
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
|
|
||||||
dim3 block(32, group_size);
|
|
||||||
paged_attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
#endif
|
#endif
|
||||||
|
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
#pragma once
|
#pragma once
|
||||||
|
#include <float.h>
|
||||||
#include <torch/extension.h>
|
#include <torch/extension.h>
|
||||||
#include <c10/cuda/CUDAGuard.h>
|
#include <c10/cuda/CUDAGuard.h>
|
||||||
#include "attn_common.h"
|
#include "attn_common.h"
|
||||||
@@ -20,11 +21,15 @@ using bf16 = __nv_bfloat16;
|
|||||||
" (supported: 32, 64, 128, 256)"); \
|
" (supported: 32, 64, 128, 256)"); \
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 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
|
||||||
|
// per-call memset (torch::zeros / torch::full) was pure overhead.
|
||||||
template<typename P>
|
template<typename P>
|
||||||
inline void alloc_split_partials(P& p) {
|
inline void alloc_split_partials(P& p) {
|
||||||
auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA);
|
auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA);
|
||||||
auto o_part = torch::empty({p.batch, p.q_head, MAX_SPLITS, p.head_dim}, fopt);
|
auto o_part = torch::empty(at::IntArrayRef{p.batch, p.q_head, MAX_SPLITS, p.head_dim}, fopt);
|
||||||
auto ml_part = torch::empty({p.batch, p.q_head, MAX_SPLITS, 2}, fopt);
|
auto ml_part = torch::empty(at::IntArrayRef{p.batch, p.q_head, MAX_SPLITS, 2}, fopt);
|
||||||
p.o_part = (float*)o_part.data_ptr();
|
p.o_part = (float*)o_part.data_ptr();
|
||||||
p.ml_part = (float*)ml_part.data_ptr();
|
p.ml_part = (float*)ml_part.data_ptr();
|
||||||
}
|
}
|
||||||
@@ -44,6 +49,9 @@ inline void extract_q_dims_and_strides(torch::Tensor& q, int64_t layout, P& p) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// ---- Shared mask packing ----
|
// ---- Shared mask packing ----
|
||||||
|
// Accepts 2D [batch, kv_len], 3D [batch, q_len, kv_len],
|
||||||
|
// or 4D [batch, n_heads, q_len, kv_len].
|
||||||
|
// Head/q dimensions with size 1 broadcast (stride set to 0).
|
||||||
template <typename P>
|
template <typename P>
|
||||||
inline void pack_mask(const c10::optional<torch::Tensor>& mask, P& p) {
|
inline void pack_mask(const c10::optional<torch::Tensor>& mask, P& p) {
|
||||||
if (p.use_mask) {
|
if (p.use_mask) {
|
||||||
@@ -54,18 +62,26 @@ inline void pack_mask(const c10::optional<torch::Tensor>& mask, P& p) {
|
|||||||
TORCH_CHECK(m.size(m.dim() - 1) == p.kv_len, "mask kv_len mismatch");
|
TORCH_CHECK(m.size(m.dim() - 1) == p.kv_len, "mask kv_len mismatch");
|
||||||
if (m.dim() == 2) {
|
if (m.dim() == 2) {
|
||||||
p.mask_b_stride = (int)m.stride(0);
|
p.mask_b_stride = (int)m.stride(0);
|
||||||
|
p.mask_h_stride = 0;
|
||||||
p.mask_q_stride = 0;
|
p.mask_q_stride = 0;
|
||||||
} else if (m.dim() == 3) {
|
} else if (m.dim() == 3) {
|
||||||
TORCH_CHECK(m.size(1) == p.q_len, "mask q_len mismatch");
|
TORCH_CHECK(m.size(1) == 1 || m.size(1) == p.q_len, "mask q_len mismatch");
|
||||||
p.mask_b_stride = (int)m.stride(0);
|
p.mask_b_stride = (int)m.stride(0);
|
||||||
p.mask_q_stride = (int)m.stride(1);
|
p.mask_h_stride = 0;
|
||||||
|
p.mask_q_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1);
|
||||||
|
} else if (m.dim() == 4) {
|
||||||
|
TORCH_CHECK(m.size(2) == 1 || m.size(2) == p.q_len, "mask q_len mismatch");
|
||||||
|
p.mask_b_stride = (int)m.stride(0);
|
||||||
|
p.mask_h_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1);
|
||||||
|
p.mask_q_stride = (m.size(2) == 1) ? 0 : (int)m.stride(2);
|
||||||
} else {
|
} else {
|
||||||
TORCH_CHECK(false, "mask must be 2D [batch, kv_len] or 3D [batch, q_len, kv_len]");
|
TORCH_CHECK(false, "mask must be 2D, 3D, or 4D");
|
||||||
}
|
}
|
||||||
p.mask = m.data_ptr<bool>();
|
p.mask = m.data_ptr<bool>();
|
||||||
} else {
|
} else {
|
||||||
p.mask = nullptr;
|
p.mask = nullptr;
|
||||||
p.mask_b_stride = 0;
|
p.mask_b_stride = 0;
|
||||||
|
p.mask_h_stride = 0;
|
||||||
p.mask_q_stride = 0;
|
p.mask_q_stride = 0;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -3,6 +3,12 @@
|
|||||||
#include <cuda_fp16.h>
|
#include <cuda_fp16.h>
|
||||||
#include <cuda_runtime.h>
|
#include <cuda_runtime.h>
|
||||||
|
|
||||||
|
// 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
|
||||||
|
|
||||||
// ============================================================================
|
// ============================================================================
|
||||||
// KernelTraits — FlashAttention-v2 style compile-time configuration bundle.
|
// KernelTraits — FlashAttention-v2 style compile-time configuration bundle.
|
||||||
//
|
//
|
||||||
@@ -192,8 +198,8 @@ __device__ inline void mma_softmax_tile(
|
|||||||
int kv0,
|
int kv0,
|
||||||
int maxc0, int maxc1,
|
int maxc0, int maxc1,
|
||||||
int qrow0, int qrow1,
|
int qrow0, int qrow1,
|
||||||
int mask_b_stride, int mask_q_stride,
|
int mask_b_stride, int mask_h_stride, int mask_q_stride,
|
||||||
int mask_batch,
|
int mask_batch, int mask_head,
|
||||||
const bool* __restrict__ mask,
|
const bool* __restrict__ mask,
|
||||||
float Sacc[Traits::NC8][4],
|
float Sacc[Traits::NC8][4],
|
||||||
float Oacc[Traits::DN8][4],
|
float Oacc[Traits::DN8][4],
|
||||||
@@ -204,8 +210,8 @@ __device__ inline void mma_softmax_tile(
|
|||||||
int tid4 = lane & 3;
|
int tid4 = lane & 3;
|
||||||
|
|
||||||
float rmax0 = -FLT_MAX, rmax1 = -FLT_MAX;
|
float rmax0 = -FLT_MAX, rmax1 = -FLT_MAX;
|
||||||
int mask_base0 = mask_batch * mask_b_stride + qrow0 * mask_q_stride;
|
int mask_base0 = mask_batch * mask_b_stride + mask_head * mask_h_stride + qrow0 * mask_q_stride;
|
||||||
int mask_base1 = mask_batch * mask_b_stride + qrow1 * mask_q_stride;
|
int mask_base1 = mask_batch * mask_b_stride + mask_head * mask_h_stride + qrow1 * mask_q_stride;
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int n8 = 0; n8 < Traits::NC8; n8++) {
|
for (int n8 = 0; n8 < Traits::NC8; n8++) {
|
||||||
int cc = kv0 + n8 * 8 + 2 * tid4;
|
int cc = kv0 + n8 * 8 + 2 * tid4;
|
||||||
|
|||||||
@@ -31,7 +31,7 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
|
|||||||
int ch_begin = split * chunks_per_split;
|
int ch_begin = split * chunks_per_split;
|
||||||
int ch_end = min(chunks_total, ch_begin + chunks_per_split);
|
int ch_end = min(chunks_total, ch_begin + chunks_per_split);
|
||||||
|
|
||||||
const int mask_base = batch * p.mask_b_stride;
|
const int mask_base = batch * p.mask_b_stride + q_head * p.mask_h_stride;
|
||||||
|
|
||||||
for (int ci = ch_begin; ci < ch_end; ci++) {
|
for (int ci = ch_begin; ci < ch_end; ci++) {
|
||||||
int chunk_start = ci * PDC_CHUNK;
|
int chunk_start = ci * PDC_CHUNK;
|
||||||
@@ -67,25 +67,32 @@ __global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p)
|
|||||||
partial = warp_reduce_sum(partial) * p.scale;
|
partial = warp_reduce_sum(partial) * p.scale;
|
||||||
|
|
||||||
int kv_idx = chunk_start + s;
|
int kv_idx = chunk_start + s;
|
||||||
|
bool masked = false;
|
||||||
if constexpr (HasMask) {
|
if constexpr (HasMask) {
|
||||||
if (!p.mask[mask_base + kv_idx])
|
if (!p.mask[mask_base + kv_idx])
|
||||||
partial = -FLT_MAX;
|
masked = true;
|
||||||
}
|
}
|
||||||
if constexpr (IsCausal) {
|
if constexpr (IsCausal) {
|
||||||
if (kv_idx > p.causal_offset)
|
if (kv_idx > p.causal_offset)
|
||||||
partial = -FLT_MAX;
|
masked = true;
|
||||||
}
|
}
|
||||||
|
if (masked)
|
||||||
|
partial = -FLT_MAX;
|
||||||
|
|
||||||
float new_m = fmaxf(m, partial);
|
float new_m = fmaxf(m, partial);
|
||||||
float alpha = expf(m - new_m);
|
float alpha = __expf(m - new_m);
|
||||||
float beta = expf(partial - new_m);
|
float beta = __expf(partial - new_m);
|
||||||
d = d * alpha + beta;
|
d = d * alpha + beta;
|
||||||
|
|
||||||
int pos = chunk_start + s;
|
int pos = chunk_start + s;
|
||||||
int logical_page = pos / p.page_size;
|
int logical_page = pos / p.page_size;
|
||||||
int page_offset = pos % p.page_size;
|
int page_offset = pos % p.page_size;
|
||||||
int phys_page = p.page_table[batch * p.max_pages + logical_page];
|
int phys_page = p.page_table[batch * p.max_pages + logical_page];
|
||||||
if (phys_page >= 0) {
|
if (masked) {
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
|
acc_reg[i] = fmaf(acc_reg[i], alpha, 0.0f);
|
||||||
|
} else if (phys_page >= 0) {
|
||||||
int64_t v_base = (int64_t)phys_page * p.page_size * p.kv_head * p.head_dim
|
int64_t v_base = (int64_t)phys_page * p.page_size * p.kv_head * p.head_dim
|
||||||
+ (int64_t)page_offset * p.kv_head * p.head_dim
|
+ (int64_t)page_offset * p.kv_head * p.head_dim
|
||||||
+ (int64_t)kv_head * p.head_dim;
|
+ (int64_t)kv_head * p.head_dim;
|
||||||
@@ -133,8 +140,8 @@ __global__ void paged_attn_decode_combine_kernel(PagedAttentionParams<bf16> p) {
|
|||||||
if (mi <= -FLT_MAX) continue;
|
if (mi <= -FLT_MAX) continue;
|
||||||
float li = mlp[s * 2 + 1];
|
float li = mlp[s * 2 + 1];
|
||||||
float nm = fmaxf(m, mi);
|
float nm = fmaxf(m, mi);
|
||||||
float corr = expf(m - nm);
|
float corr = __expf(m - nm);
|
||||||
float e = expf(mi - nm);
|
float e = __expf(mi - nm);
|
||||||
acc = fmaf(acc, corr, op[s * p.head_dim + d] * e);
|
acc = fmaf(acc, corr, op[s * p.head_dim + d] * e);
|
||||||
l = fmaf(l, corr, li * e);
|
l = fmaf(l, corr, li * e);
|
||||||
m = nm;
|
m = nm;
|
||||||
|
|||||||
@@ -31,6 +31,13 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
|
|||||||
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
|
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
|
||||||
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
|
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
|
||||||
|
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = lane; i < Traits::STAGES * Traits::BC * Traits::LD; i += 32) {
|
||||||
|
sK[i] = __float2bfloat16(0.0f);
|
||||||
|
sV[i] = __float2bfloat16(0.0f);
|
||||||
|
}
|
||||||
|
__syncwarp();
|
||||||
|
|
||||||
const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h;
|
const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h;
|
||||||
const int qra = gid;
|
const int qra = gid;
|
||||||
const int qrb = gid + 8;
|
const int qrb = gid + 8;
|
||||||
@@ -55,19 +62,24 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
|
|||||||
const int64_t head_off = (int64_t)kv_head * Traits::HEAD_DIM;
|
const int64_t head_off = (int64_t)kv_head * Traits::HEAD_DIM;
|
||||||
|
|
||||||
// ---- Load tile lambda: paged addressing ----
|
// ---- Load tile lambda: paged addressing ----
|
||||||
|
// Unified per-element page-table lookup. When page_size >= BC, all
|
||||||
|
// elements in a tile share the same page, so the lookup is redundant
|
||||||
|
// but harmless (L1-cached). This avoids a branch on page_size.
|
||||||
auto load_tile = [&](int ti, int buf) {
|
auto load_tile = [&](int ti, int buf) {
|
||||||
int kv0 = ti * Traits::BC;
|
int kv0 = ti * Traits::BC;
|
||||||
bf16* dK = sK + buf * Traits::BC * Traits::LD;
|
bf16* dK = sK + buf * Traits::BC * Traits::LD;
|
||||||
bf16* dV = sV + buf * Traits::BC * Traits::LD;
|
bf16* dV = sV + buf * Traits::BC * Traits::LD;
|
||||||
int logical_page = kv0 / p.page_size;
|
|
||||||
int phys_page = p.page_table[batch * p.max_pages + logical_page];
|
|
||||||
bool page_valid = (phys_page >= 0);
|
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int i = lane * Traits::VEC; i < Traits::TOTAL;
|
for (int i = lane * Traits::VEC; i < Traits::TOTAL;
|
||||||
i += Traits::NUM_THREADS * Traits::VEC) {
|
i += Traits::NUM_THREADS * Traits::VEC) {
|
||||||
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
|
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
|
||||||
int kc = kv0 + r;
|
int kc = kv0 + r;
|
||||||
bool valid = (kc < p.kv_len) && page_valid;
|
bool valid = (kc < p.kv_len);
|
||||||
|
if constexpr (HasMask) {
|
||||||
|
valid = valid && p.mask[batch * p.mask_b_stride + kc];
|
||||||
|
}
|
||||||
|
int phys_page = valid ? p.page_table[batch * p.max_pages + kc] : 0;
|
||||||
|
valid = valid && (phys_page >= 0);
|
||||||
int page_off = kc % p.page_size;
|
int page_off = kc % p.page_size;
|
||||||
int64_t gmem_base = (int64_t)phys_page * page_stride
|
int64_t gmem_base = (int64_t)phys_page * page_stride
|
||||||
+ (int64_t)page_off * pos_stride
|
+ (int64_t)page_off * pos_stride
|
||||||
@@ -79,25 +91,17 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
|
|||||||
cp_async_commit();
|
cp_async_commit();
|
||||||
};
|
};
|
||||||
|
|
||||||
constexpr int BUF_MASK = (Traits::STAGES > 1) ? (Traits::STAGES - 1) : 0;
|
// ---- Multi-stage cp.async pipeline ----
|
||||||
|
// Prologue loads STAGES tiles; each loop iteration waits only for the
|
||||||
if (ti_begin < ti_end) {
|
// oldest outstanding group (wait_group<STAGES-1>) so the STAGES-1 newer
|
||||||
load_tile(ti_begin, 0);
|
// tile loads stay in flight and overlap with the current tile's compute.
|
||||||
}
|
constexpr int STAGES = Traits::STAGES;
|
||||||
|
const int ntiles = ti_end - ti_begin;
|
||||||
for (int ti = ti_begin; ti < ti_end; ti++) {
|
|
||||||
int buf = (ti - ti_begin) & BUF_MASK;
|
|
||||||
|
|
||||||
cp_async_wait_group<0>();
|
|
||||||
__syncwarp();
|
|
||||||
if constexpr (Traits::STAGES > 1) {
|
|
||||||
if (ti + 1 < ti_end)
|
|
||||||
load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK);
|
|
||||||
}
|
|
||||||
|
|
||||||
|
auto process_tile = [&](int it, int buf) {
|
||||||
const bf16* bK = sK + buf * Traits::BC * Traits::LD;
|
const bf16* bK = sK + buf * Traits::BC * Traits::LD;
|
||||||
const bf16* bV = sV + buf * Traits::BC * Traits::LD;
|
const bf16* bV = sV + buf * Traits::BC * Traits::LD;
|
||||||
int kv0 = ti * Traits::BC;
|
int kv0 = (ti_begin + it) * Traits::BC;
|
||||||
|
|
||||||
float Sacc[Traits::NC8][4];
|
float Sacc[Traits::NC8][4];
|
||||||
mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
|
mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
|
||||||
@@ -110,18 +114,35 @@ __global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16>
|
|||||||
int maxc = IsCausal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
|
int maxc = IsCausal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
|
||||||
mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc,
|
mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc,
|
||||||
0, 0,
|
0, 0,
|
||||||
p.mask_b_stride, 0,
|
p.mask_b_stride, 0, 0,
|
||||||
batch,
|
batch, 0,
|
||||||
p.mask,
|
p.mask,
|
||||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||||
|
|
||||||
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
|
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
|
||||||
__syncwarp();
|
};
|
||||||
|
|
||||||
if constexpr (Traits::STAGES == 1) {
|
if (ntiles >= STAGES) {
|
||||||
if (ti + 1 < ti_end)
|
#pragma unroll
|
||||||
load_tile(ti + 1, 0);
|
for (int i = 0; i < STAGES; i++)
|
||||||
|
load_tile(ti_begin + i, i);
|
||||||
|
|
||||||
|
for (int it = 0; it < ntiles; it++) {
|
||||||
|
cp_async_wait_group<STAGES - 1>();
|
||||||
|
__syncwarp();
|
||||||
|
process_tile(it, it & (STAGES - 1));
|
||||||
|
__syncwarp();
|
||||||
|
if (it + STAGES < ntiles)
|
||||||
|
load_tile(ti_begin + it + STAGES, (it + STAGES) & (STAGES - 1));
|
||||||
}
|
}
|
||||||
|
} else {
|
||||||
|
// 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>();
|
||||||
|
__syncwarp();
|
||||||
|
for (int it = 0; it < ntiles; it++)
|
||||||
|
process_tile(it, it);
|
||||||
}
|
}
|
||||||
|
|
||||||
auto split_slot = [&](int h) -> size_t {
|
auto split_slot = [&](int h) -> size_t {
|
||||||
|
|||||||
@@ -64,7 +64,7 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
|
|||||||
|
|
||||||
// KV: stride-based base
|
// KV: stride-based base
|
||||||
int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
||||||
int mask_batch_base = batch * p.mask_b_stride;
|
int mask_batch_base = batch * p.mask_b_stride + q_head * p.mask_h_stride;
|
||||||
int tiles = (p.kv_len + P_BC - 1) / P_BC;
|
int tiles = (p.kv_len + P_BC - 1) / P_BC;
|
||||||
int tt = G * ROWS;
|
int tt = G * ROWS;
|
||||||
int lid = row * G + gpos;
|
int lid = row * G + gpos;
|
||||||
|
|||||||
@@ -114,8 +114,8 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
|
|||||||
: p.kv_len;
|
: p.kv_len;
|
||||||
mma_softmax_tile<Traits, HasMask>(kv0, maxc0, maxc1,
|
mma_softmax_tile<Traits, HasMask>(kv0, maxc0, maxc1,
|
||||||
qr0, qr1,
|
qr0, qr1,
|
||||||
p.mask_b_stride, p.mask_q_stride,
|
p.mask_b_stride, p.mask_h_stride, p.mask_q_stride,
|
||||||
batch,
|
batch, q_head,
|
||||||
p.mask,
|
p.mask,
|
||||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,87 @@
|
|||||||
|
#include <torch/extension.h>
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
|
||||||
|
__global__ void rotary_emb_kernel(
|
||||||
|
const __nv_bfloat16* __restrict__ x,
|
||||||
|
const float* __restrict__ freqs_cis,
|
||||||
|
__nv_bfloat16* __restrict__ out,
|
||||||
|
int batch,
|
||||||
|
int seq_len,
|
||||||
|
int n_heads,
|
||||||
|
int head_dim
|
||||||
|
) {
|
||||||
|
const int half_dim = head_dim >> 1;
|
||||||
|
const int total = batch * seq_len * n_heads * half_dim;
|
||||||
|
|
||||||
|
for (int idx = blockIdx.x * blockDim.x + threadIdx.x;
|
||||||
|
idx < total;
|
||||||
|
idx += gridDim.x * blockDim.x) {
|
||||||
|
|
||||||
|
int pair = idx % half_dim;
|
||||||
|
int tmp = idx / half_dim;
|
||||||
|
int head = tmp % n_heads;
|
||||||
|
tmp /= n_heads;
|
||||||
|
int seq = tmp % seq_len;
|
||||||
|
int b = tmp / seq_len;
|
||||||
|
|
||||||
|
int x_offset = ((b * seq_len + seq) * n_heads + head) * head_dim + (pair << 1);
|
||||||
|
int cs_offset = ((b * seq_len + seq) * half_dim + pair) * 2;
|
||||||
|
|
||||||
|
__nv_bfloat162 x_pair = *reinterpret_cast<const __nv_bfloat162*>(x + x_offset);
|
||||||
|
float x_even = __bfloat162float(__low2bfloat16(x_pair));
|
||||||
|
float x_odd = __bfloat162float(__high2bfloat16(x_pair));
|
||||||
|
|
||||||
|
float c = freqs_cis[cs_offset];
|
||||||
|
float s = freqs_cis[cs_offset + 1];
|
||||||
|
|
||||||
|
float out_even = x_even * c - x_odd * s;
|
||||||
|
float out_odd = x_even * s + x_odd * c;
|
||||||
|
|
||||||
|
__nv_bfloat162 out_pair = __floats2bfloat162_rn(out_even, out_odd);
|
||||||
|
*reinterpret_cast<__nv_bfloat162*>(out + x_offset) = out_pair;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
torch::Tensor rotary_emb(
|
||||||
|
torch::Tensor x,
|
||||||
|
torch::Tensor freqs_cis
|
||||||
|
) {
|
||||||
|
TORCH_CHECK(x.is_cuda(), "x must be on CUDA");
|
||||||
|
TORCH_CHECK(freqs_cis.is_cuda(), "freqs_cis must be on CUDA");
|
||||||
|
TORCH_CHECK(x.scalar_type() == torch::kBFloat16, "x must be bf16");
|
||||||
|
TORCH_CHECK(x.dim() == 4, "x must be 4D [batch, seq_len, n_heads, head_dim]");
|
||||||
|
TORCH_CHECK(x.is_contiguous(), "x must be contiguous");
|
||||||
|
TORCH_CHECK(freqs_cis.dim() == 4, "freqs_cis must be 4D [batch, seq_len, dim/2, 2]");
|
||||||
|
TORCH_CHECK(freqs_cis.is_contiguous(), "freqs_cis must be contiguous");
|
||||||
|
|
||||||
|
int batch = x.size(0);
|
||||||
|
int seq_len = x.size(1);
|
||||||
|
int n_heads = x.size(2);
|
||||||
|
int head_dim = x.size(3);
|
||||||
|
|
||||||
|
TORCH_CHECK(head_dim % 2 == 0, "head_dim must be even");
|
||||||
|
|
||||||
|
auto out = torch::empty_like(x);
|
||||||
|
|
||||||
|
int half_dim = head_dim / 2;
|
||||||
|
int total = batch * seq_len * n_heads * half_dim;
|
||||||
|
int block = 256;
|
||||||
|
int grid = std::min((total + block - 1) / block, 1024);
|
||||||
|
|
||||||
|
rotary_emb_kernel<<<grid, block>>>(
|
||||||
|
reinterpret_cast<const __nv_bfloat16*>(x.data_ptr()),
|
||||||
|
freqs_cis.data_ptr<float>(),
|
||||||
|
reinterpret_cast<__nv_bfloat16*>(out.data_ptr()),
|
||||||
|
batch, seq_len, n_heads, head_dim
|
||||||
|
);
|
||||||
|
|
||||||
|
return out;
|
||||||
|
}
|
||||||
|
|
||||||
|
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||||
|
m.def("rotary_emb", &rotary_emb,
|
||||||
|
py::arg("x"),
|
||||||
|
py::arg("freqs_cis"),
|
||||||
|
"Fused rotary embedding (bf16 x, f32 freqs_cis [b,s,d/2,2], bf16 out)"
|
||||||
|
);
|
||||||
|
}
|
||||||
@@ -103,6 +103,7 @@ inline void set_default_strides(P& p) {
|
|||||||
p.kv_stride_l = p.head_dim;
|
p.kv_stride_l = p.head_dim;
|
||||||
p.kv_stride_d = 1;
|
p.kv_stride_d = 1;
|
||||||
p.mask_b_stride = p.kv_len;
|
p.mask_b_stride = p.kv_len;
|
||||||
|
p.mask_h_stride = 0;
|
||||||
p.mask_q_stride = 0;
|
p.mask_q_stride = 0;
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -114,6 +115,7 @@ inline void set_default_paged_strides(P& p) {
|
|||||||
p.q_stride_l = p.head_dim;
|
p.q_stride_l = p.head_dim;
|
||||||
p.q_stride_d = 1;
|
p.q_stride_d = 1;
|
||||||
p.mask_b_stride = p.kv_len;
|
p.mask_b_stride = p.kv_len;
|
||||||
|
p.mask_h_stride = 0;
|
||||||
p.mask_q_stride = 0;
|
p.mask_q_stride = 0;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -3,6 +3,8 @@ services:
|
|||||||
build:
|
build:
|
||||||
context: .
|
context: .
|
||||||
dockerfile: Dockerfile
|
dockerfile: Dockerfile
|
||||||
|
args:
|
||||||
|
CUDA_TAG: ${CUDA_TAG:-cu128}
|
||||||
user: "${UID:-1000}:${GID:-1000}"
|
user: "${UID:-1000}:${GID:-1000}"
|
||||||
ports:
|
ports:
|
||||||
- "8000:8000"
|
- "8000:8000"
|
||||||
@@ -29,6 +31,8 @@ services:
|
|||||||
build:
|
build:
|
||||||
context: .
|
context: .
|
||||||
dockerfile: Dockerfile
|
dockerfile: Dockerfile
|
||||||
|
args:
|
||||||
|
CUDA_TAG: ${CUDA_TAG:-cu128}
|
||||||
user: "${UID:-1000}:${GID:-1000}"
|
user: "${UID:-1000}:${GID:-1000}"
|
||||||
ports:
|
ports:
|
||||||
- "8000:8000"
|
- "8000:8000"
|
||||||
|
|||||||
@@ -1,9 +1,9 @@
|
|||||||
<div align="center">
|
<div align="center">
|
||||||
|
|
||||||
<img src="../images/logo.png" width="auto" alt="Logo">
|
<img src="./images/logo.png" width="auto" alt="Logo">
|
||||||
|
|
||||||
<div>
|
<div>
|
||||||
<a href="../../README.md">English</a> •
|
<a href="../README.md">English</a> •
|
||||||
<a href="#chinese">中文</a>
|
<a href="#chinese">中文</a>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
@@ -23,7 +23,7 @@
|
|||||||
<br>
|
<br>
|
||||||
|
|
||||||
<div align="center">
|
<div align="center">
|
||||||
<a href="../../README.md">English</a> •
|
<a href="../README.md">English</a> •
|
||||||
<a href="#chinese">中文</a> •
|
<a href="#chinese">中文</a> •
|
||||||
<a href="https://github.com/ViperEkura/AstrAI/issues">问题追踪</a> •
|
<a href="https://github.com/ViperEkura/AstrAI/issues">问题追踪</a> •
|
||||||
<a href="https://github.com/ViperEkura/AstrAI/discussions">讨论区</a> •
|
<a href="https://github.com/ViperEkura/AstrAI/discussions">讨论区</a> •
|
||||||
@@ -219,18 +219,23 @@ curl -X POST http://localhost:8000/v1/messages \
|
|||||||
curl http://localhost:8000/health
|
curl http://localhost:8000/health
|
||||||
```
|
```
|
||||||
|
|
||||||
SSE 流式格式、错误码和统计端点详见[推理文档](./inference.md)。
|
SSE 流式格式、错误码和统计端点详见[推理文档](guides/inference.md)。
|
||||||
|
|
||||||
### 文档
|
### 文档
|
||||||
|
|
||||||
| 文档 | 说明 |
|
| 文档 | 说明 |
|
||||||
|------|------|
|
|------|------|
|
||||||
| [CLI 参考](./params.md) | 所有 CLI 工具参数(训练、服务、生成、预处理) |
|
| [快速上手](./get-started.md) | 安装与快速入门 |
|
||||||
| [架构文档](./architecture.md) | 系统架构、类图与设计模式 |
|
| [CLI 参考](./guides/params.md) | 所有 CLI 工具参数(训练、服务、生成、预处理) |
|
||||||
| [训练文档](./training.md) | 训练循环、策略与公式 |
|
| [数据预处理](./guides/preprocessing.md) | 声明式 JSON 驱动数据预处理 |
|
||||||
| [推理文档](./inference.md) | KVCache、连续批处理、采样与 HTTP API |
|
| [训练文档](./guides/training.md) | 训练循环、策略与公式 |
|
||||||
| [数据流程](./dataflow.md) | 数据管道、存储后端与数据集架构 |
|
| [推理文档](./guides/inference.md) | KVCache、连续批处理、采样与 HTTP API |
|
||||||
| [数据预处理](./preprocessing.md) | 声明式 JSON 驱动数据预处理 |
|
| [评估文档](./guides/evaluation.md) | HumanEval、MMLU、PPL、ROUGE、IFD、IFEval |
|
||||||
|
| [分布式训练](./guides/distributed.md) | 多卡 DDP / FSDP 训练 |
|
||||||
|
| [架构文档](./developer/architecture.md) | 系统架构、类图与设计模式 |
|
||||||
|
| [数据流程](./developer/dataflow.md) | 数据管道、存储后端与数据集架构 |
|
||||||
|
| [内部实现](./developer/internals.md) | 训练原理:损失公式、回调生命周期、KV Cache |
|
||||||
|
| [CUDA 内核](./developer/cuda_kernels.md) | 自定义 CUDA 注意力内核与基准测试 |
|
||||||
|
|
||||||
### 贡献
|
### 贡献
|
||||||
|
|
||||||
@@ -277,7 +277,7 @@ classDiagram
|
|||||||
+ModuleList layers
|
+ModuleList layers
|
||||||
+RMSNorm norm
|
+RMSNorm norm
|
||||||
+Linear lm_head
|
+Linear lm_head
|
||||||
+forward(input_ids, input_mask, paged_cache, position_ids) Dict[str, Tensor]
|
+forward(input_ids, input_mask, kv_cache, position_ids) Dict[str, Tensor]
|
||||||
+load_state_dict(state_dict, strict, assign)
|
+load_state_dict(state_dict, strict, assign)
|
||||||
+state_dict()
|
+state_dict()
|
||||||
}
|
}
|
||||||
@@ -299,7 +299,7 @@ classDiagram
|
|||||||
+RMSNorm input_norm
|
+RMSNorm input_norm
|
||||||
+nn.Module mlp # MLP or DeepSeekMoE via FFNFactory
|
+nn.Module mlp # MLP or DeepSeekMoE via FFNFactory
|
||||||
+RMSNorm post_attention_norm
|
+RMSNorm post_attention_norm
|
||||||
+forward(x, rotary_emb, attention_mask, paged_cache) Tensor
|
+forward(x, rotary_emb, attention_mask, kv_cache) Tensor
|
||||||
}
|
}
|
||||||
|
|
||||||
class GQA {
|
class GQA {
|
||||||
@@ -314,7 +314,7 @@ classDiagram
|
|||||||
+Linear q_proj, k_proj, v_proj, o_proj
|
+Linear q_proj, k_proj, v_proj, o_proj
|
||||||
+Linear gate # only if use_gated_attention
|
+Linear gate # only if use_gated_attention
|
||||||
+RMSNorm q_norm, k_norm # only if use_qk_norm
|
+RMSNorm q_norm, k_norm # only if use_qk_norm
|
||||||
+forward(x, rotary_emb, attn_mask, paged_cache) Tensor
|
+forward(x, rotary_emb, attn_mask, kv_cache) Tensor
|
||||||
}
|
}
|
||||||
|
|
||||||
class MLA {
|
class MLA {
|
||||||
@@ -334,7 +334,7 @@ classDiagram
|
|||||||
+Linear gate # only if use_gated_attention
|
+Linear gate # only if use_gated_attention
|
||||||
+RMSNorm kv_norm
|
+RMSNorm kv_norm
|
||||||
+RMSNorm q_norm, k_norm # only if use_qk_norm
|
+RMSNorm q_norm, k_norm # only if use_qk_norm
|
||||||
+forward(x, rotary_emb, attn_mask, paged_cache) Tensor
|
+forward(x, rotary_emb, attn_mask, kv_cache) Tensor
|
||||||
}
|
}
|
||||||
|
|
||||||
class MLP {
|
class MLP {
|
||||||
@@ -380,7 +380,9 @@ classDiagram
|
|||||||
+int max_len
|
+int max_len
|
||||||
+float base
|
+float base
|
||||||
+Optional[Dict] rope_scaling
|
+Optional[Dict] rope_scaling
|
||||||
+forward(x, position_ids=None) Tensor
|
+Tensor cos_table
|
||||||
|
+Tensor sin_table
|
||||||
|
+forward(x, position_ids=None) Tuple[Tensor, Tensor]
|
||||||
}
|
}
|
||||||
|
|
||||||
class Embedding {
|
class Embedding {
|
||||||
@@ -824,75 +826,49 @@ classDiagram
|
|||||||
+record(page_idx, token_ids, logical_page_idx)
|
+record(page_idx, token_ids, logical_page_idx)
|
||||||
}
|
}
|
||||||
|
|
||||||
class Storage {
|
class KVStorage {
|
||||||
+int page_size
|
+int size
|
||||||
+Tensor k_cache
|
+Tensor k_buffer
|
||||||
+Tensor v_cache
|
+Tensor v_buffer
|
||||||
+write(layer_id, page_table, start_pos, k, v)
|
+get_key_buffer(layer_id) Tensor
|
||||||
+gather(layer_id, page_table, total_len) Tuple[Tensor, Tensor]
|
+get_value_buffer(layer_id) Tensor
|
||||||
|
+set_kv_buffer(layer_id, loc, k, v)
|
||||||
|
}
|
||||||
|
|
||||||
|
class ReqToTokenPool {
|
||||||
|
+int size
|
||||||
|
+int max_context_len
|
||||||
|
+Tensor req_to_token
|
||||||
|
+alloc(num_reqs) List[int]
|
||||||
|
+free(req_indices)
|
||||||
|
+write(indices, values)
|
||||||
}
|
}
|
||||||
|
|
||||||
class KVCache {
|
class KVCache {
|
||||||
<<abstract>>
|
+Tensor k_buffer
|
||||||
+task_alloc(task_id, prompt_ids) bool
|
+Tensor v_buffer
|
||||||
+task_free(task_id)
|
+Tensor req_to_token
|
||||||
+task_extend(task_id, pos) bool
|
+Tensor req_pool_indices
|
||||||
+task_cached(task_id) int
|
+Tensor seq_lens
|
||||||
+task_record_hashes(task_id, prompt_ids, start_logical_page)
|
+Tensor out_cache_loc
|
||||||
+bind_tasks(task_ids, total_len, device) CacheView
|
+int max_len
|
||||||
|
+Optional[Tensor] page_table
|
||||||
|
+Optional[Tensor] decode_mask
|
||||||
}
|
}
|
||||||
|
|
||||||
class PageCache {
|
class PagePool {
|
||||||
+int page_size
|
+int page_size
|
||||||
-PagePool _pool
|
+bool contiguous
|
||||||
-Storage _storage
|
-KVStorage _storage
|
||||||
-TaskTable _table
|
-ReqToTokenPool _req_pool
|
||||||
|
-Allocator _alloc
|
||||||
|
-PrefixCache _prefix
|
||||||
+task_alloc(task_id, prompt_ids) bool
|
+task_alloc(task_id, prompt_ids) bool
|
||||||
+task_free(task_id)
|
+task_free(task_id)
|
||||||
+task_extend(task_id, pos) bool
|
+task_extend(task_id, pos) bool
|
||||||
+task_cached(task_id) int
|
+task_cached(task_id) int
|
||||||
+task_record_hashes(task_id, prompt_ids, start_logical_page)
|
+task_record_hashes(task_id, prompt_ids, start_logical_page)
|
||||||
+bind_tasks(task_ids, total_len, device) PageCacheView
|
+bind_tasks(task_ids, seq_lens, device, start_pos) KVCache
|
||||||
}
|
|
||||||
|
|
||||||
class ContiguousCache {
|
|
||||||
+int max_seq_len
|
|
||||||
+Tensor k, v
|
|
||||||
+task_alloc(task_id, prompt_ids) bool
|
|
||||||
+task_free(task_id)
|
|
||||||
+task_extend(task_id, pos) bool
|
|
||||||
+bind_tasks(task_ids, total_len, device) ContiguousCacheView
|
|
||||||
}
|
|
||||||
|
|
||||||
class CacheView {
|
|
||||||
<<abstract>>
|
|
||||||
+write(layer_id, k, v)
|
|
||||||
+gather(layer_id) Tuple[Tensor, Tensor]
|
|
||||||
}
|
|
||||||
|
|
||||||
class PageCacheView {
|
|
||||||
-Storage _storage
|
|
||||||
+Tensor _page_table
|
|
||||||
+int _total_len
|
|
||||||
+write(layer_id, k, v)
|
|
||||||
+gather(layer_id) Tuple[Tensor, Tensor]
|
|
||||||
}
|
|
||||||
|
|
||||||
class ContiguousCacheView {
|
|
||||||
-ContiguousCache _cache
|
|
||||||
+Tensor _batch_indices
|
|
||||||
+int _total_len
|
|
||||||
+write(layer_id, k, v)
|
|
||||||
+gather(layer_id) Tuple[Tensor, Tensor]
|
|
||||||
}
|
|
||||||
|
|
||||||
class TaskTable {
|
|
||||||
+set(task_id, page_table, cached)
|
|
||||||
+get(task_id) List[int]
|
|
||||||
+get_cached(task_id) int
|
|
||||||
+get_ref(task_id) List[int]
|
|
||||||
+pop(task_id) Tuple[List[int], int]
|
|
||||||
+table_tensor(task_ids, device) Tensor
|
|
||||||
}
|
}
|
||||||
|
|
||||||
class Task {
|
class Task {
|
||||||
@@ -924,7 +900,6 @@ classDiagram
|
|||||||
+AutoTokenizer tokenizer
|
+AutoTokenizer tokenizer
|
||||||
+int max_batch_size
|
+int max_batch_size
|
||||||
+int max_seq_len
|
+int max_seq_len
|
||||||
+int max_prompt_len
|
|
||||||
+Deque waiting_queue
|
+Deque waiting_queue
|
||||||
+List active_tasks
|
+List active_tasks
|
||||||
+add_task(prompt, max_tokens, temperature, top_p, top_k, stream_callback) str
|
+add_task(prompt, max_tokens, temperature, top_p, top_k, stream_callback) str
|
||||||
@@ -1196,11 +1171,6 @@ classDiagram
|
|||||||
}
|
}
|
||||||
|
|
||||||
class FSDPExecutor {
|
class FSDPExecutor {
|
||||||
-_prepare_model(model) nn.Module
|
|
||||||
+unwrap_model(model) dict
|
|
||||||
}
|
|
||||||
|
|
||||||
class FSDP2Executor {
|
|
||||||
-_prepare_model(model) nn.Module
|
-_prepare_model(model) nn.Module
|
||||||
-_no_sync(model) context manager
|
-_no_sync(model) context manager
|
||||||
+unwrap_model(model) dict
|
+unwrap_model(model) dict
|
||||||
@@ -1303,7 +1273,6 @@ classDiagram
|
|||||||
BaseExecutor <|-- NoneExecutor
|
BaseExecutor <|-- NoneExecutor
|
||||||
BaseExecutor <|-- DDPExecutor
|
BaseExecutor <|-- DDPExecutor
|
||||||
BaseExecutor <|-- FSDPExecutor
|
BaseExecutor <|-- FSDPExecutor
|
||||||
BaseExecutor <|-- FSDP2Executor
|
|
||||||
ResponseBuilder <|-- OpenAIResponseBuilder
|
ResponseBuilder <|-- OpenAIResponseBuilder
|
||||||
ResponseBuilder <|-- AnthropicResponseBuilder
|
ResponseBuilder <|-- AnthropicResponseBuilder
|
||||||
BaseToolParser <|-- SimpleJsonToolParser
|
BaseToolParser <|-- SimpleJsonToolParser
|
||||||
@@ -1321,17 +1290,13 @@ classDiagram
|
|||||||
RawRollout <|-- RolloutResult
|
RawRollout <|-- RolloutResult
|
||||||
LaunchStrategy <|-- TorchrunStrategy
|
LaunchStrategy <|-- TorchrunStrategy
|
||||||
LaunchStrategy <|-- LocalStrategy
|
LaunchStrategy <|-- LocalStrategy
|
||||||
KVCache <|-- PageCache
|
|
||||||
KVCache <|-- ContiguousCache
|
|
||||||
CacheView <|-- PageCacheView
|
|
||||||
CacheView <|-- ContiguousCacheView
|
|
||||||
|
|
||||||
%% --- Composition (strong ownership, part destroyed with whole) ---
|
%% --- Composition (strong ownership, part destroyed with whole) ---
|
||||||
PageCache *-- PagePool
|
PagePool *-- KVStorage
|
||||||
PageCache *-- Storage
|
PagePool *-- ReqToTokenPool
|
||||||
PageCache *-- TaskTable
|
PagePool *-- Allocator
|
||||||
|
PagePool *-- PrefixCache
|
||||||
InferenceEngine *-- InferenceScheduler
|
InferenceEngine *-- InferenceScheduler
|
||||||
InferenceScheduler *-- KVCache
|
InferenceScheduler *-- PagePool
|
||||||
InferenceScheduler *-- Executor
|
InferenceScheduler *-- Executor
|
||||||
InferenceScheduler *-- TaskManager
|
InferenceScheduler *-- TaskManager
|
||||||
AutoRegressiveLM *-- DecoderBlock
|
AutoRegressiveLM *-- DecoderBlock
|
||||||
@@ -1359,8 +1324,6 @@ classDiagram
|
|||||||
TrainContext o-- BaseScheduler
|
TrainContext o-- BaseScheduler
|
||||||
TrainContext o-- Checkpoint
|
TrainContext o-- Checkpoint
|
||||||
TrainContext o-- BaseExecutor
|
TrainContext o-- BaseExecutor
|
||||||
PageCacheView o-- Storage
|
|
||||||
ContiguousCacheView o-- ContiguousCache
|
|
||||||
SamplingPipeline o-- BaseSamplingStrategy
|
SamplingPipeline o-- BaseSamplingStrategy
|
||||||
BaseDataset o-- Store
|
BaseDataset o-- Store
|
||||||
Pipeline o-- PipelineConfig
|
Pipeline o-- PipelineConfig
|
||||||
@@ -1397,7 +1360,6 @@ classDiagram
|
|||||||
ExecutorFactory ..> NoneExecutor : creates
|
ExecutorFactory ..> NoneExecutor : creates
|
||||||
ExecutorFactory ..> DDPExecutor : creates
|
ExecutorFactory ..> DDPExecutor : creates
|
||||||
ExecutorFactory ..> FSDPExecutor : creates
|
ExecutorFactory ..> FSDPExecutor : creates
|
||||||
ExecutorFactory ..> FSDP2Executor : creates
|
|
||||||
ToolParserFactory ..> BaseToolParser : creates
|
ToolParserFactory ..> BaseToolParser : creates
|
||||||
TrainContextBuilder ..> ExecutorFactory : creates
|
TrainContextBuilder ..> ExecutorFactory : creates
|
||||||
Trainer ..> TrainContextBuilder : uses
|
Trainer ..> TrainContextBuilder : uses
|
||||||
@@ -1406,8 +1368,7 @@ classDiagram
|
|||||||
TrainContextBuilder ..> RDSampler : creates
|
TrainContextBuilder ..> RDSampler : creates
|
||||||
Checkpoint ..> Checkpoint : serializes
|
Checkpoint ..> Checkpoint : serializes
|
||||||
CheckpointCallback ..> Checkpoint : creates
|
CheckpointCallback ..> Checkpoint : creates
|
||||||
PageCache ..> PageCacheView : binds
|
PagePool ..> KVCache : binds
|
||||||
ContiguousCache ..> ContiguousCacheView : binds
|
|
||||||
InferenceEngine ..> GenerationRequest : uses
|
InferenceEngine ..> GenerationRequest : uses
|
||||||
InferenceEngine ..> GenerateResult : creates
|
InferenceEngine ..> GenerateResult : creates
|
||||||
OpenAIResponseBuilder ..> ChatCompletionRequest : receives
|
OpenAIResponseBuilder ..> ChatCompletionRequest : receives
|
||||||
@@ -1444,8 +1405,9 @@ classDiagram
|
|||||||
| **astrai.model** | AutoModel, AutoRegressiveLM, EmbeddingEncoder, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, RMSNorm, Linear, LoRAConfig, LoRALinear, RotaryEmbedding, Embedding | Neural network model |
|
| **astrai.model** | AutoModel, AutoRegressiveLM, EmbeddingEncoder, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, RMSNorm, Linear, LoRAConfig, LoRALinear, RotaryEmbedding, Embedding | Neural network model |
|
||||||
| **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
|
| **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
|
||||||
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy–GRPOStrategy, StrategyFactory, BaseScheduler–WSDScheduler, SchedulerFactory, TrainCallback(Protocol)–MetricCallback, CallbackFactory, RawRollout, RolloutResult, BaseRewardModel, RolloutGenerator, RolloutRunner | Training workflow |
|
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy–GRPOStrategy, StrategyFactory, BaseScheduler–WSDScheduler, SchedulerFactory, TrainCallback(Protocol)–MetricCallback, CallbackFactory, RawRollout, RolloutResult, BaseRewardModel, RolloutGenerator, RolloutRunner | Training workflow |
|
||||||
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, KVCache–ContiguousCache/PageCache, CacheView–ContiguousCacheView/PageCacheView, Allocator–Storage, Task, TaskManager, TaskStatus, StreamDecoder, GenerationRequest, GenerateResult, BaseSamplingStrategy–SamplingPipeline, FrequencyPenaltyStrategy, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, StopInfo, ChatMessage, FunctionDef, ToolDef, ChatCompletionRequest, AnthropicMessage, MessagesRequest, BaseToolParser, ToolParserFactory, SimpleJsonToolParser | Inference service |
|
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, PagePool, KVStorage, ReqToTokenPool, KVCache, Allocator, PrefixCache, Task, TaskManager, TaskStatus, StreamDecoder, GenerationRequest, GenerateResult, BaseSamplingStrategy–SamplingPipeline, FrequencyPenaltyStrategy, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, StopInfo, ChatMessage, FunctionDef, ToolDef, ChatCompletionRequest, AnthropicMessage, MessagesRequest, BaseToolParser, ToolParserFactory, SimpleJsonToolParser | Inference service |
|
||||||
| **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, FSDP2Executor, GradientState, AccumOptimizer, AccumScheduler, ParallelModel, RowParallelLinear, ColumnParallelLinear | Distributed parallel & gradient accumulation |
|
| **astrai.extension** | AttentionBackend, TorchNativeBackend, CudaBackend, attn_backend, ATTN_BACKEND, attn_decode, attn_prefill, attn_paged_decode, rotary_emb, apply_rotary_emb, rotary_backend, is_available | CUDA attention + rotary kernels, backend abstraction, auto-dispatch |
|
||||||
|
| **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.factory** | BaseFactory | Component registration |
|
||||||
| **astrai.protocols** | OptimizerProtocol, SchedulerProtocol | Structural subtyping for optimizer/scheduler wrappers |
|
| **astrai.protocols** | OptimizerProtocol, SchedulerProtocol | Structural subtyping for optimizer/scheduler wrappers |
|
||||||
|
|
||||||
@@ -1462,7 +1424,9 @@ classDiagram
|
|||||||
| **Observer** | `TrainCallback`, callback implementations | Training process monitoring |
|
| **Observer** | `TrainCallback`, callback implementations | Training process monitoring |
|
||||||
| **Context** | `TrainContext` | Unified training state bag |
|
| **Context** | `TrainContext` | Unified training state bag |
|
||||||
| **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with LRU eviction |
|
| **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with LRU eviction |
|
||||||
| **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor`, `FSDPExecutor`, `FSDP2Executor` | Gradient accumulation & model distribution |
|
| **Strategy (Attention)** | `AttentionBackend`, `TorchNativeBackend`, `CudaBackend` | Attention computation backend switching via context manager |
|
||||||
|
| **Auto-dispatch (Rotary)** | `apply_rotary_emb`, `rotary_backend.py`, `rotary_ops.py` | Rotary embedding CUDA kernel auto-dispatch with torch fallback |
|
||||||
|
| **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor`, `FSDPExecutor` | Gradient accumulation & model distribution |
|
||||||
| **Storage** | `Store`, `H5Store`, `MmapStore`, `JsonlStore` | Format-agnostic data access with multi-segment support |
|
| **Storage** | `Store`, `H5Store`, `MmapStore`, `JsonlStore` | Format-agnostic data access with multi-segment support |
|
||||||
| **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching |
|
| **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching |
|
||||||
| **AutoModel Registry** | `AutoModel`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading |
|
| **AutoModel Registry** | `AutoModel`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading |
|
||||||
@@ -1472,8 +1436,8 @@ classDiagram
|
|||||||
1. **Config → Training**: `TrainConfig` holds `model_fn`, `dataset`, `optimizer_fn`, `scheduler_fn`, `parallel_mode`, `executor_kwargs`
|
1. **Config → Training**: `TrainConfig` holds `model_fn`, `dataset`, `optimizer_fn`, `scheduler_fn`, `parallel_mode`, `executor_kwargs`
|
||||||
2. **Training Flow**: `Trainer` → `TrainContextBuilder` → `TrainContext`, uses `BaseStrategy` for loss, `BaseExecutor` for gradient accumulation + model distribution
|
2. **Training Flow**: `Trainer` → `TrainContextBuilder` → `TrainContext`, uses `BaseStrategy` for loss, `BaseExecutor` for gradient accumulation + model distribution
|
||||||
3. **Strategy Selection**: `StrategyFactory` creates strategy by `train_type`
|
3. **Strategy Selection**: `StrategyFactory` creates strategy by `train_type`
|
||||||
4. **Executor Selection**: `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)` → `NoneExecutor` / `DDPExecutor` / `FSDPExecutor` / `FSDP2Executor`
|
4. **Executor Selection**: `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)` → `NoneExecutor` / `DDPExecutor` / `FSDPExecutor`
|
||||||
5. **Inference Flow**: `InferenceEngine` → `InferenceScheduler` → `AutoRegressiveLM`, backed by `KVCache` + `SamplingPipeline`
|
5. **Inference Flow**: `InferenceEngine` → `InferenceScheduler` → `AutoRegressiveLM`, backed by `PagePool` + `KVCache` + `SamplingPipeline`. Attention backend selected via `attn_backend()` context manager (`TorchNativeBackend` default, `CudaBackend` for CUDA kernels). Rotary embedding auto-dispatches to CUDA kernel when available (inference mode), else torch complex multiply (training).
|
||||||
6. **Distributed**: `spawn_parallel_fn` + `setup_parallel` for multi-process DDP
|
6. **Distributed**: `spawn_parallel_fn` + `setup_parallel` for multi-process DDP
|
||||||
7. **Dataset Loading**: `DatasetFactory` creates datasets, `Store` (H5Store/MmapStore/JsonlStore) loads data with explicit `_length` and multi-segment `_data`
|
7. **Dataset Loading**: `DatasetFactory` creates datasets, `Store` (H5Store/MmapStore/JsonlStore) loads data with explicit `_length` and multi-segment `_data`
|
||||||
8. **Checkpoint**: `Checkpoint` saves/loads safetensors + metadata (rank-0 only), extra state saved as `{key}.pt`
|
8. **Checkpoint**: `Checkpoint` saves/loads safetensors + metadata (rank-0 only), extra state saved as `{key}.pt`
|
||||||
@@ -1481,4 +1445,4 @@ classDiagram
|
|||||||
10. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops
|
10. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops
|
||||||
11. **Protocols**: `OptimizerProtocol` / `SchedulerProtocol` — structural subtyping for `AccumOptimizer` / `AccumScheduler` wrappers
|
11. **Protocols**: `OptimizerProtocol` / `SchedulerProtocol` — structural subtyping for `AccumOptimizer` / `AccumScheduler` wrappers
|
||||||
|
|
||||||
> Document Update Time: 2026-07-20
|
> Document Update Time: 2026-07-31
|
||||||
@@ -0,0 +1,175 @@
|
|||||||
|
# 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.
|
||||||
|
|
||||||
|
## 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 |
|
||||||
|
| `rotary_emb` | `rotary_emb.cu` | Fused rotary embedding (cos/sin lookup + rotation) |
|
||||||
|
|
||||||
|
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+) |
|
||||||
|
| Paged split-KV MMA decode | `attn_paged_decode_split_kv_mma.cuh` | Paged cache + split-KV + MMA |
|
||||||
|
|
||||||
|
### Rotary Embedding Kernel
|
||||||
|
|
||||||
|
The `rotary_emb` kernel (`csrc/kernels/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
|
||||||
|
- 256-thread blocks, grid-stride loop
|
||||||
|
- Auto-dispatched via `apply_rotary_emb` in `astrai/extension/rotary_backend.py` (CUDA when available + inference mode, else torch complex-multiply fallback)
|
||||||
|
- No context-manager backend needed — rotary is backend-agnostic, both attention backends benefit
|
||||||
|
|
||||||
|
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).
|
||||||
|
|
||||||
|
## Build System
|
||||||
|
|
||||||
|
### Auto-detection
|
||||||
|
|
||||||
|
Kernels are built when **both** of these conditions are met:
|
||||||
|
1. `nvcc` is available on `PATH`
|
||||||
|
2. `torch.cuda.is_available()` returns `True`
|
||||||
|
|
||||||
|
Unless `CSRC_KERNELS=false` is set explicitly.
|
||||||
|
|
||||||
|
### Manual build
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# During install
|
||||||
|
CSRC_KERNELS=true pip install -e . --no-build-isolation
|
||||||
|
|
||||||
|
# Rebuild after editing .cu/.cuh files
|
||||||
|
CSRC_KERNELS=true python setup.py build_ext --inplace
|
||||||
|
# Output: astrai/extension/lib/*.so
|
||||||
|
```
|
||||||
|
|
||||||
|
### Architecture flags
|
||||||
|
|
||||||
|
`csrc/build.py` auto-detects the GPU compute capability and generates the appropriate `nvcc` gencode flag:
|
||||||
|
|
||||||
|
- **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
|
||||||
|
|
||||||
|
### Build configuration
|
||||||
|
|
||||||
|
```
|
||||||
|
NVCC_FLAGS = -O3 --expt-relaxed-constexpr --use_fast_math
|
||||||
|
--ptxas-options=-O3,-v --extra-device-vectorization --threads=8
|
||||||
|
```
|
||||||
|
|
||||||
|
The `REGISTRY` in `csrc/build.py` lists all registered kernels (currently 4). Each entry maps a kernel name to its source files and build flags.
|
||||||
|
|
||||||
|
## Attention Backend
|
||||||
|
|
||||||
|
`astrai/extension/attention_backend.py` provides the backend abstraction:
|
||||||
|
|
||||||
|
- **`AttentionBackend`** (ABC): `fwd_decode` / `fwd_prefill` abstract methods, `forward` dispatches by q_len
|
||||||
|
- **`TorchNativeBackend`**: SDPA with indirect KV cache gather (default)
|
||||||
|
- **`CudaBackend`**: CUDA kernel dispatch — decode via `attn_paged_decode` (page_size=1), prefill via `attn_prefill`
|
||||||
|
|
||||||
|
Select a backend via context manager (mirrors `torch.nn.attention.sdpa_kernel`):
|
||||||
|
|
||||||
|
```python
|
||||||
|
from astrai.extension import attn_backend, ATTN_BACKEND
|
||||||
|
|
||||||
|
with attn_backend(ATTN_BACKEND.CUDA):
|
||||||
|
engine.generate("hello")
|
||||||
|
```
|
||||||
|
|
||||||
|
`CudaBackend` falls back to `TorchNativeBackend` when a kernel is not available.
|
||||||
|
|
||||||
|
### Rotary Backend
|
||||||
|
|
||||||
|
`astrai/extension/rotary_backend.py` provides `apply_rotary_emb(x, (cos, sin))` with auto-dispatch:
|
||||||
|
|
||||||
|
- **CUDA path**: calls `rotary_emb` kernel directly when available, input is bf16 on CUDA, and `torch.is_grad_enabled()` is `False` (inference)
|
||||||
|
- **Torch fallback**: complex multiply (`torch.view_as_complex` → `torch.complex` multiply → `torch.view_as_real`), used during training (supports autograd) or when kernel unavailable
|
||||||
|
|
||||||
|
No context-manager switching needed — the dispatch is automatic per call.
|
||||||
|
|
||||||
|
## Python Wrappers
|
||||||
|
|
||||||
|
`astrai/extension/attention_ops.py` provides Python wrappers for each compiled attention kernel. Each wrapper calls its CUDA kernel directly and raises `RuntimeError` if the `.so` is not available. Fallback to torch SDPA is handled by the attention backend, not the wrapper functions.
|
||||||
|
|
||||||
|
`astrai/extension/rotary_ops.py` provides the wrapper for the rotary embedding kernel. Fallback to torch complex multiply is handled by `rotary_backend.py`.
|
||||||
|
|
||||||
|
Interface (all functions):
|
||||||
|
```
|
||||||
|
is_causal: True = causal mask; False = non-causal
|
||||||
|
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True=keep)
|
||||||
|
```
|
||||||
|
|
||||||
|
Layout convention: all q/k/v are `[batch, seq_len, n_heads, head_dim]` (blhd). Scale is always `1/sqrt(head_dim)`.
|
||||||
|
|
||||||
|
## Standalone Testing
|
||||||
|
|
||||||
|
Each `csrc/tests/*.cu` file has the `nvcc` compile command in its header comment. Example:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
nvcc -I csrc -arch=sm_89 -O3 --use_fast_math \
|
||||||
|
--ptxas-options=-O3,-v --extra-device-vectorization \
|
||||||
|
csrc/tests/attn_decode_test.cu -o /tmp/test && /tmp/test
|
||||||
|
```
|
||||||
|
|
||||||
|
Test files:
|
||||||
|
- `attn_decode_test.cu` — basic decode kernel
|
||||||
|
- `attn_paged_decode_test.cu` — paged decode kernel
|
||||||
|
- `attn_prefill_test.cu` — prefill kernel
|
||||||
|
|
||||||
|
## Benchmarks
|
||||||
|
|
||||||
|
Hardware: NVIDIA L20 (sm_89, 46 GB), CUDA 12.8, driver 570.86.
|
||||||
|
|
||||||
|
Reproduce:
|
||||||
|
```bash
|
||||||
|
nvcc -I csrc -arch=sm_89 -O3 --use_fast_math \
|
||||||
|
--ptxas-options=-O3,-v --extra-device-vectorization \
|
||||||
|
csrc/tests/attn_<name>_test.cu -o /tmp/test && /tmp/test
|
||||||
|
```
|
||||||
|
|
||||||
|
## Known Optimization Targets
|
||||||
|
|
||||||
|
- **Decode D=256**: spill eliminated (BC=16 + STAGES=2), but still 248 regs — further tiling could help.
|
||||||
|
- **Prefill single-batch**: bandwidth low (22 GB/s at q=kv=2048) — compute-bound at ~94 TFLOP/s (near L20 bf16 ceiling ~193 TFLOP/s for non-causal).
|
||||||
|
- **Decode single-batch**: bandwidth low (113 GB/s at kv=512, 13% of 864 GB/s theoretical) — small kv underutilizes SMs despite split-KV; scales to 757 GB/s (88%) at B=16+.
|
||||||
|
|
||||||
|
## File Layout
|
||||||
|
|
||||||
|
```
|
||||||
|
csrc/
|
||||||
|
├── build.py # Build system: REGISTRY, _arch_flags, nvcc flags
|
||||||
|
├── kernels/
|
||||||
|
│ ├── attn_common.h # Shared attention params (AttentionParams, PagedAttentionParams)
|
||||||
|
│ ├── attn_decode.cu # Basic decode kernel (registered)
|
||||||
|
│ ├── attn_prefill.cu # Basic prefill kernel (registered)
|
||||||
|
│ ├── attn_paged_decode.cu # Paged decode kernel (registered)
|
||||||
|
│ ├── rotary_emb.cu # Fused rotary embedding kernel (registered)
|
||||||
|
│ ├── attn_decode_split_kv.cuh # Split-KV variant
|
||||||
|
│ ├── attn_decode_split_kv_mma.cuh # Split-KV + MMA variant
|
||||||
|
│ ├── attn_prefill_split_q.cuh # Split-Q variant
|
||||||
|
│ ├── attn_prefill_split_q_mma.cuh # Split-Q + MMA variant
|
||||||
|
│ ├── attn_paged_decode_split_kv.cuh # Paged + split-KV variant
|
||||||
|
│ ├── attn_paged_decode_split_kv_mma.cuh # Paged + split-KV + MMA variant
|
||||||
|
│ ├── attn_dispatchers.cuh # Kernel dispatch macros
|
||||||
|
│ ├── attn_entry_utils.cuh # Entry point helpers
|
||||||
|
│ ├── attn_mma_utils.cuh # MMA utilities
|
||||||
|
│ └── attn_warp_utils.cuh # Warp-level utilities
|
||||||
|
└── tests/
|
||||||
|
├── test_utils.cuh # Shared test utilities
|
||||||
|
├── attn_decode_test.cu # Decode kernel test
|
||||||
|
├── attn_paged_decode_test.cu # Paged decode test
|
||||||
|
└── attn_prefill_test.cu # Prefill kernel test
|
||||||
|
```
|
||||||
|
|
||||||
|
Compiled `.so` files are placed in `astrai/extension/lib/`, separate from Python source files.
|
||||||
|
|
||||||
|
> Document Update Time: 2026-07-31
|
||||||
@@ -1,6 +1,6 @@
|
|||||||
# Data Flow
|
# Data Flow
|
||||||
|
|
||||||
This document describes the data pipeline: from raw text to model input tensors. For creating preprocessing configs, see [Preprocessing Guide](preprocessing.md).
|
This document describes the data pipeline: from raw text to model input tensors. For creating preprocessing configs, see [Preprocessing Guide](../guides/preprocessing.md).
|
||||||
|
|
||||||
## Contents
|
## Contents
|
||||||
|
|
||||||
@@ -33,7 +33,7 @@ Raw text is tokenized via `AutoTokenizer.encode()` and saved as HDF5 (`.h5`) or
|
|||||||
|
|
||||||
### Tokenization
|
### Tokenization
|
||||||
|
|
||||||
The `Pipeline` reads JSONL lines, applies the mask builder (see [Preprocessing](preprocessing.md)), and produces flat token sequences:
|
The `Pipeline` reads JSONL lines, applies the mask builder (see [Preprocessing](../guides/preprocessing.md)), and produces flat token sequences:
|
||||||
|
|
||||||
```python
|
```python
|
||||||
# Per JSONL line: messages → chat template → token IDs + loss mask
|
# Per JSONL line: messages → chat template → token IDs + loss mask
|
||||||
@@ -83,9 +83,11 @@ All backends normalise tensors into `Store._data[Dict[str, List[Tensor]]]` + `St
|
|||||||
## Dataset Architecture
|
## Dataset Architecture
|
||||||
|
|
||||||
```
|
```
|
||||||
DatasetFactory.load(train_type, load_path, window_size, stride=None,
|
DatasetFactory.load(
|
||||||
storage_type=None, tokenizer_path=None,
|
train_type, load_path=None, window_size=0, stride=None,
|
||||||
max_position_embeddings=2048, store=None)
|
storage_type=None, tokenizer_path=None,
|
||||||
|
max_len=2048, store=None
|
||||||
|
)
|
||||||
→ BaseDataset.load(load_path, storage_type=None)
|
→ BaseDataset.load(load_path, storage_type=None)
|
||||||
→ detect_format(load_path)
|
→ detect_format(load_path)
|
||||||
→ StoreFactory.create(storage_type)
|
→ StoreFactory.create(storage_type)
|
||||||
@@ -0,0 +1,233 @@
|
|||||||
|
# Internals
|
||||||
|
|
||||||
|
Mathematical foundations and internal algorithms for AstrAI's training, inference, and preprocessing pipelines. For practical usage guides, see [Training](../guides/training.md), [Inference](../guides/inference.md), and [Preprocessing](../guides/preprocessing.md).
|
||||||
|
|
||||||
|
## Contents
|
||||||
|
|
||||||
|
- [Autoregression & Causal Masking](#autoregression--causal-masking)
|
||||||
|
- [Rotary Position Embedding (RoPE)](#rotary-position-embedding-rope)
|
||||||
|
- [Training Loss Formulas](#training-loss-formulas)
|
||||||
|
- [Training Loop Internals](#training-loop-internals)
|
||||||
|
- [Callback Lifecycle](#callback-lifecycle)
|
||||||
|
- [KV Cache Mathematics](#kv-cache-mathematics)
|
||||||
|
- [Mask Algorithm Internals](#mask-algorithm-internals)
|
||||||
|
- [Gradient Accumulation Mechanics](#gradient-accumulation-mechanics)
|
||||||
|
|
||||||
|
## Autoregression & Causal Masking
|
||||||
|
|
||||||
|
Given a token sequence, the model predicts the probability of the next token. Each generated token is appended to the input and fed back, repeating until an end-of-sequence token or max length.
|
||||||
|
|
||||||
|
```
|
||||||
|
sequence : [[1, 2, 3, 4, 5, 6]]
|
||||||
|
input_ids: [[1, 2, 3, 4, 5]]
|
||||||
|
target_ids: [[2, 3, 4, 5, 6]]
|
||||||
|
```
|
||||||
|
|
||||||
|
A lower-triangular causal mask prevents attending to future positions:
|
||||||
|
|
||||||
|
```
|
||||||
|
[[0, -inf, -inf, -inf, -inf],
|
||||||
|
[0, 0, -inf, -inf, -inf],
|
||||||
|
[0, 0, 0, -inf, -inf],
|
||||||
|
[0, 0, 0, 0, -inf],
|
||||||
|
[0, 0, 0, 0, 0]]
|
||||||
|
```
|
||||||
|
|
||||||
|
This ensures position $i$ can only attend to positions $\leq i$, which is essential for autoregressive generation.
|
||||||
|
|
||||||
|
## Rotary Position Embedding (RoPE)
|
||||||
|
|
||||||
|
RoPE embeds position into Q/K vectors via complex rotation:
|
||||||
|
|
||||||
|
$$ q_i = R_i W_q x_i, \quad k_j = R_j W_k x_j, \quad q_i^T k_j = x_i^T W_q^T R_{i-j} W_k x_j $$
|
||||||
|
|
||||||
|
`RotaryEmbedding` pre-computes `cos_table` and `sin_table` (f32, `[max_len, dim/2]`). `forward()` returns a `(cos, sin)` tuple indexed by `position_ids`. `apply_rotary_emb` applies the rotation: during training it uses torch complex multiply (autograd-compatible); during inference it auto-dispatches to a fused CUDA kernel when available. The key property is that the dot product $q_i^T k_j$ depends only on the relative position $i - j$, not the absolute positions.
|
||||||
|
|
||||||
|
**Critical for inference**: RoPE is applied **before** KV cache write, not after. If applied after caching, position encoding drift occurs because cached K/V would have stale rotation factors.
|
||||||
|
|
||||||
|
## Training Loss Formulas
|
||||||
|
|
||||||
|
### SEQ (Pre-training)
|
||||||
|
|
||||||
|
Next-token cross-entropy with optional label smoothing:
|
||||||
|
|
||||||
|
$$ L_{\text{PT}} = -\sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta) $$
|
||||||
|
|
||||||
|
### SFT (Supervised Fine-Tuning)
|
||||||
|
|
||||||
|
Masked cross-entropy (`ignore_index=-100`) over response tokens only:
|
||||||
|
|
||||||
|
$$ L_{\text{SFT}} = -\sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta) $$
|
||||||
|
|
||||||
|
Prompt tokens are masked out via `loss_mask`; only response tokens contribute to the loss.
|
||||||
|
|
||||||
|
### DPO (Direct Preference Optimization)
|
||||||
|
|
||||||
|
Frozen reference model, preference margin via log-ratio:
|
||||||
|
|
||||||
|
$$ L_{\text{DPO}} = -\mathbb{E}\left[\log\sigma\left(\beta\log\frac{\pi_\theta(y_w\mid x)}{\pi_{\text{ref}}(y_w\mid x)} - \beta\log\frac{\pi_\theta(y_l\mid x)}{\pi_{\text{ref}}(y_l\mid x)}\right)\right] $$
|
||||||
|
|
||||||
|
Parameters: `beta=0.1`, `reduction="sum"`.
|
||||||
|
|
||||||
|
### GRPO (Group Relative Policy Optimization)
|
||||||
|
|
||||||
|
Token-level PPO with group-normalized advantages:
|
||||||
|
|
||||||
|
$$ \text{Advantage}_i = \frac{r_i - \mu}{\sigma + \epsilon} $$
|
||||||
|
|
||||||
|
$$ L_{\text{GRPO}} = -\mathbb{E}_t\left[\min\left(\rho_t A,\; \text{clip}\left(\rho_t, 1-\epsilon, 1+\epsilon\right)A\right)\right] + \lambda \cdot \mathbb{E}_t\left[\frac{\pi_{\text{ref}}}{\pi_\theta} - \log\frac{\pi_{\text{ref}}}{\pi_\theta} - 1\right] $$
|
||||||
|
|
||||||
|
Where $\rho_t = \pi_\theta(a_t|s_t) / \pi_{\text{old}}(a_t|s_t)$ is the per-token importance sampling ratio. Advantages are derived from scalar per-response rewards, group-normalized, and broadcast across all response tokens. Only response tokens contribute to the loss.
|
||||||
|
|
||||||
|
Parameters: `group_size=4`, `clip_eps=0.2`, `kl_coef=0.01`.
|
||||||
|
|
||||||
|
## Training Loop Internals
|
||||||
|
|
||||||
|
Two-level loop: **epoch** → **batch**. Optimizer step fires every `grad_accum_steps` batches.
|
||||||
|
|
||||||
|
```
|
||||||
|
on_train_begin
|
||||||
|
model.train()
|
||||||
|
on_epoch_begin
|
||||||
|
for batch in dataloader:
|
||||||
|
on_batch_begin
|
||||||
|
with executor.accumulate(model):
|
||||||
|
loss = strategy.compute_loss(batch)
|
||||||
|
context.loss = loss.item()
|
||||||
|
stand_loss = loss / executor.grad_accum_steps
|
||||||
|
executor.backward(stand_loss)
|
||||||
|
context.consumed_samples += (
|
||||||
|
context.config.batch_per_device * context.world_size
|
||||||
|
)
|
||||||
|
on_batch_end
|
||||||
|
|
||||||
|
if executor.sync_gradients:
|
||||||
|
on_optimizer_step
|
||||||
|
optimizer.step()
|
||||||
|
optimizer.zero_grad()
|
||||||
|
if scheduler:
|
||||||
|
scheduler.step()
|
||||||
|
on_epoch_end
|
||||||
|
on_train_end
|
||||||
|
```
|
||||||
|
|
||||||
|
The loss is divided by `grad_accum_steps` before `backward()`, so accumulated gradients sum to the correct mean.
|
||||||
|
|
||||||
|
## Callback Lifecycle
|
||||||
|
|
||||||
|
| Hook | Fires | Default callback |
|
||||||
|
|------|-------|-----------------|
|
||||||
|
| `on_train_begin` | Before training starts | `GradientCheckpointingCallback` |
|
||||||
|
| `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` |
|
||||||
|
| `on_batch_begin` | Every batch | — |
|
||||||
|
| `on_optimizer_step` | Every accumulation window | `GradientClippingCallback`, `MetricCallback`, `ProgressBarCallback` |
|
||||||
|
| `on_batch_end` | Every batch | `CheckpointCallback` |
|
||||||
|
| `on_epoch_end` | End of each epoch | `MetricCallback`, `ProgressBarCallback` |
|
||||||
|
| `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` |
|
||||||
|
| `on_train_end` | Training ends (always via finally) | `CheckpointCallback`, `MetricCallback`, `GradientCheckpointingCallback` |
|
||||||
|
|
||||||
|
Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric` (JSONL + validation, rank-0), `progress_bar` (tqdm), `gradient_clipping` (always registered; computes grad norm, clips only when `max_grad_norm` is not `None`).
|
||||||
|
|
||||||
|
## KV Cache Mathematics
|
||||||
|
|
||||||
|
At decode time, only the last query token matters. All previous K/V are cached to avoid recomputation:
|
||||||
|
|
||||||
|
$$ o_n = \sum_j \text{softmax}\left(\frac{q_n k_j}{\sqrt{d_k}}\right) v_j $$
|
||||||
|
|
||||||
|
The cache stores $k_j$ and $v_j$ for all previous positions. At each decode step, only $q_n$ (the current query) is computed fresh, and attention is computed against the cached K/V.
|
||||||
|
|
||||||
|
**RoPE ordering**: RoPE is applied to Q/K **before** writing to the KV cache. This is essential because:
|
||||||
|
1. The cached K values already contain the rotation for their original positions.
|
||||||
|
2. The new Q is rotated for its current position.
|
||||||
|
3. The dot product $q_n^T k_j$ then correctly depends on $n - j$ (relative position).
|
||||||
|
|
||||||
|
If RoPE were applied after caching, the rotation factors would be inconsistent between cached and new tokens.
|
||||||
|
|
||||||
|
### Cache Architecture
|
||||||
|
|
||||||
|
Three-layer separation (SGLang-inspired):
|
||||||
|
|
||||||
|
- **KVStorage**: Flat token-level buffers `[n_layers, size, n_kv_heads, head_dim]`.
|
||||||
|
- **ReqToTokenPool**: Index table `[req_idx, pos] → physical token slot`, shared across all layers.
|
||||||
|
- **Allocator + PrefixCache**: Paged-mode slot allocation with ref-counting, LRU eviction, and hash-based prefix sharing.
|
||||||
|
|
||||||
|
`PagePool` orchestrates all three. In contiguous mode (default), `req_to_token` is a trivial linear mapping. In paged mode, slots are allocated on demand with prefix caching support. `bind_tasks()` returns a `KVCache` dataclass with precomputed `page_table` and `decode_mask` fields (computed once per decode step, shared across all layers). Attention layers access buffers directly — no methods, no abstraction.
|
||||||
|
|
||||||
|
### Attention Backend
|
||||||
|
|
||||||
|
Attention computation is decoupled from the model via `AttentionBackend` ABC (`astrai/extension/attention_backend.py`):
|
||||||
|
|
||||||
|
- **`TorchNativeBackend`** (default): writes K/V to cache, gathers via `req_to_token` indirect indexing, calls `F.scaled_dot_product_attention`.
|
||||||
|
- **`CudaBackend`**: decode path uses `attn_paged_decode` with `page_size=1` (the `req_to_token` table serves as the page table, each token slot is a single-token "page"); prefill path gathers K/V then calls `attn_prefill`. Falls back to `TorchNativeBackend` when kernel unavailable.
|
||||||
|
|
||||||
|
Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/rotary_backend.py`, which auto-dispatches to the fused CUDA kernel (`rotary_emb.cu`) during inference or torch complex multiply during training (for autograd compatibility). Both attention backends share the same rotary dispatch.
|
||||||
|
|
||||||
|
Backend selection is thread-safe via `contextvars`, mirroring `torch.nn.attention.sdpa_kernel`:
|
||||||
|
|
||||||
|
```python
|
||||||
|
from astrai.extension import attn_backend, ATTN_BACKEND
|
||||||
|
|
||||||
|
with attn_backend(ATTN_BACKEND.CUDA):
|
||||||
|
engine.generate("hello")
|
||||||
|
```
|
||||||
|
|
||||||
|
Layout convention: all q/k/v are `[batch, seq_len, n_heads, head_dim]` (blhd). Scale is always `1/sqrt(head_dim)`.
|
||||||
|
|
||||||
|
## Mask Algorithm Internals
|
||||||
|
|
||||||
|
### Template mode (`template: true`)
|
||||||
|
|
||||||
|
1. Prepend BOS token (masked)
|
||||||
|
2. For each message in the field's array:
|
||||||
|
1. Render through `chat_template` for that single message
|
||||||
|
2. Encode rendered text
|
||||||
|
3. Apply mask rule for the message's role
|
||||||
|
|
||||||
|
### Non-template mode
|
||||||
|
|
||||||
|
Encode the field value as text. Mask value is 1 (train) or 0 (mask) per the section's `action`.
|
||||||
|
|
||||||
|
### Text config detection
|
||||||
|
|
||||||
|
When no section uses `template` and all sections have `action: "train"`, the builder omits `loss_mask` from the output — all tokens are trained.
|
||||||
|
|
||||||
|
### Position ID strategies
|
||||||
|
|
||||||
|
| Mode | Behavior |
|
||||||
|
|------|----------|
|
||||||
|
| `none` | No position IDs generated |
|
||||||
|
| `doc_reset` | Reset position to 0 at each document boundary in packed sequences |
|
||||||
|
| `continuous` | Continuous position IDs across packed documents |
|
||||||
|
|
||||||
|
Default is `doc_reset`, which ensures each document in a packed bin starts from position 0, preventing position encoding drift between unrelated documents.
|
||||||
|
|
||||||
|
## Gradient Accumulation Mechanics
|
||||||
|
|
||||||
|
Three cooperating layers enable gradient accumulation:
|
||||||
|
|
||||||
|
1. **`GradientState`** — tracks the micro-step counter. Fires `sync_gradients=True` every `grad_accum_steps` micro-batches. The counter is incremented at the **start** of `accumulate()`, before the forward pass.
|
||||||
|
|
||||||
|
2. **`executor._no_sync(model)`** — suppresses gradient synchronization on non-sync micro-steps:
|
||||||
|
- `NoneExecutor`: `nullcontext` (nothing to skip)
|
||||||
|
- `DDPExecutor`: `model.no_sync()` (PyTorch's built-in — skips all-reduce of gradient buckets)
|
||||||
|
- `FSDPExecutor`: `set_requires_gradient_sync(False, recurse=True)` on each `FSDPModule` (FSDP2's native mechanism)
|
||||||
|
|
||||||
|
3. **`AccumOptimizer` / `AccumScheduler`** — wrap the real optimizer/scheduler. `step()` and `zero_grad()` are gated on `sync_gradients` — they only forward to the inner optimizer when the sync flag is True.
|
||||||
|
|
||||||
|
The loss is divided by `grad_accum_steps` before `backward()`, so gradients sum to the correct mean across micro-steps. `consumed_samples` increments by `batch_per_device * world_size` every micro-batch.
|
||||||
|
|
||||||
|
### Effective batch size
|
||||||
|
|
||||||
|
$$ \text{Effective batch} = \text{nprocs} \times \text{batch\_per\_device} \times \text{grad\_accum\_steps} $$
|
||||||
|
|
||||||
|
### Total optimizer steps
|
||||||
|
|
||||||
|
```
|
||||||
|
samples_per_replica = ceil(dataset_len / nprocs)
|
||||||
|
batches_per_replica = ceil(samples_per_replica / batch_per_device)
|
||||||
|
total_steps = (batches_per_replica // grad_accum_steps) * n_epoch
|
||||||
|
```
|
||||||
|
|
||||||
|
This accounts for data-parallel sharding — each rank processes `1/nprocs` of the dataset.
|
||||||
|
|
||||||
|
> Document Update Time: 2026-07-31
|
||||||
@@ -0,0 +1,235 @@
|
|||||||
|
# Getting Started
|
||||||
|
|
||||||
|
This guide walks you through installing AstrAI, downloading a model, running inference, preprocessing data, and launching your first training job.
|
||||||
|
|
||||||
|
## Prerequisites
|
||||||
|
|
||||||
|
- **Python 3.12+**
|
||||||
|
- **PyTorch 2.11+** (CUDA 12.8 recommended for GPU support)
|
||||||
|
- NVIDIA GPU with CUDA (optional but recommended; CPU works for inference)
|
||||||
|
|
||||||
|
## 1. Install
|
||||||
|
|
||||||
|
```bash
|
||||||
|
git clone https://github.com/ViperEkura/AstrAI.git
|
||||||
|
cd AstrAI
|
||||||
|
|
||||||
|
# Basic install (pure PyTorch, no custom CUDA kernels)
|
||||||
|
pip install -e .
|
||||||
|
|
||||||
|
# With CUDA kernels (optional, for fused attention and rotary embedding)
|
||||||
|
# CSRC_KERNELS=true pip install -e . --no-build-isolation
|
||||||
|
|
||||||
|
# With dev dependencies (pytest, ruff)
|
||||||
|
# pip install -e ".[dev]"
|
||||||
|
```
|
||||||
|
|
||||||
|
> **CUDA kernels** are opt-in. They are not built by default. When built, they can be activated via `with attn_backend(ATTN_BACKEND.CUDA):` for accelerated decode/prefill, and the fused rotary embedding kernel is auto-dispatched when available. You can skip them for normal usage.
|
||||||
|
|
||||||
|
## 2. Download Model Weights
|
||||||
|
|
||||||
|
AstrAI uses HuggingFace-style model directories. Download the default 1B instruction-tuned model:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/demo/download.py
|
||||||
|
# → Downloads to params/
|
||||||
|
```
|
||||||
|
|
||||||
|
To use a different model:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/demo/download.py --repo-id <HF_REPO_ID> --local-dir ./my_model
|
||||||
|
```
|
||||||
|
|
||||||
|
The model directory contains:
|
||||||
|
- `config.json` — model architecture configuration
|
||||||
|
- `model.safetensors` — model weights
|
||||||
|
- `tokenizer.json` + `tokenizer_config.json` — tokenizer files (including chat template)
|
||||||
|
|
||||||
|
## 3. Run Inference
|
||||||
|
|
||||||
|
### Interactive Chat (Simplest)
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/demo/stream_chat.py
|
||||||
|
# Type your message after >>, type !exit to quit
|
||||||
|
```
|
||||||
|
|
||||||
|
This starts a multi-turn interactive chat session with streaming output.
|
||||||
|
|
||||||
|
### Start an HTTP Server
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# Terminal 1: start server
|
||||||
|
python scripts/tools/server.py --param_path ./params --device cuda
|
||||||
|
|
||||||
|
# Terminal 2: query (OpenAI-compatible API)
|
||||||
|
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{"messages":[{"role":"user","content":"Hello"}],"max_tokens":512}'
|
||||||
|
```
|
||||||
|
|
||||||
|
The server also supports the Anthropic API at `/v1/messages`. See [Inference Guide](guides/inference.md) for full API documentation.
|
||||||
|
|
||||||
|
### Batch Generation from a File
|
||||||
|
|
||||||
|
Create an input JSONL file (one JSON object per line):
|
||||||
|
|
||||||
|
```json
|
||||||
|
{"question": "What is machine learning?"}
|
||||||
|
{"question": "Explain gradient descent."}
|
||||||
|
```
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/tools/generate.py \
|
||||||
|
--param_path ./params \
|
||||||
|
--input_json_file input.jsonl \
|
||||||
|
--output_json_file output.jsonl
|
||||||
|
```
|
||||||
|
|
||||||
|
## 4. Preprocess Data
|
||||||
|
|
||||||
|
AstrAI uses a declarative JSON config to define the preprocessing pipeline. Create a config file for your training type:
|
||||||
|
|
||||||
|
### Pretraining (seq)
|
||||||
|
|
||||||
|
Input JSONL:
|
||||||
|
```json
|
||||||
|
{"text": "Artificial intelligence is..."}
|
||||||
|
```
|
||||||
|
|
||||||
|
Config (`pretrain.json`):
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"input": {
|
||||||
|
"sections": [{"field": "text", "action": "train"}]
|
||||||
|
},
|
||||||
|
"preprocessing": {"max_seq_len": 2048},
|
||||||
|
"output": {"storage_format": "bin"}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### SFT (Supervised Fine-Tuning)
|
||||||
|
|
||||||
|
Input JSONL:
|
||||||
|
```json
|
||||||
|
{"messages": [{"role": "user", "content": "Hi"}, {"role": "assistant", "content": "Hello!"}]}
|
||||||
|
```
|
||||||
|
|
||||||
|
Config (`sft.json`):
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"input": {
|
||||||
|
"sections": [{"field": "messages", "action": "$role", "template": true}]
|
||||||
|
},
|
||||||
|
"mask": {
|
||||||
|
"system": "mask",
|
||||||
|
"user": "mask",
|
||||||
|
"assistant": "train"
|
||||||
|
},
|
||||||
|
"mask_default": "mask",
|
||||||
|
"preprocessing": {"max_seq_len": 2048},
|
||||||
|
"output": {"storage_format": "bin", "dtype": {"loss_mask": "bool"}}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Run Preprocessing
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/tools/preprocess.py data/*.jsonl -o output/ -c pretrain.json
|
||||||
|
```
|
||||||
|
|
||||||
|
See [Preprocessing Guide](guides/preprocessing.md) for DPO/GRPO configs and all options.
|
||||||
|
|
||||||
|
## 5. Train
|
||||||
|
|
||||||
|
### Single GPU
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/tools/train.py \
|
||||||
|
--train_type=seq \
|
||||||
|
--data_root_path=/path/to/dataset \
|
||||||
|
--param_path=./params \
|
||||||
|
--batch_per_device=4 \
|
||||||
|
--grad_accum_steps=8 \
|
||||||
|
--max_lr=1e-4 \
|
||||||
|
--window_size=2048 \
|
||||||
|
--ckpt_dir=./checkpoint \
|
||||||
|
--nprocs=1 \
|
||||||
|
--parallel_mode=none
|
||||||
|
```
|
||||||
|
|
||||||
|
### Multi-GPU (DDP)
|
||||||
|
|
||||||
|
```bash
|
||||||
|
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||||
|
export NCCL_P2P_DISABLE=1
|
||||||
|
export NCCL_NET_GDR_LEVEL=0
|
||||||
|
|
||||||
|
python scripts/tools/train.py \
|
||||||
|
--train_type=seq \
|
||||||
|
--data_root_path=/path/to/dataset \
|
||||||
|
--param_path=./params \
|
||||||
|
--parallel_mode=ddp \
|
||||||
|
--nprocs=4 \
|
||||||
|
--batch_per_device=4 \
|
||||||
|
--grad_accum_steps=8 \
|
||||||
|
--max_lr=1e-4 \
|
||||||
|
--window_size=2048 \
|
||||||
|
--ckpt_dir=./checkpoint
|
||||||
|
```
|
||||||
|
|
||||||
|
### Training Types
|
||||||
|
|
||||||
|
| `--train_type` | Description | Data Keys |
|
||||||
|
|----------------|-------------|-----------|
|
||||||
|
| `seq` | Pre-training (next-token prediction) | `sequence` |
|
||||||
|
| `sft` | Supervised fine-tuning (masked loss) | `sequence`, `loss_mask` |
|
||||||
|
| `dpo` | Direct Preference Optimization | `chosen`, `rejected`, `*_mask` |
|
||||||
|
| `grpo` | Group Relative Policy Optimization | `prompts`, `responses`, `masks`, `rewards` |
|
||||||
|
|
||||||
|
See [Training Guide](guides/training.md) for loss formulas and strategies. See [Distributed Guide](guides/distributed.md) for DDP/FSDP details.
|
||||||
|
|
||||||
|
## 6. Evaluate
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# HumanEval (code generation, auto-downloads dataset)
|
||||||
|
python scripts/eval/evaluate_humaneval.py --param_path ./params --num_samples 20
|
||||||
|
|
||||||
|
# MMLU (knowledge, auto-downloads dataset)
|
||||||
|
python scripts/eval/evaluate_mmlu.py --param_path ./params --n_shot 5
|
||||||
|
|
||||||
|
# Perplexity on custom data
|
||||||
|
python scripts/eval/evaluate_ppl.py --param_path ./params --input_path data.jsonl --output_dir ppl_results/
|
||||||
|
```
|
||||||
|
|
||||||
|
See [Evaluation Guide](guides/evaluation.md) for all benchmarks.
|
||||||
|
|
||||||
|
## 7. Docker
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# Build
|
||||||
|
docker build -t astrai:latest .
|
||||||
|
|
||||||
|
# Run inference server with GPU
|
||||||
|
docker run --gpus all -p 8000:8000 astrai:latest \
|
||||||
|
python -m scripts.tools.server --port 8000 --device cuda
|
||||||
|
|
||||||
|
# Docker Compose (GPU)
|
||||||
|
docker compose up -d
|
||||||
|
```
|
||||||
|
|
||||||
|
## Next Steps
|
||||||
|
|
||||||
|
| Topic | Document |
|
||||||
|
|-------|----------|
|
||||||
|
| CLI parameters (train, server, generate, preprocess) | [CLI Reference](guides/params.md) |
|
||||||
|
| Preprocessing pipeline details | [Preprocessing Guide](guides/preprocessing.md) |
|
||||||
|
| Training loop, strategies, schedulers | [Training Guide](guides/training.md) |
|
||||||
|
| KV cache, continuous batching, HTTP API | [Inference Guide](guides/inference.md) |
|
||||||
|
| Evaluation benchmarks | [Evaluation Guide](guides/evaluation.md) |
|
||||||
|
| Multi-GPU DDP / FSDP | [Distributed Guide](guides/distributed.md) |
|
||||||
|
| System architecture | [Architecture](developer/architecture.md) |
|
||||||
|
| Data pipeline internals | [Data Flow](developer/dataflow.md) |
|
||||||
|
|
||||||
|
> Document Update Time: 2026-07-31
|
||||||
@@ -0,0 +1,255 @@
|
|||||||
|
# Distributed Training
|
||||||
|
|
||||||
|
AstrAI supports three parallel modes: **single GPU** (`none`), **Data Parallel** (`ddp`), and **Fully Sharded Data Parallel** (`fsdp`). This guide covers when to use each, how to launch multi-GPU training, and how gradient accumulation works.
|
||||||
|
|
||||||
|
## Quick Start
|
||||||
|
|
||||||
|
### Single GPU
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/tools/train.py \
|
||||||
|
--train_type=sft \
|
||||||
|
--param_path ./params \
|
||||||
|
--data_root_path ./dataset \
|
||||||
|
--parallel_mode=none \
|
||||||
|
--nprocs=1 \
|
||||||
|
--batch_per_device=4 \
|
||||||
|
--grad_accum_steps=8
|
||||||
|
```
|
||||||
|
|
||||||
|
### Multi-GPU DDP (4 GPUs)
|
||||||
|
|
||||||
|
```bash
|
||||||
|
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||||
|
export NCCL_P2P_DISABLE=1
|
||||||
|
export NCCL_NET_GDR_LEVEL=0
|
||||||
|
|
||||||
|
python scripts/tools/train.py \
|
||||||
|
--train_type=sft \
|
||||||
|
--param_path ./params \
|
||||||
|
--data_root_path ./dataset \
|
||||||
|
--parallel_mode=ddp \
|
||||||
|
--nprocs=4 \
|
||||||
|
--batch_per_device=4 \
|
||||||
|
--grad_accum_steps=8
|
||||||
|
```
|
||||||
|
|
||||||
|
### Multi-GPU FSDP (4 GPUs)
|
||||||
|
|
||||||
|
```bash
|
||||||
|
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||||
|
export NCCL_P2P_DISABLE=1
|
||||||
|
export NCCL_NET_GDR_LEVEL=0
|
||||||
|
|
||||||
|
python scripts/tools/train.py \
|
||||||
|
--train_type=sft \
|
||||||
|
--param_path ./params \
|
||||||
|
--data_root_path ./dataset \
|
||||||
|
--parallel_mode=fsdp \
|
||||||
|
--nprocs=4 \
|
||||||
|
--batch_per_device=4 \
|
||||||
|
--grad_accum_steps=8
|
||||||
|
```
|
||||||
|
|
||||||
|
> `--parallel_mode` defaults to `fsdp`. You can omit it for FSDP.
|
||||||
|
|
||||||
|
## Parallel Modes
|
||||||
|
|
||||||
|
| Mode | `--parallel_mode` | Param Layout | Memory | When to Use |
|
||||||
|
|------|-------------------|--------------|--------|-------------|
|
||||||
|
| Single GPU | `none` | Full, replicated | Highest | Small models, DPO/GRPO, debugging |
|
||||||
|
| DDP | `ddp` | Full, replicated | High | Most multi-GPU training |
|
||||||
|
| FSDP | `fsdp` | Sharded (DTensor) | Lowest | Large models that don't fit in single GPU |
|
||||||
|
|
||||||
|
### NoneExecutor
|
||||||
|
|
||||||
|
No wrapping. The model runs as-is on a single device. Gradient accumulation still works via `AccumOptimizer`/`AccumScheduler` (they gate `step()` on the sync counter). Checkpoint saving is a plain `state_dict()` call.
|
||||||
|
|
||||||
|
### DDPExecutor
|
||||||
|
|
||||||
|
Wraps the model with `torch.nn.parallel.DistributedDataParallel`. Each rank has a full copy of the model; gradients are all-reduced across ranks. Uses `gradient_as_bucket_view=True` and `broadcast_buffers=False` by default (hardcoded in `train.py`).
|
||||||
|
|
||||||
|
During gradient accumulation, non-sync micro-steps use `model.no_sync()` to skip gradient all-reduce. Only the final micro-step triggers the all-reduce.
|
||||||
|
|
||||||
|
### FSDPExecutor (FSDP2 / `fully_shard`)
|
||||||
|
|
||||||
|
Uses PyTorch's FSDP2 per-module API (`torch.distributed.fsdp.fully_shard`). Each model child (e.g., each `DecoderBlock`) is individually sharded — parameters become `DTensor`s distributed across ranks. No `FlatParameter`, original parameter names are preserved.
|
||||||
|
|
||||||
|
Key differences from DDP:
|
||||||
|
- **Lower memory**: parameters are sharded, not replicated.
|
||||||
|
- **Custom grad norm**: FSDP gradients are `DTensor`s, so `clip_grad_norm` computes the local norm, then all-reduces to get the global norm.
|
||||||
|
- **Collective checkpoint ops**: `unshard()` and `full_tensor()` are collective — all ranks must call them even though only rank-0 saves. The executor handles this via `dist.barrier()` in `checkpoint_context`.
|
||||||
|
- **Root skipped**: `fully_shard` is applied to direct children only (not the root model) due to an `ABC + Generic[T]` MRO incompatibility.
|
||||||
|
|
||||||
|
## Gradient Accumulation
|
||||||
|
|
||||||
|
Gradient accumulation lets you simulate a larger effective batch size by accumulating gradients over multiple micro-batches before calling `optimizer.step()`.
|
||||||
|
|
||||||
|
```
|
||||||
|
Effective batch = nprocs × batch_per_device × grad_accum_steps
|
||||||
|
```
|
||||||
|
|
||||||
|
Example: 4 GPUs × batch 4 × accum 8 = effective batch 256.
|
||||||
|
|
||||||
|
### How it works
|
||||||
|
|
||||||
|
Three cooperating layers:
|
||||||
|
|
||||||
|
1. **`GradientState`** — tracks the micro-step counter. Fires `sync_gradients=True` every `grad_accum_steps` micro-batches.
|
||||||
|
2. **`executor._no_sync(model)`** — suppresses gradient synchronization on non-sync micro-steps:
|
||||||
|
- `none`: `nullcontext` (nothing to skip)
|
||||||
|
- `ddp`: `model.no_sync()` (skips all-reduce)
|
||||||
|
- `fsdp`: `set_requires_gradient_sync(False)` on each `FSDPModule`
|
||||||
|
3. **`AccumOptimizer` / `AccumScheduler`** — gate `step()` and `zero_grad()` on `sync_gradients`, so the optimizer only fires on the last micro-step.
|
||||||
|
|
||||||
|
The loss is divided by `grad_accum_steps` before `backward()`, so gradients sum to the correct mean.
|
||||||
|
|
||||||
|
## Process Launching
|
||||||
|
|
||||||
|
AstrAI auto-detects the launch method:
|
||||||
|
|
||||||
|
| Detection | Strategy | Use Case |
|
||||||
|
|-----------|----------|----------|
|
||||||
|
| `torchelastic` / `torchrun` env vars | `TorchrunStrategy` | External orchestrator (torchrun, SLURM, K8s) |
|
||||||
|
| `RANK` + `WORLD_SIZE` env vars | `TorchrunStrategy` | External launch |
|
||||||
|
| Neither | `LocalStrategy` | `python scripts/tools/train.py` (in-process spawn) |
|
||||||
|
|
||||||
|
### Local (default)
|
||||||
|
|
||||||
|
When you run `python scripts/tools/train.py --nprocs=4`, AstrAI uses `torch.multiprocessing.start_processes` to spawn 4 child processes. The parent process manages signal forwarding (SIGTERM/SIGINT) and waits for all children to finish.
|
||||||
|
|
||||||
|
### Torchrun
|
||||||
|
|
||||||
|
For multi-node or SLURM environments:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
torchrun --nproc_per_node=4 scripts/tools/train.py \
|
||||||
|
--train_type=sft \
|
||||||
|
--parallel_mode=ddp \
|
||||||
|
--param_path ./params \
|
||||||
|
--data_root_path ./dataset \
|
||||||
|
--batch_per_device=4
|
||||||
|
```
|
||||||
|
|
||||||
|
When launched via torchrun, AstrAI reads `RANK`, `WORLD_SIZE`, `LOCAL_RANK` from the environment and uses `TorchrunStrategy`. The `--nprocs` flag is ignored (the orchestrator controls process count).
|
||||||
|
|
||||||
|
## NCCL Environment Variables
|
||||||
|
|
||||||
|
For multi-GPU training, you **must** set these environment variables:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
export NCCL_P2P_DISABLE=1
|
||||||
|
export NCCL_NET_GDR_LEVEL=0
|
||||||
|
```
|
||||||
|
|
||||||
|
These are required on certain hardware configurations (see `AGENTS.md`). Without them, NCCL may hang or crash during collective operations. These are set in the training shell scripts (`train-seq.sh`, `train-sft.sh`, `train-dpo.sh`) but not in Python code — you must export them before launching.
|
||||||
|
|
||||||
|
## Checkpoint Saving
|
||||||
|
|
||||||
|
Checkpoints are saved by **rank-0 only**. The flow:
|
||||||
|
|
||||||
|
1. `executor.checkpoint_context(model)` — wraps with `dist.barrier()` before and after (distributed only).
|
||||||
|
2. `executor.unwrap_model(model)` — gathers the full state dict:
|
||||||
|
- `none`: `model.state_dict()`
|
||||||
|
- `ddp`: `model.module.state_dict()`
|
||||||
|
- `fsdp`: `unshard()` → `full_tensor()` → `reshard()` (collective on all ranks, result kept only on rank-0)
|
||||||
|
3. Non-rank-0 ranks get `None` — the save is skipped.
|
||||||
|
4. Rank-0 writes `meta.json`, `config.json`, `model.safetensors`, and optional `{key}.pt` (optimizer/scheduler state).
|
||||||
|
|
||||||
|
> **FSDP note**: Even though only rank-0 saves, all ranks must participate in `unwrap_model` because `unshard()` and `full_tensor()` are collective operations. The barriers in `checkpoint_context` keep all ranks in lockstep.
|
||||||
|
|
||||||
|
## Total Steps Calculation
|
||||||
|
|
||||||
|
The scheduler's total step count accounts for data-parallel sharding:
|
||||||
|
|
||||||
|
```
|
||||||
|
samples_per_replica = ceil(dataset_len / nprocs)
|
||||||
|
batches_per_replica = ceil(samples_per_replica / batch_per_device)
|
||||||
|
total_steps = (batches_per_replica // grad_accum_steps) * n_epoch
|
||||||
|
```
|
||||||
|
|
||||||
|
This ensures the LR schedule is correctly scaled regardless of the number of GPUs.
|
||||||
|
|
||||||
|
## Real Examples
|
||||||
|
|
||||||
|
### Pretraining (seq, DDP, 4 GPUs)
|
||||||
|
|
||||||
|
```bash
|
||||||
|
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||||
|
export NCCL_P2P_DISABLE=1
|
||||||
|
export NCCL_NET_GDR_LEVEL=0
|
||||||
|
|
||||||
|
python scripts/tools/train.py \
|
||||||
|
--train_type=seq \
|
||||||
|
--param_path ./params \
|
||||||
|
--data_root_path ./dataset/cached \
|
||||||
|
--parallel_mode=ddp \
|
||||||
|
--nprocs=4 \
|
||||||
|
--n_epoch=1 \
|
||||||
|
--max_lr=2e-4 \
|
||||||
|
--schedule_type=wsd \
|
||||||
|
--warmup_ratio=0.02 \
|
||||||
|
--window_size=2048 \
|
||||||
|
--batch_per_device=4 \
|
||||||
|
--grad_accum_steps=32 \
|
||||||
|
--ckpt_interval=2000
|
||||||
|
# Effective batch = 4 × 4 × 32 = 512
|
||||||
|
```
|
||||||
|
|
||||||
|
### SFT (DDP, 4 GPUs)
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/tools/train.py \
|
||||||
|
--train_type=sft \
|
||||||
|
--param_path ./AstrAI-V1-base \
|
||||||
|
--data_root_path ./dataset/cached_sft \
|
||||||
|
--parallel_mode=ddp \
|
||||||
|
--nprocs=4 \
|
||||||
|
--n_epoch=2 \
|
||||||
|
--max_lr=2e-5 \
|
||||||
|
--schedule_type=cosine \
|
||||||
|
--warmup_ratio=0.02 \
|
||||||
|
--min_rate=0.05 \
|
||||||
|
--window_size=2048 \
|
||||||
|
--batch_per_device=4 \
|
||||||
|
--grad_accum_steps=8
|
||||||
|
# Effective batch = 4 × 4 × 8 = 128
|
||||||
|
```
|
||||||
|
|
||||||
|
### DPO (Single GPU)
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/tools/train.py \
|
||||||
|
--train_type=dpo \
|
||||||
|
--param_path ./checkpoint/epoch_1_step_6000 \
|
||||||
|
--data_root_path ./alpaca_dpo.jsonl \
|
||||||
|
--parallel_mode=none \
|
||||||
|
--nprocs=1 \
|
||||||
|
--max_lr=5e-6 \
|
||||||
|
--schedule_type=cosine \
|
||||||
|
--warmup_ratio=0.1 \
|
||||||
|
--min_rate=0.1 \
|
||||||
|
--window_size=1024 \
|
||||||
|
--batch_per_device=4 \
|
||||||
|
--grad_accum_steps=8 \
|
||||||
|
--dpo_beta=0.1 \
|
||||||
|
--max_grad_norm=50
|
||||||
|
```
|
||||||
|
|
||||||
|
## CLI Parameters
|
||||||
|
|
||||||
|
| Parameter | Default | Description |
|
||||||
|
|-----------|---------|-------------|
|
||||||
|
| `--nprocs` | 1 | Number of GPUs / processes |
|
||||||
|
| `--parallel_mode` | `fsdp` | `none`, `ddp`, or `fsdp` |
|
||||||
|
| `--start_method` | `spawn` | Multiprocessing start method (`spawn`, `fork`, `forkserver`) |
|
||||||
|
| `--backend` | `nccl` | Distributed backend (`nccl`, `gloo`) |
|
||||||
|
| `--master_addr` | `localhost` | Master node address |
|
||||||
|
| `--master_port` | `29500` | Master node port |
|
||||||
|
| `--device_type` | `cuda` | Device type |
|
||||||
|
|
||||||
|
> `--tp_size` is parsed but **not yet wired** — tensor parallelism is future work. `ColumnParallelLinear` / `RowParallelLinear` exist in `astrai/parallel/module.py` but are not used by the model.
|
||||||
|
|
||||||
|
Full parameter reference: [CLI Reference](params.md). Training loop and strategies: [Training Guide](training.md).
|
||||||
|
|
||||||
|
> Document Update Time: 2026-07-30
|
||||||
@@ -0,0 +1,252 @@
|
|||||||
|
# Evaluation
|
||||||
|
|
||||||
|
AstrAI provides 7 evaluation scripts in `scripts/eval/` covering code generation, knowledge QA, perplexity, summarization, data quality, instruction following, and weight analysis.
|
||||||
|
|
||||||
|
## Overview
|
||||||
|
|
||||||
|
| Script | Metric | Model Invocation | External Dataset |
|
||||||
|
|--------|--------|-------------------|-------------------|
|
||||||
|
| `evaluate_humaneval.py` | Code-gen pass@1/10/100 | `InferenceEngine.generate` | HF `openai/openai_humaneval` (auto-download) |
|
||||||
|
| `evaluate_mmlu.py` | MCQ accuracy (log-likelihood) | Direct `model()` forward | HF `cais/mmlu` (auto-download) |
|
||||||
|
| `evaluate_ppl.py` | Perplexity / token loss | Direct `model()` forward | User JSONL |
|
||||||
|
| `evaluate_rouge.py` | ROUGE-1/2/L | None (pure metric) | User JSONL |
|
||||||
|
| `evaluate_ifd.py` | Instruction-Following Difficulty | Direct `model()` forward | User JSONL |
|
||||||
|
| `evaluate_ifeval.py` | Instruction-following constraints | `InferenceEngine.generate` | HF `google/IFEval` (auto-download) |
|
||||||
|
| `analyze_weights.py` | SVD effective rank / weight stats | None (loads safetensors) | Checkpoint dir |
|
||||||
|
|
||||||
|
Two invocation patterns exist:
|
||||||
|
- **Generation benchmarks** (HumanEval, IFEval): use `InferenceEngine` to generate responses, then score them.
|
||||||
|
- **Scoring benchmarks** (MMLU, PPL, IFD): call `model()` directly under `torch.inference_mode()` for log-likelihood computation.
|
||||||
|
|
||||||
|
Common defaults: `--param_path` defaults to `./params`; dtype defaults to `bfloat16` on CUDA, `float32` on CPU.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## HumanEval (Code Generation)
|
||||||
|
|
||||||
|
Generates completions for 164 programming problems, executes them against hidden tests, and reports pass@k.
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/eval/evaluate_humaneval.py \
|
||||||
|
--param_path ./params \
|
||||||
|
--num_samples 20 \
|
||||||
|
--batch_size 32 \
|
||||||
|
--max_tokens 512 \
|
||||||
|
--output results/humaneval.json
|
||||||
|
```
|
||||||
|
|
||||||
|
| Parameter | Default | Description |
|
||||||
|
|-----------|---------|-------------|
|
||||||
|
| `--param_path` | `./params` | Model directory |
|
||||||
|
| `--data_path` | `./humaneval/HumanEval.jsonl` | HumanEval JSONL (auto-downloaded if missing) |
|
||||||
|
| `--output` | None | Save results JSON (also writes `_completions.json`) |
|
||||||
|
| `--test_only` | None | Test an existing completions JSON (skip generation) |
|
||||||
|
| `--generate_only` | False | Only generate, skip execution/testing |
|
||||||
|
| `--num_samples` | 200 | Completions per problem (pass@k needs >= k) |
|
||||||
|
| `--max_tokens` | 512 | Max generation length |
|
||||||
|
| `--temperature` | 0.8 | Sampling temperature |
|
||||||
|
| `--top_p` | 0.95 | Nucleus sampling threshold |
|
||||||
|
| `--top_k` | 50 | Top-k sampling |
|
||||||
|
| `--batch_size` | 32 | Generation batch size |
|
||||||
|
| `--test_workers` | 8 | ProcessPoolExecutor workers for test execution |
|
||||||
|
| `--test_timeout` | 3.0 | Per-subprocess timeout (seconds) |
|
||||||
|
| `--problems` | None | Restrict to specific problem indices |
|
||||||
|
|
||||||
|
**Output**: stdout prints `pass@1`, `pass@10`, `pass@100`. With `--output`, writes per-problem results + `_summary` aggregate and a `_completions.json` file.
|
||||||
|
|
||||||
|
**Data**: Auto-downloads `openai/openai_humaneval` from HuggingFace on first run. Each problem has `task_id`, `entry_point`, `prompt`, `test`.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## MMLU (Knowledge QA)
|
||||||
|
|
||||||
|
57-subject multiple-choice accuracy via log-likelihood comparison. Supports n-shot few-shot prompting and option permutation.
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/eval/evaluate_mmlu.py \
|
||||||
|
--param_path ./params \
|
||||||
|
--n_shot 5 \
|
||||||
|
--subjects math_algebra history_us \
|
||||||
|
--output results/mmlu.json
|
||||||
|
```
|
||||||
|
|
||||||
|
| Parameter | Default | Description |
|
||||||
|
|-----------|---------|-------------|
|
||||||
|
| `--param_path` | `./params` | Model directory |
|
||||||
|
| `--data_dir` | `./mmlu_data` | MMLU data directory (per-subject CSVs) |
|
||||||
|
| `--download` | False | Force re-download |
|
||||||
|
| `--n_shot` | 5 | Few-shot examples (0 = zero-shot) |
|
||||||
|
| `--subjects` | all 57 | Specific subjects to evaluate |
|
||||||
|
| `--output` | None | Output JSON path |
|
||||||
|
| `--split` | `test` | `test` or `val` |
|
||||||
|
| `--device` | auto | Device (`cuda` / `cpu`) |
|
||||||
|
| `--dtype` | auto | `bfloat16` on CUDA, `float32` on CPU |
|
||||||
|
| `--seed` | 0 | Seed for option permutation (0 = enabled, -1 = disabled) |
|
||||||
|
|
||||||
|
**How it works**: For each question, builds a prompt with n-shot examples, then scores each choice (A/B/C/D) by computing the summed log-likelihood of the choice token given the context. The choice with the highest log-prob is the prediction.
|
||||||
|
|
||||||
|
**Output**: stdout prints per-subject accuracy and overall. With `--output`, writes per-subject `{accuracy, correct, total}` + `_overall` aggregate.
|
||||||
|
|
||||||
|
**Data**: Auto-downloads `cais/mmlu` from HuggingFace. Stored as per-subject CSVs in `<data_dir>/<split>/` and `<data_dir>/dev/` (for few-shot).
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Perplexity (PPL)
|
||||||
|
|
||||||
|
Token-level negative-log-likelihood and perplexity on arbitrary text data. Supports streaming mode (memory-efficient) and non-streaming mode (exact per-token stats).
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/eval/evaluate_ppl.py \
|
||||||
|
--param_path ./params \
|
||||||
|
--input_path data.jsonl \
|
||||||
|
--output_dir ppl_results/ \
|
||||||
|
--batch_size 4 \
|
||||||
|
--max_length 2048
|
||||||
|
```
|
||||||
|
|
||||||
|
| Parameter | Default | Description |
|
||||||
|
|-----------|---------|-------------|
|
||||||
|
| `--param_path` | required | Model directory |
|
||||||
|
| `--input_path` | required | Input file, glob, or directory |
|
||||||
|
| `--output_dir` | required | Output directory for `summary.json` + token JSONL |
|
||||||
|
| `--text_key` | `text` | Key for the text field in input data |
|
||||||
|
| `--batch_size` | 4 | Batch size |
|
||||||
|
| `--max_length` | 2048 | Max sequence length (tokens) |
|
||||||
|
| `--token_level` | False | Store per-token log_probs + token-type analysis |
|
||||||
|
| `--max_samples` | None | Random subsample per file |
|
||||||
|
| `--device` | auto | Device |
|
||||||
|
| `--dtype` | auto | Torch dtype |
|
||||||
|
|
||||||
|
**Input**: JSONL or JSON files. Each item must have a field named by `--text_key` (default `text`). If `--input_path` is a directory, recursively collects `*.jsonl` and `*.json`.
|
||||||
|
|
||||||
|
**Output**: `summary.json` with per-file stats (tokens, mean/median loss, perplexity, p50/p90/p95/p99). With `--token_level`, also writes per-token JSONL with token IDs and log-probs.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## ROUGE
|
||||||
|
|
||||||
|
ROUGE-1/2/L (precision, recall, F1) for summarization. Self-contained implementation with no external dependencies.
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/eval/evaluate_rouge.py \
|
||||||
|
--data_path predictions.jsonl \
|
||||||
|
--output results/rouge.json
|
||||||
|
```
|
||||||
|
|
||||||
|
| Parameter | Default | Description |
|
||||||
|
|-----------|---------|-------------|
|
||||||
|
| `--data_path` | required | JSONL with `reference`/`candidate` per line |
|
||||||
|
| `--output` | None | Output JSON path |
|
||||||
|
|
||||||
|
**Input**: JSONL, one object per line:
|
||||||
|
```json
|
||||||
|
{"reference": "Ground truth text", "candidate": "Model output text"}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Output**: stdout prints `rouge-1`, `rouge-2`, `rouge-l` each as P/R/F1. With `--output`, writes JSON with `aggregate` and `per_item` scores.
|
||||||
|
|
||||||
|
Can also be imported as a library:
|
||||||
|
```python
|
||||||
|
from scripts.eval.evaluate_rouge import compute_rouge
|
||||||
|
scores = compute_rouge(reference, candidate)
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## IFD (Instruction-Following Difficulty)
|
||||||
|
|
||||||
|
Data quality metric: `IFD = L_conditional / L_unconditional`. Measures how much harder it is to predict a response given its instruction vs. without it. Useful for filtering instruction-tuning data.
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/eval/evaluate_ifd.py \
|
||||||
|
--param_path ./params \
|
||||||
|
--input_path sft_data.jsonl \
|
||||||
|
--output_dir ifd_results/ \
|
||||||
|
--format messages \
|
||||||
|
--batch_size 8
|
||||||
|
```
|
||||||
|
|
||||||
|
| Parameter | Default | Description |
|
||||||
|
|-----------|---------|-------------|
|
||||||
|
| `--param_path` | required | Model directory |
|
||||||
|
| `--input_path` | required | Input file, glob, or directory |
|
||||||
|
| `--output_dir` | required | Output directory |
|
||||||
|
| `--max_len` | 2048 | Max token length |
|
||||||
|
| `--format` | `plain` | `plain` (instruction/response fields) or `messages` (chat format) |
|
||||||
|
| `--instr_key` | `instruction` | Instruction field key (plain format) |
|
||||||
|
| `--resp_key` | `response` | Response field key (plain format) |
|
||||||
|
| `--batch_size` | 8 | Items per model-forward flush |
|
||||||
|
| `--device` | auto | Device |
|
||||||
|
| `--dtype` | auto | Torch dtype |
|
||||||
|
| `--sentinel_text` | `\n` | Prefix for unconditional pass (`""` → bos/pad fallback) |
|
||||||
|
| `--per_token` | False | Include per-token IFD breakdown |
|
||||||
|
| `--max_samples` | None | Random subsample per file |
|
||||||
|
|
||||||
|
**How it works**: Two forward passes per batch — (1) conditional: packed BFD sequence with context + response, (2) unconditional: response prefixed with a sentinel. IFD = mean_conditional_loss / mean_unconditional_loss. IFD > 1 means the instruction makes the response harder to predict (higher quality data).
|
||||||
|
|
||||||
|
**Output**: Per-file `<label>_ifd.jsonl` with IFD scores per item. `summary.json` aggregates per-file stats.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## IFEval (Instruction Following)
|
||||||
|
|
||||||
|
Google's IFEval benchmark: generates responses and verifies 27 types of constraints (keywords, format, length, case, punctuation, etc.).
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/eval/evaluate_ifeval.py \
|
||||||
|
--param_path ./params \
|
||||||
|
--num_samples 1 \
|
||||||
|
--max_tokens 512 \
|
||||||
|
--output results/ifeval.json
|
||||||
|
```
|
||||||
|
|
||||||
|
| Parameter | Default | Description |
|
||||||
|
|-----------|---------|-------------|
|
||||||
|
| `--param_path` | `./params` | Model directory |
|
||||||
|
| `--data_path` | `./ifeval/input_data.jsonl` | IFEval JSONL (auto-downloaded if missing) |
|
||||||
|
| `--output` | None | Output JSON path |
|
||||||
|
| `--max_tokens` | 512 | Max generation tokens |
|
||||||
|
| `--temperature` | 0.1 | Sampling temperature (low for instruction-following) |
|
||||||
|
| `--top_p` | 0.95 | Top-p sampling |
|
||||||
|
| `--top_k` | 50 | Top-k sampling |
|
||||||
|
| `--num_samples` | 1 | Samples per problem (best-of-n scoring) |
|
||||||
|
| `--batch_size` | 1 | Inference batch size |
|
||||||
|
| `--limit` | None | Limit to first N problems (quick testing) |
|
||||||
|
| `--dump_responses` | None | Path to dump raw responses as JSONL |
|
||||||
|
|
||||||
|
**Output**: stdout prints overall accuracy + per-constraint-type accuracy table. With `--output`, writes per-problem results + `_summary`.
|
||||||
|
|
||||||
|
**Data**: Auto-downloads `google/IFEval` from HuggingFace. Each problem has `key`, `prompt`, `instruction_id_list`, `kwargs`.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Weight Analysis
|
||||||
|
|
||||||
|
SVD-based effective rank and weight statistics for checkpoint diagnostics. Does not load the model graph or run any forward pass.
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/eval/analyze_weights.py \
|
||||||
|
--ckpt_dir ./checkpoint/epoch_1_step_6000 \
|
||||||
|
--output results/weights.json
|
||||||
|
```
|
||||||
|
|
||||||
|
| Parameter | Default | Description |
|
||||||
|
|-----------|---------|-------------|
|
||||||
|
| `--ckpt_dir` | required | Checkpoint dir with `model.safetensors` + `config.json` |
|
||||||
|
| `--compare` | None | Additional checkpoint dirs to compare |
|
||||||
|
| `--no_svd` | False | Skip SVD; show only weight stats (faster) |
|
||||||
|
| `--output` | None | Save results as JSON |
|
||||||
|
| `--device` | `cuda` | Device for SVD |
|
||||||
|
|
||||||
|
**Output**: SVD effective rank by component (ER@90/95/99%, entropic rank, condition number), per-layer effective rank grid, and weight value statistics (mean/std/min/max). Provides a utilization verdict (HIGH >0.85 / MODERATE >0.5 / LOW).
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Tips
|
||||||
|
|
||||||
|
- **Quick test**: Use `--limit` (IFEval) or `--problems` (HumanEval) to run on a small subset first.
|
||||||
|
- **Auto-download**: HumanEval, MMLU, and IFEval auto-download their datasets on first run. The other scripts expect user-provided data.
|
||||||
|
- **Output formats**: `--output` writes a single JSON for most scripts. PPL and IFD write an `--output_dir` containing `summary.json` plus per-file artifacts.
|
||||||
|
- **CPU mode**: All scripts auto-detect CUDA. To force CPU, use `--device cpu --dtype float32`.
|
||||||
|
|
||||||
|
> Document Update Time: 2026-07-30
|
||||||
@@ -4,6 +4,7 @@
|
|||||||
|
|
||||||
- [KV Cache](#kv-cache)
|
- [KV Cache](#kv-cache)
|
||||||
- [KVCache System](#kvcache-system)
|
- [KVCache System](#kvcache-system)
|
||||||
|
- [Attention Backend](#attention-backend)
|
||||||
- [Continuous Batching](#continuous-batching)
|
- [Continuous Batching](#continuous-batching)
|
||||||
- [Sampling](#sampling-strategy-pattern)
|
- [Sampling](#sampling-strategy-pattern)
|
||||||
- [Protocol Handlers](#protocol-handlers-strategy-pattern)
|
- [Protocol Handlers](#protocol-handlers-strategy-pattern)
|
||||||
@@ -23,30 +24,70 @@ RoPE is applied **before** KV cache write, not after — otherwise position enco
|
|||||||
|
|
||||||
## KVCache System
|
## KVCache System
|
||||||
|
|
||||||
Seven classes working together, with two concrete cache implementations:
|
Three-layer separation (SGLang-inspired): storage, index table, allocator.
|
||||||
|
|
||||||
### ContiguousCache (default)
|
|
||||||
|
|
||||||
```
|
```
|
||||||
ContiguousCache (simple contiguous per-slot cache)
|
PagePool (top-level manager, orchestrates all layers)
|
||||||
├── ContiguousCacheView bundles k/v tensors + slot indices for attention layers
|
├── KVStorage k_buffer / v_buffer [n_layers, size, n_kv_heads, head_dim]
|
||||||
|
├── ReqToTokenPool req_to_token [num_reqs, max_ctx_len] → physical token slot
|
||||||
|
├── Allocator bitmask-based page allocator + ref-count + LRU (paged mode only)
|
||||||
|
└── PrefixCache hash-based prefix matching (paged mode only)
|
||||||
```
|
```
|
||||||
|
|
||||||
Created by default when no cache is passed to `InferenceScheduler`. Each task occupies a fixed slot of `[max_seq_len, num_key_value_heads, head_dim]`. Simple and efficient for small-to-medium batch sizes.
|
`PagePool` supports two modes:
|
||||||
|
|
||||||
### PageCache (paged with prefix sharing)
|
- **Contiguous (default)**: pre-allocates `max_batch_size * max_seq_len` token slots. `req_to_token` is a trivial linear mapping (`slot = req_idx * max_seq_len + pos`). No dynamic allocation.
|
||||||
|
- **Paged** (`page_size=1` or `>1` with `n_tokens` set): shared token pool with on-demand allocation. Allocator + PrefixCache enable prefix sharing and LRU eviction.
|
||||||
|
|
||||||
|
`bind_tasks()` returns a `KVCache` dataclass — pure data, no methods:
|
||||||
|
|
||||||
```
|
```
|
||||||
PageCache (paged KV cache with prefix sharing, alternative)
|
KVCache
|
||||||
├── PagePool orchestrates page allocation + prefix matching
|
├── k_buffer, v_buffer [n_layers, size, n_kv_heads, head_dim]
|
||||||
│ ├── Allocator bitmask-based page allocator + ref-count + LRU
|
├── req_to_token [num_reqs, max_ctx_len]
|
||||||
│ └── PrefixCache hash-based prefix matching (page_hash via polynomial hash)
|
├── req_pool_indices [batch_size]
|
||||||
├── TaskTable maps task_id → page_table + cached token count
|
├── seq_lens [batch_size]
|
||||||
├── Storage k_cache / v_cache tensors (num_hidden_layers × n_pages × page_size × num_key_value_heads × head_dim)
|
├── out_cache_loc [batch, seq_len] — write indices for this forward
|
||||||
└── PageCacheView bundles Storage + page_table + total_len for attention layers
|
├── max_len int — max(seq_lens), avoids GPU sync in decode
|
||||||
|
├── page_table [batch, max_len] — precomputed gather indices for decode (None for prefill)
|
||||||
|
└── decode_mask [batch, max_len] bool — precomputed position validity mask (None for single-batch decode)
|
||||||
```
|
```
|
||||||
|
|
||||||
`isinstance(cache, KVCache)` checks dispatch to the correct view. Both implement the abstract `KVCache` interface used by `Executor` and `InferenceScheduler`.
|
Attention layers do raw buffer indexing: `k_buffer[layer_id, out_cache_loc] = k` to write, `k_buffer[layer_id, indices]` to gather.
|
||||||
|
|
||||||
|
## Attention Backend
|
||||||
|
|
||||||
|
Attention computation (cache I/O + SDPA/kernel dispatch) is decoupled from the model via `AttentionBackend` ABC:
|
||||||
|
|
||||||
|
```
|
||||||
|
AttentionBackend (ABC)
|
||||||
|
├── TorchNativeBackend SDPA + indirect KV cache gather (default)
|
||||||
|
└── CudaBackend CUDA kernel dispatch (attn_paged_decode, attn_prefill)
|
||||||
|
```
|
||||||
|
|
||||||
|
Select via context manager (mirrors `torch.nn.attention.sdpa_kernel`):
|
||||||
|
|
||||||
|
```python
|
||||||
|
from astrai.extension import attn_backend, ATTN_BACKEND
|
||||||
|
|
||||||
|
with attn_backend(ATTN_BACKEND.CUDA):
|
||||||
|
engine.generate("hello")
|
||||||
|
```
|
||||||
|
|
||||||
|
`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` prefill path: writes K/V, gathers full-sequence K/V via indirect indexing (same as `TorchNativeBackend`), then calls `attn_prefill`.
|
||||||
|
|
||||||
|
Fallback: `CudaBackend` delegates to `TorchNativeBackend` when a CUDA kernel is not available.
|
||||||
|
|
||||||
|
### Rotary Embedding Backend
|
||||||
|
|
||||||
|
Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/rotary_backend.py`, which auto-dispatches:
|
||||||
|
|
||||||
|
- **CUDA kernel** (`rotary_emb.cu`): fused cos/sin lookup + rotation in a single kernel, used when the kernel is available, input is on CUDA, and `torch.is_grad_enabled()` is `False` (inference mode)
|
||||||
|
- **Torch fallback**: complex multiply path (`torch.view_as_complex` → `torch.complex` multiply → `torch.view_as_real`), used during training (supports autograd backward) or when the CUDA kernel is not available
|
||||||
|
|
||||||
|
`RotaryEmbedding` stores `cos_table`/`sin_table` as f32 buffers and returns a `(cos, sin)` tuple from `forward()`. Both attention backends share the same rotary dispatch — it is backend-agnostic.
|
||||||
|
|
||||||
## Continuous Batching
|
## Continuous Batching
|
||||||
|
|
||||||
@@ -154,6 +195,10 @@ Supports `stop_sequences` and streaming via `event: content_block_delta`.
|
|||||||
| `temperature` | float | 1.0 | Sampling temperature (> 0.0) |
|
| `temperature` | float | 1.0 | Sampling temperature (> 0.0) |
|
||||||
| `max_tokens` | Optional[int] | None | Max generation length |
|
| `max_tokens` | Optional[int] | None | Max generation length |
|
||||||
| `stream` | bool | False | Stream output |
|
| `stream` | bool | False | Stream output |
|
||||||
|
| `stop` | Optional[Union[str, List[str]]] | None | Stop sequences |
|
||||||
|
| `frequency_penalty` | float | 0.0 | Frequency penalty |
|
||||||
|
| `tools` | Optional[List[dict]] | None | Tool definitions for function calling |
|
||||||
|
| `tool_choice` | Optional[str] | None | Tool selection mode |
|
||||||
|
|
||||||
### SSE Streaming Format
|
### SSE Streaming Format
|
||||||
|
|
||||||
@@ -249,4 +294,4 @@ async for token in engine.generate_async("Hello", ...): # -> AsyncGenerator[s
|
|||||||
print(token)
|
print(token)
|
||||||
```
|
```
|
||||||
|
|
||||||
> Document Update Time: 2026-07-09
|
> Document Update Time: 2026-07-31
|
||||||
@@ -28,18 +28,48 @@
|
|||||||
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
|
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
|
||||||
| `--max_grad_norm` | Maximum gradient norm for clipping (None disables) | 1.0 |
|
| `--max_grad_norm` | Maximum gradient norm for clipping (None disables) | 1.0 |
|
||||||
|
|
||||||
### Optimizer (MuonMix)
|
### Optimizer
|
||||||
|
|
||||||
Combined optimizer: matrix parameters via **Muon**, non-matrix via **AdamW** (`fused=True`).
|
The default `muon_adamw` optimizer sends matrix parameters through **Muon** and
|
||||||
|
non-matrix parameters through **AdamW** (`fused=True`).
|
||||||
|
|
||||||
| Parameter | Description | Default |
|
| Parameter | Description | Default |
|
||||||
|-----------|-------------|---------|
|
|-----------|-------------|---------|
|
||||||
|
| `--optimizer` | Built-in optimizer (`muon_adamw`, `nora_nadamw`, `mano_adamw`) | `muon_adamw` |
|
||||||
| `--weight_decay` | Weight decay (applied to Muon matrix params; non-matrix use 0) | 0.1 |
|
| `--weight_decay` | Weight decay (applied to Muon matrix params; non-matrix use 0) | 0.1 |
|
||||||
| `--muon_momentum` | Muon momentum factor | 0.95 |
|
| `--muon_momentum` | Muon momentum factor | 0.95 |
|
||||||
| `--muon_nesterov` | Enable Nesterov momentum for Muon | True |
|
| `--muon_nesterov` | Enable Nesterov momentum for Muon | True |
|
||||||
| `--muon_ns_steps` | Newton-Schulz iteration steps for Muon | 5 |
|
| `--muon_ns_steps` | Newton-Schulz iteration steps for Muon | 5 |
|
||||||
| `--muon_adjust_lr` | Muon LR adjustment strategy (`original`, `match_rms_adamw`) | `match_rms_adamw` |
|
| `--muon_adjust_lr` | Muon LR adjustment strategy (`original`, `match_rms_adamw`) | `match_rms_adamw` |
|
||||||
|
|
||||||
|
`nora_nadamw` routes internal `Linear.weight` matrices to **Nora** and
|
||||||
|
embeddings, the LM head, norms, biases, LoRA factors, and fallback parameters to
|
||||||
|
**NAdamW**. Parameters are classified by module role and identity, so tied
|
||||||
|
embedding/head weights occur in exactly one group. Nora requires complete rows
|
||||||
|
under DTensor sharding and rejects layouts sharded along the last dimension.
|
||||||
|
|
||||||
|
| Parameter | Description | Default |
|
||||||
|
|-----------|-------------|---------|
|
||||||
|
| `--nora_lr` | Nora learning rate | 5e-3 |
|
||||||
|
| `--nora_beta` | Nora momentum-buffer EMA factor | 0.95 |
|
||||||
|
| `--nora_momentum` | Nora Nesterov interpolation factor | 0.95 |
|
||||||
|
| `--nora_weight_decay` | Nora matrix weight decay | 0.0 |
|
||||||
|
|
||||||
|
`mano_adamw` routes internal `Linear.weight` matrices to **Mano** (manifold
|
||||||
|
normalized optimizer) and the remaining parameters to **NAdamW**. Mano projects
|
||||||
|
the momentum onto the tangent space of the Oblique manifold and normalizes it,
|
||||||
|
alternating the projection axis (row/column) each step — replacing Muon's
|
||||||
|
Newton-Schulz iteration with a cheaper normalization.
|
||||||
|
|
||||||
|
| Parameter | Description | Default |
|
||||||
|
|-----------|-------------|---------|
|
||||||
|
| `--mano_momentum` | Mano momentum factor | 0.95 |
|
||||||
|
| `--mano_nesterov` | Enable Nesterov momentum for Mano | True |
|
||||||
|
|
||||||
|
Optimizer identity and hyperparameters are saved in checkpoint metadata. Optimizer
|
||||||
|
states are intentionally not interchangeable: resume older MuonAdamW checkpoints
|
||||||
|
with `--optimizer=muon_adamw`.
|
||||||
|
|
||||||
### Data Loading
|
### Data Loading
|
||||||
|
|
||||||
| Parameter | Description | Default |
|
| Parameter | Description | Default |
|
||||||
@@ -84,7 +114,7 @@ Combined optimizer: matrix parameters via **Muon**, non-matrix via **AdamW** (`f
|
|||||||
| Parameter | Description | Default |
|
| Parameter | Description | Default |
|
||||||
|-----------|-------------|---------|
|
|-----------|-------------|---------|
|
||||||
| `--nprocs` | Number of GPUs / processes | 1 |
|
| `--nprocs` | Number of GPUs / processes | 1 |
|
||||||
| `--parallel_mode` | Parallel strategy (`none`, `ddp`, or `fsdp`) | none |
|
| `--parallel_mode` | Parallel strategy (`none`, `ddp`, `fsdp`) | fsdp |
|
||||||
| `--device_type` | Device type | cuda |
|
| `--device_type` | Device type | cuda |
|
||||||
| `--start_method` | Multiprocessing start method (`spawn`, `fork`, `forkserver`) | spawn |
|
| `--start_method` | Multiprocessing start method (`spawn`, `fork`, `forkserver`) | spawn |
|
||||||
| `--backend` | Distributed training backend | nccl |
|
| `--backend` | Distributed training backend | nccl |
|
||||||
@@ -164,6 +194,7 @@ nohup python scripts/tools/train.py \
|
|||||||
| `--device` | str | `cuda` | Device to load model on |
|
| `--device` | str | `cuda` | Device to load model on |
|
||||||
| `--dtype` | str | `bfloat16` | Model weights dtype (`bfloat16`, `float16`, `float32`) |
|
| `--dtype` | str | `bfloat16` | Model weights dtype (`bfloat16`, `float16`, `float32`) |
|
||||||
| `--max_batch_size` | int | `16` | Maximum batch size for continuous batching |
|
| `--max_batch_size` | int | `16` | Maximum batch size for continuous batching |
|
||||||
|
| `--max_seq_len` | int | model config `max_position_embeddings` | Maximum sequence length (KV cache size + prompt truncation) |
|
||||||
| `--reload` | flag | `False` | Enable auto-reload for development |
|
| `--reload` | flag | `False` | Enable auto-reload for development |
|
||||||
|
|
||||||
Usage:
|
Usage:
|
||||||
@@ -173,6 +204,14 @@ python scripts/tools/server.py --param_path ./params --device cuda --dtype bfloa
|
|||||||
|
|
||||||
See [Inference Guide](inference.md) for HTTP API documentation.
|
See [Inference Guide](inference.md) for HTTP API documentation.
|
||||||
|
|
||||||
|
# Preprocess
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/tools/preprocess.py data/*.jsonl -o output/ -c config.json
|
||||||
|
```
|
||||||
|
|
||||||
|
See [Preprocessing Guide](preprocessing.md) for config file format and examples.
|
||||||
|
|
||||||
## Generate (`generate.py`)
|
## Generate (`generate.py`)
|
||||||
|
|
||||||
| Parameter | Type | Default | Description |
|
| Parameter | Type | Default | Description |
|
||||||
@@ -186,7 +225,11 @@ See [Inference Guide](inference.md) for HTTP API documentation.
|
|||||||
| `--top_k` | int | `30` | Top-k filtering |
|
| `--top_k` | int | `30` | Top-k filtering |
|
||||||
| `--top_p` | float | `0.95` | Nucleus sampling threshold |
|
| `--top_p` | float | `0.95` | Nucleus sampling threshold |
|
||||||
| `--batch_size` | int | `1` | Batch size for generation |
|
| `--batch_size` | int | `1` | Batch size for generation |
|
||||||
|
| `--num_samples` | int | `1` | Responses per prompt |
|
||||||
| `--max_tokens` | int | model config `max_position_embeddings` | Maximum tokens to generate |
|
| `--max_tokens` | int | model config `max_position_embeddings` | Maximum tokens to generate |
|
||||||
|
| `--cache_len` | int | `2048` | KV cache length |
|
||||||
|
| `--frequency_penalty` | float | `0.0` | Frequency penalty |
|
||||||
|
| `--rep_window` | int | `64` | Window size for frequency penalty |
|
||||||
|
|
||||||
Usage:
|
Usage:
|
||||||
```bash
|
```bash
|
||||||
@@ -256,6 +256,7 @@ When `sources` is set, `sections` is ignored.
|
|||||||
| `min_chars` | int | `50` | Skip text-mode items shorter than this |
|
| `min_chars` | int | `50` | Skip text-mode items shorter than this |
|
||||||
| `max_chars` | int | `2000000` | Skip text-mode items longer than this |
|
| `max_chars` | int | `2000000` | Skip text-mode items longer than this |
|
||||||
| `max_items` | int or null | `null` | Stop after N documents |
|
| `max_items` | int or null | `null` | Stop after N documents |
|
||||||
|
| `batch_size` | int | `256` | Records per tokenization batch |
|
||||||
| `packing_strategy` | str | `"simple"` | Packing strategy: `"simple"`, `"bfd"`, `"bfd_split"` |
|
| `packing_strategy` | str | `"simple"` | Packing strategy: `"simple"`, `"bfd"`, `"bfd_split"` |
|
||||||
| `max_packed_len` | int | `8192` | Maximum length of a packed bin |
|
| `max_packed_len` | int | `8192` | Maximum length of a packed bin |
|
||||||
| `truncation_mode` | str | `"keep_start"` | How to truncate sequences: `"keep_start"` or `"keep_end"` |
|
| `truncation_mode` | str | `"keep_start"` | How to truncate sequences: `"keep_start"` or `"keep_end"` |
|
||||||
@@ -265,7 +266,7 @@ When `sources` is set, `sections` is ignored.
|
|||||||
| Field | Type | Default | Description |
|
| Field | Type | Default | Description |
|
||||||
|-------|------|---------|-------------|
|
|-------|------|---------|-------------|
|
||||||
| `domain_key` | str or null | `null` | JSONL key for domain grouping |
|
| `domain_key` | str or null | `null` | JSONL key for domain grouping |
|
||||||
| `storage_format` | str | `"bin"` | `"bin"` (mmap) or `"h5"` |
|
| `storage_format` | str | `"bin"` | `"bin"` (mmap). Reading also supports `"jsonl"` for on-the-fly tokenization |
|
||||||
| `max_tokens_per_shard` | int | `100000000` | Flush threshold in cumulative tokens |
|
| `max_tokens_per_shard` | int | `100000000` | Flush threshold in cumulative tokens |
|
||||||
| `dtype` | dict[str, str] | `{}` | Per-key tensor dtype override (e.g. `{"loss_mask": "bool"}`) |
|
| `dtype` | dict[str, str] | `{}` | Per-key tensor dtype override (e.g. `{"loss_mask": "bool"}`) |
|
||||||
| `position_ids_mode` | str | `"doc_reset"` | How to compute position_ids: `"none"`, `"doc_reset"`, `"continuous"` |
|
| `position_ids_mode` | str | `"doc_reset"` | How to compute position_ids: `"none"`, `"doc_reset"`, `"continuous"` |
|
||||||
@@ -41,7 +41,7 @@ RoPE embeds position into Q/K vectors via complex rotation:
|
|||||||
|
|
||||||
$$ q_i = R_i W_q x_i, \quad k_j = R_j W_k x_j, \quad q_i^T k_j = x_i^T W_q^T R_{i-j} W_k x_j $$
|
$$ q_i = R_i W_q x_i, \quad k_j = R_j W_k x_j, \quad q_i^T k_j = x_i^T W_q^T R_{i-j} W_k x_j $$
|
||||||
|
|
||||||
The complex rotation `freqs_cis` is pre-computed once (`cos, sin` pairs per position). `apply_rotary_emb` multiplies Q/K as complex numbers.
|
`RotaryEmbedding` pre-computes `cos_table` and `sin_table` (f32, `[max_len, dim/2]`). `forward()` returns a `(cos, sin)` tuple indexed by `position_ids`. `apply_rotary_emb` applies the rotation: during training it uses torch complex multiply (autograd-compatible); during inference it auto-dispatches to a fused CUDA kernel when available.
|
||||||
|
|
||||||
## Training Loop
|
## Training Loop
|
||||||
|
|
||||||
@@ -232,4 +232,4 @@ nohup python scripts/tools/train.py \
|
|||||||
|
|
||||||
Full parameter reference at [params.md](params.md).
|
Full parameter reference at [params.md](params.md).
|
||||||
|
|
||||||
> Document Update Time: 2026-07-20
|
> Document Update Time: 2026-07-31
|
||||||
|
Before Width: | Height: | Size: 281 KiB After Width: | Height: | Size: 281 KiB |
+4
-7
@@ -8,7 +8,6 @@ name = "astrai"
|
|||||||
readme = "README.md"
|
readme = "README.md"
|
||||||
requires-python = ">=3.12"
|
requires-python = ">=3.12"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"h5py==3.15.1",
|
|
||||||
"numpy==2.4.4",
|
"numpy==2.4.4",
|
||||||
"torch==2.11.0",
|
"torch==2.11.0",
|
||||||
"tokenizers==0.21.4",
|
"tokenizers==0.21.4",
|
||||||
@@ -16,10 +15,11 @@ dependencies = [
|
|||||||
"safetensors==0.5.3",
|
"safetensors==0.5.3",
|
||||||
"huggingface-hub==0.34.3",
|
"huggingface-hub==0.34.3",
|
||||||
"jinja2>=3.0.0",
|
"jinja2>=3.0.0",
|
||||||
|
"pydantic>=2.0",
|
||||||
"fastapi",
|
"fastapi",
|
||||||
"uvicorn[standard]",
|
"uvicorn[standard]",
|
||||||
"httpx",
|
"click>=8.0",
|
||||||
"requests",
|
"pyyaml>=6.0",
|
||||||
]
|
]
|
||||||
keywords = ["nlp", "datasets", "language-models", "machine-learning"]
|
keywords = ["nlp", "datasets", "language-models", "machine-learning"]
|
||||||
license = { text = "GPL-3.0" }
|
license = { text = "GPL-3.0" }
|
||||||
@@ -31,14 +31,11 @@ classifiers = [
|
|||||||
urls = { Homepage = "https://github.com/ViperEkura/AstrAI" }
|
urls = { Homepage = "https://github.com/ViperEkura/AstrAI" }
|
||||||
|
|
||||||
[project.optional-dependencies]
|
[project.optional-dependencies]
|
||||||
dev = ["pytest==9.0.2", "ruff"]
|
dev = ["pytest==9.0.2", "ruff", "httpx2"]
|
||||||
|
|
||||||
[tool.setuptools.packages.find]
|
[tool.setuptools.packages.find]
|
||||||
where = ["."]
|
where = ["."]
|
||||||
|
|
||||||
[tool.pip]
|
|
||||||
extra-index-url = "https://download.pytorch.org/whl/cu128"
|
|
||||||
|
|
||||||
[tool.setuptools.dynamic]
|
[tool.setuptools.dynamic]
|
||||||
version = { attr = "astrai.__version__" }
|
version = { attr = "astrai.__version__" }
|
||||||
|
|
||||||
|
|||||||
@@ -8,11 +8,11 @@ import safetensors.torch
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
|
||||||
def effective_rank_metrics(w: torch.Tensor) -> dict:
|
def effective_rank_metrics(w: torch.Tensor, device: str = "cpu") -> dict:
|
||||||
if w.ndim == 1:
|
if w.ndim == 1:
|
||||||
return {"shape": tuple(w.shape), "is_1d": True}
|
return {"shape": tuple(w.shape), "is_1d": True}
|
||||||
|
|
||||||
w = w.float()
|
w = w.float().to(device)
|
||||||
s = torch.linalg.svdvals(w)
|
s = torch.linalg.svdvals(w)
|
||||||
s_sq = s**2
|
s_sq = s**2
|
||||||
total = s_sq.sum()
|
total = s_sq.sum()
|
||||||
@@ -238,6 +238,12 @@ def main():
|
|||||||
default=None,
|
default=None,
|
||||||
help="Save results as JSON to this path.",
|
help="Save results as JSON to this path.",
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--device",
|
||||||
|
type=str,
|
||||||
|
default="cuda",
|
||||||
|
help="Device for SVD computation (e.g., 'cuda:0', 'cpu').",
|
||||||
|
)
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
all_results = {}
|
all_results = {}
|
||||||
@@ -277,10 +283,12 @@ def main():
|
|||||||
|
|
||||||
results = {}
|
results = {}
|
||||||
if not args.no_svd:
|
if not args.no_svd:
|
||||||
print(f"Computing SVD on {len(weight_keys)} tensors...")
|
print(
|
||||||
|
f"Computing SVD on {len(weight_keys)} tensors (device={args.device})..."
|
||||||
|
)
|
||||||
for i, k in enumerate(sorted(weight_keys)):
|
for i, k in enumerate(sorted(weight_keys)):
|
||||||
print(f" [{i + 1}/{len(weight_keys)}] {k:<60s}", end="\r")
|
print(f" [{i + 1}/{len(weight_keys)}] {k:<60s}", end="\r")
|
||||||
results[k] = effective_rank_metrics(sd[k])
|
results[k] = effective_rank_metrics(sd[k], device=args.device)
|
||||||
print()
|
print()
|
||||||
else:
|
else:
|
||||||
print(f"Computing stats on {len(weight_keys)} tensors (no SVD)...")
|
print(f"Computing stats on {len(weight_keys)} tensors (no SVD)...")
|
||||||
|
|||||||
@@ -57,6 +57,7 @@ class EvalConfig:
|
|||||||
top_p: float = 0.95
|
top_p: float = 0.95
|
||||||
top_k: int = 50
|
top_k: int = 50
|
||||||
batch_size: int = 32
|
batch_size: int = 32
|
||||||
|
max_seq_len: int = 4096
|
||||||
test_timeout: float = 3.0
|
test_timeout: float = 3.0
|
||||||
test_workers: int = 8
|
test_workers: int = 8
|
||||||
k_values: Tuple[int, ...] = (1, 10, 100)
|
k_values: Tuple[int, ...] = (1, 10, 100)
|
||||||
@@ -90,7 +91,9 @@ def save_json(path: str, data):
|
|||||||
json.dump(data, f, indent=2, ensure_ascii=False)
|
json.dump(data, f, indent=2, ensure_ascii=False)
|
||||||
|
|
||||||
|
|
||||||
def create_engine(param_path: str, batch_size: int) -> InferenceEngine:
|
def create_engine(
|
||||||
|
param_path: str, batch_size: int, max_seq_len: int
|
||||||
|
) -> InferenceEngine:
|
||||||
model = AutoModel.from_pretrained(param_path)
|
model = AutoModel.from_pretrained(param_path)
|
||||||
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
||||||
model.to(device="cuda", dtype=torch.bfloat16)
|
model.to(device="cuda", dtype=torch.bfloat16)
|
||||||
@@ -98,6 +101,7 @@ def create_engine(param_path: str, batch_size: int) -> InferenceEngine:
|
|||||||
model=model,
|
model=model,
|
||||||
tokenizer=tokenizer,
|
tokenizer=tokenizer,
|
||||||
max_batch_size=batch_size,
|
max_batch_size=batch_size,
|
||||||
|
max_seq_len=max_seq_len,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -318,7 +322,7 @@ def run_pipeline(cfg: EvalConfig) -> Dict:
|
|||||||
if cfg.problem_indices:
|
if cfg.problem_indices:
|
||||||
problems = [problems[i] for i in cfg.problem_indices if i < len(problems)]
|
problems = [problems[i] for i in cfg.problem_indices if i < len(problems)]
|
||||||
|
|
||||||
engine = create_engine(cfg.param_path, cfg.batch_size)
|
engine = create_engine(cfg.param_path, cfg.batch_size, cfg.max_seq_len)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
generated = generate_all(engine, problems, cfg)
|
generated = generate_all(engine, problems, cfg)
|
||||||
@@ -357,7 +361,8 @@ def parse_args(argv: Optional[List[str]] = None) -> EvalConfig:
|
|||||||
p.add_argument("--temperature", type=float, default=0.8)
|
p.add_argument("--temperature", type=float, default=0.8)
|
||||||
p.add_argument("--top_p", type=float, default=0.95)
|
p.add_argument("--top_p", type=float, default=0.95)
|
||||||
p.add_argument("--top_k", type=int, default=50)
|
p.add_argument("--top_k", type=int, default=50)
|
||||||
p.add_argument("--batch_size", type=int, default=32)
|
p.add_argument("--batch_size", type=int, default=64)
|
||||||
|
p.add_argument("--max_seq_len", type=int, default=4096)
|
||||||
p.add_argument("--test_workers", type=int, default=8)
|
p.add_argument("--test_workers", type=int, default=8)
|
||||||
p.add_argument("--test_timeout", type=float, default=3.0)
|
p.add_argument("--test_timeout", type=float, default=3.0)
|
||||||
p.add_argument("--problems", type=int, nargs="+", default=None)
|
p.add_argument("--problems", type=int, nargs="+", default=None)
|
||||||
@@ -375,6 +380,7 @@ def parse_args(argv: Optional[List[str]] = None) -> EvalConfig:
|
|||||||
top_p=args.top_p,
|
top_p=args.top_p,
|
||||||
top_k=args.top_k,
|
top_k=args.top_k,
|
||||||
batch_size=args.batch_size,
|
batch_size=args.batch_size,
|
||||||
|
max_seq_len=args.max_seq_len,
|
||||||
test_workers=args.test_workers,
|
test_workers=args.test_workers,
|
||||||
test_timeout=args.test_timeout,
|
test_timeout=args.test_timeout,
|
||||||
problem_indices=args.problems,
|
problem_indices=args.problems,
|
||||||
|
|||||||
@@ -9,10 +9,16 @@ v2 changelog:
|
|||||||
- Same token set: unconditional pass prefixes resp with a plain-text sentinel
|
- Same token set: unconditional pass prefixes resp with a plain-text sentinel
|
||||||
(default ``\\n``; use ``--sentinel_text ""`` for bos/pad fallback).
|
(default ``\\n``; use ``--sentinel_text ""`` for bos/pad fallback).
|
||||||
Both branches predict the identical N resp tokens.
|
Both branches predict the identical N resp tokens.
|
||||||
Single-token answers (rl=1) are now supported.
|
- Single-token answers (rl=1) are now supported.
|
||||||
- ctx_len tracked in output
|
- ctx_len tracked in output
|
||||||
- skip_reason for None samples (no more silent None)
|
- skip_reason for None samples (no more silent None)
|
||||||
- --per_token for per-token IFD breakdown
|
- --per_token for per-token IFD breakdown
|
||||||
|
|
||||||
|
v3 changelog:
|
||||||
|
- Append EOS at the end of response in both conditional and unconditional
|
||||||
|
passes (``--append_eos`` / ``--no-append_eos``, default: enabled).
|
||||||
|
The model now also predicts when the response should end, which is part
|
||||||
|
of instruction following. Falls back gracefully when tokenizer has no EOS.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import argparse
|
import argparse
|
||||||
@@ -237,6 +243,7 @@ def process_file(
|
|||||||
sentinel_ids=None,
|
sentinel_ids=None,
|
||||||
per_token=False,
|
per_token=False,
|
||||||
max_samples=None,
|
max_samples=None,
|
||||||
|
eos_ids=None,
|
||||||
):
|
):
|
||||||
"""Score a single file, write per-sample JSONL, return summary stats."""
|
"""Score a single file, write per-sample JSONL, return summary stats."""
|
||||||
if device is None:
|
if device is None:
|
||||||
@@ -245,6 +252,11 @@ def process_file(
|
|||||||
if sentinel_ids is None:
|
if sentinel_ids is None:
|
||||||
sentinel_ids = _resolve_sentinel_ids(tokenizer, "\n")
|
sentinel_ids = _resolve_sentinel_ids(tokenizer, "\n")
|
||||||
|
|
||||||
|
if eos_ids is None:
|
||||||
|
eos_ids = []
|
||||||
|
|
||||||
|
eos_len = len(eos_ids)
|
||||||
|
|
||||||
data = _load_items(input_file)
|
data = _load_items(input_file)
|
||||||
|
|
||||||
if max_samples and len(data) > max_samples:
|
if max_samples and len(data) > max_samples:
|
||||||
@@ -267,7 +279,9 @@ def process_file(
|
|||||||
ctx_text = "\n\n".join(m["content"] for m in item["messages"][:i])
|
ctx_text = "\n\n".join(m["content"] for m in item["messages"][:i])
|
||||||
ctx_ids = tokenizer.encode(ctx_text)
|
ctx_ids = tokenizer.encode(ctx_text)
|
||||||
resp_ids = tokenizer.encode(msg["content"], add_special_tokens=False)
|
resp_ids = tokenizer.encode(msg["content"], add_special_tokens=False)
|
||||||
ctx_ids, resp_ids = _trim(ctx_ids, resp_ids, max_len)
|
ctx_ids, resp_ids = _trim(ctx_ids, resp_ids, max_len - eos_len)
|
||||||
|
if eos_ids and resp_ids and resp_ids[-1:] != eos_ids:
|
||||||
|
resp_ids = resp_ids + eos_ids
|
||||||
if ctx_ids and resp_ids:
|
if ctx_ids and resp_ids:
|
||||||
turns.append((ctx_ids, resp_ids))
|
turns.append((ctx_ids, resp_ids))
|
||||||
if not turns:
|
if not turns:
|
||||||
@@ -284,7 +298,7 @@ def process_file(
|
|||||||
else:
|
else:
|
||||||
ctx_ids = tokenizer.encode(item[instr_key], add_special_tokens=False)
|
ctx_ids = tokenizer.encode(item[instr_key], add_special_tokens=False)
|
||||||
resp_ids = tokenizer.encode(item[resp_key], add_special_tokens=False)
|
resp_ids = tokenizer.encode(item[resp_key], add_special_tokens=False)
|
||||||
ctx_ids, resp_ids = _trim(ctx_ids, resp_ids, max_len)
|
ctx_ids, resp_ids = _trim(ctx_ids, resp_ids, max_len - eos_len)
|
||||||
if not ctx_ids or not resp_ids:
|
if not ctx_ids or not resp_ids:
|
||||||
results.append(
|
results.append(
|
||||||
{
|
{
|
||||||
@@ -294,6 +308,8 @@ def process_file(
|
|||||||
}
|
}
|
||||||
)
|
)
|
||||||
continue
|
continue
|
||||||
|
if eos_ids and resp_ids[-1:] != eos_ids:
|
||||||
|
resp_ids = resp_ids + eos_ids
|
||||||
buffer.append((item, [(ctx_ids, resp_ids)], "plain"))
|
buffer.append((item, [(ctx_ids, resp_ids)], "plain"))
|
||||||
|
|
||||||
if len(buffer) >= batch_size:
|
if len(buffer) >= batch_size:
|
||||||
@@ -452,6 +468,11 @@ def main():
|
|||||||
default=None,
|
default=None,
|
||||||
help="Maximum number of samples per file (random subsample). Default: all.",
|
help="Maximum number of samples per file (random subsample). Default: all.",
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--append_eos/--no-append_eos",
|
||||||
|
default=True,
|
||||||
|
help="Append EOS token at the end of response in both passes (default: enabled).",
|
||||||
|
)
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
if args.device is None:
|
if args.device is None:
|
||||||
@@ -466,6 +487,16 @@ def main():
|
|||||||
|
|
||||||
sentinel_ids = _resolve_sentinel_ids(tokenizer, args.sentinel_text)
|
sentinel_ids = _resolve_sentinel_ids(tokenizer, args.sentinel_text)
|
||||||
|
|
||||||
|
eos_ids = []
|
||||||
|
if args.append_eos:
|
||||||
|
eos_token_id = getattr(tokenizer, "eos_token_id", None)
|
||||||
|
if eos_token_id is not None:
|
||||||
|
eos_ids = [eos_token_id]
|
||||||
|
else:
|
||||||
|
print(
|
||||||
|
"Warning: --append_eos enabled but tokenizer has no EOS token; skipping."
|
||||||
|
)
|
||||||
|
|
||||||
input_files = _collect_input_files(args.input_path)
|
input_files = _collect_input_files(args.input_path)
|
||||||
if not input_files:
|
if not input_files:
|
||||||
print(f"No input files found at {args.input_path}")
|
print(f"No input files found at {args.input_path}")
|
||||||
@@ -493,6 +524,7 @@ def main():
|
|||||||
sentinel_ids=sentinel_ids,
|
sentinel_ids=sentinel_ids,
|
||||||
per_token=args.per_token,
|
per_token=args.per_token,
|
||||||
max_samples=args.max_samples,
|
max_samples=args.max_samples,
|
||||||
|
eos_ids=eos_ids,
|
||||||
)
|
)
|
||||||
all_stats[label] = stats
|
all_stats[label] = stats
|
||||||
|
|
||||||
|
|||||||
@@ -509,7 +509,10 @@ def main():
|
|||||||
help="Number of samples per problem (best-of-n scoring)",
|
help="Number of samples per problem (best-of-n scoring)",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--batch_size", type=int, default=1, help="Inference batch size"
|
"--batch_size", type=int, default=64, help="Inference batch size"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--max_seq_len", type=int, default=4096, help="Max sequence length for KV cache"
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--limit",
|
"--limit",
|
||||||
@@ -542,6 +545,7 @@ def main():
|
|||||||
model=model,
|
model=model,
|
||||||
tokenizer=tokenizer,
|
tokenizer=tokenizer,
|
||||||
max_batch_size=args.batch_size,
|
max_batch_size=args.batch_size,
|
||||||
|
max_seq_len=args.max_seq_len,
|
||||||
)
|
)
|
||||||
|
|
||||||
results = evaluate(
|
results = evaluate(
|
||||||
|
|||||||
@@ -179,31 +179,53 @@ def apply_chat(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def choice_logprob(
|
def choice_logprobs_batched(
|
||||||
model, tokenizer, context_ids: list[int], choice_letter: str, device: str
|
model,
|
||||||
) -> float:
|
tokenizer,
|
||||||
choice_text = choice_letter
|
context_ids_list: list[list[int]],
|
||||||
choice_ids = tokenizer.encode(choice_text, add_special_tokens=False)
|
device: str,
|
||||||
input_ids = context_ids + choice_ids
|
max_model_len: int,
|
||||||
max_len = model.config.max_position_embeddings
|
) -> list[dict[str, float]]:
|
||||||
if len(input_ids) > max_len:
|
"""Compute log-probs for multiple questions x 4 choices in batches.
|
||||||
overflow = len(input_ids) - max_len
|
|
||||||
input_ids = input_ids[overflow:]
|
Returns a list of dicts: [{A: score, B: score, C: score, D: score}, ...]
|
||||||
ctx_len = len(input_ids) - len(choice_ids)
|
"""
|
||||||
else:
|
letters = ("A", "B", "C", "D")
|
||||||
ctx_len = len(context_ids)
|
choice_ids_list = [tokenizer.encode(c, add_special_tokens=False) for c in letters]
|
||||||
|
|
||||||
|
all_inputs: list[tuple[int, int, list[int], int, list[int]]] = []
|
||||||
|
for qi, context_ids in enumerate(context_ids_list):
|
||||||
|
for ci, choice_ids in enumerate(choice_ids_list):
|
||||||
|
input_ids = context_ids + choice_ids
|
||||||
|
if len(input_ids) > max_model_len:
|
||||||
|
overflow = len(input_ids) - max_model_len
|
||||||
|
input_ids = input_ids[overflow:]
|
||||||
|
ctx_len = len(input_ids) - len(choice_ids)
|
||||||
|
else:
|
||||||
|
ctx_len = len(context_ids)
|
||||||
|
all_inputs.append((qi, ci, input_ids, ctx_len, choice_ids))
|
||||||
|
|
||||||
|
n = len(all_inputs)
|
||||||
|
max_input_len = max(len(x[2]) for x in all_inputs)
|
||||||
|
padded = torch.zeros(n, max_input_len, dtype=torch.long, device=device)
|
||||||
|
mask = torch.zeros(n, max_input_len, dtype=torch.bool, device=device)
|
||||||
|
for i, (_, _, ids, _, _) in enumerate(all_inputs):
|
||||||
|
padded[i, : len(ids)] = torch.tensor(ids, dtype=torch.long, device=device)
|
||||||
|
mask[i, : len(ids)] = True
|
||||||
|
|
||||||
input_tensor = torch.tensor([input_ids], device=device, dtype=torch.long)
|
|
||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
logits = model(input_tensor)["logits"][0]
|
logits = model(padded, input_mask=mask)["logits"]
|
||||||
|
|
||||||
score = 0.0
|
results = [{} for _ in range(len(context_ids_list))]
|
||||||
for i, tid in enumerate(choice_ids):
|
for i, (qi, ci, _, ctx_len, choice_ids) in enumerate(all_inputs):
|
||||||
pos = ctx_len - 1 + i
|
score = 0.0
|
||||||
if pos >= len(logits):
|
for j, tid in enumerate(choice_ids):
|
||||||
break
|
pos = ctx_len - 1 + j
|
||||||
score += F.log_softmax(logits[pos], dim=-1)[tid].item()
|
if pos >= logits.size(1):
|
||||||
return score
|
break
|
||||||
|
score += F.log_softmax(logits[i, pos].float(), dim=-1)[tid].item()
|
||||||
|
results[qi][letters[ci]] = score
|
||||||
|
return results
|
||||||
|
|
||||||
|
|
||||||
def _permute_choices(item: dict, rng: random.Random) -> tuple[dict, str]:
|
def _permute_choices(item: dict, rng: random.Random) -> tuple[dict, str]:
|
||||||
@@ -233,25 +255,42 @@ def evaluate_subject(
|
|||||||
device: str,
|
device: str,
|
||||||
n_shot: int,
|
n_shot: int,
|
||||||
seed: int = 0,
|
seed: int = 0,
|
||||||
|
batch_size: int = 16,
|
||||||
) -> tuple[float, int, int]:
|
) -> tuple[float, int, int]:
|
||||||
rng = random.Random(seed) if seed >= 0 else None
|
rng = random.Random(seed) if seed >= 0 else None
|
||||||
correct = 0
|
correct = 0
|
||||||
total = 0
|
total = 0
|
||||||
for item in tqdm.tqdm(test_data, desc=f"{subject:40s}", leave=False):
|
|
||||||
|
context_ids_list = []
|
||||||
|
answers = []
|
||||||
|
for item in test_data:
|
||||||
if rng is not None:
|
if rng is not None:
|
||||||
permuted, answer = _permute_choices(item, rng)
|
permuted, answer = _permute_choices(item, rng)
|
||||||
else:
|
else:
|
||||||
permuted, answer = item, item["answer"]
|
permuted, answer = item, item["answer"]
|
||||||
raw_prompt = build_prompt(permuted["question"], permuted, subject)
|
raw_prompt = build_prompt(permuted["question"], permuted, subject)
|
||||||
context = apply_chat(tokenizer, raw_prompt, n_shot, dev_data or [], subject)
|
context = apply_chat(tokenizer, raw_prompt, n_shot, dev_data or [], subject)
|
||||||
context_ids = tokenizer.encode(context)
|
context_ids_list.append(tokenizer.encode(context))
|
||||||
scores = {
|
answers.append(answer)
|
||||||
c: choice_logprob(model, tokenizer, context_ids, c, device)
|
|
||||||
for c in ("A", "B", "C", "D")
|
max_model_len = model.config.max_position_embeddings
|
||||||
}
|
|
||||||
if max(scores, key=scores.get) == answer:
|
num_batches = (len(context_ids_list) + batch_size - 1) // batch_size
|
||||||
correct += 1
|
for start in tqdm.tqdm(
|
||||||
total += 1
|
range(0, len(context_ids_list), batch_size),
|
||||||
|
total=num_batches,
|
||||||
|
desc=f"{subject:40s}",
|
||||||
|
leave=False,
|
||||||
|
):
|
||||||
|
batch = context_ids_list[start : start + batch_size]
|
||||||
|
batch_answers = answers[start : start + batch_size]
|
||||||
|
scores_list = choice_logprobs_batched(
|
||||||
|
model, tokenizer, batch, device, max_model_len
|
||||||
|
)
|
||||||
|
for scores, answer in zip(scores_list, batch_answers):
|
||||||
|
if max(scores, key=scores.get) == answer:
|
||||||
|
correct += 1
|
||||||
|
total += 1
|
||||||
return correct / total, correct, total
|
return correct / total, correct, total
|
||||||
|
|
||||||
|
|
||||||
@@ -290,6 +329,12 @@ def main():
|
|||||||
default=0,
|
default=0,
|
||||||
help="Seed for option permutation (0 to enable, -1 to disable)",
|
help="Seed for option permutation (0 to enable, -1 to disable)",
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--batch_size",
|
||||||
|
type=int,
|
||||||
|
default=4,
|
||||||
|
help="Number of questions per batch (4 choices each = 4*B rows)",
|
||||||
|
)
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
if args.download or not os.path.exists(args.data_dir):
|
if args.download or not os.path.exists(args.data_dir):
|
||||||
@@ -329,6 +374,7 @@ def main():
|
|||||||
device,
|
device,
|
||||||
args.n_shot,
|
args.n_shot,
|
||||||
seed=args.seed,
|
seed=args.seed,
|
||||||
|
batch_size=args.batch_size,
|
||||||
)
|
)
|
||||||
results[subject] = {"accuracy": round(acc, 4), "correct": corr, "total": tot}
|
results[subject] = {"accuracy": round(acc, 4), "correct": corr, "total": tot}
|
||||||
total_correct += corr
|
total_correct += corr
|
||||||
|
|||||||
@@ -415,7 +415,7 @@ if __name__ == "__main__":
|
|||||||
help="Key for the text field in the input data.",
|
help="Key for the text field in the input data.",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--batch_size", type=int, default=4, help="Batch size for evaluation."
|
"--batch_size", type=int, default=64, help="Batch size for evaluation."
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--max_length",
|
"--max_length",
|
||||||
|
|||||||
+214
-236
@@ -1,299 +1,277 @@
|
|||||||
"""Benchmark AutoRegressiveLM with KVCache"""
|
from pathlib import Path
|
||||||
|
from typing import Optional
|
||||||
import argparse
|
|
||||||
from dataclasses import dataclass
|
|
||||||
from typing import Any, Dict
|
|
||||||
|
|
||||||
|
import click
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from astrai import setup_logging
|
||||||
from astrai.config import AutoRegressiveLMConfig
|
from astrai.config import AutoRegressiveLMConfig
|
||||||
from astrai.inference import ContiguousCache, PageCache
|
from astrai.extension import ATTN_BACKEND, attn_backend
|
||||||
from astrai.model.transformer import AutoRegressiveLM
|
from astrai.inference.core.cache import PagePool
|
||||||
|
from astrai.model import AutoModel
|
||||||
|
|
||||||
|
_DTYPES = ["bfloat16", "float16", "float32"]
|
||||||
|
_CACHES = ["contiguous", "paged"]
|
||||||
|
DEFAULT_CKPT = str(Path(__file__).resolve().parents[2] / "ckpt_bucket" / "kami-15bt")
|
||||||
|
CACHE_MAX_SEQ = 2048
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class BenchmarkResult:
|
class BenchmarkResult:
|
||||||
total_tokens: int
|
def __init__(
|
||||||
total_time: float
|
self,
|
||||||
tokens_per_second: float
|
name: str,
|
||||||
metadata: Dict[str, Any]
|
batch_size: int,
|
||||||
|
seq_len: int,
|
||||||
|
tokens_per_second: float,
|
||||||
|
latency_ms: float,
|
||||||
|
metadata: Optional[dict] = None,
|
||||||
|
):
|
||||||
|
self.name = name
|
||||||
|
self.batch_size = batch_size
|
||||||
|
self.seq_len = seq_len
|
||||||
|
self.tokens_per_second = tokens_per_second
|
||||||
|
self.latency_ms = latency_ms
|
||||||
|
self.metadata = metadata or {}
|
||||||
|
|
||||||
|
|
||||||
class GenerationBenchmark:
|
class GenerationBenchmark:
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
|
model: AutoModel,
|
||||||
config: AutoRegressiveLMConfig,
|
config: AutoRegressiveLMConfig,
|
||||||
device: str = "cuda",
|
device: str = "cuda",
|
||||||
dtype: torch.dtype = torch.bfloat16,
|
dtype: torch.dtype = torch.bfloat16,
|
||||||
cache_type: str = "contiguous",
|
cache_type: str = "contiguous",
|
||||||
):
|
):
|
||||||
self.config = config
|
|
||||||
self.device = device
|
self.device = device
|
||||||
self.dtype = dtype
|
self.dtype = dtype
|
||||||
self.cache_type = cache_type
|
self.cache_type = cache_type
|
||||||
self.model = AutoRegressiveLM(config).to(device=device, dtype=dtype)
|
self.model = model
|
||||||
self.model.eval()
|
self.config = config
|
||||||
|
|
||||||
@torch.inference_mode()
|
def _make_pool(self, batch_size: int) -> PagePool:
|
||||||
|
return PagePool(
|
||||||
|
n_layers=self.config.num_hidden_layers,
|
||||||
|
n_kv_heads=self.config.num_key_value_heads,
|
||||||
|
head_dim=self.config.hidden_size // self.config.num_attention_heads,
|
||||||
|
max_batch_size=batch_size,
|
||||||
|
max_seq_len=CACHE_MAX_SEQ,
|
||||||
|
device=self.device,
|
||||||
|
dtype=self.dtype,
|
||||||
|
page_size=1,
|
||||||
|
n_tokens=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _run_prefill(self, pool: PagePool, batch_size: int, prompt_len: int) -> list:
|
||||||
|
input_ids = torch.randint(
|
||||||
|
0, self.config.vocab_size, (batch_size, prompt_len), device=self.device
|
||||||
|
)
|
||||||
|
position_ids = (
|
||||||
|
torch.arange(0, prompt_len, dtype=torch.long, device=self.device)
|
||||||
|
.unsqueeze(0)
|
||||||
|
.expand(batch_size, -1)
|
||||||
|
)
|
||||||
|
input_mask = position_ids.unsqueeze(-1) >= torch.arange(
|
||||||
|
prompt_len, device=self.device
|
||||||
|
)
|
||||||
|
|
||||||
|
task_ids = [f"bench_{i}" for i in range(batch_size)]
|
||||||
|
for tid in task_ids:
|
||||||
|
pool.task_alloc(tid, list(range(prompt_len)))
|
||||||
|
|
||||||
|
kv_cache = pool.bind_tasks(
|
||||||
|
task_ids, [prompt_len] * batch_size, self.device, start_pos=0
|
||||||
|
)
|
||||||
|
with torch.inference_mode(), attn_backend(ATTN_BACKEND.CUDA):
|
||||||
|
self.model(
|
||||||
|
input_ids,
|
||||||
|
input_mask=input_mask,
|
||||||
|
kv_cache=kv_cache,
|
||||||
|
position_ids=position_ids,
|
||||||
|
)
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
return task_ids
|
||||||
|
|
||||||
|
def _run_decode_step(self, pool: PagePool, task_ids: list, seq_len: int):
|
||||||
|
batch_size = len(task_ids)
|
||||||
|
input_ids = torch.randint(
|
||||||
|
0, self.config.vocab_size, (batch_size, 1), device=self.device
|
||||||
|
)
|
||||||
|
position_ids = torch.tensor(
|
||||||
|
[[seq_len] for _ in range(batch_size)], dtype=torch.long, device=self.device
|
||||||
|
)
|
||||||
|
total_len = seq_len + 1
|
||||||
|
input_mask = position_ids[:, :, None] >= torch.arange(
|
||||||
|
total_len, device=self.device
|
||||||
|
)
|
||||||
|
kv_cache = pool.bind_tasks(task_ids, [seq_len + 1] * batch_size, self.device)
|
||||||
|
with torch.inference_mode(), attn_backend(ATTN_BACKEND.CUDA):
|
||||||
|
self.model(
|
||||||
|
input_ids,
|
||||||
|
input_mask=input_mask,
|
||||||
|
kv_cache=kv_cache,
|
||||||
|
position_ids=position_ids,
|
||||||
|
)
|
||||||
|
|
||||||
def run_prefill_benchmark(
|
def run_prefill_benchmark(
|
||||||
self,
|
self,
|
||||||
batch_size: int = 1,
|
batch_size: int = 4,
|
||||||
prompt_length: int = 512,
|
prompt_length: int = 512,
|
||||||
num_trials: int = 10,
|
num_trials: int = 5,
|
||||||
) -> BenchmarkResult:
|
) -> BenchmarkResult:
|
||||||
for _ in range(3):
|
import time
|
||||||
prompt_ids = torch.randint(
|
|
||||||
0,
|
|
||||||
self.config.vocab_size,
|
|
||||||
(batch_size, prompt_length),
|
|
||||||
device=self.device,
|
|
||||||
dtype=torch.long,
|
|
||||||
)
|
|
||||||
_ = self.model(prompt_ids)
|
|
||||||
torch.cuda.synchronize()
|
|
||||||
|
|
||||||
total_time = 0.0
|
input_ids = torch.randint(
|
||||||
total_tokens = batch_size * prompt_length * num_trials
|
0, self.config.vocab_size, (batch_size, prompt_length), device=self.device
|
||||||
|
)
|
||||||
|
position_ids = (
|
||||||
|
torch.arange(0, prompt_length, dtype=torch.long, device=self.device)
|
||||||
|
.unsqueeze(0)
|
||||||
|
.expand(batch_size, -1)
|
||||||
|
)
|
||||||
|
|
||||||
for trial in range(num_trials):
|
for _ in range(3):
|
||||||
prompt_ids = torch.randint(
|
with torch.inference_mode(), attn_backend(ATTN_BACKEND.CUDA):
|
||||||
0,
|
self.model(input_ids, position_ids=position_ids)
|
||||||
self.config.vocab_size,
|
|
||||||
(batch_size, prompt_length),
|
|
||||||
device=self.device,
|
|
||||||
dtype=torch.long,
|
|
||||||
)
|
|
||||||
start = torch.cuda.Event(enable_timing=True)
|
|
||||||
end = torch.cuda.Event(enable_timing=True)
|
|
||||||
start.record()
|
|
||||||
_ = self.model(prompt_ids)
|
|
||||||
end.record()
|
|
||||||
torch.cuda.synchronize()
|
|
||||||
|
|
||||||
trial_time = start.elapsed_time(end) / 1000
|
|
||||||
total_time += trial_time
|
|
||||||
|
|
||||||
print(
|
|
||||||
f" Trial {trial + 1}/{num_trials}: {prompt_length} tokens in {trial_time:.3f}s "
|
|
||||||
f"({prompt_length / trial_time:.1f} tok/s)"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
t0 = time.perf_counter()
|
||||||
|
for _ in range(num_trials):
|
||||||
|
with torch.inference_mode(), attn_backend(ATTN_BACKEND.CUDA):
|
||||||
|
self.model(input_ids, position_ids=position_ids)
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
elapsed = time.perf_counter() - t0
|
||||||
|
tokens = batch_size * prompt_length * num_trials
|
||||||
|
tps = tokens / elapsed
|
||||||
return BenchmarkResult(
|
return BenchmarkResult(
|
||||||
total_tokens=total_tokens,
|
name="prefill",
|
||||||
total_time=total_time,
|
batch_size=batch_size,
|
||||||
tokens_per_second=total_tokens / total_time,
|
seq_len=prompt_length,
|
||||||
metadata={
|
tokens_per_second=tps,
|
||||||
"benchmark_type": "prefill",
|
latency_ms=elapsed / num_trials * 1000,
|
||||||
"batch_size": batch_size,
|
metadata={"benchmark_type": "prefill", "num_trials": num_trials},
|
||||||
"prompt_length": prompt_length,
|
|
||||||
"dtype": str(self.dtype),
|
|
||||||
"device": self.device,
|
|
||||||
"cache": "none",
|
|
||||||
},
|
|
||||||
)
|
)
|
||||||
|
|
||||||
@torch.inference_mode()
|
|
||||||
def run_decoding_benchmark(
|
def run_decoding_benchmark(
|
||||||
self,
|
self,
|
||||||
batch_size: int = 1,
|
batch_size: int = 4,
|
||||||
prompt_length: int = 512,
|
prompt_length: int = 512,
|
||||||
gen_length: int = 128,
|
gen_length: int = 128,
|
||||||
num_trials: int = 5,
|
num_trials: int = 5,
|
||||||
) -> BenchmarkResult:
|
) -> BenchmarkResult:
|
||||||
total_time = 0.0
|
import time
|
||||||
total_tokens = batch_size * gen_length * num_trials
|
|
||||||
|
|
||||||
for trial in range(num_trials):
|
pool = self._make_pool(batch_size)
|
||||||
prompt_ids = torch.randint(
|
task_ids = self._run_prefill(pool, batch_size, prompt_length)
|
||||||
0,
|
|
||||||
self.config.vocab_size,
|
|
||||||
(batch_size, prompt_length),
|
|
||||||
device=self.device,
|
|
||||||
dtype=torch.long,
|
|
||||||
)
|
|
||||||
gen_ids = torch.randint(
|
|
||||||
0,
|
|
||||||
self.config.vocab_size,
|
|
||||||
(batch_size, gen_length),
|
|
||||||
device=self.device,
|
|
||||||
dtype=torch.long,
|
|
||||||
)
|
|
||||||
|
|
||||||
head_dim = self.config.hidden_size // self.config.num_attention_heads
|
for i in range(5):
|
||||||
max_seq = prompt_length + gen_length
|
self._run_decode_step(pool, task_ids, prompt_length + i)
|
||||||
|
torch.cuda.synchronize()
|
||||||
if self.cache_type == "contiguous":
|
|
||||||
cache = ContiguousCache(
|
|
||||||
self.config.num_hidden_layers,
|
|
||||||
batch_size,
|
|
||||||
max_seq,
|
|
||||||
self.config.num_key_value_heads,
|
|
||||||
head_dim,
|
|
||||||
self.device,
|
|
||||||
self.dtype,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
page_size = 128
|
|
||||||
n_pages = (max_seq + page_size - 1) // page_size * batch_size
|
|
||||||
cache = PageCache(
|
|
||||||
self.config.num_hidden_layers,
|
|
||||||
n_pages,
|
|
||||||
page_size,
|
|
||||||
self.config.num_key_value_heads,
|
|
||||||
head_dim,
|
|
||||||
self.device,
|
|
||||||
self.dtype,
|
|
||||||
)
|
|
||||||
|
|
||||||
task_ids = [f"b{i}" for i in range(batch_size)]
|
|
||||||
for tid in task_ids:
|
|
||||||
cache.task_alloc(tid, [0] * max_seq)
|
|
||||||
for p in range(max_seq):
|
|
||||||
cache.task_extend(tid, p)
|
|
||||||
|
|
||||||
cv = cache.bind_tasks(task_ids, prompt_length, self.device)
|
|
||||||
_ = self.model(
|
|
||||||
prompt_ids,
|
|
||||||
paged_cache=cv,
|
|
||||||
position_ids=torch.arange(
|
|
||||||
prompt_length, dtype=torch.long, device=self.device
|
|
||||||
)
|
|
||||||
.unsqueeze(0)
|
|
||||||
.expand(batch_size, -1),
|
|
||||||
)
|
|
||||||
torch.cuda.synchronize()
|
|
||||||
|
|
||||||
start = torch.cuda.Event(enable_timing=True)
|
|
||||||
end = torch.cuda.Event(enable_timing=True)
|
|
||||||
start.record()
|
|
||||||
|
|
||||||
for i in range(gen_length):
|
|
||||||
pos = prompt_length + i
|
|
||||||
cv = cache.bind_tasks(task_ids, pos + 1, self.device)
|
|
||||||
_ = self.model(
|
|
||||||
gen_ids[:, i : i + 1],
|
|
||||||
paged_cache=cv,
|
|
||||||
position_ids=torch.full(
|
|
||||||
(batch_size, 1),
|
|
||||||
pos,
|
|
||||||
dtype=torch.long,
|
|
||||||
device=self.device,
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
end.record()
|
|
||||||
torch.cuda.synchronize()
|
|
||||||
|
|
||||||
for tid in task_ids:
|
|
||||||
cache.task_free(tid)
|
|
||||||
|
|
||||||
trial_time = start.elapsed_time(end) / 1000
|
|
||||||
total_time += trial_time
|
|
||||||
|
|
||||||
print(
|
|
||||||
f" Trial {trial + 1}/{num_trials}: {gen_length} tokens in {trial_time:.3f}s "
|
|
||||||
f"({gen_length / trial_time:.1f} tok/s)"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
t0 = time.perf_counter()
|
||||||
|
for i in range(gen_length * num_trials):
|
||||||
|
self._run_decode_step(pool, task_ids, prompt_length + 5 + i)
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
elapsed = time.perf_counter() - t0
|
||||||
|
tokens = batch_size * gen_length * num_trials
|
||||||
|
tps = tokens / elapsed
|
||||||
return BenchmarkResult(
|
return BenchmarkResult(
|
||||||
total_tokens=total_tokens,
|
name="decode",
|
||||||
total_time=total_time,
|
batch_size=batch_size,
|
||||||
tokens_per_second=total_tokens / total_time,
|
seq_len=gen_length,
|
||||||
|
tokens_per_second=tps,
|
||||||
|
latency_ms=elapsed / (gen_length * num_trials) * 1000,
|
||||||
metadata={
|
metadata={
|
||||||
"benchmark_type": "decoding",
|
"benchmark_type": "decode",
|
||||||
"batch_size": batch_size,
|
"num_trials": num_trials,
|
||||||
"prompt_length": prompt_length,
|
"prompt_length": prompt_length,
|
||||||
"gen_length": gen_length,
|
|
||||||
"dtype": str(self.dtype),
|
|
||||||
"device": self.device,
|
|
||||||
"cache": self.cache_type,
|
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def print_benchmark_result(result: BenchmarkResult):
|
def print_benchmark_result(result: BenchmarkResult) -> None:
|
||||||
btype = result.metadata["benchmark_type"]
|
print("-" * 80)
|
||||||
print(f"\n{' ' + btype.upper() + ' Benchmark ':-^80}")
|
print(f"{result.name.upper()} — Batch={result.batch_size}, SeqLen={result.seq_len}")
|
||||||
print(f"Total Tokens Processed: {result.total_tokens:,}")
|
print(f" Throughput : {result.tokens_per_second:.1f} tokens/s")
|
||||||
print(f"Time Consumed: {result.total_time:.3f}s")
|
print(f" Latency : {result.latency_ms:.2f} ms/step")
|
||||||
print(f"Throughput: {result.tokens_per_second:,.1f} tok/s")
|
|
||||||
for k, v in result.metadata.items():
|
for k, v in result.metadata.items():
|
||||||
if k != "benchmark_type":
|
if k != "benchmark_type":
|
||||||
print(f"{k.replace('_', ' ').title()}: {v}")
|
print(f" {k.replace('_', ' ').title()}: {v}")
|
||||||
print("-" * 80)
|
print("-" * 80)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
@click.command(name="benchmark", help="Benchmark model throughput and latency.")
|
||||||
parser = argparse.ArgumentParser(description="AutoRegressiveLM benchmark")
|
@click.option("--device", default="cuda", help="Device.")
|
||||||
parser.add_argument(
|
@click.option(
|
||||||
"--device", type=str, default="cuda", help="Device (default: cuda)"
|
"--dtype", type=click.Choice(_DTYPES), default="bfloat16", help="Data type."
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
@click.option(
|
||||||
"--dtype",
|
"--cache", type=click.Choice(_CACHES), default="contiguous", help="KV cache type."
|
||||||
type=str,
|
)
|
||||||
default="bfloat16",
|
@click.option("--batch_size", type=int, default=4, help="Batch size.")
|
||||||
choices=["bfloat16", "float16", "float32"],
|
@click.option("--prompt_length", type=int, default=512, help="Prompt length.")
|
||||||
help="Dtype",
|
@click.option("--gen_length", type=int, default=128, help="Generation length.")
|
||||||
)
|
@click.option("--num_trials", type=int, default=5, help="Number of trials.")
|
||||||
parser.add_argument(
|
@click.option("--prefill_only", is_flag=True, help="Prefill benchmark only.")
|
||||||
"--cache",
|
@click.option("--decode_only", is_flag=True, help="Decode benchmark only.")
|
||||||
type=str,
|
@click.option(
|
||||||
default="contiguous",
|
"--ckpt",
|
||||||
choices=["contiguous", "paged"],
|
default=DEFAULT_CKPT,
|
||||||
help="KV cache type",
|
help="Checkpoint directory.",
|
||||||
)
|
)
|
||||||
parser.add_argument("--batch_size", type=int, default=4, help="Batch size")
|
def benchmark_command(
|
||||||
parser.add_argument("--prompt_length", type=int, default=512, help="Prompt length")
|
device: str,
|
||||||
parser.add_argument("--gen_length", type=int, default=128, help="Generation length")
|
dtype: str,
|
||||||
parser.add_argument("--num_trials", type=int, default=5, help="Number of trials")
|
cache: str,
|
||||||
parser.add_argument(
|
batch_size: int,
|
||||||
"--prefill_only", action="store_true", help="Run prefill benchmark only"
|
prompt_length: int,
|
||||||
)
|
gen_length: int,
|
||||||
parser.add_argument(
|
num_trials: int,
|
||||||
"--decode_only", action="store_true", help="Run decoding benchmark only"
|
prefill_only: bool,
|
||||||
)
|
decode_only: bool,
|
||||||
args = parser.parse_args()
|
ckpt: str,
|
||||||
|
) -> None:
|
||||||
dtype_map = {
|
"""Benchmark model throughput and latency."""
|
||||||
|
dtype_map: dict[str, torch.dtype] = {
|
||||||
"bfloat16": torch.bfloat16,
|
"bfloat16": torch.bfloat16,
|
||||||
"float16": torch.float16,
|
"float16": torch.float16,
|
||||||
"float32": torch.float32,
|
"float32": torch.float32,
|
||||||
}
|
}
|
||||||
|
|
||||||
config = AutoRegressiveLMConfig(
|
click.echo(f"Loading model from {ckpt} ...")
|
||||||
vocab_size=10000,
|
config = AutoRegressiveLMConfig.from_file(str(Path(ckpt) / "config.json"))
|
||||||
hidden_size=1536,
|
model = AutoModel.from_pretrained(ckpt)
|
||||||
num_attention_heads=24,
|
model.to(device=device, dtype=dtype_map[dtype])
|
||||||
num_key_value_heads=4,
|
model.eval()
|
||||||
intermediate_size=6912,
|
|
||||||
max_position_embeddings=2048,
|
bench = GenerationBenchmark(
|
||||||
num_hidden_layers=24,
|
model=model,
|
||||||
rms_norm_eps=1e-5,
|
config=config,
|
||||||
|
device=device,
|
||||||
|
dtype=dtype_map[dtype],
|
||||||
|
cache_type=cache,
|
||||||
)
|
)
|
||||||
|
|
||||||
benchmark = GenerationBenchmark(
|
click.secho(f"Benchmark: device={device} dtype={dtype}", bold=True)
|
||||||
config, device=args.device, dtype=dtype_map[args.dtype], cache_type=args.cache
|
|
||||||
)
|
|
||||||
|
|
||||||
print("=" * 80)
|
if not decode_only:
|
||||||
print(
|
result = bench.run_prefill_benchmark(
|
||||||
f"Running AutoRegressiveLM Benchmark (device={args.device}, dtype={args.dtype})"
|
batch_size=batch_size,
|
||||||
)
|
prompt_length=prompt_length,
|
||||||
print("=" * 80)
|
num_trials=num_trials,
|
||||||
|
|
||||||
if not args.decode_only:
|
|
||||||
prefill_result = benchmark.run_prefill_benchmark(
|
|
||||||
batch_size=args.batch_size,
|
|
||||||
prompt_length=args.prompt_length,
|
|
||||||
num_trials=args.num_trials,
|
|
||||||
)
|
)
|
||||||
print_benchmark_result(prefill_result)
|
print_benchmark_result(result)
|
||||||
|
|
||||||
if not args.prefill_only:
|
if not prefill_only:
|
||||||
gen_result = benchmark.run_decoding_benchmark(
|
result = bench.run_decoding_benchmark(
|
||||||
batch_size=args.batch_size,
|
batch_size=batch_size,
|
||||||
prompt_length=args.prompt_length,
|
prompt_length=prompt_length,
|
||||||
gen_length=args.gen_length,
|
gen_length=gen_length,
|
||||||
num_trials=args.num_trials,
|
num_trials=num_trials,
|
||||||
)
|
)
|
||||||
print_benchmark_result(gen_result)
|
print_benchmark_result(result)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
setup_logging()
|
||||||
|
benchmark_command()
|
||||||
|
|||||||
+47
-101
@@ -1,11 +1,12 @@
|
|||||||
import argparse
|
|
||||||
import json
|
import json
|
||||||
import time
|
import time
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
|
import click
|
||||||
import torch
|
import torch
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
|
||||||
|
from astrai import setup_logging
|
||||||
from astrai.inference import InferenceEngine
|
from astrai.inference import InferenceEngine
|
||||||
from astrai.model import AutoModel
|
from astrai.model import AutoModel
|
||||||
from astrai.tokenize import AutoTokenizer
|
from astrai.tokenize import AutoTokenizer
|
||||||
@@ -20,10 +21,9 @@ def processor(
|
|||||||
top_p: float,
|
top_p: float,
|
||||||
question_key: str,
|
question_key: str,
|
||||||
response_key: str,
|
response_key: str,
|
||||||
max_tokens: Optional[int],
|
|
||||||
batch_size: int,
|
batch_size: int,
|
||||||
num_samples: int = 1,
|
num_samples: int = 1,
|
||||||
cache_len: int = 2048,
|
max_seq_len: Optional[int] = None,
|
||||||
frequency_penalty: float = 0.0,
|
frequency_penalty: float = 0.0,
|
||||||
rep_window: int = 64,
|
rep_window: int = 64,
|
||||||
):
|
):
|
||||||
@@ -38,8 +38,7 @@ def processor(
|
|||||||
model=model,
|
model=model,
|
||||||
tokenizer=tokenizer,
|
tokenizer=tokenizer,
|
||||||
max_batch_size=batch_size * num_samples,
|
max_batch_size=batch_size * num_samples,
|
||||||
max_seq_len=cache_len,
|
max_seq_len=max_seq_len,
|
||||||
max_prompt_len=cache_len,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
print(f"Reading {input_json_file} ...")
|
print(f"Reading {input_json_file} ...")
|
||||||
@@ -55,9 +54,6 @@ def processor(
|
|||||||
prompts = [item[question_key] for item in input_data]
|
prompts = [item[question_key] for item in input_data]
|
||||||
print(f" {len(prompts)} prompts loaded\n")
|
print(f" {len(prompts)} prompts loaded\n")
|
||||||
|
|
||||||
if max_tokens is None:
|
|
||||||
max_tokens = model.config.max_position_embeddings
|
|
||||||
|
|
||||||
chunk_size = max(1, batch_size)
|
chunk_size = max(1, batch_size)
|
||||||
|
|
||||||
with open(output_json_file, "w", encoding="utf-8") as f:
|
with open(output_json_file, "w", encoding="utf-8") as f:
|
||||||
@@ -74,7 +70,6 @@ def processor(
|
|||||||
resp_chunk = engine.generate(
|
resp_chunk = engine.generate(
|
||||||
prompt=chunk_expanded,
|
prompt=chunk_expanded,
|
||||||
stream=False,
|
stream=False,
|
||||||
max_tokens=max_tokens,
|
|
||||||
temperature=temperature,
|
temperature=temperature,
|
||||||
top_p=top_p,
|
top_p=top_p,
|
||||||
top_k=top_k,
|
top_k=top_k,
|
||||||
@@ -89,7 +84,6 @@ def processor(
|
|||||||
resp_chunk = engine.generate(
|
resp_chunk = engine.generate(
|
||||||
prompt=chunk,
|
prompt=chunk,
|
||||||
stream=False,
|
stream=False,
|
||||||
max_tokens=max_tokens,
|
|
||||||
temperature=temperature,
|
temperature=temperature,
|
||||||
top_p=top_p,
|
top_p=top_p,
|
||||||
top_k=top_k,
|
top_k=top_k,
|
||||||
@@ -121,95 +115,47 @@ def processor(
|
|||||||
engine.shutdown()
|
engine.shutdown()
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
@click.command(name="generate", help="Batch generation from a JSONL prompt file.")
|
||||||
parser = argparse.ArgumentParser(description="Batch generation from JSONL file.")
|
@click.option(
|
||||||
|
"--param_path",
|
||||||
parser.add_argument(
|
type=click.Path(exists=True),
|
||||||
"--param_path", type=str, required=True, help="Path to the model directory."
|
required=True,
|
||||||
)
|
help="Path to the model directory.",
|
||||||
parser.add_argument(
|
)
|
||||||
"--input_json_file",
|
@click.option(
|
||||||
type=str,
|
"--input_json_file",
|
||||||
required=True,
|
type=click.Path(exists=True),
|
||||||
help="Path to the input JSONL file.",
|
required=True,
|
||||||
)
|
help="Path to the input JSONL file.",
|
||||||
parser.add_argument(
|
)
|
||||||
"--output_json_file",
|
@click.option(
|
||||||
type=str,
|
"--output_json_file",
|
||||||
required=True,
|
type=click.Path(),
|
||||||
help="Path to the output JSONL file.",
|
required=True,
|
||||||
)
|
help="Path to the output JSONL file.",
|
||||||
parser.add_argument(
|
)
|
||||||
"--question_key",
|
@click.option(
|
||||||
type=str,
|
"--question_key", default="question", help="Key for the question in input JSON."
|
||||||
default="question",
|
)
|
||||||
help="Key for the question in the input JSON (default: question).",
|
@click.option(
|
||||||
)
|
"--response_key", default="response", help="Key for the response in output JSON."
|
||||||
parser.add_argument(
|
)
|
||||||
"--response_key",
|
@click.option("--temperature", type=float, default=0.8, help="Sampling temperature.")
|
||||||
type=str,
|
@click.option("--top_k", type=int, default=50, help="Top-k filtering.")
|
||||||
default="response",
|
@click.option("--top_p", type=float, default=0.95, help="Top-p filtering.")
|
||||||
help="Key for the response in the output JSON (default: response).",
|
@click.option("--batch_size", type=int, default=1, help="Batch size.")
|
||||||
)
|
@click.option("--num_samples", type=int, default=1, help="Responses per prompt.")
|
||||||
parser.add_argument(
|
@click.option("--max_seq_len", type=int, default=2048, help="KV cache length.")
|
||||||
"--temperature",
|
@click.option("--frequency_penalty", type=float, default=0.0, help="Frequency penalty.")
|
||||||
type=float,
|
@click.option(
|
||||||
default=0.60,
|
"--rep_window", type=int, default=64, help="Window size for frequency penalty."
|
||||||
help="Temperature for generating responses (default: 0.60).",
|
)
|
||||||
)
|
def generate_command(**kwargs):
|
||||||
parser.add_argument(
|
"""Batch generation from a JSONL prompt file."""
|
||||||
"--top_k",
|
|
||||||
type=int,
|
|
||||||
default=30,
|
|
||||||
help="Top-k value for generating responses (default: 30).",
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--top_p",
|
|
||||||
type=float,
|
|
||||||
default=0.95,
|
|
||||||
help="Top-p value for generating responses (default: 0.95).",
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--batch_size",
|
|
||||||
type=int,
|
|
||||||
default=1,
|
|
||||||
help="Batch size for generating responses (default: 1).",
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--num_samples",
|
|
||||||
type=int,
|
|
||||||
default=1,
|
|
||||||
help="Number of responses per prompt (expands batch internally, default: 1).",
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--max_tokens",
|
|
||||||
type=int,
|
|
||||||
default=None,
|
|
||||||
help=(
|
|
||||||
"Maximum tokens to generate "
|
|
||||||
"(default: model config max_position_embeddings)."
|
|
||||||
),
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--cache_len",
|
|
||||||
type=int,
|
|
||||||
default=2048,
|
|
||||||
help="KV cache & prompt truncation length (default: 2048, lower = less memory).",
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--frequency_penalty",
|
|
||||||
type=float,
|
|
||||||
default=0.0,
|
|
||||||
help="Frequency penalty to reduce repetition (default: 0.0, try 0.5-1.0).",
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--rep_window",
|
|
||||||
type=int,
|
|
||||||
default=64,
|
|
||||||
help="Window size for frequency penalty (default: 64).",
|
|
||||||
)
|
|
||||||
|
|
||||||
args = parser.parse_args()
|
|
||||||
|
|
||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
processor(**vars(args))
|
processor(**kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
setup_logging()
|
||||||
|
generate_command()
|
||||||
|
|||||||
+39
-35
@@ -1,48 +1,52 @@
|
|||||||
"""CLI: JSONL → tokenized .h5/.bin via config-driven Pipeline."""
|
"""CLI: JSONL → tokenized .bin via config-driven Pipeline."""
|
||||||
|
|
||||||
import argparse
|
import click
|
||||||
|
|
||||||
|
from astrai import setup_logging
|
||||||
from astrai.config.preprocess_config import PipelineConfig
|
from astrai.config.preprocess_config import PipelineConfig
|
||||||
from astrai.preprocessing.pipeline import Pipeline
|
from astrai.preprocessing.pipeline import Pipeline
|
||||||
|
|
||||||
|
|
||||||
def main():
|
@click.command(
|
||||||
parser = argparse.ArgumentParser(
|
name="preprocess", help="Tokenize and pack raw JSONL data into .bin format."
|
||||||
description="Raw JSONL → tokenized .h5/.bin via config-driven Pipeline"
|
)
|
||||||
)
|
@click.argument("inputs", nargs=-1, type=click.Path(exists=True), required=True)
|
||||||
parser.add_argument(
|
@click.option(
|
||||||
"inputs", nargs="+", metavar="JSONL", help="One or more JSONL files"
|
"--output_dir", "-o", type=click.Path(), required=True, help="Output directory."
|
||||||
)
|
)
|
||||||
parser.add_argument("--output_dir", "-o", required=True, help="Output directory")
|
@click.option(
|
||||||
parser.add_argument(
|
"--config",
|
||||||
"--config", "-c", required=True, help="Path to pipeline config JSON"
|
"-c",
|
||||||
)
|
"pipeline_config",
|
||||||
parser.add_argument(
|
type=click.Path(exists=True),
|
||||||
"--tokenizer_path",
|
required=True,
|
||||||
default="params",
|
help="Pipeline config JSON.",
|
||||||
help="Path to tokenizer directory (default: params)",
|
)
|
||||||
)
|
@click.option(
|
||||||
parser.add_argument(
|
"--tokenizer_path",
|
||||||
"--batch_size",
|
type=click.Path(exists=True),
|
||||||
type=int,
|
default="params",
|
||||||
default=None,
|
help="Path to tokenizer directory.",
|
||||||
help="Number of records tokenized together (default: config value)",
|
)
|
||||||
)
|
@click.option("--batch_size", type=int, default=None, help="Records per batch.")
|
||||||
args = parser.parse_args()
|
def preprocess_command(inputs, output_dir, pipeline_config, tokenizer_path, batch_size):
|
||||||
|
"""Tokenize and pack raw JSONL data into .bin format."""
|
||||||
config = PipelineConfig.from_file(args.config)
|
config = PipelineConfig.from_file(pipeline_config)
|
||||||
if args.batch_size is not None:
|
if batch_size is not None:
|
||||||
if args.batch_size < 1:
|
if batch_size < 1:
|
||||||
parser.error("--batch_size must be at least 1")
|
raise click.BadParameter("--batch_size must be at least 1")
|
||||||
config.preprocessing.batch_size = args.batch_size
|
config.preprocessing.batch_size = batch_size
|
||||||
|
|
||||||
|
click.echo(f"Preprocessing {len(inputs)} file(s) → {output_dir}")
|
||||||
Pipeline(
|
Pipeline(
|
||||||
config=config,
|
config=config,
|
||||||
input_paths=args.inputs,
|
input_paths=list(inputs),
|
||||||
output_dir=args.output_dir,
|
output_dir=output_dir,
|
||||||
tokenizer_path=args.tokenizer_path,
|
tokenizer_path=tokenizer_path,
|
||||||
).run()
|
).run()
|
||||||
|
click.echo("Done.")
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
main()
|
setup_logging()
|
||||||
|
preprocess_command()
|
||||||
|
|||||||
+50
-53
@@ -1,72 +1,69 @@
|
|||||||
import argparse
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
|
import click
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from astrai import setup_logging
|
||||||
from astrai.inference import run_server
|
from astrai.inference import run_server
|
||||||
|
|
||||||
|
_DTYPES = ["bfloat16", "float16", "float32"]
|
||||||
|
|
||||||
def main():
|
|
||||||
parser = argparse.ArgumentParser(description="Start AstrAI inference HTTP server")
|
|
||||||
parser.add_argument(
|
|
||||||
"--host", default="0.0.0.0", help="Host address (default: 0.0.0.0)"
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--port", type=int, default=8000, help="Port number (default: 8000)"
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--reload", action="store_true", help="Enable auto-reload for development"
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--param_path",
|
|
||||||
type=Path,
|
|
||||||
default=None,
|
|
||||||
help="Path to model parameters (default: project_root/params)",
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--device",
|
|
||||||
type=str,
|
|
||||||
default="cuda",
|
|
||||||
help="Device to load model on (default: cuda)",
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--dtype",
|
|
||||||
type=str,
|
|
||||||
default="bfloat16",
|
|
||||||
choices=["bfloat16", "float16", "float32"],
|
|
||||||
help="Data type for model weights (default: bfloat16)",
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--max_batch_size",
|
|
||||||
type=int,
|
|
||||||
default=16,
|
|
||||||
help="Maximum batch size for continuous batching (default: 16)",
|
|
||||||
)
|
|
||||||
args = parser.parse_args()
|
|
||||||
|
|
||||||
# Convert dtype string to torch dtype
|
@click.command(name="serve", help="Launch inference server (OpenAI-compatible API).")
|
||||||
|
@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(
|
||||||
|
"--param_path",
|
||||||
|
type=click.Path(exists=True),
|
||||||
|
default=None,
|
||||||
|
help="Path to model parameters.",
|
||||||
|
)
|
||||||
|
@click.option("--device", default="cuda", help="Device to load model on.")
|
||||||
|
@click.option(
|
||||||
|
"--dtype",
|
||||||
|
type=click.Choice(_DTYPES),
|
||||||
|
default="bfloat16",
|
||||||
|
help="Data type for model weights.",
|
||||||
|
)
|
||||||
|
@click.option(
|
||||||
|
"--max_batch_size",
|
||||||
|
type=int,
|
||||||
|
default=16,
|
||||||
|
help="Maximum batch size for continuous batching.",
|
||||||
|
)
|
||||||
|
@click.option(
|
||||||
|
"--max_seq_len",
|
||||||
|
type=int,
|
||||||
|
default=None,
|
||||||
|
help="Maximum sequence length (KV cache size + prompt truncation). Uses model config if not set.",
|
||||||
|
)
|
||||||
|
def server_command(
|
||||||
|
host, port, reload, param_path, device, dtype, max_batch_size, max_seq_len
|
||||||
|
):
|
||||||
|
"""Launch inference server (OpenAI-compatible API)."""
|
||||||
dtype_map = {
|
dtype_map = {
|
||||||
"bfloat16": torch.bfloat16,
|
"bfloat16": torch.bfloat16,
|
||||||
"float16": torch.float16,
|
"float16": torch.float16,
|
||||||
"float32": torch.float32,
|
"float32": torch.float32,
|
||||||
}
|
}
|
||||||
dtype = dtype_map[args.dtype]
|
|
||||||
|
|
||||||
project_root = Path(__file__).parent.parent.parent
|
project_root = Path(__file__).parent.parent.parent
|
||||||
param_path = args.param_path or (project_root / "params")
|
param_path = param_path or str(project_root / "params")
|
||||||
print(f"Starting AstrAI inference server on http://{args.host}:{args.port}")
|
|
||||||
print(f"Model parameters expected at: {param_path}")
|
click.echo(f"Starting server on http://{host}:{port}")
|
||||||
print(f"Device: {args.device}, Dtype: {args.dtype}")
|
click.echo(f"Model: {param_path} | Device: {device} | Dtype: {dtype}")
|
||||||
run_server(
|
run_server(
|
||||||
host=args.host,
|
host=host,
|
||||||
port=args.port,
|
port=port,
|
||||||
reload=args.reload,
|
reload=reload,
|
||||||
device=args.device,
|
device=device,
|
||||||
dtype=dtype,
|
dtype=dtype_map[dtype],
|
||||||
param_path=param_path,
|
param_path=Path(param_path),
|
||||||
max_batch_size=args.max_batch_size,
|
max_batch_size=max_batch_size,
|
||||||
|
max_seq_len=max_seq_len,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
main()
|
setup_logging()
|
||||||
|
server_command()
|
||||||
|
|||||||
+592
-383
File diff suppressed because it is too large
Load Diff
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user