Compare commits
151
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 | ||
|
|
b99485f462 | ||
|
|
20041d7aa9 | ||
|
|
59248032dc | ||
|
|
ceadc34ea9 | ||
|
|
8ab5631446 | ||
|
|
99b5d2b2da | ||
|
|
021e6f3788 | ||
|
|
4e38183e86 | ||
|
|
4eeb23e2b3 | ||
|
|
ef8783b7e3 | ||
|
|
60d7ee614a | ||
|
|
f7a16efc9d | ||
|
|
a01e8bbe98 | ||
|
|
ccf728a1b7 | ||
|
|
f1b4b05d08 | ||
|
|
0c86c89af4 | ||
|
|
d7ac66fb73 | ||
|
|
a6e920fdb0 | ||
|
|
958df58f9d | ||
|
|
e0f102c4d9 | ||
|
|
5a942527b2 | ||
|
|
37a3036934 | ||
|
|
121a7bf8b4 | ||
|
|
a5678c9185 | ||
|
|
2c50b3cf37 | ||
|
|
eee7f54789 | ||
|
|
06eeeead79 | ||
|
|
e8ff7f5321 | ||
|
|
a6e1f26cd4 | ||
|
|
95c43368ae | ||
|
|
754624acf0 | ||
|
|
0b6a17330f | ||
|
|
74b9308883 | ||
|
|
e5f9b1a3a9 | ||
|
|
31d33ccdf0 | ||
|
|
88ec786e39 | ||
|
|
663ef900fc | ||
|
|
7d478a54db | ||
|
|
f3eaaef842 | ||
|
|
d655b65027 | ||
|
|
31c22dc043 | ||
|
|
17127f8b3c | ||
|
|
d7695b40e3 | ||
|
|
fc62890e70 | ||
|
|
f433672140 | ||
|
|
7e1e5b6e6a | ||
|
|
553a42702d | ||
|
|
b133fc9c07 | ||
|
|
b33250dc28 | ||
|
|
a74e5b91a3 | ||
|
|
28886e4241 | ||
|
|
9d3ccfdffc | ||
|
|
a24a7b4da5 | ||
|
|
f7df02f9a3 | ||
|
|
ee450686f3 | ||
|
|
2565755e45 | ||
|
|
d08a92c7bd | ||
|
|
a1ea26d367 | ||
|
|
c17aa0dc54 | ||
|
|
b12b24eadc | ||
|
|
cd14d53707 | ||
|
|
e220413035 | ||
|
|
84ed2327f5 | ||
|
|
b14f301730 | ||
|
|
0654b4b916 | ||
|
|
1f0be382ad | ||
|
|
bb175fda91 | ||
|
|
13998da15a | ||
|
|
57729fd92d | ||
|
|
2c7a71a9c0 | ||
|
|
3e0007fc91 | ||
|
|
b092316385 | ||
|
|
9bcd696580 | ||
|
|
8f89c82d55 | ||
|
|
21871197d7 | ||
|
|
4c35d36146 | ||
|
|
9aca62c26c | ||
|
|
b5cdea98ad | ||
|
|
69fecaf387 | ||
|
|
fd6d25ad86 | ||
|
|
2c3cef1c87 | ||
|
|
89ece26c25 | ||
|
|
2c0b5d0b5e | ||
|
|
a4ae7d17fb | ||
|
|
8a8550184f | ||
|
|
b8b439b713 | ||
|
|
41cd40363a | ||
|
|
d923ebe38d | ||
|
|
29b0423c4e | ||
|
|
88f8dca2c2 | ||
|
|
9027fdc546 | ||
|
|
cbd140340d | ||
|
|
988e01314d | ||
|
|
7ba43a7c6f | ||
|
|
dea59f7e1d | ||
|
|
85dc771460 |
+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
|
||||||
|
|||||||
@@ -0,0 +1,100 @@
|
|||||||
|
name: Release
|
||||||
|
|
||||||
|
on:
|
||||||
|
push:
|
||||||
|
tags:
|
||||||
|
- "v*"
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
build-pure:
|
||||||
|
name: Build pure-Python wheel
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
steps:
|
||||||
|
- uses: actions/checkout@v4
|
||||||
|
- uses: actions/setup-python@v5
|
||||||
|
with:
|
||||||
|
python-version: "3.12"
|
||||||
|
|
||||||
|
- name: Build wheel (no CUDA)
|
||||||
|
run: |
|
||||||
|
pip wheel . --no-deps -w dist/
|
||||||
|
|
||||||
|
- uses: actions/upload-artifact@v4
|
||||||
|
with:
|
||||||
|
name: pure-wheel
|
||||||
|
path: dist/*.whl
|
||||||
|
if-no-files-found: error
|
||||||
|
|
||||||
|
build-cuda-linux:
|
||||||
|
name: Build CUDA wheel (Linux, ${{ matrix.cuda_tag }})
|
||||||
|
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:
|
||||||
|
- uses: actions/checkout@v4
|
||||||
|
- uses: actions/setup-python@v5
|
||||||
|
with:
|
||||||
|
python-version: "3.12"
|
||||||
|
|
||||||
|
- name: Install torch (${{ matrix.cuda_tag }})
|
||||||
|
run: |
|
||||||
|
pip install torch --index-url https://download.pytorch.org/whl/${{ matrix.cuda_tag }}
|
||||||
|
|
||||||
|
- name: Setup CUDA (${{ matrix.cuda_ver }})
|
||||||
|
uses: Jimver/cuda-toolkit@v0.2.35
|
||||||
|
with:
|
||||||
|
cuda: "${{ matrix.cuda_ver }}"
|
||||||
|
|
||||||
|
- name: Build wheel (with CUDA kernels)
|
||||||
|
run: |
|
||||||
|
CSRC_KERNELS=true pip wheel . --no-deps --no-build-isolation -w dist/
|
||||||
|
|
||||||
|
- uses: actions/upload-artifact@v4
|
||||||
|
with:
|
||||||
|
name: cuda-wheel-linux-${{ matrix.cuda_tag }}
|
||||||
|
path: dist/*.whl
|
||||||
|
if-no-files-found: error
|
||||||
|
|
||||||
|
release:
|
||||||
|
name: Attach wheels to release
|
||||||
|
needs: [build-pure, build-cuda-linux]
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
permissions:
|
||||||
|
contents: write
|
||||||
|
steps:
|
||||||
|
- name: Download pure-Python wheel
|
||||||
|
uses: actions/download-artifact@v4
|
||||||
|
with:
|
||||||
|
name: pure-wheel
|
||||||
|
path: release-assets/pure
|
||||||
|
|
||||||
|
- name: Download CUDA wheels (all variants)
|
||||||
|
uses: actions/download-artifact@v4
|
||||||
|
with:
|
||||||
|
pattern: cuda-wheel-linux-*
|
||||||
|
merge-multiple: true
|
||||||
|
path: release-assets/cuda
|
||||||
|
|
||||||
|
- name: Verify release assets
|
||||||
|
shell: bash
|
||||||
|
run: |
|
||||||
|
set -euo pipefail
|
||||||
|
pure_wheels=(release-assets/pure/*.whl)
|
||||||
|
cuda_wheels=(release-assets/cuda/*.whl)
|
||||||
|
test "${#pure_wheels[@]}" -eq 1
|
||||||
|
test "${#cuda_wheels[@]}" -ge 1
|
||||||
|
|
||||||
|
- name: Create release & upload assets
|
||||||
|
uses: softprops/action-gh-release@v2
|
||||||
|
with:
|
||||||
|
files: |
|
||||||
|
release-assets/pure/*.whl
|
||||||
|
release-assets/cuda/*.whl
|
||||||
|
tag_name: ${{ github.ref_name }}
|
||||||
|
generate_release_notes: true
|
||||||
+2
-1
@@ -11,6 +11,7 @@
|
|||||||
!csrc/**/*.py
|
!csrc/**/*.py
|
||||||
|
|
||||||
!csrc/**/*.cu
|
!csrc/**/*.cu
|
||||||
|
!csrc/**/*.h
|
||||||
!csrc/**/*.cuh
|
!csrc/**/*.cuh
|
||||||
|
|
||||||
!scripts/**/*.sh
|
!scripts/**/*.sh
|
||||||
@@ -23,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,10 +17,10 @@
|
|||||||
|
|
||||||
<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/ViperEk/">HuggingFace</a>
|
<a href="https://huggingface.co/ViperEkura">HuggingFace</a>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<br>
|
<br>
|
||||||
@@ -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
|
||||||
|
|
||||||
@@ -241,7 +246,7 @@ For major changes, please open an issue first to discuss what you would like to
|
|||||||
|
|
||||||
- **GitHub Issues**: [Issue Tracker](https://github.com/ViperEkura/AstrAI/issues)
|
- **GitHub Issues**: [Issue Tracker](https://github.com/ViperEkura/AstrAI/issues)
|
||||||
- **Discussions**: [GitHub Discussions](https://github.com/ViperEkura/AstrAI/discussions)
|
- **Discussions**: [GitHub Discussions](https://github.com/ViperEkura/AstrAI/discussions)
|
||||||
- **HuggingFace**: [Model Hub](https://huggingface.co/ViperEk)
|
- **HuggingFace**: [Model Hub](https://huggingface.co/ViperEkura)
|
||||||
|
|
||||||
### License
|
### License
|
||||||
|
|
||||||
|
|||||||
+31
-3
@@ -1,6 +1,9 @@
|
|||||||
__version__ = "1.3.8"
|
__version__ = "1.3.12"
|
||||||
__author__ = "ViperEkura"
|
__author__ = "ViperEkura"
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
|
||||||
from astrai.config import (
|
from astrai.config import (
|
||||||
AutoRegressiveLMConfig,
|
AutoRegressiveLMConfig,
|
||||||
BaseModelConfig,
|
BaseModelConfig,
|
||||||
@@ -12,7 +15,7 @@ from astrai.config import (
|
|||||||
from astrai.dataset import (
|
from astrai.dataset import (
|
||||||
BaseDataset,
|
BaseDataset,
|
||||||
DatasetFactory,
|
DatasetFactory,
|
||||||
ResumableDistributedSampler,
|
RDSampler,
|
||||||
Store,
|
Store,
|
||||||
StoreFactory,
|
StoreFactory,
|
||||||
)
|
)
|
||||||
@@ -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",
|
||||||
@@ -77,7 +104,7 @@ __all__ = [
|
|||||||
"Pipeline",
|
"Pipeline",
|
||||||
"PipelineConfig",
|
"PipelineConfig",
|
||||||
"ProtocolHandler",
|
"ProtocolHandler",
|
||||||
"ResumableDistributedSampler",
|
"RDSampler",
|
||||||
"SamplingPipeline",
|
"SamplingPipeline",
|
||||||
"SchedulerFactory",
|
"SchedulerFactory",
|
||||||
"Store",
|
"Store",
|
||||||
@@ -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:
|
||||||
|
|||||||
+105
-26
@@ -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,57 +36,126 @@ 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
|
||||||
dim: Optional[int] = None
|
hidden_size: Optional[int] = None
|
||||||
n_layers: Optional[int] = None
|
num_hidden_layers: Optional[int] = None
|
||||||
norm_eps: Optional[float] = None
|
rms_norm_eps: Optional[float] = None
|
||||||
dim_ffn: Optional[int] = None
|
intermediate_size: Optional[int] = None
|
||||||
tie_weight: Optional[bool] = None
|
tie_word_embeddings: Optional[bool] = None
|
||||||
|
max_position_embeddings: Optional[int] = None
|
||||||
max_len: 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"
|
||||||
n_heads: Optional[int] = None
|
num_attention_heads: Optional[int] = None
|
||||||
n_kv_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
|
||||||
dim: Optional[int] = None
|
hidden_size: Optional[int] = None
|
||||||
n_layers: Optional[int] = None
|
num_hidden_layers: Optional[int] = None
|
||||||
norm_eps: Optional[float] = None
|
rms_norm_eps: Optional[float] = None
|
||||||
dim_ffn: Optional[int] = None
|
intermediate_size: Optional[int] = None
|
||||||
|
max_position_embeddings: Optional[int] = None
|
||||||
max_len: 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"
|
||||||
n_heads: Optional[int] = None
|
num_attention_heads: Optional[int] = None
|
||||||
n_kv_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,63 +45,67 @@ 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).
|
|
||||||
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
|
||||||
min_chars: int = 50
|
min_chars: int = 50
|
||||||
max_chars: int = 2_000_000
|
max_chars: int = 2_000_000
|
||||||
max_items: Optional[int] = None
|
max_items: Optional[int] = None
|
||||||
|
batch_size: int = 256
|
||||||
packing_strategy: str = "simple"
|
packing_strategy: str = "simple"
|
||||||
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
|
||||||
@@ -98,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)
|
||||||
|
|||||||
+192
-128
@@ -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,142 +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: float = field(
|
|
||||||
default=1.0, metadata={"help": "Maximum gradient norm."}
|
|
||||||
)
|
|
||||||
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."}
|
|
||||||
)
|
|
||||||
|
|
||||||
# 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)."},
|
|
||||||
)
|
|
||||||
|
|
||||||
executor_kwargs: Dict[str, Any] = field(
|
random_seed: int = 3407
|
||||||
default_factory=dict,
|
num_workers: int = 0
|
||||||
metadata={"help": "Extra kwargs passed to ExecutorFactory.create()."},
|
prefetch_factor: Optional[int] = None
|
||||||
)
|
pin_memory: bool = False
|
||||||
extra_kwargs: Dict[str, Any] = field(
|
collate_fn: Optional[Callable[[List[Any]], Any]] = None
|
||||||
default_factory=dict, metadata={"help": "Other arguments."}
|
|
||||||
)
|
|
||||||
|
|
||||||
def __post_init__(self):
|
nprocs: int = 1
|
||||||
self.validate()
|
backend: str = "nccl"
|
||||||
|
master_addr: str = "localhost"
|
||||||
|
master_port: str = "29500"
|
||||||
|
parallel_mode: str = "none"
|
||||||
|
start_method: str = "spawn"
|
||||||
|
|
||||||
def validate(self):
|
device_type: str = "cuda"
|
||||||
for fld in fields(self):
|
val_dataset: Optional[Dataset] = None
|
||||||
if fld.metadata.get("required") and getattr(self, fld.name) is None:
|
val_split: Optional[float] = None
|
||||||
raise ValueError(f"TrainConfig.{fld.name} is required but got None.")
|
val_step: int = 1000
|
||||||
|
neftune_alpha: float = 0.0
|
||||||
|
|
||||||
|
rollout_interval: int = 512
|
||||||
|
rollout_temperature: float = 0.7
|
||||||
|
rollout_top_k: int = 0
|
||||||
|
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
|
||||||
|
|||||||
@@ -1,35 +1,37 @@
|
|||||||
from astrai.dataset.dataset import (
|
from astrai.dataset.dataset import (
|
||||||
BaseDataset,
|
BaseDataset,
|
||||||
DatasetFactory,
|
DatasetFactory,
|
||||||
|
dpo_collate_fn,
|
||||||
|
grpo_collate_fn,
|
||||||
)
|
)
|
||||||
from astrai.dataset.sampler import ResumableDistributedSampler
|
from astrai.dataset.sampler import RDSampler
|
||||||
from astrai.dataset.storage import (
|
from astrai.dataset.storage import (
|
||||||
H5Store,
|
|
||||||
JsonlStore,
|
JsonlStore,
|
||||||
MmapStore,
|
MmapStore,
|
||||||
|
Recordable,
|
||||||
Store,
|
Store,
|
||||||
StoreFactory,
|
StoreFactory,
|
||||||
|
Streamable,
|
||||||
detect_format,
|
detect_format,
|
||||||
)
|
)
|
||||||
from astrai.serialization import (
|
from astrai.serialization import (
|
||||||
load_bin,
|
load_bin,
|
||||||
load_h5,
|
|
||||||
save_bin,
|
save_bin,
|
||||||
save_h5,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"BaseDataset",
|
"BaseDataset",
|
||||||
"DatasetFactory",
|
"DatasetFactory",
|
||||||
|
"dpo_collate_fn",
|
||||||
|
"grpo_collate_fn",
|
||||||
"Store",
|
"Store",
|
||||||
|
"Streamable",
|
||||||
|
"Recordable",
|
||||||
"StoreFactory",
|
"StoreFactory",
|
||||||
"H5Store",
|
|
||||||
"MmapStore",
|
"MmapStore",
|
||||||
"JsonlStore",
|
"JsonlStore",
|
||||||
"detect_format",
|
"detect_format",
|
||||||
"save_h5",
|
|
||||||
"load_h5",
|
|
||||||
"save_bin",
|
"save_bin",
|
||||||
"load_bin",
|
"load_bin",
|
||||||
"ResumableDistributedSampler",
|
"RDSampler",
|
||||||
]
|
]
|
||||||
|
|||||||
+412
-183
@@ -1,7 +1,31 @@
|
|||||||
"""Dataset implementations with factory pattern for training."""
|
"""Dataset implementations for training.
|
||||||
|
|
||||||
|
Composition over inheritance — every dataset is a thin wrapper that
|
||||||
|
binds a :class:`Store` to a particular train-type's key mapping. All
|
||||||
|
sample-id → token/record indexing lives on the Store; datasets never
|
||||||
|
know about window/stride math or segment layouts.
|
||||||
|
|
||||||
|
Class hierarchy:
|
||||||
|
|
||||||
|
BaseDataset (ABC) — holds a Store, exposes __len__/keys,
|
||||||
|
overrides __getitem__
|
||||||
|
├── SEQDataset — next-token prediction (stream)
|
||||||
|
├── SFTDataset — loss-mask + position_ids (stream)
|
||||||
|
├── DPODataset — chosen/rejected pairs (record)
|
||||||
|
└── GRPODataset — prompt + response group (record)
|
||||||
|
|
||||||
|
``DatasetFactory.load(train_type, load_path, window_size, stride, …)``
|
||||||
|
builds the Store (auto-detecting format) before constructing the
|
||||||
|
matching dataset. Passing ``store=`` skips Store construction.
|
||||||
|
|
||||||
|
When a record dataset (DPO) reads from raw JSONL, a *processor*
|
||||||
|
function (pure ``record -> Dict[str, Tensor]``) is forwarded to
|
||||||
|
:class:`JsonlStore` so tokenisation happens on the fly.
|
||||||
|
"""
|
||||||
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from typing import Dict, List, Optional
|
from functools import partial
|
||||||
|
from typing import Callable, Dict, List, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
@@ -13,202 +37,401 @@ from astrai.dataset.storage import (
|
|||||||
detect_format,
|
detect_format,
|
||||||
)
|
)
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
|
||||||
|
def dpo_tokenize(
|
||||||
|
record: dict,
|
||||||
|
tokenizer,
|
||||||
|
max_len: int = 2048,
|
||||||
|
) -> Optional[dict]:
|
||||||
|
"""Tokenize one DPO record into chosen/rejected + masks.
|
||||||
|
|
||||||
|
Applies the tokenizer's chat template so token sequences match the
|
||||||
|
SFT checkpoint's format. Prompt is rendered with
|
||||||
|
``add_generation_prompt=True``; chosen/rejected are appended as a
|
||||||
|
single assistant turn.
|
||||||
|
|
||||||
|
Accepts:
|
||||||
|
|
||||||
|
- Flat: ``{"prompt": str, "chosen": str, "rejected": str}``
|
||||||
|
- Conv: ``{"prompt": [{role, content}, ...], "chosen": [...], ...}``
|
||||||
|
- Legacy: ``{"input": str, "chosen": str, "rejected": str}``
|
||||||
|
|
||||||
|
No packing, no ``position_ids`` — DPO sequences are independent.
|
||||||
|
"""
|
||||||
|
prompt = record.get("prompt") or record.get("input")
|
||||||
|
chosen = record.get("chosen")
|
||||||
|
rejected = record.get("rejected")
|
||||||
|
if prompt is None or chosen is None or rejected is None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
prompt_messages = _to_messages(prompt)
|
||||||
|
chosen_text = _extract_text(chosen)
|
||||||
|
rejected_text = _extract_text(rejected)
|
||||||
|
if chosen_text is None or rejected_text is None:
|
||||||
|
return None
|
||||||
|
chosen_messages = prompt_messages + [{"role": "assistant", "content": chosen_text}]
|
||||||
|
rejected_messages = prompt_messages + [
|
||||||
|
{"role": "assistant", "content": rejected_text}
|
||||||
|
]
|
||||||
|
|
||||||
|
prompt_ids = tokenizer.apply_chat_template(
|
||||||
|
prompt_messages, tokenize=True, add_generation_prompt=True
|
||||||
|
)
|
||||||
|
ch_ids = tokenizer.apply_chat_template(
|
||||||
|
chosen_messages, tokenize=True, add_generation_prompt=False
|
||||||
|
)
|
||||||
|
re_ids = tokenizer.apply_chat_template(
|
||||||
|
rejected_messages, tokenize=True, add_generation_prompt=False
|
||||||
|
)
|
||||||
|
|
||||||
|
full_ch = ch_ids[:max_len]
|
||||||
|
full_re = re_ids[:max_len]
|
||||||
|
|
||||||
|
prompt_len = min(len(prompt_ids), max_len)
|
||||||
|
ch_mask = [0] * prompt_len + [1] * max(0, len(full_ch) - prompt_len)
|
||||||
|
ch_mask = ch_mask[:max_len]
|
||||||
|
re_mask = [0] * prompt_len + [1] * max(0, len(full_re) - prompt_len)
|
||||||
|
re_mask = re_mask[:max_len]
|
||||||
|
|
||||||
|
return {
|
||||||
|
"chosen": full_ch,
|
||||||
|
"rejected": full_re,
|
||||||
|
"chosen_mask": ch_mask,
|
||||||
|
"rejected_mask": re_mask,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _to_messages(value) -> list:
|
||||||
|
"""Accept str or conversation list; return message list."""
|
||||||
|
if isinstance(value, str):
|
||||||
|
return [{"role": "user", "content": value}]
|
||||||
|
if isinstance(value, list):
|
||||||
|
return value
|
||||||
|
return [{"role": "user", "content": str(value)}]
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_text(value) -> Optional[str]:
|
||||||
|
"""Accept str or conversation list; return plain text."""
|
||||||
|
if value is None:
|
||||||
|
return None
|
||||||
|
if isinstance(value, str):
|
||||||
|
return value
|
||||||
|
if isinstance(value, list):
|
||||||
|
return "".join(m.get("content", "") for m in value if isinstance(m, dict))
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def dpo_processor(
|
||||||
|
record: dict,
|
||||||
|
tokenizer,
|
||||||
|
max_len: int = 2048,
|
||||||
|
) -> Dict[str, Tensor]:
|
||||||
|
"""DPO processor: wraps :func:`dpo_tokenize` and returns tensors."""
|
||||||
|
result = dpo_tokenize(record, tokenizer, max_len=max_len)
|
||||||
|
if result is None:
|
||||||
|
raise ValueError(f"Malformed DPO record: {list(record.keys())}")
|
||||||
|
return {
|
||||||
|
"chosen": torch.tensor(result["chosen"], dtype=torch.int32),
|
||||||
|
"rejected": torch.tensor(result["rejected"], dtype=torch.int32),
|
||||||
|
"chosen_mask": torch.tensor(result["chosen_mask"], dtype=torch.bool),
|
||||||
|
"rejected_mask": torch.tensor(result["rejected_mask"], dtype=torch.bool),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def dpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
|
||||||
|
"""Collate variable-length DPO samples into padded 2-D tensors.
|
||||||
|
|
||||||
|
Input: list of dicts, each with:
|
||||||
|
- chosen: [C_i]
|
||||||
|
- rejected: [R_i]
|
||||||
|
- chosen_mask: [C_i]
|
||||||
|
- rejected_mask: [R_i]
|
||||||
|
|
||||||
|
Output (padded to the max length across chosen/rejected within the batch):
|
||||||
|
- chosen: [B, S_max]
|
||||||
|
- rejected: [B, S_max]
|
||||||
|
- chosen_mask: [B, S_max]
|
||||||
|
- rejected_mask: [B, S_max]
|
||||||
|
"""
|
||||||
|
B = len(batch)
|
||||||
|
S_max = max(b["chosen"].size(0) for b in batch)
|
||||||
|
S_max = max(S_max, max(b["rejected"].size(0) for b in batch))
|
||||||
|
|
||||||
|
chosen = torch.zeros(B, S_max, dtype=torch.long)
|
||||||
|
rejected = torch.zeros(B, S_max, dtype=torch.long)
|
||||||
|
chosen_mask = torch.zeros(B, S_max, dtype=torch.bool)
|
||||||
|
rejected_mask = torch.zeros(B, S_max, dtype=torch.bool)
|
||||||
|
|
||||||
|
for i, b in enumerate(batch):
|
||||||
|
c_len = b["chosen"].size(0)
|
||||||
|
r_len = b["rejected"].size(0)
|
||||||
|
chosen[i, :c_len] = b["chosen"]
|
||||||
|
rejected[i, :r_len] = b["rejected"]
|
||||||
|
chosen_mask[i, :c_len] = b["chosen_mask"]
|
||||||
|
rejected_mask[i, :r_len] = b["rejected_mask"]
|
||||||
|
|
||||||
|
return {
|
||||||
|
"chosen": chosen,
|
||||||
|
"rejected": rejected,
|
||||||
|
"chosen_mask": chosen_mask,
|
||||||
|
"rejected_mask": rejected_mask,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
|
||||||
|
"""Collate variable-length GRPO samples into padded 3-D tensors.
|
||||||
|
|
||||||
|
Input: list of dicts, each with:
|
||||||
|
- prompts: [P_i]
|
||||||
|
- responses: list of G tensors, each [R_ij]
|
||||||
|
- masks: list of G tensors, each [R_ij]
|
||||||
|
- rewards: [G]
|
||||||
|
|
||||||
|
Output:
|
||||||
|
- prompts: [B, P_max], left-padded
|
||||||
|
- prompt_mask: [B, P_max]
|
||||||
|
- responses: [B, G, R_max]
|
||||||
|
- masks: [B, G, R_max]
|
||||||
|
- rewards: [B, G]
|
||||||
|
"""
|
||||||
|
B = len(batch)
|
||||||
|
G = len(batch[0]["responses"])
|
||||||
|
P_max = max(b["prompts"].size(0) for b in batch)
|
||||||
|
R_max = max(r.size(0) for b in batch for r in b["responses"])
|
||||||
|
|
||||||
|
prompts = torch.zeros(B, P_max, dtype=torch.long)
|
||||||
|
prompt_mask = torch.zeros(B, P_max, dtype=torch.bool)
|
||||||
|
responses = torch.zeros(B, G, R_max, dtype=torch.long)
|
||||||
|
masks = torch.zeros(B, G, R_max, dtype=torch.bool)
|
||||||
|
rewards = torch.zeros(B, G, dtype=torch.float32)
|
||||||
|
|
||||||
|
for i, b in enumerate(batch):
|
||||||
|
p_len = b["prompts"].size(0)
|
||||||
|
prompts[i, -p_len:] = b["prompts"]
|
||||||
|
prompt_mask[i, -p_len:] = True
|
||||||
|
rewards[i, : b["rewards"].size(0)] = b["rewards"]
|
||||||
|
for g in range(min(G, len(b["responses"]))):
|
||||||
|
r_len = b["responses"][g].size(0)
|
||||||
|
responses[i, g, :r_len] = b["responses"][g]
|
||||||
|
if g < len(b["masks"]):
|
||||||
|
masks[i, g, :r_len] = b["masks"][g]
|
||||||
|
|
||||||
|
return {
|
||||||
|
"prompts": prompts,
|
||||||
|
"prompt_mask": prompt_mask,
|
||||||
|
"responses": responses,
|
||||||
|
"masks": masks,
|
||||||
|
"rewards": rewards,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def validate_keys(store: Store, required: List[str]) -> None:
|
||||||
|
"""Raise ``KeyError`` if *store* is missing any *required* key."""
|
||||||
|
if not required:
|
||||||
|
return
|
||||||
|
actual = set(store.keys)
|
||||||
|
missing = [k for k in required if k not in actual]
|
||||||
|
if missing:
|
||||||
|
raise KeyError(
|
||||||
|
f"Store at {getattr(store, '_load_path', '?')} is missing required "
|
||||||
|
f"keys {missing}; available keys are {sorted(actual)}."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class BaseDataset(Dataset, ABC):
|
class BaseDataset(Dataset, ABC):
|
||||||
"""Abstract base class for all dataset types.
|
"""Abstract base class for dataset types.
|
||||||
|
|
||||||
Implements common functionality for window-based data fetching.
|
Holds a :class:`Store`. All sample-id indexing is delegated to the
|
||||||
Uses a storage abstraction for format-agnostic data loading.
|
store — this class exposes ``__len__`` as ``len(store)`` and the
|
||||||
|
``keys`` property as ``store.keys``. Subclasses implement
|
||||||
|
``__getitem__`` with the train-type-specific key mapping and any
|
||||||
|
training-only index arithmetic (e.g. the next-token ``+1`` shift).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, window_size: int, stride: int):
|
required_keys: List[str] = []
|
||||||
|
|
||||||
|
def __init__(self, store: Store):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.window_size = window_size
|
self.store: Store = store
|
||||||
self.stride = stride
|
validate_keys(store, self.required_keys)
|
||||||
self.storage: Optional[Store] = None
|
|
||||||
|
|
||||||
@property
|
def __len__(self) -> int:
|
||||||
def required_keys(self) -> List[str]:
|
return len(self.store)
|
||||||
"""Return required storage keys for this dataset type.
|
|
||||||
|
|
||||||
Subclasses should override to specify expected keys.
|
|
||||||
"""
|
|
||||||
return []
|
|
||||||
|
|
||||||
def _validate_keys(self):
|
|
||||||
if not self.required_keys:
|
|
||||||
return
|
|
||||||
actual_keys = set(self.storage.keys)
|
|
||||||
missing = [k for k in self.required_keys if k not in actual_keys]
|
|
||||||
if missing:
|
|
||||||
raise KeyError(
|
|
||||||
f"Dataset {type(self).__name__} requires keys {self.required_keys}, "
|
|
||||||
f"but storage at {self._load_path} only has {sorted(actual_keys)}. "
|
|
||||||
f"Missing: {missing}"
|
|
||||||
)
|
|
||||||
|
|
||||||
def load(self, load_path: str, storage_type: Optional[str] = None, **kwargs):
|
|
||||||
"""Load dataset from the given path.
|
|
||||||
|
|
||||||
Auto-detects the storage format if not specified.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
load_path: Path to the data directory or file
|
|
||||||
storage_type: Force a specific storage type ("h5", "bin", "jsonl"),
|
|
||||||
or None for auto-detection
|
|
||||||
**kwargs: Extra arguments forwarded to the store constructor and
|
|
||||||
to ``store.load()``.
|
|
||||||
|
|
||||||
Raises:
|
|
||||||
KeyError: If the loaded storage is missing required keys.
|
|
||||||
"""
|
|
||||||
if storage_type is None:
|
|
||||||
storage_type = detect_format(load_path)
|
|
||||||
self.storage = StoreFactory.create(storage_type, **kwargs)
|
|
||||||
self._load_path = load_path
|
|
||||||
self.storage.load(load_path, **kwargs)
|
|
||||||
self._validate_keys()
|
|
||||||
|
|
||||||
@property
|
|
||||||
def count(self) -> int:
|
|
||||||
"""Return the total number of raw elements (tokens) in the dataset."""
|
|
||||||
if self.storage is None:
|
|
||||||
return 0
|
|
||||||
return len(self.storage)
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def keys(self) -> List[str]:
|
def keys(self) -> List[str]:
|
||||||
"""Return the available data keys."""
|
return self.store.keys
|
||||||
if self.storage is None:
|
|
||||||
return []
|
|
||||||
return self.storage.keys
|
|
||||||
|
|
||||||
def get_index(self, index: int) -> tuple:
|
@property
|
||||||
"""Calculate begin and end indices for a sample.
|
def token_count(self) -> int:
|
||||||
|
return self.store.token_count
|
||||||
Args:
|
|
||||||
index: Sample index
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Tuple of (begin_idx, end_idx)
|
|
||||||
"""
|
|
||||||
if self.storage is None:
|
|
||||||
raise RuntimeError("Dataset not loaded, call load() first")
|
|
||||||
total = len(self.storage)
|
|
||||||
if total <= self.window_size:
|
|
||||||
raise ValueError(
|
|
||||||
f"Data too short: {total} tokens <= window_size {self.window_size}"
|
|
||||||
)
|
|
||||||
|
|
||||||
begin_idx = min(index * self.stride, total - 1 - self.window_size)
|
|
||||||
end_idx = min(begin_idx + self.window_size, total - 1)
|
|
||||||
|
|
||||||
return begin_idx, end_idx
|
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||||
"""Get a single sample by index.
|
|
||||||
|
|
||||||
Must be implemented by subclasses.
|
|
||||||
"""
|
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
def __len__(self) -> int:
|
|
||||||
if self.storage is None:
|
|
||||||
return 0
|
|
||||||
total = len(self.storage)
|
|
||||||
if total <= self.window_size:
|
|
||||||
return 0
|
|
||||||
return (total - 1 - self.window_size) // self.stride + 1
|
|
||||||
|
|
||||||
|
|
||||||
class DatasetFactory(BaseFactory["BaseDataset"]):
|
class DatasetFactory(BaseFactory["BaseDataset"]):
|
||||||
"""Factory class for creating dataset instances.
|
"""Factory for creating dataset instances by train-type.
|
||||||
|
|
||||||
Supports decorator-based registration for extensible dataset types.
|
Use :meth:`DatasetFactory.register("custom")` to register new
|
||||||
All default dataset types (seq, sft, dpo, grpo) are registered automatically
|
dataset classes; they must inherit from :class:`BaseDataset`.
|
||||||
when their classes are defined with the decorator.
|
|
||||||
|
|
||||||
Example usage:
|
|
||||||
@DatasetFactory.register("custom")
|
|
||||||
class CustomDataset(BaseDataset):
|
|
||||||
...
|
|
||||||
|
|
||||||
dataset = DatasetFactory.create("custom", window_size, stride)
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def load(
|
def load(
|
||||||
cls,
|
cls,
|
||||||
train_type: str,
|
train_type: str,
|
||||||
load_path: str,
|
load_path: Optional[str] = None,
|
||||||
window_size: int,
|
window_size: int = 0,
|
||||||
stride: Optional[int] = None,
|
stride: Optional[int] = None,
|
||||||
storage_type: Optional[str] = None,
|
storage_type: Optional[str] = None,
|
||||||
|
tokenizer_path: Optional[str] = None,
|
||||||
|
max_len: int = 2048,
|
||||||
|
store: Optional[Store] = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> "BaseDataset":
|
) -> "BaseDataset":
|
||||||
"""Create and load a dataset in one step.
|
"""Create and load a dataset in one step.
|
||||||
|
|
||||||
|
Two entry points:
|
||||||
|
|
||||||
|
- **store given**: bind it directly — the caller fully controls
|
||||||
|
Store construction and processor setup. *load_path*,
|
||||||
|
*storage_type*, *tokenizer_path*, *window_size*, *stride* are
|
||||||
|
ignored.
|
||||||
|
- **store is None**: build a Store from *load_path*, auto-detecting
|
||||||
|
format and constructing a processor when *tokenizer_path* is
|
||||||
|
given for a record dataset on JSONL.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
train_type: Type of training dataset
|
train_type: Registered dataset name ("seq", "sft", "dpo",
|
||||||
load_path: Path to the data file
|
"grpo", …).
|
||||||
window_size: Window size for data sampling
|
load_path: Path to the data file or directory (ignored if
|
||||||
stride: Stride between consecutive samples (default: same as window_size)
|
*store* is given).
|
||||||
storage_type: Storage type ("h5", "bin", "jsonl") or None for auto-detection
|
window_size: Stream window length — only meaningful for
|
||||||
**kwargs: Extra arguments forwarded to ``dataset.load()``.
|
stream datasets (SEQ/SFT). Record datasets ignore it.
|
||||||
|
stride: Stride between consecutive stream samples
|
||||||
|
(default: same as *window_size*).
|
||||||
|
storage_type: Storage backend ("bin", "jsonl") or
|
||||||
|
None for auto-detection.
|
||||||
|
tokenizer_path: Path to tokenizer for lazy JSONL
|
||||||
|
tokenisation (record datasets only).
|
||||||
|
max_len: Max sequence length forwarded to processors.
|
||||||
|
store: Pre-built, already-loaded Store instance.
|
||||||
|
**kwargs: Extra arguments forwarded to ``store.load()``.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Loaded dataset instance
|
Loaded dataset instance.
|
||||||
"""
|
"""
|
||||||
|
if store is not None:
|
||||||
|
return cls.create(train_type, store=store)
|
||||||
|
|
||||||
|
if load_path is None:
|
||||||
|
raise ValueError("Either load_path or store must be provided")
|
||||||
|
|
||||||
|
if storage_type is None:
|
||||||
|
storage_type = detect_format(load_path)
|
||||||
|
|
||||||
if stride is None:
|
if stride is None:
|
||||||
stride = window_size
|
stride = window_size
|
||||||
|
|
||||||
dataset = cls.create(train_type, window_size, stride)
|
processor = cls._maybe_build_processor(
|
||||||
dataset.load(load_path, storage_type=storage_type, **kwargs)
|
train_type, storage_type, tokenizer_path, max_len
|
||||||
|
)
|
||||||
|
|
||||||
return dataset
|
store_window = cls._store_window_for(train_type, window_size)
|
||||||
|
store = StoreFactory.create(
|
||||||
|
storage_type,
|
||||||
|
window_size=store_window,
|
||||||
|
stride=stride if stride else store_window,
|
||||||
|
)
|
||||||
|
if processor is not None:
|
||||||
|
store.load(load_path, processor=processor, **kwargs)
|
||||||
|
else:
|
||||||
|
load_kwargs = dict(kwargs)
|
||||||
|
if (
|
||||||
|
tokenizer_path is not None
|
||||||
|
and storage_type == "jsonl"
|
||||||
|
and train_type in ("seq", "sft")
|
||||||
|
and "tokenizer_path" not in load_kwargs
|
||||||
|
):
|
||||||
|
load_kwargs["tokenizer_path"] = tokenizer_path
|
||||||
|
store.load(load_path, **load_kwargs)
|
||||||
|
|
||||||
|
return cls.create(train_type, store=store)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _store_window_for(train_type: str, window_size: int) -> int:
|
||||||
|
"""Stream datasets consume ``window_size``; record datasets ignore it.
|
||||||
|
|
||||||
|
Record datasets (dpo/grpo) treat each record as an independent
|
||||||
|
training unit and never window, so the store is built with
|
||||||
|
``window_size=0`` and ``len(store)`` returns the record count.
|
||||||
|
"""
|
||||||
|
if train_type in ("seq", "sft"):
|
||||||
|
return window_size
|
||||||
|
return 0
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _maybe_build_processor(
|
||||||
|
train_type: str,
|
||||||
|
storage_type: str,
|
||||||
|
tokenizer_path: Optional[str],
|
||||||
|
max_len: int,
|
||||||
|
) -> Optional[Callable[[dict], Dict[str, Tensor]]]:
|
||||||
|
"""Build an on-the-fly tokenisation processor if applicable.
|
||||||
|
|
||||||
|
Only raw JSONL + record datasets (DPO/GRPO) need a processor;
|
||||||
|
pre-tokenised backends (bin) and stream datasets (SEQ/SFT)
|
||||||
|
return ``None`` so no tokenizer is loaded.
|
||||||
|
"""
|
||||||
|
if tokenizer_path is None or storage_type != "jsonl":
|
||||||
|
return None
|
||||||
|
if train_type == "dpo":
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
|
||||||
|
return partial(dpo_processor, tokenizer=tokenizer, max_len=max_len)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
@DatasetFactory.register("seq")
|
@DatasetFactory.register("seq")
|
||||||
class SEQDataset(BaseDataset):
|
class SEQDataset(BaseDataset):
|
||||||
"""Dataset for sequential next-token prediction training."""
|
"""Dataset for sequential next-token prediction training.
|
||||||
|
|
||||||
@property
|
Stream mode: ``store.fetch(begin, end, "sequence")`` returns the
|
||||||
def required_keys(self) -> List[str]:
|
input window; the +1 shifted call returns the next-token target.
|
||||||
return ["sequence"]
|
"""
|
||||||
|
|
||||||
def _fetch_data(self, begin_idx: int, end_idx: int) -> Tensor:
|
required_keys = ["sequence"]
|
||||||
return self.storage.fetch(begin_idx, end_idx, "sequence")
|
|
||||||
|
|
||||||
def __getitem__(self, index):
|
def __getitem__(self, index: int):
|
||||||
begin_idx, end_idx = self.get_index(index)
|
begin, end = self.store.sample_window(index)
|
||||||
|
x = self.store.fetch(begin, end, "sequence")
|
||||||
x = self._fetch_data(begin_idx, end_idx).to(dtype=torch.long)
|
y = self.store.fetch(begin + 1, end + 1, "sequence")
|
||||||
y = self._fetch_data(begin_idx + 1, end_idx + 1).to(dtype=torch.long)
|
return {
|
||||||
|
"input_ids": x.to(dtype=torch.long),
|
||||||
return {"input_ids": x, "target_ids": y}
|
"target_ids": y.to(dtype=torch.long),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
@DatasetFactory.register("sft")
|
@DatasetFactory.register("sft")
|
||||||
class SFTDataset(BaseDataset):
|
class SFTDataset(BaseDataset):
|
||||||
"""Dataset for supervised fine-tuning with loss masking."""
|
"""Dataset for supervised fine-tuning with loss masking.
|
||||||
|
|
||||||
@property
|
Stream mode: ``sequence``/``loss_mask``/``position_ids`` are sliced
|
||||||
def required_keys(self) -> List[str]:
|
to the window. ``loss_mask`` and ``target_ids`` use the +1 shifted
|
||||||
return ["sequence", "loss_mask", "position_ids"]
|
slice so they align with the predicted positions.
|
||||||
|
"""
|
||||||
|
|
||||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
required_keys = ["sequence", "loss_mask", "position_ids"]
|
||||||
return self.storage.fetch(begin_idx, end_idx, key)
|
|
||||||
|
|
||||||
def __getitem__(self, index):
|
|
||||||
begin_idx, end_idx = self.get_index(index)
|
|
||||||
|
|
||||||
x = self._fetch_data(begin_idx, end_idx, "sequence")
|
|
||||||
y = self._fetch_data(begin_idx + 1, end_idx + 1, "sequence")
|
|
||||||
position_ids = self._fetch_data(begin_idx, end_idx, "position_ids")
|
|
||||||
loss_mask = self._fetch_data(begin_idx + 1, end_idx + 1, "loss_mask")
|
|
||||||
|
|
||||||
|
def __getitem__(self, index: int):
|
||||||
|
begin, end = self.store.sample_window(index)
|
||||||
|
x = self.store.fetch(begin, end, "sequence")
|
||||||
|
y = self.store.fetch(begin + 1, end + 1, "sequence")
|
||||||
|
position_ids = self.store.fetch(begin, end, "position_ids")
|
||||||
|
loss_mask = self.store.fetch(begin + 1, end + 1, "loss_mask")
|
||||||
return {
|
return {
|
||||||
"input_ids": x.to(dtype=torch.long),
|
"input_ids": x.to(dtype=torch.long),
|
||||||
"target_ids": y.to(dtype=torch.long),
|
"target_ids": y.to(dtype=torch.long),
|
||||||
@@ -219,59 +442,65 @@ class SFTDataset(BaseDataset):
|
|||||||
|
|
||||||
@DatasetFactory.register("dpo")
|
@DatasetFactory.register("dpo")
|
||||||
class DPODataset(BaseDataset):
|
class DPODataset(BaseDataset):
|
||||||
"""Dataset for Direct Preference Optimization training."""
|
"""Record-structured dataset for Direct Preference Optimization.
|
||||||
|
|
||||||
@property
|
Each sample is one preference pair (chosen + rejected) and is an
|
||||||
def required_keys(self) -> List[str]:
|
independent training unit — no windowing, stride, or cross-record
|
||||||
return ["chosen", "rejected", "chosen_mask", "rejected_mask"]
|
concatenation. This keeps each sequence self-contained so attention
|
||||||
|
never leaks across preference pairs.
|
||||||
|
|
||||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
Two loading paths (handled by :class:`DatasetFactory`):
|
||||||
return self.storage.fetch(begin_idx, end_idx, key)
|
|
||||||
|
|
||||||
def __getitem__(self, index: int):
|
- **Pre-tokenized** (bin): ``store.load(path)`` reads per-record
|
||||||
begin_idx, end_idx = self.get_index(index)
|
tensors; ``__getitem__`` returns them directly.
|
||||||
|
- **Raw JSONL** (``tokenizer_path=...``): builds a lazy processor
|
||||||
|
via :func:`dpo_processor` that tokenises on the fly — no packing,
|
||||||
|
no ``position_ids``.
|
||||||
|
"""
|
||||||
|
|
||||||
chosen = self._fetch_data(begin_idx, end_idx, "chosen").to(dtype=torch.long)
|
required_keys = ["chosen", "rejected", "chosen_mask", "rejected_mask"]
|
||||||
rejected = self._fetch_data(begin_idx, end_idx, "rejected").to(dtype=torch.long)
|
|
||||||
chosen_mask = self._fetch_data(begin_idx, end_idx, "chosen_mask").to(
|
|
||||||
dtype=torch.bool
|
|
||||||
)
|
|
||||||
rejected_mask = self._fetch_data(begin_idx, end_idx, "rejected_mask").to(
|
|
||||||
dtype=torch.bool
|
|
||||||
)
|
|
||||||
|
|
||||||
|
def make_processor(self, tokenizer, max_len: int):
|
||||||
|
return partial(dpo_processor, tokenizer=tokenizer, max_len=max_len)
|
||||||
|
|
||||||
|
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||||
return {
|
return {
|
||||||
"chosen": chosen,
|
"chosen": self.store.fetch_record(index, "chosen").to(dtype=torch.long),
|
||||||
"rejected": rejected,
|
"rejected": self.store.fetch_record(index, "rejected").to(dtype=torch.long),
|
||||||
"chosen_mask": chosen_mask,
|
"chosen_mask": self.store.fetch_record(index, "chosen_mask").to(
|
||||||
"rejected_mask": rejected_mask,
|
dtype=torch.bool
|
||||||
|
),
|
||||||
|
"rejected_mask": self.store.fetch_record(index, "rejected_mask").to(
|
||||||
|
dtype=torch.bool
|
||||||
|
),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@DatasetFactory.register("grpo")
|
@DatasetFactory.register("grpo")
|
||||||
class GRPODataset(BaseDataset):
|
class GRPODataset(BaseDataset):
|
||||||
"""Dataset for Group Relative Policy Optimization training."""
|
"""Dataset for offline Group Relative Policy Optimization.
|
||||||
|
|
||||||
@property
|
Each sample is one prompt with its group of responses and scalar
|
||||||
def required_keys(self) -> List[str]:
|
rewards — an independent training unit with no windowing or stride.
|
||||||
return ["prompts", "responses", "masks", "rewards"]
|
|
||||||
|
|
||||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
Expected storage layout (produced by JsonlStore or pre-tokenized):
|
||||||
return self.storage.fetch(begin_idx, end_idx, key)
|
|
||||||
|
- ``prompts``: List[Tensor] — one 1-D token tensor per record
|
||||||
|
- ``responses``: List[List[Tensor]] — G response tensors per record
|
||||||
|
- ``masks``: List[List[Tensor]] — G mask tensors per record
|
||||||
|
- ``rewards``: List[Tensor] — one 1-D float tensor (len G) per record
|
||||||
|
"""
|
||||||
|
|
||||||
|
required_keys = ["prompts", "responses", "masks", "rewards"]
|
||||||
|
|
||||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||||
begin_idx, end_idx = self.get_index(index)
|
prompts = self.store.fetch_record(index, "prompts")
|
||||||
|
responses = self.store.fetch_record(index, "responses")
|
||||||
prompts = self._fetch_data(begin_idx, end_idx, "prompts").to(dtype=torch.long)
|
masks = self.store.fetch_record(index, "masks")
|
||||||
responses = self._fetch_data(begin_idx, end_idx, "responses").to(
|
rewards = self.store.fetch_record(index, "rewards")
|
||||||
dtype=torch.long
|
|
||||||
)
|
|
||||||
masks = self._fetch_data(begin_idx, end_idx, "masks").to(dtype=torch.bool)
|
|
||||||
rewards = self._fetch_data(begin_idx, end_idx, "rewards")
|
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"prompts": prompts,
|
"prompts": prompts.to(dtype=torch.long),
|
||||||
"responses": responses,
|
"responses": [r.to(dtype=torch.long) for r in responses],
|
||||||
"masks": masks,
|
"masks": [m.to(dtype=torch.bool) for m in masks],
|
||||||
"rewards": rewards,
|
"rewards": rewards.to(dtype=torch.float32),
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -5,7 +5,15 @@ import torch.distributed as dist
|
|||||||
from torch.utils.data import Dataset, Sampler
|
from torch.utils.data import Dataset, Sampler
|
||||||
|
|
||||||
|
|
||||||
class ResumableDistributedSampler(Sampler[int]):
|
class RDSampler(Sampler[int]):
|
||||||
|
"""Resumable Distributed Sampler.
|
||||||
|
|
||||||
|
A distributed sampler that supports checkpoint-based resume: iteration
|
||||||
|
state (epoch, position) is tracked so training can continue from the
|
||||||
|
exact sample after a restart. Shards the dataset across
|
||||||
|
``dist.world_size`` replicas with optional shuffling.
|
||||||
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
data_source: Dataset,
|
data_source: Dataset,
|
||||||
|
|||||||
+502
-188
@@ -1,20 +1,47 @@
|
|||||||
"""Storage backends for different data formats.
|
"""Storage backends for different data formats.
|
||||||
|
|
||||||
Layers:
|
Architecture (composition over inheritance):
|
||||||
- I/O layer: save_* / load_* functions, read/write raw files (HDF5/bin)
|
|
||||||
return Dict[str, List[Tensor]] — format-specific, no state
|
|
||||||
- Store (ABC): central abstraction, normalizes multi-segment into
|
|
||||||
Dict[str, List[Tensor]] per key via _normalize(),
|
|
||||||
fetch() uses bisect across segments — no forced concat
|
|
||||||
- Dataset layer: BaseDataset owns a Store, only calls store.fetch(begin, end, key)
|
|
||||||
|
|
||||||
Key properties:
|
Store (ABC) — owns _data/_cum/_offsets bookkeeping
|
||||||
- Multi-segment: segments kept as-is, no forced concatenation — safe for
|
+ window_size/stride for sample-id
|
||||||
datasets larger than RAM
|
indexing. __getitem__/__len__ produce
|
||||||
- Explicit length: _length = min(total elements across keys), set at load,
|
the smallest iterable unit so Dataset
|
||||||
__len__ returns O(1)
|
classes are pure delegators.
|
||||||
- Zero-copy mmap: MmapStore wraps np.memmap(mode="r"), all DataLoader
|
Streamable (mixin) — raw token slice fetch(begin, end, keys)
|
||||||
workers share OS page-cache pages
|
Recordable (mixin) — raw record slice fetch_record(idx, keys)
|
||||||
|
|
||||||
|
MmapStore(Store, Streamable, Recordable)
|
||||||
|
JsonlStore(Store, Streamable, Recordable)
|
||||||
|
|
||||||
|
Each mixin is a stateless trait that relies on ``self._data`` etc.
|
||||||
|
provided by :class:`Store`. Concrete stores mix in whichever access
|
||||||
|
primitives they support — ``Store`` is the sole base class, so there is
|
||||||
|
no diamond inheritance or MRO ambiguity.
|
||||||
|
|
||||||
|
Sample-id indexing lives on :class:`Store`, not on the dataset:
|
||||||
|
|
||||||
|
- **Stream mode** (``window_size > 0``): ``len(store)`` returns the number
|
||||||
|
of ``(window_size, stride)`` windows that fit in the token river;
|
||||||
|
``store[i]`` returns the *i*-th window as a dict of per-key tensors;
|
||||||
|
``store.sample_window(i)`` exposes the underlying ``(begin, end)``
|
||||||
|
token slice for callers (e.g. next-token trainers) that need a +1
|
||||||
|
shifted companion window.
|
||||||
|
- **Record mode** (``num_records > 0``): ``len(store)`` returns the
|
||||||
|
record count; ``store[i]`` returns the *i*-th record dict.
|
||||||
|
|
||||||
|
Raw token/record access via :meth:`fetch` / :meth:`fetch_record`
|
||||||
|
remains available for low-level callers that want explicit index
|
||||||
|
control. ``store.token_count`` is the total stream token count (what
|
||||||
|
``len(store)`` used to mean in the legacy stream-only API).
|
||||||
|
|
||||||
|
``segments_are_records`` (class attribute on each Store subclass)
|
||||||
|
tells ``_normalize`` whether segments are inherently per-record (JSONL)
|
||||||
|
or opaque shards (bin). Record access for bin relies on ``_offsets``
|
||||||
|
instead.
|
||||||
|
|
||||||
|
:class:`JsonlStore` supports a lazy mode (``processor=fn``) that keeps
|
||||||
|
raw records and defers tokenisation to ``fetch_record`` — used by DPO
|
||||||
|
to train directly from a ``.jsonl`` file without a pre-tokenised copy.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import bisect
|
import bisect
|
||||||
@@ -23,20 +50,18 @@ import json
|
|||||||
import logging
|
import logging
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Dict, List, Union
|
from typing import Callable, Dict, List, Optional, Tuple, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from astrai.config.preprocess_config import PipelineConfig
|
from astrai.config.preprocess_config import PipelineConfig
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
from astrai.preprocessing.builder import MaskBuilderFactory
|
from astrai.preprocessing.transform import TokenizeTransform
|
||||||
from astrai.preprocessing.position_id import PositionIdStrategyFactory
|
|
||||||
from astrai.serialization import (
|
from astrai.serialization import (
|
||||||
load_bin,
|
load_bin,
|
||||||
load_h5,
|
load_bin_offsets,
|
||||||
)
|
)
|
||||||
from astrai.tokenize import AutoTokenizer
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -48,7 +73,7 @@ def detect_format(load_path: str) -> str:
|
|||||||
load_path: Directory or file path
|
load_path: Directory or file path
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Format string ("h5", "bin", or "jsonl")
|
Format string ("h5", "bin", "jsonl", or "processed")
|
||||||
|
|
||||||
Raises:
|
Raises:
|
||||||
FileNotFoundError: If no supported data files are found
|
FileNotFoundError: If no supported data files are found
|
||||||
@@ -56,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(
|
||||||
@@ -85,228 +101,526 @@ def detect_format(load_path: str) -> str:
|
|||||||
|
|
||||||
|
|
||||||
class Store(ABC):
|
class Store(ABC):
|
||||||
"""String keys -> segmented tensors with ``fetch(begin, end, keys)``.
|
"""Common base for all storage backends.
|
||||||
|
|
||||||
Each key maps to one or more tensor segments (no forced concatenation).
|
A Store owns both its data layout AND its sample-id → token/record
|
||||||
``len(store)`` returns ``self._length`` (explicit, O(1)), the minimum
|
index translation. Datasets are thin wrappers that bind a Store
|
||||||
total element count across all keys.
|
to a particular train-type's key mapping; they never know about
|
||||||
|
window/stride math.
|
||||||
|
|
||||||
Subclasses fill ``self._data`` and ``self._cum`` during ``load()``
|
Two iteration modes:
|
||||||
via ``_normalize()``.
|
|
||||||
|
- **Stream** (``window_size > 0``): data is treated as one long
|
||||||
|
token river. ``len(store)`` returns the number of windows;
|
||||||
|
``store[i]`` slices every stream-compatible key to window ``i``;
|
||||||
|
``store.sample_window(i)`` returns the ``(begin, end)`` token
|
||||||
|
slice for callers needing a +1 shifted companion window.
|
||||||
|
- **Record** (``num_records > 0``): data is per-record.
|
||||||
|
``len(store)`` returns ``num_records``; ``store[i]`` returns
|
||||||
|
the *i*-th record as a dict.
|
||||||
|
|
||||||
|
Raw token slicing is still available via :meth:`fetch` (mixed in
|
||||||
|
by :class:`Streamable`) when a store has stream support configured.
|
||||||
|
Raw record slicing via :meth:`fetch_record` (mixed in by
|
||||||
|
:class:`Recordable`) when a store has record support.
|
||||||
|
|
||||||
|
``token_count`` exposes the raw total stream length — this is what
|
||||||
|
``len(store)`` returned in the legacy stream-only API and what
|
||||||
|
stream-bound ``fetch`` uses for its bounds check.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self):
|
segments_are_records: bool = False
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
window_size: int = 0,
|
||||||
|
stride: Optional[int] = None,
|
||||||
|
):
|
||||||
self._data: Dict[str, List[Tensor]] = {}
|
self._data: Dict[str, List[Tensor]] = {}
|
||||||
self._cum: Dict[str, List[int]] = {}
|
self._cum: Dict[str, List[int]] = {}
|
||||||
|
self._offsets: Dict[str, List[int]] = {}
|
||||||
self._length: int = 0
|
self._length: int = 0
|
||||||
|
self._num_records: int = 0
|
||||||
|
self._window_size: int = int(window_size)
|
||||||
|
self._stride: int = int(stride) if stride is not None else int(window_size)
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def load(self, path: str) -> None:
|
def load(self, path: str, **kwargs) -> None:
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def keys(self) -> List[str]:
|
def keys(self) -> List[str]:
|
||||||
return list(self._data.keys())
|
return list(self._data.keys())
|
||||||
|
|
||||||
def __len__(self) -> int:
|
@property
|
||||||
|
def window_size(self) -> int:
|
||||||
|
return self._window_size
|
||||||
|
|
||||||
|
@property
|
||||||
|
def stride(self) -> int:
|
||||||
|
return self._stride
|
||||||
|
|
||||||
|
@property
|
||||||
|
def token_count(self) -> int:
|
||||||
|
"""Total tokens across all stream segments.
|
||||||
|
|
||||||
|
Useful for the bounds-checked raw :meth:`fetch` and as the
|
||||||
|
legacy ``len(store)`` value.
|
||||||
|
"""
|
||||||
return self._length
|
return self._length
|
||||||
|
|
||||||
|
@property
|
||||||
|
def num_records(self) -> int:
|
||||||
|
"""Number of records available via :meth:`fetch_record`.
|
||||||
|
|
||||||
|
Non-zero only when the backing layout provides per-record
|
||||||
|
indexing (JSONL segments or bin ``_offsets``).
|
||||||
|
"""
|
||||||
|
return self._num_records
|
||||||
|
|
||||||
|
@property
|
||||||
|
def num_samples(self) -> int:
|
||||||
|
"""Number of items produced by ``__getitem__``.
|
||||||
|
|
||||||
|
Stream-mode wins when ``window_size > 0`` and there are tokens
|
||||||
|
to slice; otherwise falls back to ``num_records``.
|
||||||
|
"""
|
||||||
|
if self._window_size > 0 and self._length > 0:
|
||||||
|
total = self._length
|
||||||
|
w = self._window_size
|
||||||
|
if total <= w:
|
||||||
|
return 0
|
||||||
|
return (total - 1 - w) // self._stride + 1
|
||||||
|
return self._num_records
|
||||||
|
|
||||||
|
def __len__(self) -> int:
|
||||||
|
return self.num_samples
|
||||||
|
|
||||||
|
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||||
|
if index < 0:
|
||||||
|
index += self.num_samples
|
||||||
|
if not 0 <= index < self.num_samples:
|
||||||
|
raise IndexError(
|
||||||
|
f"Store index out of range: {index}, num_samples={self.num_samples}"
|
||||||
|
)
|
||||||
|
if self._window_size > 0 and self._length > 0:
|
||||||
|
begin, end = self.sample_window(index)
|
||||||
|
keys = self._stream_keys()
|
||||||
|
return {k: self.fetch(begin, end, k) for k in keys}
|
||||||
|
return self.fetch_record(index, self._record_keys())
|
||||||
|
|
||||||
|
def sample_window(self, index: int) -> Tuple[int, int]:
|
||||||
|
"""Return ``(begin, end)`` token positions for stream sample *index*.
|
||||||
|
|
||||||
|
The clipped tail keeps the last reachable window inside the
|
||||||
|
token river instead of overshooting. Caller is responsible
|
||||||
|
for staying within :attr:`num_samples`: an out-of-range index
|
||||||
|
raises ``IndexError``.
|
||||||
|
"""
|
||||||
|
if self._window_size <= 0:
|
||||||
|
raise RuntimeError("sample_window() requires window_size > 0 (stream mode)")
|
||||||
|
if self._window_size <= 0 or self._length <= self._window_size:
|
||||||
|
raise IndexError(
|
||||||
|
f"Data too short for window: token_count={self._length}, "
|
||||||
|
f"window_size={self._window_size}"
|
||||||
|
)
|
||||||
|
if not 0 <= index < self.num_samples:
|
||||||
|
raise IndexError(
|
||||||
|
f"Sample index out of range: {index}, num_samples={self.num_samples}"
|
||||||
|
)
|
||||||
|
total = self._length
|
||||||
|
begin = min(index * self._stride, total - 1 - self._window_size)
|
||||||
|
end = min(begin + self._window_size, total - 1)
|
||||||
|
return begin, end
|
||||||
|
|
||||||
|
def _stream_keys(self) -> List[str]:
|
||||||
|
out: List[str] = []
|
||||||
|
for k, tensors in self._data.items():
|
||||||
|
if tensors and isinstance(tensors[0], list):
|
||||||
|
continue
|
||||||
|
out.append(k)
|
||||||
|
return out
|
||||||
|
|
||||||
|
def _record_keys(self) -> List[str]:
|
||||||
|
return list(self._data.keys())
|
||||||
|
|
||||||
|
def _normalize(
|
||||||
|
self,
|
||||||
|
raw: Dict[str, list],
|
||||||
|
offsets: Optional[Dict[str, List[int]]] = None,
|
||||||
|
):
|
||||||
|
"""Register segments and pre-compute indices for both access modes.
|
||||||
|
|
||||||
|
Stream mode: ``_cum[key]`` accumulates per-segment lengths so
|
||||||
|
``Streamable._fetch_stream_key`` can bisect across segments
|
||||||
|
without concatenation.
|
||||||
|
|
||||||
|
Record mode: if *offsets* is provided (bin layout),
|
||||||
|
``_offsets[key]`` stores cumulative per-record offsets into the
|
||||||
|
single concatenated segment. Otherwise, when
|
||||||
|
``segments_are_records`` is True (JSONL), ``_data[key]`` is
|
||||||
|
a per-record list and ``fetch_record`` indexes it directly.
|
||||||
|
|
||||||
|
Nested keys (GRPO ``responses``/``masks`` as
|
||||||
|
``List[List[Tensor]]``) are stored as-is and excluded from both
|
||||||
|
cumulative bookkeepings — they are only accessed record-by-record.
|
||||||
|
"""
|
||||||
|
flat_lengths = []
|
||||||
|
for key, tensors in raw.items():
|
||||||
|
self._data[key] = tensors
|
||||||
|
if not tensors:
|
||||||
|
self._cum[key] = []
|
||||||
|
flat_lengths.append(0)
|
||||||
|
continue
|
||||||
|
if isinstance(tensors[0], list):
|
||||||
|
self._cum[key] = []
|
||||||
|
continue
|
||||||
|
cum = []
|
||||||
|
total = 0
|
||||||
|
for t in tensors:
|
||||||
|
total += t.shape[0]
|
||||||
|
cum.append(total)
|
||||||
|
self._cum[key] = cum
|
||||||
|
flat_lengths.append(cum[-1] if cum else 0)
|
||||||
|
self._length = min(flat_lengths) if flat_lengths else 0
|
||||||
|
|
||||||
|
valid_offsets: Dict[str, List[int]] = {}
|
||||||
|
if offsets:
|
||||||
|
for key, off in offsets.items():
|
||||||
|
segs = self._data.get(key, [])
|
||||||
|
if len(segs) == 1 and len(off) > 1:
|
||||||
|
valid_offsets[key] = off
|
||||||
|
elif len(segs) > 1:
|
||||||
|
logger.warning(
|
||||||
|
"Key '%s' has %d segments with offsets — record mode "
|
||||||
|
"disabled for this key (multi-shard bin+offsets not "
|
||||||
|
"supported). Merge shards or use JSONL.",
|
||||||
|
key,
|
||||||
|
len(segs),
|
||||||
|
)
|
||||||
|
self._offsets = valid_offsets
|
||||||
|
if valid_offsets:
|
||||||
|
record_counts = [len(v) - 1 for v in valid_offsets.values()]
|
||||||
|
self._num_records = min(record_counts) if record_counts else 0
|
||||||
|
elif self.segments_are_records:
|
||||||
|
per_record_counts = []
|
||||||
|
for key, tensors in self._data.items():
|
||||||
|
if tensors and isinstance(tensors[0], list):
|
||||||
|
continue
|
||||||
|
per_record_counts.append(len(tensors))
|
||||||
|
self._num_records = min(per_record_counts) if per_record_counts else 0
|
||||||
|
else:
|
||||||
|
self._num_records = 0
|
||||||
|
|
||||||
|
|
||||||
|
class Streamable:
|
||||||
|
"""Mixin granting raw token-stream access via :meth:`fetch`.
|
||||||
|
|
||||||
|
Stateless trait relying on ``self._data``, ``self._cum``,
|
||||||
|
``self._length`` maintained by :class:`Store`. Stream mode is
|
||||||
|
active when the owning store has ``window_size > 0``; for stores
|
||||||
|
that can also serve record access (JSONL/bin+offsets), the
|
||||||
|
``fetch_record`` API from :class:`Recordable` is used instead.
|
||||||
|
"""
|
||||||
|
|
||||||
def fetch(
|
def fetch(
|
||||||
self,
|
self,
|
||||||
begin: int,
|
begin: int,
|
||||||
end: int,
|
end: int,
|
||||||
keys: Union[str, List[str]],
|
keys: Union[str, List[str]],
|
||||||
):
|
):
|
||||||
if not self._data:
|
return _stream_fetch(self, begin, end, keys)
|
||||||
raise RuntimeError("Store not loaded")
|
|
||||||
if not (0 <= begin < self._length and 0 <= end <= self._length):
|
|
||||||
raise ValueError(
|
|
||||||
f"Index out of bounds: begin={begin}, end={end}, length={self._length}"
|
|
||||||
)
|
|
||||||
if isinstance(keys, str):
|
|
||||||
return self._fetch_key(keys, begin, end)
|
|
||||||
return {k: self._fetch_key(k, begin, end) for k in keys}
|
|
||||||
|
|
||||||
def _fetch_key(self, key: str, begin: int, end: int) -> Tensor:
|
|
||||||
"""Fetch slice [begin, end) across potentially multiple segments."""
|
|
||||||
segments = self._data[key]
|
|
||||||
cum = self._cum[key]
|
|
||||||
seg_start = bisect.bisect_right(cum, begin)
|
|
||||||
seg_end = bisect.bisect_left(cum, end)
|
|
||||||
|
|
||||||
results = []
|
def _stream_fetch(self, begin: int, end: int, keys: Union[str, List[str]]):
|
||||||
for i in range(seg_start, seg_end + 1):
|
if not getattr(self, "_data", None):
|
||||||
prev = cum[i - 1] if i > 0 else 0
|
raise RuntimeError("Store not loaded")
|
||||||
s = max(begin - prev, 0)
|
if not (0 <= begin < self._length and 0 <= end <= self._length):
|
||||||
e = min(end - prev, segments[i].shape[0])
|
raise ValueError(
|
||||||
results.append(segments[i][s:e])
|
f"Index out of bounds: begin={begin}, end={end}, length={self._length}"
|
||||||
|
|
||||||
return results[0] if len(results) == 1 else torch.cat(results, dim=0)
|
|
||||||
|
|
||||||
def _normalize(self, raw: Dict[str, List[Tensor]]):
|
|
||||||
"""Register segments and pre-compute cumulative lengths.
|
|
||||||
|
|
||||||
Does NOT concatenate — segments are kept as-is to avoid OOM on
|
|
||||||
large datasets. Sets ``self._length`` to the minimum total
|
|
||||||
element count across all keys.
|
|
||||||
"""
|
|
||||||
for key, tensors in raw.items():
|
|
||||||
self._data[key] = tensors
|
|
||||||
cum = []
|
|
||||||
total = 0
|
|
||||||
for t in tensors:
|
|
||||||
total += t.shape[0]
|
|
||||||
cum.append(total)
|
|
||||||
self._cum[key] = cum
|
|
||||||
self._length = (
|
|
||||||
min((cum[-1] if cum else 0) for cum in self._cum.values())
|
|
||||||
if self._cum
|
|
||||||
else 0
|
|
||||||
)
|
)
|
||||||
|
if isinstance(keys, str):
|
||||||
|
return _fetch_stream_key(self, keys, begin, end)
|
||||||
|
return {k: _fetch_stream_key(self, k, begin, end) for k in keys}
|
||||||
|
|
||||||
|
|
||||||
|
def _fetch_stream_key(self, key: str, begin: int, end: int) -> Tensor:
|
||||||
|
segments = self._data[key]
|
||||||
|
cum = self._cum[key]
|
||||||
|
seg_start = bisect.bisect_right(cum, begin)
|
||||||
|
seg_end = bisect.bisect_left(cum, end)
|
||||||
|
|
||||||
|
results = []
|
||||||
|
for i in range(seg_start, seg_end + 1):
|
||||||
|
prev = cum[i - 1] if i > 0 else 0
|
||||||
|
s = max(begin - prev, 0)
|
||||||
|
e = min(end - prev, segments[i].shape[0])
|
||||||
|
results.append(segments[i][s:e])
|
||||||
|
|
||||||
|
return results[0] if len(results) == 1 else torch.cat(results, dim=0)
|
||||||
|
|
||||||
|
|
||||||
|
class Recordable:
|
||||||
|
"""Mixin granting raw record access via :meth:`fetch_record`.
|
||||||
|
|
||||||
|
Stateless trait relying on ``self._data``, ``self._offsets``,
|
||||||
|
``self._num_records`` maintained by :class:`Store`.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def fetch_record(
|
||||||
|
self,
|
||||||
|
index: int,
|
||||||
|
keys: Union[str, List[str]],
|
||||||
|
):
|
||||||
|
return _record_fetch(self, index, keys)
|
||||||
|
|
||||||
|
|
||||||
|
def _record_fetch(self, index: int, keys: Union[str, List[str]]):
|
||||||
|
if not getattr(self, "_data", None) and self._num_records == 0:
|
||||||
|
raise RuntimeError("Store not loaded")
|
||||||
|
if not 0 <= index < self._num_records:
|
||||||
|
raise ValueError(
|
||||||
|
f"Record index out of bounds: {index}, num_records={self._num_records}"
|
||||||
|
)
|
||||||
|
if isinstance(keys, str):
|
||||||
|
return _fetch_record_key(self, keys, index)
|
||||||
|
return {k: _fetch_record_key(self, k, index) for k in keys}
|
||||||
|
|
||||||
|
|
||||||
|
def _fetch_record_key(self, key: str, index: int):
|
||||||
|
offsets = self._offsets.get(key)
|
||||||
|
if offsets:
|
||||||
|
start = offsets[index]
|
||||||
|
end = (
|
||||||
|
offsets[index + 1]
|
||||||
|
if index + 1 < len(offsets)
|
||||||
|
else self._data[key][0].shape[0]
|
||||||
|
)
|
||||||
|
return self._data[key][0][start:end]
|
||||||
|
return self._data[key][index]
|
||||||
|
|
||||||
|
|
||||||
class StoreFactory(BaseFactory["Store"]):
|
class StoreFactory(BaseFactory["Store"]):
|
||||||
"""Factory for creating Store instances by type name.
|
"""Factory for creating Store instances by type name."""
|
||||||
|
|
||||||
Example::
|
|
||||||
|
|
||||||
@StoreFactory.register("custom")
|
|
||||||
class CustomStore(Store):
|
|
||||||
...
|
|
||||||
"""
|
|
||||||
|
|
||||||
|
|
||||||
@StoreFactory.register("h5")
|
|
||||||
class H5Store(Store):
|
|
||||||
"""HDF5-based storage backend (pre-tokenized data)."""
|
|
||||||
|
|
||||||
def load(self, path: str):
|
|
||||||
self._normalize(load_h5(path))
|
|
||||||
|
|
||||||
|
|
||||||
@StoreFactory.register("bin")
|
@StoreFactory.register("bin")
|
||||||
class MmapStore(Store):
|
class MmapStore(Store, Streamable, Recordable):
|
||||||
"""Memory-mapped binary storage backend.
|
"""Memory-mapped binary storage backend.
|
||||||
|
|
||||||
Each key is a single .bin file backed by ``np.memmap(mode="r")``.
|
Each key is a single .bin file backed by ``np.memmap(mode="r")``.
|
||||||
No per-process memory duplication — all DataLoader workers share the
|
No per-process memory duplication — all DataLoader workers share the
|
||||||
same OS page-cache pages.
|
same OS page-cache pages.
|
||||||
|
|
||||||
Format on disk::
|
Supports both access modes:
|
||||||
|
|
||||||
data_root/
|
- **Stream**: always available via :meth:`fetch`.
|
||||||
meta.json # {key: {shape, dtype}, ...}
|
- **Record** (``fetch_record(i, key)``): only when ``meta.json``
|
||||||
<key>.bin # raw numpy array, one per key
|
contains per-record ``offsets`` (written via
|
||||||
|
``save_bin(..., record_keys=...)``). Legacy bin files without
|
||||||
|
offsets have ``num_records == 0`` and ``len(store)`` reflects the
|
||||||
|
windowed sample count when ``window_size > 0``.
|
||||||
|
|
||||||
|
``segments_are_records`` is ``False`` here (bin segments are
|
||||||
|
contiguous streams, not per-record) — record access is driven
|
||||||
|
purely by ``_offsets``.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def load(self, path: str):
|
segments_are_records = False
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
window_size: int = 0,
|
||||||
|
stride: Optional[int] = None,
|
||||||
|
):
|
||||||
|
super().__init__(window_size=window_size, stride=stride)
|
||||||
|
self._mmap_refs: List[Tensor] = []
|
||||||
|
|
||||||
|
def load(self, path: str, **kwargs):
|
||||||
self._mmap_refs = []
|
self._mmap_refs = []
|
||||||
root = Path(path)
|
root = Path(path)
|
||||||
all_raw: Dict[str, List[Tensor]] = {}
|
all_raw: Dict[str, List[Tensor]] = {}
|
||||||
|
all_offsets: Dict[str, List[int]] = {}
|
||||||
meta_paths = [
|
meta_paths = [
|
||||||
Path(p) for p in glob.glob(str(root / "**" / "meta.json"), recursive=True)
|
Path(p) for p in glob.glob(str(root / "**" / "meta.json"), recursive=True)
|
||||||
]
|
]
|
||||||
for meta_path in meta_paths:
|
for meta_path in meta_paths:
|
||||||
raw = load_bin(str(meta_path.parent))
|
raw = load_bin(str(meta_path.parent))
|
||||||
|
off = load_bin_offsets(str(meta_path.parent))
|
||||||
for key, tensors in raw.items():
|
for key, tensors in raw.items():
|
||||||
if key not in all_raw:
|
if key not in all_raw:
|
||||||
all_raw[key] = []
|
all_raw[key] = []
|
||||||
all_raw[key].extend(tensors)
|
all_raw[key].extend(tensors)
|
||||||
|
for key, o in off.items():
|
||||||
|
if key not in all_offsets:
|
||||||
|
all_offsets[key] = []
|
||||||
|
all_offsets[key].extend(o)
|
||||||
if not meta_paths:
|
if not meta_paths:
|
||||||
raise FileNotFoundError(f"No meta.json found under {path}")
|
raise FileNotFoundError(f"No meta.json found under {path}")
|
||||||
self._normalize(all_raw)
|
self._normalize(all_raw, offsets=all_offsets or None)
|
||||||
for tensors in self._data.values():
|
for tensors in self._data.values():
|
||||||
self._mmap_refs.extend(tensors)
|
self._mmap_refs.extend(tensors)
|
||||||
|
|
||||||
|
|
||||||
@StoreFactory.register("jsonl")
|
class JsonlSource:
|
||||||
class JsonlStore(Store):
|
"""Read raw JSON records from a ``.jsonl`` file or directory.
|
||||||
"""On-the-fly tokenization store for raw JSONL files.
|
|
||||||
|
|
||||||
A JSONL dataset directory contains ``*.jsonl`` files plus a
|
A thin reader used by :class:`JsonlStore` in processor mode — holds
|
||||||
``dataset_config.json`` file that follows the same schema as
|
no tokenizer, performs no tokenisation, just yields dicts.
|
||||||
:class:`PipelineConfig` with an additional ``tokenizer_path`` field.
|
"""
|
||||||
Records are tokenized when the store is loaded and concatenated into
|
|
||||||
segmented tensors matching the key layout expected by the dataset
|
def __init__(self, path: str):
|
||||||
classes (``sequence``, ``loss_mask``, ``position_ids``, ...).
|
self.path = Path(path)
|
||||||
|
self._records: Optional[List[dict]] = None
|
||||||
|
|
||||||
|
def load(self) -> List[dict]:
|
||||||
|
if self._records is None:
|
||||||
|
self._records = self._read(self.path)
|
||||||
|
return self._records
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _read(root: Path) -> List[dict]:
|
||||||
|
if root.is_file():
|
||||||
|
return JsonlSource._read_file(root)
|
||||||
|
return JsonlSource._read_dir(root)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _read_file(path: Path) -> List[dict]:
|
||||||
|
records: List[dict] = []
|
||||||
|
with open(path, "r", encoding="utf-8") as f:
|
||||||
|
for line in f:
|
||||||
|
line = line.strip()
|
||||||
|
if not line:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
records.append(json.loads(line))
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
logger.warning("Failed to parse JSON line in %s, skipping", path)
|
||||||
|
return records
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _read_dir(root: Path) -> List[dict]:
|
||||||
|
records: List[dict] = []
|
||||||
|
for jsonl_path in sorted(root.glob("*.jsonl")):
|
||||||
|
records.extend(JsonlSource._read_file(jsonl_path))
|
||||||
|
return records
|
||||||
|
|
||||||
|
|
||||||
|
@StoreFactory.register("jsonl")
|
||||||
|
class JsonlStore(Store, Streamable, Recordable):
|
||||||
|
"""JSONL reader with eager/lazy tokenisation modes.
|
||||||
|
|
||||||
|
A JSONL dataset is a ``.jsonl`` file or a directory of ``*.jsonl``
|
||||||
|
files plus (optionally) a ``dataset_config.json`` describing the
|
||||||
|
tokenization pipeline.
|
||||||
|
|
||||||
|
Three ways to supply an eager transform (first match wins):
|
||||||
|
|
||||||
|
- **Explicit** (``transform=``): caller-built
|
||||||
|
:class:`TokenizeTransform` applied eagerly.
|
||||||
|
- **Config file**: ``dataset_config.json`` alongside the ``*.jsonl``
|
||||||
|
files — loaded via :meth:`TokenizeTransform.from_config_file`.
|
||||||
|
- **Default messages** (``tokenizer_path=`` given, no config file):
|
||||||
|
a built-in chatml config that tokenises the ``messages`` field,
|
||||||
|
masking every role except ``assistant`` (loss on assistant only).
|
||||||
|
Lets SFT/SEQ train straight from a chat-style JSONL directory
|
||||||
|
without a hand-written config.
|
||||||
|
|
||||||
|
Two tokenisation modes, selected at :meth:`load` time:
|
||||||
|
|
||||||
|
- **Eager** (default): applies the transform to every record at load
|
||||||
|
time and registers per-key tensors via ``_normalize``. Both
|
||||||
|
``fetch`` (stream) and ``fetch_record`` (record) work.
|
||||||
|
- **Lazy** (``processor=fn`` passed): keeps raw records and defers
|
||||||
|
tokenisation to ``fetch_record``. Only record access works —
|
||||||
|
``len(store)`` returns ``num_records``; stream primitives raise.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
CONFIG_NAME = "dataset_config.json"
|
CONFIG_NAME = "dataset_config.json"
|
||||||
|
segments_are_records = True
|
||||||
|
|
||||||
def load(self, path: str):
|
_DEFAULT_MESSAGES_CONFIG = {
|
||||||
root = Path(path)
|
"version": 1,
|
||||||
config_path = root / self.CONFIG_NAME
|
"input": {
|
||||||
if not config_path.exists():
|
"sections": [{"field": "messages", "action": "$role", "template": True}]
|
||||||
raise FileNotFoundError(
|
},
|
||||||
f"JSONL dataset config not found: {config_path}. "
|
"mask": {"system": "mask", "user": "mask", "assistant": "train"},
|
||||||
f"Expected {self.CONFIG_NAME} alongside *.jsonl files."
|
"mask_default": "mask",
|
||||||
|
"output": {"position_ids_mode": "doc_reset"},
|
||||||
|
}
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
window_size: int = 0,
|
||||||
|
stride: Optional[int] = None,
|
||||||
|
):
|
||||||
|
super().__init__(window_size=window_size, stride=stride)
|
||||||
|
self._source: Optional[JsonlSource] = None
|
||||||
|
self._processor: Optional[Callable[[dict], Dict[str, Tensor]]] = None
|
||||||
|
self._keys_cache: Optional[List[str]] = None
|
||||||
|
|
||||||
|
def load(self, path: str, transform=None, processor=None, **kwargs):
|
||||||
|
self._source = JsonlSource(path)
|
||||||
|
records = self._source.load()
|
||||||
|
|
||||||
|
if processor is not None:
|
||||||
|
self._processor = processor
|
||||||
|
self._num_records = len(records)
|
||||||
|
return
|
||||||
|
|
||||||
|
if transform is None:
|
||||||
|
root = Path(path)
|
||||||
|
config_path = root / self.CONFIG_NAME if root.is_dir() else None
|
||||||
|
if config_path is not None and config_path.exists():
|
||||||
|
transform = TokenizeTransform.from_config_file(str(config_path))
|
||||||
|
else:
|
||||||
|
tokenizer_path = kwargs.get("tokenizer_path")
|
||||||
|
if not tokenizer_path:
|
||||||
|
raise FileNotFoundError(
|
||||||
|
f"JSONL dataset config not found. Expected "
|
||||||
|
f"{self.CONFIG_NAME} alongside *.jsonl files, pass an "
|
||||||
|
f"explicit transform, pass processor= for lazy "
|
||||||
|
f"on-the-fly tokenisation, or pass tokenizer_path= to "
|
||||||
|
f"use the built-in messages config."
|
||||||
|
)
|
||||||
|
config = PipelineConfig.from_dict(self._DEFAULT_MESSAGES_CONFIG)
|
||||||
|
transform = TokenizeTransform(config, tokenizer_path)
|
||||||
|
|
||||||
|
transformed = transform.apply(records)
|
||||||
|
self._normalize(transformed)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def keys(self) -> List[str]:
|
||||||
|
if self._processor is not None:
|
||||||
|
if self._keys_cache is None and self._num_records > 0:
|
||||||
|
sample = self._processor(self._source.load()[0])
|
||||||
|
self._keys_cache = list(sample.keys())
|
||||||
|
return self._keys_cache or []
|
||||||
|
return list(self._data.keys())
|
||||||
|
|
||||||
|
def fetch_record(self, index: int, keys: Union[str, List[str]]):
|
||||||
|
if self._processor is not None:
|
||||||
|
if not 0 <= index < self._num_records:
|
||||||
|
raise ValueError(
|
||||||
|
f"Record index out of bounds: {index}, "
|
||||||
|
f"num_records={self._num_records}"
|
||||||
|
)
|
||||||
|
record = self._source.load()[index]
|
||||||
|
data = self._processor(record)
|
||||||
|
if isinstance(keys, str):
|
||||||
|
return data[keys]
|
||||||
|
return {k: data[k] for k in keys}
|
||||||
|
return _record_fetch(self, index, keys)
|
||||||
|
|
||||||
|
def fetch(self, begin: int, end: int, keys: Union[str, List[str]]):
|
||||||
|
if self._processor is not None:
|
||||||
|
raise RuntimeError(
|
||||||
|
"JsonlStore in lazy (processor) mode does not support "
|
||||||
|
"stream fetch(); use fetch_record() instead."
|
||||||
)
|
)
|
||||||
|
return _stream_fetch(self, begin, end, keys)
|
||||||
|
|
||||||
with open(config_path, "r", encoding="utf-8") as f:
|
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||||
raw_config = json.load(f)
|
if self._processor is not None:
|
||||||
|
return self.fetch_record(index, self._record_keys())
|
||||||
tokenizer_path = raw_config.pop("tokenizer_path", None)
|
return super().__getitem__(index)
|
||||||
if tokenizer_path is None:
|
|
||||||
raise ValueError(
|
|
||||||
f"JSONL dataset config must specify 'tokenizer_path': {config_path}"
|
|
||||||
)
|
|
||||||
|
|
||||||
self.config = PipelineConfig.from_dict(raw_config)
|
|
||||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
|
|
||||||
mask_builder = MaskBuilderFactory.create("sectioned")
|
|
||||||
position_strategy = PositionIdStrategyFactory.create(
|
|
||||||
self.config.output.position_ids_mode
|
|
||||||
)
|
|
||||||
|
|
||||||
raw: Dict[str, List[Tensor]] = {}
|
|
||||||
doc_sequences: List[List[int]] = []
|
|
||||||
|
|
||||||
for jsonl_path in sorted(root.glob("*.jsonl")):
|
|
||||||
with open(jsonl_path, "r", encoding="utf-8") as f:
|
|
||||||
for line in f:
|
|
||||||
line = line.strip()
|
|
||||||
if not line:
|
|
||||||
continue
|
|
||||||
try:
|
|
||||||
item = json.loads(line)
|
|
||||||
except json.JSONDecodeError:
|
|
||||||
logger.warning(
|
|
||||||
"Failed to parse JSON line in %s, skipping", jsonl_path
|
|
||||||
)
|
|
||||||
continue
|
|
||||||
|
|
||||||
result = mask_builder.build(item, self.config, tokenizer)
|
|
||||||
if result is None:
|
|
||||||
continue
|
|
||||||
|
|
||||||
result.pop("domain", None)
|
|
||||||
primary_ids = self._primary_ids(result)
|
|
||||||
if not primary_ids:
|
|
||||||
continue
|
|
||||||
|
|
||||||
doc_sequences.append(primary_ids)
|
|
||||||
for key, ids in result.items():
|
|
||||||
if key not in raw:
|
|
||||||
raw[key] = []
|
|
||||||
raw[key].append(torch.tensor(ids, dtype=self._infer_dtype(ids)))
|
|
||||||
|
|
||||||
pos_ids = position_strategy.generate(doc_sequences)
|
|
||||||
if pos_ids:
|
|
||||||
raw["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
|
|
||||||
|
|
||||||
self._normalize(raw)
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _primary_ids(result: dict) -> List[int]:
|
|
||||||
"""Return the first integer list in *result* as the primary id sequence."""
|
|
||||||
for val in result.values():
|
|
||||||
if isinstance(val, list) and val and isinstance(val[0], int):
|
|
||||||
return val
|
|
||||||
return []
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _infer_dtype(ids: List) -> torch.dtype:
|
|
||||||
"""Infer tensor dtype from the first element of a token/value list."""
|
|
||||||
if ids and isinstance(ids[0], float):
|
|
||||||
return torch.float32
|
|
||||||
return torch.int32
|
|
||||||
|
|||||||
@@ -1,19 +1,49 @@
|
|||||||
"""CUDA attention kernel wrappers with torch fallback.
|
"""CUDA attention kernel wrappers with torch fallback.
|
||||||
|
|
||||||
Public API:
|
Public API:
|
||||||
- ``gqa_decode_attn`` — single-query decode attention
|
- ``attn_decode`` — single-query decode attention
|
||||||
- ``gqa_prefill_attn`` — multi-query prefill attention
|
- ``attn_prefill`` — multi-query prefill attention
|
||||||
|
- ``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
|
||||||
|
|
||||||
Each wrapper dispatches to its compiled CUDA kernel (``astrai.extension.gqa_*``)
|
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
|
||||||
when available, otherwise falls back to ``torch.nn.functional.scaled_dot_product_attention``.
|
(blhd). Scale is always ``1/sqrt(head_dim)``.
|
||||||
|
|
||||||
|
Each wrapper calls its compiled CUDA kernel directly. Fallback to torch
|
||||||
|
SDPA is handled by the attention backend, not the wrapper functions.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
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 gqa_decode_attn, gqa_prefill_attn
|
from astrai.extension.rotary_backend import apply_rotary_emb
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"gqa_decode_attn",
|
"ATTN_BACKEND",
|
||||||
"gqa_prefill_attn",
|
"AttentionBackend",
|
||||||
|
"CudaBackend",
|
||||||
|
"TorchNativeBackend",
|
||||||
|
"attention",
|
||||||
|
"attn_backend",
|
||||||
|
"get_backend",
|
||||||
|
"attn_decode",
|
||||||
|
"attn_paged_decode",
|
||||||
|
"attn_prefill",
|
||||||
"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 = ["gqa_decode_attn", "gqa_prefill_attn"]
|
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,86 +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.
|
|
||||||
|
|
||||||
Add new kernel wrappers here; split into per-variant files only if this file
|
|
||||||
grows large.
|
|
||||||
"""
|
|
||||||
|
|
||||||
import torch
|
|
||||||
import torch.nn.functional as F
|
|
||||||
|
|
||||||
from astrai.extension.loader import _available, _modules
|
|
||||||
|
|
||||||
|
|
||||||
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 _torch_fallback(
|
|
||||||
q: torch.Tensor,
|
|
||||||
k: torch.Tensor,
|
|
||||||
v: torch.Tensor,
|
|
||||||
mask: torch.Tensor | None,
|
|
||||||
is_causal: bool,
|
|
||||||
scale: float | None,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
"""Reference attention via ``scaled_dot_product_attention``."""
|
|
||||||
k, v = _expand_kv_heads(k, v, q.size(1))
|
|
||||||
attn_mask = mask[:, None, None, :] if mask is not None else None
|
|
||||||
return F.scaled_dot_product_attention(
|
|
||||||
q, k, v, attn_mask=attn_mask, is_causal=is_causal and mask is None, scale=scale
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def gqa_decode_attn(
|
|
||||||
q: torch.Tensor,
|
|
||||||
k: torch.Tensor,
|
|
||||||
v: torch.Tensor,
|
|
||||||
mask: torch.Tensor | None = None,
|
|
||||||
is_causal: bool = False,
|
|
||||||
causal_offset: int = 0,
|
|
||||||
scale: float | None = None,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
if _available["gqa_decode_attn"]:
|
|
||||||
return _modules["gqa_decode_attn"].gqa_decode_attn(
|
|
||||||
q,
|
|
||||||
k,
|
|
||||||
v,
|
|
||||||
mask=mask,
|
|
||||||
is_causal=is_causal,
|
|
||||||
causal_offset=causal_offset,
|
|
||||||
scale=scale,
|
|
||||||
)
|
|
||||||
return _torch_fallback(q, k, v, mask, is_causal, scale)
|
|
||||||
|
|
||||||
|
|
||||||
def gqa_prefill_attn(
|
|
||||||
q: torch.Tensor,
|
|
||||||
k: torch.Tensor,
|
|
||||||
v: torch.Tensor,
|
|
||||||
mask: torch.Tensor | None = None,
|
|
||||||
is_causal: bool = False,
|
|
||||||
causal_offset: int = 0,
|
|
||||||
scale: float | None = None,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
if _available["gqa_prefill_attn"]:
|
|
||||||
return _modules["gqa_prefill_attn"].gqa_prefill_attn(
|
|
||||||
q,
|
|
||||||
k,
|
|
||||||
v,
|
|
||||||
mask=mask,
|
|
||||||
is_causal=is_causal,
|
|
||||||
causal_offset=causal_offset,
|
|
||||||
scale=scale,
|
|
||||||
)
|
|
||||||
return _torch_fallback(q, k, v, mask, is_causal, scale)
|
|
||||||
@@ -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
-32
@@ -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,10 +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 = {}
|
||||||
cls._component_base = _resolve_type(arg, cls)
|
cls._component_base = _resolve_base_type(arg, cls)
|
||||||
return
|
return
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -79,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
|
||||||
@@ -92,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()
|
||||||
@@ -111,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."""
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ Layers:
|
|||||||
- protocols/: Response builders (OpenAI, Anthropic)
|
- protocols/: Response builders (OpenAI, Anthropic)
|
||||||
- transport/: SSE transport utilities
|
- transport/: SSE transport utilities
|
||||||
- engine.py: Facade (InferenceEngine), Value Object (GenerationRequest)
|
- engine.py: Facade (InferenceEngine), Value Object (GenerationRequest)
|
||||||
- sample.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy)
|
- sample.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy, FrequencyPenaltyStrategy)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from astrai.inference.api import (
|
from astrai.inference.api import (
|
||||||
@@ -30,26 +30,22 @@ 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
|
||||||
from astrai.inference.sample import (
|
from astrai.inference.sample import (
|
||||||
BaseSamplingStrategy,
|
BaseSamplingStrategy,
|
||||||
|
FrequencyPenaltyStrategy,
|
||||||
SamplingPipeline,
|
SamplingPipeline,
|
||||||
TemperatureStrategy,
|
TemperatureStrategy,
|
||||||
TopKStrategy,
|
TopKStrategy,
|
||||||
@@ -67,22 +63,18 @@ __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",
|
||||||
"TemperatureStrategy",
|
"TemperatureStrategy",
|
||||||
"TopKStrategy",
|
"TopKStrategy",
|
||||||
"TopPStrategy",
|
"TopPStrategy",
|
||||||
|
"FrequencyPenaltyStrategy",
|
||||||
"SamplingPipeline",
|
"SamplingPipeline",
|
||||||
"ProtocolHandler",
|
"ProtocolHandler",
|
||||||
"StopChecker",
|
"StopChecker",
|
||||||
|
|||||||
@@ -21,7 +21,6 @@ logger = logging.getLogger(__name__)
|
|||||||
_UNSUPPORTED_PARAMS = (
|
_UNSUPPORTED_PARAMS = (
|
||||||
"n",
|
"n",
|
||||||
"presence_penalty",
|
"presence_penalty",
|
||||||
"frequency_penalty",
|
|
||||||
"logit_bias",
|
"logit_bias",
|
||||||
"user",
|
"user",
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -125,6 +125,7 @@ class ProtocolHandler:
|
|||||||
temperature=self.request.temperature,
|
temperature=self.request.temperature,
|
||||||
top_p=self.request.top_p,
|
top_p=self.request.top_p,
|
||||||
top_k=self.request.top_k,
|
top_k=self.request.top_k,
|
||||||
|
frequency_penalty=getattr(self.request, "frequency_penalty", 0.0),
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.request.stream:
|
if self.request.stream:
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ Subclasses may optionally consume ``token_ids`` for token-level parsing
|
|||||||
(e.g. Harmony / VLM-style parsers).
|
(e.g. Harmony / VLM-style parsers).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
import re
|
import re
|
||||||
import uuid
|
import uuid
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
@@ -21,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"):
|
||||||
@@ -50,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
|
||||||
@@ -117,6 +113,29 @@ def _parse_tool_call_json(json_str: str, complete: bool):
|
|||||||
|
|
||||||
Returns ``(name, args, valid)``.
|
Returns ``(name, args, valid)``.
|
||||||
"""
|
"""
|
||||||
|
if complete:
|
||||||
|
try:
|
||||||
|
obj = json.loads(json_str)
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
return None, "", False
|
||||||
|
name = obj.get("name")
|
||||||
|
if not isinstance(name, str) or not name:
|
||||||
|
return None, "", False
|
||||||
|
args = obj.get("arguments")
|
||||||
|
if isinstance(args, dict):
|
||||||
|
if not args:
|
||||||
|
args = ""
|
||||||
|
else:
|
||||||
|
args = json.dumps(args, ensure_ascii=False)
|
||||||
|
args = args[1:-1].rstrip()
|
||||||
|
elif isinstance(args, list):
|
||||||
|
args = json.dumps(args, ensure_ascii=False) if args else ""
|
||||||
|
elif isinstance(args, str):
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
args = str(args) if args is not None else ""
|
||||||
|
return name, args, True
|
||||||
|
|
||||||
name_match = re.search(r'"name"\s*:\s*"([^"]*)"', json_str)
|
name_match = re.search(r'"name"\s*:\s*"([^"]*)"', json_str)
|
||||||
if not name_match:
|
if not name_match:
|
||||||
return None, "", False
|
return None, "", False
|
||||||
@@ -127,8 +146,6 @@ def _parse_tool_call_json(json_str: str, complete: bool):
|
|||||||
return name, "", True
|
return name, "", True
|
||||||
|
|
||||||
raw = args_match.group(1).rstrip()
|
raw = args_match.group(1).rstrip()
|
||||||
if complete and raw.endswith("}"):
|
|
||||||
raw = raw[:-1].rstrip()
|
|
||||||
if raw.startswith("{"):
|
if raw.startswith("{"):
|
||||||
inner = raw[1:].rstrip()
|
inner = raw[1:].rstrip()
|
||||||
if inner.endswith("}"):
|
if inner.endswith("}"):
|
||||||
@@ -156,9 +173,6 @@ def _find_tool_calls(text: str, start_pos: int = 0):
|
|||||||
break
|
break
|
||||||
|
|
||||||
json_str = text[brace:end]
|
json_str = text[brace:end]
|
||||||
if not _TOOL_CALL_HEAD_RE.search(json_str):
|
|
||||||
pos = end
|
|
||||||
continue
|
|
||||||
|
|
||||||
name, args, valid = _parse_tool_call_json(json_str, complete=True)
|
name, args, valid = _parse_tool_call_json(json_str, complete=True)
|
||||||
if not valid or name is None:
|
if not valid or name is None:
|
||||||
@@ -186,7 +200,7 @@ def _find_partial_tool_call(text: str, start_pos: int = 0):
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
json_str = text[brace:]
|
json_str = text[brace:]
|
||||||
if not _TOOL_CALL_HEAD_RE.search(json_str):
|
if '"name"' not in json_str:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
name, args, valid = _parse_tool_call_json(json_str, complete=False)
|
name, args, valid = _parse_tool_call_json(json_str, complete=False)
|
||||||
|
|||||||
@@ -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",
|
||||||
|
|||||||
+339
-337
@@ -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,392 +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, task_ids: List[str], total_len: int, device: torch.device
|
n_tokens: Total token slots for paged mode. None = contiguous mode
|
||||||
) -> CacheView: ...
|
(pre-allocates max_batch_size * max_seq_len).
|
||||||
|
"""
|
||||||
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
|
|
||||||
) -> 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
|
|
||||||
):
|
|
||||||
self._cache = cache
|
|
||||||
self._batch_indices = batch_indices
|
|
||||||
self._total_len = total_len
|
|
||||||
|
|
||||||
def write(self, layer_id: int, k: Tensor, v: Tensor):
|
|
||||||
seq_len = k.size(1)
|
|
||||||
start_pos = self._total_len - seq_len
|
|
||||||
indices = self._batch_indices
|
|
||||||
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
|
|
||||||
new_len = start_pos + seq_len
|
|
||||||
for s in indices.tolist():
|
|
||||||
cur = self._cache._slot_len.get(s, 0)
|
|
||||||
if new_len > cur:
|
|
||||||
self._cache._slot_len[s] = new_len
|
|
||||||
|
|
||||||
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]:
|
|
||||||
max_len = max(
|
|
||||||
self._cache._slot_len.get(int(s), 0) for s in self._batch_indices.tolist()
|
|
||||||
)
|
|
||||||
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:
|
||||||
|
return self._task_cached.get(task_id, 0)
|
||||||
|
|
||||||
|
def task_record_hashes(
|
||||||
|
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, task_ids: List[str], total_len: int, device: torch.device
|
self,
|
||||||
) -> ContiguousCacheView:
|
task_ids: List[str],
|
||||||
slots = [self._task_slot[tid] for tid in task_ids]
|
seq_lens: List[int],
|
||||||
batch_indices = torch.tensor(slots, dtype=torch.long, device=device)
|
device: torch.device,
|
||||||
return ContiguousCacheView(self, batch_indices, total_len)
|
start_pos: Optional[int] = None,
|
||||||
|
) -> KVCache:
|
||||||
|
req_indices = [self._task_req[tid] for tid in task_ids]
|
||||||
|
req_pool_indices = torch.tensor(req_indices, dtype=torch.long, device=device)
|
||||||
|
seq_lens_t = torch.tensor(seq_lens, dtype=torch.long, device=device)
|
||||||
|
|
||||||
|
if start_pos is not None:
|
||||||
|
seq_len = seq_lens[0]
|
||||||
|
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,
|
||||||
):
|
):
|
||||||
@@ -43,19 +43,42 @@ class Executor:
|
|||||||
)
|
)
|
||||||
|
|
||||||
task_ids = [t.task_id for t in tasks]
|
task_ids = [t.task_id for t in tasks]
|
||||||
|
position_ids = (
|
||||||
|
torch.arange(start_pos, prompt_len, dtype=torch.long, device=self.device)
|
||||||
|
.unsqueeze(0)
|
||||||
|
.expand(batch_sz, -1)
|
||||||
|
)
|
||||||
|
input_mask = position_ids.unsqueeze(-1) >= torch.arange(
|
||||||
|
prompt_len, device=self.device
|
||||||
|
)
|
||||||
|
|
||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
self.model(
|
self.model(
|
||||||
input_ids,
|
input_ids,
|
||||||
position_ids=torch.arange(
|
input_mask=input_mask,
|
||||||
start_pos, prompt_len, dtype=torch.long, device=self.device
|
position_ids=position_ids,
|
||||||
)
|
kv_cache=self.kv_cache.bind_tasks(
|
||||||
.unsqueeze(0)
|
task_ids, [prompt_len] * batch_sz, self.device, start_pos=start_pos
|
||||||
.expand(batch_sz, -1),
|
),
|
||||||
paged_cache=self.kv_cache.bind_tasks(task_ids, prompt_len, self.device),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def execute_decode(self, tasks: List[Task]) -> List[int]:
|
def execute_decode(
|
||||||
|
self, tasks: List[Task], return_logprobs: bool = False
|
||||||
|
) -> List[int]:
|
||||||
|
"""Decode next token for each task.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
return_logprobs: When ``True``, also record (and return)
|
||||||
|
the log-probability of each sampled token under the
|
||||||
|
post-strategy sampling distribution. The logprob is
|
||||||
|
appended to ``task.output_logprobs`` and the return
|
||||||
|
list becomes ``List[Tuple[int, float]]``.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
``List[int]`` of sampled token IDs, or
|
||||||
|
``List[Tuple[int, float]]`` of ``(token_id, logprob)`` when
|
||||||
|
``return_logprobs`` is ``True``.
|
||||||
|
"""
|
||||||
if not tasks:
|
if not tasks:
|
||||||
return []
|
return []
|
||||||
|
|
||||||
@@ -68,25 +91,84 @@ class Executor:
|
|||||||
position_ids = torch.tensor(
|
position_ids = torch.tensor(
|
||||||
[t.next_pos for t in tasks], dtype=torch.long, device=self.device
|
[t.next_pos for t in tasks], dtype=torch.long, device=self.device
|
||||||
)
|
)
|
||||||
total_len = position_ids.max().item() + 1
|
total_len = max(t.next_pos for t in tasks) + 1
|
||||||
|
input_mask = position_ids[:, None, None] >= torch.arange(
|
||||||
|
total_len, device=self.device
|
||||||
|
)
|
||||||
|
|
||||||
task_ids = [t.task_id for t in tasks]
|
task_ids = [t.task_id for t in tasks]
|
||||||
|
|
||||||
temperatures = torch.tensor([t.temperature for t in tasks], device=self.device)
|
temperatures = torch.tensor([t.temperature for t in tasks], device=self.device)
|
||||||
top_ks = torch.tensor([t.top_k for t in tasks], device=self.device)
|
top_ks = torch.tensor([t.top_k for t in tasks], device=self.device)
|
||||||
top_ps = torch.tensor([t.top_p for t in tasks], device=self.device)
|
top_ps = torch.tensor([t.top_p for t in tasks], device=self.device)
|
||||||
|
freq_penalties = torch.tensor(
|
||||||
|
[t.frequency_penalty for t in tasks], device=self.device
|
||||||
|
)
|
||||||
|
|
||||||
|
has_freq = bool((freq_penalties != 0).any())
|
||||||
|
if has_freq:
|
||||||
|
history_lists = []
|
||||||
|
history_lens = []
|
||||||
|
for t in tasks:
|
||||||
|
window = t.rep_window
|
||||||
|
prompt_part = t.prompt_ids[-window:]
|
||||||
|
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
|
||||||
|
padded_ids = torch.zeros(
|
||||||
|
len(tasks), max_len, dtype=torch.long, device=self.device
|
||||||
|
)
|
||||||
|
padded_mask = torch.zeros(
|
||||||
|
len(tasks), max_len, dtype=torch.bool, device=self.device
|
||||||
|
)
|
||||||
|
for i, h in enumerate(history_lists):
|
||||||
|
L = history_lens[i]
|
||||||
|
padded_ids[i, :L] = torch.as_tensor(
|
||||||
|
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),
|
||||||
paged_cache=self.kv_cache.bind_tasks(task_ids, total_len, self.device),
|
input_mask=input_mask,
|
||||||
|
kv_cache=self.kv_cache.bind_tasks(
|
||||||
|
task_ids,
|
||||||
|
[t.next_pos + 1 for t in tasks],
|
||||||
|
self.device,
|
||||||
|
),
|
||||||
position_ids=position_ids.unsqueeze(1),
|
position_ids=position_ids.unsqueeze(1),
|
||||||
)
|
)
|
||||||
logits = outputs["logits"][:, -1, :]
|
logits = outputs["logits"][:, -1, :]
|
||||||
|
|
||||||
|
if return_logprobs:
|
||||||
|
tokens, logprobs = sample(
|
||||||
|
logits,
|
||||||
|
temperature=temperatures,
|
||||||
|
top_k=top_ks,
|
||||||
|
top_p=top_ps,
|
||||||
|
frequency_penalty=freq_penalties,
|
||||||
|
input_ids=padded_ids,
|
||||||
|
input_mask=padded_mask,
|
||||||
|
return_logprobs=True,
|
||||||
|
)
|
||||||
|
tokens_list = tokens.tolist()
|
||||||
|
logprobs_list = logprobs.tolist()
|
||||||
|
for t, lp in zip(tasks, logprobs_list):
|
||||||
|
t.output_logprobs.append(float(lp))
|
||||||
|
return list(zip(tokens_list, logprobs_list))
|
||||||
|
|
||||||
return sample(
|
return sample(
|
||||||
logits,
|
logits,
|
||||||
temperature=temperatures,
|
temperature=temperatures,
|
||||||
top_k=top_ks,
|
top_k=top_ks,
|
||||||
top_p=top_ps,
|
top_p=top_ps,
|
||||||
|
frequency_penalty=freq_penalties,
|
||||||
|
input_ids=padded_ids,
|
||||||
|
input_mask=padded_mask,
|
||||||
).tolist()
|
).tolist()
|
||||||
|
|||||||
@@ -1,10 +1,11 @@
|
|||||||
import logging
|
import logging
|
||||||
import threading
|
import threading
|
||||||
|
import uuid
|
||||||
from typing import Any, Dict, List, Optional, Tuple
|
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
|
||||||
@@ -22,45 +23,43 @@ 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
|
||||||
|
|
||||||
if max_seq_len is not None:
|
if max_seq_len is not None:
|
||||||
self.max_seq_len = max_seq_len
|
self.max_seq_len = max_seq_len
|
||||||
elif config.max_len is not None:
|
elif config.max_position_embeddings is not None:
|
||||||
self.max_seq_len = config.max_len
|
self.max_seq_len = config.max_position_embeddings
|
||||||
else:
|
else:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"max_seq_len must be provided either as argument "
|
"max_seq_len must be provided either as argument "
|
||||||
"or in model config (config.max_len)"
|
"or in model config (config.max_position_embeddings)"
|
||||||
)
|
)
|
||||||
self.device = device or next(model.parameters()).device
|
self.device = device or next(model.parameters()).device
|
||||||
self.dtype = dtype or next(model.parameters()).dtype
|
self.dtype = dtype or next(model.parameters()).dtype
|
||||||
|
|
||||||
head_dim = config.dim // config.n_heads
|
head_dim = config.hidden_size // config.num_attention_heads
|
||||||
|
|
||||||
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.n_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.n_kv_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(
|
||||||
@@ -110,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)
|
||||||
]
|
]
|
||||||
@@ -138,36 +139,33 @@ class InferenceScheduler:
|
|||||||
t.task_id, t.prompt_ids, start_logical_page
|
t.task_id, t.prompt_ids, start_logical_page
|
||||||
)
|
)
|
||||||
|
|
||||||
pos_groups: Dict[int, List[Task]] = {}
|
decode_tasks = active
|
||||||
for t in self._task_mgr.get_active_tasks():
|
|
||||||
pos_groups.setdefault(t.next_pos, []).append(t)
|
|
||||||
|
|
||||||
for next_pos in sorted(pos_groups.keys()):
|
valid: List[Task] = []
|
||||||
group = sorted(pos_groups[next_pos], key=lambda t: t.task_id)
|
for t in decode_tasks:
|
||||||
|
if cache.task_extend(t.task_id, t.next_pos):
|
||||||
|
valid.append(t)
|
||||||
|
else:
|
||||||
|
t.status = TaskStatus.ABORTED
|
||||||
|
self._task_mgr.invoke_callback(t.task_id, STOP)
|
||||||
|
|
||||||
valid: List[Task] = []
|
if valid:
|
||||||
for t in group:
|
next_tokens = self._executor.execute_decode(valid)
|
||||||
if cache.task_extend(t.task_id, t.next_pos):
|
|
||||||
valid.append(t)
|
for t, ntok in zip(valid, next_tokens):
|
||||||
else:
|
t.output_ids.append(ntok)
|
||||||
t.status = TaskStatus.ABORTED
|
t.output_tokens += 1
|
||||||
|
new_text = t.decode_new_token(self._task_mgr.tokenizer)
|
||||||
|
if new_text:
|
||||||
|
self._task_mgr.invoke_callback(t.task_id, new_text)
|
||||||
|
|
||||||
|
for t in valid:
|
||||||
|
if t.is_finished(stop_ids):
|
||||||
|
remaining = t.flush_remaining(self._task_mgr.tokenizer)
|
||||||
|
if remaining:
|
||||||
|
self._task_mgr.invoke_callback(t.task_id, remaining)
|
||||||
self._task_mgr.invoke_callback(t.task_id, STOP)
|
self._task_mgr.invoke_callback(t.task_id, STOP)
|
||||||
|
|
||||||
if valid:
|
|
||||||
next_tokens = self._executor.execute_decode(valid)
|
|
||||||
|
|
||||||
for t, ntok in zip(valid, next_tokens):
|
|
||||||
t.output_ids.append(ntok)
|
|
||||||
t.output_tokens += 1
|
|
||||||
self._task_mgr.invoke_callback(
|
|
||||||
t.task_id,
|
|
||||||
self._task_mgr.tokenizer.decode([ntok]),
|
|
||||||
)
|
|
||||||
|
|
||||||
for t in valid:
|
|
||||||
if t.is_finished(stop_ids):
|
|
||||||
self._task_mgr.invoke_callback(t.task_id, STOP)
|
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
self._stop_event.set()
|
self._stop_event.set()
|
||||||
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
|
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
|
||||||
@@ -197,6 +195,117 @@ class InferenceScheduler:
|
|||||||
self._cache.task_free(task.task_id)
|
self._cache.task_free(task.task_id)
|
||||||
for task in self._task_mgr.get_waiting_tasks():
|
for task in self._task_mgr.get_waiting_tasks():
|
||||||
self._task_mgr.invoke_callback(task.task_id, STOP)
|
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||||
|
self._cache.task_free(task.task_id)
|
||||||
self._task_mgr.clear_queues()
|
self._task_mgr.clear_queues()
|
||||||
if torch.cuda.is_available():
|
if torch.cuda.is_available():
|
||||||
torch.cuda.empty_cache()
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
|
def run_batch(
|
||||||
|
self,
|
||||||
|
prompt_ids_list: List[List[int]],
|
||||||
|
*,
|
||||||
|
max_tokens: Optional[int] = None,
|
||||||
|
temperature: float = 1.0,
|
||||||
|
top_p: float = 1.0,
|
||||||
|
top_k: int = 50,
|
||||||
|
frequency_penalty: float = 0.0,
|
||||||
|
rep_window: int = 64,
|
||||||
|
return_logprobs: bool = False,
|
||||||
|
) -> List[List[int]]:
|
||||||
|
"""Synchronous batch generation without the scheduler thread.
|
||||||
|
|
||||||
|
Accepts already-tokenized prompts (no string round-trip) and runs
|
||||||
|
prefill + decode to completion on the calling thread. Designed for
|
||||||
|
RL rollout, where logprobs of the behaviour policy must be collected
|
||||||
|
alongside generated tokens.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
prompt_ids_list: ``B`` prompts, each a list of token IDs.
|
||||||
|
max_tokens: Maximum tokens to generate per prompt. ``None``
|
||||||
|
uses ``self.max_seq_len - len(prompt_ids)``.
|
||||||
|
temperature/top_p/top_k/frequency_penalty/rep_window: Sampling
|
||||||
|
parameters (uniform across the batch).
|
||||||
|
return_logprobs: If ``True``, return ``(token_ids, logprobs)``
|
||||||
|
tuples per prompt (logprobs aligned 1-to-1 with token_ids).
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
``List[List[int]]`` of generated token IDs per prompt, or —
|
||||||
|
when ``return_logprobs`` is ``True`` —
|
||||||
|
``List[Tuple[List[int], List[float]]]``.
|
||||||
|
"""
|
||||||
|
stop_ids = self._task_mgr.tokenizer.stop_ids
|
||||||
|
cache = self._cache
|
||||||
|
seq_cap = self.max_seq_len
|
||||||
|
|
||||||
|
tasks: List[Task] = []
|
||||||
|
for ids in prompt_ids_list:
|
||||||
|
if len(ids) >= seq_cap:
|
||||||
|
tasks.append(None)
|
||||||
|
continue
|
||||||
|
t_max = max_tokens
|
||||||
|
if t_max is None:
|
||||||
|
t_max = seq_cap - len(ids)
|
||||||
|
else:
|
||||||
|
t_max = min(t_max, seq_cap - len(ids))
|
||||||
|
task = Task(
|
||||||
|
task_id=f"batch_{uuid.uuid4().hex[:8]}",
|
||||||
|
prompt_ids=list(ids),
|
||||||
|
max_tokens=t_max,
|
||||||
|
temperature=temperature,
|
||||||
|
top_p=top_p,
|
||||||
|
top_k=top_k,
|
||||||
|
frequency_penalty=frequency_penalty,
|
||||||
|
rep_window=rep_window,
|
||||||
|
)
|
||||||
|
if not cache.task_alloc(task.task_id, task.prompt_ids):
|
||||||
|
tasks.append(None)
|
||||||
|
continue
|
||||||
|
task.input_tokens = len(task.prompt_ids)
|
||||||
|
tasks.append(task)
|
||||||
|
|
||||||
|
try:
|
||||||
|
live = [t for t in tasks if t is not None]
|
||||||
|
prefill_groups: Dict[Tuple[int, int], List[Task]] = {}
|
||||||
|
for t in live:
|
||||||
|
key = (len(t.prompt_ids), cache.task_cached(t.task_id))
|
||||||
|
prefill_groups.setdefault(key, []).append(t)
|
||||||
|
for (prompt_len, start_pos), group in prefill_groups.items():
|
||||||
|
self._executor.execute_prefill(group, prompt_len, start_pos)
|
||||||
|
|
||||||
|
while live:
|
||||||
|
valid: List[Task] = []
|
||||||
|
for t in sorted(live, key=lambda x: x.task_id):
|
||||||
|
if cache.task_extend(t.task_id, t.next_pos):
|
||||||
|
valid.append(t)
|
||||||
|
else:
|
||||||
|
t.status = TaskStatus.ABORTED
|
||||||
|
if not valid:
|
||||||
|
break
|
||||||
|
|
||||||
|
step_out = self._executor.execute_decode(
|
||||||
|
valid, return_logprobs=return_logprobs
|
||||||
|
)
|
||||||
|
if return_logprobs:
|
||||||
|
for t, (ntok, _lp) in zip(valid, step_out):
|
||||||
|
t.output_ids.append(ntok)
|
||||||
|
t.output_tokens += 1
|
||||||
|
else:
|
||||||
|
for t, ntok in zip(valid, step_out):
|
||||||
|
t.output_ids.append(ntok)
|
||||||
|
t.output_tokens += 1
|
||||||
|
|
||||||
|
live = [t for t in valid if not t.is_finished(stop_ids)]
|
||||||
|
finally:
|
||||||
|
for t in tasks:
|
||||||
|
if t is not None:
|
||||||
|
cache.task_free(t.task_id)
|
||||||
|
|
||||||
|
results: List[Any] = []
|
||||||
|
for t in tasks:
|
||||||
|
if t is None:
|
||||||
|
results.append(([], []) if return_logprobs else [])
|
||||||
|
elif return_logprobs:
|
||||||
|
results.append((list(t.output_ids), list(t.output_logprobs)))
|
||||||
|
else:
|
||||||
|
results.append(list(t.output_ids))
|
||||||
|
return results
|
||||||
|
|||||||
@@ -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__)
|
||||||
@@ -13,6 +15,33 @@ logger = logging.getLogger(__name__)
|
|||||||
STOP = object()
|
STOP = object()
|
||||||
|
|
||||||
|
|
||||||
|
class StreamDecoder:
|
||||||
|
"""Incremental decoder backed by the tokenizers library's DecodeStream.
|
||||||
|
|
||||||
|
Delegates to the Rust-native streaming decoder which maintains an
|
||||||
|
O(1) bounded token buffer internally (via prefix drain), avoiding
|
||||||
|
the O(n²) cost of re-decoding the full history on each step.
|
||||||
|
|
||||||
|
Multi-byte UTF-8 sequences split across token boundaries are
|
||||||
|
buffered until complete; ``push`` returns "" while the trailing
|
||||||
|
sequence is still incomplete.
|
||||||
|
"""
|
||||||
|
|
||||||
|
__slots__ = ("_stream", "_tok")
|
||||||
|
|
||||||
|
def __init__(self, tokenizer: AutoTokenizer):
|
||||||
|
self._tok = tokenizer._tokenizer
|
||||||
|
self._stream = DecodeStream(skip_special_tokens=True)
|
||||||
|
|
||||||
|
def push(self, token_id: int) -> str:
|
||||||
|
"""Append a token ID and return newly completed text.
|
||||||
|
|
||||||
|
Returns "" while a multi-byte character is still incomplete.
|
||||||
|
"""
|
||||||
|
chunk = self._stream.step(self._tok, token_id)
|
||||||
|
return chunk or ""
|
||||||
|
|
||||||
|
|
||||||
class TaskStatus(Enum):
|
class TaskStatus(Enum):
|
||||||
"""Task lifecycle states."""
|
"""Task lifecycle states."""
|
||||||
|
|
||||||
@@ -33,6 +62,8 @@ class Task:
|
|||||||
temperature: float = 1.0,
|
temperature: float = 1.0,
|
||||||
top_p: float = 1.0,
|
top_p: float = 1.0,
|
||||||
top_k: int = 50,
|
top_k: int = 50,
|
||||||
|
frequency_penalty: float = 0.0,
|
||||||
|
rep_window: int = 64,
|
||||||
):
|
):
|
||||||
self.task_id = task_id
|
self.task_id = task_id
|
||||||
self.prompt_ids = prompt_ids
|
self.prompt_ids = prompt_ids
|
||||||
@@ -40,13 +71,37 @@ class Task:
|
|||||||
self.temperature = temperature
|
self.temperature = temperature
|
||||||
self.top_p = top_p
|
self.top_p = top_p
|
||||||
self.top_k = top_k
|
self.top_k = top_k
|
||||||
|
self.frequency_penalty = frequency_penalty
|
||||||
|
self.rep_window = rep_window
|
||||||
|
|
||||||
self.status = TaskStatus.PENDING
|
self.status = TaskStatus.PENDING
|
||||||
self.output_ids: List[int] = []
|
self.output_ids: List[int] = []
|
||||||
|
self.output_logprobs: List[float] = []
|
||||||
self.input_tokens: int = 0
|
self.input_tokens: int = 0
|
||||||
self.output_tokens: int = 0
|
self.output_tokens: int = 0
|
||||||
self.arrival_time = time.time()
|
self.arrival_time = time.time()
|
||||||
self.finish_time: Optional[float] = None
|
self.finish_time: Optional[float] = None
|
||||||
|
self._decoder: Optional[StreamDecoder] = None
|
||||||
|
|
||||||
|
def decode_new_token(self, tokenizer: AutoTokenizer) -> str:
|
||||||
|
"""Decode the last appended output token, buffering incomplete
|
||||||
|
multi-byte sequences across calls.
|
||||||
|
|
||||||
|
Lazily creates a :class:`StreamDecoder` on first use.
|
||||||
|
"""
|
||||||
|
if self._decoder is None:
|
||||||
|
self._decoder = StreamDecoder(tokenizer)
|
||||||
|
return self._decoder.push(self.output_ids[-1])
|
||||||
|
|
||||||
|
def flush_remaining(self, tokenizer: AutoTokenizer) -> str:
|
||||||
|
"""Emit any text still buffered in the decoder.
|
||||||
|
|
||||||
|
With the Rust-native DecodeStream, the stream is always in a
|
||||||
|
correct state — any completed text was already emitted by the
|
||||||
|
last ``push``. A trailing incomplete multi-byte sequence has no
|
||||||
|
valid text to emit, so this is a no-op.
|
||||||
|
"""
|
||||||
|
return ""
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def next_pos(self) -> int:
|
def next_pos(self) -> int:
|
||||||
@@ -68,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] = []
|
||||||
@@ -92,14 +145,16 @@ class TaskManager:
|
|||||||
temperature: float = 1.0,
|
temperature: float = 1.0,
|
||||||
top_p: float = 1.0,
|
top_p: float = 1.0,
|
||||||
top_k: int = 50,
|
top_k: int = 50,
|
||||||
|
frequency_penalty: float = 0.0,
|
||||||
|
rep_window: int = 64,
|
||||||
stream_callback: Optional[Callable[[str], None]] = None,
|
stream_callback: Optional[Callable[[str], None]] = None,
|
||||||
) -> 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
|
||||||
@@ -116,6 +171,8 @@ class TaskManager:
|
|||||||
temperature=temperature,
|
temperature=temperature,
|
||||||
top_p=top_p,
|
top_p=top_p,
|
||||||
top_k=top_k,
|
top_k=top_k,
|
||||||
|
frequency_penalty=frequency_penalty,
|
||||||
|
rep_window=rep_window,
|
||||||
)
|
)
|
||||||
|
|
||||||
with self._lock:
|
with self._lock:
|
||||||
|
|||||||
+67
-12
@@ -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
|
||||||
@@ -74,20 +74,31 @@ class GenerationRequest:
|
|||||||
top_p: float = 1.0,
|
top_p: float = 1.0,
|
||||||
temperature: float = 1.0,
|
temperature: float = 1.0,
|
||||||
max_tokens: Optional[int] = None,
|
max_tokens: Optional[int] = None,
|
||||||
|
frequency_penalty: float = 0.0,
|
||||||
|
rep_window: int = 64,
|
||||||
stream: bool = False,
|
stream: bool = False,
|
||||||
):
|
):
|
||||||
if not (isinstance(top_k, int) and top_k >= 0):
|
if not (isinstance(top_k, int) and top_k >= 0):
|
||||||
raise ValueError("top_k must be a non-negative integer")
|
raise ValueError("top_k must be a non-negative integer")
|
||||||
if not (0.0 <= top_p <= 1.0):
|
if not (0.0 <= top_p <= 1.0):
|
||||||
raise ValueError("top_p must be a float between 0.0 and 1.0")
|
raise ValueError("top_p must be a float between 0.0 and 1.0")
|
||||||
if not (isinstance(temperature, (int, float)) and temperature > 0):
|
if not (isinstance(temperature, (int, float)) and temperature >= 0):
|
||||||
raise ValueError("temperature must be a positive number")
|
raise ValueError("temperature must be a non-negative number")
|
||||||
|
if not (
|
||||||
|
isinstance(frequency_penalty, (int, float))
|
||||||
|
and -2.0 <= frequency_penalty <= 2.0
|
||||||
|
):
|
||||||
|
raise ValueError("frequency_penalty must be between -2.0 and 2.0")
|
||||||
|
if not (isinstance(rep_window, int) and rep_window > 0):
|
||||||
|
raise ValueError("rep_window must be a positive integer")
|
||||||
|
|
||||||
self.messages = messages
|
self.messages = messages
|
||||||
self.top_k = top_k
|
self.top_k = top_k
|
||||||
self.top_p = top_p
|
self.top_p = top_p
|
||||||
self.temperature = temperature
|
self.temperature = temperature
|
||||||
self.max_tokens = max_tokens
|
self.max_tokens = max_tokens
|
||||||
|
self.frequency_penalty = frequency_penalty
|
||||||
|
self.rep_window = rep_window
|
||||||
self.stream = stream
|
self.stream = stream
|
||||||
|
|
||||||
|
|
||||||
@@ -100,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
|
||||||
@@ -111,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,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -132,17 +140,33 @@ class InferenceEngine:
|
|||||||
temperature: float = 1.0,
|
temperature: float = 1.0,
|
||||||
top_p: float = 1.0,
|
top_p: float = 1.0,
|
||||||
top_k: int = 50,
|
top_k: int = 50,
|
||||||
|
frequency_penalty: float = 0.0,
|
||||||
|
rep_window: int = 64,
|
||||||
) -> Union[Generator, str, List[str]]:
|
) -> Union[Generator, str, List[str]]:
|
||||||
is_batch = isinstance(prompt, list)
|
is_batch = isinstance(prompt, list)
|
||||||
prompts = prompt if is_batch else [prompt]
|
prompts = prompt if is_batch else [prompt]
|
||||||
|
|
||||||
if stream:
|
if stream:
|
||||||
return self._generate_streaming(
|
return self._generate_streaming(
|
||||||
prompts, is_batch, max_tokens, temperature, top_p, top_k
|
prompts,
|
||||||
|
is_batch,
|
||||||
|
max_tokens,
|
||||||
|
temperature,
|
||||||
|
top_p,
|
||||||
|
top_k,
|
||||||
|
frequency_penalty,
|
||||||
|
rep_window,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
return self._generate_non_streaming(
|
return self._generate_non_streaming(
|
||||||
prompts, is_batch, max_tokens, temperature, top_p, top_k
|
prompts,
|
||||||
|
is_batch,
|
||||||
|
max_tokens,
|
||||||
|
temperature,
|
||||||
|
top_p,
|
||||||
|
top_k,
|
||||||
|
frequency_penalty,
|
||||||
|
rep_window,
|
||||||
)
|
)
|
||||||
|
|
||||||
def generate_async(
|
def generate_async(
|
||||||
@@ -152,9 +176,18 @@ class InferenceEngine:
|
|||||||
temperature: float = 1.0,
|
temperature: float = 1.0,
|
||||||
top_p: float = 1.0,
|
top_p: float = 1.0,
|
||||||
top_k: int = 50,
|
top_k: int = 50,
|
||||||
|
frequency_penalty: float = 0.0,
|
||||||
|
rep_window: int = 64,
|
||||||
) -> AsyncGenerator[str, None]:
|
) -> AsyncGenerator[str, None]:
|
||||||
sync_gen = self._generate_streaming(
|
sync_gen = self._generate_streaming(
|
||||||
[prompt], False, max_tokens, temperature, top_p, top_k
|
[prompt],
|
||||||
|
False,
|
||||||
|
max_tokens,
|
||||||
|
temperature,
|
||||||
|
top_p,
|
||||||
|
top_k,
|
||||||
|
frequency_penalty,
|
||||||
|
rep_window,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def _agen():
|
async def _agen():
|
||||||
@@ -185,6 +218,8 @@ class InferenceEngine:
|
|||||||
temperature=request.temperature,
|
temperature=request.temperature,
|
||||||
top_p=request.top_p,
|
top_p=request.top_p,
|
||||||
top_k=request.top_k,
|
top_k=request.top_k,
|
||||||
|
frequency_penalty=request.frequency_penalty,
|
||||||
|
rep_window=request.rep_window,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _submit_tasks(
|
def _submit_tasks(
|
||||||
@@ -194,6 +229,8 @@ class InferenceEngine:
|
|||||||
temperature: float,
|
temperature: float,
|
||||||
top_p: float,
|
top_p: float,
|
||||||
top_k: int,
|
top_k: int,
|
||||||
|
frequency_penalty: float,
|
||||||
|
rep_window: int,
|
||||||
) -> Tuple[GenerateResult, List[str]]:
|
) -> Tuple[GenerateResult, List[str]]:
|
||||||
n = len(prompts)
|
n = len(prompts)
|
||||||
result = GenerateResult(count=n)
|
result = GenerateResult(count=n)
|
||||||
@@ -206,6 +243,8 @@ class InferenceEngine:
|
|||||||
temperature=temperature,
|
temperature=temperature,
|
||||||
top_p=top_p,
|
top_p=top_p,
|
||||||
top_k=top_k,
|
top_k=top_k,
|
||||||
|
frequency_penalty=frequency_penalty,
|
||||||
|
rep_window=rep_window,
|
||||||
stream_callback=cb,
|
stream_callback=cb,
|
||||||
)
|
)
|
||||||
task_ids.append(task_id)
|
task_ids.append(task_id)
|
||||||
@@ -226,9 +265,17 @@ class InferenceEngine:
|
|||||||
temperature: float,
|
temperature: float,
|
||||||
top_p: float,
|
top_p: float,
|
||||||
top_k: int,
|
top_k: int,
|
||||||
|
frequency_penalty: float,
|
||||||
|
rep_window: int,
|
||||||
) -> Generator:
|
) -> Generator:
|
||||||
result, task_ids = self._submit_tasks(
|
result, task_ids = self._submit_tasks(
|
||||||
prompts, max_tokens, temperature, top_p, top_k
|
prompts,
|
||||||
|
max_tokens,
|
||||||
|
temperature,
|
||||||
|
top_p,
|
||||||
|
top_k,
|
||||||
|
frequency_penalty,
|
||||||
|
rep_window,
|
||||||
)
|
)
|
||||||
n = len(prompts)
|
n = len(prompts)
|
||||||
remaining = n
|
remaining = n
|
||||||
@@ -262,9 +309,17 @@ class InferenceEngine:
|
|||||||
temperature: float,
|
temperature: float,
|
||||||
top_p: float,
|
top_p: float,
|
||||||
top_k: int,
|
top_k: int,
|
||||||
|
frequency_penalty: float,
|
||||||
|
rep_window: int,
|
||||||
) -> Union[str, List[str]]:
|
) -> Union[str, List[str]]:
|
||||||
result, task_ids = self._submit_tasks(
|
result, task_ids = self._submit_tasks(
|
||||||
prompts, max_tokens, temperature, top_p, top_k
|
prompts,
|
||||||
|
max_tokens,
|
||||||
|
temperature,
|
||||||
|
top_p,
|
||||||
|
top_k,
|
||||||
|
frequency_penalty,
|
||||||
|
rep_window,
|
||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|||||||
+236
-25
@@ -1,15 +1,15 @@
|
|||||||
"""Composable sampling strategies for logit transformation.
|
"""Composable sampling strategies for logit transformation.
|
||||||
|
|
||||||
Implements the Strategy pattern: each sampling technique
|
Implements the Strategy pattern: each sampling technique
|
||||||
(temperature, top-k, top-p) is a pluggable strategy that
|
(temperature, top-k, top-p, frequency penalty) is a pluggable
|
||||||
can be composed into a pipeline.
|
strategy that can be composed into a pipeline.
|
||||||
|
|
||||||
All strategies accept both scalar and per-sample tensor
|
All strategies accept both scalar and per-sample tensor
|
||||||
parameters, so a single pipeline works for any batch size.
|
parameters, so a single pipeline works for any batch size.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from typing import List, Union
|
from typing import List, Optional, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
@@ -19,12 +19,23 @@ class BaseSamplingStrategy(ABC):
|
|||||||
"""Abstract base for a logit transformation strategy."""
|
"""Abstract base for a logit transformation strategy."""
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def apply(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor:
|
def apply(
|
||||||
|
self,
|
||||||
|
logits: Tensor,
|
||||||
|
filter_value: float = -float("inf"),
|
||||||
|
input_ids: Optional[Tensor] = None,
|
||||||
|
input_mask: Optional[Tensor] = None,
|
||||||
|
) -> Tensor:
|
||||||
"""Applies the strategy to logits.
|
"""Applies the strategy to logits.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
logits: Raw logits tensor (batch, vocab_size).
|
logits: Raw logits tensor (batch, vocab_size).
|
||||||
filter_value: Value assigned to filtered-out positions.
|
filter_value: Value assigned to filtered-out positions.
|
||||||
|
input_ids: Previously generated token IDs ``[batch, seq_len]``,
|
||||||
|
padded with 0. Used by frequency penalty.
|
||||||
|
input_mask: Boolean mask ``[batch, seq_len]``, True for real
|
||||||
|
tokens, False for padding. Used to exclude padding from
|
||||||
|
penalty computation.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Transformed logits tensor.
|
Transformed logits tensor.
|
||||||
@@ -42,7 +53,13 @@ class TemperatureStrategy(BaseSamplingStrategy):
|
|||||||
def __init__(self, temperature: Union[float, Tensor] = 1.0):
|
def __init__(self, temperature: Union[float, Tensor] = 1.0):
|
||||||
self.temperature = temperature
|
self.temperature = temperature
|
||||||
|
|
||||||
def apply(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor:
|
def apply(
|
||||||
|
self,
|
||||||
|
logits: Tensor,
|
||||||
|
filter_value: float = -float("inf"),
|
||||||
|
input_ids: Optional[Tensor] = None,
|
||||||
|
input_mask: Optional[Tensor] = None,
|
||||||
|
) -> Tensor:
|
||||||
t = self.temperature
|
t = self.temperature
|
||||||
if isinstance(t, Tensor):
|
if isinstance(t, Tensor):
|
||||||
t = t.to(logits.device, non_blocking=True).view(-1, 1)
|
t = t.to(logits.device, non_blocking=True).view(-1, 1)
|
||||||
@@ -64,7 +81,13 @@ class TopKStrategy(BaseSamplingStrategy):
|
|||||||
def __init__(self, top_k: Union[int, Tensor] = 0):
|
def __init__(self, top_k: Union[int, Tensor] = 0):
|
||||||
self.top_k = top_k
|
self.top_k = top_k
|
||||||
|
|
||||||
def apply(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor:
|
def apply(
|
||||||
|
self,
|
||||||
|
logits: Tensor,
|
||||||
|
filter_value: float = -float("inf"),
|
||||||
|
input_ids: Optional[Tensor] = None,
|
||||||
|
input_mask: Optional[Tensor] = None,
|
||||||
|
) -> Tensor:
|
||||||
tk = self.top_k
|
tk = self.top_k
|
||||||
if isinstance(tk, Tensor):
|
if isinstance(tk, Tensor):
|
||||||
tk = tk.to(logits.device, non_blocking=True).long().clamp(min=0)
|
tk = tk.to(logits.device, non_blocking=True).long().clamp(min=0)
|
||||||
@@ -114,7 +137,13 @@ class TopPStrategy(BaseSamplingStrategy):
|
|||||||
logits[mask] = filter_value
|
logits[mask] = filter_value
|
||||||
return logits
|
return logits
|
||||||
|
|
||||||
def apply(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor:
|
def apply(
|
||||||
|
self,
|
||||||
|
logits: Tensor,
|
||||||
|
filter_value: float = -float("inf"),
|
||||||
|
input_ids: Optional[Tensor] = None,
|
||||||
|
input_mask: Optional[Tensor] = None,
|
||||||
|
) -> Tensor:
|
||||||
tp = self.top_p
|
tp = self.top_p
|
||||||
if isinstance(tp, Tensor):
|
if isinstance(tp, Tensor):
|
||||||
tp = tp.to(logits.device, non_blocking=True)
|
tp = tp.to(logits.device, non_blocking=True)
|
||||||
@@ -125,6 +154,84 @@ class TopPStrategy(BaseSamplingStrategy):
|
|||||||
return logits
|
return logits
|
||||||
|
|
||||||
|
|
||||||
|
class FrequencyPenaltyStrategy(BaseSamplingStrategy):
|
||||||
|
"""Penalizes tokens based on how many times they appeared in history.
|
||||||
|
|
||||||
|
Subtracts ``penalty * count(token)`` from each token's logit, where
|
||||||
|
``count(token)`` is the number of occurrences in the generation history
|
||||||
|
(prompt + output). A penalty of ``0.0`` disables the strategy.
|
||||||
|
|
||||||
|
Unlike repetition penalty (which only checks *presence*), frequency
|
||||||
|
penalty scales linearly with occurrence count: the first use is
|
||||||
|
penalized once, the third use three times. This allows natural
|
||||||
|
repetition of common words while suppressing degenerate loops.
|
||||||
|
|
||||||
|
Reference: OpenAI API ``frequency_penalty`` parameter.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
penalty: Scalar or ``[batch]`` tensor (0.0 disables, range -2.0~2.0).
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, penalty: Union[float, Tensor] = 0.0):
|
||||||
|
self.penalty = penalty
|
||||||
|
|
||||||
|
def apply(
|
||||||
|
self,
|
||||||
|
logits: Tensor,
|
||||||
|
filter_value: float = -float("inf"),
|
||||||
|
input_ids: Optional[Tensor] = None,
|
||||||
|
input_mask: Optional[Tensor] = None,
|
||||||
|
) -> Tensor:
|
||||||
|
if input_ids is None:
|
||||||
|
return logits
|
||||||
|
|
||||||
|
p = self.penalty
|
||||||
|
if isinstance(p, Tensor):
|
||||||
|
p = p.to(logits.device, non_blocking=True).view(-1, 1)
|
||||||
|
if (p == 0.0).all():
|
||||||
|
return logits
|
||||||
|
elif p == 0.0:
|
||||||
|
return logits
|
||||||
|
|
||||||
|
input_ids = input_ids.to(logits.device, non_blocking=True)
|
||||||
|
|
||||||
|
if input_mask is not None:
|
||||||
|
input_mask = input_mask.to(logits.device, non_blocking=True)
|
||||||
|
masked_ids = input_ids.clone()
|
||||||
|
masked_ids[~input_mask] = -1
|
||||||
|
else:
|
||||||
|
masked_ids = input_ids
|
||||||
|
|
||||||
|
batch_sz, seq_len = masked_ids.shape
|
||||||
|
vocab_size = logits.size(-1)
|
||||||
|
|
||||||
|
if isinstance(p, Tensor):
|
||||||
|
penalty_per_row = p.expand(batch_sz, 1)
|
||||||
|
else:
|
||||||
|
penalty_per_row = torch.full(
|
||||||
|
(batch_sz, 1), float(p), device=logits.device, dtype=logits.dtype
|
||||||
|
)
|
||||||
|
|
||||||
|
counts = torch.zeros(
|
||||||
|
batch_sz, vocab_size, device=logits.device, dtype=logits.dtype
|
||||||
|
)
|
||||||
|
valid_mask = masked_ids >= 0
|
||||||
|
if valid_mask.any():
|
||||||
|
valid_ids = masked_ids[valid_mask]
|
||||||
|
row_indices = (
|
||||||
|
torch.arange(batch_sz, device=logits.device)
|
||||||
|
.unsqueeze(1)
|
||||||
|
.expand_as(masked_ids)[valid_mask]
|
||||||
|
)
|
||||||
|
counts.index_put_(
|
||||||
|
(row_indices, valid_ids),
|
||||||
|
torch.ones_like(valid_ids, dtype=logits.dtype),
|
||||||
|
accumulate=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
return logits - penalty_per_row * counts
|
||||||
|
|
||||||
|
|
||||||
class SamplingPipeline(BaseSamplingStrategy):
|
class SamplingPipeline(BaseSamplingStrategy):
|
||||||
"""Composes multiple sampling strategies into a single transformation.
|
"""Composes multiple sampling strategies into a single transformation.
|
||||||
|
|
||||||
@@ -145,25 +252,76 @@ class SamplingPipeline(BaseSamplingStrategy):
|
|||||||
def __init__(self, strategies: List[BaseSamplingStrategy]):
|
def __init__(self, strategies: List[BaseSamplingStrategy]):
|
||||||
self.strategies = strategies
|
self.strategies = strategies
|
||||||
|
|
||||||
def apply(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor:
|
def apply(
|
||||||
|
self,
|
||||||
|
logits: Tensor,
|
||||||
|
filter_value: float = -float("inf"),
|
||||||
|
input_ids: Optional[Tensor] = None,
|
||||||
|
input_mask: Optional[Tensor] = None,
|
||||||
|
) -> Tensor:
|
||||||
for strategy in self.strategies:
|
for strategy in self.strategies:
|
||||||
logits = strategy.apply(logits, filter_value)
|
logits = strategy.apply(logits, filter_value, input_ids, input_mask)
|
||||||
return logits
|
return logits
|
||||||
|
|
||||||
@torch.no_grad()
|
@staticmethod
|
||||||
def sample(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor:
|
def _is_greedy(temperature: Union[float, Tensor]) -> bool:
|
||||||
|
if isinstance(temperature, Tensor):
|
||||||
|
return temperature.numel() == 1 and temperature.item() == 0
|
||||||
|
return temperature == 0
|
||||||
|
|
||||||
|
@torch.inference_mode()
|
||||||
|
def sample(
|
||||||
|
self,
|
||||||
|
logits: Tensor,
|
||||||
|
filter_value: float = -float("inf"),
|
||||||
|
input_ids: Optional[Tensor] = None,
|
||||||
|
input_mask: Optional[Tensor] = None,
|
||||||
|
return_logprobs: bool = False,
|
||||||
|
):
|
||||||
"""Apply strategies then sample (softmax + multinomial).
|
"""Apply strategies then sample (softmax + multinomial).
|
||||||
|
|
||||||
|
Short-circuits to ``argmax`` when temperature is exactly 0
|
||||||
|
(deterministic / greedy decode).
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
logits: Raw logits ``[batch, vocab_size]``.
|
logits: Raw logits ``[batch, vocab_size]``.
|
||||||
|
input_ids: Previously generated token IDs ``[batch, seq_len]``.
|
||||||
|
input_mask: Boolean mask for ``input_ids`` padding.
|
||||||
|
return_logprobs: If ``True``, return ``(tokens, logprobs)``
|
||||||
|
where ``logprobs[i]`` is the log-probability of
|
||||||
|
``tokens[i]`` under the (post-strategy) sampling
|
||||||
|
distribution.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Sampled token IDs ``[batch]``.
|
Sampled token IDs ``[batch]``, or — when ``return_logprobs``
|
||||||
|
is ``True`` — a ``(token_ids, chosen_logprobs)`` tuple.
|
||||||
"""
|
"""
|
||||||
return torch.multinomial(
|
if self._is_greedy_pipeline():
|
||||||
torch.softmax(self.apply(logits, filter_value), dim=-1),
|
tokens = logits.argmax(dim=-1)
|
||||||
num_samples=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
|
||||||
|
|
||||||
|
transformed = self.apply(logits, filter_value, input_ids, input_mask)
|
||||||
|
log_probs = torch.log_softmax(transformed.float(), dim=-1)
|
||||||
|
tokens = torch.multinomial(
|
||||||
|
torch.softmax(transformed, dim=-1), num_samples=1
|
||||||
).squeeze(-1)
|
).squeeze(-1)
|
||||||
|
if not return_logprobs:
|
||||||
|
return tokens
|
||||||
|
chosen = torch.gather(log_probs, -1, tokens.unsqueeze(-1)).squeeze(-1)
|
||||||
|
return tokens, chosen
|
||||||
|
|
||||||
|
def _is_greedy_pipeline(self) -> bool:
|
||||||
|
"""True if the first strategy is greedy temperature (temp=0)."""
|
||||||
|
if not self.strategies:
|
||||||
|
return False
|
||||||
|
first = self.strategies[0]
|
||||||
|
return isinstance(first, TemperatureStrategy) and self._is_greedy(
|
||||||
|
first.temperature
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@torch.inference_mode()
|
@torch.inference_mode()
|
||||||
@@ -172,22 +330,75 @@ def sample(
|
|||||||
temperature: Union[float, Tensor] = 1.0,
|
temperature: Union[float, Tensor] = 1.0,
|
||||||
top_k: Union[int, Tensor] = 0,
|
top_k: Union[int, Tensor] = 0,
|
||||||
top_p: Union[float, Tensor] = 1.0,
|
top_p: Union[float, Tensor] = 1.0,
|
||||||
|
frequency_penalty: Union[float, Tensor] = 0.0,
|
||||||
|
input_ids: Optional[Tensor] = None,
|
||||||
|
input_mask: Optional[Tensor] = None,
|
||||||
filter_value: float = -float("inf"),
|
filter_value: float = -float("inf"),
|
||||||
) -> Tensor:
|
return_logprobs: bool = False,
|
||||||
|
):
|
||||||
"""Apply sampling strategies then sample (softmax + multinomial).
|
"""Apply sampling strategies then sample (softmax + multinomial).
|
||||||
|
|
||||||
Shortcut for ``SamplingPipeline(...).sample(logits)``.
|
Shortcut for ``SamplingPipeline(...).sample(logits, return_logprobs=)``.
|
||||||
|
|
||||||
|
When **temperature** is exactly 0 (scalar or single-element tensor)
|
||||||
|
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
|
||||||
|
(0.0 disables, range -2.0~2.0).
|
||||||
|
input_ids: Previously generated token IDs ``[batch, seq_len]``.
|
||||||
|
input_mask: Boolean mask for ``input_ids`` padding.
|
||||||
|
return_logprobs: If ``True``, also return the log-probability
|
||||||
|
of each sampled token under the (post-strategy) sampling
|
||||||
|
distribution — useful for RL rollout (PPO/GRPO importance
|
||||||
|
ratios).
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Sampled token IDs ``[batch]``.
|
Sampled token IDs ``[batch]``, or — when ``return_logprobs`` is
|
||||||
|
``True`` — a ``(token_ids, chosen_logprobs)`` tuple where
|
||||||
|
``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
|
||||||
]
|
)
|
||||||
).sample(logits, filter_value)
|
if isinstance(temperature, Tensor)
|
||||||
|
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,
|
||||||
|
filter_value=filter_value,
|
||||||
|
input_ids=input_ids,
|
||||||
|
input_mask=input_mask,
|
||||||
|
return_logprobs=return_logprobs,
|
||||||
|
)
|
||||||
|
|||||||
@@ -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,10 +65,9 @@ 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,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
is_causal = attn_mask is None
|
|
||||||
|
|
||||||
q = self._split_heads(self.q_proj(x), self.n_heads)
|
q = self._split_heads(self.q_proj(x), self.n_heads)
|
||||||
k = self._split_heads(self.k_proj(x), self.n_kv_heads)
|
k = self._split_heads(self.k_proj(x), self.n_kv_heads)
|
||||||
v = self._split_heads(self.v_proj(x), self.n_kv_heads)
|
v = self._split_heads(self.v_proj(x), self.n_kv_heads)
|
||||||
@@ -87,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))
|
||||||
@@ -162,10 +139,10 @@ 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,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
bsz, seq_len, _ = x.size()
|
bsz, seq_len, _ = x.size()
|
||||||
is_causal = attn_mask is None
|
|
||||||
|
|
||||||
q = self.q_proj(x)
|
q = self.q_proj(x)
|
||||||
q = q.view(bsz, seq_len, self.n_heads, self.head_dim)
|
q = q.view(bsz, seq_len, self.n_heads, self.head_dim)
|
||||||
@@ -194,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
|
||||||
@@ -14,10 +14,18 @@ class DecoderBlock(nn.Module):
|
|||||||
def __init__(self, config, layer_id: int):
|
def __init__(self, config, layer_id: int):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
cfg = asdict(config)
|
cfg = asdict(config)
|
||||||
cfg["down_init_std"] = 0.02 / (2 * config.n_layers) ** 0.5
|
cfg.update(
|
||||||
|
dim=config.hidden_size,
|
||||||
|
dim_ffn=config.intermediate_size,
|
||||||
|
n_layers=config.num_hidden_layers,
|
||||||
|
n_heads=config.num_attention_heads,
|
||||||
|
n_kv_heads=config.num_key_value_heads,
|
||||||
|
norm_eps=config.rms_norm_eps,
|
||||||
|
down_init_std=0.02 / (2 * config.num_hidden_layers) ** 0.5,
|
||||||
|
)
|
||||||
self.attention = AttnFactory.create(config.attn_type, **cfg, layer_id=layer_id)
|
self.attention = AttnFactory.create(config.attn_type, **cfg, layer_id=layer_id)
|
||||||
self.input_norm = RMSNorm(config.dim, config.norm_eps)
|
self.input_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
||||||
self.post_attention_norm = RMSNorm(config.dim, config.norm_eps)
|
self.post_attention_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
||||||
self.mlp = FFNFactory.create(config.ffn_type, **cfg)
|
self.mlp = FFNFactory.create(config.ffn_type, **cfg)
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
@@ -25,13 +33,15 @@ 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,
|
||||||
) -> 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,
|
||||||
)
|
)
|
||||||
x = attn_output + x
|
x = attn_output + x
|
||||||
x = self.mlp(self.post_attention_norm(x)) + x
|
x = self.mlp(self.post_attention_norm(x)) + 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 (
|
||||||
@@ -39,8 +40,12 @@ class LoRALinear(nn.Module):
|
|||||||
|
|
||||||
self.r = r
|
self.r = r
|
||||||
self.scaling = alpha / r
|
self.scaling = alpha / r
|
||||||
self.lora_A = nn.Parameter(torch.randn(r, self.weight.shape[1]) / r)
|
device = self.weight.device
|
||||||
self.lora_B = nn.Parameter(torch.zeros(self.weight.shape[0], r))
|
dtype = self.weight.dtype
|
||||||
|
lora_a = torch.randn(r, self.weight.shape[1], device=device, dtype=dtype) / r
|
||||||
|
lora_b = torch.zeros(self.weight.shape[0], r, device=device, dtype=dtype)
|
||||||
|
self.lora_A = nn.Parameter(lora_a)
|
||||||
|
self.lora_B = nn.Parameter(lora_b)
|
||||||
self._merged = False
|
self._merged = False
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
|
|||||||
@@ -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)
|
|
||||||
|
|||||||
+17
-9
@@ -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,25 +13,33 @@ 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)
|
||||||
self.config = config
|
self.config = config
|
||||||
rope_dim = config.dim // config.n_heads
|
rope_dim = config.hidden_size // config.num_attention_heads
|
||||||
rope_base = config.rope_theta if config.rope_theta is not None else 10000
|
rope_base = config.rope_theta if config.rope_theta is not None else 10000
|
||||||
self.rotary_embedding = RotaryEmbedding(
|
self.rotary_embedding = RotaryEmbedding(
|
||||||
rope_dim, config.max_len, rope_base, rope_scaling=config.rope_scaling
|
rope_dim,
|
||||||
|
config.max_position_embeddings,
|
||||||
|
rope_base,
|
||||||
|
rope_scaling=config.rope_scaling,
|
||||||
)
|
)
|
||||||
self.embed_tokens = Embedding(
|
self.embed_tokens = Embedding(
|
||||||
config.vocab_size, config.dim, neftune_alpha=config.neftune_alpha
|
config.vocab_size,
|
||||||
|
config.hidden_size,
|
||||||
|
neftune_alpha=config.neftune_alpha,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.layers = nn.ModuleList(
|
self.layers = nn.ModuleList(
|
||||||
[DecoderBlock(config, layer_id) for layer_id in range(config.n_layers)]
|
[
|
||||||
|
DecoderBlock(config, layer_id)
|
||||||
|
for layer_id in range(config.num_hidden_layers)
|
||||||
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
self.norm = RMSNorm(config.dim, config.norm_eps)
|
self.norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
||||||
|
|
||||||
self.pooling_type = config.pooling_type or "mean"
|
self.pooling_type = config.pooling_type or "mean"
|
||||||
self.normalize_embeddings = config.normalize_embeddings or False
|
self.normalize_embeddings = config.normalize_embeddings or False
|
||||||
@@ -59,10 +67,10 @@ class EmbeddingEncoder(AutoModel):
|
|||||||
x = self.embed_tokens(input_ids)
|
x = self.embed_tokens(input_ids)
|
||||||
|
|
||||||
rotary_emb = self.rotary_embedding(x, position_ids)
|
rotary_emb = self.rotary_embedding(x, position_ids)
|
||||||
attn_mask = process_attention_mask(x, position_ids, input_mask, is_causal=False)
|
attn_mask = process_attention_mask(input_mask)
|
||||||
|
|
||||||
for layer in self.layers:
|
for layer in self.layers:
|
||||||
x = layer(x, rotary_emb, attn_mask, paged_cache=None)
|
x = layer(x, rotary_emb, attn_mask)
|
||||||
|
|
||||||
hidden_states = self.norm(x)
|
hidden_states = self.norm(x)
|
||||||
|
|
||||||
|
|||||||
+31
-39
@@ -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
|
||||||
@@ -15,35 +15,18 @@ from astrai.model.components.rope import RotaryEmbedding
|
|||||||
|
|
||||||
|
|
||||||
def process_attention_mask(
|
def process_attention_mask(
|
||||||
input_tensor: Tensor,
|
input_mask: Optional[Tensor],
|
||||||
position_ids: Optional[Tensor],
|
|
||||||
input_mask: Optional[Tensor] = None,
|
|
||||||
is_causal: bool = False,
|
|
||||||
) -> Optional[Tensor]:
|
) -> Optional[Tensor]:
|
||||||
if position_ids is None:
|
|
||||||
return None
|
|
||||||
if input_mask is not None and input_mask.dim() > 2:
|
|
||||||
return input_mask
|
|
||||||
|
|
||||||
device = input_tensor.device
|
|
||||||
B = input_tensor.size(0)
|
|
||||||
T = position_ids.max().item() + 1
|
|
||||||
|
|
||||||
if input_mask is None:
|
if input_mask is None:
|
||||||
if position_ids.min().item() == 0 and is_causal:
|
return None
|
||||||
return None
|
if input_mask.dim() == 2:
|
||||||
attend = torch.ones(B, 1, T, dtype=torch.bool, device=device)
|
return input_mask[:, None, None, :]
|
||||||
else:
|
if input_mask.dim() == 3:
|
||||||
attend = input_mask[:, :T].to(device=device, dtype=torch.bool).unsqueeze(1)
|
return input_mask[:, None, :, :]
|
||||||
|
return input_mask
|
||||||
if is_causal:
|
|
||||||
causal = position_ids.unsqueeze(-1) >= torch.arange(T, device=device)
|
|
||||||
attend = attend & causal
|
|
||||||
|
|
||||||
return attend.unsqueeze(1)
|
|
||||||
|
|
||||||
|
|
||||||
@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."""
|
||||||
|
|
||||||
@@ -53,24 +36,32 @@ class AutoRegressiveLM(AutoModel):
|
|||||||
rope_dim = (
|
rope_dim = (
|
||||||
config.qk_rope_head_dim
|
config.qk_rope_head_dim
|
||||||
if config.attn_type == "mla"
|
if config.attn_type == "mla"
|
||||||
else config.dim // config.n_heads
|
else config.hidden_size // config.num_attention_heads
|
||||||
)
|
)
|
||||||
rope_base = config.rope_theta if config.rope_theta is not None else 10000
|
rope_base = config.rope_theta if config.rope_theta is not None else 10000
|
||||||
self.rotary_embedding = RotaryEmbedding(
|
self.rotary_embedding = RotaryEmbedding(
|
||||||
rope_dim, config.max_len, rope_base, rope_scaling=config.rope_scaling
|
rope_dim,
|
||||||
|
config.max_position_embeddings,
|
||||||
|
rope_base,
|
||||||
|
rope_scaling=config.rope_scaling,
|
||||||
)
|
)
|
||||||
self.embed_tokens = Embedding(
|
self.embed_tokens = Embedding(
|
||||||
config.vocab_size, config.dim, neftune_alpha=config.neftune_alpha
|
config.vocab_size,
|
||||||
|
config.hidden_size,
|
||||||
|
neftune_alpha=config.neftune_alpha,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.layers = nn.ModuleList(
|
self.layers = nn.ModuleList(
|
||||||
[DecoderBlock(config, layer_id) for layer_id in range(config.n_layers)]
|
[
|
||||||
|
DecoderBlock(config, layer_id)
|
||||||
|
for layer_id in range(config.num_hidden_layers)
|
||||||
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
self.norm = RMSNorm(config.dim, config.norm_eps)
|
self.norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
||||||
self.lm_head = Linear(config.dim, config.vocab_size)
|
self.lm_head = Linear(config.hidden_size, config.vocab_size)
|
||||||
|
|
||||||
if self.config.tie_weight is True:
|
if self.config.tie_word_embeddings is True:
|
||||||
self.lm_head.weight = self.embed_tokens.weight
|
self.lm_head.weight = self.embed_tokens.weight
|
||||||
|
|
||||||
self.apply(self._init_weights)
|
self.apply(self._init_weights)
|
||||||
@@ -85,7 +76,7 @@ class AutoRegressiveLM(AutoModel):
|
|||||||
|
|
||||||
state_dict = dict(state_dict)
|
state_dict = dict(state_dict)
|
||||||
|
|
||||||
if self.config.tie_weight is True:
|
if self.config.tie_word_embeddings is True:
|
||||||
# same tensor for embed and lm_head
|
# same tensor for embed and lm_head
|
||||||
if embed_key in state_dict:
|
if embed_key in state_dict:
|
||||||
state_dict[lm_head_key] = state_dict[embed_key]
|
state_dict[lm_head_key] = state_dict[embed_key]
|
||||||
@@ -101,7 +92,7 @@ class AutoRegressiveLM(AutoModel):
|
|||||||
destination=destination, prefix=prefix, keep_vars=keep_vars
|
destination=destination, prefix=prefix, keep_vars=keep_vars
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.config.tie_weight is True:
|
if self.config.tie_word_embeddings is True:
|
||||||
lm_head_key = prefix + "lm_head.weight"
|
lm_head_key = prefix + "lm_head.weight"
|
||||||
if lm_head_key in state_dict:
|
if lm_head_key in state_dict:
|
||||||
del state_dict[lm_head_key]
|
del state_dict[lm_head_key]
|
||||||
@@ -112,17 +103,18 @@ 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
|
||||||
|
|
||||||
x = self.embed_tokens(input_ids)
|
x = self.embed_tokens(input_ids)
|
||||||
rotary_emb = self.rotary_embedding(x, position_ids)
|
rotary_emb = self.rotary_embedding(x, position_ids)
|
||||||
attn_mask = process_attention_mask(x, position_ids, input_mask, is_causal=True)
|
attn_mask = process_attention_mask(input_mask)
|
||||||
|
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)
|
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])
|
||||||
@@ -7,8 +7,9 @@ from astrai.parallel.executor import (
|
|||||||
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,
|
||||||
@@ -25,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",
|
||||||
@@ -35,4 +34,6 @@ __all__ = [
|
|||||||
"NoneExecutor",
|
"NoneExecutor",
|
||||||
"DDPExecutor",
|
"DDPExecutor",
|
||||||
"FSDPExecutor",
|
"FSDPExecutor",
|
||||||
|
"create_ref_model",
|
||||||
|
"broadcast_state_dict",
|
||||||
]
|
]
|
||||||
|
|||||||
+212
-70
@@ -4,16 +4,19 @@ import contextlib
|
|||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from typing import Optional, Tuple
|
from typing import Any, Callable, Dict, Optional, Tuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
import torch.distributed as dist
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from torch.distributed.fsdp import FullStateDictConfig, StateDictType
|
from torch.distributed.fsdp import (
|
||||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
FSDPModule,
|
||||||
|
fully_shard,
|
||||||
|
)
|
||||||
|
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
|
||||||
from torch.optim.lr_scheduler import LRScheduler
|
from torch.optim.lr_scheduler import LRScheduler
|
||||||
from torch.utils.data import DataLoader
|
|
||||||
|
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
from astrai.parallel.setup import get_rank, get_world_size
|
from astrai.parallel.setup import get_rank, get_world_size
|
||||||
@@ -21,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)
|
||||||
@@ -85,19 +164,28 @@ class BaseExecutor:
|
|||||||
|
|
||||||
def prepare(
|
def prepare(
|
||||||
self,
|
self,
|
||||||
model: nn.Module,
|
model_fn: Callable[[], nn.Module],
|
||||||
optimizer: Optional[Optimizer] = None,
|
optimizer_fn: Optional[Callable[[nn.Module], Optimizer]] = None,
|
||||||
dataloader: Optional[DataLoader] = None,
|
scheduler_fn: Optional[Callable[[Optimizer], LRScheduler]] = None,
|
||||||
scheduler: Optional[LRScheduler] = None,
|
before_wrap: Optional[Callable[[nn.Module], nn.Module]] = None,
|
||||||
) -> Tuple[
|
after_wrap: Optional[Callable[[nn.Module], nn.Module]] = None,
|
||||||
nn.Module, Optional[Optimizer], Optional[DataLoader], Optional[LRScheduler]
|
) -> Tuple[nn.Module, Optional[Optimizer], Optional[LRScheduler]]:
|
||||||
]:
|
model = model_fn()
|
||||||
|
if before_wrap is not None:
|
||||||
|
model = before_wrap(model)
|
||||||
model = self._prepare_model(model)
|
model = self._prepare_model(model)
|
||||||
if optimizer is not None:
|
if after_wrap is not None:
|
||||||
|
model = after_wrap(model)
|
||||||
|
optimizer = None
|
||||||
|
scheduler = None
|
||||||
|
if optimizer_fn is not None:
|
||||||
|
optimizer = optimizer_fn(model)
|
||||||
|
if scheduler_fn is not None:
|
||||||
|
scheduler = scheduler_fn(optimizer)
|
||||||
optimizer = AccumOptimizer(optimizer, self.gradient_state)
|
optimizer = AccumOptimizer(optimizer, self.gradient_state)
|
||||||
if scheduler is not None:
|
if scheduler is not None:
|
||||||
scheduler = AccumScheduler(scheduler, self.gradient_state)
|
scheduler = AccumScheduler(scheduler, self.gradient_state)
|
||||||
return model, optimizer, dataloader, scheduler
|
return model, optimizer, scheduler
|
||||||
|
|
||||||
def _prepare_model(self, model: nn.Module) -> nn.Module:
|
def _prepare_model(self, model: nn.Module) -> nn.Module:
|
||||||
return model
|
return model
|
||||||
@@ -120,6 +208,21 @@ class BaseExecutor:
|
|||||||
def unwrap_model(self, model: nn.Module):
|
def unwrap_model(self, model: nn.Module):
|
||||||
return model.state_dict()
|
return model.state_dict()
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def checkpoint_context(self, model: nn.Module):
|
||||||
|
if self.use_distributed:
|
||||||
|
dist.barrier()
|
||||||
|
state_dict = self._gather_state_dict(model)
|
||||||
|
yield state_dict
|
||||||
|
if self.use_distributed:
|
||||||
|
dist.barrier()
|
||||||
|
|
||||||
|
def _gather_state_dict(self, model: nn.Module):
|
||||||
|
state_dict = self.unwrap_model(model)
|
||||||
|
if self.use_distributed and get_rank() != 0:
|
||||||
|
return None
|
||||||
|
return state_dict
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def use_distributed(self) -> bool:
|
def use_distributed(self) -> bool:
|
||||||
return get_world_size() > 1
|
return get_world_size() > 1
|
||||||
@@ -211,76 +314,115 @@ class DDPExecutor(BaseExecutor):
|
|||||||
|
|
||||||
@ExecutorFactory.register("fsdp")
|
@ExecutorFactory.register("fsdp")
|
||||||
class FSDPExecutor(BaseExecutor):
|
class FSDPExecutor(BaseExecutor):
|
||||||
|
"""FSDP executor using `torch.distributed.fsdp.fully_shard` (per-module API).
|
||||||
|
|
||||||
|
Wraps each child module individually via ``fully_shard``.
|
||||||
|
Skips the root model because ``ABC + Generic[T]`` in the MRO makes
|
||||||
|
``fully_shard``'s dynamic ``__class__`` assignment fail at the CPython level.
|
||||||
|
Original ``Parameter`` objects are preserved (as DTensors) — no
|
||||||
|
``FlatParameter``, no ``use_orig_params=True`` hack.
|
||||||
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
grad_accum_steps: int = 1,
|
grad_accum_steps: int = 1,
|
||||||
process_group=None,
|
mesh: Optional[Any] = None,
|
||||||
sharding_strategy=None,
|
mp_policy: Optional[Any] = None,
|
||||||
cpu_offload=None,
|
reshard_after_forward: bool = False,
|
||||||
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)
|
super().__init__(grad_accum_steps=grad_accum_steps)
|
||||||
self._fsdp_kwargs = {
|
self._mesh = mesh
|
||||||
k: v
|
self._mp_policy = mp_policy
|
||||||
for k, v in dict(
|
self._reshard_after_forward = reshard_after_forward
|
||||||
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:
|
def _prepare_model(self, model: nn.Module) -> nn.Module:
|
||||||
if not self.use_distributed:
|
if not self.use_distributed:
|
||||||
logger.warning("FSDP 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
|
||||||
self._original_model = model
|
|
||||||
device_id = torch.device("cuda", get_rank())
|
kwargs = dict(
|
||||||
model = FSDP(model, device_id=device_id, **self._fsdp_kwargs)
|
mesh=self._mesh,
|
||||||
logger.info("Model wrapped with FSDP (world_size=%d)", get_world_size())
|
mp_policy=self._mp_policy,
|
||||||
|
reshard_after_forward=self._reshard_after_forward,
|
||||||
|
)
|
||||||
|
kwargs = {k: v for k, v in kwargs.items() if v is not None}
|
||||||
|
|
||||||
|
for child in model.children():
|
||||||
|
if isinstance(child, nn.ModuleList):
|
||||||
|
for sub in child:
|
||||||
|
fully_shard(sub, **kwargs)
|
||||||
|
else:
|
||||||
|
fully_shard(child, **kwargs)
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"FSDP wrapping applied to %d direct children (root skipped for ABC compat)",
|
||||||
|
len(list(model.children())),
|
||||||
|
)
|
||||||
return model
|
return model
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
def _no_sync(self, model: nn.Module):
|
def _no_sync(self, model: nn.Module):
|
||||||
if isinstance(model, FSDP):
|
fsdp_modules = [m for m in model.modules() if isinstance(m, FSDPModule)]
|
||||||
return model.no_sync()
|
if fsdp_modules:
|
||||||
return contextlib.nullcontext()
|
for m in fsdp_modules:
|
||||||
|
m.set_requires_gradient_sync(False, recurse=True)
|
||||||
|
try:
|
||||||
|
yield
|
||||||
|
finally:
|
||||||
|
for m in fsdp_modules:
|
||||||
|
m.set_requires_gradient_sync(True, recurse=True)
|
||||||
|
else:
|
||||||
|
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 isinstance(model, FSDP) and self.use_distributed:
|
if not self.use_distributed:
|
||||||
total_norm = model.clip_grad_norm_(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 isinstance(model, FSDP) and self.use_distributed:
|
if not self.use_distributed:
|
||||||
with FSDP.state_dict_type(
|
return model.state_dict()
|
||||||
model,
|
|
||||||
StateDictType.FULL_STATE_DICT,
|
|
||||||
FullStateDictConfig(offload_to_cpu=True, rank0_only=False),
|
|
||||||
):
|
|
||||||
return model.state_dict()
|
|
||||||
|
|
||||||
return model.state_dict()
|
# unshard() and full_tensor() are collective ops — all ranks must
|
||||||
|
# participate. Non-rank-0 ranks still call them but discard results.
|
||||||
|
for module in model.modules():
|
||||||
|
if isinstance(module, FSDPModule):
|
||||||
|
module.unshard()
|
||||||
|
|
||||||
|
state_dict = model.state_dict()
|
||||||
|
result = {}
|
||||||
|
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():
|
||||||
|
if isinstance(module, FSDPModule):
|
||||||
|
module.reshard()
|
||||||
|
|
||||||
|
if get_rank() != 0:
|
||||||
|
return None
|
||||||
|
|
||||||
|
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)
|
|
||||||
@@ -1,13 +1,27 @@
|
|||||||
|
import logging
|
||||||
import os
|
import os
|
||||||
|
import signal
|
||||||
|
import socket
|
||||||
|
import threading
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from functools import wraps
|
from functools import wraps
|
||||||
from typing import Callable
|
from typing import Callable, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.distributed as dist
|
import torch.distributed as dist
|
||||||
import torch.multiprocessing as mp
|
import torch.multiprocessing as mp
|
||||||
|
|
||||||
|
from astrai.signal_handler import install_early_signal_handlers
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def find_free_port() -> str:
|
||||||
|
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
||||||
|
s.bind(("", 0))
|
||||||
|
return str(s.getsockname()[1])
|
||||||
|
|
||||||
|
|
||||||
def get_current_device():
|
def get_current_device():
|
||||||
return os.environ["LOCAL_DEVICE"]
|
return os.environ["LOCAL_DEVICE"]
|
||||||
@@ -108,6 +122,7 @@ def _run_single_rank(
|
|||||||
func: Callable,
|
func: Callable,
|
||||||
kwargs: dict,
|
kwargs: dict,
|
||||||
):
|
):
|
||||||
|
install_early_signal_handlers()
|
||||||
with setup_parallel(
|
with setup_parallel(
|
||||||
rank=rank,
|
rank=rank,
|
||||||
world_size=world_size,
|
world_size=world_size,
|
||||||
@@ -148,6 +163,7 @@ class TorchrunStrategy(LaunchStrategy):
|
|||||||
"""External orchestrator (torchrun, SLURM, K8s) — env vars pre-set."""
|
"""External orchestrator (torchrun, SLURM, K8s) — env vars pre-set."""
|
||||||
|
|
||||||
def launch(self, func: Callable, **kwargs):
|
def launch(self, func: Callable, **kwargs):
|
||||||
|
install_early_signal_handlers()
|
||||||
rank = int(os.environ["RANK"])
|
rank = int(os.environ["RANK"])
|
||||||
world_size = int(os.environ["WORLD_SIZE"])
|
world_size = int(os.environ["WORLD_SIZE"])
|
||||||
local_rank = int(os.environ.get("LOCAL_RANK", rank))
|
local_rank = int(os.environ.get("LOCAL_RANK", rank))
|
||||||
@@ -181,6 +197,7 @@ class LocalStrategy(LaunchStrategy):
|
|||||||
_run_single_rank(0, *args)
|
_run_single_rank(0, *args)
|
||||||
return
|
return
|
||||||
|
|
||||||
|
install_early_signal_handlers()
|
||||||
ctx = mp.start_processes(
|
ctx = mp.start_processes(
|
||||||
_run_single_rank,
|
_run_single_rank,
|
||||||
args=args,
|
args=args,
|
||||||
@@ -188,14 +205,46 @@ class LocalStrategy(LaunchStrategy):
|
|||||||
start_method=self.start_method,
|
start_method=self.start_method,
|
||||||
join=False,
|
join=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
parent_stop = threading.Event()
|
||||||
|
original_handlers = {}
|
||||||
|
|
||||||
|
def _parent_handler(signum, frame):
|
||||||
|
sig = signal.Signals(signum)
|
||||||
|
logger.warning(
|
||||||
|
"Parent (pid=%d) received %s, forwarding to children...",
|
||||||
|
os.getpid(),
|
||||||
|
sig.name,
|
||||||
|
)
|
||||||
|
parent_stop.set()
|
||||||
|
for p in ctx.processes:
|
||||||
|
if p.is_alive():
|
||||||
|
p.terminate()
|
||||||
|
|
||||||
|
for sig in (signal.SIGTERM, signal.SIGINT):
|
||||||
|
prev = signal.signal(sig, _parent_handler)
|
||||||
|
if prev not in (signal.SIG_DFL, signal.SIG_IGN, None, _parent_handler):
|
||||||
|
original_handlers[sig] = prev
|
||||||
|
|
||||||
try:
|
try:
|
||||||
while not ctx.join():
|
while not ctx.join() and not parent_stop.is_set():
|
||||||
pass
|
pass
|
||||||
except BaseException:
|
except BaseException:
|
||||||
|
logger.warning(
|
||||||
|
"Parent received unexpected exception, terminating children..."
|
||||||
|
)
|
||||||
for p in ctx.processes:
|
for p in ctx.processes:
|
||||||
p.terminate()
|
if p.is_alive():
|
||||||
ctx.join()
|
p.terminate()
|
||||||
raise
|
raise
|
||||||
|
finally:
|
||||||
|
for sig, handler in original_handlers.items():
|
||||||
|
signal.signal(sig, handler)
|
||||||
|
|
||||||
|
for p in ctx.processes:
|
||||||
|
p.join()
|
||||||
|
|
||||||
|
ctx.join()
|
||||||
|
|
||||||
|
|
||||||
def _detect_launcher() -> str:
|
def _detect_launcher() -> str:
|
||||||
@@ -217,11 +266,13 @@ def spawn_parallel_fn(
|
|||||||
world_size: int,
|
world_size: int,
|
||||||
backend: str = "nccl",
|
backend: str = "nccl",
|
||||||
master_addr: str = "localhost",
|
master_addr: str = "localhost",
|
||||||
master_port: str = "29500",
|
master_port: Optional[str] = None,
|
||||||
device_type: str = "cuda",
|
device_type: str = "cuda",
|
||||||
start_method: str = "spawn",
|
start_method: str = "spawn",
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
|
if master_port is None:
|
||||||
|
master_port = find_free_port()
|
||||||
launcher = _detect_launcher()
|
launcher = _detect_launcher()
|
||||||
if launcher in ("torchelastic", "torchrun", "external"):
|
if launcher in ("torchelastic", "torchrun", "external"):
|
||||||
strategy = TorchrunStrategy(
|
strategy = TorchrunStrategy(
|
||||||
|
|||||||
@@ -8,12 +8,14 @@ from astrai.preprocessing.builder import (
|
|||||||
from astrai.preprocessing.packing import (
|
from astrai.preprocessing.packing import (
|
||||||
PackingStrategy,
|
PackingStrategy,
|
||||||
PackingStrategyFactory,
|
PackingStrategyFactory,
|
||||||
|
plan_bfd,
|
||||||
)
|
)
|
||||||
from astrai.preprocessing.pipeline import Pipeline, filter_by_length
|
from astrai.preprocessing.pipeline import Pipeline, filter_by_length
|
||||||
from astrai.preprocessing.position_id import (
|
from astrai.preprocessing.position_id import (
|
||||||
PositionIdStrategy,
|
PositionIdStrategy,
|
||||||
PositionIdStrategyFactory,
|
PositionIdStrategyFactory,
|
||||||
)
|
)
|
||||||
|
from astrai.preprocessing.transform import TokenizeTransform
|
||||||
from astrai.preprocessing.writer import (
|
from astrai.preprocessing.writer import (
|
||||||
StoreWriter,
|
StoreWriter,
|
||||||
StoreWriterFactory,
|
StoreWriterFactory,
|
||||||
@@ -32,5 +34,7 @@ __all__ = [
|
|||||||
"SingleOutputMaskBuilder",
|
"SingleOutputMaskBuilder",
|
||||||
"StoreWriter",
|
"StoreWriter",
|
||||||
"StoreWriterFactory",
|
"StoreWriterFactory",
|
||||||
|
"TokenizeTransform",
|
||||||
"filter_by_length",
|
"filter_by_length",
|
||||||
|
"plan_bfd",
|
||||||
]
|
]
|
||||||
|
|||||||
+234
-21
@@ -94,9 +94,107 @@ class SectionRenderer:
|
|||||||
|
|
||||||
return all_ids, loss_mask
|
return all_ids, loss_mask
|
||||||
|
|
||||||
|
def process_sections_batch(
|
||||||
|
self,
|
||||||
|
items: list[dict],
|
||||||
|
sections: list,
|
||||||
|
config,
|
||||||
|
tokenizer,
|
||||||
|
*,
|
||||||
|
is_top_level=False,
|
||||||
|
filter_text=True,
|
||||||
|
):
|
||||||
|
"""Render and tokenize a group of records with batched Rust tokenization."""
|
||||||
|
has_template = any(s.get("template") for s in sections)
|
||||||
|
is_text_config = not has_template and all(
|
||||||
|
s["action"] == "train" for s in sections
|
||||||
|
)
|
||||||
|
plans: list[list[tuple[str, str, bool]]] = []
|
||||||
|
|
||||||
|
for item in items:
|
||||||
|
plan: list[tuple[str, str, bool]] = []
|
||||||
|
first_section = True
|
||||||
|
for sec in sections:
|
||||||
|
field = sec["field"]
|
||||||
|
action = sec["action"]
|
||||||
|
use_template = sec.get("template", False)
|
||||||
|
add_special = sec.get(
|
||||||
|
"add_special_tokens", not use_template and first_section
|
||||||
|
)
|
||||||
|
|
||||||
|
if use_template:
|
||||||
|
messages = item.get(field)
|
||||||
|
if not isinstance(messages, list) or not messages:
|
||||||
|
continue
|
||||||
|
for msg in messages:
|
||||||
|
role = msg.get("role", "")
|
||||||
|
rendered = tokenizer.apply_chat_template(
|
||||||
|
[msg], tokenize=False, add_generation_prompt=False
|
||||||
|
)
|
||||||
|
plan.append(
|
||||||
|
(rendered, _resolve_action(action, role, config), False)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
text = str(item.get(field, ""))
|
||||||
|
if not text.strip():
|
||||||
|
continue
|
||||||
|
if is_text_config and filter_text:
|
||||||
|
pp = config.preprocessing
|
||||||
|
if pp.min_chars > 0 and len(text) < pp.min_chars:
|
||||||
|
continue
|
||||||
|
if len(text) > pp.max_chars:
|
||||||
|
continue
|
||||||
|
plan.append((text, action, add_special))
|
||||||
|
|
||||||
|
first_section = False
|
||||||
|
plans.append(plan)
|
||||||
|
|
||||||
|
encoded: dict[tuple[int, int], list[int]] = {}
|
||||||
|
for add_special in (False, True):
|
||||||
|
refs = [
|
||||||
|
(item_idx, unit_idx, text)
|
||||||
|
for item_idx, plan in enumerate(plans)
|
||||||
|
for unit_idx, (text, _, add) in enumerate(plan)
|
||||||
|
if add == add_special
|
||||||
|
]
|
||||||
|
if not refs:
|
||||||
|
continue
|
||||||
|
ids_batch = tokenizer.encode(
|
||||||
|
[text for _, _, text in refs], add_special_tokens=add_special
|
||||||
|
)
|
||||||
|
for (item_idx, unit_idx, _), ids in zip(refs, ids_batch):
|
||||||
|
encoded[(item_idx, unit_idx)] = ids
|
||||||
|
|
||||||
|
outputs = []
|
||||||
|
max_len = config.preprocessing.max_seq_len
|
||||||
|
for item_idx, plan in enumerate(plans):
|
||||||
|
all_ids = []
|
||||||
|
loss_mask = []
|
||||||
|
if is_top_level and has_template and tokenizer.bos_token_id is not None:
|
||||||
|
all_ids.append(tokenizer.bos_token_id)
|
||||||
|
loss_mask.append(0)
|
||||||
|
for unit_idx, (_, action, _) in enumerate(plan):
|
||||||
|
ids = encoded[(item_idx, unit_idx)]
|
||||||
|
all_ids.extend(ids)
|
||||||
|
loss_mask.extend([1 if action == "train" else 0] * len(ids))
|
||||||
|
all_ids = all_ids[:max_len]
|
||||||
|
loss_mask = loss_mask[: len(all_ids)]
|
||||||
|
if not all_ids or (is_top_level and has_template and len(all_ids) <= 1):
|
||||||
|
outputs.append((None, None))
|
||||||
|
else:
|
||||||
|
outputs.append((all_ids, loss_mask))
|
||||||
|
return outputs
|
||||||
|
|
||||||
def process_list_field(self, item: dict, sections: list, config, tokenizer):
|
def process_list_field(self, item: dict, sections: list, config, tokenizer):
|
||||||
all_ids: list[int] = []
|
"""Tokenize a list-valued field, preserving per-element boundaries.
|
||||||
loss_mask: list[int] = []
|
|
||||||
|
Returns ``(list_of_id_lists, list_of_mask_lists)`` where each
|
||||||
|
inner list corresponds to one element of the source list. This
|
||||||
|
is critical for GRPO where each response must stay a separate
|
||||||
|
sequence so the strategy can form a ``[G, R]`` tensor.
|
||||||
|
"""
|
||||||
|
per_item_ids: list[list[int]] = []
|
||||||
|
per_item_masks: list[list[int]] = []
|
||||||
|
|
||||||
for sec in sections:
|
for sec in sections:
|
||||||
field = sec["field"]
|
field = sec["field"]
|
||||||
@@ -108,17 +206,13 @@ class SectionRenderer:
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
for val in values:
|
for val in values:
|
||||||
|
ids: list[int] = []
|
||||||
|
mask: list[int] = []
|
||||||
if use_template:
|
if use_template:
|
||||||
if isinstance(val, list):
|
if isinstance(val, list):
|
||||||
wrapper = {field: val}
|
wrapper = {field: val}
|
||||||
self._append_template(
|
self._append_template(
|
||||||
wrapper,
|
wrapper, field, action, tokenizer, config, ids, mask
|
||||||
field,
|
|
||||||
action,
|
|
||||||
tokenizer,
|
|
||||||
config,
|
|
||||||
all_ids,
|
|
||||||
loss_mask,
|
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
wrapper = {field: str(val)}
|
wrapper = {field: str(val)}
|
||||||
@@ -130,17 +224,55 @@ class SectionRenderer:
|
|||||||
False,
|
False,
|
||||||
False,
|
False,
|
||||||
config,
|
config,
|
||||||
all_ids,
|
ids,
|
||||||
loss_mask,
|
mask,
|
||||||
)
|
)
|
||||||
|
if ids:
|
||||||
|
max_len = config.preprocessing.max_seq_len
|
||||||
|
ids = ids[:max_len]
|
||||||
|
mask = mask[: len(ids)]
|
||||||
|
per_item_ids.append(ids)
|
||||||
|
per_item_masks.append(mask)
|
||||||
|
|
||||||
max_len = config.preprocessing.max_seq_len
|
if not per_item_ids:
|
||||||
all_ids = all_ids[:max_len]
|
|
||||||
loss_mask = loss_mask[: len(all_ids)]
|
|
||||||
|
|
||||||
if not all_ids:
|
|
||||||
return None, None
|
return None, None
|
||||||
return all_ids, loss_mask
|
return per_item_ids, per_item_masks
|
||||||
|
|
||||||
|
def process_list_field_batch(self, items, sections, config, tokenizer):
|
||||||
|
per_item_ids = [[] for _ in items]
|
||||||
|
per_item_masks = [[] for _ in items]
|
||||||
|
|
||||||
|
for sec in sections:
|
||||||
|
wrappers = []
|
||||||
|
owners = []
|
||||||
|
field = sec["field"]
|
||||||
|
for item_idx, item in enumerate(items):
|
||||||
|
values = item.get(field)
|
||||||
|
if not isinstance(values, list):
|
||||||
|
continue
|
||||||
|
for val in values:
|
||||||
|
if sec.get("template", False) and not isinstance(val, list):
|
||||||
|
continue
|
||||||
|
wrappers.append({field: val if isinstance(val, list) else str(val)})
|
||||||
|
owners.append(item_idx)
|
||||||
|
|
||||||
|
rendered = self.process_sections_batch(
|
||||||
|
wrappers,
|
||||||
|
[sec],
|
||||||
|
config,
|
||||||
|
tokenizer,
|
||||||
|
is_top_level=False,
|
||||||
|
filter_text=False,
|
||||||
|
)
|
||||||
|
for owner, (ids, mask) in zip(owners, rendered):
|
||||||
|
if ids:
|
||||||
|
per_item_ids[owner].append(ids)
|
||||||
|
per_item_masks[owner].append(mask)
|
||||||
|
|
||||||
|
return [
|
||||||
|
(ids, masks) if ids else (None, None)
|
||||||
|
for ids, masks in zip(per_item_ids, per_item_masks)
|
||||||
|
]
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def is_value_section(sections: list) -> bool:
|
def is_value_section(sections: list) -> bool:
|
||||||
@@ -209,6 +341,9 @@ class BaseMaskBuilder(ABC):
|
|||||||
@abstractmethod
|
@abstractmethod
|
||||||
def build(self, item: dict, config, tokenizer) -> Optional[dict]: ...
|
def build(self, item: dict, config, tokenizer) -> Optional[dict]: ...
|
||||||
|
|
||||||
|
def build_batch(self, items: list[dict], config, tokenizer) -> list[Optional[dict]]:
|
||||||
|
return [self.build(item, config, tokenizer) for item in items]
|
||||||
|
|
||||||
|
|
||||||
class MaskBuilderFactory(BaseFactory["BaseMaskBuilder"]):
|
class MaskBuilderFactory(BaseFactory["BaseMaskBuilder"]):
|
||||||
pass
|
pass
|
||||||
@@ -243,6 +378,27 @@ class SingleOutputMaskBuilder(BaseMaskBuilder):
|
|||||||
result["loss_mask"] = mask
|
result["loss_mask"] = mask
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
def build_batch(self, items, config, tokenizer):
|
||||||
|
sections = config.input.sections
|
||||||
|
if not sections:
|
||||||
|
return [None] * len(items)
|
||||||
|
rendered = self.renderer.process_sections_batch(
|
||||||
|
items, sections, config, tokenizer, is_top_level=True
|
||||||
|
)
|
||||||
|
results = []
|
||||||
|
for item, (ids, mask) in zip(items, rendered):
|
||||||
|
if ids is None:
|
||||||
|
results.append(None)
|
||||||
|
continue
|
||||||
|
result = {
|
||||||
|
"sequence": ids,
|
||||||
|
"domain": _extract_domain(item, config.output.domain_key),
|
||||||
|
}
|
||||||
|
if not all(m == 1 for m in mask):
|
||||||
|
result["loss_mask"] = mask
|
||||||
|
results.append(result)
|
||||||
|
return results
|
||||||
|
|
||||||
|
|
||||||
@MaskBuilderFactory.register("multi")
|
@MaskBuilderFactory.register("multi")
|
||||||
class MultiOutputMaskBuilder(BaseMaskBuilder):
|
class MultiOutputMaskBuilder(BaseMaskBuilder):
|
||||||
@@ -282,10 +438,18 @@ class MultiOutputMaskBuilder(BaseMaskBuilder):
|
|||||||
ids, mask = self.renderer.process_list_field(
|
ids, mask = self.renderer.process_list_field(
|
||||||
item, sections, config, tokenizer
|
item, sections, config, tokenizer
|
||||||
)
|
)
|
||||||
else:
|
if ids is None:
|
||||||
ids, mask = self.renderer.process_sections(
|
continue
|
||||||
item, sections, config, tokenizer, is_top_level=True
|
# ids is List[List[int]] — preserve per-response structure
|
||||||
)
|
result[output_key] = ids
|
||||||
|
if mask is not None:
|
||||||
|
result[mask_key] = mask
|
||||||
|
any_output = True
|
||||||
|
continue
|
||||||
|
|
||||||
|
ids, mask = self.renderer.process_sections(
|
||||||
|
item, sections, config, tokenizer, is_top_level=True
|
||||||
|
)
|
||||||
|
|
||||||
if ids is None:
|
if ids is None:
|
||||||
continue
|
continue
|
||||||
@@ -304,6 +468,49 @@ class MultiOutputMaskBuilder(BaseMaskBuilder):
|
|||||||
result["domain"] = _extract_domain(item, config.output.domain_key)
|
result["domain"] = _extract_domain(item, config.output.domain_key)
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
def build_batch(self, items, config, tokenizer):
|
||||||
|
sources_spec = getattr(config.input, "sources", None)
|
||||||
|
if not sources_spec:
|
||||||
|
return [None] * len(items)
|
||||||
|
|
||||||
|
results = [{} for _ in items]
|
||||||
|
for output_key, spec in sources_spec.items():
|
||||||
|
sections = spec.get("sections", [])
|
||||||
|
if not sections:
|
||||||
|
continue
|
||||||
|
if self.renderer.is_value_section(sections):
|
||||||
|
for item, result in zip(items, results):
|
||||||
|
value = self.renderer.extract_raw_value(item, sections)
|
||||||
|
if value is not None:
|
||||||
|
result[output_key] = value
|
||||||
|
continue
|
||||||
|
|
||||||
|
mask_key = spec.get("mask_key", f"{output_key}_mask")
|
||||||
|
if spec.get("list_field", False):
|
||||||
|
rendered = self.renderer.process_list_field_batch(
|
||||||
|
items, sections, config, tokenizer
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
rendered = self.renderer.process_sections_batch(
|
||||||
|
items, sections, config, tokenizer, is_top_level=True
|
||||||
|
)
|
||||||
|
|
||||||
|
for result, (ids, mask) in zip(results, rendered):
|
||||||
|
if ids is None:
|
||||||
|
continue
|
||||||
|
result[output_key] = ids
|
||||||
|
if spec.get("list_field", False) or not all(m == 1 for m in mask):
|
||||||
|
result[mask_key] = mask
|
||||||
|
elif "mask_key" in spec:
|
||||||
|
result[mask_key] = mask
|
||||||
|
|
||||||
|
return [
|
||||||
|
({**result, "domain": _extract_domain(item, config.output.domain_key)})
|
||||||
|
if result
|
||||||
|
else None
|
||||||
|
for item, result in zip(items, results)
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
@MaskBuilderFactory.register("sectioned")
|
@MaskBuilderFactory.register("sectioned")
|
||||||
class SectionedMaskBuilder(BaseMaskBuilder):
|
class SectionedMaskBuilder(BaseMaskBuilder):
|
||||||
@@ -322,3 +529,9 @@ class SectionedMaskBuilder(BaseMaskBuilder):
|
|||||||
if sources_spec:
|
if sources_spec:
|
||||||
return self._multi.build(item, config, tokenizer)
|
return self._multi.build(item, config, tokenizer)
|
||||||
return self._single.build(item, config, tokenizer)
|
return self._single.build(item, config, tokenizer)
|
||||||
|
|
||||||
|
def build_batch(self, items, config, tokenizer):
|
||||||
|
sources_spec = getattr(config.input, "sources", None)
|
||||||
|
if sources_spec:
|
||||||
|
return self._multi.build_batch(items, config, tokenizer)
|
||||||
|
return self._single.build_batch(items, config, tokenizer)
|
||||||
|
|||||||
@@ -0,0 +1,124 @@
|
|||||||
|
"""Shared preprocessing kernel used by both :class:`Pipeline` and
|
||||||
|
:class:`TokenizeTransform`.
|
||||||
|
|
||||||
|
The two entry points previously duplicated ~60 % of their logic:
|
||||||
|
record iteration, mask-builder invocation, primary-id extraction,
|
||||||
|
per-key accumulation, dtype inference and position-id generation.
|
||||||
|
This module factors out the common core as pure functions so that
|
||||||
|
the online (``TokenizeTransform``) and offline (``Pipeline``) paths
|
||||||
|
stay in lockstep.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from itertools import chain
|
||||||
|
from typing import Dict, Iterator, List, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.config.preprocess_config import PipelineConfig
|
||||||
|
from astrai.preprocessing.builder import MaskBuilderFactory
|
||||||
|
from astrai.preprocessing.position_id import PositionIdStrategyFactory
|
||||||
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
|
||||||
|
def build_preprocessing_components(config: PipelineConfig, tokenizer_path: str):
|
||||||
|
"""Load tokenizer, mask builder and position-id strategy together.
|
||||||
|
|
||||||
|
Both ``Pipeline`` and ``TokenizeTransform`` need the same triple;
|
||||||
|
centralising the construction avoids drift (e.g. one path forgetting
|
||||||
|
to create the position-id strategy).
|
||||||
|
"""
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
|
||||||
|
mask_builder = MaskBuilderFactory.create("sectioned")
|
||||||
|
position_strategy = PositionIdStrategyFactory.create(
|
||||||
|
config.output.position_ids_mode
|
||||||
|
)
|
||||||
|
return tokenizer, mask_builder, position_strategy
|
||||||
|
|
||||||
|
|
||||||
|
def primary_ids(result: dict) -> List[int]:
|
||||||
|
"""Return the first flat int-list value in *result*.
|
||||||
|
|
||||||
|
Used for token counting and position-id generation when the
|
||||||
|
primary key name is not known (DPO uses ``chosen``, GRPO uses
|
||||||
|
``prompts``, SFT uses ``sequence``).
|
||||||
|
"""
|
||||||
|
for val in result.values():
|
||||||
|
if isinstance(val, list) and val and isinstance(val[0], int):
|
||||||
|
return val
|
||||||
|
return []
|
||||||
|
|
||||||
|
|
||||||
|
def infer_dtype(ids: List) -> torch.dtype:
|
||||||
|
"""Float values become float32, everything else int32."""
|
||||||
|
if ids and isinstance(ids[0], float):
|
||||||
|
return torch.float32
|
||||||
|
return torch.int32
|
||||||
|
|
||||||
|
|
||||||
|
def iter_raw_records(
|
||||||
|
records: List[dict],
|
||||||
|
mask_builder,
|
||||||
|
config: PipelineConfig,
|
||||||
|
tokenizer,
|
||||||
|
) -> Iterator[dict]:
|
||||||
|
"""Yield mask-builder output dicts for each record, skipping failures.
|
||||||
|
|
||||||
|
Drops ``domain`` from the result (callers that need it should read
|
||||||
|
it before calling this). Each yielded dict maps a key
|
||||||
|
(``sequence``, ``chosen``, ``responses``…) to either a flat
|
||||||
|
``List[int]`` or a nested ``List[List[int]]`` (GRPO responses/masks).
|
||||||
|
"""
|
||||||
|
for item in records:
|
||||||
|
result = mask_builder.build(item, config, tokenizer)
|
||||||
|
if result is None:
|
||||||
|
continue
|
||||||
|
result.pop("domain", None)
|
||||||
|
if not primary_ids(result):
|
||||||
|
continue
|
||||||
|
yield result
|
||||||
|
|
||||||
|
|
||||||
|
def to_per_record_tensors(
|
||||||
|
raw: Dict[str, list],
|
||||||
|
) -> Dict[str, List[torch.Tensor]]:
|
||||||
|
"""Convert an accumulated ``{key: [per-record ids]}`` dict to tensors.
|
||||||
|
|
||||||
|
Handles three shapes transparently:
|
||||||
|
|
||||||
|
- ``List[int]`` per record (``sequence``, ``chosen``…) → one tensor per record.
|
||||||
|
- ``List[List[int]]`` per record (GRPO ``responses``/``masks``) → one
|
||||||
|
``List[Tensor]`` per record (nested), preserving the per-response
|
||||||
|
boundary so downstream code can index responses individually.
|
||||||
|
- ``List[int]`` for the whole shard (pre-packed keys) → single tensor.
|
||||||
|
|
||||||
|
The detection mirrors the previous inline logic in
|
||||||
|
``Pipeline._flush`` and ``TokenizeTransform.apply``.
|
||||||
|
"""
|
||||||
|
tensors: Dict[str, List[torch.Tensor]] = {}
|
||||||
|
for key, ids_list in raw.items():
|
||||||
|
if ids_list and isinstance(ids_list[0], list):
|
||||||
|
tensors[key] = [
|
||||||
|
[torch.tensor(sub, dtype=infer_dtype(sub)) for sub in ids]
|
||||||
|
if ids and isinstance(ids[0], list)
|
||||||
|
else torch.tensor(ids, dtype=infer_dtype(ids))
|
||||||
|
for ids in ids_list
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
tensors[key] = [
|
||||||
|
torch.tensor(list(chain.from_iterable(ids_list)), dtype=torch.int32)
|
||||||
|
]
|
||||||
|
return tensors
|
||||||
|
|
||||||
|
|
||||||
|
def build_position_ids(
|
||||||
|
sequences: List[List[int]],
|
||||||
|
strategy,
|
||||||
|
) -> Optional[List[int]]:
|
||||||
|
"""Generate position ids for *sequences* using *strategy*.
|
||||||
|
|
||||||
|
Returns ``None`` when the strategy produces no ids (e.g. ``none``
|
||||||
|
mode), so callers can skip attaching the key instead of storing
|
||||||
|
an empty list.
|
||||||
|
"""
|
||||||
|
pos_ids = strategy.generate(sequences)
|
||||||
|
return pos_ids or None
|
||||||
@@ -19,6 +19,43 @@ def _truncate(seq: List[int], max_len: int, mode: str) -> List[int]:
|
|||||||
return seq[:max_len]
|
return seq[:max_len]
|
||||||
|
|
||||||
|
|
||||||
|
def plan_bfd(
|
||||||
|
sequences: List[List[int]], max_packed_len: int, truncation_mode: str = "keep_start"
|
||||||
|
) -> List[List[int]]:
|
||||||
|
"""Best-Fit Decreasing bin packing of *sequences* into bins.
|
||||||
|
|
||||||
|
Returns a list of bins, each bin a list of original indices into
|
||||||
|
*sequences*. Bin capacities are respected on the *truncated*
|
||||||
|
length of each sequence (so a sequence longer than
|
||||||
|
*max_packed_len* counts at *max_packed_len*).
|
||||||
|
|
||||||
|
Pure index-based so callers can apply the same plan to any
|
||||||
|
aligned key (``loss_mask``, ``position_ids``…).
|
||||||
|
"""
|
||||||
|
n = len(sequences)
|
||||||
|
order = sorted(range(n), key=lambda i: len(sequences[i]), reverse=True)
|
||||||
|
bins: List[List[int]] = []
|
||||||
|
bin_lengths: List[int] = []
|
||||||
|
|
||||||
|
for orig_idx in order:
|
||||||
|
seq_len = len(_truncate(sequences[orig_idx], max_packed_len, truncation_mode))
|
||||||
|
best_bin = None
|
||||||
|
best_remain = max_packed_len + 1
|
||||||
|
for i, bl in enumerate(bin_lengths):
|
||||||
|
remain = max_packed_len - bl
|
||||||
|
if seq_len <= remain < best_remain:
|
||||||
|
best_remain = remain
|
||||||
|
best_bin = i
|
||||||
|
if best_bin is not None:
|
||||||
|
bins[best_bin].append(orig_idx)
|
||||||
|
bin_lengths[best_bin] += seq_len
|
||||||
|
else:
|
||||||
|
bins.append([orig_idx])
|
||||||
|
bin_lengths.append(seq_len)
|
||||||
|
|
||||||
|
return bins
|
||||||
|
|
||||||
|
|
||||||
class PackingStrategy(ABC):
|
class PackingStrategy(ABC):
|
||||||
"""Reorder and truncate sequences within a shard."""
|
"""Reorder and truncate sequences within a shard."""
|
||||||
|
|
||||||
@@ -70,7 +107,7 @@ class BFDPacking(PackingStrategy):
|
|||||||
sequences = keys.get("sequence", [])
|
sequences = keys.get("sequence", [])
|
||||||
if not sequences:
|
if not sequences:
|
||||||
return keys
|
return keys
|
||||||
bins = self._plan(sequences, max_packed_len, truncation_mode)
|
bins = plan_bfd(sequences, max_packed_len, truncation_mode)
|
||||||
|
|
||||||
packed: Dict[str, List[List[int]]] = {}
|
packed: Dict[str, List[List[int]]] = {}
|
||||||
for k, vals in keys.items():
|
for k, vals in keys.items():
|
||||||
@@ -91,31 +128,49 @@ class BFDPacking(PackingStrategy):
|
|||||||
result.extend(vals[i])
|
result.extend(vals[i])
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
@PackingStrategyFactory.register("bfd_split")
|
||||||
|
class BFDSplitPacking(BFDPacking):
|
||||||
|
"""BFD packing with over-length sequences split into chunks.
|
||||||
|
|
||||||
|
Sequences longer than *max_packed_len* are split into consecutive
|
||||||
|
chunks of at most *max_packed_len* tokens instead of being
|
||||||
|
truncated. Each chunk becomes an independent sequence that enters
|
||||||
|
BFD planning. All keys (``loss_mask``, ``position_ids``, …) are
|
||||||
|
split in lockstep so per-token alignment is preserved.
|
||||||
|
|
||||||
|
Note: because each chunk is treated as a separate document, the
|
||||||
|
second chunk of a split sequence loses the preceding context.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def apply(
|
||||||
|
self,
|
||||||
|
keys: Dict[str, List[List[int]]],
|
||||||
|
max_packed_len: int,
|
||||||
|
truncation_mode: str,
|
||||||
|
) -> Dict[str, List[List[int]]]:
|
||||||
|
sequences = keys.get("sequence", [])
|
||||||
|
if not sequences:
|
||||||
|
return keys
|
||||||
|
if max_packed_len <= 0:
|
||||||
|
return super().apply(keys, max_packed_len, truncation_mode)
|
||||||
|
|
||||||
|
split_keys = self._split_all(keys, max_packed_len)
|
||||||
|
return super().apply(split_keys, max_packed_len, truncation_mode)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _plan(
|
def _split_all(
|
||||||
sequences: List[List[int]], max_packed_len: int, truncation_mode: str
|
keys: Dict[str, List[List[int]]], max_packed_len: int
|
||||||
) -> List[List[int]]:
|
) -> Dict[str, List[List[int]]]:
|
||||||
n = len(sequences)
|
"""Split every sequence exceeding *max_packed_len* into chunks,
|
||||||
order = sorted(range(n), key=lambda i: len(sequences[i]), reverse=True)
|
applying the same chunk boundaries to all keys."""
|
||||||
bins: List[List[int]] = []
|
sequences = keys["sequence"]
|
||||||
bin_lengths: List[int] = []
|
chunk_bounds = [list(range(0, len(s), max_packed_len)) for s in sequences]
|
||||||
|
result: Dict[str, List[List[int]]] = {}
|
||||||
for orig_idx in order:
|
for key, vals in keys.items():
|
||||||
seq_len = len(
|
split_vals: List[List[int]] = []
|
||||||
_truncate(sequences[orig_idx], max_packed_len, truncation_mode)
|
for val, starts in zip(vals, chunk_bounds):
|
||||||
)
|
for start in starts:
|
||||||
best_bin = None
|
split_vals.append(val[start : start + max_packed_len])
|
||||||
best_remain = max_packed_len + 1
|
result[key] = split_vals
|
||||||
for i, bl in enumerate(bin_lengths):
|
return result
|
||||||
remain = max_packed_len - bl
|
|
||||||
if seq_len <= remain < best_remain:
|
|
||||||
best_remain = remain
|
|
||||||
best_bin = i
|
|
||||||
if best_bin is not None:
|
|
||||||
bins[best_bin].append(orig_idx)
|
|
||||||
bin_lengths[best_bin] += seq_len
|
|
||||||
else:
|
|
||||||
bins.append([orig_idx])
|
|
||||||
bin_lengths.append(seq_len)
|
|
||||||
|
|
||||||
return bins
|
|
||||||
|
|||||||
@@ -1,9 +1,13 @@
|
|||||||
"""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.
|
||||||
|
|
||||||
|
Record iteration, mask building, primary-id extraction and per-key
|
||||||
|
accumulation are shared with :class:`TokenizeTransform` via the
|
||||||
|
:mod:`astrai.preprocessing.core` helpers.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import json
|
import json
|
||||||
@@ -17,11 +21,12 @@ import torch
|
|||||||
import tqdm
|
import tqdm
|
||||||
|
|
||||||
from astrai.config.preprocess_config import PipelineConfig
|
from astrai.config.preprocess_config import PipelineConfig
|
||||||
from astrai.preprocessing.builder import MaskBuilderFactory
|
from astrai.preprocessing.core import (
|
||||||
|
build_preprocessing_components,
|
||||||
|
primary_ids,
|
||||||
|
)
|
||||||
from astrai.preprocessing.packing import PackingStrategyFactory
|
from astrai.preprocessing.packing import PackingStrategyFactory
|
||||||
from astrai.preprocessing.position_id import PositionIdStrategyFactory
|
|
||||||
from astrai.preprocessing.writer import StoreWriterFactory
|
from astrai.preprocessing.writer import StoreWriterFactory
|
||||||
from astrai.tokenize import AutoTokenizer
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -64,20 +69,21 @@ class Pipeline:
|
|||||||
self.output_dir = output_dir
|
self.output_dir = output_dir
|
||||||
self.tokenizer_path = tokenizer_path
|
self.tokenizer_path = tokenizer_path
|
||||||
|
|
||||||
self.mask_builder = MaskBuilderFactory.create("sectioned")
|
self.tokenizer, self.mask_builder, self._position_id = (
|
||||||
|
build_preprocessing_components(config, tokenizer_path)
|
||||||
|
)
|
||||||
self._packer = PackingStrategyFactory.create(
|
self._packer = PackingStrategyFactory.create(
|
||||||
config.preprocessing.packing_strategy
|
config.preprocessing.packing_strategy
|
||||||
)
|
)
|
||||||
self._position_id = PositionIdStrategyFactory.create(
|
|
||||||
config.output.position_ids_mode
|
|
||||||
)
|
|
||||||
self._writer = StoreWriterFactory.create(config.output.storage_format)
|
self._writer = StoreWriterFactory.create(config.output.storage_format)
|
||||||
|
|
||||||
def transform(self, item: dict) -> Optional[dict]:
|
def transform(self, item: dict) -> Optional[dict]:
|
||||||
return self.mask_builder.build(item, self.config, self._tokenizer)
|
return self.mask_builder.build(item, self.config, self.tokenizer)
|
||||||
|
|
||||||
|
def transform_batch(self, items: list[dict]) -> list[Optional[dict]]:
|
||||||
|
return self.mask_builder.build_batch(items, self.config, self.tokenizer)
|
||||||
|
|
||||||
def run(self):
|
def run(self):
|
||||||
self._tokenizer = AutoTokenizer.from_pretrained(self.tokenizer_path)
|
|
||||||
domains: dict = defaultdict(lambda: defaultdict(list))
|
domains: dict = defaultdict(lambda: defaultdict(list))
|
||||||
total_tokens = 0
|
total_tokens = 0
|
||||||
shard_idx: dict[str, int] = defaultdict(int)
|
shard_idx: dict[str, int] = defaultdict(int)
|
||||||
@@ -85,59 +91,59 @@ class Pipeline:
|
|||||||
|
|
||||||
pp = self.config.preprocessing
|
pp = self.config.preprocessing
|
||||||
|
|
||||||
for item in tqdm.tqdm(
|
progress = tqdm.tqdm(desc="Tokenizing", unit="docs", mininterval=0.5)
|
||||||
self._iter_items(), desc="Tokenizing", unit="docs", mininterval=0.5
|
stop = False
|
||||||
):
|
for items in self._iter_batches(pp.batch_size):
|
||||||
if pp.max_items and count >= pp.max_items:
|
progress.update(len(items))
|
||||||
break
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
result = self.transform(item)
|
results = self.transform_batch(items)
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Failed to process item #%d, skipping", count + 1, exc_info=True
|
"Failed to process batch, retrying records individually",
|
||||||
|
exc_info=True,
|
||||||
)
|
)
|
||||||
continue
|
results = []
|
||||||
if result is None:
|
for item in items:
|
||||||
continue
|
try:
|
||||||
|
results.append(self.transform(item))
|
||||||
|
except Exception:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to process item, skipping", exc_info=True
|
||||||
|
)
|
||||||
|
results.append(None)
|
||||||
|
|
||||||
domain = result.pop("domain", "__default__")
|
for result in results:
|
||||||
|
if pp.max_items and count >= pp.max_items:
|
||||||
|
stop = True
|
||||||
|
break
|
||||||
|
if result is None:
|
||||||
|
continue
|
||||||
|
|
||||||
is_multi = bool(getattr(self.config.input, "sources", None))
|
domain = result.pop("domain", "__default__")
|
||||||
if is_multi:
|
ids = primary_ids(result)
|
||||||
ids = self._primary_ids(result)
|
if not ids:
|
||||||
else:
|
continue
|
||||||
ids = result.pop("sequence")
|
|
||||||
result["sequence"] = ids
|
|
||||||
|
|
||||||
if not ids:
|
bucket = domains[domain]
|
||||||
continue
|
self._align_bucket(bucket, result, ids)
|
||||||
|
for key, val in result.items():
|
||||||
|
bucket[key].append(val)
|
||||||
|
|
||||||
bucket = domains[domain]
|
count += 1
|
||||||
self._align_bucket(bucket, result, ids)
|
total_tokens += len(ids)
|
||||||
for key, val in result.items():
|
|
||||||
bucket[key].append(val)
|
|
||||||
|
|
||||||
count += 1
|
if total_tokens >= self.config.output.max_tokens_per_shard:
|
||||||
total_tokens += len(ids)
|
self._flush(domains, shard_idx)
|
||||||
|
domains.clear()
|
||||||
|
total_tokens = 0
|
||||||
|
if stop:
|
||||||
|
break
|
||||||
|
|
||||||
if total_tokens >= self.config.output.max_tokens_per_shard:
|
progress.close()
|
||||||
self._flush(domains, shard_idx)
|
|
||||||
domains.clear()
|
|
||||||
total_tokens = 0
|
|
||||||
|
|
||||||
if total_tokens > 0:
|
if total_tokens > 0:
|
||||||
self._flush(domains, shard_idx)
|
self._flush(domains, shard_idx)
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _primary_ids(result: dict) -> list:
|
|
||||||
"""Return the first list-valued entry in *result* as the primary id
|
|
||||||
sequence for token counting."""
|
|
||||||
for val in result.values():
|
|
||||||
if isinstance(val, list) and val and isinstance(val[0], int):
|
|
||||||
return val
|
|
||||||
return []
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _align_bucket(bucket: dict, result: dict, ids: list):
|
def _align_bucket(bucket: dict, result: dict, ids: list):
|
||||||
"""Pad previously-accumulated keys that are missing from *result*."""
|
"""Pad previously-accumulated keys that are missing from *result*."""
|
||||||
@@ -149,11 +155,29 @@ class Pipeline:
|
|||||||
def _iter_items(self):
|
def _iter_items(self):
|
||||||
for path in self.paths:
|
for path in self.paths:
|
||||||
with open(path, "r", encoding="utf-8") as f:
|
with open(path, "r", encoding="utf-8") as f:
|
||||||
for line in f:
|
if path.endswith(".json"):
|
||||||
line = line.strip()
|
data = json.load(f)
|
||||||
if not line:
|
if isinstance(data, dict):
|
||||||
continue
|
yield data
|
||||||
yield json.loads(line)
|
elif isinstance(data, list):
|
||||||
|
yield from data
|
||||||
|
else:
|
||||||
|
for line in f:
|
||||||
|
line = line.strip()
|
||||||
|
if not line:
|
||||||
|
continue
|
||||||
|
yield json.loads(line)
|
||||||
|
|
||||||
|
def _iter_batches(self, batch_size: int):
|
||||||
|
batch_size = max(1, batch_size)
|
||||||
|
batch = []
|
||||||
|
for item in self._iter_items():
|
||||||
|
batch.append(item)
|
||||||
|
if len(batch) >= batch_size:
|
||||||
|
yield batch
|
||||||
|
batch = []
|
||||||
|
if batch:
|
||||||
|
yield batch
|
||||||
|
|
||||||
def _flush(self, domains, shard_idx):
|
def _flush(self, domains, shard_idx):
|
||||||
for domain, keys in domains.items():
|
for domain, keys in domains.items():
|
||||||
@@ -163,24 +187,12 @@ class Pipeline:
|
|||||||
original_sequences = keys.get("sequence", [])
|
original_sequences = keys.get("sequence", [])
|
||||||
mode = self.config.output.position_ids_mode
|
mode = self.config.output.position_ids_mode
|
||||||
|
|
||||||
if mode == "doc_reset" and original_sequences:
|
keys = self._inject_doc_reset_position_ids(keys, mode, original_sequences)
|
||||||
keys["position_ids"] = [list(range(len(s))) for s in original_sequences]
|
|
||||||
|
|
||||||
keys = self._packer.apply(dict(keys), pp.max_packed_len, pp.truncation_mode)
|
keys = self._packer.apply(dict(keys), pp.max_packed_len, pp.truncation_mode)
|
||||||
|
tensors = self._to_tensors(keys)
|
||||||
tensors: Dict[str, List[torch.Tensor]] = {}
|
tensors = self._inject_continuous_position_ids(
|
||||||
for key, ids_list in keys.items():
|
tensors, mode, keys.get("sequence", [])
|
||||||
dt = _STR_TO_DTYPE.get(
|
)
|
||||||
self.config.output.dtype.get(key, "int32"), torch.int32
|
|
||||||
)
|
|
||||||
tensors[key] = [
|
|
||||||
torch.tensor(list(chain.from_iterable(ids_list)), dtype=dt)
|
|
||||||
]
|
|
||||||
|
|
||||||
if mode == "continuous" and original_sequences:
|
|
||||||
pos_ids = self._position_id.generate(keys.get("sequence", []))
|
|
||||||
if pos_ids:
|
|
||||||
tensors["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
|
|
||||||
|
|
||||||
self._writer.save(self.output_dir, domain, idx, tensors)
|
self._writer.save(self.output_dir, domain, idx, tensors)
|
||||||
shard_idx[domain] = idx + 1
|
shard_idx[domain] = idx + 1
|
||||||
@@ -190,3 +202,76 @@ class Pipeline:
|
|||||||
f" saved {domain}/shard_{idx:04d} "
|
f" saved {domain}/shard_{idx:04d} "
|
||||||
f"({tensors[first_key][0].numel():,} tokens)"
|
f"({tensors[first_key][0].numel():,} tokens)"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _inject_doc_reset_position_ids(
|
||||||
|
self,
|
||||||
|
keys: Dict[str, list],
|
||||||
|
mode: str,
|
||||||
|
original_sequences: List[List[int]],
|
||||||
|
) -> Dict[str, list]:
|
||||||
|
"""Attach per-document position_ids before packing (``doc_reset``).
|
||||||
|
|
||||||
|
``doc_reset`` position ids must enter the packer so that each
|
||||||
|
packed bin concatenates the per-doc ranges in bin order. The
|
||||||
|
per-record structure ``[range(len(s)) for s in seqs]`` is required
|
||||||
|
by the packer (it concatenates per-record lists per bin); the
|
||||||
|
``PositionIdStrategy.generate`` flattens, so it cannot be used
|
||||||
|
directly here — it is only consulted for the ``continuous``
|
||||||
|
post-packing path.
|
||||||
|
"""
|
||||||
|
if mode != "doc_reset" or not original_sequences:
|
||||||
|
return keys
|
||||||
|
keys["position_ids"] = [list(range(len(s))) for s in original_sequences]
|
||||||
|
return keys
|
||||||
|
|
||||||
|
def _inject_continuous_position_ids(
|
||||||
|
self,
|
||||||
|
tensors: Dict[str, List[torch.Tensor]],
|
||||||
|
mode: str,
|
||||||
|
packed_sequences: List[List[int]],
|
||||||
|
) -> Dict[str, List[torch.Tensor]]:
|
||||||
|
"""Attach a single continuous position_ids tensor after packing.
|
||||||
|
|
||||||
|
``continuous`` mode spans the whole shard (post-packing), so it
|
||||||
|
cannot participate in bin packing — it is computed from the
|
||||||
|
packed sequences and appended directly to the tensor dict.
|
||||||
|
"""
|
||||||
|
if mode != "continuous" or not packed_sequences:
|
||||||
|
return tensors
|
||||||
|
pos_ids = self._position_id.generate(packed_sequences)
|
||||||
|
if pos_ids:
|
||||||
|
tensors["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
|
||||||
|
return tensors
|
||||||
|
|
||||||
|
def _to_tensors(self, keys: Dict[str, list]) -> Dict[str, List[torch.Tensor]]:
|
||||||
|
"""Convert packed per-key id lists to tensors.
|
||||||
|
|
||||||
|
Honours ``config.output.dtype`` overrides per key; falls back to
|
||||||
|
``int32``. Handles three shapes (see
|
||||||
|
:func:`astrai.preprocessing.core.to_per_record_tensors` for the
|
||||||
|
equivalent online-path helper):
|
||||||
|
- ``List[int]`` per record → one tensor per record.
|
||||||
|
- ``List[List[int]]`` per record (GRPO responses/masks) → one tensor
|
||||||
|
per record, inner lists flattened.
|
||||||
|
- ``List[int]`` for the whole shard (pre-packed keys) → single tensor.
|
||||||
|
"""
|
||||||
|
tensors: Dict[str, List[torch.Tensor]] = {}
|
||||||
|
for key, ids_list in keys.items():
|
||||||
|
dt = _STR_TO_DTYPE.get(
|
||||||
|
self.config.output.dtype.get(key, "int32"), torch.int32
|
||||||
|
)
|
||||||
|
if ids_list and isinstance(ids_list[0], list):
|
||||||
|
tensors[key] = [
|
||||||
|
torch.tensor(
|
||||||
|
list(chain.from_iterable(ids))
|
||||||
|
if ids and isinstance(ids[0], list)
|
||||||
|
else ids,
|
||||||
|
dtype=dt,
|
||||||
|
)
|
||||||
|
for ids in ids_list
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
tensors[key] = [
|
||||||
|
torch.tensor(list(chain.from_iterable(ids_list)), dtype=dt)
|
||||||
|
]
|
||||||
|
return tensors
|
||||||
|
|||||||
@@ -0,0 +1,92 @@
|
|||||||
|
"""Tokenization transform for JSONL record streams.
|
||||||
|
|
||||||
|
Bridges the Reader layer (``JsonlStore`` reads raw JSON records) and the
|
||||||
|
Dataset layer (expects per-record tensors). Holds the tokenizer,
|
||||||
|
mask-builder and position-id strategy together so that I/O code stays
|
||||||
|
free of model dependencies.
|
||||||
|
|
||||||
|
The record-processing core (mask building, primary-id extraction,
|
||||||
|
per-key tensorisation, position-id generation) is shared with
|
||||||
|
:class:`astrai.preprocessing.pipeline.Pipeline` via the
|
||||||
|
:mod:`astrai.preprocessing.core` helpers.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Dict, List
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.config.preprocess_config import PipelineConfig
|
||||||
|
from astrai.preprocessing.core import (
|
||||||
|
build_position_ids,
|
||||||
|
build_preprocessing_components,
|
||||||
|
iter_raw_records,
|
||||||
|
to_per_record_tensors,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TokenizeTransform:
|
||||||
|
"""Tokenize raw JSONL record dicts into per-key tensor lists.
|
||||||
|
|
||||||
|
Owns the three preprocessing concerns that were previously inlined in
|
||||||
|
``JsonlStore``: tokenization, loss-mask construction and position-id
|
||||||
|
generation. Constructing it loads the tokenizer, so it is intentionally
|
||||||
|
cheap to pass around once built.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config: Pipeline config describing sections / masks / position mode.
|
||||||
|
tokenizer_path: Path passed to ``AutoTokenizer.from_pretrained``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, config: PipelineConfig, tokenizer_path: str):
|
||||||
|
self.config = config
|
||||||
|
self.tokenizer, self.mask_builder, self.position_strategy = (
|
||||||
|
build_preprocessing_components(config, tokenizer_path)
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_config_file(cls, config_path: str) -> "TokenizeTransform":
|
||||||
|
"""Build from a ``dataset_config.json`` file path.
|
||||||
|
|
||||||
|
The config file follows :class:`PipelineConfig` schema with an
|
||||||
|
extra ``tokenizer_path`` field. When omitted, the config's
|
||||||
|
parent directory is used as the tokenizer path.
|
||||||
|
"""
|
||||||
|
root = Path(config_path).parent
|
||||||
|
with open(config_path, "r", encoding="utf-8") as f:
|
||||||
|
raw_config = json.load(f)
|
||||||
|
tokenizer_path = raw_config.pop("tokenizer_path", None) or str(root)
|
||||||
|
config = PipelineConfig.from_dict(raw_config)
|
||||||
|
return cls(config, tokenizer_path)
|
||||||
|
|
||||||
|
def apply(self, records: List[dict]) -> Dict[str, list]:
|
||||||
|
"""Tokenize a list of raw record dicts.
|
||||||
|
|
||||||
|
Returns a dict mapping key (``sequence``, ``chosen``, ``responses``,
|
||||||
|
…) to a list of per-record tensors (or nested tensor lists for
|
||||||
|
multi-response keys such as GRPO ``responses``).
|
||||||
|
"""
|
||||||
|
raw: Dict[str, list] = {}
|
||||||
|
doc_sequences: List[List[int]] = []
|
||||||
|
|
||||||
|
for result in iter_raw_records(
|
||||||
|
records, self.mask_builder, self.config, self.tokenizer
|
||||||
|
):
|
||||||
|
primary = None
|
||||||
|
for val in result.values():
|
||||||
|
if isinstance(val, list) and val and isinstance(val[0], int):
|
||||||
|
primary = val
|
||||||
|
break
|
||||||
|
if primary is not None:
|
||||||
|
doc_sequences.append(primary)
|
||||||
|
for key, ids in result.items():
|
||||||
|
raw.setdefault(key, []).append(ids)
|
||||||
|
|
||||||
|
tensors = to_per_record_tensors(raw)
|
||||||
|
|
||||||
|
pos_ids = build_position_ids(doc_sequences, self.position_strategy)
|
||||||
|
if pos_ids is not None:
|
||||||
|
tensors["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
|
||||||
|
|
||||||
|
return tensors
|
||||||
@@ -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
|
|
||||||
|
|||||||
@@ -19,9 +19,8 @@ from astrai.serialization.checkpoint import (
|
|||||||
)
|
)
|
||||||
from astrai.serialization.dataset import (
|
from astrai.serialization.dataset import (
|
||||||
load_bin,
|
load_bin,
|
||||||
load_h5,
|
load_bin_offsets,
|
||||||
save_bin,
|
save_bin,
|
||||||
save_h5,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
@@ -37,7 +36,6 @@ __all__ = [
|
|||||||
"save_safetensors",
|
"save_safetensors",
|
||||||
"save_torch",
|
"save_torch",
|
||||||
"load_bin",
|
"load_bin",
|
||||||
"load_h5",
|
"load_bin_offsets",
|
||||||
"save_bin",
|
"save_bin",
|
||||||
"save_h5",
|
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -2,7 +2,6 @@
|
|||||||
|
|
||||||
import io
|
import io
|
||||||
import json
|
import json
|
||||||
import os
|
|
||||||
import time
|
import time
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
@@ -148,9 +147,6 @@ class Checkpoint:
|
|||||||
save_path = Path(save_dir)
|
save_path = Path(save_dir)
|
||||||
save_path.mkdir(parents=True, exist_ok=True)
|
save_path.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
if get_rank() != 0:
|
|
||||||
return
|
|
||||||
|
|
||||||
meta = {
|
meta = {
|
||||||
"epoch": self.epoch,
|
"epoch": self.epoch,
|
||||||
"consumed_samples": self.consumed_samples,
|
"consumed_samples": self.consumed_samples,
|
||||||
@@ -181,6 +177,7 @@ class Checkpoint:
|
|||||||
epoch=meta.get("epoch", 0),
|
epoch=meta.get("epoch", 0),
|
||||||
consumed_samples=meta.get("consumed_samples", 0),
|
consumed_samples=meta.get("consumed_samples", 0),
|
||||||
extra=extra,
|
extra=extra,
|
||||||
|
meta=meta,
|
||||||
config=config,
|
config=config,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -1,61 +1,51 @@
|
|||||||
"""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 Dict, List
|
|
||||||
|
|
||||||
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]]):
|
def save_bin(
|
||||||
os.makedirs(file_path, exist_ok=True)
|
file_path: str,
|
||||||
full_file_path = os.path.join(file_path, f"{file_name}.h5")
|
tensor_group: Dict[str, List[Tensor]],
|
||||||
with h5py.File(full_file_path, "w") as f:
|
record_keys: Optional[List[str]] = None,
|
||||||
for key, tensors in tensor_group.items():
|
):
|
||||||
grp = f.create_group(key)
|
"""Save tensors as memory-mapped binary files.
|
||||||
for idx, tensor in enumerate(tensors):
|
|
||||||
arr = tensor.cpu().numpy()
|
When *record_keys* is provided, those keys are written with per-record
|
||||||
grp.create_dataset(f"data_{idx}", data=arr)
|
cumulative offsets in ``meta.json`` so that ``MmapStore.fetch_record``
|
||||||
|
can slice individual records from the concatenated binary without
|
||||||
|
cross-record concatenation. Keys not in *record_keys* (e.g. SEQ
|
||||||
def load_h5(file_path: str, share_memory=True) -> Dict[str, List[Tensor]]:
|
``sequence``) are written as a single contiguous stream without
|
||||||
tensor_group: Dict[str, List[Tensor]] = {}
|
offsets, preserving backward compatibility.
|
||||||
|
|
||||||
root_path = Path(file_path)
|
Nested keys (``List[List[Tensor]]`` such as GRPO ``responses``) are
|
||||||
if root_path.is_file() and root_path.suffix in (".h5", ".hdf5"):
|
not supported in bin format — use JSONL for those.
|
||||||
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(file_path: str, tensor_group: Dict[str, List[Tensor]]):
|
|
||||||
os.makedirs(file_path, exist_ok=True)
|
os.makedirs(file_path, exist_ok=True)
|
||||||
|
record_keys = set(record_keys or [])
|
||||||
meta = {}
|
meta = {}
|
||||||
for key, tensors in tensor_group.items():
|
for key, tensors in tensor_group.items():
|
||||||
|
if tensors and isinstance(tensors[0], list):
|
||||||
|
raise ValueError(
|
||||||
|
f"Nested key '{key}' (List[List[Tensor]]) is not supported "
|
||||||
|
f"in bin format. Use JSONL storage instead."
|
||||||
|
)
|
||||||
cat = torch.cat(tensors, dim=0)
|
cat = torch.cat(tensors, dim=0)
|
||||||
meta[key] = {"shape": list(cat.shape), "dtype": str(cat.dtype).split(".")[-1]}
|
entry: Dict[str, Any] = {
|
||||||
|
"shape": list(cat.shape),
|
||||||
|
"dtype": str(cat.dtype).split(".")[-1],
|
||||||
|
}
|
||||||
|
if key in record_keys:
|
||||||
|
offsets = [0]
|
||||||
|
for t in tensors:
|
||||||
|
offsets.append(offsets[-1] + t.shape[0])
|
||||||
|
entry["offsets"] = offsets
|
||||||
|
meta[key] = entry
|
||||||
np.asarray(cat.cpu().numpy()).tofile(os.path.join(file_path, f"{key}.bin"))
|
np.asarray(cat.cpu().numpy()).tofile(os.path.join(file_path, f"{key}.bin"))
|
||||||
with open(os.path.join(file_path, "meta.json"), "w") as f:
|
with open(os.path.join(file_path, "meta.json"), "w") as f:
|
||||||
json.dump(meta, f)
|
json.dump(meta, f)
|
||||||
@@ -69,8 +59,24 @@ def load_bin(file_path: str) -> Dict[str, List[Tensor]]:
|
|||||||
arr = np.memmap(
|
arr = np.memmap(
|
||||||
os.path.join(file_path, f"{key}.bin"),
|
os.path.join(file_path, f"{key}.bin"),
|
||||||
dtype=info["dtype"],
|
dtype=info["dtype"],
|
||||||
mode="r+",
|
mode="c",
|
||||||
shape=tuple(info["shape"]),
|
shape=tuple(info["shape"]),
|
||||||
)
|
)
|
||||||
segments[key] = [torch.from_numpy(arr)]
|
segments[key] = [torch.from_numpy(arr)]
|
||||||
return segments
|
return segments
|
||||||
|
|
||||||
|
|
||||||
|
def load_bin_offsets(file_path: str) -> Dict[str, List[int]]:
|
||||||
|
"""Read per-record cumulative offsets from ``meta.json``.
|
||||||
|
|
||||||
|
Returns an empty dict when no key has offsets (legacy bin files),
|
||||||
|
in which case record-mode access falls back to per-record segment
|
||||||
|
indexing (JSONL layout).
|
||||||
|
"""
|
||||||
|
with open(os.path.join(file_path, "meta.json"), "r") as f:
|
||||||
|
meta = json.load(f)
|
||||||
|
offsets: Dict[str, List[int]] = {}
|
||||||
|
for key, info in meta.items():
|
||||||
|
if "offsets" in info:
|
||||||
|
offsets[key] = info["offsets"]
|
||||||
|
return offsets
|
||||||
|
|||||||
@@ -0,0 +1,53 @@
|
|||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import signal
|
||||||
|
import threading
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_early_stop = threading.Event()
|
||||||
|
_active_context = None
|
||||||
|
|
||||||
|
|
||||||
|
def _early_handler(signum: int, frame):
|
||||||
|
sig = signal.Signals(signum)
|
||||||
|
logger.warning(
|
||||||
|
"Received %s (pid=%d), requesting graceful training stop...",
|
||||||
|
sig.name,
|
||||||
|
os.getpid(),
|
||||||
|
)
|
||||||
|
_early_stop.set()
|
||||||
|
if _active_context is not None:
|
||||||
|
_active_context.request_stop()
|
||||||
|
|
||||||
|
|
||||||
|
def install_early_signal_handlers():
|
||||||
|
for sig in (signal.SIGTERM, signal.SIGINT):
|
||||||
|
signal.signal(sig, _early_handler)
|
||||||
|
_unblock_signals()
|
||||||
|
|
||||||
|
|
||||||
|
def _unblock_signals():
|
||||||
|
try:
|
||||||
|
mask = signal.pthread_sigmask(signal.SIG_BLOCK, set())
|
||||||
|
blocked = {signal.SIGTERM, signal.SIGINT} & mask
|
||||||
|
if blocked:
|
||||||
|
signal.pthread_sigmask(signal.SIG_UNBLOCK, blocked)
|
||||||
|
except (AttributeError, OSError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def register_signal_handlers(context):
|
||||||
|
global _active_context
|
||||||
|
_active_context = context
|
||||||
|
for sig in (signal.SIGTERM, signal.SIGINT):
|
||||||
|
signal.signal(sig, _early_handler)
|
||||||
|
if _early_stop.is_set():
|
||||||
|
context.request_stop()
|
||||||
|
logger.warning("Signal was received during initialization, stopping...")
|
||||||
|
|
||||||
|
|
||||||
|
def unregister_signal_handlers():
|
||||||
|
global _active_context
|
||||||
|
_active_context = None
|
||||||
|
_early_stop.clear()
|
||||||
@@ -1,8 +1,10 @@
|
|||||||
from astrai.tokenize.chat_template import ChatTemplate, MessageType
|
from astrai.tokenize.chat_template import ChatTemplate, MessageType
|
||||||
from astrai.tokenize.tokenizer import AutoTokenizer
|
from astrai.tokenize.tokenizer import AutoTokenizer, Message, Messages
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"AutoTokenizer",
|
"AutoTokenizer",
|
||||||
"ChatTemplate",
|
"ChatTemplate",
|
||||||
"MessageType",
|
"MessageType",
|
||||||
|
"Message",
|
||||||
|
"Messages",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
from functools import cached_property
|
||||||
from typing import Any, Dict, List, Optional
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
from jinja2 import Template
|
from jinja2 import Template
|
||||||
@@ -29,7 +30,34 @@ class ChatTemplate:
|
|||||||
self.description = description
|
self.description = description
|
||||||
self.default_variables = default_variables or {}
|
self.default_variables = default_variables or {}
|
||||||
self.special_tokens = special_tokens or {}
|
self.special_tokens = special_tokens or {}
|
||||||
self._compiled: Template = Template(template_str)
|
|
||||||
|
@cached_property
|
||||||
|
def _compiled(self) -> Template:
|
||||||
|
"""Lazy-compiled Jinja2 template, cached on first access.
|
||||||
|
|
||||||
|
The compiled :class:`~jinja2.Template` holds a dynamically-generated
|
||||||
|
``root`` render function whose ``__module__`` is ``None``; under
|
||||||
|
``pickle`` it falls back to ``__main__`` and breaks ``spawn``-based
|
||||||
|
multiprocessing. :meth:`__getstate__` drops the cached template so
|
||||||
|
that pickle serialises only ``template_str``; each worker rebuilds
|
||||||
|
the cache on first render.
|
||||||
|
"""
|
||||||
|
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(
|
||||||
|
|||||||
@@ -10,12 +10,16 @@ from tokenizers import Tokenizer
|
|||||||
|
|
||||||
from astrai.tokenize.chat_template import ChatTemplate
|
from astrai.tokenize.chat_template import ChatTemplate
|
||||||
|
|
||||||
|
Message = Dict[str, str]
|
||||||
|
"""Single chat message with ``role`` and ``content`` keys."""
|
||||||
|
|
||||||
|
Messages = List[Message]
|
||||||
|
"""Single conversation — a list of messages."""
|
||||||
|
|
||||||
|
|
||||||
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,
|
||||||
@@ -102,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]],
|
||||||
@@ -120,7 +113,16 @@ class AutoTokenizer:
|
|||||||
is_pretokenized: bool = False,
|
is_pretokenized: bool = False,
|
||||||
add_special_tokens: bool = True,
|
add_special_tokens: bool = True,
|
||||||
) -> List:
|
) -> List:
|
||||||
"""Encode text to tokens or token IDs."""
|
"""Encode text to token IDs.
|
||||||
|
|
||||||
|
Accepts both single strings and batches:
|
||||||
|
|
||||||
|
- ``encode("hello")`` → ``[123, 456]``
|
||||||
|
- ``encode(["hello", "world"])`` → ``[[123, 456], [789]]``
|
||||||
|
|
||||||
|
Batches are tokenised in parallel via the Rust backend's
|
||||||
|
``encode_batch`` (uses all available CPU cores).
|
||||||
|
"""
|
||||||
if self._tokenizer is None:
|
if self._tokenizer is None:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
"Tokenizer not initialized. Load or create a tokenizer first."
|
"Tokenizer not initialized. Load or create a tokenizer first."
|
||||||
@@ -133,15 +135,13 @@ class AutoTokenizer:
|
|||||||
add_special_tokens=add_special_tokens,
|
add_special_tokens=add_special_tokens,
|
||||||
)
|
)
|
||||||
return encoded.ids if out_ids else encoded.tokens
|
return encoded.ids if out_ids else encoded.tokens
|
||||||
else:
|
|
||||||
encoded_list = self._tokenizer.encode_batch(
|
encoded_list = self._tokenizer.encode_batch(
|
||||||
tokens,
|
tokens,
|
||||||
is_pretokenized=is_pretokenized,
|
is_pretokenized=is_pretokenized,
|
||||||
add_special_tokens=add_special_tokens,
|
add_special_tokens=add_special_tokens,
|
||||||
)
|
)
|
||||||
return [
|
return [encoded.ids if out_ids else encoded.tokens for encoded in encoded_list]
|
||||||
encoded.ids if out_ids else encoded.tokens for encoded in encoded_list
|
|
||||||
]
|
|
||||||
|
|
||||||
def decode(self, tokens: List[int], skip_special_tokens: bool = True) -> str:
|
def decode(self, tokens: List[int], skip_special_tokens: bool = True) -> str:
|
||||||
"""Decode token IDs to text."""
|
"""Decode token IDs to text."""
|
||||||
@@ -164,7 +164,14 @@ class AutoTokenizer:
|
|||||||
- tokenizer.bos_token → returns string
|
- tokenizer.bos_token → returns string
|
||||||
- tokenizer.bos_token_id → returns corresponding integer ID
|
- tokenizer.bos_token_id → returns corresponding integer ID
|
||||||
- tokenizer.stop_ids → returns list of corresponding integer IDs for all special tokens
|
- tokenizer.stop_ids → returns list of corresponding integer IDs for all special tokens
|
||||||
|
|
||||||
|
Internal/private attrs are not intercepted: during unpickle
|
||||||
|
``__dict__`` is empty, so probing ``self._special_token_map``
|
||||||
|
would recurse infinitely.
|
||||||
"""
|
"""
|
||||||
|
if key.startswith("_"):
|
||||||
|
raise AttributeError(key)
|
||||||
|
|
||||||
# Handle stop_ids - return IDs for all special tokens
|
# Handle stop_ids - return IDs for all special tokens
|
||||||
if key == "stop_ids":
|
if key == "stop_ids":
|
||||||
stop_ids = []
|
stop_ids = []
|
||||||
@@ -220,45 +227,63 @@ class AutoTokenizer:
|
|||||||
|
|
||||||
def apply_chat_template(
|
def apply_chat_template(
|
||||||
self,
|
self,
|
||||||
messages: List[Dict[str, str]],
|
messages: Union[Messages, List[Messages]],
|
||||||
system_prompt: Optional[str] = None,
|
system_prompt: Optional[str] = None,
|
||||||
tokenize: bool = True,
|
tokenize: bool = True,
|
||||||
add_generation_prompt: bool = True,
|
add_generation_prompt: bool = True,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> Union[str, List[int]]:
|
) -> Union[str, List[int], List[str], List[List[int]]]:
|
||||||
"""
|
"""Apply the chat template and optionally tokenize.
|
||||||
Apply the chat template to messages and optionally tokenize the result.
|
|
||||||
|
Accepts both single conversations and batches:
|
||||||
|
|
||||||
|
- ``apply_chat_template([msg1, msg2])`` → ``"..."`` or ``[ids]``
|
||||||
|
- ``apply_chat_template([[msg1, msg2], [msg3]])`` → ``["..", ".."]``
|
||||||
|
or ``[[ids], [ids]]``
|
||||||
|
|
||||||
|
Batches render each conversation list and tokenise all at once via
|
||||||
|
:meth:`encode` (``List[str]`` → Rust parallel ``encode_batch``).
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
messages: List of message dicts with 'role' and 'content'.
|
messages: Single conversation (``Messages``) or batch of
|
||||||
system_prompt: Optional system prompt string (auto-converted to first message).
|
conversations (``BatchMessages``).
|
||||||
|
system_prompt: Optional system prompt prepended (single mode only).
|
||||||
tokenize: Whether to return token IDs (True) or raw string (False).
|
tokenize: Whether to return token IDs (True) or raw string (False).
|
||||||
add_generation_prompt: Whether to add the generation prompt (default: True).
|
add_generation_prompt: Whether to add the generation prompt.
|
||||||
**kwargs: Additional variables to pass to the template.
|
**kwargs: Additional template variables.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Either the rendered string or list of token IDs.
|
Single mode: ``str`` or ``List[int]``.
|
||||||
|
Batch mode: ``List[str]`` or ``List[List[int]]``.
|
||||||
Raises:
|
|
||||||
RuntimeError: If chat template is not set.
|
|
||||||
"""
|
"""
|
||||||
if self._chat_template is None:
|
if self._chat_template is None:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
"Chat template not set. Use set_chat_template() to set a template first."
|
"Chat template not set. Use set_chat_template() to set a template first."
|
||||||
)
|
)
|
||||||
|
|
||||||
# Auto-convert system_prompt to first message if provided
|
is_batch = bool(messages) and isinstance(messages[0], list)
|
||||||
|
|
||||||
|
if is_batch:
|
||||||
|
rendered = [
|
||||||
|
self._chat_template.render(
|
||||||
|
messages=msgs,
|
||||||
|
add_generation_prompt=add_generation_prompt,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
for msgs in messages
|
||||||
|
]
|
||||||
|
if tokenize:
|
||||||
|
return self.encode(rendered) # List[str] → batch encode
|
||||||
|
return rendered
|
||||||
|
|
||||||
|
# Single conversation
|
||||||
if system_prompt:
|
if system_prompt:
|
||||||
messages = [{"role": "system", "content": system_prompt}] + list(messages)
|
messages = [{"role": "system", "content": system_prompt}] + list(messages)
|
||||||
|
|
||||||
# Render the template
|
|
||||||
rendered = self._chat_template.render(
|
rendered = self._chat_template.render(
|
||||||
messages=messages,
|
messages=messages,
|
||||||
add_generation_prompt=add_generation_prompt,
|
add_generation_prompt=add_generation_prompt,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
if tokenize:
|
if tokenize:
|
||||||
return self.encode(rendered)
|
return self.encode(rendered)
|
||||||
|
|
||||||
return rendered
|
return rendered
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -0,0 +1,421 @@
|
|||||||
|
"""Online rollout runner for RL training.
|
||||||
|
|
||||||
|
Provides:
|
||||||
|
- :class:`RawRollout` — generation output container (no reward yet)
|
||||||
|
- :class:`RolloutResult` — a :class:`RawRollout` with rewards attached
|
||||||
|
- :class:`BaseRewardModel` — pluggable reward interface
|
||||||
|
- :class:`RolloutGenerator` — KV-cache-backed generation of grouped
|
||||||
|
responses + decoding (no reward); delegates the generation loop to
|
||||||
|
:class:`~astrai.inference.core.scheduler.InferenceScheduler.run_batch`
|
||||||
|
so rollout and the production inference server share one code path
|
||||||
|
- :class:`RolloutRunner` — orchestrates generation + scoring with a
|
||||||
|
step-driven cache; its ``__call__`` returns ``(RolloutResult, is_fresh)``
|
||||||
|
so callers do not need to rely on object identity to detect refreshes.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Dict, List, Optional, Tuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.inference.core.scheduler import InferenceScheduler
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(kw_only=True)
|
||||||
|
class RawRollout:
|
||||||
|
"""Generation output before reward scoring.
|
||||||
|
|
||||||
|
Produced by :class:`RolloutGenerator`; consumed by :class:`RolloutRunner`
|
||||||
|
to assemble a :class:`RolloutResult` once rewards are attached.
|
||||||
|
|
||||||
|
Fields are designed to cover all common RL algorithms:
|
||||||
|
GRPO, PPO, Online DPO, Rejection Sampling, etc.
|
||||||
|
|
||||||
|
Fields:
|
||||||
|
prompts: Tokenized prompts, shape ``[B, P_len]``.
|
||||||
|
prompt_mask: Boolean mask for real prompt tokens, shape ``[B, P_len]``.
|
||||||
|
responses: Generated response token IDs, shape ``[B, G, R_max]``.
|
||||||
|
response_mask: Boolean mask for real (non-pad) response tokens,
|
||||||
|
shape ``[B, G, R_max]``.
|
||||||
|
logprobs_old: Per-token log-probs under the behaviour policy,
|
||||||
|
shape ``[B, G, R_max]``.
|
||||||
|
prompt_texts: Decoded prompt strings (for reward models that
|
||||||
|
need text).
|
||||||
|
response_texts: Decoded response strings, shape ``[B, G]``
|
||||||
|
(for reward models).
|
||||||
|
"""
|
||||||
|
|
||||||
|
prompts: Tensor
|
||||||
|
prompt_mask: Tensor
|
||||||
|
responses: Tensor
|
||||||
|
response_mask: Tensor
|
||||||
|
logprobs_old: Tensor
|
||||||
|
prompt_texts: List[str] = field(default_factory=list)
|
||||||
|
response_texts: List[List[str]] = field(default_factory=list)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(kw_only=True)
|
||||||
|
class RolloutResult(RawRollout):
|
||||||
|
"""A :class:`RawRollout` with reward scoring attached.
|
||||||
|
|
||||||
|
Produced by :class:`RolloutRunner` once the :class:`BaseRewardModel`
|
||||||
|
has scored the decoded responses.
|
||||||
|
|
||||||
|
Fields:
|
||||||
|
rewards: Reward per response, shape ``[B, G]``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
rewards: Tensor
|
||||||
|
|
||||||
|
|
||||||
|
class BaseRewardModel(ABC):
|
||||||
|
"""Pluggable reward model interface.
|
||||||
|
|
||||||
|
Subclasses should implement ``score()`` to return a ``[B, G]`` float
|
||||||
|
tensor of rewards. Implementations can be:
|
||||||
|
* A loaded reward model (e.g. ArmoRM, Skywork-Reward)
|
||||||
|
* An external API call
|
||||||
|
* A rule-based function (format, length, keyword matching)
|
||||||
|
"""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def score(self, prompts: List[str], responses: List[List[str]]) -> Tensor:
|
||||||
|
"""Score each generated response.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
prompts: Raw prompt strings, length ``B``.
|
||||||
|
responses: Generated response strings, shape ``[B, G]``.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Float tensor of shape ``[B, G]``.
|
||||||
|
"""
|
||||||
|
...
|
||||||
|
|
||||||
|
|
||||||
|
_PAD = 0
|
||||||
|
|
||||||
|
|
||||||
|
class RolloutGenerator:
|
||||||
|
"""Pure generation + decoding for a group of responses per prompt.
|
||||||
|
|
||||||
|
Delegates the prefill/decode loop to
|
||||||
|
:meth:`~astrai.inference.core.scheduler.InferenceScheduler.run_batch`,
|
||||||
|
which uses a real KV cache (no O(n²) recompute). Has no dependency
|
||||||
|
on any reward model; can be reused in isolation for offline
|
||||||
|
generation, qualitative sampling, or eval pipelines.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
scheduler: InferenceScheduler,
|
||||||
|
tokenizer,
|
||||||
|
max_tokens: int = 1024,
|
||||||
|
group_size: int = 8,
|
||||||
|
temperature: float = 1.0,
|
||||||
|
top_k: int = 0,
|
||||||
|
top_p: float = 1.0,
|
||||||
|
frequency_penalty: float = 0.0,
|
||||||
|
rep_window: int = 64,
|
||||||
|
):
|
||||||
|
self.scheduler = scheduler
|
||||||
|
self.tokenizer = tokenizer
|
||||||
|
self.max_tokens = max_tokens
|
||||||
|
self.group_size = group_size
|
||||||
|
self.temperature = temperature
|
||||||
|
self.top_k = top_k
|
||||||
|
self.top_p = top_p
|
||||||
|
self.frequency_penalty = frequency_penalty
|
||||||
|
self.rep_window = rep_window
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def generate(self, batch: Dict) -> RawRollout:
|
||||||
|
"""Expand prompts by ``group_size`` and generate one response each.
|
||||||
|
|
||||||
|
Accepted batch formats (per sample, repeated B times):
|
||||||
|
|
||||||
|
- **messages**: ``{"messages": [{"role": "user", "content": "..."}, ...]}``
|
||||||
|
- **instruction + input + output**: ``{"instruction": "...",
|
||||||
|
"input": "...", "output": "..."}`` — mapped to ``system`` /
|
||||||
|
``user`` / ``assistant`` messages; ``input`` and ``output``
|
||||||
|
are optional and skipped when empty.
|
||||||
|
|
||||||
|
Both are rendered through the tokenizer's chat template with
|
||||||
|
``add_generation_prompt=True`` so rollout prompts match the
|
||||||
|
format the policy was SFT-trained on.
|
||||||
|
"""
|
||||||
|
model = self.scheduler._executor.model
|
||||||
|
was_training = model.training
|
||||||
|
model.eval()
|
||||||
|
try:
|
||||||
|
return self._generate_eval(batch)
|
||||||
|
finally:
|
||||||
|
model.train(was_training)
|
||||||
|
|
||||||
|
def _generate_eval(self, batch: Dict) -> RawRollout:
|
||||||
|
prompt_texts, flat_prompt_ids = self._prepare_prompts(batch)
|
||||||
|
B = len(prompt_texts)
|
||||||
|
G = self.group_size
|
||||||
|
# Re-expand flat list to G copies per prompt for run_batch.
|
||||||
|
expanded_prompt_ids: List[List[int]] = []
|
||||||
|
for ids in flat_prompt_ids:
|
||||||
|
expanded_prompt_ids.extend([list(ids)] * G)
|
||||||
|
|
||||||
|
results = self.scheduler.run_batch(
|
||||||
|
expanded_prompt_ids,
|
||||||
|
max_tokens=self.max_tokens,
|
||||||
|
temperature=self.temperature,
|
||||||
|
top_k=self.top_k,
|
||||||
|
top_p=self.top_p,
|
||||||
|
frequency_penalty=self.frequency_penalty,
|
||||||
|
rep_window=self.rep_window,
|
||||||
|
return_logprobs=True,
|
||||||
|
)
|
||||||
|
if len(results) != B * G:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Rollout scheduler returned {len(results)} results, expected {B * G}"
|
||||||
|
)
|
||||||
|
for token_ids, logprobs in results:
|
||||||
|
if len(token_ids) != len(logprobs):
|
||||||
|
raise RuntimeError(
|
||||||
|
"Rollout scheduler returned misaligned token IDs and logprobs"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Each element is (token_ids, logprobs); pad to max length.
|
||||||
|
max_len = 0
|
||||||
|
for token_ids, _lp in results:
|
||||||
|
max_len = max(max_len, len(token_ids))
|
||||||
|
max_len = max(max_len, 1)
|
||||||
|
|
||||||
|
device = self.scheduler.device
|
||||||
|
P_len = max(len(ids) for ids in flat_prompt_ids)
|
||||||
|
prompts_tensor = torch.zeros(B, P_len, dtype=torch.long, device=device)
|
||||||
|
prompt_mask = torch.zeros(B, P_len, dtype=torch.bool, device=device)
|
||||||
|
for i, ids in enumerate(flat_prompt_ids):
|
||||||
|
prompts_tensor[i, -len(ids) :] = torch.tensor(
|
||||||
|
ids, dtype=torch.long, device=device
|
||||||
|
)
|
||||||
|
prompt_mask[i, -len(ids) :] = True
|
||||||
|
|
||||||
|
responses = torch.full((B, G, max_len), _PAD, dtype=torch.long, device=device)
|
||||||
|
response_mask = torch.zeros((B, G, max_len), dtype=torch.bool, device=device)
|
||||||
|
logprobs_old = torch.zeros((B, G, max_len), dtype=torch.float, device=device)
|
||||||
|
|
||||||
|
flat_idx = 0
|
||||||
|
response_texts: List[List[str]] = [[] for _ in range(B)]
|
||||||
|
for i in range(B):
|
||||||
|
for g in range(G):
|
||||||
|
token_ids, lps = results[flat_idx]
|
||||||
|
flat_idx += 1
|
||||||
|
n = len(token_ids)
|
||||||
|
if n:
|
||||||
|
responses[i, g, :n] = torch.tensor(
|
||||||
|
token_ids, dtype=torch.long, device=device
|
||||||
|
)
|
||||||
|
response_mask[i, g, :n] = True
|
||||||
|
logprobs_old[i, g, :n] = torch.tensor(
|
||||||
|
lps, dtype=torch.float, device=device
|
||||||
|
)
|
||||||
|
response_texts[i].append(
|
||||||
|
self.tokenizer.decode(token_ids, skip_special_tokens=True)
|
||||||
|
)
|
||||||
|
|
||||||
|
return RawRollout(
|
||||||
|
prompts=prompts_tensor,
|
||||||
|
prompt_mask=prompt_mask,
|
||||||
|
responses=responses,
|
||||||
|
response_mask=response_mask,
|
||||||
|
logprobs_old=logprobs_old,
|
||||||
|
prompt_texts=prompt_texts,
|
||||||
|
response_texts=response_texts,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _prepare_prompts(self, batch: Dict) -> Tuple[List[str], List[List[int]]]:
|
||||||
|
"""Render batch prompts to ``(texts, token_id_lists)``.
|
||||||
|
|
||||||
|
Returns two parallel lists of length B (number of prompts in
|
||||||
|
the batch). Dispatches by batch keys:
|
||||||
|
|
||||||
|
- ``"messages"``: treated as a pre-built message list per sample.
|
||||||
|
- ``"instruction"`` (optionally ``"input"`` and ``"output"``): mapped
|
||||||
|
to ``system`` / ``user`` / ``assistant`` messages respectively.
|
||||||
|
|
||||||
|
Both paths go through the tokenizer's chat template with
|
||||||
|
``add_generation_prompt=True``.
|
||||||
|
"""
|
||||||
|
if "messages" in batch:
|
||||||
|
messages_list = batch["messages"]
|
||||||
|
elif "instruction" in batch:
|
||||||
|
instructions = batch["instruction"]
|
||||||
|
B = len(instructions)
|
||||||
|
inputs = batch.get("input") or [""] * B
|
||||||
|
outputs = batch.get("output") or [""] * B
|
||||||
|
messages_list = [
|
||||||
|
self._instruction_to_messages(i, u, o)
|
||||||
|
for i, u, o in zip(instructions, inputs, outputs)
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
raise ValueError(
|
||||||
|
"Rollout batch must contain either 'messages' or "
|
||||||
|
"'instruction' (optionally 'input'/'output'); got keys: "
|
||||||
|
f"{list(batch.keys())}"
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
prompt_texts = self.tokenizer.apply_chat_template(
|
||||||
|
messages_list, tokenize=False, add_generation_prompt=True
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
not isinstance(prompt_texts, list)
|
||||||
|
or len(prompt_texts) != len(messages_list)
|
||||||
|
or not all(isinstance(text, str) for text in prompt_texts)
|
||||||
|
):
|
||||||
|
raise TypeError("Tokenizer does not support batched chat templates")
|
||||||
|
flat_prompt_ids = self.tokenizer.encode(prompt_texts)
|
||||||
|
if len(flat_prompt_ids) != len(messages_list) or not all(
|
||||||
|
isinstance(ids, list) for ids in flat_prompt_ids
|
||||||
|
):
|
||||||
|
raise TypeError("Tokenizer does not support batched encoding")
|
||||||
|
except (TypeError, IndexError, KeyError):
|
||||||
|
# Keep compatibility with lightweight tokenizer adapters that only
|
||||||
|
# implement the single-conversation template API.
|
||||||
|
prompt_texts = []
|
||||||
|
flat_prompt_ids = []
|
||||||
|
for messages in messages_list:
|
||||||
|
text = self.tokenizer.apply_chat_template(
|
||||||
|
messages, tokenize=False, add_generation_prompt=True
|
||||||
|
)
|
||||||
|
ids = self.tokenizer.apply_chat_template(
|
||||||
|
messages, tokenize=True, add_generation_prompt=True
|
||||||
|
)
|
||||||
|
prompt_texts.append(text)
|
||||||
|
flat_prompt_ids.append(list(ids))
|
||||||
|
return prompt_texts, flat_prompt_ids
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _instruction_to_messages(
|
||||||
|
instruction: str, inp: str = "", output: str = ""
|
||||||
|
) -> List[Dict[str, str]]:
|
||||||
|
"""Map instruction/input/output to chat messages.
|
||||||
|
|
||||||
|
Role mapping follows the convention used throughout the
|
||||||
|
preprocessing pipeline: ``instruction`` → system, ``input`` →
|
||||||
|
user, ``output`` → assistant. Empty fields are skipped so a
|
||||||
|
bare instruction produces a ``[system]`` list and the chat
|
||||||
|
template's ``add_generation_prompt`` adds the assistant header
|
||||||
|
for sampling.
|
||||||
|
"""
|
||||||
|
messages: List[Dict[str, str]] = []
|
||||||
|
if instruction:
|
||||||
|
messages.append({"role": "system", "content": instruction})
|
||||||
|
if inp:
|
||||||
|
messages.append({"role": "user", "content": inp})
|
||||||
|
if output:
|
||||||
|
messages.append({"role": "assistant", "content": output})
|
||||||
|
return messages
|
||||||
|
|
||||||
|
|
||||||
|
class RolloutRunner:
|
||||||
|
"""Produces :class:`RolloutResult` from a prompt batch.
|
||||||
|
|
||||||
|
Composes a :class:`RolloutGenerator` (generation + decoding) with a
|
||||||
|
:class:`BaseRewardModel` (scoring). Maintains an internal cache so
|
||||||
|
the same batch prompt can be replayed for multiple gradient steps.
|
||||||
|
A new rollout is triggered every ``rollout_interval`` calls to
|
||||||
|
:meth:`step` (or after :meth:`clear_cache`).
|
||||||
|
|
||||||
|
The ``__call__`` contract returns a ``(RolloutResult, is_fresh)``
|
||||||
|
tuple — callers must use the boolean to detect a refreshed rollout
|
||||||
|
rather than relying on object identity.
|
||||||
|
|
||||||
|
Usage::
|
||||||
|
|
||||||
|
generator = RolloutGenerator(policy, tokenizer, pipeline, ...)
|
||||||
|
runner = RolloutRunner(generator, reward_model, rollout_interval=512)
|
||||||
|
result, is_fresh = runner(prompt_batch)
|
||||||
|
if is_fresh:
|
||||||
|
... # e.g. sync behaviour policy
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
generator: RolloutGenerator,
|
||||||
|
reward_model: BaseRewardModel,
|
||||||
|
rollout_interval: int = 512,
|
||||||
|
):
|
||||||
|
self.generator = generator
|
||||||
|
self.reward_model = reward_model
|
||||||
|
self.rollout_interval = rollout_interval
|
||||||
|
|
||||||
|
self._cache: Optional[RolloutResult] = None
|
||||||
|
self._cache_key = None
|
||||||
|
self._steps_since_rollout: int = 0
|
||||||
|
|
||||||
|
def step(self):
|
||||||
|
"""Advance the internal counter (call once per optimizer step)."""
|
||||||
|
self._steps_since_rollout += 1
|
||||||
|
|
||||||
|
def clear_cache(self):
|
||||||
|
"""Force next call to re-run rollout."""
|
||||||
|
self._cache = None
|
||||||
|
self._cache_key = None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _batch_key(batch: Dict):
|
||||||
|
"""Build a stable key for the prompt fields accepted by the generator."""
|
||||||
|
|
||||||
|
def freeze(value):
|
||||||
|
if isinstance(value, dict):
|
||||||
|
return tuple(sorted((key, freeze(val)) for key, val in value.items()))
|
||||||
|
if isinstance(value, (list, tuple)):
|
||||||
|
return tuple(freeze(item) for item in value)
|
||||||
|
return value
|
||||||
|
|
||||||
|
fields = ("messages", "instruction", "input", "output")
|
||||||
|
return tuple(
|
||||||
|
(field, freeze(batch[field])) for field in fields if field in batch
|
||||||
|
)
|
||||||
|
|
||||||
|
def _score(self, raw: RawRollout) -> RolloutResult:
|
||||||
|
rewards = self.reward_model.score(raw.prompt_texts, raw.response_texts)
|
||||||
|
if not isinstance(rewards, Tensor):
|
||||||
|
rewards = torch.as_tensor(rewards, dtype=torch.float32)
|
||||||
|
expected_shape = raw.responses.shape[:2]
|
||||||
|
if rewards.shape != expected_shape:
|
||||||
|
raise ValueError(
|
||||||
|
f"Reward model returned shape {tuple(rewards.shape)}, "
|
||||||
|
f"expected {tuple(expected_shape)}"
|
||||||
|
)
|
||||||
|
if not torch.isfinite(rewards).all():
|
||||||
|
raise ValueError("Reward model returned non-finite values")
|
||||||
|
device = raw.prompts.device
|
||||||
|
return RolloutResult(
|
||||||
|
prompts=raw.prompts,
|
||||||
|
prompt_mask=raw.prompt_mask,
|
||||||
|
responses=raw.responses,
|
||||||
|
response_mask=raw.response_mask,
|
||||||
|
rewards=rewards.to(device=device),
|
||||||
|
logprobs_old=raw.logprobs_old,
|
||||||
|
prompt_texts=raw.prompt_texts,
|
||||||
|
response_texts=raw.response_texts,
|
||||||
|
)
|
||||||
|
|
||||||
|
def __call__(self, batch: Dict[str, Tensor]) -> Tuple[RolloutResult, bool]:
|
||||||
|
"""Return ``(cached or fresh) RolloutResult`` plus an ``is_fresh`` flag.
|
||||||
|
|
||||||
|
Triggers a new rollout when ``_steps_since_rollout >= rollout_interval``
|
||||||
|
or when the cache is empty.
|
||||||
|
"""
|
||||||
|
cache_key = self._batch_key(batch)
|
||||||
|
if (
|
||||||
|
self._cache is None
|
||||||
|
or cache_key != self._cache_key
|
||||||
|
or self._steps_since_rollout >= self.rollout_interval
|
||||||
|
):
|
||||||
|
raw = self.generator.generate(batch)
|
||||||
|
self._cache = self._score(raw)
|
||||||
|
self._cache_key = cache_key
|
||||||
|
self._steps_since_rollout = 0
|
||||||
|
return self._cache, True
|
||||||
|
return self._cache, False
|
||||||
+225
-62
@@ -9,17 +9,8 @@ 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
|
||||||
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]:
|
||||||
@@ -28,9 +19,10 @@ def move_to_device(batch: Dict[str, Tensor], device: str) -> Dict[str, Tensor]:
|
|||||||
|
|
||||||
|
|
||||||
def get_logprobs(
|
def get_logprobs(
|
||||||
model: Union[nn.Module, Callable[..., Dict[str, Tensor]]],
|
model: nn.Module,
|
||||||
input_ids: Tensor,
|
input_ids: Tensor,
|
||||||
mask: Tensor,
|
attn_mask: Tensor,
|
||||||
|
loss_mask: Tensor,
|
||||||
reduction: str,
|
reduction: str,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
"""Compute token-wise log probabilities from model outputs.
|
"""Compute token-wise log probabilities from model outputs.
|
||||||
@@ -38,7 +30,8 @@ def get_logprobs(
|
|||||||
Args:
|
Args:
|
||||||
model: The language model
|
model: The language model
|
||||||
input_ids: Input token IDs of shape [batch_size, seq_len]
|
input_ids: Input token IDs of shape [batch_size, seq_len]
|
||||||
mask: Attention mask of shape [batch_size, seq_len]
|
attn_mask: Attention mask passed to the model (may include causal).
|
||||||
|
loss_mask: Per-token mask for loss reduction.
|
||||||
reduction: How to reduce over sequence dimension ("mean", "sum", "none")
|
reduction: How to reduce over sequence dimension ("mean", "sum", "none")
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
@@ -51,9 +44,12 @@ def get_logprobs(
|
|||||||
)
|
)
|
||||||
|
|
||||||
shifted_input_ids = input_ids[:, 1:]
|
shifted_input_ids = input_ids[:, 1:]
|
||||||
shifted_mask = mask[:, 1:]
|
shifted_loss_mask = loss_mask[:, 1:]
|
||||||
|
|
||||||
logits = model(input_ids[:, :-1], mask[:, :-1])["logits"]
|
logits = model(
|
||||||
|
input_ids[:, :-1],
|
||||||
|
attn_mask[:, :, :-1, :-1] if attn_mask.dim() == 4 else attn_mask[:, :-1],
|
||||||
|
)["logits"]
|
||||||
log_probs = torch.log_softmax(logits.float(), dim=-1)
|
log_probs = torch.log_softmax(logits.float(), dim=-1)
|
||||||
|
|
||||||
token_logprobs = torch.gather(
|
token_logprobs = torch.gather(
|
||||||
@@ -61,13 +57,13 @@ def get_logprobs(
|
|||||||
).squeeze(-1)
|
).squeeze(-1)
|
||||||
|
|
||||||
if reduction == "mean":
|
if reduction == "mean":
|
||||||
return (token_logprobs * shifted_mask).sum(dim=-1) / shifted_mask.sum(
|
return (token_logprobs * shifted_loss_mask).sum(dim=-1) / shifted_loss_mask.sum(
|
||||||
dim=-1
|
dim=-1
|
||||||
).clamp(min=1.0)
|
).clamp(min=1.0)
|
||||||
elif reduction == "sum":
|
elif reduction == "sum":
|
||||||
return (token_logprobs * shifted_mask).sum(dim=-1)
|
return (token_logprobs * shifted_loss_mask).sum(dim=-1)
|
||||||
else:
|
else:
|
||||||
return token_logprobs * shifted_mask
|
return token_logprobs * shifted_loss_mask
|
||||||
|
|
||||||
|
|
||||||
def make_doc_boundary_mask(position_ids: Tensor) -> Tensor:
|
def make_doc_boundary_mask(position_ids: Tensor) -> Tensor:
|
||||||
@@ -87,7 +83,15 @@ def make_doc_boundary_mask(position_ids: Tensor) -> Tensor:
|
|||||||
|
|
||||||
|
|
||||||
class BaseStrategy(ABC):
|
class BaseStrategy(ABC):
|
||||||
"""Abstract base class for training strategies."""
|
"""Abstract base class for training strategies.
|
||||||
|
|
||||||
|
When a :class:`~astrai.trainer.rollout.RolloutRunner` is injected via
|
||||||
|
:meth:`set_rollout_runner`, the strategy transparently switches to
|
||||||
|
online mode: each ``__call__`` produces a :class:`RolloutResult`,
|
||||||
|
converts it to a training batch via :meth:`prepare_from_rollout`, and
|
||||||
|
then computes the loss. Without a runner the strategy runs in
|
||||||
|
offline mode and consumes the batch directly.
|
||||||
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -98,8 +102,8 @@ class BaseStrategy(ABC):
|
|||||||
self.model = model
|
self.model = model
|
||||||
self.device = device
|
self.device = device
|
||||||
self.executor = kwargs.pop("executor", None)
|
self.executor = kwargs.pop("executor", None)
|
||||||
self.model_fn = kwargs.pop("model_fn", None)
|
|
||||||
self.extra_kwargs = kwargs
|
self.extra_kwargs = kwargs
|
||||||
|
self._rollout_runner = None
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||||
@@ -113,9 +117,53 @@ class BaseStrategy(ABC):
|
|||||||
"""
|
"""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def supports_online(self) -> bool:
|
||||||
|
"""Whether this strategy can operate with a rollout runner.
|
||||||
|
|
||||||
|
Base implementation returns ``False``; strategies that implement
|
||||||
|
:meth:`prepare_from_rollout` should override to return ``True``.
|
||||||
|
"""
|
||||||
|
return False
|
||||||
|
|
||||||
|
def set_rollout_runner(self, runner):
|
||||||
|
"""Inject a :class:`RolloutRunner` to enable online rollout mode."""
|
||||||
|
self._rollout_runner = runner
|
||||||
|
|
||||||
|
def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]:
|
||||||
|
"""Map a :class:`RolloutResult` to the batch layout expected by
|
||||||
|
:meth:`compute_loss`.
|
||||||
|
|
||||||
|
Strategies that return ``True`` from :meth:`supports_online` must
|
||||||
|
override this. Default raises :class:`NotImplementedError`.
|
||||||
|
"""
|
||||||
|
raise NotImplementedError(
|
||||||
|
f"{type(self).__name__} does not support online rollout"
|
||||||
|
)
|
||||||
|
|
||||||
|
def _on_rollout_refresh(self):
|
||||||
|
"""Hook fired when a fresh rollout result is produced.
|
||||||
|
|
||||||
|
Override to refresh stale state (e.g. syncing the behaviour
|
||||||
|
policy). Default is a no-op.
|
||||||
|
"""
|
||||||
|
pass
|
||||||
|
|
||||||
|
def on_optimizer_step(self):
|
||||||
|
"""Advance online rollout state after a successful optimizer step."""
|
||||||
|
if self._rollout_runner is not None:
|
||||||
|
self._rollout_runner.step()
|
||||||
|
|
||||||
def __call__(self, batch: Dict[str, Tensor]) -> Tensor:
|
def __call__(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||||
"""Allow calling strategy directly as a callable."""
|
"""Run offline or online forward depending on runner injection."""
|
||||||
return self.compute_loss(batch)
|
if self._rollout_runner is None:
|
||||||
|
return self.compute_loss(batch)
|
||||||
|
|
||||||
|
result, is_fresh = self._rollout_runner(batch)
|
||||||
|
if is_fresh:
|
||||||
|
self._on_rollout_refresh()
|
||||||
|
|
||||||
|
train_batch = self.prepare_from_rollout(result)
|
||||||
|
return self.compute_loss(train_batch)
|
||||||
|
|
||||||
|
|
||||||
class StrategyFactory(BaseFactory["BaseStrategy"]):
|
class StrategyFactory(BaseFactory["BaseStrategy"]):
|
||||||
@@ -223,14 +271,13 @@ class DPOStrategy(BaseStrategy):
|
|||||||
self,
|
self,
|
||||||
model: nn.Module,
|
model: nn.Module,
|
||||||
device: str,
|
device: str,
|
||||||
|
ref_model: nn.Module,
|
||||||
beta: float = 0.1,
|
beta: float = 0.1,
|
||||||
reduction: str = "mean",
|
reduction: str = "sum",
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
super().__init__(model, device, **kwargs)
|
super().__init__(model, device, **kwargs)
|
||||||
self.ref_model = create_ref_model(
|
self.ref_model = ref_model
|
||||||
self.model_fn, self.executor.unwrap_model(model)
|
|
||||||
).to(device=self.device)
|
|
||||||
self.beta = beta
|
self.beta = beta
|
||||||
self.reduction = reduction
|
self.reduction = reduction
|
||||||
|
|
||||||
@@ -240,13 +287,31 @@ class DPOStrategy(BaseStrategy):
|
|||||||
chosen_mask, rejected_mask = batch["chosen_mask"], batch["rejected_mask"]
|
chosen_mask, rejected_mask = batch["chosen_mask"], batch["rejected_mask"]
|
||||||
|
|
||||||
concat_ids = torch.cat([chosen_ids, rejected_ids], dim=0)
|
concat_ids = torch.cat([chosen_ids, rejected_ids], dim=0)
|
||||||
concat_mask = torch.cat([chosen_mask, rejected_mask], dim=0)
|
concat_loss_mask = torch.cat([chosen_mask, rejected_mask], dim=0)
|
||||||
|
|
||||||
log_pi = get_logprobs(self.model, concat_ids, concat_mask, self.reduction)
|
# Build full attention mask: key-padding + causal
|
||||||
|
key_pad = concat_ids.bool()[:, None, None, :] # [B*2, 1, 1, S]
|
||||||
|
S = key_pad.shape[-1]
|
||||||
|
causal = torch.tril(
|
||||||
|
torch.ones(S, S, dtype=torch.bool, device=concat_ids.device)
|
||||||
|
)[None, None, :, :] # [1, 1, S, S]
|
||||||
|
full_mask = key_pad & causal # [B*2, 1, S, S] — composed
|
||||||
|
|
||||||
|
log_pi = get_logprobs(
|
||||||
|
self.model,
|
||||||
|
concat_ids,
|
||||||
|
full_mask,
|
||||||
|
concat_loss_mask,
|
||||||
|
self.reduction,
|
||||||
|
)
|
||||||
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
log_ref = get_logprobs(
|
log_ref = get_logprobs(
|
||||||
self.ref_model, concat_ids, concat_mask, self.reduction
|
self.ref_model,
|
||||||
|
concat_ids,
|
||||||
|
full_mask,
|
||||||
|
concat_loss_mask,
|
||||||
|
self.reduction,
|
||||||
)
|
)
|
||||||
|
|
||||||
log_pi_chosen = log_pi[: chosen_ids.shape[0]]
|
log_pi_chosen = log_pi[: chosen_ids.shape[0]]
|
||||||
@@ -262,47 +327,77 @@ class DPOStrategy(BaseStrategy):
|
|||||||
|
|
||||||
return dpo_loss
|
return dpo_loss
|
||||||
|
|
||||||
|
def supports_online(self) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]:
|
||||||
|
"""Pick best/worst response per prompt by reward as chosen/rejected."""
|
||||||
|
rewards = result.rewards
|
||||||
|
responses = result.responses
|
||||||
|
masks = result.response_mask
|
||||||
|
best = rewards.argmax(dim=-1)
|
||||||
|
worst = rewards.argmin(dim=-1)
|
||||||
|
B = responses.shape[0]
|
||||||
|
idx = torch.arange(B, device=responses.device)
|
||||||
|
chosen = responses[idx, best]
|
||||||
|
chosen_mask = masks[idx, best].float()
|
||||||
|
rejected = responses[idx, worst]
|
||||||
|
rejected_mask = masks[idx, worst].float()
|
||||||
|
return {
|
||||||
|
"chosen": chosen,
|
||||||
|
"chosen_mask": chosen_mask,
|
||||||
|
"rejected": rejected,
|
||||||
|
"rejected_mask": rejected_mask,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
@StrategyFactory.register("grpo")
|
@StrategyFactory.register("grpo")
|
||||||
class GRPOStrategy(BaseStrategy):
|
class GRPOStrategy(BaseStrategy):
|
||||||
"""Group Relative Policy Optimization strategy.
|
"""Group Relative Policy Optimization strategy.
|
||||||
|
|
||||||
On-policy GRPO following DeepSeek-R1: the policy model is updated while
|
Implements GRPO following DeepSeek-R1 with token-level PPO clipping.
|
||||||
a frozen ref_model stores the old-policy log-probs. ratio = exp(logπ_θ - logπ_ref),
|
Advantages are group-normalized from scalar per-response rewards and
|
||||||
clipped PPO objective. Call ``sync_ref_model()`` after each data-generation round.
|
broadcast across all response tokens. The loss is computed **only on
|
||||||
|
response tokens** — prompt tokens are masked out.
|
||||||
|
|
||||||
|
Three model roles are distinguished:
|
||||||
|
|
||||||
|
* **Policy** ``self.model`` — the model being trained.
|
||||||
|
* **Old policy** ``self.old_model`` — the behaviour policy that generated
|
||||||
|
the responses. Used for the importance sampling ratio
|
||||||
|
``ρ = π_θ / π_old``. Synced externally after each data-generation round.
|
||||||
|
* **Reference model** ``self.ref_model`` — a frozen copy of the initial
|
||||||
|
policy (typically the SFT checkpoint) used **only** for the KL
|
||||||
|
regularisation term. It is never updated during training.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
model: nn.Module,
|
model: nn.Module,
|
||||||
device: str,
|
device: str,
|
||||||
|
old_model: nn.Module,
|
||||||
|
ref_model: nn.Module,
|
||||||
clip_eps: float = 0.2,
|
clip_eps: float = 0.2,
|
||||||
kl_coef: float = 0.01,
|
kl_coef: float = 0.01,
|
||||||
group_size: int = 4,
|
group_size: int = 4,
|
||||||
reduction: str = "mean",
|
|
||||||
sync_interval: int = 200,
|
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
super().__init__(model, device, **kwargs)
|
super().__init__(model, device, **kwargs)
|
||||||
self.ref_model = create_ref_model(
|
self.old_model = old_model
|
||||||
self.model_fn, self.executor.unwrap_model(model)
|
self.ref_model = ref_model
|
||||||
).to(device=self.device)
|
|
||||||
self.clip_eps = clip_eps
|
self.clip_eps = clip_eps
|
||||||
self.kl_coef = kl_coef
|
self.kl_coef = kl_coef
|
||||||
self.group_size = group_size
|
self.group_size = group_size
|
||||||
self.reduction = reduction
|
|
||||||
self.sync_interval = sync_interval
|
|
||||||
self._step = 0
|
|
||||||
|
|
||||||
def sync_ref_model(self):
|
def sync_old_model(self):
|
||||||
"""Copy current model weights to ref model."""
|
"""Copy current policy weights to old model."""
|
||||||
self.ref_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:
|
||||||
self._step += 1
|
|
||||||
if self._step % self.sync_interval == 0:
|
|
||||||
self.sync_ref_model()
|
|
||||||
|
|
||||||
batch = move_to_device(batch, self.device)
|
batch = move_to_device(batch, self.device)
|
||||||
prompts = batch["prompts"]
|
prompts = batch["prompts"]
|
||||||
responses = batch["responses"]
|
responses = batch["responses"]
|
||||||
@@ -313,33 +408,101 @@ class GRPOStrategy(BaseStrategy):
|
|||||||
responses_flat = responses.view(-1, response_len)
|
responses_flat = responses.view(-1, response_len)
|
||||||
masks_flat = masks.view(-1, response_len)
|
masks_flat = masks.view(-1, response_len)
|
||||||
prompt_expanded = prompts.unsqueeze(1).repeat(1, group_size, 1).flatten(0, 1)
|
prompt_expanded = prompts.unsqueeze(1).repeat(1, group_size, 1).flatten(0, 1)
|
||||||
|
prompt_mask = batch.get("prompt_mask")
|
||||||
|
if prompt_mask is None:
|
||||||
|
prompt_mask = prompts.ne(0)
|
||||||
|
prompt_mask_expanded = (
|
||||||
|
prompt_mask.unsqueeze(1).expand(-1, group_size, -1).flatten(0, 1)
|
||||||
|
)
|
||||||
|
prompt_len = prompt_expanded.size(1)
|
||||||
|
|
||||||
full_sequences = torch.cat([prompt_expanded, responses_flat], dim=-1)
|
full_sequences = torch.cat([prompt_expanded, responses_flat], dim=-1)
|
||||||
full_masks = torch.cat([torch.ones_like(prompt_expanded), masks_flat], dim=-1)
|
# Prompt tokens are masked out (0) so logprobs are computed only for
|
||||||
|
# response tokens. get_logprobs shifts the mask by one position, so
|
||||||
log_probs_policy = get_logprobs(
|
# the first response token's logprob (predicted from the last prompt
|
||||||
self.model, full_sequences, full_masks, self.reduction
|
# token) is correctly included.
|
||||||
|
full_masks = torch.cat(
|
||||||
|
[torch.zeros_like(prompt_expanded, dtype=torch.bool), masks_flat], dim=-1
|
||||||
)
|
)
|
||||||
log_probs_policy = log_probs_policy.view(batch_size, group_size)
|
|
||||||
|
|
||||||
|
# Build full attention mask: key-padding + causal
|
||||||
|
key_pad = torch.cat([prompt_mask_expanded, masks_flat.bool()], dim=-1)[
|
||||||
|
:, None, None, :
|
||||||
|
]
|
||||||
|
S = key_pad.shape[-1]
|
||||||
|
causal = torch.tril(
|
||||||
|
torch.ones(S, S, dtype=torch.bool, device=full_sequences.device)
|
||||||
|
)[None, None, :, :]
|
||||||
|
attn_mask = key_pad & causal
|
||||||
|
|
||||||
|
# get_logprobs returns [B*G, S-1] (S = prompt_len + response_len).
|
||||||
|
# Response token logprobs occupy the last ``response_len`` positions
|
||||||
|
# (the first response token is predicted from the last prompt token).
|
||||||
|
token_log_probs_policy = get_logprobs(
|
||||||
|
self.model, full_sequences, attn_mask, full_masks, "none"
|
||||||
|
)[:, prompt_len - 1 :]
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
log_probs_ref = get_logprobs(
|
token_log_probs_old = get_logprobs(
|
||||||
self.ref_model, full_sequences, full_masks, self.reduction
|
self.old_model, full_sequences, attn_mask, full_masks, "none"
|
||||||
)
|
)[:, prompt_len - 1 :]
|
||||||
log_probs_ref = log_probs_ref.view(batch_size, group_size)
|
token_log_probs_ref = get_logprobs(
|
||||||
|
self.ref_model, full_sequences, attn_mask, full_masks, "none"
|
||||||
|
)[:, prompt_len - 1 :]
|
||||||
|
|
||||||
eps = torch.finfo(log_probs_policy.dtype).eps
|
# Reshape to [B, G, response_len]
|
||||||
|
token_log_probs_policy = token_log_probs_policy.view(batch_size, group_size, -1)
|
||||||
|
token_log_probs_old = token_log_probs_old.view(batch_size, group_size, -1)
|
||||||
|
token_log_probs_ref = token_log_probs_ref.view(batch_size, group_size, -1)
|
||||||
|
token_masks = masks_flat.view(batch_size, group_size, -1).float()
|
||||||
|
|
||||||
|
# Group-normalized advantages from scalar per-response rewards.
|
||||||
|
eps = 1e-8
|
||||||
mean = rewards.mean(dim=-1, keepdim=True)
|
mean = rewards.mean(dim=-1, keepdim=True)
|
||||||
std = rewards.std(dim=-1, keepdim=True)
|
std = rewards.std(dim=-1, keepdim=True, unbiased=False)
|
||||||
advantages = (rewards - mean) / (std + eps)
|
advantages = (rewards - mean) / (std + eps)
|
||||||
|
# Broadcast scalar advantage to every response token: [B, G, 1]
|
||||||
|
advantages = advantages.unsqueeze(-1)
|
||||||
|
|
||||||
ratio = torch.exp(log_probs_policy - log_probs_ref)
|
# Token-level ratio (π_θ / π_old) and PPO clipping.
|
||||||
|
log_ratio = token_log_probs_policy - token_log_probs_old
|
||||||
|
ratio = torch.exp(log_ratio)
|
||||||
|
|
||||||
surr1 = ratio * advantages
|
surr1 = ratio * advantages
|
||||||
surr2 = torch.clamp(ratio, 1 - self.clip_eps, 1 + self.clip_eps) * advantages
|
surr2 = torch.clamp(ratio, 1 - self.clip_eps, 1 + self.clip_eps) * advantages
|
||||||
|
per_token_policy_loss = -torch.min(surr1, surr2)
|
||||||
|
token_count = token_masks.sum().clamp(min=1.0)
|
||||||
|
policy_loss = (per_token_policy_loss * token_masks).sum() / token_count
|
||||||
|
|
||||||
|
# KL penalty to frozen reference model with k1 estimator (non-negative):
|
||||||
|
# k1 = π_ref / π_θ - log(π_ref / π_θ) - 1, where π_ref / π_θ = exp(log_ref - log_policy).
|
||||||
|
log_ref_ratio = token_log_probs_ref - token_log_probs_policy
|
||||||
|
r = torch.exp(log_ref_ratio)
|
||||||
|
kl_per_token = r - torch.log(r + eps) - 1.0
|
||||||
|
kl_penalty = self.kl_coef * (kl_per_token * token_masks).sum() / token_count
|
||||||
|
|
||||||
policy_loss = -torch.min(surr1, surr2).mean()
|
|
||||||
kl_penalty = self.kl_coef * (log_probs_policy - log_probs_ref).square().mean()
|
|
||||||
total_loss = policy_loss + kl_penalty
|
total_loss = policy_loss + kl_penalty
|
||||||
|
|
||||||
return total_loss
|
return total_loss
|
||||||
|
|
||||||
|
def supports_online(self) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]:
|
||||||
|
return {
|
||||||
|
"prompts": result.prompts,
|
||||||
|
"prompt_mask": result.prompt_mask,
|
||||||
|
"responses": result.responses,
|
||||||
|
"masks": result.response_mask,
|
||||||
|
"rewards": result.rewards,
|
||||||
|
}
|
||||||
|
|
||||||
|
def _on_rollout_refresh(self):
|
||||||
|
"""Sync the behaviour policy whenever a fresh rollout arrives."""
|
||||||
|
self.sync_old_model()
|
||||||
|
|
||||||
|
|
||||||
|
# Factory aliases: online variants use the same strategy class; the
|
||||||
|
# ``RolloutRunner`` is injected by ``TrainContextBuilder`` to enable
|
||||||
|
# online mode, so no separate subclass is needed.
|
||||||
|
StrategyFactory.register("online_grpo")(GRPOStrategy)
|
||||||
|
StrategyFactory.register("online_dpo")(DPOStrategy)
|
||||||
|
|||||||
@@ -14,10 +14,11 @@ from tqdm import tqdm
|
|||||||
|
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
from astrai.parallel import only_on_rank
|
from astrai.parallel import only_on_rank
|
||||||
from astrai.parallel.setup import get_current_device, get_rank
|
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,
|
||||||
@@ -139,28 +140,31 @@ class CheckpointCallback(TrainCallback):
|
|||||||
self.interval = interval
|
self.interval = interval
|
||||||
self.weight_only = weight_only
|
self.weight_only = weight_only
|
||||||
self.save_extra_fn = save_extra_fn or CheckpointCallback.save_extra
|
self.save_extra_fn = save_extra_fn or CheckpointCallback.save_extra
|
||||||
self.last_ckpt_step = 0
|
self.last_ckpt_step = None
|
||||||
|
|
||||||
def _save_checkpoint(self, context: TrainContext):
|
def on_train_begin(self, context: TrainContext):
|
||||||
state_dict = context.executor.unwrap_model(context.model)
|
|
||||||
self.last_ckpt_step = context.optimizer_step
|
self.last_ckpt_step = context.optimizer_step
|
||||||
|
|
||||||
if get_rank() == 0:
|
def _save_checkpoint(self, context: TrainContext):
|
||||||
save_path = os.path.join(
|
self.last_ckpt_step = context.optimizer_step
|
||||||
self.save_dir,
|
|
||||||
f"epoch_{context.epoch}_step_{context.optimizer_step}",
|
with context.executor.checkpoint_context(context.model) as state_dict:
|
||||||
)
|
if state_dict is not None:
|
||||||
extra = self.save_extra_fn(context)
|
save_path = os.path.join(
|
||||||
meta = context.config.to_dict()
|
self.save_dir,
|
||||||
context.checkpoint = Checkpoint(
|
f"epoch_{context.epoch}_step_{context.optimizer_step}",
|
||||||
state_dict=state_dict,
|
)
|
||||||
epoch=context.epoch,
|
extra = self.save_extra_fn(context)
|
||||||
consumed_samples=context.consumed_samples,
|
meta = context.config.to_dict()
|
||||||
config=context.model_config,
|
context.checkpoint = Checkpoint(
|
||||||
extra=extra,
|
state_dict=state_dict,
|
||||||
meta=meta,
|
epoch=context.epoch,
|
||||||
)
|
consumed_samples=context.consumed_samples,
|
||||||
context.checkpoint.save(save_path)
|
config=context.model_config,
|
||||||
|
extra=extra,
|
||||||
|
meta=meta,
|
||||||
|
)
|
||||||
|
context.checkpoint.save(save_path)
|
||||||
|
|
||||||
def on_batch_end(self, context: TrainContext):
|
def on_batch_end(self, context: TrainContext):
|
||||||
if context.optimizer_step - self.last_ckpt_step >= self.interval:
|
if context.optimizer_step - self.last_ckpt_step >= self.interval:
|
||||||
@@ -210,7 +214,7 @@ class ProgressBarCallback(TrainCallback):
|
|||||||
@only_on_rank(0)
|
@only_on_rank(0)
|
||||||
def on_optimizer_step(self, context: TrainContext):
|
def on_optimizer_step(self, context: TrainContext):
|
||||||
postfix = {
|
postfix = {
|
||||||
"step": context.optimizer_step,
|
"step": f"{context.optimizer_step:d}",
|
||||||
"loss": f"{context.loss:.4f}",
|
"loss": f"{context.loss:.4f}",
|
||||||
"lr": f"{context.optimizer.param_groups[-1]['lr']:.2e}",
|
"lr": f"{context.optimizer.param_groups[-1]['lr']:.2e}",
|
||||||
}
|
}
|
||||||
@@ -232,19 +236,18 @@ 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,
|
||||||
):
|
):
|
||||||
self.last_log_flush_step = 0
|
self.last_log_flush_step = None
|
||||||
self.save_interval = save_interval
|
self.save_interval = save_interval
|
||||||
self.metrics = metrics or ["loss", "lr"]
|
self.metrics = metrics or ["loss", "lr"]
|
||||||
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 = []
|
||||||
|
|
||||||
@@ -253,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):
|
||||||
@@ -298,15 +302,20 @@ class MetricCallback(TrainCallback):
|
|||||||
context.model.train()
|
context.model.train()
|
||||||
return avg_loss
|
return avg_loss
|
||||||
|
|
||||||
|
def on_train_begin(self, context: TrainContext):
|
||||||
|
self.last_log_flush_step = context.optimizer_step
|
||||||
|
|
||||||
@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
|
||||||
@@ -327,8 +336,12 @@ class MetricCallback(TrainCallback):
|
|||||||
self._append("epoch", context)
|
self._append("epoch", context)
|
||||||
|
|
||||||
def on_train_end(self, context):
|
def on_train_end(self, context):
|
||||||
if context.optimizer_step != self.last_log_flush_step:
|
if (
|
||||||
|
self.last_log_flush_step is None
|
||||||
|
or context.optimizer_step != self.last_log_flush_step
|
||||||
|
):
|
||||||
self._flush(context.epoch, context.optimizer_step)
|
self._flush(context.epoch, context.optimizer_step)
|
||||||
|
self.last_log_flush_step = context.optimizer_step
|
||||||
|
|
||||||
def on_error(self, context):
|
def on_error(self, context):
|
||||||
self._flush(context.epoch, context.optimizer_step)
|
self._flush(context.epoch, context.optimizer_step)
|
||||||
|
|||||||
+152
-50
@@ -1,3 +1,5 @@
|
|||||||
|
import logging
|
||||||
|
import threading
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Dict, Optional, Self
|
from typing import Any, Dict, Optional, Self
|
||||||
@@ -7,14 +9,20 @@ import torch.nn as nn
|
|||||||
from torch.utils.data import DataLoader, random_split
|
from torch.utils.data import DataLoader, random_split
|
||||||
|
|
||||||
from astrai.config.train_config import TrainConfig
|
from astrai.config.train_config import TrainConfig
|
||||||
from astrai.dataset import ResumableDistributedSampler
|
from astrai.dataset import RDSampler
|
||||||
|
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.trainer.metric_util import GradSNRTracker
|
||||||
|
from astrai.trainer.rollout import RolloutGenerator, RolloutRunner
|
||||||
from astrai.trainer.strategy import BaseStrategy, StrategyFactory
|
from astrai.trainer.strategy import BaseStrategy, StrategyFactory
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class TrainContext:
|
class TrainContext:
|
||||||
@@ -27,11 +35,11 @@ class TrainContext:
|
|||||||
config: TrainConfig = field(default=None)
|
config: TrainConfig = field(default=None)
|
||||||
model_config: dict = field(default_factory=dict)
|
model_config: dict = field(default_factory=dict)
|
||||||
executor: BaseExecutor = field(default=None)
|
executor: BaseExecutor = field(default=None)
|
||||||
|
|
||||||
epoch: int = field(default=0)
|
epoch: int = field(default=0)
|
||||||
consumed_samples: int = field(default=0)
|
consumed_samples: int = field(default=0)
|
||||||
loss: float = field(default=0.0)
|
loss: float = field(default=0.0)
|
||||||
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)
|
||||||
|
|
||||||
@@ -39,6 +47,15 @@ class TrainContext:
|
|||||||
rank: int = field(default=0)
|
rank: int = field(default=0)
|
||||||
kwargs: Dict[str, Any] = field(default_factory=dict)
|
kwargs: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
_stop_event: threading.Event = field(default_factory=threading.Event)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def stop_requested(self) -> bool:
|
||||||
|
return self._stop_event.is_set()
|
||||||
|
|
||||||
|
def request_stop(self) -> None:
|
||||||
|
self._stop_event.set()
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def optimizer_step(self) -> int:
|
def optimizer_step(self) -> int:
|
||||||
return self.consumed_samples // (
|
return self.consumed_samples // (
|
||||||
@@ -54,10 +71,12 @@ class TrainContextBuilder:
|
|||||||
config: TrainConfig,
|
config: TrainConfig,
|
||||||
):
|
):
|
||||||
self.config = config
|
self.config = config
|
||||||
self._resume_dir: Optional[str] = None
|
self._param_path: Optional[str] = None
|
||||||
|
self._resume: bool = False
|
||||||
|
|
||||||
def with_resume_dir(self, resume_dir: Optional[str]) -> Self:
|
def with_param_path(self, param_path: Optional[str], resume: bool = False) -> Self:
|
||||||
self._resume_dir = resume_dir
|
self._param_path = param_path
|
||||||
|
self._resume = resume
|
||||||
return self
|
return self
|
||||||
|
|
||||||
def build(self) -> TrainContext:
|
def build(self) -> TrainContext:
|
||||||
@@ -70,50 +89,72 @@ class TrainContextBuilder:
|
|||||||
**cfg.executor_kwargs,
|
**cfg.executor_kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
model = cfg.model_fn()
|
|
||||||
model = model.to(device=device)
|
|
||||||
|
|
||||||
model_config = {}
|
model_config = {}
|
||||||
if self._resume_dir:
|
if self._param_path:
|
||||||
config_path = Path(self._resume_dir) / "config.json"
|
config_path = Path(self._param_path) / "config.json"
|
||||||
if config_path.exists():
|
if config_path.exists():
|
||||||
model_config = load_json(config_path)
|
model_config = load_json(config_path)
|
||||||
|
|
||||||
if not model_config and hasattr(model, "config"):
|
preloaded_state_dict = None
|
||||||
model_config = model.config.to_dict()
|
preloaded_epoch = cfg.start_epoch
|
||||||
|
preloaded_consumed = cfg.start_samples * get_world_size()
|
||||||
|
preloaded_checkpoint = None
|
||||||
|
if self._param_path:
|
||||||
|
checkpoint = Checkpoint.load_any(self._param_path)
|
||||||
|
if checkpoint is not None:
|
||||||
|
preloaded_state_dict = checkpoint.state_dict
|
||||||
|
if checkpoint.config:
|
||||||
|
model_config = checkpoint.config
|
||||||
|
if self._resume:
|
||||||
|
preloaded_epoch = checkpoint.epoch
|
||||||
|
per_step = (
|
||||||
|
cfg.batch_per_device * get_world_size() * cfg.grad_accum_steps
|
||||||
|
)
|
||||||
|
preloaded_consumed = (
|
||||||
|
checkpoint.consumed_samples // per_step
|
||||||
|
) * per_step
|
||||||
|
preloaded_checkpoint = checkpoint
|
||||||
|
|
||||||
|
if not model_config and hasattr(cfg.model_fn(), "config"):
|
||||||
|
model_config = cfg.model_fn().config.to_dict()
|
||||||
|
|
||||||
|
def _before_wrap(m):
|
||||||
|
m = m.to(device=device)
|
||||||
|
if cfg.lora is not None:
|
||||||
|
inject_lora(
|
||||||
|
m,
|
||||||
|
r=cfg.lora.r,
|
||||||
|
alpha=cfg.lora.alpha,
|
||||||
|
target_modules=set(cfg.lora.target_modules),
|
||||||
|
)
|
||||||
|
if preloaded_state_dict is not None:
|
||||||
|
m.load_state_dict(preloaded_state_dict, strict=False)
|
||||||
|
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(
|
||||||
model=model,
|
|
||||||
world_size=get_world_size(),
|
world_size=get_world_size(),
|
||||||
rank=get_rank(),
|
rank=get_rank(),
|
||||||
config=cfg,
|
config=cfg,
|
||||||
model_config=model_config,
|
model_config=model_config,
|
||||||
executor=executor,
|
executor=executor,
|
||||||
|
epoch=preloaded_epoch,
|
||||||
|
consumed_samples=preloaded_consumed,
|
||||||
|
checkpoint=preloaded_checkpoint,
|
||||||
)
|
)
|
||||||
|
|
||||||
if self._resume_dir:
|
context.model, context.optimizer, context.scheduler = executor.prepare(
|
||||||
checkpoint = Checkpoint.load_any(self._resume_dir)
|
cfg.model_fn,
|
||||||
if checkpoint is not None:
|
cfg.optimizer_fn,
|
||||||
model.load_state_dict(checkpoint.state_dict, strict=False)
|
cfg.scheduler_fn,
|
||||||
if checkpoint.config:
|
before_wrap=_before_wrap,
|
||||||
context.model_config = checkpoint.config
|
after_wrap=_after_wrap,
|
||||||
context.epoch = checkpoint.epoch or cfg.start_epoch
|
)
|
||||||
if checkpoint.consumed_samples > 0:
|
|
||||||
context.consumed_samples = checkpoint.consumed_samples
|
|
||||||
else:
|
|
||||||
context.consumed_samples = cfg.start_samples * context.world_size
|
|
||||||
context.checkpoint = checkpoint
|
|
||||||
|
|
||||||
if cfg.lora is not None:
|
|
||||||
inject_lora(
|
|
||||||
model,
|
|
||||||
r=cfg.lora.r,
|
|
||||||
alpha=cfg.lora.alpha,
|
|
||||||
target_modules=set(cfg.lora.target_modules),
|
|
||||||
)
|
|
||||||
|
|
||||||
context.optimizer = cfg.optimizer_fn(model)
|
|
||||||
context.scheduler = cfg.scheduler_fn(context.optimizer)
|
|
||||||
|
|
||||||
train_dataset = cfg.dataset
|
train_dataset = cfg.dataset
|
||||||
val_dataset = cfg.val_dataset
|
val_dataset = cfg.val_dataset
|
||||||
@@ -128,7 +169,16 @@ class TrainContextBuilder:
|
|||||||
)
|
)
|
||||||
|
|
||||||
sampler_offset = context.consumed_samples // context.world_size
|
sampler_offset = context.consumed_samples // context.world_size
|
||||||
sampler = ResumableDistributedSampler(
|
|
||||||
|
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(
|
||||||
data_source=train_dataset,
|
data_source=train_dataset,
|
||||||
start_epoch=context.epoch,
|
start_epoch=context.epoch,
|
||||||
start_iter=sampler_offset,
|
start_iter=sampler_offset,
|
||||||
@@ -141,10 +191,11 @@ class TrainContextBuilder:
|
|||||||
num_workers=cfg.num_workers,
|
num_workers=cfg.num_workers,
|
||||||
pin_memory=cfg.pin_memory,
|
pin_memory=cfg.pin_memory,
|
||||||
prefetch_factor=cfg.prefetch_factor,
|
prefetch_factor=cfg.prefetch_factor,
|
||||||
|
collate_fn=cfg.collate_fn,
|
||||||
)
|
)
|
||||||
|
|
||||||
if val_dataset is not None:
|
if val_dataset is not None:
|
||||||
val_sampler = ResumableDistributedSampler(
|
val_sampler = RDSampler(
|
||||||
data_source=val_dataset,
|
data_source=val_dataset,
|
||||||
start_epoch=0,
|
start_epoch=0,
|
||||||
start_iter=0,
|
start_iter=0,
|
||||||
@@ -158,17 +209,9 @@ class TrainContextBuilder:
|
|||||||
num_workers=cfg.num_workers,
|
num_workers=cfg.num_workers,
|
||||||
pin_memory=cfg.pin_memory,
|
pin_memory=cfg.pin_memory,
|
||||||
prefetch_factor=cfg.prefetch_factor,
|
prefetch_factor=cfg.prefetch_factor,
|
||||||
|
collate_fn=cfg.collate_fn,
|
||||||
)
|
)
|
||||||
|
|
||||||
context.model, context.optimizer, context.dataloader, context.scheduler = (
|
|
||||||
executor.prepare(
|
|
||||||
model,
|
|
||||||
context.optimizer,
|
|
||||||
context.dataloader,
|
|
||||||
context.scheduler,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
if context.checkpoint and context.checkpoint.extra:
|
if context.checkpoint and context.checkpoint.extra:
|
||||||
extra = context.checkpoint.extra
|
extra = context.checkpoint.extra
|
||||||
for name in ("optimizer", "scheduler"):
|
for name in ("optimizer", "scheduler"):
|
||||||
@@ -177,13 +220,72 @@ class TrainContextBuilder:
|
|||||||
if obj is not None:
|
if obj is not None:
|
||||||
obj.load_state_dict(extra[name])
|
obj.load_state_dict(extra[name])
|
||||||
|
|
||||||
|
strategy_kwargs = dict(cfg.extra_kwargs)
|
||||||
|
|
||||||
|
needs_ref = cfg.strategy in (
|
||||||
|
"dpo",
|
||||||
|
"grpo",
|
||||||
|
"online_grpo",
|
||||||
|
"online_dpo",
|
||||||
|
)
|
||||||
|
needs_old = cfg.strategy in ("grpo", "online_grpo")
|
||||||
|
|
||||||
|
if needs_ref:
|
||||||
|
strategy_kwargs["ref_model"] = create_ref_model(
|
||||||
|
cfg.model_fn, executor=executor, model=context.model, device=device
|
||||||
|
)
|
||||||
|
|
||||||
|
if needs_old:
|
||||||
|
strategy_kwargs["old_model"] = create_ref_model(
|
||||||
|
cfg.model_fn, executor=executor, model=context.model, device=device
|
||||||
|
)
|
||||||
|
|
||||||
context.strategy = StrategyFactory.create(
|
context.strategy = StrategyFactory.create(
|
||||||
cfg.strategy,
|
cfg.strategy,
|
||||||
model=context.model,
|
model=context.model,
|
||||||
device=device,
|
device=device,
|
||||||
executor=executor,
|
executor=executor,
|
||||||
model_fn=cfg.model_fn,
|
**strategy_kwargs,
|
||||||
**cfg.extra_kwargs,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Enable online rollout when the train_type is an ``online_*`` variant.
|
||||||
|
is_online = cfg.strategy.startswith("online_")
|
||||||
|
if is_online:
|
||||||
|
if not context.strategy.supports_online():
|
||||||
|
raise ValueError(
|
||||||
|
f"Strategy '{cfg.strategy}' does not support online rollout"
|
||||||
|
)
|
||||||
|
if cfg.reward_model_fn is None:
|
||||||
|
raise ValueError("reward_model_fn is required for online RL strategies")
|
||||||
|
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(self._param_path)
|
||||||
|
reward_model = cfg.reward_model_fn()
|
||||||
|
|
||||||
|
group_size = strategy_kwargs.get("group_size", 1)
|
||||||
|
rollout_batch_size = group_size * max(1, cfg.batch_per_device)
|
||||||
|
max_seq_len = getattr(context.model.config, "max_position_embeddings", None)
|
||||||
|
|
||||||
|
scheduler = InferenceScheduler(
|
||||||
|
model=context.model,
|
||||||
|
tokenizer=tokenizer,
|
||||||
|
max_batch_size=rollout_batch_size,
|
||||||
|
max_seq_len=max_seq_len,
|
||||||
|
)
|
||||||
|
|
||||||
|
generator = RolloutGenerator(
|
||||||
|
scheduler=scheduler,
|
||||||
|
tokenizer=tokenizer,
|
||||||
|
max_tokens=cfg.rollout_max_tokens,
|
||||||
|
group_size=group_size,
|
||||||
|
temperature=cfg.rollout_temperature,
|
||||||
|
top_k=cfg.rollout_top_k,
|
||||||
|
top_p=cfg.rollout_top_p,
|
||||||
|
)
|
||||||
|
runner = RolloutRunner(
|
||||||
|
generator=generator,
|
||||||
|
reward_model=reward_model,
|
||||||
|
rollout_interval=cfg.rollout_interval,
|
||||||
|
)
|
||||||
|
context.strategy.set_rollout_runner(runner)
|
||||||
|
|
||||||
return context
|
return context
|
||||||
|
|||||||
@@ -1,8 +1,14 @@
|
|||||||
import logging
|
import logging
|
||||||
from typing import List, Optional
|
from typing import List, Optional
|
||||||
|
|
||||||
|
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.signal_handler import (
|
||||||
|
register_signal_handlers,
|
||||||
|
unregister_signal_handlers,
|
||||||
|
)
|
||||||
from astrai.trainer.train_callback import (
|
from astrai.trainer.train_callback import (
|
||||||
CallbackFactory,
|
CallbackFactory,
|
||||||
TrainCallback,
|
TrainCallback,
|
||||||
@@ -36,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,
|
||||||
@@ -52,10 +58,13 @@ class Trainer:
|
|||||||
if method:
|
if method:
|
||||||
method(context)
|
method(context)
|
||||||
|
|
||||||
def _trainer_loop(self, resume_dir: Optional[str] = None):
|
def _trainer_loop(self, param_path: Optional[str] = None, resume: bool = False):
|
||||||
context = (
|
context = (
|
||||||
TrainContextBuilder(self.train_config).with_resume_dir(resume_dir).build()
|
TrainContextBuilder(self.train_config)
|
||||||
|
.with_param_path(param_path, resume=resume)
|
||||||
|
.build()
|
||||||
)
|
)
|
||||||
|
register_signal_handlers(context)
|
||||||
executor = context.executor
|
executor = context.executor
|
||||||
self._call_callbacks("on_train_begin", context)
|
self._call_callbacks("on_train_begin", context)
|
||||||
|
|
||||||
@@ -63,10 +72,14 @@ class Trainer:
|
|||||||
context.model.train()
|
context.model.train()
|
||||||
|
|
||||||
for epoch in range(context.epoch, context.config.n_epoch):
|
for epoch in range(context.epoch, context.config.n_epoch):
|
||||||
|
if context.stop_requested:
|
||||||
|
break
|
||||||
context.epoch = epoch
|
context.epoch = epoch
|
||||||
self._call_callbacks("on_epoch_begin", context)
|
self._call_callbacks("on_epoch_begin", context)
|
||||||
|
|
||||||
for batch in context.dataloader:
|
for batch in context.dataloader:
|
||||||
|
if context.stop_requested:
|
||||||
|
break
|
||||||
with executor.accumulate(context.model):
|
with executor.accumulate(context.model):
|
||||||
self._call_callbacks("on_batch_begin", context)
|
self._call_callbacks("on_batch_begin", context)
|
||||||
loss = context.strategy(batch)
|
loss = context.strategy(batch)
|
||||||
@@ -81,6 +94,7 @@ class Trainer:
|
|||||||
if executor.sync_gradients:
|
if executor.sync_gradients:
|
||||||
self._call_callbacks("on_optimizer_step", context)
|
self._call_callbacks("on_optimizer_step", context)
|
||||||
context.optimizer.step()
|
context.optimizer.step()
|
||||||
|
context.strategy.on_optimizer_step()
|
||||||
context.optimizer.zero_grad()
|
context.optimizer.zero_grad()
|
||||||
|
|
||||||
if context.scheduler:
|
if context.scheduler:
|
||||||
@@ -88,14 +102,23 @@ class Trainer:
|
|||||||
|
|
||||||
self._call_callbacks("on_epoch_end", context)
|
self._call_callbacks("on_epoch_end", context)
|
||||||
|
|
||||||
|
if context.stop_requested:
|
||||||
|
logger.warning(
|
||||||
|
"Training interrupted by signal, saving emergency checkpoint..."
|
||||||
|
)
|
||||||
|
self._call_callbacks("on_error", context)
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error("Training failed: %s", str(e), exc_info=True)
|
logger.error("Training failed: %s", str(e), exc_info=True)
|
||||||
self._call_callbacks("on_error", context)
|
self._call_callbacks("on_error", context)
|
||||||
raise
|
raise
|
||||||
finally:
|
finally:
|
||||||
self._call_callbacks("on_train_end", context)
|
self._call_callbacks("on_train_end", context)
|
||||||
|
if executor.use_distributed and dist.is_initialized():
|
||||||
|
dist.barrier()
|
||||||
|
unregister_signal_handlers()
|
||||||
|
|
||||||
def train(self, resume_dir: Optional[str] = None):
|
def train(self, param_path: Optional[str] = None, resume: bool = False):
|
||||||
cfg = self.train_config
|
cfg = self.train_config
|
||||||
spawn_parallel_fn(
|
spawn_parallel_fn(
|
||||||
self._trainer_loop,
|
self._trainer_loop,
|
||||||
@@ -105,5 +128,6 @@ class Trainer:
|
|||||||
master_port=cfg.master_port,
|
master_port=cfg.master_port,
|
||||||
device_type=cfg.device_type,
|
device_type=cfg.device_type,
|
||||||
start_method=cfg.start_method,
|
start_method=cfg.start_method,
|
||||||
resume_dir=resume_dir,
|
param_path=param_path,
|
||||||
|
resume=resume,
|
||||||
)
|
)
|
||||||
|
|||||||
+32
-3
@@ -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
|
||||||
|
|
||||||
@@ -20,13 +46,14 @@ def _arch_flags() -> list[str]:
|
|||||||
_kernels_dir = Path("csrc/kernels")
|
_kernels_dir = Path("csrc/kernels")
|
||||||
REGISTRY: dict[str, dict] = {}
|
REGISTRY: dict[str, dict] = {}
|
||||||
|
|
||||||
CXX_FLAGS = ["-O3", "-march=native", "-funroll-loops"]
|
CXX_FLAGS = ["-O3", "-funroll-loops"]
|
||||||
NVCC_FLAGS = [
|
NVCC_FLAGS = [
|
||||||
"-O3",
|
"-O3",
|
||||||
"--expt-relaxed-constexpr",
|
"--expt-relaxed-constexpr",
|
||||||
"--use_fast_math",
|
"--use_fast_math",
|
||||||
"--ptxas-options=-O3,-v",
|
"--ptxas-options=-O3,-v",
|
||||||
"--extra-device-vectorization",
|
"--extra-device-vectorization",
|
||||||
|
"--threads=16",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
@@ -42,5 +69,7 @@ def register(name: str, sources: list[str] | None = None, **kwargs):
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
register("gqa_decode_attn")
|
register("attn_decode")
|
||||||
register("gqa_prefill_attn")
|
register("attn_prefill")
|
||||||
|
register("attn_paged_decode")
|
||||||
|
register("rotary_emb")
|
||||||
|
|||||||
@@ -0,0 +1,71 @@
|
|||||||
|
#pragma once
|
||||||
|
|
||||||
|
|
||||||
|
template<typename T, typename AT = float>
|
||||||
|
struct AttentionParams {
|
||||||
|
int batch;
|
||||||
|
int q_head;
|
||||||
|
int kv_head;
|
||||||
|
int q_len;
|
||||||
|
int kv_len;
|
||||||
|
int head_dim;
|
||||||
|
int use_mask;
|
||||||
|
int causal_offset; // -1 = non-causal; >=0 = absolute position of first Q token
|
||||||
|
int num_splits;
|
||||||
|
float scale;
|
||||||
|
|
||||||
|
// Q strides (element offsets for each dim — layout-agnostic)
|
||||||
|
int q_stride_b, q_stride_h, q_stride_l, q_stride_d;
|
||||||
|
// 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;
|
||||||
|
|
||||||
|
// Mask: 2D [batch, kv_len], 3D [batch, q_len, kv_len],
|
||||||
|
// or 4D [batch, n_heads, q_len, kv_len] (head dim broadcasts when stride=0)
|
||||||
|
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__ k;
|
||||||
|
const T* __restrict__ v;
|
||||||
|
const bool* __restrict__ mask;
|
||||||
|
|
||||||
|
T* __restrict__ o;
|
||||||
|
AT* __restrict__ o_part;
|
||||||
|
AT* __restrict__ ml_part;
|
||||||
|
};
|
||||||
|
|
||||||
|
template<typename T, typename AT = float>
|
||||||
|
struct PagedAttentionParams {
|
||||||
|
int batch;
|
||||||
|
int q_head;
|
||||||
|
int kv_head;
|
||||||
|
int q_len;
|
||||||
|
int kv_len;
|
||||||
|
int head_dim;
|
||||||
|
int use_mask;
|
||||||
|
int causal_offset;
|
||||||
|
float scale;
|
||||||
|
|
||||||
|
int num_splits;
|
||||||
|
int page_size;
|
||||||
|
int max_pages;
|
||||||
|
|
||||||
|
// Q strides (layout-agnostic)
|
||||||
|
int q_stride_b, q_stride_h, q_stride_l, q_stride_d;
|
||||||
|
|
||||||
|
// Mask strides (2D, 3D, or 4D)
|
||||||
|
int mask_b_stride;
|
||||||
|
int mask_h_stride;
|
||||||
|
int mask_q_stride;
|
||||||
|
|
||||||
|
const T* __restrict__ q;
|
||||||
|
const T* __restrict__ k_cache;
|
||||||
|
const T* __restrict__ v_cache;
|
||||||
|
const bool* __restrict__ mask;
|
||||||
|
const int64_t* __restrict__ page_table;
|
||||||
|
|
||||||
|
T* __restrict__ o;
|
||||||
|
AT* __restrict__ o_part;
|
||||||
|
AT* __restrict__ ml_part;
|
||||||
|
};
|
||||||
@@ -0,0 +1,37 @@
|
|||||||
|
#include "attn_dispatchers.cuh"
|
||||||
|
#include "attn_entry_utils.cuh"
|
||||||
|
|
||||||
|
torch::Tensor attn_decode(
|
||||||
|
torch::Tensor q,
|
||||||
|
torch::Tensor k,
|
||||||
|
torch::Tensor v,
|
||||||
|
c10::optional<torch::Tensor> mask,
|
||||||
|
int64_t causal_offset,
|
||||||
|
double scale,
|
||||||
|
int64_t layout
|
||||||
|
) {
|
||||||
|
AttentionParams<bf16> p;
|
||||||
|
attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p);
|
||||||
|
TORCH_CHECK(p.q_len == 1, "Q seq_len must be 1");
|
||||||
|
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
|
||||||
|
|
||||||
|
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
|
||||||
|
auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
|
||||||
|
p.o = (bf16*)O_view.data_ptr();
|
||||||
|
|
||||||
|
alloc_split_partials(p);
|
||||||
|
DISPATCH_HEAD_DIM(p.head_dim, dispatch_decode, p);
|
||||||
|
return O;
|
||||||
|
}
|
||||||
|
|
||||||
|
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||||
|
m.def("attn_decode", &attn_decode,
|
||||||
|
py::arg("q"),
|
||||||
|
py::arg("k"),
|
||||||
|
py::arg("v"),
|
||||||
|
py::arg("mask") = py::none(),
|
||||||
|
py::arg("causal_offset") = -1,
|
||||||
|
py::arg("scale") = 0.0,
|
||||||
|
py::arg("layout") = 0,
|
||||||
|
"GQA decode (tensor-core head-packing on sm_80+, scalar fallback)");
|
||||||
|
}
|
||||||
@@ -0,0 +1,129 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
#include <float.h>
|
||||||
|
#include "attn_common.h"
|
||||||
|
#include "attn_warp_utils.cuh"
|
||||||
|
constexpr int DC_CHUNK = 64;
|
||||||
|
|
||||||
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
|
__global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
|
||||||
|
int batch = blockIdx.x / p.kv_head;
|
||||||
|
int kv_head = blockIdx.x % p.kv_head;
|
||||||
|
int split = blockIdx.z;
|
||||||
|
int group_size = blockDim.y;
|
||||||
|
int q_head = kv_head * group_size + threadIdx.y;
|
||||||
|
int lane = threadIdx.x;
|
||||||
|
int hd_per_thread = p.head_dim / 32;
|
||||||
|
|
||||||
|
// Q: [batch, q_head, q_len=1, head_dim] — stride-based
|
||||||
|
float q_reg[8];
|
||||||
|
int q_off = batch * p.q_stride_b + q_head * p.q_stride_h
|
||||||
|
+ lane * hd_per_thread * p.q_stride_d;
|
||||||
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
|
q_reg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]);
|
||||||
|
|
||||||
|
// 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 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};
|
||||||
|
|
||||||
|
extern __shared__ __align__(16) bf16 k_smem[];
|
||||||
|
|
||||||
|
// Split-KV: each split processes a contiguous subset of chunks
|
||||||
|
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
|
||||||
|
int chunks_per_split = (chunks_total + p.num_splits - 1) / p.num_splits;
|
||||||
|
int ch_begin = split * chunks_per_split;
|
||||||
|
int ch_end = min(chunks_total, ch_begin + chunks_per_split);
|
||||||
|
|
||||||
|
for (int ci = ch_begin; ci < ch_end; ci++) {
|
||||||
|
int chunk_start = ci * DC_CHUNK;
|
||||||
|
int this_chunk = min(DC_CHUNK, p.kv_len - chunk_start);
|
||||||
|
|
||||||
|
// Load K into shared memory (gather from strided global)
|
||||||
|
int total = this_chunk * p.head_dim;
|
||||||
|
for (int i = threadIdx.y * 32 + lane; i < total;
|
||||||
|
i += blockDim.x * blockDim.y) {
|
||||||
|
int s = i / p.head_dim;
|
||||||
|
int d_dim = i % p.head_dim;
|
||||||
|
int kv_idx = chunk_start + s;
|
||||||
|
int g_off = kv_base + kv_idx * p.kv_stride_l + d_dim * p.kv_stride_d;
|
||||||
|
k_smem[i] = p.k[g_off];
|
||||||
|
}
|
||||||
|
__syncthreads();
|
||||||
|
|
||||||
|
for (int s = 0; s < this_chunk; s++) {
|
||||||
|
float partial = 0.0f;
|
||||||
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
|
partial += q_reg[i] * __bfloat162float(
|
||||||
|
k_smem[s * p.head_dim + lane * hd_per_thread + i]);
|
||||||
|
partial = warp_reduce_sum(partial) * p.scale;
|
||||||
|
|
||||||
|
int kv_idx = chunk_start + s;
|
||||||
|
if constexpr (HasMask) {
|
||||||
|
if (!p.mask[mask_base + kv_idx])
|
||||||
|
partial = -FLT_MAX;
|
||||||
|
}
|
||||||
|
if constexpr (IsCausal) {
|
||||||
|
if (kv_idx > p.causal_offset)
|
||||||
|
partial = -FLT_MAX;
|
||||||
|
}
|
||||||
|
|
||||||
|
float new_m = fmaxf(m, partial);
|
||||||
|
float alpha = __expf(m - new_m);
|
||||||
|
float beta = __expf(partial - new_m);
|
||||||
|
d = d * alpha + beta;
|
||||||
|
|
||||||
|
int v_off = kv_base + kv_idx * p.kv_stride_l
|
||||||
|
+ lane * hd_per_thread * p.kv_stride_d;
|
||||||
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
|
acc_reg[i] = fmaf(acc_reg[i], alpha,
|
||||||
|
__bfloat162float(p.v[v_off + i * p.kv_stride_d]) * beta);
|
||||||
|
m = new_m;
|
||||||
|
}
|
||||||
|
__syncthreads();
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- write UN-normalised partials for this split ----
|
||||||
|
size_t bh = (size_t)batch * p.q_head + q_head;
|
||||||
|
size_t slot = bh * MAX_SPLITS + split;
|
||||||
|
int d0 = lane * hd_per_thread;
|
||||||
|
for (int i = 0; i < hd_per_thread; i++) {
|
||||||
|
int dd = d0 + i;
|
||||||
|
p.o_part[slot * p.head_dim + dd] = acc_reg[i];
|
||||||
|
}
|
||||||
|
if (lane == 0) {
|
||||||
|
p.ml_part[slot * 2] = m;
|
||||||
|
p.ml_part[slot * 2 + 1] = d;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
__global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
|
||||||
|
int bh = blockIdx.x;
|
||||||
|
int d = threadIdx.x;
|
||||||
|
if (d >= p.head_dim) return;
|
||||||
|
|
||||||
|
int batch = bh / p.q_head;
|
||||||
|
int q_head = bh % p.q_head;
|
||||||
|
|
||||||
|
size_t split_base = (size_t)bh * MAX_SPLITS;
|
||||||
|
const float* mlp = p.ml_part + split_base * 2;
|
||||||
|
const float* op = p.o_part + split_base * p.head_dim;
|
||||||
|
|
||||||
|
float m = -FLT_MAX, l = 0.0f, acc = 0.0f;
|
||||||
|
for (int s = 0; s < p.num_splits; s++) {
|
||||||
|
float mi = mlp[s * 2];
|
||||||
|
if (mi <= -FLT_MAX) continue;
|
||||||
|
float li = mlp[s * 2 + 1];
|
||||||
|
float nm = fmaxf(m, mi);
|
||||||
|
float corr = __expf(m - nm);
|
||||||
|
float e = __expf(mi - nm);
|
||||||
|
acc = fmaf(acc, corr, op[s * p.head_dim + d] * e);
|
||||||
|
l = fmaf(l, corr, li * e);
|
||||||
|
m = nm;
|
||||||
|
}
|
||||||
|
|
||||||
|
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
|
||||||
|
int o_off = batch * p.q_stride_b + q_head * p.q_stride_h + d * p.q_stride_d;
|
||||||
|
p.o[o_off] = __float2bfloat16(acc * inv);
|
||||||
|
}
|
||||||
@@ -0,0 +1,169 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <cfloat>
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
#include "attn_common.h"
|
||||||
|
#include "attn_mma_utils.cuh"
|
||||||
|
#include "attn_warp_utils.cuh"
|
||||||
|
|
||||||
|
// Split-K (FlashDecoding) tensor-core decode via GQA head-packing.
|
||||||
|
// Decode has q_len == 1, so we pack G = q_head/kv_head query heads into the
|
||||||
|
// M=16 rows of mma.sync.m16n8k16, turning G independent GEMVs into a single
|
||||||
|
// GEMM that reuses each loaded K/V tile across all G heads.
|
||||||
|
//
|
||||||
|
// IsCausal and HasMask are compile-time bools — no runtime branch in the
|
||||||
|
// inner compute loop.
|
||||||
|
//
|
||||||
|
// Traits = KernelTraits<HEAD_DIM, BC=32, WARPS=1, STAGES=<2 or 1>>.
|
||||||
|
template <typename Traits, bool IsCausal, bool HasMask>
|
||||||
|
__global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
||||||
|
const int lane = threadIdx.x;
|
||||||
|
const int gid = lane >> 2;
|
||||||
|
const int tid4 = lane & 3;
|
||||||
|
|
||||||
|
const int pass = blockIdx.x / p.kv_head;
|
||||||
|
const int kv_head = blockIdx.x % p.kv_head;
|
||||||
|
const int batch = blockIdx.y;
|
||||||
|
const int split = blockIdx.z;
|
||||||
|
|
||||||
|
constexpr int MAX_G = 16;
|
||||||
|
const int G_total = p.q_head / p.kv_head;
|
||||||
|
const int g_begin = pass * MAX_G;
|
||||||
|
const int G = min(MAX_G, G_total - g_begin);
|
||||||
|
const int q_head0 = kv_head * G_total + g_begin;
|
||||||
|
|
||||||
|
// Double-buffered shared memory for K/V (no sQ needed)
|
||||||
|
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
|
||||||
|
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
|
||||||
|
|
||||||
|
// Load Q directly from global into mma A-operand registers.
|
||||||
|
// stride_row = p.q_stride_h for decode (q_len=1).
|
||||||
|
const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h;
|
||||||
|
const int qra = gid;
|
||||||
|
const int qrb = gid + 8;
|
||||||
|
const bool va = qra < G, vb = qrb < G;
|
||||||
|
unsigned Qa[Traits::KD][4];
|
||||||
|
load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
|
||||||
|
qra, qrb, va, vb, tid4, Qa);
|
||||||
|
|
||||||
|
float Oacc[Traits::DN8][4];
|
||||||
|
#pragma unroll
|
||||||
|
for (int j = 0; j < Traits::DN8; j++)
|
||||||
|
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
||||||
|
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
||||||
|
|
||||||
|
const int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
||||||
|
const int tiles_total = (p.kv_len + Traits::BC - 1) / Traits::BC;
|
||||||
|
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
|
||||||
|
const int ti_begin = split * tiles_per_split;
|
||||||
|
const int ti_end = min(tiles_total, ti_begin + tiles_per_split);
|
||||||
|
|
||||||
|
// ---- Load tile lambda: predicated cp.async ----
|
||||||
|
auto load_tile = [&](int ti, int buf) {
|
||||||
|
int kv0 = ti * Traits::BC;
|
||||||
|
bf16* dK = sK + buf * Traits::BC * Traits::LD;
|
||||||
|
bf16* dV = sV + buf * Traits::BC * Traits::LD;
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = lane * Traits::VEC; i < Traits::TOTAL;
|
||||||
|
i += Traits::NUM_THREADS * Traits::VEC) {
|
||||||
|
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
|
||||||
|
int kc = kv0 + r;
|
||||||
|
bool valid = kc < p.kv_len;
|
||||||
|
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
|
||||||
|
int g_off = kv_base + kc * p.kv_stride_l + d * p.kv_stride_d;
|
||||||
|
cp_async_16_pred(&dK[off], &p.k[g_off], valid);
|
||||||
|
cp_async_16_pred(&dV[off], &p.v[g_off], valid);
|
||||||
|
}
|
||||||
|
cp_async_commit();
|
||||||
|
};
|
||||||
|
|
||||||
|
// ---- Multi-stage cp.async pipeline ----
|
||||||
|
// Prologue loads STAGES tiles; each loop iteration waits only for the
|
||||||
|
// oldest outstanding group (wait_group<STAGES-1>) so the STAGES-1 newer
|
||||||
|
// tile loads stay in flight and overlap with the current tile's compute.
|
||||||
|
constexpr int STAGES = Traits::STAGES;
|
||||||
|
const int ntiles = ti_end - ti_begin;
|
||||||
|
|
||||||
|
auto process_tile = [&](int it, int buf) {
|
||||||
|
const bf16* bK = sK + buf * Traits::BC * Traits::LD;
|
||||||
|
const bf16* bV = sV + buf * Traits::BC * Traits::LD;
|
||||||
|
int kv0 = (ti_begin + it) * Traits::BC;
|
||||||
|
|
||||||
|
float Sacc[Traits::NC8][4];
|
||||||
|
mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
|
||||||
|
|
||||||
|
#pragma unroll
|
||||||
|
for (int n8 = 0; n8 < Traits::NC8; n8++)
|
||||||
|
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
||||||
|
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
||||||
|
|
||||||
|
// Decode: q_len=1, so qrow0=qrow1=0
|
||||||
|
int maxc = IsCausal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
|
||||||
|
mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc,
|
||||||
|
0, 0,
|
||||||
|
p.mask_b_stride, 0, 0,
|
||||||
|
batch, 0,
|
||||||
|
p.mask,
|
||||||
|
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||||
|
|
||||||
|
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
|
||||||
|
};
|
||||||
|
|
||||||
|
if (ntiles >= STAGES) {
|
||||||
|
#pragma unroll
|
||||||
|
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 ----
|
||||||
|
auto split_slot = [&](int h) -> size_t {
|
||||||
|
size_t bh = (size_t)batch * p.q_head + h;
|
||||||
|
return bh * MAX_SPLITS + split;
|
||||||
|
};
|
||||||
|
#pragma unroll
|
||||||
|
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
|
||||||
|
int d = dn8 * 8 + 2 * tid4;
|
||||||
|
int r0 = gid, r1 = gid + 8;
|
||||||
|
if (r0 < G) {
|
||||||
|
int h = q_head0 + r0;
|
||||||
|
float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
|
||||||
|
op[d] = Oacc[dn8][0];
|
||||||
|
op[d + 1] = Oacc[dn8][1];
|
||||||
|
}
|
||||||
|
if (r1 < G) {
|
||||||
|
int h = q_head0 + r1;
|
||||||
|
float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
|
||||||
|
op[d] = Oacc[dn8][2];
|
||||||
|
op[d + 1] = Oacc[dn8][3];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if (tid4 == 0) {
|
||||||
|
int r0 = gid, r1 = gid + 8;
|
||||||
|
if (r0 < G) {
|
||||||
|
int h = q_head0 + r0;
|
||||||
|
float* mp = p.ml_part + split_slot(h) * 2;
|
||||||
|
mp[0] = m0; mp[1] = l0;
|
||||||
|
}
|
||||||
|
if (r1 < G) {
|
||||||
|
int h = q_head0 + r1;
|
||||||
|
float* mp = p.ml_part + split_slot(h) * 2;
|
||||||
|
mp[0] = m1; mp[1] = l1;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,195 @@
|
|||||||
|
#pragma once
|
||||||
|
// Shared attention dispatchers — used by both production .cu and test .cu.
|
||||||
|
// No torch dependency; pure CUDA.
|
||||||
|
|
||||||
|
#include <cuda_runtime.h>
|
||||||
|
#include <algorithm>
|
||||||
|
#include "attn_warp_utils.cuh"
|
||||||
|
#include "attn_prefill_split_q.cuh"
|
||||||
|
#include "attn_decode_split_kv.cuh"
|
||||||
|
#include "attn_paged_decode_split_kv.cuh"
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
#include "attn_prefill_split_q_mma.cuh"
|
||||||
|
#include "attn_decode_split_kv_mma.cuh"
|
||||||
|
#include "attn_paged_decode_split_kv_mma.cuh"
|
||||||
|
#endif
|
||||||
|
|
||||||
|
// Split-KV: compute number of splits to fill all SMs for small-batch decode.
|
||||||
|
// Caps splits so each split processes at least `min_tiles_per_split` tiles,
|
||||||
|
// avoiding excessive loop/prologue overhead when tiles are small.
|
||||||
|
inline int compute_num_splits(int base_blocks, int tiles_total,
|
||||||
|
int min_tiles_per_split = 1) {
|
||||||
|
int sm_count = 0;
|
||||||
|
cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0);
|
||||||
|
int n = (2 * sm_count + base_blocks - 1) / base_blocks;
|
||||||
|
int max_by_work = tiles_total / min_tiles_per_split;
|
||||||
|
return std::max(1, std::min(n, std::min(max_by_work, MAX_SPLITS)));
|
||||||
|
}
|
||||||
|
|
||||||
|
// ======================================================================
|
||||||
|
// Prefill
|
||||||
|
// ======================================================================
|
||||||
|
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
|
static inline void launch_prefill_mma(AttentionParams<bf16>& p) {
|
||||||
|
constexpr int WARPS = 4;
|
||||||
|
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
|
||||||
|
using Traits = KernelTraits<HEAD_DIM, BC, WARPS, 2>;
|
||||||
|
dim3 grid((p.q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS), p.q_head, p.batch);
|
||||||
|
dim3 block(Traits::NUM_THREADS);
|
||||||
|
attn_prefill_split_q_mma_kernel<Traits, IsCausal, HasMask><<<grid, block>>>(p);
|
||||||
|
}
|
||||||
|
#endif
|
||||||
|
|
||||||
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
|
static inline void launch_prefill_scalar(AttentionParams<bf16>& p) {
|
||||||
|
constexpr int G = 8, ROWS = 32, P_BC = 32;
|
||||||
|
dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch);
|
||||||
|
dim3 block(G, ROWS);
|
||||||
|
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC, IsCausal, HasMask><<<grid, block>>>(p);
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
static inline void dispatch_prefill(AttentionParams<bf16>& p) {
|
||||||
|
bool is_causal = (p.causal_offset >= 0);
|
||||||
|
bool has_mask = (p.use_mask && p.mask);
|
||||||
|
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
if (is_causal) {
|
||||||
|
if (has_mask) launch_prefill_mma<HEAD_DIM, true, true>(p);
|
||||||
|
else launch_prefill_mma<HEAD_DIM, true, false>(p);
|
||||||
|
} else {
|
||||||
|
if (has_mask) launch_prefill_mma<HEAD_DIM, false, true>(p);
|
||||||
|
else launch_prefill_mma<HEAD_DIM, false, false>(p);
|
||||||
|
}
|
||||||
|
#else
|
||||||
|
if (is_causal) {
|
||||||
|
if (has_mask) launch_prefill_scalar<HEAD_DIM, true, true>(p);
|
||||||
|
else launch_prefill_scalar<HEAD_DIM, true, false>(p);
|
||||||
|
} else {
|
||||||
|
if (has_mask) launch_prefill_scalar<HEAD_DIM, false, true>(p);
|
||||||
|
else launch_prefill_scalar<HEAD_DIM, false, false>(p);
|
||||||
|
}
|
||||||
|
#endif
|
||||||
|
}
|
||||||
|
|
||||||
|
// ======================================================================
|
||||||
|
// Decode
|
||||||
|
// ======================================================================
|
||||||
|
|
||||||
|
#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>
|
||||||
|
static inline void launch_decode_mma(AttentionParams<bf16>& p, int group_size) {
|
||||||
|
int G = p.q_head / p.kv_head;
|
||||||
|
constexpr int MAX_G = 16;
|
||||||
|
int num_passes = (G + MAX_G - 1) / MAX_G;
|
||||||
|
constexpr int BC = 16;
|
||||||
|
int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||||
|
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total, 2);
|
||||||
|
constexpr int STAGES = 2;
|
||||||
|
using Traits = KernelTraits<HEAD_DIM, BC, 1, STAGES>;
|
||||||
|
dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
|
||||||
|
attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask><<<grid, 32>>>(p);
|
||||||
|
}
|
||||||
|
#endif
|
||||||
|
|
||||||
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
|
static inline void launch_decode_scalar(AttentionParams<bf16>& p, int group_size) {
|
||||||
|
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
|
||||||
|
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||||
|
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
|
||||||
|
int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit
|
||||||
|
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
|
||||||
|
dim3 block(32, g);
|
||||||
|
attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
static inline void dispatch_decode(AttentionParams<bf16>& p) {
|
||||||
|
bool is_causal = (p.causal_offset >= 0);
|
||||||
|
bool has_mask = (p.use_mask && p.mask);
|
||||||
|
int group_size = p.q_head / p.kv_head;
|
||||||
|
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
if (is_causal) {
|
||||||
|
if (has_mask) launch_decode_mma<HEAD_DIM, true, true>(p, group_size);
|
||||||
|
else launch_decode_mma<HEAD_DIM, true, false>(p, group_size);
|
||||||
|
} else {
|
||||||
|
if (has_mask) launch_decode_mma<HEAD_DIM, false, true>(p, group_size);
|
||||||
|
else launch_decode_mma<HEAD_DIM, false, false>(p, group_size);
|
||||||
|
}
|
||||||
|
#else
|
||||||
|
if (is_causal) {
|
||||||
|
if (has_mask) launch_decode_scalar<HEAD_DIM, true, true>(p, group_size);
|
||||||
|
else launch_decode_scalar<HEAD_DIM, true, false>(p, group_size);
|
||||||
|
} else {
|
||||||
|
if (has_mask) launch_decode_scalar<HEAD_DIM, false, true>(p, group_size);
|
||||||
|
else launch_decode_scalar<HEAD_DIM, false, false>(p, group_size);
|
||||||
|
}
|
||||||
|
#endif
|
||||||
|
|
||||||
|
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||||
|
}
|
||||||
|
|
||||||
|
// ======================================================================
|
||||||
|
// Paged Decode
|
||||||
|
// ======================================================================
|
||||||
|
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
|
static inline void launch_paged_decode_mma(PagedAttentionParams<bf16>& p, int group_size) {
|
||||||
|
int G = p.q_head / p.kv_head;
|
||||||
|
constexpr int MAX_G = 16;
|
||||||
|
constexpr int BC = 16;
|
||||||
|
int num_passes = (G + MAX_G - 1) / MAX_G;
|
||||||
|
int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||||
|
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total, 2);
|
||||||
|
constexpr int STAGES = 2;
|
||||||
|
using Traits = KernelTraits<HEAD_DIM, BC, 1, STAGES>;
|
||||||
|
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);
|
||||||
|
}
|
||||||
|
#endif
|
||||||
|
|
||||||
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
|
static inline void launch_paged_decode_scalar(PagedAttentionParams<bf16>& p, int group_size) {
|
||||||
|
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);
|
||||||
|
int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit
|
||||||
|
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
|
||||||
|
dim3 block(32, g);
|
||||||
|
paged_attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
static inline void dispatch_paged_decode(PagedAttentionParams<bf16>& p) {
|
||||||
|
bool is_causal = (p.causal_offset >= 0);
|
||||||
|
bool has_mask = (p.use_mask && p.mask);
|
||||||
|
int group_size = p.q_head / p.kv_head;
|
||||||
|
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
if (is_causal) {
|
||||||
|
if (has_mask) launch_paged_decode_mma<HEAD_DIM, true, true>(p, group_size);
|
||||||
|
else launch_paged_decode_mma<HEAD_DIM, true, false>(p, group_size);
|
||||||
|
} else {
|
||||||
|
if (has_mask) launch_paged_decode_mma<HEAD_DIM, false, true>(p, group_size);
|
||||||
|
else launch_paged_decode_mma<HEAD_DIM, false, false>(p, group_size);
|
||||||
|
}
|
||||||
|
#else
|
||||||
|
if (is_causal) {
|
||||||
|
if (has_mask) launch_paged_decode_scalar<HEAD_DIM, true, true>(p, group_size);
|
||||||
|
else launch_paged_decode_scalar<HEAD_DIM, true, false>(p, group_size);
|
||||||
|
} else {
|
||||||
|
if (has_mask) launch_paged_decode_scalar<HEAD_DIM, false, true>(p, group_size);
|
||||||
|
else launch_paged_decode_scalar<HEAD_DIM, false, false>(p, group_size);
|
||||||
|
}
|
||||||
|
#endif
|
||||||
|
|
||||||
|
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||||
|
}
|
||||||
@@ -0,0 +1,187 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <float.h>
|
||||||
|
#include <torch/extension.h>
|
||||||
|
#include <c10/cuda/CUDAGuard.h>
|
||||||
|
#include "attn_common.h"
|
||||||
|
#include "attn_warp_utils.cuh"
|
||||||
|
|
||||||
|
using bf16 = __nv_bfloat16;
|
||||||
|
|
||||||
|
// Dispatch head_dim: shared macro — avoids C++20 lambda template syntax.
|
||||||
|
// Usage: DISPATCH_HEAD_DIM(hd, fn, arg)
|
||||||
|
// Expands to: fn<32>(arg); fn<64>(arg); etc.
|
||||||
|
#define DISPATCH_HEAD_DIM(hd, fn, arg) \
|
||||||
|
switch (hd) { \
|
||||||
|
case 32: fn<32>(arg); break; \
|
||||||
|
case 64: fn<64>(arg); break; \
|
||||||
|
case 128: fn<128>(arg); break; \
|
||||||
|
case 256: fn<256>(arg); break; \
|
||||||
|
default: \
|
||||||
|
TORCH_CHECK(false, "unsupported head_dim ", hd, \
|
||||||
|
" (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>
|
||||||
|
inline void alloc_split_partials(P& p) {
|
||||||
|
auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA);
|
||||||
|
auto o_part = torch::empty(at::IntArrayRef{p.batch, p.q_head, MAX_SPLITS, p.head_dim}, 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.ml_part = (float*)ml_part.data_ptr();
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- Shared Q-dims + strides extraction ----
|
||||||
|
template <typename P>
|
||||||
|
inline void extract_q_dims_and_strides(torch::Tensor& q, int64_t layout, P& p) {
|
||||||
|
if (layout == 1) q = q.transpose(1, 2);
|
||||||
|
p.batch = (int)q.size(0);
|
||||||
|
p.q_head = (int)q.size(1);
|
||||||
|
p.q_len = (int)q.size(2);
|
||||||
|
p.head_dim = (int)q.size(3);
|
||||||
|
p.q_stride_b = (int)q.stride(0);
|
||||||
|
p.q_stride_h = (int)q.stride(1);
|
||||||
|
p.q_stride_l = (int)q.stride(2);
|
||||||
|
p.q_stride_d = (int)q.stride(3);
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- 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>
|
||||||
|
inline void pack_mask(const c10::optional<torch::Tensor>& mask, P& p) {
|
||||||
|
if (p.use_mask) {
|
||||||
|
auto m = mask.value();
|
||||||
|
TORCH_CHECK(m.is_cuda(), "mask must be on CUDA");
|
||||||
|
TORCH_CHECK(m.dtype() == torch::kBool, "mask must be bool");
|
||||||
|
TORCH_CHECK(m.size(0) == p.batch, "mask batch mismatch");
|
||||||
|
TORCH_CHECK(m.size(m.dim() - 1) == p.kv_len, "mask kv_len mismatch");
|
||||||
|
if (m.dim() == 2) {
|
||||||
|
p.mask_b_stride = (int)m.stride(0);
|
||||||
|
p.mask_h_stride = 0;
|
||||||
|
p.mask_q_stride = 0;
|
||||||
|
} else if (m.dim() == 3) {
|
||||||
|
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_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 {
|
||||||
|
TORCH_CHECK(false, "mask must be 2D, 3D, or 4D");
|
||||||
|
}
|
||||||
|
p.mask = m.data_ptr<bool>();
|
||||||
|
} else {
|
||||||
|
p.mask = nullptr;
|
||||||
|
p.mask_b_stride = 0;
|
||||||
|
p.mask_h_stride = 0;
|
||||||
|
p.mask_q_stride = 0;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- attn_pack_params (contiguous KV) ----
|
||||||
|
template<typename T>
|
||||||
|
inline void attn_pack_params(
|
||||||
|
torch::Tensor q,
|
||||||
|
torch::Tensor k,
|
||||||
|
torch::Tensor v,
|
||||||
|
c10::optional<torch::Tensor> mask,
|
||||||
|
int64_t causal_offset,
|
||||||
|
double scale,
|
||||||
|
int64_t layout,
|
||||||
|
AttentionParams<T>& p
|
||||||
|
) {
|
||||||
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||||
|
|
||||||
|
TORCH_CHECK(q.is_cuda() && k.is_cuda() && v.is_cuda());
|
||||||
|
TORCH_CHECK(q.dtype() == torch::kBFloat16);
|
||||||
|
TORCH_CHECK(k.dtype() == torch::kBFloat16);
|
||||||
|
TORCH_CHECK(v.dtype() == torch::kBFloat16);
|
||||||
|
TORCH_CHECK(k.sizes() == v.sizes(), "K and V must have identical shapes");
|
||||||
|
TORCH_CHECK(q.dim() == 4 && k.dim() == 4, "Q/K/V must be 4D");
|
||||||
|
|
||||||
|
extract_q_dims_and_strides(q, layout, p);
|
||||||
|
|
||||||
|
if (layout == 1) k = k.transpose(1, 2), v = v.transpose(1, 2);
|
||||||
|
|
||||||
|
p.kv_head = (int)k.size(1);
|
||||||
|
p.kv_len = (int)k.size(2);
|
||||||
|
TORCH_CHECK(k.size(3) == p.head_dim, "K/V head_dim must match Q");
|
||||||
|
|
||||||
|
p.kv_stride_b = (int)k.stride(0);
|
||||||
|
p.kv_stride_h = (int)k.stride(1);
|
||||||
|
p.kv_stride_l = (int)k.stride(2);
|
||||||
|
p.kv_stride_d = (int)k.stride(3);
|
||||||
|
|
||||||
|
p.causal_offset = (int)causal_offset;
|
||||||
|
p.use_mask = mask.has_value() ? 1 : 0;
|
||||||
|
p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
|
||||||
|
|
||||||
|
p.q = (const T*)q.data_ptr();
|
||||||
|
p.k = (const T*)k.data_ptr();
|
||||||
|
p.v = (const T*)v.data_ptr();
|
||||||
|
p.o = nullptr;
|
||||||
|
p.o_part = nullptr;
|
||||||
|
p.ml_part = nullptr;
|
||||||
|
|
||||||
|
pack_mask(mask, p);
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- attn_pack_paged_params ----
|
||||||
|
template<typename T>
|
||||||
|
inline void attn_pack_paged_params(
|
||||||
|
torch::Tensor q,
|
||||||
|
torch::Tensor page_table,
|
||||||
|
torch::Tensor k_cache,
|
||||||
|
torch::Tensor v_cache,
|
||||||
|
int64_t page_size,
|
||||||
|
int64_t kv_len,
|
||||||
|
c10::optional<torch::Tensor> mask,
|
||||||
|
int64_t causal_offset,
|
||||||
|
double scale,
|
||||||
|
int64_t layout,
|
||||||
|
PagedAttentionParams<T>& p
|
||||||
|
) {
|
||||||
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||||
|
|
||||||
|
TORCH_CHECK(q.is_cuda() && page_table.is_cuda() && k_cache.is_cuda() && v_cache.is_cuda());
|
||||||
|
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16");
|
||||||
|
TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bf16");
|
||||||
|
TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bf16");
|
||||||
|
TORCH_CHECK(page_table.dtype() == torch::kLong, "page_table must be int64");
|
||||||
|
TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must have identical shapes");
|
||||||
|
|
||||||
|
extract_q_dims_and_strides(q, layout, p);
|
||||||
|
|
||||||
|
p.kv_head = (int)k_cache.size(2);
|
||||||
|
p.kv_len = (int)kv_len;
|
||||||
|
p.page_size = (int)page_size;
|
||||||
|
p.max_pages = (int)page_table.size(1);
|
||||||
|
|
||||||
|
TORCH_CHECK(q.size(2) == 1, "Q seq_len must be 1 (decode)");
|
||||||
|
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
|
||||||
|
TORCH_CHECK(k_cache.size(1) == page_size,
|
||||||
|
"k_cache dim 1 must equal page_size, got ",
|
||||||
|
k_cache.size(1), " vs ", page_size);
|
||||||
|
|
||||||
|
p.causal_offset = (int)causal_offset;
|
||||||
|
p.use_mask = (mask.has_value() && mask.value().defined()) ? 1 : 0;
|
||||||
|
p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
|
||||||
|
|
||||||
|
p.page_table = page_table.data_ptr<int64_t>();
|
||||||
|
p.k_cache = (const T*)k_cache.data_ptr();
|
||||||
|
p.v_cache = (const T*)v_cache.data_ptr();
|
||||||
|
p.q = (const T*)q.data_ptr();
|
||||||
|
p.o = nullptr;
|
||||||
|
p.o_part = nullptr;
|
||||||
|
p.ml_part = nullptr;
|
||||||
|
|
||||||
|
pack_mask(mask, p);
|
||||||
|
}
|
||||||
@@ -0,0 +1,297 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <cfloat>
|
||||||
|
#include <cuda_fp16.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.
|
||||||
|
//
|
||||||
|
// Bundles all dimension-dependent constants so device functions only need a
|
||||||
|
// single Traits template parameter rather than scattered <KD, NC8, KT2, ...>.
|
||||||
|
// ============================================================================
|
||||||
|
template <int HEAD_DIM_, int BC_, int WARPS_, int STAGES_>
|
||||||
|
struct KernelTraits {
|
||||||
|
static constexpr int HEAD_DIM = HEAD_DIM_;
|
||||||
|
static constexpr int BC = BC_; // K/V tile size along seq dim
|
||||||
|
static constexpr int WARPS = WARPS_; // warps per block
|
||||||
|
static constexpr int STAGES = STAGES_; // double-buffer stages (1 or 2)
|
||||||
|
|
||||||
|
static constexpr int BR = 16; // Q rows per warp (mma M=16)
|
||||||
|
|
||||||
|
// Derived: mma.sync.m16n8k16 tile counts
|
||||||
|
static constexpr int KD = HEAD_DIM / 16; // Q/K k-slides
|
||||||
|
static constexpr int NC8 = BC / 8; // S n-tiles (N=8)
|
||||||
|
static constexpr int KT2 = BC / 16; // P k-tiles (K=16)
|
||||||
|
static constexpr int DN8 = HEAD_DIM / 8; // O n-tiles (N=8)
|
||||||
|
|
||||||
|
static constexpr int LD = HEAD_DIM; // smem leading dim
|
||||||
|
|
||||||
|
// XOR swizzle chunk bits for ldmatrix bank-conflict avoidance.
|
||||||
|
// mask = log2(LD/8) bits, clamped to stay within LD.
|
||||||
|
static constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1);
|
||||||
|
|
||||||
|
static constexpr int NUM_THREADS = WARPS * 32;
|
||||||
|
static constexpr int VEC = 8; // bf16 per cp.async unit (16 bytes)
|
||||||
|
static constexpr int TOTAL = BC * HEAD_DIM; // total elements per tile
|
||||||
|
};
|
||||||
|
|
||||||
|
// ---- PTX wrappers ----
|
||||||
|
using bf16 = __nv_bfloat16;
|
||||||
|
|
||||||
|
__device__ __forceinline__ void mma16816(float* d, const unsigned* a,
|
||||||
|
const unsigned* b, const float* c) {
|
||||||
|
asm volatile(
|
||||||
|
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
|
||||||
|
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%10,%11,%12,%13};"
|
||||||
|
: "=f"(d[0]), "=f"(d[1]), "=f"(d[2]), "=f"(d[3])
|
||||||
|
: "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]),
|
||||||
|
"f"(c[0]), "f"(c[1]), "f"(c[2]), "f"(c[3]));
|
||||||
|
}
|
||||||
|
|
||||||
|
// read two adjacent bf16 from smem as one packed .b32 (elem0 low, elem1 high)
|
||||||
|
__device__ __forceinline__ unsigned ld2(const bf16* p) {
|
||||||
|
return *reinterpret_cast<const unsigned*>(p);
|
||||||
|
}
|
||||||
|
|
||||||
|
// pack two floats into one bf16x2 as .b32
|
||||||
|
__device__ __forceinline__ unsigned pk2(float a, float b) {
|
||||||
|
__nv_bfloat162 v = __floats2bfloat162_rn(a, b);
|
||||||
|
return *reinterpret_cast<unsigned*>(&v);
|
||||||
|
}
|
||||||
|
|
||||||
|
// pack two (non-contiguous) bf16 into one .b32
|
||||||
|
__device__ __forceinline__ unsigned pkb(bf16 a, bf16 b) {
|
||||||
|
__nv_bfloat162 v;
|
||||||
|
v.x = a;
|
||||||
|
v.y = b;
|
||||||
|
return *reinterpret_cast<unsigned*>(&v);
|
||||||
|
}
|
||||||
|
|
||||||
|
// ldmatrix: cooperatively load mma fragments from smem (one instruction per
|
||||||
|
// 16x16 / 16x8 tile) with the exact register layout mma expects.
|
||||||
|
__device__ __forceinline__ void ldmatrix_x4(unsigned* r, const bf16* p) {
|
||||||
|
unsigned a = __cvta_generic_to_shared(p);
|
||||||
|
asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];"
|
||||||
|
: "=r"(r[0]), "=r"(r[1]), "=r"(r[2]), "=r"(r[3])
|
||||||
|
: "r"(a));
|
||||||
|
}
|
||||||
|
__device__ __forceinline__ void ldmatrix_x2(unsigned* r, const bf16* p) {
|
||||||
|
unsigned a = __cvta_generic_to_shared(p);
|
||||||
|
asm volatile("ldmatrix.sync.aligned.m8n8.x2.shared.b16 {%0,%1}, [%2];"
|
||||||
|
: "=r"(r[0]), "=r"(r[1])
|
||||||
|
: "r"(a));
|
||||||
|
}
|
||||||
|
__device__ __forceinline__ void ldmatrix_x2_trans(unsigned* r, const bf16* p) {
|
||||||
|
unsigned a = __cvta_generic_to_shared(p);
|
||||||
|
asm volatile("ldmatrix.sync.aligned.m8n8.x2.trans.shared.b16 {%0,%1}, [%2];"
|
||||||
|
: "=r"(r[0]), "=r"(r[1])
|
||||||
|
: "r"(a));
|
||||||
|
}
|
||||||
|
|
||||||
|
// XOR swizzle for shared-memory column at 8-bf16 chunk granularity.
|
||||||
|
__device__ __forceinline__ int swiz_col(int d, int r, int mask = 7) {
|
||||||
|
return ((d >> 3) ^ (r & mask)) << 3 | (d & 7);
|
||||||
|
}
|
||||||
|
|
||||||
|
// cp.async: copy 16 bytes (8 bf16) from global to shared memory directly.
|
||||||
|
__device__ __forceinline__ void cp_async_16(bf16* smem_ptr, const void* gmem_ptr) {
|
||||||
|
unsigned smem_addr = __cvta_generic_to_shared(smem_ptr);
|
||||||
|
asm volatile("cp.async.ca.shared.global [%0], [%1], 16;"
|
||||||
|
:: "r"(smem_addr), "l"(gmem_ptr));
|
||||||
|
}
|
||||||
|
|
||||||
|
// Predicated cp.async: copy 16 bytes when `pred`, otherwise zero-fill.
|
||||||
|
// src_size=0 → no bytes read from src, so out-of-bounds src address is safe.
|
||||||
|
__device__ __forceinline__ void cp_async_16_pred(bf16* smem_ptr,
|
||||||
|
const void* gmem_ptr,
|
||||||
|
bool pred) {
|
||||||
|
unsigned smem_addr = __cvta_generic_to_shared(smem_ptr);
|
||||||
|
int src_size = pred ? 16 : 0;
|
||||||
|
asm volatile("cp.async.ca.shared.global [%0], [%1], 16, %2;"
|
||||||
|
:: "r"(smem_addr), "l"(gmem_ptr), "r"(src_size));
|
||||||
|
}
|
||||||
|
|
||||||
|
__device__ __forceinline__ void cp_async_commit() {
|
||||||
|
asm volatile("cp.async.commit_group;");
|
||||||
|
}
|
||||||
|
|
||||||
|
__device__ __forceinline__ void cp_async_wait_all() {
|
||||||
|
asm volatile("cp.async.wait_all;");
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int N>
|
||||||
|
__device__ __forceinline__ void cp_async_wait_group() {
|
||||||
|
asm volatile("cp.async.wait_group %0;" :: "n"(N));
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
// Q-load: load query rows directly from global memory into mma A-operand
|
||||||
|
// register layout. One call replaces ~15 duplicated lines in each MMA kernel.
|
||||||
|
// stride_row is p.q_stride_h for decode (q_len=1, G heads) or
|
||||||
|
// p.q_stride_l for prefill (multi-q rows).
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
template <int KD>
|
||||||
|
__device__ inline void load_q_mma_frags(
|
||||||
|
const bf16* __restrict__ q,
|
||||||
|
int stride_row,
|
||||||
|
int stride_d,
|
||||||
|
int qra, int qrb,
|
||||||
|
bool va, bool vb,
|
||||||
|
int tid4,
|
||||||
|
unsigned Qa[KD][4])
|
||||||
|
{
|
||||||
|
#pragma unroll
|
||||||
|
for (int kt = 0; kt < KD; kt++) {
|
||||||
|
int c = kt * 16 + tid4 * 2;
|
||||||
|
const unsigned* pau = reinterpret_cast<const unsigned*>(
|
||||||
|
&q[qra * stride_row + c * stride_d]);
|
||||||
|
const unsigned* pbu = reinterpret_cast<const unsigned*>(
|
||||||
|
&q[qrb * stride_row + c * stride_d]);
|
||||||
|
Qa[kt][0] = va ? pau[0] : 0u;
|
||||||
|
Qa[kt][1] = vb ? pbu[0] : 0u;
|
||||||
|
Qa[kt][2] = va ? pau[4] : 0u;
|
||||||
|
Qa[kt][3] = vb ? pbu[4] : 0u;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
// S = Q @ K^T (Qa pre-loaded by the caller; scale applied post-mma in the
|
||||||
|
// caller to avoid bf16 precision loss).
|
||||||
|
// Traits provides KD, NC8, LD, and SWIZ_MASK.
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
template <typename Traits>
|
||||||
|
__device__ inline void mma_compute_scores(
|
||||||
|
const unsigned Qa[Traits::KD][4],
|
||||||
|
const bf16* __restrict__ sK,
|
||||||
|
int lane,
|
||||||
|
float Sacc[Traits::NC8][4])
|
||||||
|
{
|
||||||
|
#pragma unroll
|
||||||
|
for (int n8 = 0; n8 < Traits::NC8; n8++) {
|
||||||
|
Sacc[n8][0] = Sacc[n8][1] = Sacc[n8][2] = Sacc[n8][3] = 0.0f;
|
||||||
|
int krow_l = n8 * 8 + (lane & 7);
|
||||||
|
int kcol_h = (lane & 8) ? 8 : 0;
|
||||||
|
#pragma unroll
|
||||||
|
for (int kt = 0; kt < Traits::KD; kt++) {
|
||||||
|
unsigned b[2];
|
||||||
|
ldmatrix_x2(b, &sK[krow_l * Traits::LD
|
||||||
|
+ swiz_col(kt * 16 + kcol_h, krow_l, Traits::SWIZ_MASK)]);
|
||||||
|
mma16816(Sacc[n8], Qa[kt], b, Sacc[n8]);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
// Online softmax + Oacc rescale for one K/V tile.
|
||||||
|
//
|
||||||
|
// HasMask is a compile-time template bool: when false, the mask branch is
|
||||||
|
// entirely dead-code-eliminated from the inner unrolled loop.
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
template <typename Traits, bool HasMask>
|
||||||
|
__device__ inline void mma_softmax_tile(
|
||||||
|
int kv0,
|
||||||
|
int maxc0, int maxc1,
|
||||||
|
int qrow0, int qrow1,
|
||||||
|
int mask_b_stride, int mask_h_stride, int mask_q_stride,
|
||||||
|
int mask_batch, int mask_head,
|
||||||
|
const bool* __restrict__ mask,
|
||||||
|
float Sacc[Traits::NC8][4],
|
||||||
|
float Oacc[Traits::DN8][4],
|
||||||
|
float& m0, float& m1,
|
||||||
|
float& l0, float& l1,
|
||||||
|
int lane)
|
||||||
|
{
|
||||||
|
int tid4 = lane & 3;
|
||||||
|
|
||||||
|
float rmax0 = -FLT_MAX, rmax1 = -FLT_MAX;
|
||||||
|
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 + mask_head * mask_h_stride + qrow1 * mask_q_stride;
|
||||||
|
#pragma unroll
|
||||||
|
for (int n8 = 0; n8 < Traits::NC8; n8++) {
|
||||||
|
int cc = kv0 + n8 * 8 + 2 * tid4;
|
||||||
|
int c1 = cc + 1;
|
||||||
|
bool b0 = (cc >= maxc0) || (HasMask && !mask[mask_base0 + cc]);
|
||||||
|
bool b1 = (c1 >= maxc0) || (HasMask && !mask[mask_base0 + c1]);
|
||||||
|
bool b2 = (cc >= maxc1) || (HasMask && !mask[mask_base1 + cc]);
|
||||||
|
bool b3 = (c1 >= maxc1) || (HasMask && !mask[mask_base1 + c1]);
|
||||||
|
float s0 = b0 ? -FLT_MAX : Sacc[n8][0];
|
||||||
|
float s1 = b1 ? -FLT_MAX : Sacc[n8][1];
|
||||||
|
float s2 = b2 ? -FLT_MAX : Sacc[n8][2];
|
||||||
|
float s3 = b3 ? -FLT_MAX : Sacc[n8][3];
|
||||||
|
Sacc[n8][0] = s0; Sacc[n8][1] = s1;
|
||||||
|
Sacc[n8][2] = s2; Sacc[n8][3] = s3;
|
||||||
|
rmax0 = fmaxf(rmax0, fmaxf(s0, s1));
|
||||||
|
rmax1 = fmaxf(rmax1, fmaxf(s2, s3));
|
||||||
|
}
|
||||||
|
rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 1));
|
||||||
|
rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 2));
|
||||||
|
rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 1));
|
||||||
|
rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 2));
|
||||||
|
|
||||||
|
float nm0 = fmaxf(m0, rmax0), nm1 = fmaxf(m1, rmax1);
|
||||||
|
float corr0 = __expf(m0 - nm0);
|
||||||
|
float corr1 = __expf(m1 - nm1);
|
||||||
|
float pn0 = (nm0 == -FLT_MAX) ? 0.0f : 1.0f;
|
||||||
|
float pn1 = (nm1 == -FLT_MAX) ? 0.0f : 1.0f;
|
||||||
|
|
||||||
|
float rsum0 = 0.0f, rsum1 = 0.0f;
|
||||||
|
#pragma unroll
|
||||||
|
for (int n8 = 0; n8 < Traits::NC8; n8++) {
|
||||||
|
float p0 = pn0 * __expf(Sacc[n8][0] - nm0);
|
||||||
|
float p1 = pn0 * __expf(Sacc[n8][1] - nm0);
|
||||||
|
float p2 = pn1 * __expf(Sacc[n8][2] - nm1);
|
||||||
|
float p3 = pn1 * __expf(Sacc[n8][3] - nm1);
|
||||||
|
Sacc[n8][0] = p0; Sacc[n8][1] = p1;
|
||||||
|
Sacc[n8][2] = p2; Sacc[n8][3] = p3;
|
||||||
|
rsum0 += p0 + p1;
|
||||||
|
rsum1 += p2 + p3;
|
||||||
|
}
|
||||||
|
rsum0 += __shfl_xor_sync(0xFFFFFFFF, rsum0, 1);
|
||||||
|
rsum0 += __shfl_xor_sync(0xFFFFFFFF, rsum0, 2);
|
||||||
|
rsum1 += __shfl_xor_sync(0xFFFFFFFF, rsum1, 1);
|
||||||
|
rsum1 += __shfl_xor_sync(0xFFFFFFFF, rsum1, 2);
|
||||||
|
l0 = l0 * corr0 + rsum0;
|
||||||
|
l1 = l1 * corr1 + rsum1;
|
||||||
|
m0 = nm0; m1 = nm1;
|
||||||
|
|
||||||
|
#pragma unroll
|
||||||
|
for (int j = 0; j < Traits::DN8; j++) {
|
||||||
|
Oacc[j][0] *= corr0; Oacc[j][1] *= corr0;
|
||||||
|
Oacc[j][2] *= corr1; Oacc[j][3] *= corr1;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
// O += P @ V (Sacc must contain P = attention weights after softmax).
|
||||||
|
// Traits provides DN8, KT2, LD, and SWIZ_MASK.
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
template <typename Traits>
|
||||||
|
__device__ inline void mma_pv_accumulate(
|
||||||
|
float Sacc[][4],
|
||||||
|
const bf16* __restrict__ sV,
|
||||||
|
int lane,
|
||||||
|
float Oacc[Traits::DN8][4])
|
||||||
|
{
|
||||||
|
#pragma unroll
|
||||||
|
for (int kt2 = 0; kt2 < Traits::KT2; kt2++) {
|
||||||
|
unsigned Pa[4];
|
||||||
|
Pa[0] = pk2(Sacc[kt2 * 2][0], Sacc[kt2 * 2][1]);
|
||||||
|
Pa[1] = pk2(Sacc[kt2 * 2][2], Sacc[kt2 * 2][3]);
|
||||||
|
Pa[2] = pk2(Sacc[kt2 * 2 + 1][0], Sacc[kt2 * 2 + 1][1]);
|
||||||
|
Pa[3] = pk2(Sacc[kt2 * 2 + 1][2], Sacc[kt2 * 2 + 1][3]);
|
||||||
|
int vrow_l = kt2 * 16 + (lane & 15);
|
||||||
|
#pragma unroll
|
||||||
|
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
|
||||||
|
unsigned b[2];
|
||||||
|
ldmatrix_x2_trans(b, &sV[vrow_l * Traits::LD
|
||||||
|
+ swiz_col(dn8 * 8, vrow_l, Traits::SWIZ_MASK)]);
|
||||||
|
mma16816(Oacc[dn8], Pa, b, Oacc[dn8]);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,42 @@
|
|||||||
|
#include "attn_dispatchers.cuh"
|
||||||
|
#include "attn_entry_utils.cuh"
|
||||||
|
|
||||||
|
torch::Tensor attn_paged_decode(
|
||||||
|
torch::Tensor q,
|
||||||
|
torch::Tensor page_table,
|
||||||
|
torch::Tensor k_cache,
|
||||||
|
torch::Tensor v_cache,
|
||||||
|
int64_t page_size,
|
||||||
|
int64_t kv_len,
|
||||||
|
c10::optional<torch::Tensor> mask,
|
||||||
|
int64_t causal_offset,
|
||||||
|
double scale,
|
||||||
|
int64_t layout
|
||||||
|
) {
|
||||||
|
PagedAttentionParams<bf16> p;
|
||||||
|
attn_pack_paged_params(q, page_table, k_cache, v_cache,
|
||||||
|
page_size, kv_len, mask, causal_offset, scale, layout, p);
|
||||||
|
|
||||||
|
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
|
||||||
|
auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
|
||||||
|
p.o = (bf16*)O_view.data_ptr();
|
||||||
|
|
||||||
|
alloc_split_partials(p);
|
||||||
|
DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_decode, p);
|
||||||
|
return O;
|
||||||
|
}
|
||||||
|
|
||||||
|
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||||
|
m.def("attn_paged_decode", &attn_paged_decode,
|
||||||
|
py::arg("q"),
|
||||||
|
py::arg("page_table"),
|
||||||
|
py::arg("k_cache"),
|
||||||
|
py::arg("v_cache"),
|
||||||
|
py::arg("page_size"),
|
||||||
|
py::arg("kv_len"),
|
||||||
|
py::arg("mask") = py::none(),
|
||||||
|
py::arg("causal_offset") = -1,
|
||||||
|
py::arg("scale") = 0.0,
|
||||||
|
py::arg("layout") = 0,
|
||||||
|
"Paged GQA decode — split-KV with direct page-table access.");
|
||||||
|
}
|
||||||
@@ -0,0 +1,153 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
#include <float.h>
|
||||||
|
#include "attn_common.h"
|
||||||
|
#include "attn_warp_utils.cuh"
|
||||||
|
constexpr int PDC_CHUNK = 64;
|
||||||
|
|
||||||
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
|
__global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p) {
|
||||||
|
int batch = blockIdx.x / p.kv_head;
|
||||||
|
int kv_head = blockIdx.x % p.kv_head;
|
||||||
|
int split = blockIdx.z;
|
||||||
|
int group_size = blockDim.y;
|
||||||
|
int q_head = kv_head * group_size + threadIdx.y;
|
||||||
|
int lane = threadIdx.x;
|
||||||
|
int hd_per_thread = p.head_dim / 32;
|
||||||
|
|
||||||
|
float q_reg[8];
|
||||||
|
int q_off = batch * p.q_stride_b + q_head * p.q_stride_h
|
||||||
|
+ lane * hd_per_thread * p.q_stride_d;
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
|
q_reg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]);
|
||||||
|
|
||||||
|
float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f};
|
||||||
|
|
||||||
|
extern __shared__ __align__(16) bf16 k_smem[];
|
||||||
|
|
||||||
|
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
|
||||||
|
int chunks_per_split = (chunks_total + p.num_splits - 1) / p.num_splits;
|
||||||
|
int ch_begin = split * chunks_per_split;
|
||||||
|
int ch_end = min(chunks_total, ch_begin + chunks_per_split);
|
||||||
|
|
||||||
|
const int mask_base = batch * p.mask_b_stride + q_head * p.mask_h_stride;
|
||||||
|
|
||||||
|
for (int ci = ch_begin; ci < ch_end; ci++) {
|
||||||
|
int chunk_start = ci * PDC_CHUNK;
|
||||||
|
int this_chunk = min(PDC_CHUNK, p.kv_len - chunk_start);
|
||||||
|
|
||||||
|
int total = this_chunk * p.head_dim;
|
||||||
|
for (int i = threadIdx.y * 32 + lane; i < total;
|
||||||
|
i += blockDim.x * blockDim.y) {
|
||||||
|
int s = i / p.head_dim;
|
||||||
|
int d_dim = i % p.head_dim;
|
||||||
|
int pos = chunk_start + s;
|
||||||
|
int logical_page = pos / p.page_size;
|
||||||
|
int page_offset = pos % p.page_size;
|
||||||
|
int phys_page = p.page_table[batch * p.max_pages + logical_page];
|
||||||
|
if (phys_page >= 0) {
|
||||||
|
int64_t off = (int64_t)phys_page * p.page_size * p.kv_head * p.head_dim
|
||||||
|
+ (int64_t)page_offset * p.kv_head * p.head_dim
|
||||||
|
+ (int64_t)kv_head * p.head_dim
|
||||||
|
+ d_dim;
|
||||||
|
k_smem[i] = p.k_cache[off];
|
||||||
|
} else {
|
||||||
|
k_smem[i] = __float2bfloat16(0.0f);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
__syncthreads();
|
||||||
|
|
||||||
|
for (int s = 0; s < this_chunk; s++) {
|
||||||
|
float partial = 0.0f;
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
|
partial += q_reg[i] * __bfloat162float(
|
||||||
|
k_smem[s * p.head_dim + lane * hd_per_thread + i]);
|
||||||
|
partial = warp_reduce_sum(partial) * p.scale;
|
||||||
|
|
||||||
|
int kv_idx = chunk_start + s;
|
||||||
|
bool masked = false;
|
||||||
|
if constexpr (HasMask) {
|
||||||
|
if (!p.mask[mask_base + kv_idx])
|
||||||
|
masked = true;
|
||||||
|
}
|
||||||
|
if constexpr (IsCausal) {
|
||||||
|
if (kv_idx > p.causal_offset)
|
||||||
|
masked = true;
|
||||||
|
}
|
||||||
|
if (masked)
|
||||||
|
partial = -FLT_MAX;
|
||||||
|
|
||||||
|
float new_m = fmaxf(m, partial);
|
||||||
|
float alpha = __expf(m - new_m);
|
||||||
|
float beta = __expf(partial - new_m);
|
||||||
|
d = d * alpha + beta;
|
||||||
|
|
||||||
|
int pos = chunk_start + s;
|
||||||
|
int logical_page = pos / p.page_size;
|
||||||
|
int page_offset = pos % p.page_size;
|
||||||
|
int phys_page = p.page_table[batch * p.max_pages + logical_page];
|
||||||
|
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)page_offset * p.kv_head * p.head_dim
|
||||||
|
+ (int64_t)kv_head * p.head_dim;
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
|
acc_reg[i] = fmaf(acc_reg[i], alpha,
|
||||||
|
__bfloat162float(p.v_cache[v_base + lane * hd_per_thread + i]) * beta);
|
||||||
|
} else {
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
|
acc_reg[i] = fmaf(acc_reg[i], alpha, 0.0f);
|
||||||
|
}
|
||||||
|
m = new_m;
|
||||||
|
}
|
||||||
|
__syncthreads();
|
||||||
|
}
|
||||||
|
|
||||||
|
size_t bh = (size_t)batch * p.q_head + q_head;
|
||||||
|
size_t slot = bh * MAX_SPLITS + split;
|
||||||
|
int d0 = lane * hd_per_thread;
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
|
p.o_part[slot * p.head_dim + (d0 + i)] = acc_reg[i];
|
||||||
|
if (lane == 0) {
|
||||||
|
p.ml_part[slot * 2] = m;
|
||||||
|
p.ml_part[slot * 2 + 1] = d;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
__global__ void paged_attn_decode_combine_kernel(PagedAttentionParams<bf16> p) {
|
||||||
|
int bh = blockIdx.x;
|
||||||
|
int d = threadIdx.x;
|
||||||
|
if (d >= p.head_dim) return;
|
||||||
|
|
||||||
|
int batch = bh / p.q_head;
|
||||||
|
int q_head = bh % p.q_head;
|
||||||
|
|
||||||
|
size_t split_base = (size_t)bh * MAX_SPLITS;
|
||||||
|
const float* mlp = p.ml_part + split_base * 2;
|
||||||
|
const float* op = p.o_part + split_base * p.head_dim;
|
||||||
|
|
||||||
|
float m = -FLT_MAX, l = 0.0f, acc = 0.0f;
|
||||||
|
for (int s = 0; s < p.num_splits; s++) {
|
||||||
|
float mi = mlp[s * 2];
|
||||||
|
if (mi <= -FLT_MAX) continue;
|
||||||
|
float li = mlp[s * 2 + 1];
|
||||||
|
float nm = fmaxf(m, mi);
|
||||||
|
float corr = __expf(m - nm);
|
||||||
|
float e = __expf(mi - nm);
|
||||||
|
acc = fmaf(acc, corr, op[s * p.head_dim + d] * e);
|
||||||
|
l = fmaf(l, corr, li * e);
|
||||||
|
m = nm;
|
||||||
|
}
|
||||||
|
|
||||||
|
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
|
||||||
|
int o_off = batch * p.q_stride_b + q_head * p.q_stride_h + d * p.q_stride_d;
|
||||||
|
p.o[o_off] = __float2bfloat16(acc * inv);
|
||||||
|
}
|
||||||
@@ -0,0 +1,182 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <cfloat>
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
#include "attn_common.h"
|
||||||
|
#include "attn_mma_utils.cuh"
|
||||||
|
#include "attn_warp_utils.cuh"
|
||||||
|
|
||||||
|
// Paged split-KV tensor-core decode via GQA head-packing.
|
||||||
|
// Reads K/V directly from the page pool through a page table — one tile
|
||||||
|
// (BC=32) fits within a single page (page_size >= 32), so the page-table
|
||||||
|
// lookup happens once per tile for cp.async.
|
||||||
|
//
|
||||||
|
// IsCausal and HasMask are compile-time bools.
|
||||||
|
template <typename Traits, bool IsCausal, bool HasMask>
|
||||||
|
__global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16> p) {
|
||||||
|
const int lane = threadIdx.x;
|
||||||
|
const int gid = lane >> 2;
|
||||||
|
const int tid4 = lane & 3;
|
||||||
|
|
||||||
|
const int pass = blockIdx.x / p.kv_head;
|
||||||
|
const int kv_head = blockIdx.x % p.kv_head;
|
||||||
|
const int batch = blockIdx.y;
|
||||||
|
const int split = blockIdx.z;
|
||||||
|
|
||||||
|
constexpr int MAX_G = 16;
|
||||||
|
const int G_total = p.q_head / p.kv_head;
|
||||||
|
const int g_begin = pass * MAX_G;
|
||||||
|
const int G = min(MAX_G, G_total - g_begin);
|
||||||
|
const int q_head0 = kv_head * G_total + g_begin;
|
||||||
|
|
||||||
|
__shared__ __align__(16) bf16 sK[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 qra = gid;
|
||||||
|
const int qrb = gid + 8;
|
||||||
|
const bool va = qra < G, vb = qrb < G;
|
||||||
|
unsigned Qa[Traits::KD][4];
|
||||||
|
load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
|
||||||
|
qra, qrb, va, vb, tid4, Qa);
|
||||||
|
|
||||||
|
float Oacc[Traits::DN8][4];
|
||||||
|
#pragma unroll
|
||||||
|
for (int j = 0; j < Traits::DN8; j++)
|
||||||
|
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
||||||
|
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
||||||
|
|
||||||
|
const int tiles_total = (p.kv_len + Traits::BC - 1) / Traits::BC;
|
||||||
|
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
|
||||||
|
const int ti_begin = split * tiles_per_split;
|
||||||
|
const int ti_end = min(tiles_total, ti_begin + tiles_per_split);
|
||||||
|
|
||||||
|
const int64_t page_stride = (int64_t)p.page_size * p.kv_head * Traits::HEAD_DIM;
|
||||||
|
const int64_t pos_stride = (int64_t)p.kv_head * Traits::HEAD_DIM;
|
||||||
|
const int64_t head_off = (int64_t)kv_head * Traits::HEAD_DIM;
|
||||||
|
|
||||||
|
// ---- 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) {
|
||||||
|
int kv0 = ti * Traits::BC;
|
||||||
|
bf16* dK = sK + buf * Traits::BC * Traits::LD;
|
||||||
|
bf16* dV = sV + buf * Traits::BC * Traits::LD;
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = lane * Traits::VEC; i < Traits::TOTAL;
|
||||||
|
i += Traits::NUM_THREADS * Traits::VEC) {
|
||||||
|
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
|
||||||
|
int kc = kv0 + r;
|
||||||
|
bool valid = (kc < 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;
|
||||||
|
int64_t gmem_base = (int64_t)phys_page * page_stride
|
||||||
|
+ (int64_t)page_off * pos_stride
|
||||||
|
+ head_off;
|
||||||
|
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
|
||||||
|
cp_async_16_pred(&dK[off], &p.k_cache[gmem_base + d], valid);
|
||||||
|
cp_async_16_pred(&dV[off], &p.v_cache[gmem_base + d], valid);
|
||||||
|
}
|
||||||
|
cp_async_commit();
|
||||||
|
};
|
||||||
|
|
||||||
|
// ---- Multi-stage cp.async pipeline ----
|
||||||
|
// Prologue loads STAGES tiles; each loop iteration waits only for the
|
||||||
|
// oldest outstanding group (wait_group<STAGES-1>) so the STAGES-1 newer
|
||||||
|
// tile loads stay in flight and overlap with the current tile's compute.
|
||||||
|
constexpr int STAGES = Traits::STAGES;
|
||||||
|
const int ntiles = ti_end - ti_begin;
|
||||||
|
|
||||||
|
auto process_tile = [&](int it, int buf) {
|
||||||
|
const bf16* bK = sK + buf * Traits::BC * Traits::LD;
|
||||||
|
const bf16* bV = sV + buf * Traits::BC * Traits::LD;
|
||||||
|
int kv0 = (ti_begin + it) * Traits::BC;
|
||||||
|
|
||||||
|
float Sacc[Traits::NC8][4];
|
||||||
|
mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
|
||||||
|
|
||||||
|
#pragma unroll
|
||||||
|
for (int n8 = 0; n8 < Traits::NC8; n8++)
|
||||||
|
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
||||||
|
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
||||||
|
|
||||||
|
int maxc = IsCausal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
|
||||||
|
mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc,
|
||||||
|
0, 0,
|
||||||
|
p.mask_b_stride, 0, 0,
|
||||||
|
batch, 0,
|
||||||
|
p.mask,
|
||||||
|
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||||
|
|
||||||
|
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
|
||||||
|
};
|
||||||
|
|
||||||
|
if (ntiles >= STAGES) {
|
||||||
|
#pragma unroll
|
||||||
|
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 {
|
||||||
|
size_t bh = (size_t)batch * p.q_head + h;
|
||||||
|
return bh * MAX_SPLITS + split;
|
||||||
|
};
|
||||||
|
#pragma unroll
|
||||||
|
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
|
||||||
|
int d = dn8 * 8 + 2 * tid4;
|
||||||
|
int r0 = gid, r1 = gid + 8;
|
||||||
|
if (r0 < G) {
|
||||||
|
int h = q_head0 + r0;
|
||||||
|
float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
|
||||||
|
op[d] = Oacc[dn8][0];
|
||||||
|
op[d + 1] = Oacc[dn8][1];
|
||||||
|
}
|
||||||
|
if (r1 < G) {
|
||||||
|
int h = q_head0 + r1;
|
||||||
|
float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
|
||||||
|
op[d] = Oacc[dn8][2];
|
||||||
|
op[d + 1] = Oacc[dn8][3];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if (tid4 == 0) {
|
||||||
|
int r0 = gid, r1 = gid + 8;
|
||||||
|
if (r0 < G) {
|
||||||
|
int h = q_head0 + r0;
|
||||||
|
float* mp = p.ml_part + split_slot(h) * 2;
|
||||||
|
mp[0] = m0; mp[1] = l0;
|
||||||
|
}
|
||||||
|
if (r1 < G) {
|
||||||
|
int h = q_head0 + r1;
|
||||||
|
float* mp = p.ml_part + split_slot(h) * 2;
|
||||||
|
mp[0] = m1; mp[1] = l1;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,35 @@
|
|||||||
|
#include "attn_dispatchers.cuh"
|
||||||
|
#include "attn_entry_utils.cuh"
|
||||||
|
|
||||||
|
torch::Tensor attn_prefill(
|
||||||
|
torch::Tensor q,
|
||||||
|
torch::Tensor k,
|
||||||
|
torch::Tensor v,
|
||||||
|
c10::optional<torch::Tensor> mask,
|
||||||
|
int64_t causal_offset,
|
||||||
|
double scale,
|
||||||
|
int64_t layout
|
||||||
|
) {
|
||||||
|
AttentionParams<bf16> p;
|
||||||
|
attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p);
|
||||||
|
TORCH_CHECK(p.head_dim % 16 == 0, "head_dim must be multiple of 16");
|
||||||
|
|
||||||
|
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
|
||||||
|
auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
|
||||||
|
p.o = (bf16*)O_view.data_ptr();
|
||||||
|
|
||||||
|
DISPATCH_HEAD_DIM(p.head_dim, dispatch_prefill, p);
|
||||||
|
return O;
|
||||||
|
}
|
||||||
|
|
||||||
|
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||||
|
m.def("attn_prefill", &attn_prefill,
|
||||||
|
py::arg("q"),
|
||||||
|
py::arg("k"),
|
||||||
|
py::arg("v"),
|
||||||
|
py::arg("mask") = py::none(),
|
||||||
|
py::arg("causal_offset") = -1,
|
||||||
|
py::arg("scale") = 0.0,
|
||||||
|
py::arg("layout") = 0,
|
||||||
|
"GQA prefill (tensor-core mma on sm_80+, scalar fallback)");
|
||||||
|
}
|
||||||
@@ -1,13 +1,14 @@
|
|||||||
#pragma once
|
#pragma once
|
||||||
#include "gqa_common.cuh"
|
#include <cfloat>
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
#include "attn_common.h"
|
||||||
|
|
||||||
|
using bf16 = __nv_bfloat16;
|
||||||
|
|
||||||
// v9: group-split register blocking. G threads cooperate on one query row,
|
// v9: group-split register blocking. G threads cooperate on one query row,
|
||||||
// each owning HEAD_DIM/G dims of qreg[]/acc[]. Small per-thread footprint keeps
|
// each owning HEAD_DIM/G dims of qreg[]/acc[]. IsCausal and HasMask are
|
||||||
// occupancy high; the S dot product is reduced across the G-lane group with a
|
// compile-time bools — the compiler eliminates dead branches.
|
||||||
// short shuffle chain (log2(G) shuffles) instead of a full 32-lane warp reduce.
|
// Templated on <HEAD_DIM, G, ROWS, P_BC, IsCausal, HasMask>.
|
||||||
// Online (per-kv) softmax — cheap because acc[] is only HEAD_DIM/G long.
|
|
||||||
// Templated on <HEAD_DIM, G, ROWS, P_BC>. Block = (G, ROWS). G power-of-two,
|
|
||||||
// G*ROWS a multiple of 32 with groups warp-aligned.
|
|
||||||
|
|
||||||
template <int G>
|
template <int G>
|
||||||
__device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
|
__device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
|
||||||
@@ -17,8 +18,7 @@ __device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
|
|||||||
return v;
|
return v;
|
||||||
}
|
}
|
||||||
|
|
||||||
// load 8 contiguous bf16 from (16-byte aligned) smem as one float4, unpack to
|
// load 8 contiguous bf16 from (16-byte aligned) smem as one float4
|
||||||
// 8 floats — cuts shared-load instructions 8x vs scalar bf16 loads.
|
|
||||||
__device__ __forceinline__ void ld8(const bf16* p, float* o) {
|
__device__ __forceinline__ void ld8(const bf16* p, float* o) {
|
||||||
float4 raw = *reinterpret_cast<const float4*>(p);
|
float4 raw = *reinterpret_cast<const float4*>(p);
|
||||||
const __nv_bfloat162* h = reinterpret_cast<const __nv_bfloat162*>(&raw);
|
const __nv_bfloat162* h = reinterpret_cast<const __nv_bfloat162*>(&raw);
|
||||||
@@ -30,8 +30,8 @@ __device__ __forceinline__ void ld8(const bf16* p, float* o) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
template <int HEAD_DIM, int G, int ROWS, int P_BC>
|
template <int HEAD_DIM, int G, int ROWS, int P_BC, bool IsCausal, bool HasMask>
|
||||||
__global__ void gqa_prefill_attn_kernel_t(GQAParams p) {
|
__global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
|
||||||
constexpr int DPT = HEAD_DIM / G;
|
constexpr int DPT = HEAD_DIM / G;
|
||||||
|
|
||||||
int q_tile = blockIdx.x;
|
int q_tile = blockIdx.x;
|
||||||
@@ -43,16 +43,17 @@ __global__ void gqa_prefill_attn_kernel_t(GQAParams p) {
|
|||||||
|
|
||||||
int kv_head = q_head / (p.q_head / p.kv_head);
|
int kv_head = q_head / (p.q_head / p.kv_head);
|
||||||
|
|
||||||
extern __shared__ __align__(16) bf16 smem[];
|
__shared__ __align__(16) bf16 sK[P_BC * HEAD_DIM];
|
||||||
bf16* sK = smem;
|
__shared__ __align__(16) bf16 sV[P_BC * HEAD_DIM];
|
||||||
bf16* sV = sK + P_BC * HEAD_DIM;
|
|
||||||
|
|
||||||
|
// Q: stride-based load [batch, q_head, q_len, head_dim]
|
||||||
float qreg[DPT];
|
float qreg[DPT];
|
||||||
if (q_row < p.q_len) {
|
if (q_row < p.q_len) {
|
||||||
int q_off = ((batch * p.q_head + q_head) * p.q_len + q_row) * HEAD_DIM + gpos * DPT;
|
int q_off = batch * p.q_stride_b + q_head * p.q_stride_h
|
||||||
|
+ q_row * p.q_stride_l + gpos * DPT * p.q_stride_d;
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int i = 0; i < DPT; i++)
|
for (int i = 0; i < DPT; i++)
|
||||||
qreg[i] = __bfloat162float(p.q[q_off + i]) * p.scale;
|
qreg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]);
|
||||||
}
|
}
|
||||||
|
|
||||||
float m = -FLT_MAX, l = 0.0f;
|
float m = -FLT_MAX, l = 0.0f;
|
||||||
@@ -61,13 +62,13 @@ __global__ void gqa_prefill_attn_kernel_t(GQAParams p) {
|
|||||||
for (int i = 0; i < DPT; i++)
|
for (int i = 0; i < DPT; i++)
|
||||||
acc[i] = 0.0f;
|
acc[i] = 0.0f;
|
||||||
|
|
||||||
int kv_base = ((batch * p.kv_head + kv_head) * p.kv_len) * HEAD_DIM;
|
// KV: stride-based base
|
||||||
|
int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
||||||
|
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;
|
||||||
|
|
||||||
// per-group shuffle mask: only the G lanes of this row's group participate,
|
|
||||||
// so causal masking (differing loop bounds across rows in a warp) is safe.
|
|
||||||
int lane_in_warp = lid & 31;
|
int lane_in_warp = lid & 31;
|
||||||
unsigned gmask = (G == 32) ? 0xFFFFFFFFu
|
unsigned gmask = (G == 32) ? 0xFFFFFFFFu
|
||||||
: (((1u << G) - 1u) << (lane_in_warp & ~(G - 1)));
|
: (((1u << G) - 1u) << (lane_in_warp & ~(G - 1)));
|
||||||
@@ -76,22 +77,29 @@ __global__ void gqa_prefill_attn_kernel_t(GQAParams p) {
|
|||||||
int kv0 = ti * P_BC;
|
int kv0 = ti * P_BC;
|
||||||
int tlen = min(P_BC, p.kv_len - kv0);
|
int tlen = min(P_BC, p.kv_len - kv0);
|
||||||
|
|
||||||
|
// Load K/V into shared memory from strided global
|
||||||
for (int i = lid; i < tlen * HEAD_DIM; i += tt) {
|
for (int i = lid; i < tlen * HEAD_DIM; i += tt) {
|
||||||
int gidx = kv_base + (kv0 + i / HEAD_DIM) * HEAD_DIM + (i % HEAD_DIM);
|
int s = i / HEAD_DIM;
|
||||||
sK[i] = p.k[gidx];
|
int d_dim = i % HEAD_DIM;
|
||||||
sV[i] = p.v[gidx];
|
int kv_idx = kv0 + s;
|
||||||
|
int g_off = kv_base + kv_idx * p.kv_stride_l + d_dim * p.kv_stride_d;
|
||||||
|
sK[i] = p.k[g_off];
|
||||||
|
sV[i] = p.v[g_off];
|
||||||
}
|
}
|
||||||
__syncthreads();
|
__syncthreads();
|
||||||
|
|
||||||
int lim = tlen;
|
int lim = tlen;
|
||||||
if (p.is_causal && q_row < p.q_len) {
|
if constexpr (IsCausal) {
|
||||||
int ep = q_row + p.causal_offset + 1;
|
if (q_row < p.q_len) {
|
||||||
if (kv0 >= ep)
|
int ep = q_row + p.causal_offset + 1;
|
||||||
lim = 0;
|
if (kv0 >= ep)
|
||||||
else if (kv0 + tlen > ep)
|
lim = 0;
|
||||||
lim = ep - kv0;
|
else if (kv0 + tlen > ep)
|
||||||
|
lim = ep - kv0;
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
int mask_row_base = mask_batch_base + q_row * p.mask_q_stride;
|
||||||
for (int s = 0; s < lim; s++) {
|
for (int s = 0; s < lim; s++) {
|
||||||
const bf16* kr = sK + s * HEAD_DIM + gpos * DPT;
|
const bf16* kr = sK + s * HEAD_DIM + gpos * DPT;
|
||||||
float part = 0.0f;
|
float part = 0.0f;
|
||||||
@@ -103,10 +111,13 @@ __global__ void gqa_prefill_attn_kernel_t(GQAParams p) {
|
|||||||
for (int j = 0; j < 8; j++)
|
for (int j = 0; j < 8; j++)
|
||||||
part = fmaf(qreg[i + j], k8[j], part);
|
part = fmaf(qreg[i + j], k8[j], part);
|
||||||
}
|
}
|
||||||
float dot = group_reduce_sum<G>(part, gmask);
|
float dot = group_reduce_sum<G>(part, gmask) * p.scale;
|
||||||
|
|
||||||
if (p.use_mask && p.mask && !p.mask[batch * p.kv_len + kv0 + s])
|
int kv_idx = kv0 + s;
|
||||||
dot = -FLT_MAX;
|
if constexpr (HasMask) {
|
||||||
|
if (!p.mask[mask_row_base + kv_idx])
|
||||||
|
dot = -FLT_MAX;
|
||||||
|
}
|
||||||
|
|
||||||
float nm = fmaxf(m, dot);
|
float nm = fmaxf(m, dot);
|
||||||
float al = __expf(m - nm);
|
float al = __expf(m - nm);
|
||||||
@@ -128,10 +139,11 @@ __global__ void gqa_prefill_attn_kernel_t(GQAParams p) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if (q_row < p.q_len) {
|
if (q_row < p.q_len) {
|
||||||
int o_off = ((batch * p.q_head + q_head) * p.q_len + q_row) * HEAD_DIM + gpos * DPT;
|
int o_off = batch * p.q_stride_b + q_head * p.q_stride_h
|
||||||
float rl = (l > 1e-10f) ? (1.0f / l) : 0.0f;
|
+ q_row * p.q_stride_l + gpos * DPT * p.q_stride_d;
|
||||||
|
float rl = (l > 1e-20f) ? (1.0f / l) : 0.0f;
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int i = 0; i < DPT; i++)
|
for (int i = 0; i < DPT; i++)
|
||||||
p.o[o_off + i] = __float2bfloat16(acc[i] * rl);
|
p.o[o_off + i * p.q_stride_d] = __float2bfloat16(acc[i] * rl);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -0,0 +1,146 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <cfloat>
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
#include "attn_common.h"
|
||||||
|
#include "attn_mma_utils.cuh"
|
||||||
|
|
||||||
|
// Tensor-core prefill flash attention (raw mma.sync PTX).
|
||||||
|
// One warp owns BR=16 query rows. S = Q@K^T and O = P@V run on bf16 tensor
|
||||||
|
// cores via mma.sync.m16n8k16 (f32 accumulate).
|
||||||
|
//
|
||||||
|
// IsCausal and HasMask are compile-time bools — the compiler eliminates all
|
||||||
|
// dead branches in the inner compute loop (FA2-style).
|
||||||
|
//
|
||||||
|
// Traits = KernelTraits<HEAD_DIM, BC, WARPS=4, STAGES=2>.
|
||||||
|
template <typename Traits, bool IsCausal, bool HasMask>
|
||||||
|
__global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
|
||||||
|
const int warp = threadIdx.x / 32;
|
||||||
|
const int lane = threadIdx.x % 32;
|
||||||
|
const int gid = lane >> 2; // 0..7
|
||||||
|
const int tid4 = lane & 3; // 0..3
|
||||||
|
|
||||||
|
const int q_head = blockIdx.y;
|
||||||
|
const int batch = blockIdx.z;
|
||||||
|
const int kv_head = q_head / (p.q_head / p.kv_head);
|
||||||
|
const int qrow0 = (blockIdx.x * Traits::WARPS + warp) * Traits::BR;
|
||||||
|
|
||||||
|
// Static shared memory: double-buffered K/V (no sQ — Q goes direct
|
||||||
|
// to registers in mma A-operand layout).
|
||||||
|
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
|
||||||
|
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
|
||||||
|
|
||||||
|
// Load Q fragments straight from global into mma A-operand layout.
|
||||||
|
const int q_base = batch * p.q_stride_b + q_head * p.q_stride_h;
|
||||||
|
const int qra = qrow0 + gid;
|
||||||
|
const int qrb = qrow0 + gid + 8;
|
||||||
|
const bool va = qra < p.q_len, vb = qrb < p.q_len;
|
||||||
|
unsigned Qa[Traits::KD][4];
|
||||||
|
load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_l, p.q_stride_d,
|
||||||
|
qra, qrb, va, vb, tid4, Qa);
|
||||||
|
|
||||||
|
float Oacc[Traits::DN8][4];
|
||||||
|
#pragma unroll
|
||||||
|
for (int j = 0; j < Traits::DN8; j++)
|
||||||
|
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
||||||
|
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
||||||
|
|
||||||
|
// KV: stride-based base
|
||||||
|
const int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
||||||
|
const int tiles = (p.kv_len + Traits::BC - 1) / Traits::BC;
|
||||||
|
const int qr0 = qrow0 + gid;
|
||||||
|
const int qr1 = qrow0 + gid + 8;
|
||||||
|
|
||||||
|
// Causal tile-skip bounds (dead code when IsCausal == false)
|
||||||
|
const int max_kv = qrow0 + Traits::BR - 1 + p.causal_offset;
|
||||||
|
const int block_max_kv =
|
||||||
|
blockIdx.x * Traits::WARPS * Traits::BR + Traits::WARPS * Traits::BR - 1
|
||||||
|
+ p.causal_offset;
|
||||||
|
|
||||||
|
int t_end = tiles - 1;
|
||||||
|
if constexpr (IsCausal) {
|
||||||
|
int bt = block_max_kv / Traits::BC;
|
||||||
|
if (bt < t_end) t_end = bt;
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- Load tile lambda: predicated cp.async ----
|
||||||
|
auto load_tile = [&](int ti, int buf) {
|
||||||
|
int kv0 = ti * Traits::BC;
|
||||||
|
bf16* dK = sK + buf * Traits::BC * Traits::LD;
|
||||||
|
bf16* dV = sV + buf * Traits::BC * Traits::LD;
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = threadIdx.x * Traits::VEC; i < Traits::TOTAL;
|
||||||
|
i += Traits::NUM_THREADS * Traits::VEC) {
|
||||||
|
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
|
||||||
|
int kc = kv0 + r;
|
||||||
|
bool valid = kc < p.kv_len;
|
||||||
|
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
|
||||||
|
int g_off = kv_base + kc * p.kv_stride_l + d * p.kv_stride_d;
|
||||||
|
cp_async_16_pred(&dK[off], &p.k[g_off], valid);
|
||||||
|
cp_async_16_pred(&dV[off], &p.v[g_off], valid);
|
||||||
|
}
|
||||||
|
cp_async_commit();
|
||||||
|
};
|
||||||
|
|
||||||
|
// ---- Prologue: issue first tile load ----
|
||||||
|
load_tile(0, 0);
|
||||||
|
|
||||||
|
for (int ti = 0; ti <= t_end; ti++) {
|
||||||
|
int buf = ti & 1;
|
||||||
|
|
||||||
|
// Wait for current tile, then publish cross-warp + guard buffer reuse.
|
||||||
|
cp_async_wait_group<0>();
|
||||||
|
__syncthreads();
|
||||||
|
if (ti < t_end) load_tile(ti + 1, (ti + 1) & 1);
|
||||||
|
|
||||||
|
const bf16* bK = sK + buf * Traits::BC * Traits::LD;
|
||||||
|
const bf16* bV = sV + buf * Traits::BC * Traits::LD;
|
||||||
|
int kv0 = ti * Traits::BC;
|
||||||
|
|
||||||
|
// Warp-level causal skip (dead branch eliminated when IsCausal == false)
|
||||||
|
if (!IsCausal || kv0 <= max_kv) {
|
||||||
|
|
||||||
|
float Sacc[Traits::NC8][4];
|
||||||
|
mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
|
||||||
|
|
||||||
|
// Post-multiply scale in float (no bf16 precision loss)
|
||||||
|
#pragma unroll
|
||||||
|
for (int n8 = 0; n8 < Traits::NC8; n8++)
|
||||||
|
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
||||||
|
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
||||||
|
|
||||||
|
int maxc0 = IsCausal ? min(p.kv_len, qr0 + p.causal_offset + 1)
|
||||||
|
: p.kv_len;
|
||||||
|
int maxc1 = IsCausal ? min(p.kv_len, qr1 + p.causal_offset + 1)
|
||||||
|
: p.kv_len;
|
||||||
|
mma_softmax_tile<Traits, HasMask>(kv0, maxc0, maxc1,
|
||||||
|
qr0, qr1,
|
||||||
|
p.mask_b_stride, p.mask_h_stride, p.mask_q_stride,
|
||||||
|
batch, q_head,
|
||||||
|
p.mask,
|
||||||
|
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||||
|
|
||||||
|
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- write output: packed bf16x2 stores ----
|
||||||
|
float rl0 = (l0 > 1e-20f) ? (1.0f / l0) : 0.0f;
|
||||||
|
float rl1 = (l1 > 1e-20f) ? (1.0f / l1) : 0.0f;
|
||||||
|
const int o_base = batch * p.q_stride_b + q_head * p.q_stride_h;
|
||||||
|
#pragma unroll
|
||||||
|
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
|
||||||
|
int d = dn8 * 8 + 2 * tid4;
|
||||||
|
if (qr0 < p.q_len) {
|
||||||
|
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][0] * rl0,
|
||||||
|
Oacc[dn8][1] * rl0);
|
||||||
|
*reinterpret_cast<__nv_bfloat162*>(
|
||||||
|
&p.o[o_base + qr0 * p.q_stride_l + d * p.q_stride_d]) = v;
|
||||||
|
}
|
||||||
|
if (qr1 < p.q_len) {
|
||||||
|
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][2] * rl1,
|
||||||
|
Oacc[dn8][3] * rl1);
|
||||||
|
*reinterpret_cast<__nv_bfloat162*>(
|
||||||
|
&p.o[o_base + qr1 * p.q_stride_l + d * p.q_stride_d]) = v;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,13 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
|
||||||
|
using bf16 = __nv_bfloat16;
|
||||||
|
|
||||||
|
static constexpr int MAX_SPLITS = 32;
|
||||||
|
|
||||||
|
__device__ inline float warp_reduce_sum(float val) {
|
||||||
|
#pragma unroll
|
||||||
|
for (int offset = 16; offset > 0; offset >>= 1)
|
||||||
|
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
|
||||||
|
return val;
|
||||||
|
}
|
||||||
@@ -1,35 +0,0 @@
|
|||||||
#pragma once
|
|
||||||
#include <cuda_bf16.h>
|
|
||||||
#include <cuda_runtime.h>
|
|
||||||
#include <cfloat>
|
|
||||||
#include <algorithm>
|
|
||||||
|
|
||||||
using bf16 = __nv_bfloat16;
|
|
||||||
using std::min;
|
|
||||||
|
|
||||||
constexpr int DC_CHUNK = 64;
|
|
||||||
constexpr int Br = 32, Bc = 64;
|
|
||||||
|
|
||||||
__device__ inline float warp_reduce_sum(float val) {
|
|
||||||
for (int offset = 16; offset > 0; offset >>= 1)
|
|
||||||
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
|
|
||||||
return val;
|
|
||||||
}
|
|
||||||
|
|
||||||
struct GQAParams {
|
|
||||||
int batch;
|
|
||||||
int q_head;
|
|
||||||
int kv_head;
|
|
||||||
int q_len;
|
|
||||||
int kv_len;
|
|
||||||
int head_dim;
|
|
||||||
int use_mask;
|
|
||||||
int is_causal;
|
|
||||||
int causal_offset;
|
|
||||||
float scale;
|
|
||||||
const bf16* __restrict__ q;
|
|
||||||
const bf16* __restrict__ k;
|
|
||||||
const bf16* __restrict__ v;
|
|
||||||
const bool* __restrict__ mask;
|
|
||||||
bf16* __restrict__ o;
|
|
||||||
};
|
|
||||||
@@ -1,114 +0,0 @@
|
|||||||
#include "gqa_decode_attn.cuh"
|
|
||||||
#include <torch/extension.h>
|
|
||||||
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
|
||||||
#include "gqa_decode_attn_mma.cuh"
|
|
||||||
#endif
|
|
||||||
|
|
||||||
template <int HEAD_DIM>
|
|
||||||
static void dispatch_decode(GQAParams& p) {
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
|
||||||
constexpr int BC = 32, BR = 16, LD = HEAD_DIM; // XOR swizzle → no padding
|
|
||||||
int G = p.q_head / p.kv_head;
|
|
||||||
// head-packing tensor-core path needs 1 < G <= 16 (MMA M dim) and no mask;
|
|
||||||
// everything else uses the scalar kernel
|
|
||||||
if (!p.use_mask && G > 1 && G <= 16) {
|
|
||||||
dim3 grid(p.kv_head, p.batch, 1);
|
|
||||||
dim3 block(32, 1, 1);
|
|
||||||
// sK + sV + sQ, each BC/BR * LD (single buffer for high occupancy)
|
|
||||||
int smem = (2 * BC * LD + BR * LD) * (int)sizeof(bf16);
|
|
||||||
cudaFuncSetAttribute(gqa_decode_attn_mma_kernel<HEAD_DIM, BC>,
|
|
||||||
cudaFuncAttributeMaxDynamicSharedMemorySize, smem);
|
|
||||||
gqa_decode_attn_mma_kernel<HEAD_DIM, BC><<<grid, block, smem>>>(p);
|
|
||||||
return;
|
|
||||||
}
|
|
||||||
// scalar fallback (per-KV-head, one warp per query head)
|
|
||||||
int group_size = p.q_head / p.kv_head;
|
|
||||||
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
|
|
||||||
dim3 block(32, group_size);
|
|
||||||
dim3 grid(p.batch * p.kv_head);
|
|
||||||
gqa_decode_attn_kernel<<<grid, block, smem>>>(p);
|
|
||||||
#else
|
|
||||||
// scalar fallback (per-KV-head, one warp per query head)
|
|
||||||
int group_size = p.q_head / p.kv_head;
|
|
||||||
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
|
|
||||||
dim3 block(32, group_size);
|
|
||||||
dim3 grid(p.batch * p.kv_head);
|
|
||||||
gqa_decode_attn_kernel<<<grid, block, smem>>>(p);
|
|
||||||
#endif
|
|
||||||
}
|
|
||||||
|
|
||||||
torch::Tensor gqa_decode_attn(
|
|
||||||
torch::Tensor q,
|
|
||||||
torch::Tensor k,
|
|
||||||
torch::Tensor v,
|
|
||||||
c10::optional<torch::Tensor> mask,
|
|
||||||
bool is_causal = false,
|
|
||||||
int64_t causal_offset = 0,
|
|
||||||
c10::optional<double> scale = c10::nullopt
|
|
||||||
) {
|
|
||||||
TORCH_CHECK(q.is_cuda() && k.is_cuda() && v.is_cuda());
|
|
||||||
TORCH_CHECK(q.dtype() == torch::kBFloat16);
|
|
||||||
TORCH_CHECK(k.dtype() == torch::kBFloat16);
|
|
||||||
TORCH_CHECK(v.dtype() == torch::kBFloat16);
|
|
||||||
TORCH_CHECK(q.size(2) == 1, "Q seq_len must be 1");
|
|
||||||
|
|
||||||
GQAParams p;
|
|
||||||
p.batch = q.size(0);
|
|
||||||
p.q_head = q.size(1);
|
|
||||||
p.kv_head = k.size(1);
|
|
||||||
p.q_len = 1;
|
|
||||||
p.kv_len = k.size(2);
|
|
||||||
p.head_dim = q.size(3);
|
|
||||||
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
|
|
||||||
p.use_mask = mask.has_value();
|
|
||||||
p.is_causal = (int)is_causal;
|
|
||||||
p.causal_offset = (int)causal_offset;
|
|
||||||
p.scale = scale.has_value() ? (float)scale.value() : 1.0f / sqrtf((float)p.head_dim);
|
|
||||||
p.q = (const bf16*)q.data_ptr();
|
|
||||||
p.k = (const bf16*)k.data_ptr();
|
|
||||||
p.v = (const bf16*)v.data_ptr();
|
|
||||||
if (p.use_mask) {
|
|
||||||
TORCH_CHECK(mask.value().dtype() == torch::kBool);
|
|
||||||
TORCH_CHECK(mask.value().dim() == 2);
|
|
||||||
TORCH_CHECK(mask.value().size(0) == p.batch);
|
|
||||||
TORCH_CHECK(mask.value().size(1) == p.kv_len);
|
|
||||||
p.mask = mask.value().data_ptr<bool>();
|
|
||||||
} else {
|
|
||||||
p.mask = nullptr;
|
|
||||||
}
|
|
||||||
|
|
||||||
auto O = torch::empty_like(q);
|
|
||||||
p.o = (bf16*)O.data_ptr();
|
|
||||||
|
|
||||||
switch (p.head_dim) {
|
|
||||||
case 32:
|
|
||||||
dispatch_decode<32>(p);
|
|
||||||
break;
|
|
||||||
case 64:
|
|
||||||
dispatch_decode<64>(p);
|
|
||||||
break;
|
|
||||||
case 128:
|
|
||||||
dispatch_decode<128>(p);
|
|
||||||
break;
|
|
||||||
case 256:
|
|
||||||
dispatch_decode<256>(p);
|
|
||||||
break;
|
|
||||||
default:
|
|
||||||
TORCH_CHECK(false, "decode: unsupported head_dim ", p.head_dim,
|
|
||||||
" (supported: 32, 64, 128, 256)");
|
|
||||||
}
|
|
||||||
return O;
|
|
||||||
}
|
|
||||||
|
|
||||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
|
||||||
m.def("gqa_decode_attn", &gqa_decode_attn,
|
|
||||||
py::arg("q"),
|
|
||||||
py::arg("k"),
|
|
||||||
py::arg("v"),
|
|
||||||
py::arg("mask") = py::none(),
|
|
||||||
py::arg("is_causal") = false,
|
|
||||||
py::arg("causal_offset") = 0,
|
|
||||||
py::arg("scale") = py::none(),
|
|
||||||
"GQA decode (tensor-core head-packing on sm_80+, scalar fallback)");
|
|
||||||
}
|
|
||||||
@@ -1,59 +0,0 @@
|
|||||||
#pragma once
|
|
||||||
#include "gqa_common.cuh"
|
|
||||||
|
|
||||||
__global__ void gqa_decode_attn_kernel(GQAParams p) {
|
|
||||||
int batch = blockIdx.x / p.kv_head;
|
|
||||||
int kv_head = blockIdx.x % p.kv_head;
|
|
||||||
int group_size = blockDim.y;
|
|
||||||
int q_head = kv_head * group_size + threadIdx.y;
|
|
||||||
int lane = threadIdx.x;
|
|
||||||
int hd_per_thread = p.head_dim / 32;
|
|
||||||
|
|
||||||
float q_reg[8];
|
|
||||||
int q_off = ((batch * p.q_head + q_head) * 1) * p.head_dim + lane * hd_per_thread;
|
|
||||||
for (int i = 0; i < hd_per_thread; i++)
|
|
||||||
q_reg[i] = __bfloat162float(p.q[q_off + i]);
|
|
||||||
|
|
||||||
int kv_base = ((batch * p.kv_head + kv_head) * p.kv_len) * p.head_dim;
|
|
||||||
int mask_base = batch * p.kv_len;
|
|
||||||
|
|
||||||
float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f};
|
|
||||||
|
|
||||||
extern __shared__ __align__(16) bf16 k_smem[];
|
|
||||||
|
|
||||||
for (int chunk_start = 0; chunk_start < p.kv_len; chunk_start += DC_CHUNK) {
|
|
||||||
int this_chunk = min(DC_CHUNK, p.kv_len - chunk_start);
|
|
||||||
|
|
||||||
int total = this_chunk * p.head_dim;
|
|
||||||
for (int i = threadIdx.y * 32 + lane; i < total; i += blockDim.x * blockDim.y)
|
|
||||||
k_smem[i] = p.k[kv_base + chunk_start * p.head_dim + i];
|
|
||||||
__syncthreads();
|
|
||||||
|
|
||||||
for (int s = 0; s < this_chunk; s++) {
|
|
||||||
float partial = 0.0f;
|
|
||||||
for (int i = 0; i < hd_per_thread; i++)
|
|
||||||
partial += q_reg[i] * __bfloat162float(k_smem[s * p.head_dim + lane * hd_per_thread + i]);
|
|
||||||
partial = warp_reduce_sum(partial) * p.scale;
|
|
||||||
|
|
||||||
if (p.use_mask && p.mask && !p.mask[mask_base + chunk_start + s])
|
|
||||||
partial = -FLT_MAX;
|
|
||||||
if (p.is_causal && (chunk_start + s) > p.causal_offset)
|
|
||||||
partial = -FLT_MAX;
|
|
||||||
|
|
||||||
float new_m = fmaxf(m, partial);
|
|
||||||
float alpha = expf(m - new_m);
|
|
||||||
float beta = expf(partial - new_m);
|
|
||||||
d = d * alpha + beta;
|
|
||||||
|
|
||||||
int v_off = kv_base + (chunk_start + s) * p.head_dim + lane * hd_per_thread;
|
|
||||||
for (int i = 0; i < hd_per_thread; i++)
|
|
||||||
acc_reg[i] = acc_reg[i] * alpha + __bfloat162float(p.v[v_off + i]) * beta;
|
|
||||||
m = new_m;
|
|
||||||
}
|
|
||||||
__syncthreads();
|
|
||||||
}
|
|
||||||
|
|
||||||
int out_off = ((batch * p.q_head + q_head) * 1) * p.head_dim + lane * hd_per_thread;
|
|
||||||
for (int i = 0; i < hd_per_thread; i++)
|
|
||||||
p.o[out_off + i] = __float2bfloat16(acc_reg[i] / d);
|
|
||||||
}
|
|
||||||
@@ -1,219 +0,0 @@
|
|||||||
#pragma once
|
|
||||||
#include "gqa_common.cuh"
|
|
||||||
#include "gqa_mma_utils.cuh"
|
|
||||||
|
|
||||||
// Tensor-core decode via GQA head-packing with cp.async loads.
|
|
||||||
//
|
|
||||||
// Decode has q_len == 1, so S = q @ K^T is a GEMV per head — no tensor-core work
|
|
||||||
// on its own. But GQA gives us G = q_head / kv_head query heads that all share
|
|
||||||
// one kv_head. We pack those G heads into the M=16 rows of mma.sync.m16n8k16,
|
|
||||||
// turning G independent GEMVs into a single GEMM that reuses each loaded K/V tile
|
|
||||||
// across all G heads (K/V load is the decode bottleneck, so the reuse is the win,
|
|
||||||
// not the flops). Fragment layout is identical to the prefill mma kernel; the
|
|
||||||
// only differences are (1) the M rows come from different heads at position 0
|
|
||||||
// instead of different sequence positions of one head, and (2) causal masking is
|
|
||||||
// a single scalar bound shared by every row. One warp owns one (batch, kv_head);
|
|
||||||
// requires G <= 16.
|
|
||||||
//
|
|
||||||
// Optimizations:
|
|
||||||
// - cp.async global→shared for K/V (bypasses registers, cuts instruction count)
|
|
||||||
// - XOR swizzle (swiz_col): LD=HEAD_DIM, zero waste, no bank conflicts
|
|
||||||
// - pre-scaled Q: Q scaled during load, softmax skips per-tile multiply
|
|
||||||
// - single-buffer: keeps smem small for high occupancy
|
|
||||||
|
|
||||||
template <int HEAD_DIM, int BC>
|
|
||||||
__global__ void gqa_decode_attn_mma_kernel(GQAParams p) {
|
|
||||||
constexpr int BR = 16;
|
|
||||||
constexpr int KD = HEAD_DIM / 16; // Q/K k-tiles
|
|
||||||
constexpr int NC8 = BC / 8; // S n-tiles (N=8 each)
|
|
||||||
constexpr int KT2 = BC / 16; // P k-tiles (K=16 each)
|
|
||||||
constexpr int DN8 = HEAD_DIM / 8; // O n-tiles (N=8 each)
|
|
||||||
constexpr int LD = HEAD_DIM; // XOR swizzle handles bank conflicts, zero waste
|
|
||||||
constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1);
|
|
||||||
|
|
||||||
const int lane = threadIdx.x; // single warp
|
|
||||||
const int gid = lane >> 2; // 0..7 → rows gid, gid+8
|
|
||||||
const int tid4 = lane & 3;
|
|
||||||
|
|
||||||
const int kv_head = blockIdx.x;
|
|
||||||
const int batch = blockIdx.y;
|
|
||||||
const int G = p.q_head / p.kv_head;
|
|
||||||
const int q_head0 = kv_head * G;
|
|
||||||
|
|
||||||
extern __shared__ __align__(16) bf16 smem[];
|
|
||||||
bf16* sK = smem; // [BC][LD]
|
|
||||||
bf16* sV = sK + BC * LD; // [BC][LD]
|
|
||||||
bf16* sQ = sV + BC * LD; // [BR][LD]
|
|
||||||
|
|
||||||
// ---- stage Q into shared (pre-scaled, swizzled) ----
|
|
||||||
bf16 scale_bf16 = __float2bfloat16(p.scale);
|
|
||||||
for (int i = lane; i < BR * HEAD_DIM; i += 32) {
|
|
||||||
int r = i / HEAD_DIM, d = i % HEAD_DIM;
|
|
||||||
bf16 val = __float2bfloat16(0.0f);
|
|
||||||
if (r < G) {
|
|
||||||
int qh = q_head0 + r;
|
|
||||||
val = p.q[(batch * p.q_head + qh) * HEAD_DIM + d]; // q_len == 1
|
|
||||||
}
|
|
||||||
sQ[r * LD + swiz_col(d, r, SWIZ_MASK)] = __hmul(val, scale_bf16);
|
|
||||||
}
|
|
||||||
__syncwarp();
|
|
||||||
|
|
||||||
// Q resident A-fragments
|
|
||||||
unsigned Qa[KD][4];
|
|
||||||
int qrow_l = (lane & 7) + (lane & 8);
|
|
||||||
int qcol_l = (lane & 16) ? 8 : 0;
|
|
||||||
#pragma unroll
|
|
||||||
for (int kt = 0; kt < KD; kt++)
|
|
||||||
ldmatrix_x4(Qa[kt], &sQ[qrow_l * LD + swiz_col(kt * 16 + qcol_l, qrow_l, SWIZ_MASK)]);
|
|
||||||
|
|
||||||
float Oacc[DN8][4];
|
|
||||||
#pragma unroll
|
|
||||||
for (int j = 0; j < DN8; j++)
|
|
||||||
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
|
||||||
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
|
||||||
|
|
||||||
const int kv_base = (batch * p.kv_head + kv_head) * p.kv_len * HEAD_DIM;
|
|
||||||
const int mask_base = batch * p.kv_len;
|
|
||||||
const int tiles = (p.kv_len + BC - 1) / BC;
|
|
||||||
const int has_mask = p.use_mask && p.mask;
|
|
||||||
|
|
||||||
for (int ti = 0; ti < tiles; ti++) {
|
|
||||||
int kv0 = ti * BC;
|
|
||||||
|
|
||||||
// ---- load K/V tile to shared (cp.async on full tiles) ----
|
|
||||||
bool full_tile = (kv0 + BC <= p.kv_len);
|
|
||||||
if (full_tile) {
|
|
||||||
constexpr int VEC = 8; // 8 bf16 = 16 bytes per cp.async
|
|
||||||
int total = BC * HEAD_DIM;
|
|
||||||
#pragma unroll
|
|
||||||
for (int i = lane * VEC; i < total; i += 32 * VEC) {
|
|
||||||
int r = i / HEAD_DIM, d = i % HEAD_DIM;
|
|
||||||
int kc = kv0 + r;
|
|
||||||
cp_async_16(&sK[r * LD + swiz_col(d, r, SWIZ_MASK)],
|
|
||||||
&p.k[kv_base + kc * HEAD_DIM + d]);
|
|
||||||
cp_async_16(&sV[r * LD + swiz_col(d, r, SWIZ_MASK)],
|
|
||||||
&p.v[kv_base + kc * HEAD_DIM + d]);
|
|
||||||
}
|
|
||||||
cp_async_commit();
|
|
||||||
cp_async_wait_all();
|
|
||||||
} else {
|
|
||||||
for (int i = lane; i < BC * HEAD_DIM; i += 32) {
|
|
||||||
int r = i / HEAD_DIM, d = i % HEAD_DIM;
|
|
||||||
int kc = kv0 + r;
|
|
||||||
bf16 z = __float2bfloat16(0.0f);
|
|
||||||
sK[r * LD + swiz_col(d, r, SWIZ_MASK)] =
|
|
||||||
(kc < p.kv_len) ? p.k[kv_base + kc * HEAD_DIM + d] : z;
|
|
||||||
sV[r * LD + swiz_col(d, r, SWIZ_MASK)] =
|
|
||||||
(kc < p.kv_len) ? p.v[kv_base + kc * HEAD_DIM + d] : z;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
__syncwarp();
|
|
||||||
|
|
||||||
// S = Q @ K^T (Q already pre-scaled, so Sacc includes scale)
|
|
||||||
float Sacc[NC8][4];
|
|
||||||
#pragma unroll
|
|
||||||
for (int n8 = 0; n8 < NC8; n8++) {
|
|
||||||
Sacc[n8][0] = Sacc[n8][1] = Sacc[n8][2] = Sacc[n8][3] = 0.0f;
|
|
||||||
int krow_l = n8 * 8 + (lane & 7);
|
|
||||||
int kcol_h = (lane & 8) ? 8 : 0;
|
|
||||||
#pragma unroll
|
|
||||||
for (int kt = 0; kt < KD; kt++) {
|
|
||||||
unsigned b[2];
|
|
||||||
ldmatrix_x2(b, &sK[krow_l * LD + swiz_col(kt * 16 + kcol_h, krow_l, SWIZ_MASK)]);
|
|
||||||
mma16816(Sacc[n8], Qa[kt], b, Sacc[n8]);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ---- online softmax (Q pre-scaled → no per-tile scale multiply) ----
|
|
||||||
float rmax0 = -FLT_MAX, rmax1 = -FLT_MAX;
|
|
||||||
#pragma unroll
|
|
||||||
for (int n8 = 0; n8 < NC8; n8++) {
|
|
||||||
int cc = kv0 + n8 * 8 + 2 * tid4;
|
|
||||||
bool bc0 = (cc >= p.kv_len) ||
|
|
||||||
(has_mask && !p.mask[mask_base + cc]);
|
|
||||||
bool bc1 = (cc + 1 >= p.kv_len) ||
|
|
||||||
(has_mask && !p.mask[mask_base + cc + 1]);
|
|
||||||
bool cz = p.is_causal;
|
|
||||||
int off = p.causal_offset;
|
|
||||||
bool bad0 = bc0 || (cz && cc > off);
|
|
||||||
bool bad1 = bc1 || (cz && (cc + 1) > off);
|
|
||||||
float s0 = bad0 ? -FLT_MAX : Sacc[n8][0];
|
|
||||||
float s1 = bad1 ? -FLT_MAX : Sacc[n8][1];
|
|
||||||
float s2 = bad0 ? -FLT_MAX : Sacc[n8][2];
|
|
||||||
float s3 = bad1 ? -FLT_MAX : Sacc[n8][3];
|
|
||||||
Sacc[n8][0] = s0; Sacc[n8][1] = s1; Sacc[n8][2] = s2; Sacc[n8][3] = s3;
|
|
||||||
rmax0 = fmaxf(rmax0, fmaxf(s0, s1));
|
|
||||||
rmax1 = fmaxf(rmax1, fmaxf(s2, s3));
|
|
||||||
}
|
|
||||||
rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 1));
|
|
||||||
rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 2));
|
|
||||||
rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 1));
|
|
||||||
rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 2));
|
|
||||||
|
|
||||||
float nm0 = fmaxf(m0, rmax0), nm1 = fmaxf(m1, rmax1);
|
|
||||||
float corr0 = (nm0 == -FLT_MAX) ? 1.0f : __expf(m0 - nm0);
|
|
||||||
float corr1 = (nm1 == -FLT_MAX) ? 1.0f : __expf(m1 - nm1);
|
|
||||||
|
|
||||||
float rsum0 = 0.0f, rsum1 = 0.0f;
|
|
||||||
#pragma unroll
|
|
||||||
for (int n8 = 0; n8 < NC8; n8++) {
|
|
||||||
float p0 = (Sacc[n8][0] == -FLT_MAX) ? 0.0f : __expf(Sacc[n8][0] - nm0);
|
|
||||||
float p1 = (Sacc[n8][1] == -FLT_MAX) ? 0.0f : __expf(Sacc[n8][1] - nm0);
|
|
||||||
float p2 = (Sacc[n8][2] == -FLT_MAX) ? 0.0f : __expf(Sacc[n8][2] - nm1);
|
|
||||||
float p3 = (Sacc[n8][3] == -FLT_MAX) ? 0.0f : __expf(Sacc[n8][3] - nm1);
|
|
||||||
Sacc[n8][0] = p0; Sacc[n8][1] = p1; Sacc[n8][2] = p2; Sacc[n8][3] = p3;
|
|
||||||
rsum0 += p0 + p1;
|
|
||||||
rsum1 += p2 + p3;
|
|
||||||
}
|
|
||||||
rsum0 += __shfl_xor_sync(0xFFFFFFFF, rsum0, 1);
|
|
||||||
rsum0 += __shfl_xor_sync(0xFFFFFFFF, rsum0, 2);
|
|
||||||
rsum1 += __shfl_xor_sync(0xFFFFFFFF, rsum1, 1);
|
|
||||||
rsum1 += __shfl_xor_sync(0xFFFFFFFF, rsum1, 2);
|
|
||||||
l0 = l0 * corr0 + rsum0;
|
|
||||||
l1 = l1 * corr1 + rsum1;
|
|
||||||
m0 = nm0; m1 = nm1;
|
|
||||||
|
|
||||||
#pragma unroll
|
|
||||||
for (int j = 0; j < DN8; j++) {
|
|
||||||
Oacc[j][0] *= corr0; Oacc[j][1] *= corr0;
|
|
||||||
Oacc[j][2] *= corr1; Oacc[j][3] *= corr1;
|
|
||||||
}
|
|
||||||
|
|
||||||
// O += P @ V
|
|
||||||
#pragma unroll
|
|
||||||
for (int kt2 = 0; kt2 < KT2; kt2++) {
|
|
||||||
unsigned Pa[4];
|
|
||||||
Pa[0] = pk2(Sacc[kt2 * 2][0], Sacc[kt2 * 2][1]);
|
|
||||||
Pa[1] = pk2(Sacc[kt2 * 2][2], Sacc[kt2 * 2][3]);
|
|
||||||
Pa[2] = pk2(Sacc[kt2 * 2 + 1][0], Sacc[kt2 * 2 + 1][1]);
|
|
||||||
Pa[3] = pk2(Sacc[kt2 * 2 + 1][2], Sacc[kt2 * 2 + 1][3]);
|
|
||||||
int vrow_l = kt2 * 16 + (lane & 15);
|
|
||||||
#pragma unroll
|
|
||||||
for (int dn8 = 0; dn8 < DN8; dn8++) {
|
|
||||||
unsigned b[2];
|
|
||||||
ldmatrix_x2_trans(b, &sV[vrow_l * LD + swiz_col(dn8 * 8, vrow_l, SWIZ_MASK)]);
|
|
||||||
mma16816(Oacc[dn8], Pa, b, Oacc[dn8]);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
__syncwarp(); // sK/sV reused next tile
|
|
||||||
}
|
|
||||||
|
|
||||||
// ---- write output ----
|
|
||||||
float rl0 = (l0 > 1e-20f) ? (1.0f / l0) : 0.0f;
|
|
||||||
float rl1 = (l1 > 1e-20f) ? (1.0f / l1) : 0.0f;
|
|
||||||
#pragma unroll
|
|
||||||
for (int dn8 = 0; dn8 < DN8; dn8++) {
|
|
||||||
int d = dn8 * 8 + 2 * tid4;
|
|
||||||
int r0 = gid, r1 = gid + 8;
|
|
||||||
if (r0 < G) {
|
|
||||||
int o_off = (batch * p.q_head + q_head0 + r0) * HEAD_DIM + d;
|
|
||||||
p.o[o_off] = __float2bfloat16(Oacc[dn8][0] * rl0);
|
|
||||||
p.o[o_off + 1] = __float2bfloat16(Oacc[dn8][1] * rl0);
|
|
||||||
}
|
|
||||||
if (r1 < G) {
|
|
||||||
int o_off = (batch * p.q_head + q_head0 + r1) * HEAD_DIM + d;
|
|
||||||
p.o[o_off] = __float2bfloat16(Oacc[dn8][2] * rl1);
|
|
||||||
p.o[o_off + 1] = __float2bfloat16(Oacc[dn8][3] * rl1);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,93 +0,0 @@
|
|||||||
#pragma once
|
|
||||||
|
|
||||||
// Shared MMA utilities for tensor-core GQA kernels.
|
|
||||||
// mma.sync.m16n8k16 PTX wrappers, ldmatrix helpers, and bf16 packing.
|
|
||||||
|
|
||||||
// mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32
|
|
||||||
__device__ __forceinline__ void mma16816(float* d, const unsigned* a,
|
|
||||||
const unsigned* b, const float* c) {
|
|
||||||
asm volatile(
|
|
||||||
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
|
|
||||||
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%10,%11,%12,%13};"
|
|
||||||
: "=f"(d[0]), "=f"(d[1]), "=f"(d[2]), "=f"(d[3])
|
|
||||||
: "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]),
|
|
||||||
"f"(c[0]), "f"(c[1]), "f"(c[2]), "f"(c[3]));
|
|
||||||
}
|
|
||||||
|
|
||||||
// read two adjacent bf16 from smem as one packed .b32 (elem0 low, elem1 high)
|
|
||||||
__device__ __forceinline__ unsigned ld2(const bf16* p) {
|
|
||||||
return *reinterpret_cast<const unsigned*>(p);
|
|
||||||
}
|
|
||||||
|
|
||||||
// pack two floats into one bf16x2 as .b32
|
|
||||||
__device__ __forceinline__ unsigned pk2(float a, float b) {
|
|
||||||
__nv_bfloat162 v = __floats2bfloat162_rn(a, b);
|
|
||||||
return *reinterpret_cast<unsigned*>(&v);
|
|
||||||
}
|
|
||||||
|
|
||||||
// pack two (non-contiguous) bf16 into one .b32
|
|
||||||
__device__ __forceinline__ unsigned pkb(bf16 a, bf16 b) {
|
|
||||||
__nv_bfloat162 v;
|
|
||||||
v.x = a;
|
|
||||||
v.y = b;
|
|
||||||
return *reinterpret_cast<unsigned*>(&v);
|
|
||||||
}
|
|
||||||
|
|
||||||
// ldmatrix: cooperatively load mma fragments from smem (one instruction per
|
|
||||||
// 16x16 / 16x8 tile) with the exact register layout mma expects — replaces the
|
|
||||||
// scalar per-thread fragment packing, cutting shared-load instructions and bank
|
|
||||||
// conflicts. Each lane supplies the shared address of one 8-wide row.
|
|
||||||
__device__ __forceinline__ void ldmatrix_x4(unsigned* r, const bf16* p) {
|
|
||||||
unsigned a = __cvta_generic_to_shared(p);
|
|
||||||
asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];"
|
|
||||||
: "=r"(r[0]), "=r"(r[1]), "=r"(r[2]), "=r"(r[3])
|
|
||||||
: "r"(a));
|
|
||||||
}
|
|
||||||
__device__ __forceinline__ void ldmatrix_x2(unsigned* r, const bf16* p) {
|
|
||||||
unsigned a = __cvta_generic_to_shared(p);
|
|
||||||
asm volatile("ldmatrix.sync.aligned.m8n8.x2.shared.b16 {%0,%1}, [%2];"
|
|
||||||
: "=r"(r[0]), "=r"(r[1])
|
|
||||||
: "r"(a));
|
|
||||||
}
|
|
||||||
__device__ __forceinline__ void ldmatrix_x2_trans(unsigned* r, const bf16* p) {
|
|
||||||
unsigned a = __cvta_generic_to_shared(p);
|
|
||||||
asm volatile("ldmatrix.sync.aligned.m8n8.x2.trans.shared.b16 {%0,%1}, [%2];"
|
|
||||||
: "=r"(r[0]), "=r"(r[1])
|
|
||||||
: "r"(a));
|
|
||||||
}
|
|
||||||
|
|
||||||
// XOR swizzle for shared-memory column at 8-bf16 chunk granularity.
|
|
||||||
// Eliminates ldmatrix bank conflicts without LD padding: consecutive rows
|
|
||||||
// land in distinct bank groups. swiz_col(d, r, mask) = ((d>>3)^(r&mask))<<3 | (d&7).
|
|
||||||
// mask must cover log2(HEAD_DIM/8) chunk bits but stay within LD: use 7 for
|
|
||||||
// HEAD_DIM>=64 (8+ chunks), 3 for HEAD_DIM=32 (4 chunks). Default 7 keeps
|
|
||||||
// existing HEAD_DIM>=64 call sites working unchanged.
|
|
||||||
__device__ __forceinline__ int swiz_col(int d, int r, int mask = 7) {
|
|
||||||
return ((d >> 3) ^ (r & mask)) << 3 | (d & 7);
|
|
||||||
}
|
|
||||||
|
|
||||||
// cp.async: copy 16 bytes (8 bf16) from global to shared memory directly,
|
|
||||||
// bypassing registers. Eliminates shared-store bank conflicts and cuts
|
|
||||||
// load-loop instruction count in half (1 cp.async vs 1 LDG + 1 STS).
|
|
||||||
// Requires sm_80+.
|
|
||||||
__device__ __forceinline__ void cp_async_16(bf16* smem_ptr, const void* gmem_ptr) {
|
|
||||||
unsigned smem_addr = __cvta_generic_to_shared(smem_ptr);
|
|
||||||
asm volatile("cp.async.ca.shared.global [%0], [%1], 16;"
|
|
||||||
:: "r"(smem_addr), "l"(gmem_ptr));
|
|
||||||
}
|
|
||||||
|
|
||||||
__device__ __forceinline__ void cp_async_commit() {
|
|
||||||
asm volatile("cp.async.commit_group;");
|
|
||||||
}
|
|
||||||
|
|
||||||
__device__ __forceinline__ void cp_async_wait_all() {
|
|
||||||
asm volatile("cp.async.wait_all;");
|
|
||||||
}
|
|
||||||
|
|
||||||
// Wait until at most N commit groups are still in flight. Used for
|
|
||||||
// double-buffered pipelining: wait_group<1> lets the next tile's cp.async
|
|
||||||
// continue while ensuring the current tile's data is ready.
|
|
||||||
template <int N>
|
|
||||||
__device__ __forceinline__ void cp_async_wait_group() {
|
|
||||||
asm volatile("cp.async.wait_group %0;" :: "n"(N));
|
|
||||||
}
|
|
||||||
@@ -1,100 +0,0 @@
|
|||||||
#include "gqa_prefill_attn.cuh"
|
|
||||||
#include <torch/extension.h>
|
|
||||||
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
|
||||||
#include "gqa_prefill_attn_mma.cuh"
|
|
||||||
#endif
|
|
||||||
|
|
||||||
template <int HEAD_DIM>
|
|
||||||
static void dispatch_prefill(GQAParams& p) {
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
|
||||||
constexpr int WARPS = 4, BC = 32, BR = 16, LD = HEAD_DIM;
|
|
||||||
dim3 grid((p.q_len + BR * WARPS - 1) / (BR * WARPS), p.q_head, p.batch);
|
|
||||||
dim3 block(WARPS * 32, 1, 1);
|
|
||||||
// sK + sV (each BC*LD) + shared sQ staging (BR*LD)
|
|
||||||
int smem = (2 * BC * LD + BR * LD) * (int)sizeof(bf16);
|
|
||||||
cudaFuncSetAttribute(gqa_prefill_attn_mma_kernel<HEAD_DIM, WARPS, BC>,
|
|
||||||
cudaFuncAttributeMaxDynamicSharedMemorySize, smem);
|
|
||||||
gqa_prefill_attn_mma_kernel<HEAD_DIM, WARPS, BC><<<grid, block, smem>>>(p);
|
|
||||||
#else
|
|
||||||
constexpr int G = 8, ROWS = 32, P_BC = 32;
|
|
||||||
dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch);
|
|
||||||
dim3 block(G, ROWS, 1);
|
|
||||||
size_t smem = 2 * P_BC * HEAD_DIM * sizeof(bf16);
|
|
||||||
gqa_prefill_attn_kernel_t<HEAD_DIM, G, ROWS, P_BC><<<grid, block, smem>>>(p);
|
|
||||||
#endif
|
|
||||||
}
|
|
||||||
|
|
||||||
torch::Tensor gqa_prefill_attn(
|
|
||||||
torch::Tensor q,
|
|
||||||
torch::Tensor k,
|
|
||||||
torch::Tensor v,
|
|
||||||
c10::optional<torch::Tensor> mask,
|
|
||||||
bool is_causal = false,
|
|
||||||
int64_t causal_offset = 0,
|
|
||||||
c10::optional<double> scale = c10::nullopt
|
|
||||||
) {
|
|
||||||
TORCH_CHECK(q.is_cuda() && k.is_cuda() && v.is_cuda());
|
|
||||||
TORCH_CHECK(q.dtype() == torch::kBFloat16);
|
|
||||||
TORCH_CHECK(k.dtype() == torch::kBFloat16);
|
|
||||||
TORCH_CHECK(v.dtype() == torch::kBFloat16);
|
|
||||||
|
|
||||||
GQAParams p;
|
|
||||||
p.batch = q.size(0);
|
|
||||||
p.q_head = q.size(1);
|
|
||||||
p.kv_head = k.size(1);
|
|
||||||
p.q_len = q.size(2);
|
|
||||||
p.kv_len = k.size(2);
|
|
||||||
p.head_dim = q.size(3);
|
|
||||||
TORCH_CHECK(p.head_dim % 16 == 0, "head_dim must be multiple of 16");
|
|
||||||
p.use_mask = mask.has_value();
|
|
||||||
p.is_causal = (int)is_causal;
|
|
||||||
p.causal_offset = (int)causal_offset;
|
|
||||||
p.scale = scale.has_value() ? (float)scale.value() : 1.0f / sqrtf((float)p.head_dim);
|
|
||||||
p.q = (const bf16*)q.data_ptr();
|
|
||||||
p.k = (const bf16*)k.data_ptr();
|
|
||||||
p.v = (const bf16*)v.data_ptr();
|
|
||||||
if (p.use_mask) {
|
|
||||||
TORCH_CHECK(mask.value().dtype() == torch::kBool);
|
|
||||||
TORCH_CHECK(mask.value().dim() == 2);
|
|
||||||
TORCH_CHECK(mask.value().size(0) == p.batch);
|
|
||||||
TORCH_CHECK(mask.value().size(1) == p.kv_len);
|
|
||||||
p.mask = mask.value().data_ptr<bool>();
|
|
||||||
} else {
|
|
||||||
p.mask = nullptr;
|
|
||||||
}
|
|
||||||
|
|
||||||
auto O = torch::empty_like(q);
|
|
||||||
p.o = (bf16*)O.data_ptr();
|
|
||||||
|
|
||||||
switch (p.head_dim) {
|
|
||||||
case 32:
|
|
||||||
dispatch_prefill<32>(p);
|
|
||||||
break;
|
|
||||||
case 64:
|
|
||||||
dispatch_prefill<64>(p);
|
|
||||||
break;
|
|
||||||
case 128:
|
|
||||||
dispatch_prefill<128>(p);
|
|
||||||
break;
|
|
||||||
case 256:
|
|
||||||
dispatch_prefill<256>(p);
|
|
||||||
break;
|
|
||||||
default:
|
|
||||||
TORCH_CHECK(false, "prefill: unsupported head_dim ", p.head_dim,
|
|
||||||
" (supported: 32,64,128,256)");
|
|
||||||
}
|
|
||||||
return O;
|
|
||||||
}
|
|
||||||
|
|
||||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
|
||||||
m.def("gqa_prefill_attn", &gqa_prefill_attn,
|
|
||||||
py::arg("q"),
|
|
||||||
py::arg("k"),
|
|
||||||
py::arg("v"),
|
|
||||||
py::arg("mask") = py::none(),
|
|
||||||
py::arg("is_causal") = false,
|
|
||||||
py::arg("causal_offset") = 0,
|
|
||||||
py::arg("scale") = py::none(),
|
|
||||||
"GQA prefill (tensor-core mma on sm_80+, scalar fallback)");
|
|
||||||
}
|
|
||||||
@@ -1,246 +0,0 @@
|
|||||||
#pragma once
|
|
||||||
#include "gqa_common.cuh"
|
|
||||||
#include "gqa_mma_utils.cuh"
|
|
||||||
|
|
||||||
// Tensor-core prefill, register-resident flash attention (raw mma.sync PTX).
|
|
||||||
// One warp owns BR=16 query rows. S = Q@K^T and O = P@V run on bf16 tensor
|
|
||||||
// cores via mma.sync.m16n8k16 (f32 accumulate). Q stays resident in registers;
|
|
||||||
// S, O, and the online-softmax stats (m, l) live in registers too — nothing is
|
|
||||||
// staged through shared memory except the cooperatively-loaded K/V tiles. The
|
|
||||||
// mma fragment layout is used directly: the S accumulator (f32) maps element-
|
|
||||||
// for-element onto the P matrix_a (bf16) operand, so softmax needs no shuffle
|
|
||||||
// repack; row reductions fold across the 4-lane thread group. Templated on
|
|
||||||
// <HEAD_DIM, WARPS, BC> with BC a multiple of 16.
|
|
||||||
//
|
|
||||||
// Optimizations: shared sQ staging (single area, serialized per-warp load)
|
|
||||||
// → cuts smem; pre-scale Q by attention scale during Q load; cp.async global→
|
|
||||||
// shared for K/V; scalar fallback only for the last partial tile; causal tile
|
|
||||||
// skipping (block-level early break + warp-level skip); XOR swizzle (swiz_col)
|
|
||||||
// → eliminates ldmatrix bank conflicts without LD padding (LD=HEAD_DIM).
|
|
||||||
|
|
||||||
template <int HEAD_DIM, int WARPS, int BC>
|
|
||||||
__global__ void gqa_prefill_attn_mma_kernel(GQAParams p) {
|
|
||||||
constexpr int BR = 16;
|
|
||||||
constexpr int KD = HEAD_DIM / 16; // Q/K k-tiles
|
|
||||||
constexpr int NC8 = BC / 8; // S n-tiles (N=8 each)
|
|
||||||
constexpr int KT2 = BC / 16; // P k-tiles (K=16 each)
|
|
||||||
constexpr int DN8 = HEAD_DIM / 8; // O n-tiles (N=8 each)
|
|
||||||
constexpr int LD = HEAD_DIM; // XOR swizzle (swiz_col) handles bank conflicts
|
|
||||||
constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1); // chunk bits, stay within LD
|
|
||||||
|
|
||||||
const int warp = threadIdx.x / 32;
|
|
||||||
const int lane = threadIdx.x % 32;
|
|
||||||
const int gid = lane >> 2; // 0..7 → rows gid, gid+8
|
|
||||||
const int tid4 = lane & 3; // 0..3
|
|
||||||
const int nthreads = WARPS * 32;
|
|
||||||
|
|
||||||
const int q_head = blockIdx.y;
|
|
||||||
const int batch = blockIdx.z;
|
|
||||||
const int kv_head = q_head / (p.q_head / p.kv_head);
|
|
||||||
const int qrow0 = (blockIdx.x * WARPS + warp) * BR;
|
|
||||||
|
|
||||||
extern __shared__ __align__(16) bf16 smem[];
|
|
||||||
bf16* sK = smem; // [BC][LD]
|
|
||||||
bf16* sV = sK + BC * LD; // [BC][LD]
|
|
||||||
bf16* sQ = sV + BC * LD; // shared staging [BR][LD]
|
|
||||||
|
|
||||||
// Q resident A-fragments (loaded once per warp via shared staging).
|
|
||||||
// Pre-scale by attention scale so softmax doesn't need to multiply later.
|
|
||||||
const int q_base = ((batch * p.q_head + q_head) * p.q_len) * HEAD_DIM;
|
|
||||||
unsigned Qa[KD][4];
|
|
||||||
bf16 scale_bf16 = __float2bfloat16(p.scale);
|
|
||||||
int qrow_l = (lane & 7) + (lane & 8); // 0..15
|
|
||||||
int qcol_l = (lane & 16) ? 8 : 0;
|
|
||||||
for (int w = 0; w < WARPS; w++) {
|
|
||||||
if (warp == w) {
|
|
||||||
for (int i = lane; i < BR * HEAD_DIM; i += 32) {
|
|
||||||
int r = i / HEAD_DIM, d = i % HEAD_DIM;
|
|
||||||
int qr = qrow0 + r;
|
|
||||||
bf16 qv = (qr < p.q_len) ? p.q[q_base + qr * HEAD_DIM + d]
|
|
||||||
: __float2bfloat16(0.0f);
|
|
||||||
sQ[r * LD + swiz_col(d, r, SWIZ_MASK)] = __hmul(qv, scale_bf16);
|
|
||||||
}
|
|
||||||
__syncwarp();
|
|
||||||
#pragma unroll
|
|
||||||
for (int kt = 0; kt < KD; kt++)
|
|
||||||
ldmatrix_x4(Qa[kt], &sQ[qrow_l * LD + swiz_col(kt * 16 + qcol_l, qrow_l, SWIZ_MASK)]);
|
|
||||||
}
|
|
||||||
__syncthreads(); // prevent next warp from overwriting sQ prematurely
|
|
||||||
}
|
|
||||||
|
|
||||||
float Oacc[DN8][4];
|
|
||||||
#pragma unroll
|
|
||||||
for (int j = 0; j < DN8; j++)
|
|
||||||
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
|
||||||
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
|
||||||
|
|
||||||
const int kv_base = ((batch * p.kv_head + kv_head) * p.kv_len) * HEAD_DIM;
|
|
||||||
const int tiles = (p.kv_len + BC - 1) / BC;
|
|
||||||
const int qr0 = qrow0 + gid; // row for c0/c1
|
|
||||||
const int qr1 = qrow0 + gid + 8; // row for c2/c3
|
|
||||||
|
|
||||||
// Causal tile-skip bounds (no-op when is_causal == 0)
|
|
||||||
const int use_skip = p.is_causal;
|
|
||||||
const int max_kv = qrow0 + BR - 1 + p.causal_offset;
|
|
||||||
const int block_max_kv =
|
|
||||||
blockIdx.x * WARPS * BR + WARPS * BR - 1 + p.causal_offset;
|
|
||||||
const int has_mask = p.use_mask && p.mask;
|
|
||||||
const int mb = batch * p.kv_len;
|
|
||||||
|
|
||||||
for (int ti = 0; ti < tiles; ti++) {
|
|
||||||
int kv0 = ti * BC;
|
|
||||||
|
|
||||||
// Block-level causal early break
|
|
||||||
if (use_skip && kv0 > block_max_kv) break;
|
|
||||||
|
|
||||||
// ---- load K/V tile to shared memory (cp.async on full tiles) ----
|
|
||||||
bool full_tile = (kv0 + BC <= p.kv_len);
|
|
||||||
if (full_tile) {
|
|
||||||
constexpr int VEC = 8; // bf16 per cp.async unit (16 bytes)
|
|
||||||
int total = BC * HEAD_DIM;
|
|
||||||
#pragma unroll
|
|
||||||
for (int i = threadIdx.x * VEC; i < total; i += nthreads * VEC) {
|
|
||||||
int r = i / HEAD_DIM;
|
|
||||||
int d = i % HEAD_DIM;
|
|
||||||
int kc = kv0 + r;
|
|
||||||
cp_async_16(&sK[r * LD + swiz_col(d, r, SWIZ_MASK)], &p.k[kv_base + kc * HEAD_DIM + d]);
|
|
||||||
cp_async_16(&sV[r * LD + swiz_col(d, r, SWIZ_MASK)], &p.v[kv_base + kc * HEAD_DIM + d]);
|
|
||||||
}
|
|
||||||
cp_async_commit();
|
|
||||||
cp_async_wait_all();
|
|
||||||
} else {
|
|
||||||
for (int i = threadIdx.x; i < BC * HEAD_DIM; i += nthreads) {
|
|
||||||
int r = i / HEAD_DIM, d = i % HEAD_DIM;
|
|
||||||
int kc = kv0 + r;
|
|
||||||
bf16 z = __float2bfloat16(0.0f);
|
|
||||||
sK[r * LD + swiz_col(d, r, SWIZ_MASK)] = (kc < p.kv_len)
|
|
||||||
? p.k[kv_base + kc * HEAD_DIM + d] : z;
|
|
||||||
sV[r * LD + swiz_col(d, r, SWIZ_MASK)] = (kc < p.kv_len)
|
|
||||||
? p.v[kv_base + kc * HEAD_DIM + d] : z;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
__syncthreads();
|
|
||||||
|
|
||||||
// Warp-level causal skip
|
|
||||||
if (!use_skip || kv0 <= max_kv) {
|
|
||||||
|
|
||||||
// S = Q @ K^T → Sacc[n8][0..3] (n8: 8 kv cols each)
|
|
||||||
float Sacc[NC8][4];
|
|
||||||
#pragma unroll
|
|
||||||
for (int n8 = 0; n8 < NC8; n8++) {
|
|
||||||
Sacc[n8][0] = Sacc[n8][1] = Sacc[n8][2] = Sacc[n8][3] = 0.0f;
|
|
||||||
int krow_l = n8 * 8 + (lane & 7);
|
|
||||||
int kcol_h = (lane & 8) ? 8 : 0;
|
|
||||||
#pragma unroll
|
|
||||||
for (int kt = 0; kt < KD; kt++) {
|
|
||||||
unsigned b[2];
|
|
||||||
ldmatrix_x2(b, &sK[krow_l * LD + swiz_col(kt * 16 + kcol_h, krow_l, SWIZ_MASK)]);
|
|
||||||
mma16816(Sacc[n8], Qa[kt], b, Sacc[n8]);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ---- online softmax (in registers) ----
|
|
||||||
// Q is pre-scaled, so Sacc already includes the attention scale.
|
|
||||||
int maxc0 = p.is_causal ? min(p.kv_len, qr0 + p.causal_offset + 1)
|
|
||||||
: p.kv_len;
|
|
||||||
int maxc1 = p.is_causal ? min(p.kv_len, qr1 + p.causal_offset + 1)
|
|
||||||
: p.kv_len;
|
|
||||||
float rmax0 = -FLT_MAX, rmax1 = -FLT_MAX;
|
|
||||||
#pragma unroll
|
|
||||||
for (int n8 = 0; n8 < NC8; n8++) {
|
|
||||||
int cc = kv0 + n8 * 8 + 2 * tid4;
|
|
||||||
int c1 = cc + 1;
|
|
||||||
bool b0 = (cc >= maxc0) || (has_mask && !p.mask[mb + cc]);
|
|
||||||
bool b1 = (c1 >= maxc0) || (has_mask && !p.mask[mb + c1]);
|
|
||||||
bool b2 = (cc >= maxc1) || (has_mask && !p.mask[mb + cc]);
|
|
||||||
bool b3 = (c1 >= maxc1) || (has_mask && !p.mask[mb + c1]);
|
|
||||||
float s0 = b0 ? -FLT_MAX : Sacc[n8][0];
|
|
||||||
float s1 = b1 ? -FLT_MAX : Sacc[n8][1];
|
|
||||||
float s2 = b2 ? -FLT_MAX : Sacc[n8][2];
|
|
||||||
float s3 = b3 ? -FLT_MAX : Sacc[n8][3];
|
|
||||||
Sacc[n8][0] = s0; Sacc[n8][1] = s1;
|
|
||||||
Sacc[n8][2] = s2; Sacc[n8][3] = s3;
|
|
||||||
rmax0 = fmaxf(rmax0, fmaxf(s0, s1));
|
|
||||||
rmax1 = fmaxf(rmax1, fmaxf(s2, s3));
|
|
||||||
}
|
|
||||||
rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 1));
|
|
||||||
rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 2));
|
|
||||||
rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 1));
|
|
||||||
rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 2));
|
|
||||||
|
|
||||||
float nm0 = fmaxf(m0, rmax0), nm1 = fmaxf(m1, rmax1);
|
|
||||||
float corr0 = (nm0 == -FLT_MAX) ? 1.0f : __expf(m0 - nm0);
|
|
||||||
float corr1 = (nm1 == -FLT_MAX) ? 1.0f : __expf(m1 - nm1);
|
|
||||||
|
|
||||||
float rsum0 = 0.0f, rsum1 = 0.0f;
|
|
||||||
#pragma unroll
|
|
||||||
for (int n8 = 0; n8 < NC8; n8++) {
|
|
||||||
float p0 = (Sacc[n8][0] == -FLT_MAX) ? 0.0f
|
|
||||||
: __expf(Sacc[n8][0] - nm0);
|
|
||||||
float p1 = (Sacc[n8][1] == -FLT_MAX) ? 0.0f
|
|
||||||
: __expf(Sacc[n8][1] - nm0);
|
|
||||||
float p2 = (Sacc[n8][2] == -FLT_MAX) ? 0.0f
|
|
||||||
: __expf(Sacc[n8][2] - nm1);
|
|
||||||
float p3 = (Sacc[n8][3] == -FLT_MAX) ? 0.0f
|
|
||||||
: __expf(Sacc[n8][3] - nm1);
|
|
||||||
Sacc[n8][0] = p0; Sacc[n8][1] = p1;
|
|
||||||
Sacc[n8][2] = p2; Sacc[n8][3] = p3;
|
|
||||||
rsum0 += p0 + p1;
|
|
||||||
rsum1 += p2 + p3;
|
|
||||||
}
|
|
||||||
rsum0 += __shfl_xor_sync(0xFFFFFFFF, rsum0, 1);
|
|
||||||
rsum0 += __shfl_xor_sync(0xFFFFFFFF, rsum0, 2);
|
|
||||||
rsum1 += __shfl_xor_sync(0xFFFFFFFF, rsum1, 1);
|
|
||||||
rsum1 += __shfl_xor_sync(0xFFFFFFFF, rsum1, 2);
|
|
||||||
l0 = l0 * corr0 + rsum0;
|
|
||||||
l1 = l1 * corr1 + rsum1;
|
|
||||||
m0 = nm0; m1 = nm1;
|
|
||||||
|
|
||||||
// rescale O accumulator by per-row correction
|
|
||||||
#pragma unroll
|
|
||||||
for (int j = 0; j < DN8; j++) {
|
|
||||||
Oacc[j][0] *= corr0; Oacc[j][1] *= corr0;
|
|
||||||
Oacc[j][2] *= corr1; Oacc[j][3] *= corr1;
|
|
||||||
}
|
|
||||||
|
|
||||||
// O += P @ V
|
|
||||||
#pragma unroll
|
|
||||||
for (int kt2 = 0; kt2 < KT2; kt2++) {
|
|
||||||
unsigned Pa[4];
|
|
||||||
Pa[0] = pk2(Sacc[kt2 * 2][0], Sacc[kt2 * 2][1]);
|
|
||||||
Pa[1] = pk2(Sacc[kt2 * 2][2], Sacc[kt2 * 2][3]);
|
|
||||||
Pa[2] = pk2(Sacc[kt2 * 2 + 1][0], Sacc[kt2 * 2 + 1][1]);
|
|
||||||
Pa[3] = pk2(Sacc[kt2 * 2 + 1][2], Sacc[kt2 * 2 + 1][3]);
|
|
||||||
int vrow_l = kt2 * 16 + (lane & 15);
|
|
||||||
#pragma unroll
|
|
||||||
for (int dn8 = 0; dn8 < DN8; dn8++) {
|
|
||||||
unsigned b[2];
|
|
||||||
ldmatrix_x2_trans(b, &sV[vrow_l * LD + swiz_col(dn8 * 8, vrow_l, SWIZ_MASK)]);
|
|
||||||
mma16816(Oacc[dn8], Pa, b, Oacc[dn8]);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
} // if active (warp-level causal skip)
|
|
||||||
__syncthreads();
|
|
||||||
}
|
|
||||||
|
|
||||||
// ---- write output ----
|
|
||||||
float rl0 = (l0 > 1e-20f) ? (1.0f / l0) : 0.0f;
|
|
||||||
float rl1 = (l1 > 1e-20f) ? (1.0f / l1) : 0.0f;
|
|
||||||
const int o_base = ((batch * p.q_head + q_head) * p.q_len) * HEAD_DIM;
|
|
||||||
#pragma unroll
|
|
||||||
for (int dn8 = 0; dn8 < DN8; dn8++) {
|
|
||||||
int d = dn8 * 8 + 2 * tid4;
|
|
||||||
if (qr0 < p.q_len) {
|
|
||||||
p.o[o_base + qr0 * HEAD_DIM + d] =
|
|
||||||
__float2bfloat16(Oacc[dn8][0] * rl0);
|
|
||||||
p.o[o_base + qr0 * HEAD_DIM + d + 1] =
|
|
||||||
__float2bfloat16(Oacc[dn8][1] * rl0);
|
|
||||||
}
|
|
||||||
if (qr1 < p.q_len) {
|
|
||||||
p.o[o_base + qr1 * HEAD_DIM + d] =
|
|
||||||
__float2bfloat16(Oacc[dn8][2] * rl1);
|
|
||||||
p.o[o_base + qr1 * HEAD_DIM + d + 1] =
|
|
||||||
__float2bfloat16(Oacc[dn8][3] * rl1);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -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)"
|
||||||
|
);
|
||||||
|
}
|
||||||
@@ -0,0 +1,185 @@
|
|||||||
|
/*
|
||||||
|
Pure-C test — uses shared dispatcher.
|
||||||
|
nvcc -I csrc -arch=sm_89 -O3 \
|
||||||
|
--use_fast_math --ptxas-options=-O3 --extra-device-vectorization \
|
||||||
|
csrc/tests/attn_decode_test.cu -o test && ./test
|
||||||
|
*/
|
||||||
|
|
||||||
|
#include "test_utils.cuh"
|
||||||
|
#include "../kernels/attn_dispatchers.cuh"
|
||||||
|
|
||||||
|
// Split-K scratch (torch-free)
|
||||||
|
struct DecodeScratch {
|
||||||
|
float* o_part = nullptr;
|
||||||
|
float* ml_part = nullptr;
|
||||||
|
};
|
||||||
|
|
||||||
|
static void setup_scratch(AttentionParams<bf16>& p, DecodeScratch& sc) {
|
||||||
|
int max_splits = 32;
|
||||||
|
cudaMalloc(&sc.o_part, (size_t)p.batch * p.q_head * max_splits * p.head_dim * sizeof(float));
|
||||||
|
cudaMalloc(&sc.ml_part, (size_t)p.batch * p.q_head * max_splits * 2 * sizeof(float));
|
||||||
|
}
|
||||||
|
|
||||||
|
static void free_scratch(DecodeScratch& sc) {
|
||||||
|
cudaFree(sc.o_part); cudaFree(sc.ml_part);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Warmed-up, CUDA-event timed sweep over the production decode MMA path.
|
||||||
|
static void bench() {
|
||||||
|
const int cfgs[][5] = {
|
||||||
|
{1, 32, 4, 512, 128},
|
||||||
|
{1, 32, 4, 1024, 128},
|
||||||
|
{1, 32, 4, 2048, 128},
|
||||||
|
{1, 32, 4, 4096, 128},
|
||||||
|
{16, 32, 4, 2048, 128},
|
||||||
|
{32, 32, 4, 1024, 128},
|
||||||
|
};
|
||||||
|
const int WARMUP = 10, ITERS = 100;
|
||||||
|
printf("\n===== DECODE BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS);
|
||||||
|
print_bench_header();
|
||||||
|
|
||||||
|
for (int ci = 0; ci < 6; ci++) {
|
||||||
|
int B = cfgs[ci][0], Hq = cfgs[ci][1], Hk = cfgs[ci][2];
|
||||||
|
int sl = cfgs[ci][3], D = cfgs[ci][4];
|
||||||
|
size_t nQ = (size_t)B * Hq * D;
|
||||||
|
size_t nKV = (size_t)B * Hk * sl * D;
|
||||||
|
|
||||||
|
bf16 *dQ, *dK, *dV, *dO;
|
||||||
|
cudaMalloc(&dQ, nQ*2); cudaMalloc(&dK, nKV*2);
|
||||||
|
cudaMalloc(&dV, nKV*2); cudaMalloc(&dO, nQ*2);
|
||||||
|
size_t big = nQ > nKV ? nQ : nKV; bf16* tmp = new bf16[big];
|
||||||
|
for (size_t i = 0; i < nQ; i++) tmp[i] = f2bf(randf());
|
||||||
|
cudaMemcpy(dQ, tmp, nQ*2, cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf());
|
||||||
|
cudaMemcpy(dK, tmp, nKV*2, cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf());
|
||||||
|
cudaMemcpy(dV, tmp, nKV*2, cudaMemcpyHostToDevice);
|
||||||
|
delete[] tmp;
|
||||||
|
|
||||||
|
AttentionParams<bf16> p;
|
||||||
|
p.batch = B; p.q_head = Hq; p.kv_head = Hk; p.q_len = 1; p.kv_len = sl;
|
||||||
|
p.head_dim = D; p.use_mask = 0; p.causal_offset = -1;
|
||||||
|
p.scale = 1.0f / sqrtf((float)D);
|
||||||
|
set_default_strides(p);
|
||||||
|
p.q = dQ; p.k = dK; p.v = dV; p.mask = nullptr; p.o = dO;
|
||||||
|
|
||||||
|
DecodeScratch sc;
|
||||||
|
setup_scratch(p, sc);
|
||||||
|
p.o_part = sc.o_part; p.ml_part = sc.ml_part;
|
||||||
|
|
||||||
|
auto launch = [&]() { dispatch_by_head_dim(D, [&]<int H>() { dispatch_decode<H>(p); }); };
|
||||||
|
double flops = 4.0 * B * Hq * (double)sl * D;
|
||||||
|
double bytes = 2.0 * (2.0 * nKV * sizeof(bf16));
|
||||||
|
BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops, bytes);
|
||||||
|
|
||||||
|
char cfg[64];
|
||||||
|
snprintf(cfg, sizeof(cfg),
|
||||||
|
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d causal=%d",
|
||||||
|
B, Hq, Hk, 1, sl, D, 0);
|
||||||
|
print_bench_row(cfg, r);
|
||||||
|
|
||||||
|
cudaFree(dQ); cudaFree(dK); cudaFree(dV); cudaFree(dO);
|
||||||
|
free_scratch(sc);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
static int run_test(int B, int Hq, int Hk, int sl, int D, int causal) {
|
||||||
|
int gs = Hq / Hk;
|
||||||
|
printf("=== B=%d Hq=%d Hk=%d seq=%d D=%d gs=%d causal=%d ===\n",
|
||||||
|
B,Hq,Hk,sl,D,gs,causal);
|
||||||
|
|
||||||
|
size_t nQ = B*Hq*1*D, nKV = B*Hk*sl*D;
|
||||||
|
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
|
||||||
|
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
|
||||||
|
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
|
||||||
|
|
||||||
|
bool* hMask=new bool[B*sl];
|
||||||
|
for (int i=0;i<B*sl;i++) hMask[i]=true;
|
||||||
|
|
||||||
|
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
||||||
|
bool* dMask;
|
||||||
|
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
||||||
|
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
||||||
|
cudaMalloc(&dMask,B*sl);
|
||||||
|
|
||||||
|
tmp=new bf16[max(nQ,nKV)];
|
||||||
|
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
|
||||||
|
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
|
||||||
|
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
|
||||||
|
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||||
|
cudaMemcpy(dMask,hMask,B*sl,cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
AttentionParams<bf16> p;
|
||||||
|
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=1; p.kv_len=sl; p.head_dim=D;
|
||||||
|
p.use_mask=0; p.causal_offset=causal?0:-1;
|
||||||
|
p.scale=1.0f/sqrtf((float)D);
|
||||||
|
set_default_strides(p);
|
||||||
|
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||||
|
|
||||||
|
DecodeScratch sc;
|
||||||
|
setup_scratch(p, sc);
|
||||||
|
p.o_part = sc.o_part; p.ml_part = sc.ml_part;
|
||||||
|
|
||||||
|
double t0=now_ms();
|
||||||
|
dispatch_by_head_dim(D, [&]<int H>() { dispatch_decode<H>(p); });
|
||||||
|
cudaDeviceSynchronize();
|
||||||
|
double kms=now_ms()-t0;
|
||||||
|
cudaError_t err=cudaGetLastError();
|
||||||
|
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
|
||||||
|
|
||||||
|
bf16* hOut=new bf16[nQ];
|
||||||
|
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
|
||||||
|
|
||||||
|
float* ref=new float[nQ];
|
||||||
|
cpu_attention_ref(hQ, hK, hV, hMask, ref, B, Hq, Hk, 1, sl, D, causal ? 0 : -1);
|
||||||
|
|
||||||
|
float max_abs_err=0, max_rel_err=0;
|
||||||
|
for (size_t i=0;i<nQ;i++){
|
||||||
|
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
||||||
|
if(err>max_abs_err) max_abs_err=err;
|
||||||
|
float rel=err/fmaxf(fabsf(ref[i]), 1e-8f);
|
||||||
|
if(rel>max_rel_err) max_rel_err=rel;
|
||||||
|
}
|
||||||
|
const float atol=0.01f, rtol=0.01f;
|
||||||
|
bool pass=true;
|
||||||
|
for (size_t i=0;i<nQ;i++){
|
||||||
|
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
||||||
|
if (err > atol + rtol * fabsf(ref[i])) { pass=false; break; }
|
||||||
|
}
|
||||||
|
printf("kernel: %.3f ms max_abs_err: %.6e max_rel_err: %.6e %s\n\n",
|
||||||
|
kms, max_abs_err, max_rel_err, pass?"PASS":"FAIL");
|
||||||
|
|
||||||
|
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);cudaFree(dMask);
|
||||||
|
free_scratch(sc);
|
||||||
|
delete[]hQ;delete[]hK;delete[]hV;delete[]hMask;delete[]hOut;delete[]ref;delete[]tmp;
|
||||||
|
|
||||||
|
return pass ? 0 : 1;
|
||||||
|
}
|
||||||
|
|
||||||
|
int main() {
|
||||||
|
const int configs[][6] = {
|
||||||
|
{1, 2, 1, 64, 32, 0},
|
||||||
|
{1, 32, 4, 512, 128, 0},
|
||||||
|
{1, 32, 4, 1024, 128, 0},
|
||||||
|
{1, 32, 4, 512, 128, 1},
|
||||||
|
};
|
||||||
|
int n_cfgs = sizeof(configs) / sizeof(configs[0]);
|
||||||
|
int fail = 0;
|
||||||
|
|
||||||
|
for (int ci = 0; ci < n_cfgs; ci++) {
|
||||||
|
int B = configs[ci][0], Hq = configs[ci][1], Hk = configs[ci][2];
|
||||||
|
int sl = configs[ci][3], D = configs[ci][4], causal = configs[ci][5];
|
||||||
|
fail += run_test(B, Hq, Hk, sl, D, causal);
|
||||||
|
if (fail) break;
|
||||||
|
}
|
||||||
|
|
||||||
|
if (fail) {
|
||||||
|
printf("FAILED\n");
|
||||||
|
return fail;
|
||||||
|
}
|
||||||
|
printf("All tests passed!\n");
|
||||||
|
bench();
|
||||||
|
return 0;
|
||||||
|
}
|
||||||
@@ -0,0 +1,308 @@
|
|||||||
|
// Compile:
|
||||||
|
// nvcc -I csrc -arch=sm_89 -O3 --use_fast_math --ptxas-options=-O3 \
|
||||||
|
// --extra-device-vectorization csrc/tests/attn_paged_decode_test.cu \
|
||||||
|
// -o /tmp/test_paged && /tmp/test_paged
|
||||||
|
|
||||||
|
#include <cstring>
|
||||||
|
#include "test_utils.cuh"
|
||||||
|
#include "../kernels/attn_dispatchers.cuh"
|
||||||
|
|
||||||
|
static void gather_kv_cpu(
|
||||||
|
const bf16* h_k_pool, const bf16* h_v_pool,
|
||||||
|
const int64_t* h_pt, int B, int Hkv, int kv_len,
|
||||||
|
int page_size, int head_dim,
|
||||||
|
bf16* h_k, bf16* h_v)
|
||||||
|
{
|
||||||
|
int max_pages = (kv_len + page_size - 1) / page_size;
|
||||||
|
size_t page_stride = (size_t)page_size * Hkv * head_dim;
|
||||||
|
for (int b = 0; b < B; b++) {
|
||||||
|
for (int pos = 0; pos < kv_len; pos++) {
|
||||||
|
int log_pg = pos / page_size;
|
||||||
|
int pg_off = pos % page_size;
|
||||||
|
int phys = (int)h_pt[b * max_pages + log_pg];
|
||||||
|
for (int h = 0; h < Hkv; h++) {
|
||||||
|
size_t src_base = (size_t)phys * page_stride
|
||||||
|
+ (size_t)pg_off * Hkv * head_dim
|
||||||
|
+ h * head_dim;
|
||||||
|
size_t dst_base = ((size_t)b * Hkv + h) * kv_len * head_dim
|
||||||
|
+ (size_t)pos * head_dim;
|
||||||
|
memcpy(h_k + dst_base, h_k_pool + src_base, head_dim * sizeof(bf16));
|
||||||
|
memcpy(h_v + dst_base, h_v_pool + src_base, head_dim * sizeof(bf16));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int causal, int seed) {
|
||||||
|
printf("B=%d Hq=%d Hkv=%d kv_len=%d page_sz=%d head_dim=%d causal=%d ... ",
|
||||||
|
B, Hq, Hkv, kv_len, page_size, HEAD_DIM, causal);
|
||||||
|
fflush(stdout);
|
||||||
|
|
||||||
|
int max_pages = (kv_len + page_size - 1) / page_size;
|
||||||
|
int n_phys_pages = B * max_pages;
|
||||||
|
int max_splits = 32;
|
||||||
|
|
||||||
|
size_t sz_q = (size_t)B * Hq * 1 * HEAD_DIM * sizeof(bf16);
|
||||||
|
size_t sz_o = sz_q;
|
||||||
|
size_t sz_kv = (size_t)n_phys_pages * page_size * Hkv * HEAD_DIM * sizeof(bf16);
|
||||||
|
size_t sz_pt = (size_t)B * max_pages * sizeof(int64_t);
|
||||||
|
size_t sz_op = (size_t)B * Hq * max_splits * HEAD_DIM * sizeof(float);
|
||||||
|
size_t sz_ml = (size_t)B * Hq * max_splits * 2 * sizeof(float);
|
||||||
|
|
||||||
|
bf16 *d_q, *d_o_paged;
|
||||||
|
bf16 *d_k_pool, *d_v_pool;
|
||||||
|
int64_t* d_pt;
|
||||||
|
float *d_op, *d_ml;
|
||||||
|
|
||||||
|
cudaMalloc(&d_q, sz_q);
|
||||||
|
cudaMalloc(&d_o_paged, sz_o);
|
||||||
|
cudaMalloc(&d_k_pool, sz_kv);
|
||||||
|
cudaMalloc(&d_v_pool, sz_kv);
|
||||||
|
cudaMalloc(&d_pt, sz_pt);
|
||||||
|
cudaMalloc(&d_op, sz_op);
|
||||||
|
cudaMalloc(&d_ml, sz_ml);
|
||||||
|
|
||||||
|
srand(seed);
|
||||||
|
auto rnd = [&]() { return (rand() / (float)RAND_MAX) * 2.0f - 1.0f; };
|
||||||
|
|
||||||
|
bf16* h_q = (bf16*)malloc(sz_q);
|
||||||
|
for (int i = 0; i < B * Hq * HEAD_DIM; i++)
|
||||||
|
h_q[i] = __float2bfloat16(rnd());
|
||||||
|
cudaMemcpy(d_q, h_q, sz_q, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
bf16* h_k_pool = (bf16*)malloc(sz_kv);
|
||||||
|
bf16* h_v_pool = (bf16*)malloc(sz_kv);
|
||||||
|
size_t ps = (size_t)page_size * Hkv * HEAD_DIM;
|
||||||
|
for (int pg = 0; pg < n_phys_pages; pg++) {
|
||||||
|
for (int off = 0; off < page_size; off++) {
|
||||||
|
for (int h = 0; h < Hkv; h++) {
|
||||||
|
for (int d = 0; d < HEAD_DIM; d++) {
|
||||||
|
float v = sinf((float)(pg * 7919 + off * 1049 + h * 331 + d));
|
||||||
|
size_t idx = (size_t)pg * ps + (size_t)off * Hkv * HEAD_DIM
|
||||||
|
+ h * HEAD_DIM + d;
|
||||||
|
h_k_pool[idx] = __float2bfloat16(v);
|
||||||
|
h_v_pool[idx] = __float2bfloat16(v * 0.3f);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
|
||||||
|
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
int64_t* h_pt = (int64_t*)malloc(sz_pt);
|
||||||
|
int next_pg = 0;
|
||||||
|
for (int b = 0; b < B; b++)
|
||||||
|
for (int p = 0; p < max_pages; p++)
|
||||||
|
h_pt[b * max_pages + p] = next_pg++;
|
||||||
|
cudaMemcpy(d_pt, h_pt, sz_pt, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
bf16* h_k_cont = (bf16*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(bf16));
|
||||||
|
bf16* h_v_cont = (bf16*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(bf16));
|
||||||
|
gather_kv_cpu(h_k_pool, h_v_pool, h_pt, B, Hkv, kv_len, page_size, HEAD_DIM, h_k_cont, h_v_cont);
|
||||||
|
|
||||||
|
float* h_q_f = (float*)malloc((size_t)B * Hq * HEAD_DIM * sizeof(float));
|
||||||
|
float* h_k_f = (float*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(float));
|
||||||
|
float* h_v_f = (float*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(float));
|
||||||
|
for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_q_f[i] = bf2f(h_q[i]);
|
||||||
|
for (int i = 0; i < B * kv_len * Hkv * HEAD_DIM; i++) {
|
||||||
|
h_k_f[i] = bf2f(h_k_cont[i]);
|
||||||
|
h_v_f[i] = bf2f(h_v_cont[i]);
|
||||||
|
}
|
||||||
|
|
||||||
|
float* h_o_ref = (float*)calloc(B * Hq * HEAD_DIM, sizeof(float));
|
||||||
|
cpu_attention_ref(h_q_f, h_k_f, h_v_f, nullptr, h_o_ref, B, Hq, Hkv,
|
||||||
|
1, kv_len, HEAD_DIM, causal ? 0 : -1);
|
||||||
|
|
||||||
|
PagedAttentionParams<bf16> p;
|
||||||
|
p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.q_len = 1;
|
||||||
|
p.kv_len = kv_len; p.head_dim = HEAD_DIM;
|
||||||
|
p.use_mask = 0; p.causal_offset = causal ? 0 : -1;
|
||||||
|
set_default_paged_strides(p);
|
||||||
|
p.scale = 1.0f / sqrtf((float)HEAD_DIM);
|
||||||
|
p.page_size = page_size; p.max_pages = max_pages;
|
||||||
|
p.page_table = d_pt;
|
||||||
|
p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
||||||
|
p.q = d_q; p.mask = nullptr; p.o = d_o_paged;
|
||||||
|
p.o_part = d_op; p.ml_part = d_ml;
|
||||||
|
|
||||||
|
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_decode<H>(p); });
|
||||||
|
cudaDeviceSynchronize();
|
||||||
|
|
||||||
|
bf16* h_o_bf16 = (bf16*)malloc(sz_o);
|
||||||
|
cudaMemcpy(h_o_bf16, d_o_paged, sz_o, cudaMemcpyDeviceToHost);
|
||||||
|
float* h_o_paged = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float));
|
||||||
|
for (int i = 0; i < B * Hq * HEAD_DIM; i++)
|
||||||
|
h_o_paged[i] = __bfloat162float(h_o_bf16[i]);
|
||||||
|
|
||||||
|
float max_abs_err = 0.0f, max_rel_err = 0.0f;
|
||||||
|
int bad_idx = -1;
|
||||||
|
for (int i = 0; i < B * Hq * HEAD_DIM; i++) {
|
||||||
|
float e = fabsf(h_o_paged[i] - h_o_ref[i]);
|
||||||
|
if (e > max_abs_err) { max_abs_err = e; bad_idx = i; }
|
||||||
|
float rel = e / fmaxf(fabsf(h_o_ref[i]), 1e-8f);
|
||||||
|
if (rel > max_rel_err) max_rel_err = rel;
|
||||||
|
}
|
||||||
|
|
||||||
|
const float atol = 0.01f, rtol = 0.01f;
|
||||||
|
bool pass = true;
|
||||||
|
for (int i = 0; i < B * Hq * HEAD_DIM; i++) {
|
||||||
|
float e = fabsf(h_o_paged[i] - h_o_ref[i]);
|
||||||
|
if (e > atol + rtol * fabsf(h_o_ref[i])) { pass = false; break; }
|
||||||
|
}
|
||||||
|
|
||||||
|
if (pass) {
|
||||||
|
printf("PASS (max_abs_err=%.4e max_rel_err=%.4e)\n", max_abs_err, max_rel_err);
|
||||||
|
} else {
|
||||||
|
int b = bad_idx / (Hq * HEAD_DIM);
|
||||||
|
int h = (bad_idx / HEAD_DIM) % Hq;
|
||||||
|
int d = bad_idx % HEAD_DIM;
|
||||||
|
printf("FAIL (max_abs_err=%.4e max_rel_err=%.4e at [%d,%d,%d]: ref=%.4f got=%.4f)\n",
|
||||||
|
max_abs_err, max_rel_err, b, h, d, h_o_ref[bad_idx], h_o_paged[bad_idx]);
|
||||||
|
printf(" ref[0..7]:");
|
||||||
|
for (int i = 0; i < 8 && i < HEAD_DIM; i++)
|
||||||
|
printf(" %.4f", h_o_ref[i]);
|
||||||
|
printf("\n got[0..7]:");
|
||||||
|
for (int i = 0; i < 8 && i < HEAD_DIM; i++)
|
||||||
|
printf(" %.4f", h_o_paged[i]);
|
||||||
|
printf("\n");
|
||||||
|
}
|
||||||
|
|
||||||
|
free(h_q); free(h_k_pool); free(h_v_pool); free(h_pt);
|
||||||
|
free(h_k_cont); free(h_v_cont);
|
||||||
|
free(h_q_f); free(h_k_f); free(h_v_f);
|
||||||
|
free(h_o_ref); free(h_o_bf16); free(h_o_paged);
|
||||||
|
cudaFree(d_q); cudaFree(d_o_paged);
|
||||||
|
cudaFree(d_k_pool); cudaFree(d_v_pool); cudaFree(d_pt);
|
||||||
|
cudaFree(d_op); cudaFree(d_ml);
|
||||||
|
|
||||||
|
return pass ? 0 : 1;
|
||||||
|
}
|
||||||
|
|
||||||
|
struct TestCase {
|
||||||
|
int head_dim;
|
||||||
|
int B, Hq, Hkv, kv_len, page_size, causal, seed;
|
||||||
|
};
|
||||||
|
|
||||||
|
static const TestCase TESTS[] = {
|
||||||
|
{128, 1, 1, 1, 8, 128, 0, 1},
|
||||||
|
{128, 1, 4, 4, 128, 128, 0, 2},
|
||||||
|
{128, 2, 4, 4, 256, 128, 0, 3},
|
||||||
|
{128, 1, 4, 1, 64, 64, 0, 4},
|
||||||
|
{128, 1, 8, 2, 64, 128, 0, 5},
|
||||||
|
{128, 2, 16, 4, 128, 128, 0, 6},
|
||||||
|
{64, 1, 4, 2, 32, 128, 0, 7},
|
||||||
|
{256, 1, 2, 1, 16, 128, 0, 8},
|
||||||
|
{32, 1, 4, 2, 32, 64, 0, 9},
|
||||||
|
{128, 3, 8, 2, 256, 128, 0, 10},
|
||||||
|
{128, 2, 32, 8, 512, 128, 0, 11},
|
||||||
|
{128, 1, 16, 2, 256, 128, 0, 12},
|
||||||
|
{128, 2, 32, 4, 512, 128, 0, 13},
|
||||||
|
{128, 2, 8, 2, 128, 128, 1, 14}, // causal
|
||||||
|
};
|
||||||
|
|
||||||
|
static int dispatch_test(const TestCase& tc) {
|
||||||
|
int r = 0;
|
||||||
|
dispatch_by_head_dim(tc.head_dim, [&]<int D>() {
|
||||||
|
r = run_test<D>(tc.B, tc.Hq, tc.Hkv, tc.kv_len, tc.page_size, tc.causal, tc.seed);
|
||||||
|
});
|
||||||
|
return r;
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
static void bench_config(int B, int Hq, int Hkv, int kv_len, int page_size) {
|
||||||
|
int max_pages = (kv_len + page_size - 1) / page_size;
|
||||||
|
int n_phys_pages = B * max_pages;
|
||||||
|
int max_splits = 32;
|
||||||
|
|
||||||
|
size_t sz_q = (size_t)B * Hq * 1 * HEAD_DIM * sizeof(bf16);
|
||||||
|
size_t sz_kv = (size_t)n_phys_pages * page_size * Hkv * HEAD_DIM * sizeof(bf16);
|
||||||
|
size_t sz_pt = (size_t)B * max_pages * sizeof(int64_t);
|
||||||
|
size_t sz_op = (size_t)B * Hq * max_splits * HEAD_DIM * sizeof(float);
|
||||||
|
size_t sz_ml = (size_t)B * Hq * max_splits * 2 * sizeof(float);
|
||||||
|
|
||||||
|
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
|
||||||
|
int64_t* d_pt;
|
||||||
|
float *d_op, *d_ml;
|
||||||
|
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
|
||||||
|
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
|
||||||
|
cudaMalloc(&d_pt, sz_pt);
|
||||||
|
cudaMalloc(&d_op, sz_op); cudaMalloc(&d_ml, sz_ml);
|
||||||
|
|
||||||
|
bf16* tmp = (bf16*)malloc(sz_kv > sz_q ? sz_kv : sz_q);
|
||||||
|
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) tmp[i] = f2bf(randf());
|
||||||
|
cudaMemcpy(d_q, tmp, sz_q, cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) tmp[i] = f2bf(randf());
|
||||||
|
cudaMemcpy(d_k_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
|
||||||
|
cudaMemcpy(d_v_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
int64_t* h_pt = (int64_t*)malloc(sz_pt);
|
||||||
|
int next_pg = 0;
|
||||||
|
for (int b = 0; b < B; b++)
|
||||||
|
for (int p = 0; p < max_pages; p++)
|
||||||
|
h_pt[b * max_pages + p] = next_pg++;
|
||||||
|
cudaMemcpy(d_pt, h_pt, sz_pt, cudaMemcpyHostToDevice);
|
||||||
|
free(h_pt);
|
||||||
|
|
||||||
|
PagedAttentionParams<bf16> pa;
|
||||||
|
pa.batch = B; pa.q_head = Hq; pa.kv_head = Hkv; pa.q_len = 1;
|
||||||
|
pa.kv_len = kv_len; pa.head_dim = HEAD_DIM;
|
||||||
|
pa.use_mask = 0; pa.causal_offset = -1;
|
||||||
|
set_default_paged_strides(pa);
|
||||||
|
pa.scale = 1.0f / sqrtf((float)HEAD_DIM);
|
||||||
|
pa.page_size = page_size; pa.max_pages = max_pages;
|
||||||
|
pa.page_table = d_pt;
|
||||||
|
pa.k_cache = d_k_pool; pa.v_cache = d_v_pool;
|
||||||
|
pa.q = d_q; pa.mask = nullptr; pa.o = d_o;
|
||||||
|
pa.o_part = d_op; pa.ml_part = d_ml;
|
||||||
|
|
||||||
|
const int WARMUP = 10, ITERS = 100;
|
||||||
|
auto launch = [&]() {
|
||||||
|
dispatch_by_head_dim(HEAD_DIM, [&]<int H>() { dispatch_paged_decode<H>(pa); });
|
||||||
|
};
|
||||||
|
double flops = 4.0 * B * Hq * (double)kv_len * HEAD_DIM;
|
||||||
|
size_t nKV = (size_t)B * Hkv * kv_len * HEAD_DIM;
|
||||||
|
double bytes = 2.0 * (2.0 * nKV * sizeof(bf16));
|
||||||
|
BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops, bytes);
|
||||||
|
|
||||||
|
char cfg[64];
|
||||||
|
snprintf(cfg, sizeof(cfg),
|
||||||
|
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d page=%3d",
|
||||||
|
B, Hq, Hkv, 1, kv_len, HEAD_DIM, page_size);
|
||||||
|
print_bench_row(cfg, r);
|
||||||
|
|
||||||
|
free(tmp);
|
||||||
|
cudaFree(d_q); cudaFree(d_o);
|
||||||
|
cudaFree(d_k_pool); cudaFree(d_v_pool); cudaFree(d_pt);
|
||||||
|
cudaFree(d_op); cudaFree(d_ml);
|
||||||
|
}
|
||||||
|
|
||||||
|
static void bench() {
|
||||||
|
printf("\n===== PAGED DECODE BENCH =====\n");
|
||||||
|
print_bench_header();
|
||||||
|
bench_config<128>(1, 32, 4, 512, 128);
|
||||||
|
bench_config<128>(1, 32, 4, 1024, 128);
|
||||||
|
bench_config<128>(1, 32, 4, 2048, 128);
|
||||||
|
bench_config<128>(1, 32, 4, 4096, 128);
|
||||||
|
bench_config<128>(16, 32, 4, 2048, 128);
|
||||||
|
bench_config<128>(32, 32, 4, 1024, 128);
|
||||||
|
}
|
||||||
|
|
||||||
|
int main() {
|
||||||
|
int n = sizeof(TESTS) / sizeof(TESTS[0]);
|
||||||
|
int fail = 0;
|
||||||
|
printf("=== Paged Decode vs CPU reference (%d cases) ===\n\n", n);
|
||||||
|
|
||||||
|
for (int i = 0; i < n; i++) {
|
||||||
|
fail += dispatch_test(TESTS[i]);
|
||||||
|
if (fail) break;
|
||||||
|
}
|
||||||
|
|
||||||
|
if (fail) {
|
||||||
|
printf("\nFAILED (%d/%d tests failed)\n", fail, n);
|
||||||
|
return fail;
|
||||||
|
}
|
||||||
|
printf("\nAll %d tests passed!\n", n);
|
||||||
|
bench();
|
||||||
|
return 0;
|
||||||
|
}
|
||||||
@@ -0,0 +1,169 @@
|
|||||||
|
/*
|
||||||
|
Pure-C test — uses shared dispatcher.
|
||||||
|
nvcc -I csrc -arch=sm_89 -O3 \
|
||||||
|
--use_fast_math --ptxas-options=-O3 --extra-device-vectorization \
|
||||||
|
csrc/tests/attn_prefill_test.cu -o test && ./test
|
||||||
|
*/
|
||||||
|
|
||||||
|
#include "test_utils.cuh"
|
||||||
|
#include "../kernels/attn_dispatchers.cuh"
|
||||||
|
|
||||||
|
// Warmed-up, CUDA-event timed throughput sweep over the production MMA path.
|
||||||
|
static void bench() {
|
||||||
|
const int cfgs[][7] = {
|
||||||
|
{1,32,4,512,512,128,0},
|
||||||
|
{1,32,4,1024,1024,128,0},
|
||||||
|
{1,32,4,2048,2048,128,0},
|
||||||
|
{1,32,4,2048,2048,128,1},
|
||||||
|
{4,32,4,2048,2048,128,1},
|
||||||
|
{1,32,4,4096,4096,128,1},
|
||||||
|
};
|
||||||
|
int n = sizeof(cfgs)/sizeof(cfgs[0]);
|
||||||
|
const int WARMUP = 10, ITERS = 50;
|
||||||
|
printf("\n===== PREFILL BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS);
|
||||||
|
printf("%-46s | %10s | %10s | %10s\n",
|
||||||
|
"config", "latency", "bandwidth", "throughput");
|
||||||
|
printf("---------------------------------------------------------------"
|
||||||
|
"----------------------------\n");
|
||||||
|
|
||||||
|
for (int ci = 0; ci < n; ci++) {
|
||||||
|
int B=cfgs[ci][0], Hq=cfgs[ci][1], Hk=cfgs[ci][2];
|
||||||
|
int ql=cfgs[ci][3], kl=cfgs[ci][4], D=cfgs[ci][5], causal=cfgs[ci][6];
|
||||||
|
size_t nQ=(size_t)B*Hq*ql*D, nKV=(size_t)B*Hk*kl*D;
|
||||||
|
|
||||||
|
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
||||||
|
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
||||||
|
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
||||||
|
size_t big = nQ>nKV?nQ:nKV; tmp=new bf16[big];
|
||||||
|
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(randf());
|
||||||
|
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(randf());
|
||||||
|
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(randf());
|
||||||
|
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
AttentionParams<bf16> p;
|
||||||
|
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D;
|
||||||
|
p.use_mask=0; p.causal_offset=causal?0:-1;
|
||||||
|
set_default_strides(p);
|
||||||
|
p.scale=1.0f/sqrtf((float)D);
|
||||||
|
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||||
|
|
||||||
|
auto launch = [&]() { dispatch_by_head_dim(D, [&]<int H>() { dispatch_prefill<H>(p); }); };
|
||||||
|
for (int i=0;i<WARMUP;i++) launch();
|
||||||
|
cudaDeviceSynchronize();
|
||||||
|
cudaError_t err=cudaGetLastError();
|
||||||
|
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return;}
|
||||||
|
|
||||||
|
cudaEvent_t s,e; cudaEventCreate(&s); cudaEventCreate(&e);
|
||||||
|
cudaEventRecord(s);
|
||||||
|
for (int i=0;i<ITERS;i++) launch();
|
||||||
|
cudaEventRecord(e); cudaEventSynchronize(e);
|
||||||
|
float ms=0; cudaEventElapsedTime(&ms,s,e); ms/=ITERS;
|
||||||
|
|
||||||
|
double flops = 4.0*B*Hq*(double)ql*kl*D;
|
||||||
|
if (causal) flops *= 0.5;
|
||||||
|
double tflops = flops/(ms*1e-3)/1e12;
|
||||||
|
double bytes = 2.0 * (2.0*nQ + 2.0*nKV);
|
||||||
|
double gbps = bytes/(ms*1e-3)/1e9;
|
||||||
|
|
||||||
|
char cfg[64];
|
||||||
|
snprintf(cfg, sizeof(cfg),
|
||||||
|
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d causal=%d",
|
||||||
|
B,Hq,Hk,ql,kl,D,causal);
|
||||||
|
printf("%-46s | %7.4f ms | %7.1f GB/s | %6.2f TFLOP/s\n",
|
||||||
|
cfg, ms, gbps, tflops);
|
||||||
|
|
||||||
|
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
|
||||||
|
delete[]tmp; cudaEventDestroy(s); cudaEventDestroy(e);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
static int run_test(int B, int Hq, int Hk, int ql, int kl, int D, int causal) {
|
||||||
|
printf("=== B=%d Hq=%d Hk=%d q=%d kv=%d D=%d causal=%d ===\n",
|
||||||
|
B,Hq,Hk,ql,kl,D,causal);
|
||||||
|
|
||||||
|
size_t nQ = B*Hq*ql*D, nKV = B*Hk*kl*D;
|
||||||
|
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
|
||||||
|
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
|
||||||
|
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
|
||||||
|
|
||||||
|
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
||||||
|
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
||||||
|
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
||||||
|
tmp=new bf16[max(nQ,nKV)];
|
||||||
|
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
|
||||||
|
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
|
||||||
|
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||||
|
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
|
||||||
|
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
|
AttentionParams<bf16> p;
|
||||||
|
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D;
|
||||||
|
p.use_mask=0; p.causal_offset=causal?0:-1;
|
||||||
|
set_default_strides(p);
|
||||||
|
p.scale=1.0f/sqrtf((float)D);
|
||||||
|
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||||
|
|
||||||
|
double t0=now_ms();
|
||||||
|
dispatch_by_head_dim(D, [&]<int H>() { dispatch_prefill<H>(p); });
|
||||||
|
cudaDeviceSynchronize();
|
||||||
|
double kms=now_ms()-t0;
|
||||||
|
cudaError_t err=cudaGetLastError();
|
||||||
|
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
|
||||||
|
|
||||||
|
bf16* hOut=new bf16[nQ];
|
||||||
|
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
|
||||||
|
|
||||||
|
float* ref=new float[nQ];
|
||||||
|
cpu_attention_ref(hQ, hK, hV, nullptr, ref, B, Hq, Hk, ql, kl, D, causal ? 0 : -1);
|
||||||
|
|
||||||
|
float max_abs_err=0, max_rel_err=0;
|
||||||
|
for (size_t i=0;i<nQ;i++) {
|
||||||
|
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
||||||
|
if(err>max_abs_err) max_abs_err=err;
|
||||||
|
float rel=err/fmaxf(fabsf(ref[i]), 1e-8f);
|
||||||
|
if(rel>max_rel_err) max_rel_err=rel;
|
||||||
|
}
|
||||||
|
const float atol=0.01f, rtol=0.01f;
|
||||||
|
bool pass=true;
|
||||||
|
for (size_t i=0;i<nQ;i++) {
|
||||||
|
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
||||||
|
if (err > atol + rtol * fabsf(ref[i])) { pass=false; break; }
|
||||||
|
}
|
||||||
|
printf("kernel: %.3f ms max_abs_err: %.6e max_rel_err: %.6e %s\n\n",
|
||||||
|
kms, max_abs_err, max_rel_err, pass?"PASS":"FAIL");
|
||||||
|
|
||||||
|
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
|
||||||
|
delete[]hQ;delete[]hK;delete[]hV;delete[]hOut;delete[]ref;delete[]tmp;
|
||||||
|
|
||||||
|
return pass ? 0 : 1;
|
||||||
|
}
|
||||||
|
|
||||||
|
int main() {
|
||||||
|
const int configs[][7] = {
|
||||||
|
{1,2,1,64,128,64,0}, // tiny: B,Hq,Hk,q,kv,D,causal
|
||||||
|
{1,32,4,512,512,128,0}, // standard
|
||||||
|
{1,32,4,128,256,128,0}, // medium
|
||||||
|
{1,4,2,256,256,128,1}, // causal
|
||||||
|
};
|
||||||
|
int n_configs = sizeof(configs) / sizeof(configs[0]);
|
||||||
|
int fail = 0;
|
||||||
|
|
||||||
|
for (int ci = 0; ci < n_configs; ci++) {
|
||||||
|
int B=configs[ci][0], Hq=configs[ci][1], Hk=configs[ci][2];
|
||||||
|
int ql=configs[ci][3], kl=configs[ci][4], D=configs[ci][5];
|
||||||
|
int causal=configs[ci][6];
|
||||||
|
fail += run_test(B, Hq, Hk, ql, kl, D, causal);
|
||||||
|
if (fail) break;
|
||||||
|
}
|
||||||
|
|
||||||
|
if (fail) {
|
||||||
|
printf("FAILED\n");
|
||||||
|
return fail;
|
||||||
|
}
|
||||||
|
printf("All tests passed!\n");
|
||||||
|
bench();
|
||||||
|
return 0;
|
||||||
|
}
|
||||||
@@ -1,130 +0,0 @@
|
|||||||
/*
|
|
||||||
Pure-C test:
|
|
||||||
nvcc -I csrc -arch=sm_89 -O3 \
|
|
||||||
--use_fast_math --ptxas-options=-O3 --extra-device-vectorization \
|
|
||||||
csrc/tests/gqa_decode_test.cu -o test && ./test
|
|
||||||
*/
|
|
||||||
|
|
||||||
#include <cstdio>
|
|
||||||
#include <cstdlib>
|
|
||||||
#include <cmath>
|
|
||||||
#include <sys/time.h>
|
|
||||||
#include "../kernels/gqa_decode_attn.cuh"
|
|
||||||
|
|
||||||
static double now_ms() {
|
|
||||||
struct timeval tv;
|
|
||||||
gettimeofday(&tv, NULL);
|
|
||||||
return tv.tv_sec * 1000.0 + tv.tv_usec / 1000.0;
|
|
||||||
}
|
|
||||||
|
|
||||||
static void cpu_decode(const float* Q, const float* K, const float* V,
|
|
||||||
const bool* mask, float* O,
|
|
||||||
int B, int Hq, int Hk, int seq_len, int D) {
|
|
||||||
float scale = 1.0f / sqrtf((float)D);
|
|
||||||
int n_rep = Hq / Hk;
|
|
||||||
for (int b = 0; b < B; b++) {
|
|
||||||
for (int h = 0; h < Hq; h++) {
|
|
||||||
int kv_h = h / n_rep;
|
|
||||||
float mv = -INFINITY, sv = 0.0f;
|
|
||||||
float accum[256] = {0};
|
|
||||||
for (int s = 0; s < seq_len; s++) {
|
|
||||||
if (!mask[b * seq_len + s]) continue;
|
|
||||||
float dot = 0.0f;
|
|
||||||
for (int d = 0; d < D; d++)
|
|
||||||
dot += Q[((b * Hq + h) * 1 + 0) * D + d]
|
|
||||||
* K[((b * Hk + kv_h) * seq_len + s) * D + d];
|
|
||||||
dot *= scale;
|
|
||||||
float nm = fmaxf(mv, dot);
|
|
||||||
float al = expf(mv - nm);
|
|
||||||
float be = expf(dot - nm);
|
|
||||||
sv = sv * al + be;
|
|
||||||
for (int d = 0; d < D; d++)
|
|
||||||
accum[d] = accum[d] * al
|
|
||||||
+ V[((b * Hk + kv_h) * seq_len + s) * D + d] * be;
|
|
||||||
mv = nm;
|
|
||||||
}
|
|
||||||
float inv = 1.0f / sv;
|
|
||||||
for (int d = 0; d < D; d++)
|
|
||||||
O[((b * Hq + h) * 1 + 0) * D + d] = accum[d] * inv;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
static bf16 f2bf(float x) { return __float2bfloat16(x); }
|
|
||||||
static float bf2f(bf16 x) { return __bfloat162float(x); }
|
|
||||||
static float randf() { return (float)rand() / (float)RAND_MAX - 0.5f; }
|
|
||||||
|
|
||||||
int main() {
|
|
||||||
const int configs[][5] = {
|
|
||||||
{1, 2, 1, 64, 32}, // B,Hq,Hk,seq_len,D
|
|
||||||
{1, 32, 4, 512, 128},
|
|
||||||
{1, 32, 4, 1024, 128},
|
|
||||||
};
|
|
||||||
int n_cfgs = sizeof(configs) / sizeof(configs[0]);
|
|
||||||
|
|
||||||
for (int ci = 0; ci < n_cfgs; ci++) {
|
|
||||||
int B = configs[ci][0], Hq = configs[ci][1], Hk = configs[ci][2];
|
|
||||||
int sl = configs[ci][3], D = configs[ci][4], gs = Hq / Hk;
|
|
||||||
printf("=== B=%d Hq=%d Hk=%d seq=%d D=%d gs=%d ===\n", B,Hq,Hk,sl,D,gs);
|
|
||||||
|
|
||||||
size_t nQ = B*Hq*1*D, nKV = B*Hk*sl*D;
|
|
||||||
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
|
|
||||||
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
|
|
||||||
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
|
|
||||||
|
|
||||||
bool* hMask=new bool[B*sl];
|
|
||||||
for (int i=0;i<B*sl;i++) hMask[i]=true;
|
|
||||||
|
|
||||||
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
|
||||||
bool* dMask;
|
|
||||||
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
|
||||||
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
|
||||||
cudaMalloc(&dMask,B*sl);
|
|
||||||
|
|
||||||
tmp=new bf16[max(nQ,nKV)];
|
|
||||||
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
|
|
||||||
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
|
||||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
|
|
||||||
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
|
||||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
|
|
||||||
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
|
||||||
cudaMemcpy(dMask,hMask,B*sl,cudaMemcpyHostToDevice);
|
|
||||||
|
|
||||||
GQAParams p;
|
|
||||||
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=1; p.kv_len=sl; p.head_dim=D;
|
|
||||||
p.use_mask=1; p.is_causal=0; p.causal_offset=0;
|
|
||||||
p.scale=1.0f/sqrtf((float)D);
|
|
||||||
p.q=dQ; p.k=dK; p.v=dV; p.mask=dMask; p.o=dO;
|
|
||||||
|
|
||||||
size_t smem=DC_CHUNK*D*sizeof(bf16);
|
|
||||||
dim3 block(32, gs);
|
|
||||||
dim3 grid(B*Hk);
|
|
||||||
printf("grid=(%d,1,1) block=(%d,%d,1) smem=%zu\n",
|
|
||||||
grid.x, block.x, block.y, smem);
|
|
||||||
|
|
||||||
double t0=now_ms();
|
|
||||||
gqa_decode_attn_kernel<<<grid,block,smem>>>(p);
|
|
||||||
cudaDeviceSynchronize();
|
|
||||||
double kms=now_ms()-t0;
|
|
||||||
cudaError_t err=cudaGetLastError();
|
|
||||||
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
|
|
||||||
|
|
||||||
bf16* hOut=new bf16[nQ];
|
|
||||||
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
|
|
||||||
|
|
||||||
float* ref=new float[nQ];
|
|
||||||
cpu_decode(hQ,hK,hV,hMask,ref,B,Hq,Hk,sl,D);
|
|
||||||
|
|
||||||
float max_err=0;
|
|
||||||
for (size_t i=0;i<nQ;i++){
|
|
||||||
float d=fabsf(bf2f(hOut[i])-ref[i]);
|
|
||||||
if(d>max_err) max_err=d;
|
|
||||||
}
|
|
||||||
printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err);
|
|
||||||
|
|
||||||
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);cudaFree(dMask);
|
|
||||||
delete[]hQ;delete[]hK;delete[]hV;delete[]hMask;delete[]hOut;delete[]ref;delete[]tmp;
|
|
||||||
}
|
|
||||||
printf("All tests passed!\n");
|
|
||||||
return 0;
|
|
||||||
}
|
|
||||||
@@ -1,133 +0,0 @@
|
|||||||
/*
|
|
||||||
Pure-C test:
|
|
||||||
nvcc -I csrc -arch=sm_89 -O3 \
|
|
||||||
--use_fast_math --ptxas-options=-O3 --extra-device-vectorization \
|
|
||||||
csrc/tests/gqa_prefill_test.cu -o test && ./test
|
|
||||||
*/
|
|
||||||
|
|
||||||
#include <cstdio>
|
|
||||||
#include <cstdlib>
|
|
||||||
#include <cmath>
|
|
||||||
#include <sys/time.h>
|
|
||||||
#include "../kernels/gqa_prefill_attn.cuh"
|
|
||||||
|
|
||||||
static double now_ms() {
|
|
||||||
struct timeval tv;
|
|
||||||
gettimeofday(&tv, NULL);
|
|
||||||
return tv.tv_sec * 1000.0 + tv.tv_usec / 1000.0;
|
|
||||||
}
|
|
||||||
|
|
||||||
static void cpu_attention(const float* Q, const float* K, const float* V, float* O,
|
|
||||||
int B, int Hq, int Hk, int q_len, int kv_len, int D,
|
|
||||||
int is_causal, int causal_off) {
|
|
||||||
float scale = 1.0f / sqrtf((float)D);
|
|
||||||
int n_rep = Hq / Hk;
|
|
||||||
for (int b = 0; b < B; b++) {
|
|
||||||
for (int h = 0; h < Hq; h++) {
|
|
||||||
for (int qi = 0; qi < q_len; qi++) {
|
|
||||||
int kv_h = h / n_rep;
|
|
||||||
float mv = -INFINITY, sv = 0.0f;
|
|
||||||
float accum[256] = {0};
|
|
||||||
int lim = is_causal ? min(kv_len, qi + causal_off + 1) : kv_len;
|
|
||||||
for (int kj = 0; kj < lim; kj++) {
|
|
||||||
float dot = 0.0f;
|
|
||||||
for (int d = 0; d < D; d++)
|
|
||||||
dot += Q[((b*Hq + h)*q_len + qi)*D + d]
|
|
||||||
* K[((b*Hk + kv_h)*kv_len + kj)*D + d];
|
|
||||||
dot *= scale;
|
|
||||||
float nm = fmaxf(mv, dot);
|
|
||||||
float al = expf(mv - nm);
|
|
||||||
float be = expf(dot - nm);
|
|
||||||
sv = sv * al + be;
|
|
||||||
for (int d = 0; d < D; d++)
|
|
||||||
accum[d] = accum[d] * al
|
|
||||||
+ V[((b*Hk + kv_h)*kv_len + kj)*D + d] * be;
|
|
||||||
mv = nm;
|
|
||||||
}
|
|
||||||
float inv = 1.0f / sv;
|
|
||||||
for (int d = 0; d < D; d++)
|
|
||||||
O[((b*Hq + h)*q_len + qi)*D + d] = accum[d] * inv;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
static __nv_bfloat16 f2bf(float x) { return __float2bfloat16(x); }
|
|
||||||
static float bf2f(__nv_bfloat16 x) { return __bfloat162float(x); }
|
|
||||||
static float randf() { return (float)rand() / (float)RAND_MAX - 0.5f; }
|
|
||||||
|
|
||||||
int main() {
|
|
||||||
const int configs[][7] = {
|
|
||||||
{1,2,1,64,128,64,0}, // tiny: B,Hq,Hk,q,kv,D,causal
|
|
||||||
{1,32,4,512,512,128,0}, // standard
|
|
||||||
{1,32,4,128,256,128,0}, // medium
|
|
||||||
{1,4,2,256,256,128,1}, // causal
|
|
||||||
};
|
|
||||||
int n_configs = sizeof(configs) / sizeof(configs[0]);
|
|
||||||
|
|
||||||
for (int ci = 0; ci < n_configs; ci++) {
|
|
||||||
int B=configs[ci][0], Hq=configs[ci][1], Hk=configs[ci][2];
|
|
||||||
int ql=configs[ci][3], kl=configs[ci][4], D=configs[ci][5];
|
|
||||||
int causal=configs[ci][6];
|
|
||||||
printf("=== B=%d Hq=%d Hk=%d q=%d kv=%d D=%d causal=%d ===\n",
|
|
||||||
B,Hq,Hk,ql,kl,D,causal);
|
|
||||||
|
|
||||||
size_t nQ = B*Hq*ql*D, nKV = B*Hk*kl*D;
|
|
||||||
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
|
|
||||||
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
|
|
||||||
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
|
|
||||||
|
|
||||||
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
|
||||||
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
|
||||||
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
|
||||||
tmp=new bf16[max(nQ,nKV)];
|
|
||||||
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
|
|
||||||
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
|
||||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
|
|
||||||
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
|
||||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
|
|
||||||
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
|
||||||
|
|
||||||
GQAParams p;
|
|
||||||
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D;
|
|
||||||
p.use_mask=0; p.is_causal=causal; p.causal_offset=0;
|
|
||||||
p.scale=1.0f/sqrtf((float)D);
|
|
||||||
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
|
||||||
|
|
||||||
constexpr int G=8, ROWS=32, P_BC=32;
|
|
||||||
dim3 grid((ql+ROWS-1)/ROWS, Hq, B);
|
|
||||||
dim3 block(G, ROWS, 1);
|
|
||||||
size_t smem=2*P_BC*D*sizeof(bf16);
|
|
||||||
printf("grid=(%d,%d,%d) block=(%d,%d,%d) smem=%zu\n",
|
|
||||||
grid.x,grid.y,grid.z, block.x,block.y,block.z, smem);
|
|
||||||
|
|
||||||
double t0=now_ms();
|
|
||||||
switch (D) {
|
|
||||||
case 64: gqa_prefill_attn_kernel_t<64, G,ROWS,P_BC><<<grid,block,smem>>>(p); break;
|
|
||||||
case 128: gqa_prefill_attn_kernel_t<128,G,ROWS,P_BC><<<grid,block,smem>>>(p); break;
|
|
||||||
default: printf("unsupported D=%d\n",D); return 1;
|
|
||||||
}
|
|
||||||
cudaDeviceSynchronize();
|
|
||||||
double kms=now_ms()-t0;
|
|
||||||
cudaError_t err=cudaGetLastError();
|
|
||||||
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
|
|
||||||
|
|
||||||
bf16* hOut=new bf16[nQ];
|
|
||||||
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
|
|
||||||
|
|
||||||
float* ref=new float[nQ];
|
|
||||||
cpu_attention(hQ,hK,hV,ref,B,Hq,Hk,ql,kl,D,causal,0);
|
|
||||||
|
|
||||||
float max_err=0;
|
|
||||||
for (size_t i=0;i<nQ;i++) {
|
|
||||||
float d=fabsf(bf2f(hOut[i])-ref[i]);
|
|
||||||
if(d>max_err) max_err=d;
|
|
||||||
}
|
|
||||||
printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err);
|
|
||||||
|
|
||||||
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
|
|
||||||
delete[]hQ;delete[]hK;delete[]hV;delete[]hOut;delete[]ref;delete[]tmp;
|
|
||||||
}
|
|
||||||
printf("All tests passed!\n");
|
|
||||||
return 0;
|
|
||||||
}
|
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user