Compare commits
158
Commits
v1.3.11
...
6ac3b51496
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
6ac3b51496 | ||
|
|
3406157431 | ||
|
|
0dd9a417b7 | ||
|
|
a01c1fd427 | ||
|
|
f8d9ab344d | ||
|
|
3fb4b8ab13 | ||
|
|
b5afe3d7a4 | ||
|
|
69f35c46e0 | ||
|
|
0378e62e17 | ||
|
|
5244f1a8fc | ||
|
|
5104638447 | ||
|
|
a711d9f478 | ||
|
|
15862d4b56 | ||
|
|
f9efb705b8 | ||
|
|
c6a82a5029 | ||
|
|
a5b238dd86 | ||
|
|
da6d94492d | ||
|
|
71b6e3aaaf | ||
|
|
f95722a277 | ||
|
|
9f48cb8928 | ||
|
|
9b58fef222 | ||
|
|
c5fba9c238 | ||
|
|
cd31f1f62f | ||
|
|
a5a3cc1fc2 | ||
|
|
d565d44c43 | ||
|
|
596c35fd71 | ||
|
|
47b3ed4e44 | ||
|
|
c1d05ae11d | ||
|
|
cf4f5ab9f6 | ||
|
|
3416f98c58 | ||
|
|
d28552f878 | ||
|
|
be90dfe2bd | ||
|
|
a33ca04f60 | ||
|
|
7f0e8bb8c2 | ||
|
|
0c1b7664c1 | ||
|
|
3fa7e66676 | ||
|
|
ca50fe4721 | ||
|
|
d9240ab149 | ||
|
|
d7cd69fef5 | ||
|
|
9bff61fb91 | ||
|
|
0b661bae85 | ||
|
|
ae9fd546ef | ||
|
|
e3ea850dc9 | ||
|
|
6e5088cc7d | ||
|
|
cbc584470d | ||
|
|
cb60713a72 | ||
|
|
c52a2487ae | ||
|
|
49aaa9a714 | ||
|
|
056c1382ff | ||
|
|
f163520fff | ||
|
|
1b1f1a0707 | ||
|
|
184fbbce5c | ||
|
|
02469887f5 | ||
|
|
05739629fc | ||
|
|
e0f7fa8e13 | ||
|
|
af25833fab | ||
|
|
6572be4f98 | ||
|
|
81788faef4 | ||
|
|
0e7fe57d96 | ||
|
|
55ee258e95 | ||
|
|
ef1bb6f401 | ||
|
|
6f49738991 | ||
|
|
a59ae8f32e | ||
|
|
6054b8dbd4 | ||
|
|
6f67ba8942 | ||
|
|
d0c5debbab | ||
|
|
4f2e03880b | ||
|
|
5c180cfa90 | ||
|
|
6f09b1d2ee | ||
|
|
b2230fefd8 | ||
|
|
654e6eb0d1 | ||
|
|
a317a4756b | ||
|
|
9b7e6c205f | ||
|
|
602b5ce216 | ||
|
|
8152760b5f | ||
|
|
8c052c99ee | ||
|
|
2667b8116d | ||
|
|
6dffb0305a | ||
|
|
49a9c6b3d2 | ||
|
|
cdf9145ecf | ||
|
|
85f0461b3b | ||
|
|
9f0e9195f7 | ||
|
|
88751d0b08 | ||
|
|
d0e5d910de | ||
|
|
a03504a280 | ||
|
|
d033b2ef0f | ||
|
|
8447f88f61 | ||
|
|
b1b65a657e | ||
|
|
3439e3104e | ||
|
|
288ba20db1 | ||
|
|
020e2eff4e | ||
|
|
1c7369f293 | ||
|
|
0fc1b1bd46 | ||
|
|
d7db37a70f | ||
|
|
6d98bb4f9f | ||
|
|
925cbedc93 | ||
|
|
fda82ee232 | ||
|
|
4b25664c79 | ||
|
|
a27c8a819d | ||
|
|
91acaf4b0b | ||
|
|
41dcf0feb9 | ||
|
|
9960f79920 | ||
|
|
7feeb0b93e | ||
|
|
3639b50b4a | ||
|
|
d855c09cf3 | ||
|
|
d6bfb09863 | ||
|
|
6db276f37a | ||
|
|
6c76c16480 | ||
|
|
11073bd1d2 | ||
|
|
25c9e81b2b | ||
|
|
ffbd9b57c9 | ||
|
|
04899a2b15 | ||
|
|
530d280e33 | ||
|
|
21ddead238 | ||
|
|
7aa5ed09d9 | ||
|
|
75411ce0cc | ||
|
|
9f83d982ec | ||
|
|
3e67b4f88d | ||
|
|
50cfd0d555 | ||
|
|
5756054d38 | ||
|
|
738cb8f128 | ||
|
|
28d1bd07cf | ||
|
|
02625739fe | ||
|
|
f688cd9c5a | ||
|
|
8055027df7 | ||
|
|
3067a8e1a6 | ||
|
|
97114b95a4 | ||
|
|
32fd03a025 | ||
|
|
21bf37dd83 | ||
|
|
5b67d5865a | ||
|
|
df979b4469 | ||
|
|
deb2d7e127 | ||
|
|
fc47319240 | ||
|
|
22cf798d81 | ||
|
|
164be9708b | ||
|
|
6a97524db4 | ||
|
|
c8b1e40f71 | ||
|
|
bcaa2d1ae0 | ||
|
|
8206afefd9 | ||
|
|
646b1b0f46 | ||
|
|
8150ab6c32 | ||
|
|
0b0693a0a2 | ||
|
|
115192c67c | ||
|
|
c2b04d8458 | ||
|
|
db487ab48b | ||
|
|
a95794d3db | ||
|
|
39f84f3b4c | ||
|
|
9f7cf50c56 | ||
|
|
d9a0c72149 | ||
|
|
5ab18bec48 | ||
|
|
2e29ed45d3 | ||
|
|
5ba21f4eb3 | ||
|
|
c26a47b0df | ||
|
|
b1a87b22bb | ||
|
|
07625057f2 | ||
|
|
53c804e233 | ||
|
|
05c7432964 | ||
|
|
4de42d83c2 |
+3
-1
@@ -4,6 +4,8 @@
|
||||
# Allow necessary files
|
||||
!astrai/
|
||||
!scripts/
|
||||
!assets/
|
||||
!docs/
|
||||
!csrc/
|
||||
!setup.py
|
||||
!pyproject.toml
|
||||
!README.md
|
||||
|
||||
@@ -26,30 +26,41 @@ jobs:
|
||||
if-no-files-found: error
|
||||
|
||||
build-cuda-linux:
|
||||
name: Build CUDA wheel (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 (CUDA 12.8)
|
||||
- name: Install torch (${{ matrix.cuda_tag }})
|
||||
run: |
|
||||
pip install torch --index-url https://download.pytorch.org/whl/cu128
|
||||
pip install torch --index-url https://download.pytorch.org/whl/${{ matrix.cuda_tag }}
|
||||
|
||||
- name: Setup CUDA
|
||||
- name: Setup CUDA (${{ matrix.cuda_ver }})
|
||||
uses: Jimver/cuda-toolkit@v0.2.35
|
||||
with:
|
||||
cuda: "12.8.0"
|
||||
cuda: "${{ matrix.cuda_ver }}"
|
||||
|
||||
- name: Build wheel (with CUDA kernels)
|
||||
run: |
|
||||
CSRC_KERNELS=true pip wheel . --no-deps --no-build-isolation -w dist/
|
||||
for f in dist/*.whl; do
|
||||
mv "$f" "dist/$(basename "$f" .whl)+${{ matrix.cuda_tag }}.whl"
|
||||
done
|
||||
|
||||
- uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: cuda-wheel-linux
|
||||
name: cuda-wheel-linux-${{ matrix.cuda_tag }}
|
||||
path: dist/*.whl
|
||||
if-no-files-found: error
|
||||
|
||||
@@ -66,10 +77,11 @@ jobs:
|
||||
name: pure-wheel
|
||||
path: release-assets/pure
|
||||
|
||||
- name: Download CUDA wheel
|
||||
- name: Download CUDA wheels (all variants)
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
name: cuda-wheel-linux
|
||||
pattern: cuda-wheel-linux-*
|
||||
merge-multiple: true
|
||||
path: release-assets/cuda
|
||||
|
||||
- name: Verify release assets
|
||||
@@ -79,8 +91,7 @@ jobs:
|
||||
pure_wheels=(release-assets/pure/*.whl)
|
||||
cuda_wheels=(release-assets/cuda/*.whl)
|
||||
test "${#pure_wheels[@]}" -eq 1
|
||||
test "${#cuda_wheels[@]}" -eq 1
|
||||
test "$(basename "${pure_wheels[0]}")" != "$(basename "${cuda_wheels[0]}")"
|
||||
test "${#cuda_wheels[@]}" -ge 1
|
||||
|
||||
- name: Create release & upload assets
|
||||
uses: softprops/action-gh-release@v2
|
||||
|
||||
+2
-1
@@ -9,6 +9,7 @@
|
||||
!scripts/**/*.py
|
||||
!tests/**/*.py
|
||||
!csrc/**/*.py
|
||||
!csrc/CMakeLists.txt
|
||||
|
||||
!csrc/**/*.cu
|
||||
!csrc/**/*.h
|
||||
@@ -24,7 +25,7 @@
|
||||
!/.dockerignore
|
||||
!/Dockerfile
|
||||
!/docker-compose.yml
|
||||
!/assets/**
|
||||
!/docs/**
|
||||
!/CONTRIBUTING.md
|
||||
!/LICENSE
|
||||
!/pyproject.toml
|
||||
|
||||
+10
-8
@@ -20,9 +20,6 @@ Run the following checks **in order** — CI will reject if any fail.
|
||||
ruff format .
|
||||
```
|
||||
|
||||
> **Note**: `ruff format` may rename parameters (e.g. `mask` → `attn_mask`).
|
||||
> Always review the diff after formatting.
|
||||
|
||||
### 2. Import sorting
|
||||
|
||||
```bash
|
||||
@@ -44,7 +41,7 @@ python -u -m pytest tests/ -v
|
||||
|
||||
> Failed tests may leave orphan tempdirs under `%TEMP%`. Clean them manually if needed.
|
||||
|
||||
### 4. (Optional) Full pre-commit check
|
||||
### 4. (Optional) Full pre-commit check script
|
||||
|
||||
If you have Git Bash available:
|
||||
|
||||
@@ -52,12 +49,17 @@ If you have Git Bash available:
|
||||
bash scripts/pre_commit.sh
|
||||
```
|
||||
|
||||
This runs format check, import sort check, and tests in one go.
|
||||
The script installs development dependencies by default, then runs the format
|
||||
check, import sort check, and tests. If dependencies are already installed, use:
|
||||
|
||||
```bash
|
||||
bash scripts/pre_commit.sh --skip-deps
|
||||
```
|
||||
|
||||
## Commit Style
|
||||
|
||||
```
|
||||
fix/feat/chore/docs/refactor/perf/test/style/ci/build/revert : short description (~50 chars)
|
||||
type: short description (~50 chars)
|
||||
|
||||
- bullet point body (each ~60 chars)
|
||||
```
|
||||
@@ -73,7 +75,7 @@ fix/feat/chore/docs/refactor/perf/test/style/ci/build/revert : short description
|
||||
|---------|-------|-----|
|
||||
| `ruff check --select I` fails | Wrong import order | `ruff check . --select I --fix .` then `ruff format .` |
|
||||
| `ruff format` changed many files | Not formatted before commit | Review diff carefully before staging |
|
||||
| Pre-commit hook rejects | Tests or lint failed | Fix individually, do not `--no-verify` |
|
||||
| Pre-commit check script fails | Dependency install, tests, or lint failed | Fix the failing step; use `--skip-deps` only when dependencies are already installed |
|
||||
| Tests fail with tempdir left | Test crash | Clean `%TEMP%` manually |
|
||||
|
||||
## Submitting Changes
|
||||
@@ -93,7 +95,7 @@ fix/feat/chore/docs/refactor/perf/test/style/ci/build/revert : short description
|
||||
|
||||
## License
|
||||
|
||||
By contributing, you agree that your contributions will be licensed under the [GPL-3.0 License](LICENSE).
|
||||
By contributing, you agree that your contributions will be licensed under the [Apache-2.0 License](LICENSE).
|
||||
|
||||
---
|
||||
|
||||
|
||||
+19
-4
@@ -1,8 +1,16 @@
|
||||
# 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
|
||||
FROM ubuntu:24.04 AS builder
|
||||
|
||||
ARG CUDA_TAG=cu128
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# 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 astrai/ ./astrai/
|
||||
COPY csrc/ ./csrc/
|
||||
COPY setup.py .
|
||||
COPY pyproject.toml .
|
||||
RUN pip install --no-cache-dir --upgrade pip \
|
||||
&& 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
|
||||
FROM ubuntu:24.04 AS production
|
||||
@@ -43,12 +53,17 @@ ENV PATH="/opt/venv/bin:$PATH"
|
||||
# Copy application code
|
||||
COPY astrai/ ./astrai/
|
||||
COPY scripts/ ./scripts/
|
||||
COPY assets/ ./assets/
|
||||
COPY docs/ ./docs/
|
||||
COPY pyproject.toml .
|
||||
COPY README.md .
|
||||
|
||||
# Create non-root user
|
||||
RUN useradd -m astrai && chown -R astrai:astrai /app
|
||||
# Create non-root user matching the host uid/gid (passed via build args)
|
||||
ARG USER_UID=1000
|
||||
ARG USER_GID=1000
|
||||
RUN groupadd -g "${USER_GID}" astrai \
|
||||
&& useradd -m -u "${USER_UID}" -g astrai astrai \
|
||||
&& chown -R astrai:astrai /app
|
||||
ENV HOME=/home/astrai
|
||||
USER astrai
|
||||
|
||||
ENV PYTHONUNBUFFERED=1 \
|
||||
|
||||
@@ -1,674 +1,201 @@
|
||||
GNU GENERAL PUBLIC LICENSE
|
||||
Version 3, 29 June 2007
|
||||
|
||||
Copyright (C) 2007 Free Software Foundation, Inc. <https://fsf.org/>
|
||||
Everyone is permitted to copy and distribute verbatim copies
|
||||
of this license document, but changing it is not allowed.
|
||||
|
||||
Preamble
|
||||
|
||||
The GNU General Public License is a free, copyleft license for
|
||||
software and other kinds of works.
|
||||
|
||||
The licenses for most software and other practical works are designed
|
||||
to take away your freedom to share and change the works. By contrast,
|
||||
the GNU General Public License is intended to guarantee your freedom to
|
||||
share and change all versions of a program--to make sure it remains free
|
||||
software for all its users. We, the Free Software Foundation, use the
|
||||
GNU General Public License for most of our software; it applies also to
|
||||
any other work released this way by its authors. You can apply it to
|
||||
your programs, too.
|
||||
|
||||
When we speak of free software, we are referring to freedom, not
|
||||
price. Our General Public Licenses are designed to make sure that you
|
||||
have the freedom to distribute copies of free software (and charge for
|
||||
them if you wish), that you receive source code or can get it if you
|
||||
want it, that you can change the software or use pieces of it in new
|
||||
free programs, and that you know you can do these things.
|
||||
|
||||
To protect your rights, we need to prevent others from denying you
|
||||
these rights or asking you to surrender the rights. Therefore, you have
|
||||
certain responsibilities if you distribute copies of the software, or if
|
||||
you modify it: responsibilities to respect the freedom of others.
|
||||
|
||||
For example, if you distribute copies of such a program, whether
|
||||
gratis or for a fee, you must pass on to the recipients the same
|
||||
freedoms that you received. You must make sure that they, too, receive
|
||||
or can get the source code. And you must show them these terms so they
|
||||
know their rights.
|
||||
|
||||
Developers that use the GNU GPL protect your rights with two steps:
|
||||
(1) assert copyright on the software, and (2) offer you this License
|
||||
giving you legal permission to copy, distribute and/or modify it.
|
||||
|
||||
For the developers' and authors' protection, the GPL clearly explains
|
||||
that there is no warranty for this free software. For both users' and
|
||||
authors' sake, the GPL requires that modified versions be marked as
|
||||
changed, so that their problems will not be attributed erroneously to
|
||||
authors of previous versions.
|
||||
|
||||
Some devices are designed to deny users access to install or run
|
||||
modified versions of the software inside them, although the manufacturer
|
||||
can do so. This is fundamentally incompatible with the aim of
|
||||
protecting users' freedom to change the software. The systematic
|
||||
pattern of such abuse occurs in the area of products for individuals to
|
||||
use, which is precisely where it is most unacceptable. Therefore, we
|
||||
have designed this version of the GPL to prohibit the practice for those
|
||||
products. If such problems arise substantially in other domains, we
|
||||
stand ready to extend this provision to those domains in future versions
|
||||
of the GPL, as needed to protect the freedom of users.
|
||||
|
||||
Finally, every program is threatened constantly by software patents.
|
||||
States should not allow patents to restrict development and use of
|
||||
software on general-purpose computers, but in those that do, we wish to
|
||||
avoid the special danger that patents applied to a free program could
|
||||
make it effectively proprietary. To prevent this, the GPL assures that
|
||||
patents cannot be used to render the program non-free.
|
||||
|
||||
The precise terms and conditions for copying, distribution and
|
||||
modification follow.
|
||||
|
||||
TERMS AND CONDITIONS
|
||||
|
||||
0. Definitions.
|
||||
|
||||
"This License" refers to version 3 of the GNU General Public License.
|
||||
|
||||
"Copyright" also means copyright-like laws that apply to other kinds of
|
||||
works, such as semiconductor masks.
|
||||
|
||||
"The Program" refers to any copyrightable work licensed under this
|
||||
License. Each licensee is addressed as "you". "Licensees" and
|
||||
"recipients" may be individuals or organizations.
|
||||
|
||||
To "modify" a work means to copy from or adapt all or part of the work
|
||||
in a fashion requiring copyright permission, other than the making of an
|
||||
exact copy. The resulting work is called a "modified version" of the
|
||||
earlier work or a work "based on" the earlier work.
|
||||
|
||||
A "covered work" means either the unmodified Program or a work based
|
||||
on the Program.
|
||||
|
||||
To "propagate" a work means to do anything with it that, without
|
||||
permission, would make you directly or secondarily liable for
|
||||
infringement under applicable copyright law, except executing it on a
|
||||
computer or modifying a private copy. Propagation includes copying,
|
||||
distribution (with or without modification), making available to the
|
||||
public, and in some countries other activities as well.
|
||||
|
||||
To "convey" a work means any kind of propagation that enables other
|
||||
parties to make or receive copies. Mere interaction with a user through
|
||||
a computer network, with no transfer of a copy, is not conveying.
|
||||
|
||||
An interactive user interface displays "Appropriate Legal Notices"
|
||||
to the extent that it includes a convenient and prominently visible
|
||||
feature that (1) displays an appropriate copyright notice, and (2)
|
||||
tells the user that there is no warranty for the work (except to the
|
||||
extent that warranties are provided), that licensees may convey the
|
||||
work under this License, and how to view a copy of this License. If
|
||||
the interface presents a list of user commands or options, such as a
|
||||
menu, a prominent item in the list meets this criterion.
|
||||
|
||||
1. Source Code.
|
||||
|
||||
The "source code" for a work means the preferred form of the work
|
||||
for making modifications to it. "Object code" means any non-source
|
||||
form of a work.
|
||||
|
||||
A "Standard Interface" means an interface that either is an official
|
||||
standard defined by a recognized standards body, or, in the case of
|
||||
interfaces specified for a particular programming language, one that
|
||||
is widely used among developers working in that language.
|
||||
|
||||
The "System Libraries" of an executable work include anything, other
|
||||
than the work as a whole, that (a) is included in the normal form of
|
||||
packaging a Major Component, but which is not part of that Major
|
||||
Component, and (b) serves only to enable use of the work with that
|
||||
Major Component, or to implement a Standard Interface for which an
|
||||
implementation is available to the public in source code form. A
|
||||
"Major Component", in this context, means a major essential component
|
||||
(kernel, window system, and so on) of the specific operating system
|
||||
(if any) on which the executable work runs, or a compiler used to
|
||||
produce the work, or an object code interpreter used to run it.
|
||||
|
||||
The "Corresponding Source" for a work in object code form means all
|
||||
the source code needed to generate, install, and (for an executable
|
||||
work) run the object code and to modify the work, including scripts to
|
||||
control those activities. However, it does not include the work's
|
||||
System Libraries, or general-purpose tools or generally available free
|
||||
programs which are used unmodified in performing those activities but
|
||||
which are not part of the work. For example, Corresponding Source
|
||||
includes interface definition files associated with source files for
|
||||
the work, and the source code for shared libraries and dynamically
|
||||
linked subprograms that the work is specifically designed to require,
|
||||
such as by intimate data communication or control flow between those
|
||||
subprograms and other parts of the work.
|
||||
|
||||
The Corresponding Source need not include anything that users
|
||||
can regenerate automatically from other parts of the Corresponding
|
||||
Source.
|
||||
|
||||
The Corresponding Source for a work in source code form is that
|
||||
same work.
|
||||
|
||||
2. Basic Permissions.
|
||||
|
||||
All rights granted under this License are granted for the term of
|
||||
copyright on the Program, and are irrevocable provided the stated
|
||||
conditions are met. This License explicitly affirms your unlimited
|
||||
permission to run the unmodified Program. The output from running a
|
||||
covered work is covered by this License only if the output, given its
|
||||
content, constitutes a covered work. This License acknowledges your
|
||||
rights of fair use or other equivalent, as provided by copyright law.
|
||||
|
||||
You may make, run and propagate covered works that you do not
|
||||
convey, without conditions so long as your license otherwise remains
|
||||
in force. You may convey covered works to others for the sole purpose
|
||||
of having them make modifications exclusively for you, or provide you
|
||||
with facilities for running those works, provided that you comply with
|
||||
the terms of this License in conveying all material for which you do
|
||||
not control copyright. Those thus making or running the covered works
|
||||
for you must do so exclusively on your behalf, under your direction
|
||||
and control, on terms that prohibit them from making any copies of
|
||||
your copyrighted material outside their relationship with you.
|
||||
|
||||
Conveying under any other circumstances is permitted solely under
|
||||
the conditions stated below. Sublicensing is not allowed; section 10
|
||||
makes it unnecessary.
|
||||
|
||||
3. Protecting Users' Legal Rights From Anti-Circumvention Law.
|
||||
|
||||
No covered work shall be deemed part of an effective technological
|
||||
measure under any applicable law fulfilling obligations under article
|
||||
11 of the WIPO copyright treaty adopted on 20 December 1996, or
|
||||
similar laws prohibiting or restricting circumvention of such
|
||||
measures.
|
||||
|
||||
When you convey a covered work, you waive any legal power to forbid
|
||||
circumvention of technological measures to the extent such circumvention
|
||||
is effected by exercising rights under this License with respect to
|
||||
the covered work, and you disclaim any intention to limit operation or
|
||||
modification of the work as a means of enforcing, against the work's
|
||||
users, your or third parties' legal rights to forbid circumvention of
|
||||
technological measures.
|
||||
|
||||
4. Conveying Verbatim Copies.
|
||||
|
||||
You may convey verbatim copies of the Program's source code as you
|
||||
receive it, in any medium, provided that you conspicuously and
|
||||
appropriately publish on each copy an appropriate copyright notice;
|
||||
keep intact all notices stating that this License and any
|
||||
non-permissive terms added in accord with section 7 apply to the code;
|
||||
keep intact all notices of the absence of any warranty; and give all
|
||||
recipients a copy of this License along with the Program.
|
||||
|
||||
You may charge any price or no price for each copy that you convey,
|
||||
and you may offer support or warranty protection for a fee.
|
||||
|
||||
5. Conveying Modified Source Versions.
|
||||
|
||||
You may convey a work based on the Program, or the modifications to
|
||||
produce it from the Program, in the form of source code under the
|
||||
terms of section 4, provided that you also meet all of these conditions:
|
||||
|
||||
a) The work must carry prominent notices stating that you modified
|
||||
it, and giving a relevant date.
|
||||
|
||||
b) The work must carry prominent notices stating that it is
|
||||
released under this License and any conditions added under section
|
||||
7. This requirement modifies the requirement in section 4 to
|
||||
"keep intact all notices".
|
||||
|
||||
c) You must license the entire work, as a whole, under this
|
||||
License to anyone who comes into possession of a copy. This
|
||||
License will therefore apply, along with any applicable section 7
|
||||
additional terms, to the whole of the work, and all its parts,
|
||||
regardless of how they are packaged. This License gives no
|
||||
permission to license the work in any other way, but it does not
|
||||
invalidate such permission if you have separately received it.
|
||||
|
||||
d) If the work has interactive user interfaces, each must display
|
||||
Appropriate Legal Notices; however, if the Program has interactive
|
||||
interfaces that do not display Appropriate Legal Notices, your
|
||||
work need not make them do so.
|
||||
|
||||
A compilation of a covered work with other separate and independent
|
||||
works, which are not by their nature extensions of the covered work,
|
||||
and which are not combined with it such as to form a larger program,
|
||||
in or on a volume of a storage or distribution medium, is called an
|
||||
"aggregate" if the compilation and its resulting copyright are not
|
||||
used to limit the access or legal rights of the compilation's users
|
||||
beyond what the individual works permit. Inclusion of a covered work
|
||||
in an aggregate does not cause this License to apply to the other
|
||||
parts of the aggregate.
|
||||
|
||||
6. Conveying Non-Source Forms.
|
||||
|
||||
You may convey a covered work in object code form under the terms
|
||||
of sections 4 and 5, provided that you also convey the
|
||||
machine-readable Corresponding Source under the terms of this License,
|
||||
in one of these ways:
|
||||
|
||||
a) Convey the object code in, or embodied in, a physical product
|
||||
(including a physical distribution medium), accompanied by the
|
||||
Corresponding Source fixed on a durable physical medium
|
||||
customarily used for software interchange.
|
||||
|
||||
b) Convey the object code in, or embodied in, a physical product
|
||||
(including a physical distribution medium), accompanied by a
|
||||
written offer, valid for at least three years and valid for as
|
||||
long as you offer spare parts or customer support for that product
|
||||
model, to give anyone who possesses the object code either (1) a
|
||||
copy of the Corresponding Source for all the software in the
|
||||
product that is covered by this License, on a durable physical
|
||||
medium customarily used for software interchange, for a price no
|
||||
more than your reasonable cost of physically performing this
|
||||
conveying of source, or (2) access to copy the
|
||||
Corresponding Source from a network server at no charge.
|
||||
|
||||
c) Convey individual copies of the object code with a copy of the
|
||||
written offer to provide the Corresponding Source. This
|
||||
alternative is allowed only occasionally and noncommercially, and
|
||||
only if you received the object code with such an offer, in accord
|
||||
with subsection 6b.
|
||||
|
||||
d) Convey the object code by offering access from a designated
|
||||
place (gratis or for a charge), and offer equivalent access to the
|
||||
Corresponding Source in the same way through the same place at no
|
||||
further charge. You need not require recipients to copy the
|
||||
Corresponding Source along with the object code. If the place to
|
||||
copy the object code is a network server, the Corresponding Source
|
||||
may be on a different server (operated by you or a third party)
|
||||
that supports equivalent copying facilities, provided you maintain
|
||||
clear directions next to the object code saying where to find the
|
||||
Corresponding Source. Regardless of what server hosts the
|
||||
Corresponding Source, you remain obligated to ensure that it is
|
||||
available for as long as needed to satisfy these requirements.
|
||||
|
||||
e) Convey the object code using peer-to-peer transmission, provided
|
||||
you inform other peers where the object code and Corresponding
|
||||
Source of the work are being offered to the general public at no
|
||||
charge under subsection 6d.
|
||||
|
||||
A separable portion of the object code, whose source code is excluded
|
||||
from the Corresponding Source as a System Library, need not be
|
||||
included in conveying the object code work.
|
||||
|
||||
A "User Product" is either (1) a "consumer product", which means any
|
||||
tangible personal property which is normally used for personal, family,
|
||||
or household purposes, or (2) anything designed or sold for incorporation
|
||||
into a dwelling. In determining whether a product is a consumer product,
|
||||
doubtful cases shall be resolved in favor of coverage. For a particular
|
||||
product received by a particular user, "normally used" refers to a
|
||||
typical or common use of that class of product, regardless of the status
|
||||
of the particular user or of the way in which the particular user
|
||||
actually uses, or expects or is expected to use, the product. A product
|
||||
is a consumer product regardless of whether the product has substantial
|
||||
commercial, industrial or non-consumer uses, unless such uses represent
|
||||
the only significant mode of use of the product.
|
||||
|
||||
"Installation Information" for a User Product means any methods,
|
||||
procedures, authorization keys, or other information required to install
|
||||
and execute modified versions of a covered work in that User Product from
|
||||
a modified version of its Corresponding Source. The information must
|
||||
suffice to ensure that the continued functioning of the modified object
|
||||
code is in no case prevented or interfered with solely because
|
||||
modification has been made.
|
||||
|
||||
If you convey an object code work under this section in, or with, or
|
||||
specifically for use in, a User Product, and the conveying occurs as
|
||||
part of a transaction in which the right of possession and use of the
|
||||
User Product is transferred to the recipient in perpetuity or for a
|
||||
fixed term (regardless of how the transaction is characterized), the
|
||||
Corresponding Source conveyed under this section must be accompanied
|
||||
by the Installation Information. But this requirement does not apply
|
||||
if neither you nor any third party retains the ability to install
|
||||
modified object code on the User Product (for example, the work has
|
||||
been installed in ROM).
|
||||
|
||||
The requirement to provide Installation Information does not include a
|
||||
requirement to continue to provide support service, warranty, or updates
|
||||
for a work that has been modified or installed by the recipient, or for
|
||||
the User Product in which it has been modified or installed. Access to a
|
||||
network may be denied when the modification itself materially and
|
||||
adversely affects the operation of the network or violates the rules and
|
||||
protocols for communication across the network.
|
||||
|
||||
Corresponding Source conveyed, and Installation Information provided,
|
||||
in accord with this section must be in a format that is publicly
|
||||
documented (and with an implementation available to the public in
|
||||
source code form), and must require no special password or key for
|
||||
unpacking, reading or copying.
|
||||
|
||||
7. Additional Terms.
|
||||
|
||||
"Additional permissions" are terms that supplement the terms of this
|
||||
License by making exceptions from one or more of its conditions.
|
||||
Additional permissions that are applicable to the entire Program shall
|
||||
be treated as though they were included in this License, to the extent
|
||||
that they are valid under applicable law. If additional permissions
|
||||
apply only to part of the Program, that part may be used separately
|
||||
under those permissions, but the entire Program remains governed by
|
||||
this License without regard to the additional permissions.
|
||||
|
||||
When you convey a copy of a covered work, you may at your option
|
||||
remove any additional permissions from that copy, or from any part of
|
||||
it. (Additional permissions may be written to require their own
|
||||
removal in certain cases when you modify the work.) You may place
|
||||
additional permissions on material, added by you to a covered work,
|
||||
for which you have or can give appropriate copyright permission.
|
||||
|
||||
Notwithstanding any other provision of this License, for material you
|
||||
add to a covered work, you may (if authorized by the copyright holders of
|
||||
that material) supplement the terms of this License with terms:
|
||||
|
||||
a) Disclaiming warranty or limiting liability differently from the
|
||||
terms of sections 15 and 16 of this License; or
|
||||
|
||||
b) Requiring preservation of specified reasonable legal notices or
|
||||
author attributions in that material or in the Appropriate Legal
|
||||
Notices displayed by works containing it; or
|
||||
|
||||
c) Prohibiting misrepresentation of the origin of that material, or
|
||||
requiring that modified versions of such material be marked in
|
||||
reasonable ways as different from the original version; or
|
||||
|
||||
d) Limiting the use for publicity purposes of names of licensors or
|
||||
authors of the material; or
|
||||
|
||||
e) Declining to grant rights under trademark law for use of some
|
||||
trade names, trademarks, or service marks; or
|
||||
|
||||
f) Requiring indemnification of licensors and authors of that
|
||||
material by anyone who conveys the material (or modified versions of
|
||||
it) with contractual assumptions of liability to the recipient, for
|
||||
any liability that these contractual assumptions directly impose on
|
||||
those licensors and authors.
|
||||
|
||||
All other non-permissive additional terms are considered "further
|
||||
restrictions" within the meaning of section 10. If the Program as you
|
||||
received it, or any part of it, contains a notice stating that it is
|
||||
governed by this License along with a term that is a further
|
||||
restriction, you may remove that term. If a license document contains
|
||||
a further restriction but permits relicensing or conveying under this
|
||||
License, you may add to a covered work material governed by the terms
|
||||
of that license document, provided that the further restriction does
|
||||
not survive such relicensing or conveying.
|
||||
|
||||
If you add terms to a covered work in accord with this section, you
|
||||
must place, in the relevant source files, a statement of the
|
||||
additional terms that apply to those files, or a notice indicating
|
||||
where to find the applicable terms.
|
||||
|
||||
Additional terms, permissive or non-permissive, may be stated in the
|
||||
form of a separately written license, or stated as exceptions;
|
||||
the above requirements apply either way.
|
||||
|
||||
8. Termination.
|
||||
|
||||
You may not propagate or modify a covered work except as expressly
|
||||
provided under this License. Any attempt otherwise to propagate or
|
||||
modify it is void, and will automatically terminate your rights under
|
||||
this License (including any patent licenses granted under the third
|
||||
paragraph of section 11).
|
||||
|
||||
However, if you cease all violation of this License, then your
|
||||
license from a particular copyright holder is reinstated (a)
|
||||
provisionally, unless and until the copyright holder explicitly and
|
||||
finally terminates your license, and (b) permanently, if the copyright
|
||||
holder fails to notify you of the violation by some reasonable means
|
||||
prior to 60 days after the cessation.
|
||||
|
||||
Moreover, your license from a particular copyright holder is
|
||||
reinstated permanently if the copyright holder notifies you of the
|
||||
violation by some reasonable means, this is the first time you have
|
||||
received notice of violation of this License (for any work) from that
|
||||
copyright holder, and you cure the violation prior to 30 days after
|
||||
your receipt of the notice.
|
||||
|
||||
Termination of your rights under this section does not terminate the
|
||||
licenses of parties who have received copies or rights from you under
|
||||
this License. If your rights have been terminated and not permanently
|
||||
reinstated, you do not qualify to receive new licenses for the same
|
||||
material under section 10.
|
||||
|
||||
9. Acceptance Not Required for Having Copies.
|
||||
|
||||
You are not required to accept this License in order to receive or
|
||||
run a copy of the Program. Ancillary propagation of a covered work
|
||||
occurring solely as a consequence of using peer-to-peer transmission
|
||||
to receive a copy likewise does not require acceptance. However,
|
||||
nothing other than this License grants you permission to propagate or
|
||||
modify any covered work. These actions infringe copyright if you do
|
||||
not accept this License. Therefore, by modifying or propagating a
|
||||
covered work, you indicate your acceptance of this License to do so.
|
||||
|
||||
10. Automatic Licensing of Downstream Recipients.
|
||||
|
||||
Each time you convey a covered work, the recipient automatically
|
||||
receives a license from the original licensors, to run, modify and
|
||||
propagate that work, subject to this License. You are not responsible
|
||||
for enforcing compliance by third parties with this License.
|
||||
|
||||
An "entity transaction" is a transaction transferring control of an
|
||||
organization, or substantially all assets of one, or subdividing an
|
||||
organization, or merging organizations. If propagation of a covered
|
||||
work results from an entity transaction, each party to that
|
||||
transaction who receives a copy of the work also receives whatever
|
||||
licenses to the work the party's predecessor in interest had or could
|
||||
give under the previous paragraph, plus a right to possession of the
|
||||
Corresponding Source of the work from the predecessor in interest, if
|
||||
the predecessor has it or can get it with reasonable efforts.
|
||||
|
||||
You may not impose any further restrictions on the exercise of the
|
||||
rights granted or affirmed under this License. For example, you may
|
||||
not impose a license fee, royalty, or other charge for exercise of
|
||||
rights granted under this License, and you may not initiate litigation
|
||||
(including a cross-claim or counterclaim in a lawsuit) alleging that
|
||||
any patent claim is infringed by making, using, selling, offering for
|
||||
sale, or importing the Program or any portion of it.
|
||||
|
||||
11. Patents.
|
||||
|
||||
A "contributor" is a copyright holder who authorizes use under this
|
||||
License of the Program or a work on which the Program is based. The
|
||||
work thus licensed is called the contributor's "contributor version".
|
||||
|
||||
A contributor's "essential patent claims" are all patent claims
|
||||
owned or controlled by the contributor, whether already acquired or
|
||||
hereafter acquired, that would be infringed by some manner, permitted
|
||||
by this License, of making, using, or selling its contributor version,
|
||||
but do not include claims that would be infringed only as a
|
||||
consequence of further modification of the contributor version. For
|
||||
purposes of this definition, "control" includes the right to grant
|
||||
patent sublicenses in a manner consistent with the requirements of
|
||||
this License.
|
||||
|
||||
Each contributor grants you a non-exclusive, worldwide, royalty-free
|
||||
patent license under the contributor's essential patent claims, to
|
||||
make, use, sell, offer for sale, import and otherwise run, modify and
|
||||
propagate the contents of its contributor version.
|
||||
|
||||
In the following three paragraphs, a "patent license" is any express
|
||||
agreement or commitment, however denominated, not to enforce a patent
|
||||
(such as an express permission to practice a patent or covenant not to
|
||||
sue for patent infringement). To "grant" such a patent license to a
|
||||
party means to make such an agreement or commitment not to enforce a
|
||||
patent against the party.
|
||||
|
||||
If you convey a covered work, knowingly relying on a patent license,
|
||||
and the Corresponding Source of the work is not available for anyone
|
||||
to copy, free of charge and under the terms of this License, through a
|
||||
publicly available network server or other readily accessible means,
|
||||
then you must either (1) cause the Corresponding Source to be so
|
||||
available, or (2) arrange to deprive yourself of the benefit of the
|
||||
patent license for this particular work, or (3) arrange, in a manner
|
||||
consistent with the requirements of this License, to extend the patent
|
||||
license to downstream recipients. "Knowingly relying" means you have
|
||||
actual knowledge that, but for the patent license, your conveying the
|
||||
covered work in a country, or your recipient's use of the covered work
|
||||
in a country, would infringe one or more identifiable patents in that
|
||||
country that you have reason to believe are valid.
|
||||
|
||||
If, pursuant to or in connection with a single transaction or
|
||||
arrangement, you convey, or propagate by procuring conveyance of, a
|
||||
covered work, and grant a patent license to some of the parties
|
||||
receiving the covered work authorizing them to use, propagate, modify
|
||||
or convey a specific copy of the covered work, then the patent license
|
||||
you grant is automatically extended to all recipients of the covered
|
||||
work and works based on it.
|
||||
|
||||
A patent license is "discriminatory" if it does not include within
|
||||
the scope of its coverage, prohibits the exercise of, or is
|
||||
conditioned on the non-exercise of one or more of the rights that are
|
||||
specifically granted under this License. You may not convey a covered
|
||||
work if you are a party to an arrangement with a third party that is
|
||||
in the business of distributing software, under which you make payment
|
||||
to the third party based on the extent of your activity of conveying
|
||||
the work, and under which the third party grants, to any of the
|
||||
parties who would receive the covered work from you, a discriminatory
|
||||
patent license (a) in connection with copies of the covered work
|
||||
conveyed by you (or copies made from those copies), or (b) primarily
|
||||
for and in connection with specific products or compilations that
|
||||
contain the covered work, unless you entered into that arrangement,
|
||||
or that patent license was granted, prior to 28 March 2007.
|
||||
|
||||
Nothing in this License shall be construed as excluding or limiting
|
||||
any implied license or other defenses to infringement that may
|
||||
otherwise be available to you under applicable patent law.
|
||||
|
||||
12. No Surrender of Others' Freedom.
|
||||
|
||||
If conditions are imposed on you (whether by court order, agreement or
|
||||
otherwise) that contradict the conditions of this License, they do not
|
||||
excuse you from the conditions of this License. If you cannot convey a
|
||||
covered work so as to satisfy simultaneously your obligations under this
|
||||
License and any other pertinent obligations, then as a consequence you may
|
||||
not convey it at all. For example, if you agree to terms that obligate you
|
||||
to collect a royalty for further conveying from those to whom you convey
|
||||
the Program, the only way you could satisfy both those terms and this
|
||||
License would be to refrain entirely from conveying the Program.
|
||||
|
||||
13. Use with the GNU Affero General Public License.
|
||||
|
||||
Notwithstanding any other provision of this License, you have
|
||||
permission to link or combine any covered work with a work licensed
|
||||
under version 3 of the GNU Affero General Public License into a single
|
||||
combined work, and to convey the resulting work. The terms of this
|
||||
License will continue to apply to the part which is the covered work,
|
||||
but the special requirements of the GNU Affero General Public License,
|
||||
section 13, concerning interaction through a network will apply to the
|
||||
combination as such.
|
||||
|
||||
14. Revised Versions of this License.
|
||||
|
||||
The Free Software Foundation may publish revised and/or new versions of
|
||||
the GNU General Public License from time to time. Such new versions will
|
||||
be similar in spirit to the present version, but may differ in detail to
|
||||
address new problems or concerns.
|
||||
|
||||
Each version is given a distinguishing version number. If the
|
||||
Program specifies that a certain numbered version of the GNU General
|
||||
Public License "or any later version" applies to it, you have the
|
||||
option of following the terms and conditions either of that numbered
|
||||
version or of any later version published by the Free Software
|
||||
Foundation. If the Program does not specify a version number of the
|
||||
GNU General Public License, you may choose any version ever published
|
||||
by the Free Software Foundation.
|
||||
|
||||
If the Program specifies that a proxy can decide which future
|
||||
versions of the GNU General Public License can be used, that proxy's
|
||||
public statement of acceptance of a version permanently authorizes you
|
||||
to choose that version for the Program.
|
||||
|
||||
Later license versions may give you additional or different
|
||||
permissions. However, no additional obligations are imposed on any
|
||||
author or copyright holder as a result of your choosing to follow a
|
||||
later version.
|
||||
|
||||
15. Disclaimer of Warranty.
|
||||
|
||||
THERE IS NO WARRANTY FOR THE PROGRAM, TO THE EXTENT PERMITTED BY
|
||||
APPLICABLE LAW. EXCEPT WHEN OTHERWISE STATED IN WRITING THE COPYRIGHT
|
||||
HOLDERS AND/OR OTHER PARTIES PROVIDE THE PROGRAM "AS IS" WITHOUT WARRANTY
|
||||
OF ANY KIND, EITHER EXPRESSED OR IMPLIED, INCLUDING, BUT NOT LIMITED TO,
|
||||
THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR
|
||||
PURPOSE. THE ENTIRE RISK AS TO THE QUALITY AND PERFORMANCE OF THE PROGRAM
|
||||
IS WITH YOU. SHOULD THE PROGRAM PROVE DEFECTIVE, YOU ASSUME THE COST OF
|
||||
ALL NECESSARY SERVICING, REPAIR OR CORRECTION.
|
||||
|
||||
16. Limitation of Liability.
|
||||
|
||||
IN NO EVENT UNLESS REQUIRED BY APPLICABLE LAW OR AGREED TO IN WRITING
|
||||
WILL ANY COPYRIGHT HOLDER, OR ANY OTHER PARTY WHO MODIFIES AND/OR CONVEYS
|
||||
THE PROGRAM AS PERMITTED ABOVE, BE LIABLE TO YOU FOR DAMAGES, INCLUDING ANY
|
||||
GENERAL, SPECIAL, INCIDENTAL OR CONSEQUENTIAL DAMAGES ARISING OUT OF THE
|
||||
USE OR INABILITY TO USE THE PROGRAM (INCLUDING BUT NOT LIMITED TO LOSS OF
|
||||
DATA OR DATA BEING RENDERED INACCURATE OR LOSSES SUSTAINED BY YOU OR THIRD
|
||||
PARTIES OR A FAILURE OF THE PROGRAM TO OPERATE WITH ANY OTHER PROGRAMS),
|
||||
EVEN IF SUCH HOLDER OR OTHER PARTY HAS BEEN ADVISED OF THE POSSIBILITY OF
|
||||
SUCH DAMAGES.
|
||||
|
||||
17. Interpretation of Sections 15 and 16.
|
||||
|
||||
If the disclaimer of warranty and limitation of liability provided
|
||||
above cannot be given local legal effect according to their terms,
|
||||
reviewing courts shall apply local law that most closely approximates
|
||||
an absolute waiver of all civil liability in connection with the
|
||||
Program, unless a warranty or assumption of liability accompanies a
|
||||
copy of the Program in return for a fee.
|
||||
|
||||
END OF TERMS AND CONDITIONS
|
||||
|
||||
How to Apply These Terms to Your New Programs
|
||||
|
||||
If you develop a new program, and you want it to be of the greatest
|
||||
possible use to the public, the best way to achieve this is to make it
|
||||
free software which everyone can redistribute and change under these terms.
|
||||
|
||||
To do so, attach the following notices to the program. It is safest
|
||||
to attach them to the start of each source file to most effectively
|
||||
state the exclusion of warranty; and each file should have at least
|
||||
the "copyright" line and a pointer to where the full notice is found.
|
||||
|
||||
<one line to give the program's name and a brief idea of what it does.>
|
||||
Copyright (C) <year> <name of author>
|
||||
|
||||
This program is free software: you can redistribute it and/or modify
|
||||
it under the terms of the GNU General Public License as published by
|
||||
the Free Software Foundation, either version 3 of the License, or
|
||||
(at your option) any later version.
|
||||
|
||||
This program is distributed in the hope that it will be useful,
|
||||
but WITHOUT ANY WARRANTY; without even the implied warranty of
|
||||
MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
|
||||
GNU General Public License for more details.
|
||||
|
||||
You should have received a copy of the GNU General Public License
|
||||
along with this program. If not, see <https://www.gnu.org/licenses/>.
|
||||
|
||||
Also add information on how to contact you by electronic and paper mail.
|
||||
|
||||
If the program does terminal interaction, make it output a short
|
||||
notice like this when it starts in an interactive mode:
|
||||
|
||||
<program> Copyright (C) <year> <name of author>
|
||||
This program comes with ABSOLUTELY NO WARRANTY; for details type `show w'.
|
||||
This is free software, and you are welcome to redistribute it
|
||||
under certain conditions; type `show c' for details.
|
||||
|
||||
The hypothetical commands `show w' and `show c' should show the appropriate
|
||||
parts of the General Public License. Of course, your program's commands
|
||||
might be different; for a GUI interface, you would use an "about box".
|
||||
|
||||
You should also get your employer (if you work as a programmer) or school,
|
||||
if any, to sign a "copyright disclaimer" for the program, if necessary.
|
||||
For more information on this, and how to apply and follow the GNU GPL, see
|
||||
<https://www.gnu.org/licenses/>.
|
||||
|
||||
The GNU General Public License does not permit incorporating your program
|
||||
into proprietary programs. If your program is a subroutine library, you
|
||||
may consider it more useful to permit linking proprietary applications with
|
||||
the library. If this is what you want to do, use the GNU Lesser General
|
||||
Public License instead of this License. But first, please read
|
||||
<https://www.gnu.org/licenses/why-not-lgpl.html>.
|
||||
Apache License
|
||||
Version 2.0, January 2004
|
||||
http://www.apache.org/licenses/
|
||||
|
||||
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
||||
|
||||
1. Definitions.
|
||||
|
||||
"License" shall mean the terms and conditions for use, reproduction,
|
||||
and distribution as defined by Sections 1 through 9 of this document.
|
||||
|
||||
"Licensor" shall mean the copyright owner or entity authorized by
|
||||
the copyright owner that is granting the License.
|
||||
|
||||
"Legal Entity" shall mean the union of the acting entity and all
|
||||
other entities that control, are controlled by, or are under common
|
||||
control with that entity. For the purposes of this definition,
|
||||
"control" means (i) the power, direct or indirect, to cause the
|
||||
direction or management of such entity, whether by contract or
|
||||
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
||||
outstanding shares, or (iii) beneficial ownership of such entity.
|
||||
|
||||
"You" (or "Your") shall mean an individual or Legal Entity
|
||||
exercising permissions granted by this License.
|
||||
|
||||
"Source" form shall mean the preferred form for making modifications,
|
||||
including but not limited to software source code, documentation
|
||||
source, and configuration files.
|
||||
|
||||
"Object" form shall mean any form resulting from mechanical
|
||||
transformation or translation of a Source form, including but
|
||||
not limited to compiled object code, generated documentation,
|
||||
and conversions to other media types.
|
||||
|
||||
"Work" shall mean the work of authorship, whether in Source or
|
||||
Object form, made available under the License, as indicated by a
|
||||
copyright notice that is included in or attached to the work
|
||||
(an example is provided in the Appendix below).
|
||||
|
||||
"Derivative Works" shall mean any work, whether in Source or Object
|
||||
form, that is based on (or derived from) the Work and for which the
|
||||
editorial revisions, annotations, elaborations, or other modifications
|
||||
represent, as a whole, an original work of authorship. For the purposes
|
||||
of this License, Derivative Works shall not include works that remain
|
||||
separable from, or merely link (or bind by name) to the interfaces of,
|
||||
the Work and Derivative Works thereof.
|
||||
|
||||
"Contribution" shall mean any work of authorship, including
|
||||
the original version of the Work and any modifications or additions
|
||||
to that Work or Derivative Works thereof, that is intentionally
|
||||
submitted to Licensor for inclusion in the Work by the copyright owner
|
||||
or by an individual or Legal Entity authorized to submit on behalf of
|
||||
the copyright owner. For the purposes of this definition, "submitted"
|
||||
means any form of electronic, verbal, or written communication sent
|
||||
to the Licensor or its representatives, including but not limited to
|
||||
communication on electronic mailing lists, source code control systems,
|
||||
and issue tracking systems that are managed by, or on behalf of, the
|
||||
Licensor for the purpose of discussing and improving the Work, but
|
||||
excluding communication that is conspicuously marked or otherwise
|
||||
designated in writing by the copyright owner as "Not a Contribution."
|
||||
|
||||
"Contributor" shall mean Licensor and any individual or Legal Entity
|
||||
on behalf of whom a Contribution has been received by Licensor and
|
||||
subsequently incorporated within the Work.
|
||||
|
||||
2. Grant of Copyright License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
copyright license to reproduce, prepare Derivative Works of,
|
||||
publicly display, publicly perform, sublicense, and distribute the
|
||||
Work and such Derivative Works in Source or Object form.
|
||||
|
||||
3. Grant of Patent License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
(except as stated in this section) patent license to make, have made,
|
||||
use, offer to sell, sell, import, and otherwise transfer the Work,
|
||||
where such license applies only to those patent claims licensable
|
||||
by such Contributor that are necessarily infringed by their
|
||||
Contribution(s) alone or by combination of their Contribution(s)
|
||||
with the Work to which such Contribution(s) was submitted. If You
|
||||
institute patent litigation against any entity (including a
|
||||
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
||||
or a Contribution incorporated within the Work constitutes direct
|
||||
or contributory patent infringement, then any patent licenses
|
||||
granted to You under this License for that Work shall terminate
|
||||
as of the date such litigation is filed.
|
||||
|
||||
4. Redistribution. You may reproduce and distribute copies of the
|
||||
Work or Derivative Works thereof in any medium, with or without
|
||||
modifications, and in Source or Object form, provided that You
|
||||
meet the following conditions:
|
||||
|
||||
(a) You must give any other recipients of the Work or
|
||||
Derivative Works a copy of this License; and
|
||||
|
||||
(b) You must cause any modified files to carry prominent notices
|
||||
stating that You changed the files; and
|
||||
|
||||
(c) You must retain, in the Source form of any Derivative Works
|
||||
that You distribute, all copyright, patent, trademark, and
|
||||
attribution notices from the Source form of the Work,
|
||||
excluding those notices that do not pertain to any part of
|
||||
the Derivative Works; and
|
||||
|
||||
(d) If the Work includes a "NOTICE" text file as part of its
|
||||
distribution, then any Derivative Works that You distribute must
|
||||
include a readable copy of the attribution notices contained
|
||||
within such NOTICE file, excluding those notices that do not
|
||||
pertain to any part of the Derivative Works, in at least one
|
||||
of the following places: within a NOTICE text file distributed
|
||||
as part of the Derivative Works; within the Source form or
|
||||
documentation, if provided along with the Derivative Works; or,
|
||||
within a display generated by the Derivative Works, if and
|
||||
wherever such third-party notices normally appear. The contents
|
||||
of the NOTICE file are for informational purposes only and
|
||||
do not modify the License. You may add Your own attribution
|
||||
notices within Derivative Works that You distribute, alongside
|
||||
or as an addendum to the NOTICE text from the Work, provided
|
||||
that such additional attribution notices cannot be construed
|
||||
as modifying the License.
|
||||
|
||||
You may add Your own copyright statement to Your modifications and
|
||||
may provide additional or different license terms and conditions
|
||||
for use, reproduction, or distribution of Your modifications, or
|
||||
for any such Derivative Works as a whole, provided Your use,
|
||||
reproduction, and distribution of the Work otherwise complies with
|
||||
the conditions stated in this License.
|
||||
|
||||
5. Submission of Contributions. Unless You explicitly state otherwise,
|
||||
any Contribution intentionally submitted for inclusion in the Work
|
||||
by You to the Licensor shall be under the terms and conditions of
|
||||
this License, without any additional terms or conditions.
|
||||
Notwithstanding the above, nothing herein shall supersede or modify
|
||||
the terms of any separate license agreement you may have executed
|
||||
with Licensor regarding such Contributions.
|
||||
|
||||
6. Trademarks. This License does not grant permission to use the trade
|
||||
names, trademarks, service marks, or product names of the Licensor,
|
||||
except as required for reasonable and customary use in describing the
|
||||
origin of the Work and reproducing the content of the NOTICE file.
|
||||
|
||||
7. Disclaimer of Warranty. Unless required by applicable law or
|
||||
agreed to in writing, Licensor provides the Work (and each
|
||||
Contributor provides its Contributions) on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
||||
implied, including, without limitation, any warranties or conditions
|
||||
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
||||
PARTICULAR PURPOSE. You are solely responsible for determining the
|
||||
appropriateness of using or redistributing the Work and assume any
|
||||
risks associated with Your exercise of permissions under this License.
|
||||
|
||||
8. Limitation of Liability. In no event and under no legal theory,
|
||||
whether in tort (including negligence), contract, or otherwise,
|
||||
unless required by applicable law (such as deliberate and grossly
|
||||
negligent acts) or agreed to in writing, shall any Contributor be
|
||||
liable to You for damages, including any direct, indirect, special,
|
||||
incidental, or consequential damages of any character arising as a
|
||||
result of this License or out of the use or inability to use the
|
||||
Work (including but not limited to damages for loss of goodwill,
|
||||
work stoppage, computer failure or malfunction, or any and all
|
||||
other commercial damages or losses), even if such Contributor
|
||||
has been advised of the possibility of such damages.
|
||||
|
||||
9. Accepting Warranty or Additional Liability. While redistributing
|
||||
the Work or Derivative Works thereof, You may choose to offer,
|
||||
and charge a fee for, acceptance of support, warranty, indemnity,
|
||||
or other liability obligations and/or rights consistent with this
|
||||
License. However, in accepting such obligations, You may act only
|
||||
on Your own behalf and on Your sole responsibility, not on behalf
|
||||
of any other Contributor, and only if You agree to indemnify,
|
||||
defend, and hold each Contributor harmless for any liability
|
||||
incurred by, or claims asserted against, such Contributor by reason
|
||||
of your accepting any such warranty or additional liability.
|
||||
|
||||
END OF TERMS AND CONDITIONS
|
||||
|
||||
APPENDIX: How to apply the Apache License to your work.
|
||||
|
||||
To apply the Apache License to your work, attach the following
|
||||
boilerplate notice, with the fields enclosed by brackets "[]"
|
||||
replaced with your own identifying information. (Don't include
|
||||
the brackets!) The text should be enclosed in the appropriate
|
||||
comment syntax for the file format. We also recommend that a
|
||||
file or class name and description of purpose be included on the
|
||||
same "printed page" as the copyright notice for easier
|
||||
identification within third-party archives.
|
||||
|
||||
Copyright [yyyy] [name of copyright owner]
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
<div align="center">
|
||||
|
||||
<img src="assets/images/logo.png" width="auto" alt="Logo">
|
||||
<img src="docs/images/logo.png" width="auto" alt="Logo">
|
||||
<p>
|
||||
<strong>A lightweight Transformer training & inference framework</strong>
|
||||
</p>
|
||||
@@ -8,7 +8,7 @@
|
||||
|
||||
<div align="center">
|
||||
<img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python">
|
||||
<img src="https://img.shields.io/badge/license-GPL--3.0-blue.svg" alt="license">
|
||||
<img src="https://img.shields.io/badge/license-Apache--2.0-blue.svg" alt="license">
|
||||
<img src="https://img.shields.io/github/v/tag/ViperEkura/AstrAI?label=Release&color=76bad9" alt="release">
|
||||
<img src="https://img.shields.io/github/stars/ViperEkura/AstrAI?style=flat&label=Stars&color=76bad9" alt="stars">
|
||||
<img src="https://img.shields.io/github/forks/ViperEkura/AstrAI?style=flat&label=Forks&color=76bad9" alt="forks">
|
||||
@@ -17,7 +17,7 @@
|
||||
|
||||
<div align="center">
|
||||
<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/discussions">Discussions</a> •
|
||||
<a href="https://huggingface.co/ViperEkura">HuggingFace</a>
|
||||
@@ -27,7 +27,7 @@
|
||||
|
||||
## 📖 Table of Contents
|
||||
|
||||
- [Features](#features)
|
||||
- [Overview](#overview)
|
||||
- [Getting Started](#getting-started)
|
||||
- [Demo](#demo)
|
||||
- [Documentation](#documentation)
|
||||
@@ -40,15 +40,19 @@
|
||||
<a id="english"></a>
|
||||
## English
|
||||
|
||||
### Features
|
||||
### Overview
|
||||
|
||||
- 🚀 **High Performance**: Optimized for both training and inference with efficient parallelization.
|
||||
- 🔧 **Flexible**: Support for seq/sft/dpo/grpo training, customizable model architectures.
|
||||
- 💡 **Easy to Use**: Simple API with comprehensive examples and demos.
|
||||
- 📦 **Lightweight**: Minimal dependencies, easy to deploy.
|
||||
- 🔬 **Research‑Friendly**: Modular design, easy to experiment with new ideas.
|
||||
- 🤗 **HuggingFace-Style API**: AutoModel/AutoTokenizer APIs inspired by HuggingFace for easy model and tokenizer loading.
|
||||
- 🔌 **Dual API Compatibility**: Supports both OpenAI and Anthropic chat completion APIs out of the box.
|
||||
AstrAI is an end-to-end Transformer framework for building, training, evaluating, and serving models. It provides a compact PyTorch codebase for the complete model lifecycle, from declarative data preprocessing and distributed training to continuous-batching inference and OpenAI/Anthropic-compatible APIs.
|
||||
|
||||
| Area | Capabilities |
|
||||
|---|---|
|
||||
| **Models** | Autoregressive language models and embedding models with GQA, MLA, MoE, RoPE, and extensible attention/FFN components |
|
||||
| **Training** | Pre-training (`seq`), supervised fine-tuning (`sft`), DPO, and GRPO with gradient accumulation, checkpointing, DDP, and FSDP |
|
||||
| **Data** | Declarative JSON preprocessing, configurable masking and packing, binary/JSONL storage, and streaming datasets |
|
||||
| **Inference** | Continuous batching, paged KV cache, radix prefix caching, streaming generation, and Torch/CUDA/FlashAttention backends |
|
||||
| **Serving** | FastAPI server with OpenAI and Anthropic chat completion protocols, including SSE streaming and tool calls |
|
||||
| **Evaluation** | Perplexity, MMLU, HumanEval, IFEval, IFD, and ROUGE evaluation tools |
|
||||
| **Extensibility** | Factory and registry architecture for models, datasets, training strategies, callbacks, kernels, and protocol components |
|
||||
|
||||
### Getting Started
|
||||
|
||||
@@ -56,6 +60,8 @@ End-to-end walkthrough in 5 steps:
|
||||
|
||||
**1. Install**
|
||||
|
||||
AstrAI requires Python 3.12+ and pins PyTorch exactly to `2.11.0`. Training, `scripts/tools/generate.py`, generation evaluations, and the generation demos require CUDA; CPU support is limited to components with an explicit CPU device path, such as the HTTP server and direct-scoring evaluations.
|
||||
|
||||
```bash
|
||||
git clone https://github.com/ViperEkura/AstrAI.git
|
||||
cd AstrAI
|
||||
@@ -132,7 +138,7 @@ Check out the demos in the `scripts/demo/` folder:
|
||||
# Download model weights (required before running demos)
|
||||
python scripts/demo/download.py # model → params/
|
||||
|
||||
# Interactive streaming chat (multi-turn, maintains history)
|
||||
# Single-turn interactive streaming prompt loop (no conversation history)
|
||||
python scripts/demo/stream_chat.py
|
||||
# Type your message after >>, type !exit to quit
|
||||
|
||||
@@ -183,7 +189,7 @@ docker run --gpus all -v /path/to/data:/data -it astrai:latest
|
||||
# Docker Compose (GPU, default)
|
||||
docker compose up -d
|
||||
|
||||
# Docker Compose (CPU only)
|
||||
# Docker Compose CPU server profile (CUDA-only generation scripts/demos are unavailable)
|
||||
docker compose --profile cpu up -d
|
||||
```
|
||||
|
||||
@@ -213,18 +219,23 @@ curl -X POST http://localhost:8000/v1/messages \
|
||||
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
|
||||
|
||||
| Document | Description |
|
||||
|----------|-------------|
|
||||
| [CLI Reference](./assets/docs/params.md) | Parameters for all CLI tools (train, server, generate, preprocess) |
|
||||
| [Architecture](./assets/docs/architecture.md) | System architecture, class diagram & design patterns |
|
||||
| [Training](./assets/docs/training.md) | Training loop, strategies & formulas |
|
||||
| [Inference](./assets/docs/inference.md) | KVCache, continuous batching, sampling & HTTP API |
|
||||
| [Data Flow](./assets/docs/dataflow.md) | Data pipeline, storage backends & dataset architecture |
|
||||
| [Preprocessing](./assets/docs/preprocessing.md) | Declarative JSON-driven data preprocessing |
|
||||
| [Get Started](./docs/get-started.md) | Installation and quickstart |
|
||||
| [CLI Reference](./docs/guides/params.md) | Parameters for all CLI tools (train, server, generate, preprocess) |
|
||||
| [Preprocessing](./docs/guides/preprocessing.md) | Declarative JSON-driven data preprocessing |
|
||||
| [Training](./docs/guides/training.md) | Training loop, strategies & formulas |
|
||||
| [Inference](./docs/guides/inference.md) | KVCache, continuous batching, sampling & HTTP API |
|
||||
| [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
|
||||
|
||||
@@ -245,10 +256,10 @@ For major changes, please open an issue first to discuss what you would like to
|
||||
|
||||
### License
|
||||
|
||||
This project is licensed under the [GPL-3.0 License](LICENSE).
|
||||
This project is licensed under the [Apache-2.0 License](LICENSE).
|
||||
|
||||
---
|
||||
|
||||
<div align="center">
|
||||
<em>A lightweight Transformer framework designed for both high performance and ease of use.</em>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
@@ -1,130 +0,0 @@
|
||||
# Data Flow
|
||||
|
||||
This document describes the data pipeline: from raw text to model input tensors. For creating preprocessing configs, see [Preprocessing Guide](preprocessing.md).
|
||||
|
||||
## Contents
|
||||
|
||||
- [Overview](#overview)
|
||||
- [Data Preparation](#data-preparation) — tokenization, format detection, backends
|
||||
- [Data Keys by Training Type](#data-keys-by-training-type)
|
||||
- [Dataset Architecture](#dataset-architecture)
|
||||
- [Sampler](#sampler)
|
||||
- [DataLoader](#dataloader)
|
||||
|
||||
## Overview
|
||||
|
||||
```
|
||||
JSONL Lines → Pipeline (mask builder) → Tokenized Tensors
|
||||
↓
|
||||
.h5 or .bin storage
|
||||
↓
|
||||
Store.load()
|
||||
↓
|
||||
Store.fetch(begin, end, keys)
|
||||
↓
|
||||
BaseDataset.__getitem__(idx)
|
||||
↓
|
||||
Sampler → DataLoader → Training / Inference
|
||||
```
|
||||
|
||||
## Data Preparation
|
||||
|
||||
Raw text is tokenized via `AutoTokenizer.encode()` and saved as HDF5 (`.h5`) or binary (`.bin` + `meta.json`) files with keyed tensor groups.
|
||||
|
||||
### Tokenization
|
||||
|
||||
The `Pipeline` reads JSONL lines, applies the mask builder (see [Preprocessing](preprocessing.md)), and produces flat token sequences:
|
||||
|
||||
```python
|
||||
# Per JSONL line: messages → chat template → token IDs + loss mask
|
||||
tokens = tokenizer.encode(rendered_text) # List[int]
|
||||
loss_mask = [0, 0, 0, 1, 1, 1, 1, 1, 1] # 0=masked, 1=train
|
||||
# Stored as flat tensors, packed with other lines by packing strategy
|
||||
```
|
||||
|
||||
The output `meta.json` records the storage format, key names, dtype, total token count, and tensor shapes for each shard.
|
||||
|
||||
### Format Detection
|
||||
|
||||
`detect_format(load_path)` inspects the path:
|
||||
|
||||
- If `load_path` is a file: checks suffix — `.h5`/`.hdf5` → `"h5"`, `.jsonl` → `"jsonl"`, unknown suffix raises `ValueError`
|
||||
- If `load_path` is a directory: recursively globs for `*.h5`/`*.hdf5` files → `"h5"`, `*.bin` + `**/meta.json` → `"bin"`, or `*.jsonl` + `dataset_config.json` → `"jsonl"`
|
||||
|
||||
### Store Backends
|
||||
|
||||
Storage format is auto-detected by `detect_format()`; backends are dispatched via registry:
|
||||
|
||||
```
|
||||
StoreFactory.create("h5") → H5Store
|
||||
StoreFactory.create("bin") → MmapStore
|
||||
StoreFactory.create("jsonl") → JsonlStore
|
||||
```
|
||||
|
||||
All three inherit `Store` (base, owns `_data`/`_cum`/`_offsets`/`_normalize`) plus the `Streamable` and `Recordable` mixins, so every backend supports both `fetch(begin, end, keys)` (stream) and `fetch_record(index, keys)` (record) APIs.
|
||||
|
||||
**H5Store**: Reads HDF5 files. Tensors are loaded into host memory and normalized into segmented storage. `segments_are_records=True` — each `data_i` dataset is one record.
|
||||
|
||||
**MmapStore**: Memory-maps `.bin` files. OS page cache sharing is native — no explicit `share_memory_()` needed. Uses `torch.from_numpy(np.memmap(...))`. `segments_are_records=False` — bin segments are contiguous streams; record access is driven by `_offsets` (written when `save_bin(..., record_keys=...)` was used at preprocessing time).
|
||||
|
||||
**JsonlStore**: On-the-fly tokenization of raw JSONL files at load time. Requires a `dataset_config.json` alongside the `.jsonl` files following the same `PipelineConfig` schema with an additional `tokenizer_path` field. Two modes: eager (default, applies `TokenizeTransform` to all records at load) and lazy (`processor=fn` given, defers tokenisation to `fetch_record` — used by DPO/GRPO).
|
||||
|
||||
All backends normalise tensors into `Store._data[Dict[str, List[Tensor]]]` + `Store._cum[Dict[str, List[int]]]` (cumulative lengths for bisect-based stream indexing) + `Store._offsets[Dict[str, List[int]]]` (per-record offsets for record-mode indexing). Nested keys (GRPO `responses`/`masks` as `List[List[Tensor]]`) are stored as-is and excluded from both bookkeepings — they are only accessed record-by-record.
|
||||
|
||||
## Data Keys by Training Type
|
||||
|
||||
| Type | Storage Keys | Access Mode |
|
||||
|------|-------------|-------------|
|
||||
| `seq` | `sequence` (→ input_ids, target_ids via offset-by-1) | stream (`fetch`) |
|
||||
| `sft` | `sequence`, `loss_mask`, `position_ids` | stream (`fetch`) |
|
||||
| `dpo` | `chosen`, `rejected`, `chosen_mask`, `rejected_mask` | record (`fetch_record`) |
|
||||
| `grpo` | `prompts`, `responses`, `masks`, `rewards` | record (`fetch_record`) |
|
||||
|
||||
## Dataset Architecture
|
||||
|
||||
```
|
||||
DatasetFactory.load(train_type, load_path, window_size, stride=None,
|
||||
storage_type=None, tokenizer_path=None,
|
||||
max_position_embeddings=2048, store=None)
|
||||
→ BaseDataset.load(load_path, storage_type=None)
|
||||
→ detect_format(load_path)
|
||||
→ StoreFactory.create(storage_type)
|
||||
→ Store.load(load_path)
|
||||
→ _normalize(raw) # base Store, shared by both backends
|
||||
→ Store._data[Dict[str, List[Tensor]]]
|
||||
+ _cum[Dict[str, List[int]]] (stream mode)
|
||||
+ _offsets[Dict[str, List[int]]] (record mode)
|
||||
|
||||
Stream datasets (SEQ/SFT):
|
||||
BaseDataset.__getitem__(idx)
|
||||
→ get_index(idx) → [begin, end)
|
||||
→ Store.fetch(begin, end, keys) → Tensor / Dict[str, Tensor]
|
||||
|
||||
Record datasets (DPO/GRPO via RecordDataset):
|
||||
RecordDataset.__getitem__(idx)
|
||||
→ Store.fetch_record(idx, keys) → Tensor / Dict[str, Tensor]
|
||||
```
|
||||
|
||||
Class hierarchy: `BaseDataset` ← `SEQDataset` / `SFTDataset` (stream); `BaseDataset` ← `RecordDataset` ← `DPODataset` / `GRPODataset` (record).
|
||||
|
||||
`window_size` = max input length, `stride` = step between consecutive samples (defaults to `window_size`, optional). Only meaningful for stream datasets — record datasets ignore both. `storage_type` defaults to `None` (auto-detect via `detect_format`).
|
||||
|
||||
`tokenizer_path` triggers lazy on-the-fly tokenisation for record datasets on raw JSONL (DPO builds a `dpo_processor`; SEQ/SFT/pre-tokenised backends ignore it). `store` (pre-built `Store`) bypasses `load_path`/`storage_type`/`tokenizer_path` entirely — the caller controls Store construction.
|
||||
|
||||
`Store.fetch(begin, end, keys)` (stream mode, on `Streamable`): accepts a single key (`str`) returning a `Tensor`, or a list of keys returning `Dict[str, Tensor]`. Internally uses `bisect` across multi-segment tensors. Raises `RuntimeError("Store not loaded")` if called before `load()`.
|
||||
|
||||
`Store.fetch_record(index, keys)` (record mode, on `Recordable`): same key API. Uses `_offsets[key]` when present (bin layout with per-record offsets), otherwise indexes `_data[key]` directly (H5/JSONL where each segment is one record).
|
||||
|
||||
## Sampler
|
||||
|
||||
`ResumableDistributedSampler` supports checkpoint-aware distributed sampling:
|
||||
|
||||
- Tracks `start_epoch` / `start_iter` for resume
|
||||
- Shuffle via `torch.Generator(seed + epoch)`
|
||||
- Per-replica index slicing for DDP
|
||||
|
||||
## DataLoader
|
||||
|
||||
Standard PyTorch `DataLoader` with configurable `batch_size`, `num_workers`, `pin_memory`, `prefetch_factor`. Sampler produces indices; dataloader fetches tensor batches via `__getitem__`.
|
||||
|
||||
> Document Update Time: 2026-07-19
|
||||
@@ -1,252 +0,0 @@
|
||||
# Inference
|
||||
|
||||
## Contents
|
||||
|
||||
- [KV Cache](#kv-cache)
|
||||
- [KVCache System](#kvcache-system)
|
||||
- [Continuous Batching](#continuous-batching)
|
||||
- [Sampling](#sampling-strategy-pattern)
|
||||
- [Protocol Handlers](#protocol-handlers-strategy-pattern)
|
||||
- [Engine & GenerateResult](#engine--generateresult)
|
||||
- [HTTP API](#http-api) — endpoints, SSE, errors, stats
|
||||
- [Engine API](#engine-api)
|
||||
|
||||
## KV Cache
|
||||
|
||||
At decode time, only the last query token matters. All previous K/V are cached to avoid recomputation:
|
||||
|
||||
$$
|
||||
o_n = \sum_j \text{softmax}\left(\frac{q_n k_j}{\sqrt{d_k}}\right) v_j
|
||||
$$
|
||||
|
||||
RoPE is applied **before** KV cache write, not after — otherwise position encoding drift occurs.
|
||||
|
||||
## KVCache System
|
||||
|
||||
Seven classes working together, with two concrete cache implementations:
|
||||
|
||||
### ContiguousCache (default)
|
||||
|
||||
```
|
||||
ContiguousCache (simple contiguous per-slot cache)
|
||||
├── ContiguousCacheView bundles k/v tensors + slot indices for attention layers
|
||||
```
|
||||
|
||||
Created by default when no cache is passed to `InferenceScheduler`. Each task occupies a fixed slot of `[max_seq_len, num_key_value_heads, head_dim]`. Simple and efficient for small-to-medium batch sizes.
|
||||
|
||||
### PageCache (paged with prefix sharing)
|
||||
|
||||
```
|
||||
PageCache (paged KV cache with prefix sharing, alternative)
|
||||
├── PagePool orchestrates page allocation + prefix matching
|
||||
│ ├── Allocator bitmask-based page allocator + ref-count + LRU
|
||||
│ └── PrefixCache hash-based prefix matching (page_hash via polynomial hash)
|
||||
├── TaskTable maps task_id → page_table + cached token count
|
||||
├── Storage k_cache / v_cache tensors (num_hidden_layers × n_pages × page_size × num_key_value_heads × head_dim)
|
||||
└── PageCacheView bundles Storage + page_table + total_len for attention layers
|
||||
```
|
||||
|
||||
`isinstance(cache, KVCache)` checks dispatch to the correct view. Both implement the abstract `KVCache` interface used by `Executor` and `InferenceScheduler`.
|
||||
|
||||
## Continuous Batching
|
||||
|
||||
`InferenceScheduler` runs a daemon thread with a 4-phase loop:
|
||||
|
||||
```
|
||||
1. Cleanup → Remove finished tasks, free KV cache slots/pages
|
||||
2. Refill → Pop from waiting_queue, task_alloc resources, activate
|
||||
3. Prefill → Group by (prompt_len, start_pos), run full forward
|
||||
4. Decode → Run single-token forward for each same-position group
|
||||
```
|
||||
|
||||
## Sampling (Strategy Pattern)
|
||||
|
||||
```
|
||||
BaseSamplingStrategy (ABC)
|
||||
├── TemperatureStrategy
|
||||
├── TopKStrategy
|
||||
├── TopPStrategy
|
||||
└── SamplingPipeline
|
||||
```
|
||||
|
||||
`SamplingPipeline` composes them: Temperature → Top-K → Top-P → softmax → multinomial.
|
||||
`sample()` is a convenience shortcut for one-shot usage.
|
||||
|
||||
## Protocol Handlers (Strategy Pattern)
|
||||
|
||||
```python
|
||||
class ProtocolHandler: # concrete orchestrator
|
||||
def __init__(self, request, engine, builder): ...
|
||||
async def handle(self):
|
||||
prompt, ctx, stops = builder.prepare(request, engine)
|
||||
agen = engine.generate_async(prompt, ...)
|
||||
if stream: self._handle_stream(agen, ctx, stops)
|
||||
else: return await self._handle_non_stream(agen, ctx, stops)
|
||||
```
|
||||
|
||||
`ResponseBuilder` (ABC): `prepare()`, `format_stream_start()`, `format_chunk()`, `format_stream_end()`, `format_response()`.
|
||||
|
||||
`OpenAIResponseBuilder` → `/v1/chat/completions`, `AnthropicResponseBuilder` → `/v1/messages`.
|
||||
|
||||
Adding a protocol = one builder file, no handler subclassing needed.
|
||||
|
||||
## Engine & GenerateResult
|
||||
|
||||
```
|
||||
InferenceEngine
|
||||
├── generate(prompt, stream, ...) → str | List[str] | Generator
|
||||
├── generate_with_request(req) → same
|
||||
├── generate_async(prompt, ...) → AsyncGenerator
|
||||
├── get_stats() → Dict
|
||||
└── shutdown()
|
||||
```
|
||||
|
||||
`GenerateResult` uses `Condition` for non-streaming (`wait_completion()`) and `Event` for streaming (`wait()`). Stream callback is `cb(token)`.
|
||||
|
||||
## HTTP API
|
||||
|
||||
```
|
||||
POST /v1/chat/completions OpenAI
|
||||
POST /v1/messages Anthropic
|
||||
GET /health {"status":"ok","model_loaded":true}
|
||||
GET /stats scheduler statistics
|
||||
```
|
||||
|
||||
### OpenAI
|
||||
|
||||
```bash
|
||||
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"messages":[{"role":"user","content":"Hello"}],"max_tokens":512}'
|
||||
```
|
||||
|
||||
Response:
|
||||
```json
|
||||
{
|
||||
"id": "chatcmpl-abc123",
|
||||
"object": "chat.completion",
|
||||
"created": 1717000000,
|
||||
"model": "astrai",
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": "Hello!"}, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 5, "completion_tokens": 10, "total_tokens": 15}
|
||||
}
|
||||
```
|
||||
|
||||
Streaming SSE: `object: "chat.completion.chunk"` — starts with role delta, then token chunks, ends with finish chunk + usage stats, then `data: [DONE]`.
|
||||
|
||||
### Anthropic
|
||||
|
||||
```bash
|
||||
curl -X POST http://localhost:8000/v1/messages \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"model":"astrai","system":"You are helpful.","messages":[{"role":"user","content":"Hello"}],"max_tokens":512}'
|
||||
```
|
||||
|
||||
Supports `stop_sequences` and streaming via `event: content_block_delta`.
|
||||
|
||||
### GenerationRequest Parameters
|
||||
|
||||
| Param | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `messages` | List[dict] | required | Chat messages (role, content) |
|
||||
| `top_k` | int | 50 | Top-k count |
|
||||
| `top_p` | float | 1.0 | Nucleus threshold |
|
||||
| `temperature` | float | 1.0 | Sampling temperature (> 0.0) |
|
||||
| `max_tokens` | Optional[int] | None | Max generation length |
|
||||
| `stream` | bool | False | Stream output |
|
||||
|
||||
### SSE Streaming Format
|
||||
|
||||
**OpenAI** (`/v1/chat/completions`, `stream=true`):
|
||||
|
||||
```
|
||||
data: {"id":"chatcmpl-...","object":"chat.completion.chunk","created":...,"model":"astrai",
|
||||
"choices":[{"index":0,"delta":{"role":"assistant"},"finish_reason":null}]}
|
||||
|
||||
data: {"id":"chatcmpl-...","object":"chat.completion.chunk","created":0,"model":"astrai",
|
||||
"choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]}
|
||||
|
||||
data: {"id":"chatcmpl-...","object":"chat.completion.chunk","created":...,"model":"astrai",
|
||||
"choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}
|
||||
|
||||
data: {"prompt_tokens":5,"completion_tokens":1,"total_tokens":6}
|
||||
|
||||
data: [DONE]
|
||||
```
|
||||
|
||||
**Anthropic** (`/v1/messages`, `stream=true`):
|
||||
|
||||
```
|
||||
event: message_start
|
||||
data: {"type":"message_start","message":{"id":"msg_...","model":"astrai","role":"assistant",
|
||||
"content":[],"usage":{"input_tokens":0}}}
|
||||
|
||||
event: content_block_start
|
||||
data: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}
|
||||
|
||||
event: content_block_delta
|
||||
data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hello"}}
|
||||
|
||||
event: content_block_stop
|
||||
data: {"type":"content_block_stop","index":0}
|
||||
|
||||
event: message_delta
|
||||
data: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":null},"usage":{...}}
|
||||
|
||||
event: message_stop
|
||||
data: {"type":"message_stop"}
|
||||
```
|
||||
|
||||
### Error Responses
|
||||
|
||||
The server returns standard HTTP status codes. Pydantic validation errors (e.g. missing required fields)
|
||||
are handled automatically by FastAPI with 422 status. The only application-level error is engine initialization:
|
||||
|
||||
| Status | Meaning |
|
||||
|--------|---------|
|
||||
| 200 | Success |
|
||||
| 422 | Unprocessable entity (Pydantic validation) |
|
||||
| 503 | Service unavailable (model not loaded, engine not ready) |
|
||||
|
||||
Error response body (503):
|
||||
|
||||
```json
|
||||
{
|
||||
"detail": "Engine not initialized"
|
||||
}
|
||||
```
|
||||
|
||||
### Stats Endpoint
|
||||
|
||||
```
|
||||
GET /stats
|
||||
```
|
||||
|
||||
Response:
|
||||
|
||||
```json
|
||||
{
|
||||
"total_tasks": 128,
|
||||
"total_tokens": 10240,
|
||||
"active_tasks": 3,
|
||||
"waiting_queue": 2
|
||||
}
|
||||
```
|
||||
|
||||
## Engine API
|
||||
|
||||
```python
|
||||
# Non-streaming
|
||||
engine.generate("Hello", stream=False) # -> str
|
||||
engine.generate(["A", "B"], stream=False) # -> List[str]
|
||||
|
||||
# Streaming
|
||||
engine.generate("Hello", stream=True) # -> Generator[str]
|
||||
engine.generate(["A", "B"], stream=True) # -> Generator[Tuple[int, str]]
|
||||
|
||||
# Async
|
||||
async for token in engine.generate_async("Hello", ...): # -> AsyncGenerator[str]
|
||||
print(token)
|
||||
```
|
||||
|
||||
> Document Update Time: 2026-07-09
|
||||
+5
-3
@@ -1,4 +1,4 @@
|
||||
__version__ = "1.3.11"
|
||||
__version__ = "1.3.13"
|
||||
__author__ = "ViperEkura"
|
||||
|
||||
from astrai.config import (
|
||||
@@ -18,7 +18,6 @@ from astrai.dataset import (
|
||||
)
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.inference import (
|
||||
GenerationRequest,
|
||||
InferenceEngine,
|
||||
ProtocolHandler,
|
||||
SamplingPipeline,
|
||||
@@ -26,6 +25,7 @@ from astrai.inference import (
|
||||
run_server,
|
||||
sample,
|
||||
)
|
||||
from astrai.logging import setup_logging
|
||||
from astrai.model import (
|
||||
AutoModel,
|
||||
AutoRegressiveLM,
|
||||
@@ -71,7 +71,6 @@ __all__ = [
|
||||
"EmbeddingEncoder",
|
||||
"EncoderConfig",
|
||||
"ExecutorFactory",
|
||||
"GenerationRequest",
|
||||
"InferenceEngine",
|
||||
"LoRAConfig",
|
||||
"Pipeline",
|
||||
@@ -94,5 +93,8 @@ __all__ = [
|
||||
"only_on_rank",
|
||||
"run_server",
|
||||
"sample",
|
||||
"setup_logging",
|
||||
"spawn_parallel_fn",
|
||||
]
|
||||
|
||||
setup_logging()
|
||||
|
||||
+20
-80
@@ -1,92 +1,32 @@
|
||||
import json
|
||||
from dataclasses import MISSING, dataclass, fields
|
||||
from dataclasses import asdict
|
||||
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:
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
d = {}
|
||||
for fld in fields(self):
|
||||
v = getattr(self, fld.name)
|
||||
if isinstance(v, (str, int, float, bool)):
|
||||
d[fld.name] = v
|
||||
elif v is None:
|
||||
d[fld.name] = None
|
||||
elif isinstance(v, (dict, list, tuple)):
|
||||
try:
|
||||
val = list(v) if isinstance(v, tuple) else v
|
||||
json.dumps(val)
|
||||
d[fld.name] = val
|
||||
except (TypeError, ValueError):
|
||||
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
|
||||
result = {}
|
||||
for k, v in asdict(self).items():
|
||||
if isinstance(v, tuple):
|
||||
v = list(v)
|
||||
try:
|
||||
json.dumps(v)
|
||||
result[k] = v
|
||||
except (TypeError, ValueError):
|
||||
# Skip non-serializable runtime objects (e.g. model_fn, dataset).
|
||||
# TrainConfig mixes hyperparams with callables/datasets; only the
|
||||
# JSON-serializable subset is written to checkpoint meta.
|
||||
pass
|
||||
return result
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, d: Dict[str, Any]) -> Self:
|
||||
hints = get_type_hints(cls)
|
||||
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
|
||||
return cls(**d)
|
||||
|
||||
@classmethod
|
||||
def from_file(cls, path: Union[str, Path]) -> Self:
|
||||
|
||||
+107
-11
@@ -1,9 +1,14 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from pydantic import field_validator
|
||||
from pydantic.dataclasses import dataclass
|
||||
|
||||
from astrai.config.base import BaseConfig
|
||||
from astrai.factory import BaseFactory
|
||||
|
||||
_ATTN_TYPES = frozenset({"gqa", "mla"})
|
||||
_FFN_TYPES = frozenset({"mlp", "moe"})
|
||||
|
||||
|
||||
class ConfigFactory(BaseFactory[BaseConfig]):
|
||||
"""Factory that dispatches config classes by ``model_type``."""
|
||||
@@ -17,7 +22,12 @@ class ConfigFactory(BaseFactory[BaseConfig]):
|
||||
|
||||
@dataclass
|
||||
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
|
||||
neftune_alpha: float = 0.0
|
||||
@@ -26,7 +36,39 @@ class BaseModelConfig(BaseConfig):
|
||||
@dataclass
|
||||
@ConfigFactory.register("autoregressive_lm")
|
||||
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.
|
||||
moe_intermediate_size (Optional[int]): Expert hidden dim, defaults to intermediate_size if None. MoE only.
|
||||
shared_expert_intermediate_size (Optional[int]): Shared expert hidden dim, defaults to intermediate_size if None. MoE only.
|
||||
norm_topk_prob (bool): Normalize top-k routing probabilities. Defaults to True.
|
||||
decoder_sparse_step (int): Frequency of MoE layers, 1=every layer. Defaults to 1.
|
||||
mlp_only_layers (Optional[list[int]]): Layer indices using dense MLP instead of MoE. Defaults to None.
|
||||
"""
|
||||
|
||||
vocab_size: Optional[int] = None
|
||||
hidden_size: Optional[int] = None
|
||||
@@ -34,49 +76,103 @@ class AutoRegressiveLMConfig(BaseModelConfig):
|
||||
rms_norm_eps: Optional[float] = None
|
||||
intermediate_size: Optional[int] = None
|
||||
tie_word_embeddings: Optional[bool] = None
|
||||
|
||||
max_position_embeddings: Optional[int] = None
|
||||
rope_theta: Optional[float] = None
|
||||
rope_scaling: Optional[dict] = None
|
||||
|
||||
attn_type: str = "gqa"
|
||||
num_attention_heads: Optional[int] = None
|
||||
num_key_value_heads: Optional[int] = None
|
||||
use_qk_norm: Optional[bool] = None
|
||||
use_gated_attention: Optional[bool] = None
|
||||
|
||||
kv_lora_rank: Optional[int] = None
|
||||
qk_nope_head_dim: Optional[int] = None
|
||||
qk_rope_head_dim: Optional[int] = None
|
||||
|
||||
ffn_type: str = "mlp"
|
||||
n_routed_experts: Optional[int] = None
|
||||
n_shared_experts: Optional[int] = None
|
||||
n_activated_experts: Optional[int] = None
|
||||
topk_method: Optional[str] = None
|
||||
moe_intermediate_size: Optional[int] = None
|
||||
shared_expert_intermediate_size: Optional[int] = None
|
||||
norm_topk_prob: bool = True
|
||||
decoder_sparse_step: int = 1
|
||||
mlp_only_layers: Optional[list[int]] = None
|
||||
moe_aux_loss_coef: float = 0.01
|
||||
|
||||
@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
|
||||
|
||||
@field_validator("decoder_sparse_step")
|
||||
def _validate_decoder_sparse_step(cls, v: int) -> int:
|
||||
if v < 1:
|
||||
raise ValueError(f"decoder_sparse_step must be at least 1, got {v}")
|
||||
return v
|
||||
|
||||
|
||||
@dataclass
|
||||
@ConfigFactory.register("embedding")
|
||||
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
|
||||
hidden_size: Optional[int] = None
|
||||
num_hidden_layers: Optional[int] = None
|
||||
rms_norm_eps: Optional[float] = None
|
||||
intermediate_size: Optional[int] = None
|
||||
|
||||
max_position_embeddings: Optional[int] = None
|
||||
rope_theta: Optional[float] = None
|
||||
rope_scaling: Optional[dict] = None
|
||||
|
||||
attn_type: str = "gqa"
|
||||
num_attention_heads: Optional[int] = None
|
||||
num_key_value_heads: Optional[int] = None
|
||||
use_qk_norm: Optional[bool] = None
|
||||
use_gated_attention: Optional[bool] = None
|
||||
|
||||
ffn_type: str = "mlp"
|
||||
pooling_type: Optional[str] = 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``.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from dataclasses import field
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from pydantic import field_validator
|
||||
from pydantic.dataclasses import dataclass
|
||||
|
||||
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
|
||||
class InputConfig(BaseConfig):
|
||||
@@ -25,6 +33,10 @@ class InputConfig(BaseConfig):
|
||||
"chosen": {"sections": [{"field": "chosen", ...}]},
|
||||
"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
|
||||
@@ -33,34 +45,17 @@ class InputConfig(BaseConfig):
|
||||
|
||||
@dataclass
|
||||
class ProcessingConfig(BaseConfig):
|
||||
"""Processing configuration.
|
||||
"""Processing configuration for tokenization and packing.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
max_seq_len : int
|
||||
Maximum sequence length (default: 2048).
|
||||
min_chars : int
|
||||
Minimum number of characters to keep (default: 50).
|
||||
max_chars : int
|
||||
Maximum number of characters to keep (default: 2_000_000).
|
||||
max_items : Optional[int]
|
||||
Maximum number of items to process (default: None, unlimited).
|
||||
batch_size : int
|
||||
Number of records tokenized together (default: 256).
|
||||
packing_strategy : str
|
||||
How to pack sequences into a contiguous stream.
|
||||
|
||||
- ``"simple"``: sequential concatenation (default, backward compatible).
|
||||
- ``"bfd"``: best-fit decreasing bin packing, minimises wasted tokens.
|
||||
- ``"bfd_split"``: BFD with over-length sequences split into chunks.
|
||||
max_packed_len : int
|
||||
Maximum length of a packed bin. Sequences longer than this are
|
||||
truncated or split depending on ``packing_strategy`` (default: 8192).
|
||||
truncation_mode : str
|
||||
How to truncate sequences longer than ``max_packed_len``.
|
||||
|
||||
- ``"keep_start"``: keep the first ``max_packed_len`` tokens (default).
|
||||
- ``"keep_end"``: keep the last ``max_packed_len`` tokens.
|
||||
Args:
|
||||
max_seq_len (int): Maximum sequence length. Defaults to 2048.
|
||||
min_chars (int): Minimum number of characters to keep. Defaults to 50.
|
||||
max_chars (int): Maximum number of characters to keep. Defaults to 2_000_000.
|
||||
max_items (Optional[int]): Maximum number of items to process, None=unlimited. Defaults to None.
|
||||
batch_size (int): Number of records tokenized together. Defaults to 256.
|
||||
packing_strategy (str): How to pack sequences: 'simple', 'bfd', or 'bfd_split'. Defaults to "simple".
|
||||
max_packed_len (int): Maximum length of a packed bin. Defaults to 8192.
|
||||
truncation_mode (str): How to truncate over-length sequences: 'keep_start' or 'keep_end'. Defaults to "keep_start".
|
||||
"""
|
||||
|
||||
max_seq_len: int = 2048
|
||||
@@ -72,27 +67,45 @@ class ProcessingConfig(BaseConfig):
|
||||
max_packed_len: int = 8192
|
||||
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
|
||||
class OutputConfig(BaseConfig):
|
||||
"""Output configuration.
|
||||
"""Output configuration for storage.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
domain_key : Optional[str]
|
||||
Domain key for the output store (default: None).
|
||||
storage_format : str
|
||||
Storage format, one of ``"bin"``, ``"jsonl"`` (default: ``"bin"``).
|
||||
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).
|
||||
Args:
|
||||
domain_key (Optional[str]): Domain key for the output store. Defaults to None.
|
||||
storage_format (str): Storage format: 'bin' or 'jsonl'. Defaults to "bin".
|
||||
max_tokens_per_shard (int): Maximum tokens per shard before splitting. Defaults to 100_000_000.
|
||||
dtype (Dict[str, str]): Per-key dtype overrides, e.g. {"input_ids": "int32"}. Defaults to {}.
|
||||
position_ids_mode (str): Position ids mode: 'none', 'doc_reset', or 'continuous'. Defaults to "doc_reset".
|
||||
"""
|
||||
|
||||
domain_key: Optional[str] = None
|
||||
@@ -101,9 +114,36 @@ class OutputConfig(BaseConfig):
|
||||
dtype: Dict[str, str] = field(default_factory=dict)
|
||||
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
|
||||
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
|
||||
input: InputConfig = field(default_factory=InputConfig)
|
||||
mask: Dict[str, str] = field(default_factory=dict)
|
||||
|
||||
+197
-158
@@ -1,7 +1,9 @@
|
||||
from dataclasses import dataclass, field, fields
|
||||
from dataclasses import field
|
||||
from typing import Any, Callable, Dict, List, Optional
|
||||
|
||||
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.lr_scheduler import LRScheduler
|
||||
from torch.utils.data import Dataset
|
||||
@@ -9,173 +11,210 @@ from torch.utils.data import Dataset
|
||||
from astrai.config.base import BaseConfig
|
||||
from astrai.model.components.lora import LoRAConfig
|
||||
|
||||
|
||||
def required(**kw):
|
||||
return {"required": True, **kw}
|
||||
_TRAIN_TYPES = frozenset({"seq", "sft", "dpo", "grpo", "online_grpo", "online_dpo"})
|
||||
_PARALLEL_MODES = frozenset({"none", "ddp", "fsdp"})
|
||||
_BACKENDS = frozenset({"nccl", "gloo"})
|
||||
_START_METHODS = frozenset({"spawn", "fork", "forkserver"})
|
||||
_COMPILE_MODES = frozenset({"default", "reduce-overhead", "max-autotune"})
|
||||
|
||||
|
||||
@dataclass
|
||||
@dataclass(config=ConfigDict(arbitrary_types_allowed=True))
|
||||
class TrainConfig(BaseConfig):
|
||||
# basic setting
|
||||
model_fn: Callable[[], nn.Module] = field(
|
||||
default=None, metadata=required(help="Model factory for training.")
|
||||
)
|
||||
strategy: str = field(default=None, metadata=required(help="Training strategy."))
|
||||
dataset: Dataset = field(
|
||||
default=None, metadata=required(help="Dataset for training.")
|
||||
)
|
||||
optimizer_fn: Callable[[nn.Module], Optimizer] = field(
|
||||
default=None, metadata=required(help="Optimizer factory for training.")
|
||||
)
|
||||
scheduler_fn: Callable[[Optimizer], LRScheduler] = field(
|
||||
default=None, metadata=required(help="Scheduler factory for training.")
|
||||
)
|
||||
n_epoch: int = field(default=1, metadata={"help": "Number of epochs for training."})
|
||||
batch_per_device: int = field(
|
||||
default=4, metadata={"help": "Batch size per device."}
|
||||
)
|
||||
grad_accum_steps: int = field(
|
||||
default=1, metadata={"help": "Number of iterations between steps."}
|
||||
)
|
||||
max_grad_norm: Optional[float] = field(
|
||||
default=1.0,
|
||||
metadata={"help": "Maximum gradient norm. None disables clipping."},
|
||||
)
|
||||
gradient_checkpointing_modules: List[str] = field(
|
||||
default_factory=list,
|
||||
metadata={"help": "Module types to enable activation checkpointing for."},
|
||||
)
|
||||
"""Training configuration.
|
||||
|
||||
# checkpoint setting
|
||||
start_epoch: int = field(default=0, metadata={"help": "Start epoch for training."})
|
||||
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."},
|
||||
)
|
||||
Combines hyperparameters with runtime objects (model_fn, dataset, etc.).
|
||||
Only JSON-serializable fields are written to checkpoint meta via to_dict().
|
||||
|
||||
# lora setting
|
||||
lora: Optional[LoRAConfig] = field(
|
||||
default=None,
|
||||
metadata={"help": "LoRA config. None means full fine-tuning."},
|
||||
)
|
||||
Args:
|
||||
model_fn (Callable[[], nn.Module]): Model factory for training.
|
||||
strategy (str): Training strategy (seq, sft, dpo, grpo, online_*).
|
||||
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.
|
||||
persistent_workers (bool): Keep DataLoader workers alive between epochs. Defaults to False.
|
||||
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.
|
||||
moe_aux_loss_coef (float): Weight applied to the MoE load-balancing loss. Defaults to 0.01.
|
||||
rollout_interval (int): Number of optimizer steps between online rollouts. Defaults to 512.
|
||||
rollout_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
|
||||
log_dir: str = field(
|
||||
default="./checkpoint/logs", metadata={"help": "Directory for metric logs."}
|
||||
)
|
||||
metrics: List[str] = field(
|
||||
default_factory=lambda: ["loss", "lr", "grad_norm"],
|
||||
metadata={"help": "Metrics to record during training."},
|
||||
)
|
||||
model_fn: Callable[[], nn.Module]
|
||||
strategy: str
|
||||
dataset: Dataset
|
||||
optimizer_fn: Callable[[nn.Module], Optimizer]
|
||||
scheduler_fn: Callable[[Optimizer], LRScheduler]
|
||||
optimizer_name: Optional[str] = None
|
||||
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
|
||||
random_seed: int = field(default=3407, metadata={"help": "Random seed."})
|
||||
num_workers: int = field(
|
||||
default=0, metadata={"help": "Number of workers for dataloader."}
|
||||
)
|
||||
prefetch_factor: Optional[int] = field(
|
||||
default=None, metadata={"help": "Prefetch factor for dataloader."}
|
||||
)
|
||||
pin_memory: bool = field(
|
||||
default=False, metadata={"help": "Pin memory for dataloader."}
|
||||
)
|
||||
collate_fn: Optional[Callable[[List[Any]], Any]] = field(
|
||||
default=None,
|
||||
metadata={"help": "Collate function for dataloader (e.g. dpo_collate_fn)."},
|
||||
)
|
||||
start_epoch: int = 0
|
||||
start_samples: int = 0
|
||||
ckpt_dir: str = "./checkpoint"
|
||||
ckpt_interval: int = 5000
|
||||
|
||||
# distributed training
|
||||
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)."},
|
||||
)
|
||||
lora: Optional[LoRAConfig] = None
|
||||
|
||||
# others
|
||||
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)."},
|
||||
)
|
||||
metrics: List[str] = field(default_factory=lambda: ["loss", "lr", "grad_norm"])
|
||||
|
||||
# online rollout
|
||||
rollout_interval: int = field(
|
||||
default=512,
|
||||
metadata={"help": "Number of optimizer steps between online rollouts."},
|
||||
)
|
||||
rollout_temperature: float = field(
|
||||
default=0.7, metadata={"help": "Sampling temperature for online rollout."}
|
||||
)
|
||||
rollout_top_k: int = field(
|
||||
default=0, metadata={"help": "Top-k filtering for online rollout (0=disable)."}
|
||||
)
|
||||
rollout_top_p: float = field(
|
||||
default=0.9,
|
||||
metadata={"help": "Top-p (nucleus) filtering for online rollout."},
|
||||
)
|
||||
rollout_max_tokens: int = field(
|
||||
default=1024,
|
||||
metadata={"help": "Maximum generated tokens per response in rollout."},
|
||||
)
|
||||
reward_model_fn: Optional[Callable] = field(
|
||||
default=None,
|
||||
metadata={
|
||||
"help": "Factory for reward model (required for online RL strategies)."
|
||||
},
|
||||
)
|
||||
random_seed: int = 3407
|
||||
num_workers: int = 0
|
||||
prefetch_factor: Optional[int] = None
|
||||
persistent_workers: bool = False
|
||||
pin_memory: bool = False
|
||||
collate_fn: Optional[Callable[[List[Any]], Any]] = None
|
||||
|
||||
executor_kwargs: Dict[str, Any] = field(
|
||||
default_factory=dict,
|
||||
metadata={"help": "Extra kwargs passed to ExecutorFactory.create()."},
|
||||
)
|
||||
extra_kwargs: Dict[str, Any] = field(
|
||||
default_factory=dict, metadata={"help": "Other arguments."}
|
||||
)
|
||||
nprocs: int = 1
|
||||
backend: str = "nccl"
|
||||
master_addr: str = "localhost"
|
||||
master_port: str = "29500"
|
||||
parallel_mode: str = "none"
|
||||
start_method: str = "spawn"
|
||||
|
||||
def __post_init__(self):
|
||||
self.validate()
|
||||
device_type: str = "cuda"
|
||||
val_dataset: Optional[Dataset] = None
|
||||
val_split: Optional[float] = None
|
||||
val_step: int = 1000
|
||||
neftune_alpha: float = 0.0
|
||||
moe_aux_loss_coef: float = 0.01
|
||||
|
||||
def validate(self):
|
||||
for fld in fields(self):
|
||||
if fld.metadata.get("required") and getattr(self, fld.name) is None:
|
||||
raise ValueError(f"TrainConfig.{fld.name} is required but got None.")
|
||||
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", "moe_aux_loss_coef"
|
||||
)
|
||||
def _validate_non_negative(cls, v):
|
||||
if v < 0:
|
||||
raise ValueError(f"must be non-negative, got {v}")
|
||||
return v
|
||||
|
||||
@field_validator("max_grad_norm")
|
||||
def _validate_max_grad_norm(cls, v: Optional[float]) -> Optional[float]:
|
||||
if v is not None and v <= 0:
|
||||
raise ValueError(f"max_grad_norm must be positive or None, got {v}")
|
||||
return v
|
||||
|
||||
@field_validator("val_split")
|
||||
def _validate_val_split(cls, v: Optional[float]) -> Optional[float]:
|
||||
if v is not None and not 0 < v < 1:
|
||||
raise ValueError(f"val_split must be in (0, 1) or None, got {v}")
|
||||
return v
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _validate_online_strategy(self) -> "TrainConfig":
|
||||
if self.strategy.startswith("online_") and self.reward_model_fn is None:
|
||||
raise ValueError(
|
||||
f"reward_model_fn is required for online RL strategy {self.strategy!r}"
|
||||
)
|
||||
return self
|
||||
|
||||
@@ -6,7 +6,6 @@ from astrai.dataset.dataset import (
|
||||
)
|
||||
from astrai.dataset.sampler import RDSampler
|
||||
from astrai.dataset.storage import (
|
||||
H5Store,
|
||||
JsonlStore,
|
||||
MmapStore,
|
||||
Recordable,
|
||||
@@ -17,9 +16,7 @@ from astrai.dataset.storage import (
|
||||
)
|
||||
from astrai.serialization import (
|
||||
load_bin,
|
||||
load_h5,
|
||||
save_bin,
|
||||
save_h5,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
@@ -31,12 +28,9 @@ __all__ = [
|
||||
"Streamable",
|
||||
"Recordable",
|
||||
"StoreFactory",
|
||||
"H5Store",
|
||||
"MmapStore",
|
||||
"JsonlStore",
|
||||
"detect_format",
|
||||
"save_h5",
|
||||
"load_h5",
|
||||
"save_bin",
|
||||
"load_bin",
|
||||
"RDSampler",
|
||||
|
||||
+44
-12
@@ -25,20 +25,50 @@ function (pure ``record -> Dict[str, Tensor]``) is forwarded to
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from functools import partial
|
||||
from pathlib import Path
|
||||
from typing import Callable, Dict, List, Optional
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
from astrai.config.preprocess_config import PipelineConfig
|
||||
from astrai.dataset.storage import (
|
||||
Store,
|
||||
StoreFactory,
|
||||
detect_format,
|
||||
)
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.preprocessing.transform import TokenizeTransform
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
_DEFAULT_MESSAGES_CONFIG = {
|
||||
"version": 1,
|
||||
"input": {"sections": [{"field": "messages", "action": "$role", "template": True}]},
|
||||
"mask": {"system": "mask", "user": "mask", "assistant": "train"},
|
||||
"mask_default": "mask",
|
||||
"output": {"position_ids_mode": "doc_reset"},
|
||||
}
|
||||
|
||||
|
||||
def _build_jsonl_transform(
|
||||
path: str, tokenizer_path: Optional[str] = None
|
||||
) -> Optional["TokenizeTransform"]:
|
||||
"""Auto-build a TokenizeTransform for JSONL eager loading.
|
||||
|
||||
Reads ``dataset_config.json`` from the data dir if present, or
|
||||
falls back to the built-in chatml SFT config when *tokenizer_path*
|
||||
is provided.
|
||||
"""
|
||||
root = Path(path)
|
||||
config_path = root / "dataset_config.json" if root.is_dir() else None
|
||||
if config_path is not None and config_path.exists():
|
||||
return TokenizeTransform.from_config_file(str(config_path))
|
||||
if tokenizer_path:
|
||||
config = PipelineConfig.from_dict(_DEFAULT_MESSAGES_CONFIG)
|
||||
return TokenizeTransform(config, tokenizer_path)
|
||||
return None
|
||||
|
||||
|
||||
def dpo_tokenize(
|
||||
record: dict,
|
||||
@@ -314,7 +344,7 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
|
||||
stream datasets (SEQ/SFT). Record datasets ignore it.
|
||||
stride: Stride between consecutive stream samples
|
||||
(default: same as *window_size*).
|
||||
storage_type: Storage backend ("h5", "bin", "jsonl") or
|
||||
storage_type: Storage backend ("bin", "jsonl") or
|
||||
None for auto-detection.
|
||||
tokenizer_path: Path to tokenizer for lazy JSONL
|
||||
tokenisation (record datasets only).
|
||||
@@ -349,16 +379,18 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
|
||||
)
|
||||
if processor is not None:
|
||||
store.load(load_path, processor=processor, **kwargs)
|
||||
elif storage_type == "jsonl":
|
||||
transform = _build_jsonl_transform(load_path, tokenizer_path)
|
||||
if transform is None:
|
||||
raise FileNotFoundError(
|
||||
f"JSONL dataset config not found. Expected "
|
||||
f"dataset_config.json alongside *.jsonl files, pass "
|
||||
f"tokenizer_path= for the built-in messages config, or "
|
||||
f"use processor= for lazy on-the-fly tokenisation."
|
||||
)
|
||||
store.load(load_path, transform=transform, **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)
|
||||
store.load(load_path, **kwargs)
|
||||
|
||||
return cls.create(train_type, store=store)
|
||||
|
||||
@@ -384,7 +416,7 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
|
||||
"""Build an on-the-fly tokenisation processor if applicable.
|
||||
|
||||
Only raw JSONL + record datasets (DPO/GRPO) need a processor;
|
||||
pre-tokenised backends (H5/bin) and stream datasets (SEQ/SFT)
|
||||
pre-tokenised backends (bin) and stream datasets (SEQ/SFT)
|
||||
return ``None`` so no tokenizer is loaded.
|
||||
"""
|
||||
if tokenizer_path is None or storage_type != "jsonl":
|
||||
@@ -451,7 +483,7 @@ class DPODataset(BaseDataset):
|
||||
|
||||
Two loading paths (handled by :class:`DatasetFactory`):
|
||||
|
||||
- **Pre-tokenized** (H5/bin): ``store.load(path)`` reads per-record
|
||||
- **Pre-tokenized** (bin): ``store.load(path)`` reads per-record
|
||||
tensors; ``__getitem__`` returns them directly.
|
||||
- **Raw JSONL** (``tokenizer_path=...``): builds a lazy processor
|
||||
via :func:`dpo_processor` that tokenises on the fly — no packing,
|
||||
|
||||
+11
-74
@@ -10,7 +10,6 @@ Architecture (composition over inheritance):
|
||||
Streamable (mixin) — raw token slice fetch(begin, end, keys)
|
||||
Recordable (mixin) — raw record slice fetch_record(idx, keys)
|
||||
|
||||
H5Store(Store, Streamable, Recordable)
|
||||
MmapStore(Store, Streamable, Recordable)
|
||||
JsonlStore(Store, Streamable, Recordable)
|
||||
|
||||
@@ -36,9 +35,9 @@ control. ``store.token_count`` is the total stream token count (what
|
||||
``len(store)`` used to mean in the legacy stream-only API).
|
||||
|
||||
``segments_are_records`` (class attribute on each Store subclass)
|
||||
tells ``_normalize`` whether segments are inherently per-record (H5/
|
||||
JSONL) or opaque shards (bin). Record access for bin relies on
|
||||
``_offsets`` instead.
|
||||
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
|
||||
@@ -56,13 +55,10 @@ from typing import Callable, Dict, List, Optional, Tuple, Union
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.config.preprocess_config import PipelineConfig
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.preprocessing.transform import TokenizeTransform
|
||||
from astrai.serialization import (
|
||||
load_bin,
|
||||
load_bin_offsets,
|
||||
load_h5,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -83,19 +79,10 @@ def detect_format(load_path: str) -> str:
|
||||
root = Path(load_path)
|
||||
if root.is_file():
|
||||
suffix = root.suffix.lower()
|
||||
if suffix in (".h5", ".hdf5"):
|
||||
return "h5"
|
||||
if suffix == ".jsonl":
|
||||
return "jsonl"
|
||||
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)]
|
||||
if bin_files:
|
||||
has_meta = (root / "meta.json").exists() or len(
|
||||
@@ -185,7 +172,7 @@ class Store(ABC):
|
||||
"""Number of records available via :meth:`fetch_record`.
|
||||
|
||||
Non-zero only when the backing layout provides per-record
|
||||
indexing (H5/JSONL segments or bin ``_offsets``).
|
||||
indexing (JSONL segments or bin ``_offsets``).
|
||||
"""
|
||||
return self._num_records
|
||||
|
||||
@@ -269,7 +256,7 @@ class Store(ABC):
|
||||
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 (H5/JSONL), ``_data[key]`` is
|
||||
``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
|
||||
@@ -305,7 +292,7 @@ class Store(ABC):
|
||||
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 H5/JSONL.",
|
||||
"supported). Merge shards or use JSONL.",
|
||||
key,
|
||||
len(segs),
|
||||
)
|
||||
@@ -330,7 +317,7 @@ class Streamable:
|
||||
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 (H5/JSONL/bin+offsets), the
|
||||
that can also serve record access (JSONL/bin+offsets), the
|
||||
``fetch_record`` API from :class:`Recordable` is used instead.
|
||||
"""
|
||||
|
||||
@@ -415,33 +402,6 @@ class StoreFactory(BaseFactory["Store"]):
|
||||
"""Factory for creating Store instances by type name."""
|
||||
|
||||
|
||||
@StoreFactory.register("h5")
|
||||
class H5Store(Store, Streamable, Recordable):
|
||||
"""HDF5-based storage backend (pre-tokenized data).
|
||||
|
||||
Each key is stored as a group of per-record datasets (``data_0``,
|
||||
``data_1``, …). Supports both access modes:
|
||||
|
||||
- **Stream**: ``fetch(begin, end, key)`` and ``store[i]`` slice
|
||||
across concatenated records via ``_cum`` — used by SEQ/SFT.
|
||||
- **Record**: ``fetch_record(i, key)`` and ``store[i]`` (when
|
||||
``window_size == 0``) index ``_data[key]`` directly — used by
|
||||
DPO/GRPO.
|
||||
"""
|
||||
|
||||
segments_are_records = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
window_size: int = 0,
|
||||
stride: Optional[int] = None,
|
||||
):
|
||||
super().__init__(window_size=window_size, stride=stride)
|
||||
|
||||
def load(self, path: str, **kwargs):
|
||||
self._normalize(load_h5(path))
|
||||
|
||||
|
||||
@StoreFactory.register("bin")
|
||||
class MmapStore(Store, Streamable, Recordable):
|
||||
"""Memory-mapped binary storage backend.
|
||||
@@ -574,19 +534,8 @@ class JsonlStore(Store, Streamable, Recordable):
|
||||
``len(store)`` returns ``num_records``; stream primitives raise.
|
||||
"""
|
||||
|
||||
CONFIG_NAME = "dataset_config.json"
|
||||
segments_are_records = True
|
||||
|
||||
_DEFAULT_MESSAGES_CONFIG = {
|
||||
"version": 1,
|
||||
"input": {
|
||||
"sections": [{"field": "messages", "action": "$role", "template": True}]
|
||||
},
|
||||
"mask": {"system": "mask", "user": "mask", "assistant": "train"},
|
||||
"mask_default": "mask",
|
||||
"output": {"position_ids_mode": "doc_reset"},
|
||||
}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
window_size: int = 0,
|
||||
@@ -607,22 +556,10 @@ class JsonlStore(Store, Streamable, Recordable):
|
||||
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)
|
||||
raise ValueError(
|
||||
"JsonlStore eager mode requires transform=. "
|
||||
"Use DatasetFactory.load() which auto-constructs it."
|
||||
)
|
||||
|
||||
transformed = transform.apply(records)
|
||||
self._normalize(transformed)
|
||||
|
||||
@@ -4,27 +4,52 @@ Public API:
|
||||
- ``attn_decode`` — single-query decode 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
|
||||
|
||||
Interface (shared by all wrappers):
|
||||
causal_offset: -1 = non-causal; >=0 = absolute position of first Q token
|
||||
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True = keep)
|
||||
scale: 0.0 = auto (1/sqrt(head_dim)); >0 = explicit
|
||||
layout: "bhld" (default) or "blhd"
|
||||
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
|
||||
(blhd). Scale is always ``1/sqrt(head_dim)``.
|
||||
|
||||
Causal and mask can coexist — both are applied simultaneously.
|
||||
|
||||
Each wrapper dispatches to its compiled CUDA kernel (``astrai.extension.attn_*``)
|
||||
when available, otherwise falls back to ``torch.nn.functional.scaled_dot_product_attention``.
|
||||
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.backend import (
|
||||
ATTN_BACKEND,
|
||||
AttentionBackend,
|
||||
AttentionBackendFactory,
|
||||
CudaBackend,
|
||||
FlashAttnBackend,
|
||||
TorchNativeBackend,
|
||||
apply_rotary_emb,
|
||||
attention,
|
||||
attn_backend,
|
||||
get_backend,
|
||||
)
|
||||
from astrai.extension.loader import KERNEL_NAMES, is_available
|
||||
from astrai.extension.ops import attention, attn_decode, attn_paged_decode, attn_prefill
|
||||
from astrai.extension.ops import (
|
||||
TensorLayout,
|
||||
attn_decode,
|
||||
attn_paged_decode,
|
||||
attn_prefill,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"ATTN_BACKEND",
|
||||
"AttentionBackend",
|
||||
"AttentionBackendFactory",
|
||||
"CudaBackend",
|
||||
"TorchNativeBackend",
|
||||
"FlashAttnBackend",
|
||||
"TensorLayout",
|
||||
"attention",
|
||||
"attn_backend",
|
||||
"get_backend",
|
||||
"attn_decode",
|
||||
"attn_paged_decode",
|
||||
"attn_prefill",
|
||||
"attention",
|
||||
"is_available",
|
||||
"KERNEL_NAMES",
|
||||
"apply_rotary_emb",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
"""Backend selection, fallbacks, and execution policies."""
|
||||
|
||||
from astrai.extension.backend.attention import (
|
||||
ATTN_BACKEND,
|
||||
AttentionBackend,
|
||||
AttentionBackendFactory,
|
||||
CudaBackend,
|
||||
FlashAttnBackend,
|
||||
TorchNativeBackend,
|
||||
attention,
|
||||
attn_backend,
|
||||
get_backend,
|
||||
)
|
||||
from astrai.extension.backend.rotary import apply_rotary_emb
|
||||
|
||||
__all__ = [
|
||||
"ATTN_BACKEND",
|
||||
"AttentionBackend",
|
||||
"AttentionBackendFactory",
|
||||
"CudaBackend",
|
||||
"FlashAttnBackend",
|
||||
"TorchNativeBackend",
|
||||
"apply_rotary_emb",
|
||||
"attention",
|
||||
"attn_backend",
|
||||
"get_backend",
|
||||
]
|
||||
@@ -0,0 +1,724 @@
|
||||
"""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 default (cuda > flash > torch, overridable via
|
||||
``ASTR_BACKEND``).
|
||||
|
||||
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
|
||||
import functools
|
||||
import os
|
||||
import threading
|
||||
from abc import ABC, abstractmethod
|
||||
from contextlib import contextmanager
|
||||
from typing import TYPE_CHECKING, Optional, Union
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.extension.loader import is_available
|
||||
from astrai.extension.ops.attention import (
|
||||
attn_paged_decode,
|
||||
attn_paged_prefill,
|
||||
)
|
||||
from astrai.factory import BaseFactory
|
||||
|
||||
try:
|
||||
import flash_attn as _flash_attn
|
||||
except Exception:
|
||||
_flash_attn = None
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from astrai.inference.cache import KVCache
|
||||
|
||||
|
||||
_default_backend: Optional["AttentionBackend"] = None
|
||||
_default_backend_lock = threading.Lock()
|
||||
_env_backend_name: Optional[str] = None
|
||||
_env_backend: Optional["AttentionBackend"] = None
|
||||
_current_backend: contextvars.ContextVar[Optional["AttentionBackend"]] = (
|
||||
contextvars.ContextVar("attn_backend", default=None)
|
||||
)
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=1)
|
||||
def flash_attn_available() -> bool:
|
||||
if not torch.cuda.is_available():
|
||||
return False
|
||||
fa = _flash_attn
|
||||
if fa is None:
|
||||
return False
|
||||
|
||||
try:
|
||||
major = int(fa.__version__.split(".")[0])
|
||||
cc = torch.cuda.get_device_capability()
|
||||
cc_num = cc[0] * 10 + cc[1]
|
||||
except Exception:
|
||||
major, cc_num = 0, 0
|
||||
if (major >= 3 and cc_num < 90) or (major < 3 and 0 < cc_num < 70):
|
||||
return False
|
||||
|
||||
try:
|
||||
if not hasattr(fa, "flash_attn_func"):
|
||||
return False
|
||||
x = torch.zeros(1, 1, 1, 64, device="cuda", dtype=torch.bfloat16)
|
||||
out = fa.flash_attn_func(x, x, x, causal=True)
|
||||
return bool(torch.isfinite(out).all().item())
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
class ATTN_BACKEND(enum.Enum):
|
||||
"""Backend selector enum, mirroring ``torch.nn.attention.SDPBackend``."""
|
||||
|
||||
TORCH_NATIVE = "torch_native"
|
||||
CUDA = "cuda"
|
||||
FLASH = "flash"
|
||||
|
||||
|
||||
def _priority_backends() -> list["AttentionBackend"]:
|
||||
"""Available backends in priority order: cuda -> flash -> torch."""
|
||||
backends: list[AttentionBackend] = []
|
||||
if is_available("attn_paged_decode") and is_available("attn_paged_prefill"):
|
||||
backends.append(CudaBackend())
|
||||
if flash_attn_available():
|
||||
backends.append(FlashAttnBackend())
|
||||
backends.append(TorchNativeBackend())
|
||||
return backends
|
||||
|
||||
|
||||
def _backend_supports(
|
||||
backend: "AttentionBackend",
|
||||
q: Tensor,
|
||||
kv_cache: Optional["KVCache"],
|
||||
attn_mask: Optional[Tensor],
|
||||
is_causal: bool,
|
||||
fwd: Optional[str],
|
||||
) -> bool:
|
||||
"""Whether ``backend`` can run this attention call.
|
||||
|
||||
The CUDA kernels are bf16-only, support head_dim in 32/64/128/256, and
|
||||
need a KV cache (decode/prefill); everything else falls back to torch.
|
||||
"""
|
||||
if isinstance(backend, CudaBackend):
|
||||
return (
|
||||
fwd in ("prefill", "decode")
|
||||
and kv_cache is not None
|
||||
and q.ndim == 3
|
||||
and q.dtype == torch.bfloat16
|
||||
and q.size(-1) in (32, 64, 128, 256)
|
||||
and is_available(f"attn_paged_{fwd}")
|
||||
)
|
||||
if isinstance(backend, FlashAttnBackend):
|
||||
if not flash_attn_available():
|
||||
return False
|
||||
if q.dtype not in (torch.float16, torch.bfloat16):
|
||||
return False
|
||||
if fwd is not None:
|
||||
return q.ndim == 3 and hasattr(_flash_attn, "flash_attn_varlen_func")
|
||||
if attn_mask is None or is_causal:
|
||||
return True
|
||||
return attn_mask.dim() == 4
|
||||
return True
|
||||
|
||||
|
||||
def _resolve_default_backend() -> "AttentionBackend":
|
||||
"""Pick the highest-priority available backend (cuda -> flash -> torch).
|
||||
|
||||
Resolved lazily on first ``get_backend()`` and cached. Per-call
|
||||
capability fallback happens in ``attention()``, so the default is
|
||||
safe for training and fp32 models.
|
||||
"""
|
||||
return _priority_backends()[0]
|
||||
|
||||
|
||||
def _environment_backend() -> Optional["AttentionBackend"]:
|
||||
"""Resolve the process-wide ``ASTR_BACKEND`` override, if configured."""
|
||||
global _env_backend, _env_backend_name
|
||||
name = os.environ.get("ASTR_BACKEND", "").strip().lower()
|
||||
if not name:
|
||||
return None
|
||||
if name != _env_backend_name:
|
||||
with _default_backend_lock:
|
||||
if name != _env_backend_name:
|
||||
try:
|
||||
_env_backend = AttentionBackendFactory.create(name)
|
||||
except (ValueError, RuntimeError):
|
||||
_env_backend = None
|
||||
_env_backend_name = name
|
||||
return _env_backend
|
||||
|
||||
|
||||
def _resolve_backend(
|
||||
backend: Optional[Union[str, ATTN_BACKEND, "AttentionBackend", type]] = None,
|
||||
) -> "AttentionBackend":
|
||||
"""Resolve a backend configuration, defaulting to the process policy."""
|
||||
if backend is not None:
|
||||
if isinstance(backend, ATTN_BACKEND):
|
||||
return AttentionBackendFactory.create(backend.value)
|
||||
if isinstance(backend, str):
|
||||
return AttentionBackendFactory.create(backend)
|
||||
if isinstance(backend, type) and issubclass(backend, AttentionBackend):
|
||||
return backend()
|
||||
if isinstance(backend, AttentionBackend):
|
||||
return backend
|
||||
raise TypeError(
|
||||
f"expected a registered name, ATTN_BACKEND, AttentionBackend type, "
|
||||
f"or instance, got {type(backend).__name__}"
|
||||
)
|
||||
|
||||
global _default_backend
|
||||
if _default_backend is None:
|
||||
with _default_backend_lock:
|
||||
if _default_backend is None:
|
||||
_default_backend = _resolve_default_backend()
|
||||
return _default_backend
|
||||
|
||||
|
||||
def get_backend(
|
||||
use_default: bool = True,
|
||||
) -> Optional["AttentionBackend"]:
|
||||
"""Return the context override, optionally falling back to the process default.
|
||||
|
||||
``ASTR_BACKEND`` is a process-wide override and takes precedence over the
|
||||
context value. Pass ``use_default=False`` at request submission to retain
|
||||
only an environment override or the caller's :func:`attn_backend` value.
|
||||
"""
|
||||
return (
|
||||
_environment_backend()
|
||||
or _current_backend.get()
|
||||
or (_resolve_backend() if use_default else None)
|
||||
)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def attn_backend(backend: Union[str, ATTN_BACKEND, "AttentionBackend", type]):
|
||||
"""Context manager to select an attention backend.
|
||||
|
||||
Mirrors ``torch.nn.attention.sdpa_kernel``. Accepts an
|
||||
registered name, ``ATTN_BACKEND`` enum value, backend class, or instance.
|
||||
|
||||
Examples::
|
||||
|
||||
with attn_backend(ATTN_BACKEND.TORCH_NATIVE):
|
||||
...
|
||||
with attn_backend(TorchNativeBackend):
|
||||
...
|
||||
with attn_backend(TorchNativeBackend()):
|
||||
...
|
||||
"""
|
||||
instance = _resolve_backend(backend)
|
||||
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."""
|
||||
if n_rep == 1:
|
||||
return x
|
||||
n_heads, head_dim = x.shape[-2:]
|
||||
return (
|
||||
x.unsqueeze(-2)
|
||||
.expand(*x.shape[:-2], n_heads, n_rep, head_dim)
|
||||
.reshape(*x.shape[:-2], n_heads * n_rep, head_dim)
|
||||
)
|
||||
|
||||
|
||||
def _write_and_gather_kv(
|
||||
kv_cache: "KVCache",
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
layer_id: int,
|
||||
q: Tensor,
|
||||
attn_mask: Optional[Tensor],
|
||||
) -> tuple[Tensor, Tensor]:
|
||||
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]
|
||||
if q.size(1) == 1 and attn_mask is not None and attn_mask.dim() == 4:
|
||||
pos_mask = attn_mask[:, 0, 0]
|
||||
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))
|
||||
return kv_cache.k_buffer[layer_id, indices], kv_cache.v_buffer[layer_id, indices]
|
||||
|
||||
|
||||
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,
|
||||
fwd: Optional[str] = None,
|
||||
) -> 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]
|
||||
"""
|
||||
explicit = get_backend(use_default=False)
|
||||
backend = get_backend()
|
||||
if fwd is None and explicit is None:
|
||||
backend = TorchNativeBackend()
|
||||
if not _backend_supports(backend, q, kv_cache, attn_mask, is_causal, fwd):
|
||||
if explicit is not None:
|
||||
raise RuntimeError(
|
||||
f"Explicitly-set backend {type(backend).__name__} cannot "
|
||||
f"handle this attention call (shape={q.shape}, "
|
||||
f"dtype={q.dtype}, kv_cache={'none' if kv_cache is None else 'present'}, "
|
||||
f"attn_mask={'none' if attn_mask is None else 'present'}). "
|
||||
f"Remove the attn_backend() context or switch to a compatible backend."
|
||||
)
|
||||
for candidate in _priority_backends():
|
||||
if isinstance(candidate, type(backend)):
|
||||
continue
|
||||
if _backend_supports(candidate, q, kv_cache, attn_mask, is_causal, fwd):
|
||||
backend = candidate
|
||||
break
|
||||
return backend.forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal, fwd)
|
||||
|
||||
|
||||
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,
|
||||
fwd: Optional[str] = None,
|
||||
) -> 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 fwd == "decode":
|
||||
return self.fwd_decode(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||
if fwd == "prefill" or fwd is None:
|
||||
return self.fwd_prefill(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||
raise ValueError(f"unsupported attention forward mode: {fwd}")
|
||||
|
||||
@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."""
|
||||
|
||||
@staticmethod
|
||||
def supports_graph() -> bool:
|
||||
"""Return True if this backend supports CUDA-graph capture.
|
||||
|
||||
Override in subclasses that can run under ``torch.cuda.graph``.
|
||||
|
||||
Called on the *active* backend instance (or its class) — a cheap
|
||||
boolean check with no side-effects.
|
||||
"""
|
||||
return False
|
||||
|
||||
|
||||
class AttentionBackendFactory(BaseFactory[AttentionBackend]):
|
||||
"""Factory for registered attention backends."""
|
||||
|
||||
|
||||
@AttentionBackendFactory.register(ATTN_BACKEND.TORCH_NATIVE.value)
|
||||
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.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def supports(**kwargs) -> bool:
|
||||
return True
|
||||
|
||||
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 q.ndim == 4:
|
||||
n_rep = q.size(2) // k.size(2)
|
||||
if n_rep > 1:
|
||||
k = repeat_kv(k, n_rep)
|
||||
v = repeat_kv(v, n_rep)
|
||||
return (
|
||||
F.scaled_dot_product_attention(
|
||||
q.permute(0, 2, 1, 3),
|
||||
k.permute(0, 2, 1, 3),
|
||||
v.permute(0, 2, 1, 3),
|
||||
attn_mask,
|
||||
is_causal=is_causal,
|
||||
)
|
||||
.permute(0, 2, 1, 3)
|
||||
.contiguous()
|
||||
)
|
||||
|
||||
if kv_cache is None or kv_cache.qo_indptr is None:
|
||||
raise ValueError("packed attention requires KV cache metadata")
|
||||
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
||||
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
||||
outputs = []
|
||||
n_rep = q.size(1) // k.size(1)
|
||||
for i in range(kv_cache.req_pool_indices.numel()):
|
||||
q_start = int(kv_cache.qo_indptr[i])
|
||||
q_end = int(kv_cache.qo_indptr[i + 1])
|
||||
indices = kv_cache.req_to_token[
|
||||
kv_cache.req_pool_indices[i], : kv_cache.seq_lens[i]
|
||||
]
|
||||
k_i = kv_cache.k_buffer[layer_id, indices]
|
||||
v_i = kv_cache.v_buffer[layer_id, indices]
|
||||
if n_rep > 1:
|
||||
k_i = repeat_kv(k_i, n_rep)
|
||||
v_i = repeat_kv(v_i, n_rep)
|
||||
q_len = q_end - q_start
|
||||
kv_len = k_i.size(0)
|
||||
q_pos = torch.arange(kv_len - q_len, kv_len, device=q.device)
|
||||
causal_mask = q_pos[:, None] >= torch.arange(kv_len, device=q.device)
|
||||
out = F.scaled_dot_product_attention(
|
||||
q[q_start:q_end].transpose(0, 1).unsqueeze(0),
|
||||
k_i.transpose(0, 1).unsqueeze(0),
|
||||
v_i.transpose(0, 1).unsqueeze(0),
|
||||
attn_mask=causal_mask,
|
||||
)
|
||||
outputs.append(out.squeeze(0).transpose(0, 1))
|
||||
return torch.cat(outputs)
|
||||
|
||||
|
||||
@AttentionBackendFactory.register(ATTN_BACKEND.CUDA.value)
|
||||
class CudaBackend(AttentionBackend):
|
||||
"""CUDA kernel backend with direct KV cache access.
|
||||
|
||||
Decode path: writes K/V to the flat pool, then calls
|
||||
``attn_paged_decode`` with req_to_token + kv_indptr.
|
||||
|
||||
Prefill path: writes K/V to the flat pool, then calls
|
||||
``attn_paged_prefill`` with ragged-batch support via qo_indptr +
|
||||
kv_indptr.
|
||||
|
||||
``kv_cache is None`` (training) raises — the per-call fallback to
|
||||
torch SDPA for training / fp32 / unsupported head_dim happens in the
|
||||
``attention()`` entry point.
|
||||
|
||||
Raises ``RuntimeError`` if the required kernel is not available.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def supports(**kwargs) -> bool:
|
||||
head_dim = kwargs.get("head_dim", -1)
|
||||
return (
|
||||
torch.cuda.is_available()
|
||||
and head_dim in (32, 64, 128, 256)
|
||||
and is_available("attn_paged_decode")
|
||||
and is_available("attn_paged_prefill")
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def supports_graph() -> bool:
|
||||
return True
|
||||
|
||||
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:
|
||||
raise RuntimeError("CudaBackend does not support training (kv_cache=None)")
|
||||
|
||||
loc = kv_cache.out_cache_loc
|
||||
kv_cache.k_buffer[layer_id, loc] = k
|
||||
kv_cache.v_buffer[layer_id, loc] = v
|
||||
|
||||
kv_indptr = kv_cache.kv_indptr
|
||||
|
||||
out = attn_paged_decode(
|
||||
q,
|
||||
kv_cache.k_buffer[layer_id],
|
||||
kv_cache.v_buffer[layer_id],
|
||||
kv_cache.req_to_token,
|
||||
kv_cache.req_pool_indices,
|
||||
kv_indptr,
|
||||
is_causal=True,
|
||||
o_part_buf=kv_cache.decode_o_part,
|
||||
ml_part_buf=kv_cache.decode_ml_part,
|
||||
out_buf=kv_cache.decode_out,
|
||||
)
|
||||
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:
|
||||
raise RuntimeError("CudaBackend does not support training (kv_cache=None)")
|
||||
|
||||
loc = kv_cache.out_cache_loc
|
||||
kv_cache.k_buffer[layer_id, loc] = k
|
||||
kv_cache.v_buffer[layer_id, loc] = v
|
||||
|
||||
out = attn_paged_prefill(
|
||||
q,
|
||||
kv_cache.k_buffer[layer_id],
|
||||
kv_cache.v_buffer[layer_id],
|
||||
kv_cache.req_to_token,
|
||||
kv_cache.req_pool_indices,
|
||||
kv_cache.kv_indptr,
|
||||
kv_cache.qo_indptr,
|
||||
attn_mask,
|
||||
is_causal=is_causal,
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
@AttentionBackendFactory.register(ATTN_BACKEND.FLASH.value)
|
||||
class FlashAttnBackend(AttentionBackend):
|
||||
"""FlashAttention backend via the optional ``flash-attn`` package.
|
||||
|
||||
Decode (q_len=1, contiguous cache): uses ``flash_attn_with_kvcache``,
|
||||
which reads K/V directly from the flat pool via cache_batch_idx +
|
||||
cache_seqlens — no materialized KV gather.
|
||||
|
||||
Prefill / non-contiguous decode: falls back to KV gather +
|
||||
``flash_attn_func``.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def supports(**kwargs) -> bool:
|
||||
return flash_attn_available()
|
||||
|
||||
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_packed(q, k, v, kv_cache, layer_id)
|
||||
|
||||
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 q.ndim == 3:
|
||||
return self._forward_packed(q, k, v, kv_cache, layer_id)
|
||||
return self._forward_dense(q, k, v, attn_mask, is_causal)
|
||||
|
||||
def _forward_dense(
|
||||
self,
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
attn_mask: Optional[Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
n_rep = q.size(2) // k.size(2)
|
||||
if n_rep > 1:
|
||||
k = repeat_kv(k, n_rep)
|
||||
v = repeat_kv(v, n_rep)
|
||||
|
||||
if attn_mask is not None and not is_causal and attn_mask.dim() != 4:
|
||||
raise ValueError(
|
||||
"FlashAttnBackend does not support a custom attention mask; "
|
||||
"use a causal mask or select TorchNativeBackend."
|
||||
)
|
||||
fa = _flash_attn
|
||||
if fa is None:
|
||||
raise RuntimeError(
|
||||
"FlashAttnBackend requires the optional 'flash-attn' package. "
|
||||
"Install with `pip install flash-attn`."
|
||||
)
|
||||
out = fa.flash_attn_func(
|
||||
q.contiguous(),
|
||||
k.contiguous(),
|
||||
v.contiguous(),
|
||||
causal=is_causal or (attn_mask is not None and attn_mask.dim() == 4),
|
||||
)
|
||||
return out.contiguous()
|
||||
|
||||
def _forward_packed(
|
||||
self,
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
kv_cache: "KVCache",
|
||||
layer_id: int,
|
||||
) -> Tensor:
|
||||
fa = _flash_attn
|
||||
if fa is None or not hasattr(fa, "flash_attn_varlen_func"):
|
||||
raise RuntimeError("packed inference requires flash_attn_varlen_func")
|
||||
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
||||
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
||||
page_table = kv_cache.req_to_token[
|
||||
kv_cache.req_pool_indices, : kv_cache.max_len
|
||||
]
|
||||
positions = torch.arange(kv_cache.max_len, device=q.device)
|
||||
indices = page_table[positions.unsqueeze(0) < kv_cache.seq_lens.unsqueeze(1)]
|
||||
k_flat = kv_cache.k_buffer[layer_id, indices].contiguous()
|
||||
v_flat = kv_cache.v_buffer[layer_id, indices].contiguous()
|
||||
out = fa.flash_attn_varlen_func(
|
||||
q.contiguous(),
|
||||
k_flat,
|
||||
v_flat,
|
||||
kv_cache.qo_indptr,
|
||||
kv_cache.kv_indptr,
|
||||
int((kv_cache.qo_indptr[1:] - kv_cache.qo_indptr[:-1]).max()),
|
||||
int(kv_cache.seq_lens.max()),
|
||||
dropout_p=0.0,
|
||||
causal=True,
|
||||
)
|
||||
return out
|
||||
@@ -0,0 +1,53 @@
|
||||
"""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
|
||||
from astrai.extension.ops.rotary import rotary_emb as _cuda_rotary
|
||||
|
||||
_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
|
||||
):
|
||||
return _cuda_rotary(x, freqs_cis)
|
||||
return _torch_apply(x, freqs_cis)
|
||||
@@ -0,0 +1,334 @@
|
||||
"""FP8 training: scaling state and aten::linear dispatch.
|
||||
|
||||
Layered (see also ``ops/fp8.py`` for the CUDA interface adapter):
|
||||
|
||||
1. Kernel interface: ``ops.fp8`` - the only module touching the pybind.
|
||||
2. Training state (this module): per-tensor scales, amax history, delayed
|
||||
scaling, and the ``fp8_autocast`` context (TE-style, like
|
||||
``torch.autocast``).
|
||||
3. aten::linear integration (this module): registers the CUDA impl and the
|
||||
M/N alignment guard.
|
||||
|
||||
Usage::
|
||||
|
||||
from astrai.extension.fp8 import fp8_autocast
|
||||
|
||||
with fp8_autocast(enabled=True):
|
||||
logits = model(input_ids)
|
||||
loss.backward()
|
||||
|
||||
Importing this module registers the aten::linear CUDA implementation.
|
||||
"""
|
||||
|
||||
from contextlib import contextmanager
|
||||
|
||||
import torch
|
||||
from torch.library import Library
|
||||
|
||||
from astrai.extension.ops.fp8 import (
|
||||
linear_backward_scaled,
|
||||
linear_forward_scaled,
|
||||
)
|
||||
|
||||
E4M3_MAX = 448.0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Layer 2: training state (scales, amax history, delayed scaling, autocast)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class FP8TensorMeta:
|
||||
"""Scales + amax state for one weight tensor and its paired activations.
|
||||
|
||||
- weight: delayed scale from a 16-step amax history window (TE style)
|
||||
- x/g: delayed one step, reuse the quantize kernel's free atomic amax
|
||||
"""
|
||||
|
||||
__slots__ = (
|
||||
"scale",
|
||||
"scale_inv",
|
||||
"amax_history",
|
||||
"idx",
|
||||
"x_scale",
|
||||
"x_scale_inv",
|
||||
"x_history",
|
||||
"x_idx",
|
||||
"g_scale",
|
||||
"g_scale_inv",
|
||||
"g_history",
|
||||
"g_idx",
|
||||
"w_init",
|
||||
"x_init",
|
||||
"g_init",
|
||||
)
|
||||
|
||||
def __init__(self, device: torch.device, update_interval: int):
|
||||
self.scale = torch.ones(1, device=device, dtype=torch.float32)
|
||||
self.scale_inv = torch.ones(1, device=device, dtype=torch.float32)
|
||||
self.amax_history = torch.ones(
|
||||
update_interval, device=device, dtype=torch.float32
|
||||
)
|
||||
self.idx = 0
|
||||
self.x_scale = torch.ones(1, device=device, dtype=torch.float32)
|
||||
self.x_scale_inv = torch.ones(1, device=device, dtype=torch.float32)
|
||||
self.x_history = torch.ones(update_interval, device=device, dtype=torch.float32)
|
||||
self.x_idx = 0
|
||||
self.g_scale = torch.ones(1, device=device, dtype=torch.float32)
|
||||
self.g_scale_inv = torch.ones(1, device=device, dtype=torch.float32)
|
||||
self.g_history = torch.ones(update_interval, device=device, dtype=torch.float32)
|
||||
self.g_idx = 0
|
||||
self.w_init = False
|
||||
self.x_init = False
|
||||
self.g_init = False
|
||||
|
||||
def init_scale(self, t: torch.Tensor) -> None:
|
||||
"""Immediate scale from the current amax; used on the first call.
|
||||
|
||||
A scale of 1 would underflow small activations/gradients (e4m3 min
|
||||
normal is 2^-6); initialize from the actual amax once, then delayed
|
||||
updates take over.
|
||||
"""
|
||||
amax = t.abs().amax().to(torch.float32).clamp_min(1e-12)
|
||||
self.scale.copy_(amax / E4M3_MAX)
|
||||
self.scale_inv.copy_(E4M3_MAX / amax)
|
||||
self.record(amax)
|
||||
|
||||
def push_x_scale(self, amax: torch.Tensor) -> None:
|
||||
"""Window update for the activation scale (delayed, TE style)."""
|
||||
self.x_history[self.x_idx] = amax.reshape(())
|
||||
self.x_idx = (self.x_idx + 1) % self.x_history.numel()
|
||||
m = self.x_history.max()
|
||||
self.x_scale.copy_(m / E4M3_MAX)
|
||||
self.x_scale_inv.copy_(E4M3_MAX / m)
|
||||
|
||||
def push_g_scale(self, amax: torch.Tensor) -> None:
|
||||
"""Window update for the gradient scale (delayed, TE style)."""
|
||||
self.g_history[self.g_idx] = amax.reshape(())
|
||||
self.g_idx = (self.g_idx + 1) % self.g_history.numel()
|
||||
m = self.g_history.max()
|
||||
self.g_scale.copy_(m / E4M3_MAX)
|
||||
self.g_scale_inv.copy_(E4M3_MAX / m)
|
||||
|
||||
def record(self, amax: torch.Tensor) -> None:
|
||||
"""Push the latest amax into the ring buffer (device-side copy, no sync)."""
|
||||
self.amax_history[self.idx] = amax.reshape(())
|
||||
self.idx = (self.idx + 1) % self.amax_history.numel()
|
||||
|
||||
def refresh(self) -> None:
|
||||
"""Recompute scale from the amax history window (delayed scaling)."""
|
||||
amax = self.amax_history.max()
|
||||
if amax > 0:
|
||||
self.scale.copy_(amax / E4M3_MAX)
|
||||
self.scale_inv.copy_(E4M3_MAX / amax)
|
||||
|
||||
|
||||
class FP8State:
|
||||
"""Global fp8 training state, TE-style."""
|
||||
|
||||
def __init__(self, update_interval: int = 16):
|
||||
self.enabled = False
|
||||
self.update_interval = update_interval
|
||||
self.step_count = 0
|
||||
self._metas: dict[tuple, FP8TensorMeta] = {}
|
||||
self._last_device: torch.device | None = None
|
||||
|
||||
def _get_device(self, t: torch.Tensor) -> torch.device:
|
||||
if self._last_device is None:
|
||||
self._last_device = t.device
|
||||
return t.device
|
||||
|
||||
def get_weight_meta(self, w: torch.Tensor) -> FP8TensorMeta:
|
||||
key = (w.data_ptr(), w.shape, w.dtype)
|
||||
meta = self._metas.get(key)
|
||||
if meta is None:
|
||||
meta = FP8TensorMeta(self._get_device(w), self.update_interval)
|
||||
self._metas[key] = meta
|
||||
return meta
|
||||
|
||||
def step(self) -> None:
|
||||
"""Advance the counter and refresh all weight scales every N steps."""
|
||||
self.step_count += 1
|
||||
if self.step_count % self.update_interval == 0:
|
||||
for meta in self._metas.values():
|
||||
meta.refresh()
|
||||
|
||||
def reset(self) -> None:
|
||||
self.enabled = False
|
||||
self.step_count = 0
|
||||
self._metas.clear()
|
||||
self._last_device = None
|
||||
|
||||
|
||||
# Global singleton: autograd backward runs on the engine worker threads, so
|
||||
# thread-local state would lose the fp8 flag during loss.backward(). The GIL
|
||||
# protects Python-side mutation; the CUDA kernels take their own mutex.
|
||||
_state = FP8State()
|
||||
|
||||
|
||||
def fp8_state() -> FP8State:
|
||||
return _state
|
||||
|
||||
|
||||
@contextmanager
|
||||
def fp8_autocast(enabled: bool = True, update_interval: int = 16):
|
||||
"""Autocast-style context: fp8 linear dispatch on this thread.
|
||||
|
||||
Usage::
|
||||
|
||||
with fp8_autocast(enabled=True):
|
||||
logits = model(input_ids) # aten::linear -> fp8 path
|
||||
loss.backward()
|
||||
|
||||
The scale-update counter advances once per ``enter`` (one training step),
|
||||
refreshing weight scales from their amax history every ``update_interval``.
|
||||
"""
|
||||
state = fp8_state()
|
||||
prev_enabled = state.enabled
|
||||
prev_interval = state.update_interval
|
||||
state.enabled = enabled
|
||||
state.update_interval = update_interval
|
||||
try:
|
||||
if enabled:
|
||||
state.step()
|
||||
yield
|
||||
finally:
|
||||
state.enabled = prev_enabled
|
||||
state.update_interval = prev_interval
|
||||
|
||||
|
||||
def fp8_linear_forward(x: torch.Tensor, w: torch.Tensor, bias=None):
|
||||
"""TE-style scaled fp8 linear forward (called from the aten::linear impl).
|
||||
|
||||
x uses the delayed scale of its paired weight meta (amax from the previous
|
||||
forward of this linear); the quantize kernel emits the current amax for the
|
||||
next step. No extra abs/max reduce.
|
||||
"""
|
||||
if bias is None:
|
||||
bias = torch.empty(0, device=x.device, dtype=x.dtype)
|
||||
state = fp8_state()
|
||||
meta = state.get_weight_meta(w)
|
||||
if not meta.w_init:
|
||||
meta.init_scale(w)
|
||||
meta.w_init = True
|
||||
if not meta.x_init:
|
||||
amax = x.abs().amax().to(torch.float32).clamp_min(1e-12)
|
||||
meta.x_history.fill_(amax)
|
||||
meta.x_scale.copy_(amax / E4M3_MAX)
|
||||
meta.x_scale_inv.copy_(E4M3_MAX / amax)
|
||||
meta.x_init = True
|
||||
amax_x = torch.empty(1, device=x.device, dtype=torch.float32)
|
||||
amax_w = torch.empty(1, device=x.device, dtype=torch.float32)
|
||||
out = linear_forward_scaled(
|
||||
x,
|
||||
w,
|
||||
bias,
|
||||
meta.x_scale,
|
||||
meta.scale,
|
||||
meta.x_scale_inv,
|
||||
meta.scale_inv,
|
||||
amax_x,
|
||||
amax_w,
|
||||
)
|
||||
meta.record(amax_w)
|
||||
meta.push_x_scale(amax_x)
|
||||
return out
|
||||
|
||||
|
||||
def fp8_linear_backward(g, x, w, masks):
|
||||
"""TE-style scaled fp8 linear backward (called from aten::linear_backward)."""
|
||||
state = fp8_state()
|
||||
meta = state.get_weight_meta(w)
|
||||
if not meta.g_init:
|
||||
amax = g.abs().amax().to(torch.float32).clamp_min(1e-12)
|
||||
meta.g_history.fill_(amax)
|
||||
meta.g_scale.copy_(amax / E4M3_MAX)
|
||||
meta.g_scale_inv.copy_(E4M3_MAX / amax)
|
||||
meta.g_init = True
|
||||
amax_g = torch.empty(1, device=g.device, dtype=torch.float32)
|
||||
out = linear_backward_scaled(
|
||||
g,
|
||||
x,
|
||||
w,
|
||||
masks,
|
||||
meta.g_scale,
|
||||
meta.scale,
|
||||
meta.x_scale,
|
||||
meta.g_scale_inv,
|
||||
meta.scale_inv,
|
||||
meta.x_scale_inv,
|
||||
amax_g,
|
||||
)
|
||||
meta.push_g_scale(amax_g)
|
||||
return out
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Layer 3: aten::linear integration
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def fp8_linear_enable(enabled: bool = True) -> None:
|
||||
"""Toggle fp8 dispatch for aten::linear (global; backward runs on engine
|
||||
worker threads, so a thread-local flag would be lost during backward)."""
|
||||
fp8_state().enabled = enabled
|
||||
|
||||
|
||||
def fp8_linear_enabled() -> bool:
|
||||
return fp8_state().enabled
|
||||
|
||||
|
||||
def _fp8_supported(x: torch.Tensor, w: torch.Tensor) -> bool:
|
||||
"""cuBLASLt fp8 requires M % 16 == 0 and N % 16 == 0 (K is padded)."""
|
||||
m = x.numel() // x.size(-1)
|
||||
return m % 16 == 0 and w.size(0) % 16 == 0
|
||||
|
||||
|
||||
def _linear_cuda_impl(x: torch.Tensor, w: torch.Tensor, bias=None):
|
||||
if (
|
||||
fp8_linear_enabled()
|
||||
and x.dtype == torch.bfloat16
|
||||
and w.dtype == torch.bfloat16
|
||||
and _fp8_supported(x, w)
|
||||
):
|
||||
return fp8_linear_forward(x, w, bias)
|
||||
return torch.ops.aten.linear.default.redispatch(
|
||||
torch._C.DispatchKeySet(torch._C.DispatchKey.CompositeImplicitAutograd),
|
||||
x,
|
||||
w,
|
||||
bias,
|
||||
)
|
||||
|
||||
|
||||
def _linear_backward_cuda_impl(input_tensor, grad_output, weight, output_mask):
|
||||
if (
|
||||
fp8_linear_enabled()
|
||||
and weight.dtype == torch.bfloat16
|
||||
and _fp8_supported(grad_output, weight)
|
||||
):
|
||||
return fp8_linear_backward(grad_output, input_tensor, weight, list(output_mask))
|
||||
compute_dtype = weight.dtype
|
||||
grad = grad_output.to(compute_dtype)
|
||||
grad_2d = grad.reshape(-1, weight.size(0))
|
||||
input_2d = input_tensor.reshape(-1, input_tensor.size(-1)).to(compute_dtype)
|
||||
grad_input = (
|
||||
torch.mm(grad_2d, weight)
|
||||
if output_mask[0]
|
||||
else torch.empty(0, device=input_tensor.device, dtype=input_tensor.dtype)
|
||||
)
|
||||
grad_weight = (
|
||||
torch.mm(grad_2d.t(), input_2d)
|
||||
if output_mask[1]
|
||||
else torch.empty(0, device=input_tensor.device, dtype=input_tensor.dtype)
|
||||
)
|
||||
grad_bias = (
|
||||
grad.sum(dim=0)
|
||||
if output_mask[2]
|
||||
else torch.empty(0, device=input_tensor.device, dtype=input_tensor.dtype)
|
||||
)
|
||||
return grad_input.reshape_as(input_tensor), grad_weight, grad_bias
|
||||
|
||||
|
||||
_lib = Library("aten", "IMPL", "CUDA")
|
||||
_lib.impl("linear", _linear_cuda_impl)
|
||||
_lib.impl("linear_backward", _linear_backward_cuda_impl)
|
||||
@@ -0,0 +1 @@
|
||||
"""Compiled CUDA kernel modules (``*.so``) live here, kept separate from Python source."""
|
||||
@@ -11,14 +11,21 @@ import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
KERNEL_NAMES = ["attn_decode", "attn_prefill", "attn_paged_decode"]
|
||||
KERNEL_NAMES = [
|
||||
"attn_decode",
|
||||
"attn_prefill",
|
||||
"attn_paged_decode",
|
||||
"attn_paged_prefill",
|
||||
"rotary_emb",
|
||||
"fp8_mm",
|
||||
]
|
||||
|
||||
_available: dict[str, bool] = {}
|
||||
_modules: dict[str, object] = {}
|
||||
|
||||
for _name in KERNEL_NAMES:
|
||||
try:
|
||||
_mod = importlib.import_module(f".{_name}", package=__package__)
|
||||
_mod = importlib.import_module(f".lib.{_name}", package=__package__)
|
||||
_available[_name] = True
|
||||
_modules[_name] = _mod
|
||||
except ImportError:
|
||||
|
||||
@@ -1,298 +0,0 @@
|
||||
"""GQA attention wrapper functions — one entry point per compiled kernel.
|
||||
|
||||
Each wrapper dispatches to its CUDA kernel (loaded in ``loader.py``) when
|
||||
available, otherwise falls back to ``torch`` SDPA.
|
||||
|
||||
Interface (all functions):
|
||||
causal_offset: -1 = non-causal; >=0 = absolute position of first Q token
|
||||
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool)
|
||||
scale: 0.0 = auto (1/sqrt(head_dim)); >0 = explicit
|
||||
layout: "bhld" (default) or "blhd"
|
||||
|
||||
Add new kernel wrappers here; split into per-variant files only if this file
|
||||
grows large.
|
||||
"""
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from astrai.extension.loader import _available, _modules
|
||||
|
||||
_LAYOUT_CODES: dict[str, int] = {"bhld": 0, "blhd": 1}
|
||||
|
||||
|
||||
def _parse_layout(layout: str | int) -> int:
|
||||
if isinstance(layout, int):
|
||||
return layout
|
||||
code = _LAYOUT_CODES.get(layout.lower())
|
||||
if code is None:
|
||||
raise ValueError(
|
||||
f"unknown layout '{layout}', expected one of {list(_LAYOUT_CODES)}"
|
||||
)
|
||||
return code
|
||||
|
||||
|
||||
def _to_bhld(t: torch.Tensor, layout: int) -> torch.Tensor:
|
||||
"""Normalize to b h l d view. Zero-copy transpose if layout==1 (b l h d)."""
|
||||
if layout == 1:
|
||||
return t.transpose(1, 2)
|
||||
return t
|
||||
|
||||
|
||||
def _expand_kv_heads(
|
||||
k: torch.Tensor, v: torch.Tensor, q_head: int
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Expand K/V heads to match Q heads for GQA fallback."""
|
||||
kv_head = k.size(1)
|
||||
if kv_head == q_head:
|
||||
return k, v
|
||||
group = q_head // kv_head
|
||||
k = k.repeat_interleave(group, dim=1)
|
||||
v = v.repeat_interleave(group, dim=1)
|
||||
return k, v
|
||||
|
||||
|
||||
def _build_attn_mask(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
mask: torch.Tensor | None,
|
||||
causal_offset: int,
|
||||
scale: float,
|
||||
) -> tuple[torch.Tensor | None, float]:
|
||||
"""Build SDPA-compatible attn_mask + resolved scale.
|
||||
|
||||
q and k must already be in b h l d layout.
|
||||
Causal and mask can coexist: causal sets -inf above the diagonal, mask
|
||||
sets -inf for padded positions. Both are OR'd into a single bool mask.
|
||||
"""
|
||||
q_len = q.size(2)
|
||||
kv_len = k.size(2)
|
||||
head_dim = q.size(3)
|
||||
resolved_scale = scale if scale and scale > 0 else 1.0 / math.sqrt(head_dim)
|
||||
|
||||
attn_mask = None
|
||||
|
||||
if mask is not None:
|
||||
if mask.dim() == 2:
|
||||
# [batch, kv_len] → [batch, 1, 1, kv_len]
|
||||
attn_mask = mask[:, None, None, :]
|
||||
elif mask.dim() == 3:
|
||||
# [batch, q_len, kv_len] → [batch, 1, q_len, kv_len]
|
||||
attn_mask = mask[:, None, :, :]
|
||||
else:
|
||||
raise ValueError(f"mask must be 2D or 3D, got {mask.dim()}D")
|
||||
|
||||
if causal_offset >= 0:
|
||||
batch = q.size(0)
|
||||
# q row i attends to kv cols 0..(causal_offset + i)
|
||||
q_idx = torch.arange(q_len, device=q.device).unsqueeze(1) # [q_len, 1]
|
||||
kv_idx = torch.arange(kv_len, device=q.device).unsqueeze(0) # [1, kv_len]
|
||||
causal_bool = kv_idx > (causal_offset + q_idx) # True = masked out
|
||||
causal_mask = causal_bool.unsqueeze(0).expand(
|
||||
batch, -1, -1
|
||||
) # [batch, q_len, kv_len]
|
||||
causal_mask = causal_mask[:, None, :, :] # [batch, 1, q_len, kv_len]
|
||||
|
||||
if attn_mask is not None:
|
||||
attn_mask = attn_mask | causal_mask
|
||||
else:
|
||||
attn_mask = causal_mask
|
||||
|
||||
return attn_mask, resolved_scale
|
||||
|
||||
|
||||
def _torch_fallback(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
mask: torch.Tensor | None,
|
||||
causal_offset: int,
|
||||
scale: float,
|
||||
q_layout: int,
|
||||
kv_layout: int | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Reference attention via ``scaled_dot_product_attention``.
|
||||
|
||||
q_layout / kv_layout: 0 = b h l d, 1 = b l h d.
|
||||
If kv_layout is None, uses q_layout (Q and K/V share the same layout).
|
||||
"""
|
||||
if kv_layout is None:
|
||||
kv_layout = q_layout
|
||||
q = _to_bhld(q, q_layout)
|
||||
k = _to_bhld(k, kv_layout)
|
||||
v = _to_bhld(v, kv_layout)
|
||||
k, v = _expand_kv_heads(k, v, q.size(1))
|
||||
attn_mask, resolved_scale = _build_attn_mask(q, k, mask, causal_offset, scale)
|
||||
out = F.scaled_dot_product_attention(
|
||||
q, k, v, attn_mask=attn_mask, is_causal=False, scale=resolved_scale
|
||||
)
|
||||
# Restore Q's original layout
|
||||
if q_layout == 1:
|
||||
out = out.transpose(1, 2)
|
||||
return out
|
||||
|
||||
|
||||
def _gather_kv_from_pages(
|
||||
page_table: torch.Tensor,
|
||||
k_cache: torch.Tensor,
|
||||
v_cache: torch.Tensor,
|
||||
page_size: int,
|
||||
kv_len: int,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Gather contiguous K/V from paged cache for torch SDPA fallback.
|
||||
|
||||
Shapes:
|
||||
page_table : [batch, max_pages] (int64)
|
||||
k_cache : [n_pages, page_size, n_kv_heads, head_dim]
|
||||
v_cache : same as k_cache
|
||||
Returns:
|
||||
k, v : [batch, kv_len, n_kv_heads, head_dim] (b l h d)
|
||||
"""
|
||||
batch, max_pages = page_table.shape
|
||||
_, ps, n_kv_heads, head_dim = k_cache.shape
|
||||
if ps != page_size:
|
||||
raise ValueError(f"k_cache page_size mismatch: {ps} vs {page_size}")
|
||||
|
||||
# Vectorized gather: build physical page + offset indices, then advanced-index
|
||||
positions = torch.arange(kv_len, device=page_table.device)
|
||||
logical_pages = positions // page_size # [kv_len]
|
||||
page_offsets = positions % page_size # [kv_len]
|
||||
|
||||
phys_pages = page_table[:, logical_pages] # [batch, kv_len]
|
||||
# k_cache[phys_pages, page_offsets] → [batch, kv_len, n_kv_heads, head_dim] (b l h d)
|
||||
k = k_cache[phys_pages, page_offsets]
|
||||
v = v_cache[phys_pages, page_offsets]
|
||||
return k, v
|
||||
|
||||
|
||||
def attn_decode(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
mask: torch.Tensor | None = None,
|
||||
causal_offset: int = -1,
|
||||
scale: float = 0.0,
|
||||
layout: str = "bhld",
|
||||
) -> torch.Tensor:
|
||||
li = _parse_layout(layout)
|
||||
if _available["attn_decode"]:
|
||||
return _modules["attn_decode"].attn_decode(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
mask=mask,
|
||||
causal_offset=causal_offset,
|
||||
scale=scale,
|
||||
layout=li,
|
||||
)
|
||||
return _torch_fallback(q, k, v, mask, causal_offset, scale, q_layout=li)
|
||||
|
||||
|
||||
def attn_prefill(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
mask: torch.Tensor | None = None,
|
||||
causal_offset: int = -1,
|
||||
scale: float = 0.0,
|
||||
layout: str = "bhld",
|
||||
) -> torch.Tensor:
|
||||
li = _parse_layout(layout)
|
||||
if _available["attn_prefill"]:
|
||||
return _modules["attn_prefill"].attn_prefill(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
mask=mask,
|
||||
causal_offset=causal_offset,
|
||||
scale=scale,
|
||||
layout=li,
|
||||
)
|
||||
return _torch_fallback(q, k, v, mask, causal_offset, scale, q_layout=li)
|
||||
|
||||
|
||||
def attn_paged_decode(
|
||||
q: torch.Tensor,
|
||||
page_table: torch.Tensor,
|
||||
k_cache: torch.Tensor,
|
||||
v_cache: torch.Tensor,
|
||||
page_size: int,
|
||||
kv_len: int,
|
||||
mask: torch.Tensor | None = None,
|
||||
causal_offset: int = -1,
|
||||
scale: float = 0.0,
|
||||
layout: str = "bhld",
|
||||
) -> torch.Tensor:
|
||||
li = _parse_layout(layout)
|
||||
if _available["attn_paged_decode"]:
|
||||
return _modules["attn_paged_decode"].attn_paged_decode(
|
||||
q,
|
||||
page_table,
|
||||
k_cache,
|
||||
v_cache,
|
||||
page_size,
|
||||
kv_len,
|
||||
mask=mask,
|
||||
causal_offset=causal_offset,
|
||||
scale=scale,
|
||||
layout=li,
|
||||
)
|
||||
# Gathered K/V are always b l h d
|
||||
k, v = _gather_kv_from_pages(page_table, k_cache, v_cache, page_size, kv_len)
|
||||
return _torch_fallback(
|
||||
q, k, v, mask, causal_offset, scale, q_layout=li, kv_layout=1
|
||||
)
|
||||
|
||||
|
||||
def attention(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
mask: torch.Tensor | None = None,
|
||||
causal_offset: int = -1,
|
||||
scale: float = 0.0,
|
||||
layout: str = "bhld",
|
||||
) -> torch.Tensor:
|
||||
"""Dispatch to decode or prefill attention based on the query length.
|
||||
|
||||
A query length of one is the decode case; longer queries use prefill.
|
||||
The paged-cache decode path cannot be selected here because its page-table
|
||||
arguments are not part of this interface.
|
||||
"""
|
||||
li = _parse_layout(layout)
|
||||
|
||||
if q.ndim not in (2, 3, 4) or k.ndim != q.ndim or v.ndim != q.ndim:
|
||||
raise ValueError(
|
||||
"q, k, and v must all have the same rank in {2, 3, 4}, "
|
||||
f"got {q.ndim}D, {k.ndim}D, {v.ndim}D"
|
||||
)
|
||||
if k.shape != v.shape:
|
||||
raise ValueError(
|
||||
f"k and v must have the same shape, got {k.shape} and {v.shape}"
|
||||
)
|
||||
|
||||
original_ndim = q.ndim
|
||||
if original_ndim == 2:
|
||||
# [L, D] -> [1, 1, L, D] or [1, L, 1, D]
|
||||
q = q.unsqueeze(0).unsqueeze(1 if li == 0 else 2)
|
||||
k = k.unsqueeze(0).unsqueeze(1 if li == 0 else 2)
|
||||
v = v.unsqueeze(0).unsqueeze(1 if li == 0 else 2)
|
||||
elif original_ndim == 3:
|
||||
# [B, L, D] -> single-head 4D input.
|
||||
q = q.unsqueeze(1 if li == 0 else 2)
|
||||
k = k.unsqueeze(1 if li == 0 else 2)
|
||||
v = v.unsqueeze(1 if li == 0 else 2)
|
||||
|
||||
q_len = q.size(2 if li == 0 else 1)
|
||||
if q_len == 1:
|
||||
out = attn_decode(q, k, v, mask, causal_offset, scale, layout)
|
||||
else:
|
||||
out = attn_prefill(q, k, v, mask, causal_offset, scale, layout)
|
||||
|
||||
if original_ndim == 2:
|
||||
return out.squeeze(0).squeeze(0 if li == 0 else 1)
|
||||
if original_ndim == 3:
|
||||
return out.squeeze(1 if li == 0 else 2)
|
||||
return out
|
||||
@@ -0,0 +1,19 @@
|
||||
"""Stateless wrappers around compiled extension kernels."""
|
||||
|
||||
from astrai.extension.ops.attention import (
|
||||
TensorLayout,
|
||||
attn_decode,
|
||||
attn_paged_decode,
|
||||
attn_paged_prefill,
|
||||
attn_prefill,
|
||||
)
|
||||
from astrai.extension.ops.rotary import rotary_emb
|
||||
|
||||
__all__ = [
|
||||
"TensorLayout",
|
||||
"attn_decode",
|
||||
"attn_paged_decode",
|
||||
"attn_paged_prefill",
|
||||
"attn_prefill",
|
||||
"rotary_emb",
|
||||
]
|
||||
@@ -0,0 +1,188 @@
|
||||
"""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 enum
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.extension.loader import _available, _modules
|
||||
|
||||
|
||||
class TensorLayout(enum.IntEnum):
|
||||
"""Q/K/V tensor layout, mirrors the C++ ``TensorLayout`` enum in ``attn_common.h``.
|
||||
|
||||
Kernels internally operate on BHLD; BLHD inputs are transposed at entry.
|
||||
"""
|
||||
|
||||
BHLD = 0 # [batch, n_heads, seq_len, head_dim]
|
||||
BLHD = 1 # [batch, seq_len, n_heads, head_dim]
|
||||
|
||||
|
||||
def _check_available(name: str):
|
||||
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: Optional[torch.Tensor] = 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=TensorLayout.BLHD
|
||||
)
|
||||
|
||||
|
||||
def attn_prefill(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
mask: Optional[torch.Tensor] = 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=TensorLayout.BLHD
|
||||
)
|
||||
|
||||
|
||||
def attn_paged_decode(
|
||||
q: torch.Tensor,
|
||||
k_cache: torch.Tensor,
|
||||
v_cache: torch.Tensor,
|
||||
req_to_token: torch.Tensor,
|
||||
req_pool_indices: torch.Tensor,
|
||||
kv_indptr: torch.Tensor,
|
||||
mask: Optional[torch.Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
o_part_buf: Optional[torch.Tensor] = None,
|
||||
ml_part_buf: Optional[torch.Tensor] = None,
|
||||
out_buf: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
"""SGLang-style paged decode (q_len == 1, flat KV pool).
|
||||
|
||||
Reads K/V directly from a flat pool [size, kv_head, head_dim] via
|
||||
req_to_token indirect indexing. Each request has its own seq_len
|
||||
(from kv_indptr), eliminating padding waste.
|
||||
|
||||
Args:
|
||||
q: [batch, n_heads, head_dim] (bf16, 3D — no seq dim)
|
||||
k_cache: [pool_size, n_kv_heads, head_dim] (bf16, flat)
|
||||
v_cache: same as k_cache
|
||||
req_to_token: [num_reqs, max_context_len] (int32) — token -> slot
|
||||
req_pool_indices: [batch] (int32) — rows into req_to_token
|
||||
kv_indptr: [batch+1] (int32) — prefix sum of per-request seq_lens
|
||||
mask: 2D [batch, max_context_len] (bool, True=keep) or None
|
||||
is_causal: apply causal mask
|
||||
o_part_buf: pre-allocated split-KV o partial buffer (workflow bypass)
|
||||
ml_part_buf: pre-allocated split-KV m/l buffer (workflow bypass)
|
||||
out_buf: pre-allocated output buffer [batch, n_heads, head_dim] (graph-safe)
|
||||
|
||||
Returns:
|
||||
[batch, n_heads, head_dim] (bf16, 3D)
|
||||
"""
|
||||
_check_available("attn_paged_decode")
|
||||
causal_offset = 0 if is_causal else -1
|
||||
return _modules["attn_paged_decode"].attn_paged_decode(
|
||||
q,
|
||||
k_cache,
|
||||
v_cache,
|
||||
req_to_token,
|
||||
req_pool_indices,
|
||||
kv_indptr,
|
||||
mask=mask,
|
||||
causal_offset=causal_offset,
|
||||
o_part_buf=o_part_buf,
|
||||
ml_part_buf=ml_part_buf,
|
||||
out_buf=out_buf,
|
||||
)
|
||||
|
||||
|
||||
def attn_paged_prefill(
|
||||
q: torch.Tensor,
|
||||
k_cache: torch.Tensor,
|
||||
v_cache: torch.Tensor,
|
||||
req_to_token: torch.Tensor,
|
||||
req_pool_indices: torch.Tensor,
|
||||
kv_indptr: torch.Tensor,
|
||||
qo_indptr: torch.Tensor,
|
||||
mask: Optional[torch.Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""SGLang-style paged prefill (ragged batch, flat KV pool).
|
||||
|
||||
Reads K/V directly from a flat pool [size, kv_head, head_dim] via
|
||||
req_to_token. Supports ragged batches: each request has its own
|
||||
q_len and kv_len, addressed via qo_indptr and kv_indptr.
|
||||
|
||||
Args:
|
||||
q: [total_q, n_heads, head_dim] (bf16, 3D — flattened across requests)
|
||||
k_cache: [pool_size, n_kv_heads, head_dim] (bf16, flat)
|
||||
v_cache: same as k_cache
|
||||
req_to_token: [num_reqs, max_context_len] (int32)
|
||||
req_pool_indices: [batch] (int32)
|
||||
kv_indptr: [batch+1] (int32) — prefix sum of per-request kv_lens
|
||||
qo_indptr: [batch+1] (int32) — prefix sum of per-request q_lens
|
||||
mask: 4D [batch, 1, q_len, kv_len] (bool, True=keep) or None
|
||||
is_causal: apply causal mask
|
||||
|
||||
Returns:
|
||||
[total_q, n_heads, head_dim] (bf16, 3D)
|
||||
"""
|
||||
_check_available("attn_paged_prefill")
|
||||
causal_offset = 0 if is_causal else -1
|
||||
return _modules["attn_paged_prefill"].attn_paged_prefill(
|
||||
q,
|
||||
k_cache,
|
||||
v_cache,
|
||||
req_to_token,
|
||||
req_pool_indices,
|
||||
kv_indptr,
|
||||
qo_indptr,
|
||||
mask,
|
||||
causal_offset=causal_offset,
|
||||
)
|
||||
@@ -0,0 +1,73 @@
|
||||
"""FP8 CUDA kernel interface adapter (the only module touching the pybind.
|
||||
|
||||
Isolates the ``fp8_mm`` CUDA extension behind stable Python functions:
|
||||
- availability / dtype checks and clear errors
|
||||
- torch.library ``custom::fp8_mm`` registration (meta + CPU fallback)
|
||||
- quantize-in-GEMM primitives used by ``fp8.py`` training state
|
||||
|
||||
Policy (scales, amax history, delayed scaling, autocast) lives in ``fp8.py``;
|
||||
this module is stateless.
|
||||
"""
|
||||
|
||||
import torch
|
||||
from torch.library import custom_op
|
||||
|
||||
from astrai.extension.loader import get_module, is_available
|
||||
|
||||
|
||||
def _mod():
|
||||
if not is_available("fp8_mm"):
|
||||
raise RuntimeError(
|
||||
"CUDA kernel 'fp8_mm' is not available. Build with CSRC_KERNELS=true."
|
||||
)
|
||||
return get_module("fp8_mm")
|
||||
|
||||
|
||||
@custom_op("custom::fp8_mm", mutates_args=())
|
||||
def fp8_mm(
|
||||
a: torch.Tensor, b: torch.Tensor, sx: torch.Tensor, sw: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
"""FP8 e4m3 GEMM: a[M,K] x b[N,K] -> bf16[M,N] (pre-scaled inputs)."""
|
||||
|
||||
|
||||
@fp8_mm.register_fake
|
||||
def _fp8_mm_fake(a, b, sx, sw):
|
||||
return torch.empty((a.size(0), b.size(1)), device=a.device, dtype=torch.bfloat16)
|
||||
|
||||
|
||||
@fp8_mm.register_kernel("cuda")
|
||||
def _fp8_mm_cuda(a, b, sx, sw):
|
||||
return _mod().fp8_mm(a, b)
|
||||
|
||||
|
||||
@fp8_mm.register_kernel("cpu")
|
||||
def _fp8_mm_cpu(a, b, sx, sw):
|
||||
return torch.mm(a.float(), b.float().t()).to(torch.bfloat16)
|
||||
|
||||
|
||||
def linear_forward_scaled(x, w, bias, sx, sw, sx_inv, sw_inv, amax_x, amax_w):
|
||||
"""Quantize x/w with per-tensor scales + cuBLASLt GEMM + bias -> bf16.
|
||||
|
||||
x/w: [..., K] / [N, K] bf16; sx/sw: f32 scale tensors (device scalars);
|
||||
sx_inv/sw_inv: 1/scale; amax_x/amax_w: f32 buffers receiving max-abs.
|
||||
"""
|
||||
if not (x.dtype == torch.bfloat16 and w.dtype == torch.bfloat16):
|
||||
raise TypeError(f"fp8 forward requires bf16 inputs, got {x.dtype}/{w.dtype}")
|
||||
return _mod().fp8_linear_forward_scaled(
|
||||
x, w, bias, sx, sw, sx_inv, sw_inv, amax_x, amax_w
|
||||
)
|
||||
|
||||
|
||||
def linear_backward_scaled(g, x, w, masks, sg, sw, sx, sg_inv, sw_inv, sx_inv, amax_g):
|
||||
"""dX = g @ W, dW = g^T @ X, dB = sum(g) with per-tensor scales."""
|
||||
if not (
|
||||
g.dtype == torch.bfloat16
|
||||
and x.dtype == torch.bfloat16
|
||||
and w.dtype == torch.bfloat16
|
||||
):
|
||||
raise TypeError(
|
||||
f"fp8 backward requires bf16 inputs, got {g.dtype}/{x.dtype}/{w.dtype}"
|
||||
)
|
||||
return _mod().fp8_linear_backward_scaled(
|
||||
g, x, w, masks, sg, sw, sx, sg_inv, sw_inv, sx_inv, amax_g
|
||||
)
|
||||
@@ -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.backend.rotary.apply_rotary_emb``.
|
||||
|
||||
Layout: x is packed [tokens, n_heads, head_dim] or dense
|
||||
[batch, seq_len, n_heads, head_dim]. ``freqs_cis`` has matching token axes.
|
||||
"""
|
||||
|
||||
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: packed 3D or dense 4D bf16 tensor.
|
||||
freqs_cis: matching token axes followed by [head_dim/2, 2].
|
||||
|
||||
Returns:
|
||||
Tensor with the same shape as ``x``.
|
||||
"""
|
||||
_check_available()
|
||||
if not x.is_contiguous():
|
||||
x = x.contiguous()
|
||||
if not freqs_cis.is_contiguous():
|
||||
freqs_cis = freqs_cis.contiguous()
|
||||
return _modules["rotary_emb"].rotary_emb(x, freqs_cis)
|
||||
+41
-35
@@ -13,41 +13,63 @@ from typing import (
|
||||
Type,
|
||||
TypeVar,
|
||||
Union,
|
||||
get_args,
|
||||
get_origin,
|
||||
)
|
||||
from typing import get_args as _get_args
|
||||
from typing import get_origin as _get_origin
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
def _resolve_type(
|
||||
def _resolve_base_type(
|
||||
arg: Union[Type, str, ForwardRef], factory_cls: type
|
||||
) -> Optional[Type]:
|
||||
"""Resolve a generic type-arg (str forward-ref, ForwardRef, or class)."""
|
||||
if not isinstance(arg, (str, ForwardRef)):
|
||||
"""Resolve the generic type-arg T to a concrete class.
|
||||
|
||||
- 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
|
||||
|
||||
name = arg if isinstance(arg, str) else arg.__forward_arg__
|
||||
if name == factory_cls.__name__:
|
||||
return factory_cls
|
||||
if isinstance(arg, str):
|
||||
name = arg
|
||||
elif isinstance(arg, ForwardRef):
|
||||
name = arg.__forward_arg__
|
||||
else:
|
||||
return None
|
||||
|
||||
mod = sys.modules.get(factory_cls.__module__)
|
||||
if mod is 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]):
|
||||
"""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]):
|
||||
pass
|
||||
|
||||
Register components with the ``register`` decorator::
|
||||
|
||||
@MyFactory.register("custom")
|
||||
class CustomComponent(MyBase):
|
||||
...
|
||||
@@ -64,13 +86,10 @@ class BaseFactory(ABC, Generic[T]):
|
||||
def __init_subclass__(cls, **kwargs):
|
||||
super().__init_subclass__(**kwargs)
|
||||
for orig_base in getattr(cls, "__orig_bases__", ()):
|
||||
if _get_origin(orig_base) is BaseFactory:
|
||||
(arg,) = _get_args(orig_base)
|
||||
if get_origin(orig_base) is BaseFactory:
|
||||
(arg,) = get_args(orig_base)
|
||||
cls._entries = {}
|
||||
try:
|
||||
cls._component_base = _resolve_type(arg, cls)
|
||||
except Exception:
|
||||
cls._component_base = None
|
||||
cls._component_base = _resolve_base_type(arg, cls)
|
||||
return
|
||||
|
||||
@classmethod
|
||||
@@ -82,7 +101,7 @@ class BaseFactory(ABC, Generic[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:
|
||||
raise ValueError(f"Component '{name}' is already registered")
|
||||
cls._entries[name] = component_cls
|
||||
@@ -95,12 +114,11 @@ class BaseFactory(ABC, Generic[T]):
|
||||
"""Create a component instance by name, filtering kwargs to match
|
||||
the component's ``__init__`` signature.
|
||||
"""
|
||||
entry = cls._entries.get(name)
|
||||
if entry is None:
|
||||
component_cls = cls._entries.get(name)
|
||||
if component_cls is None:
|
||||
raise ValueError(
|
||||
f"Unknown component: '{name}'. Supported types: {sorted(cls._entries)}"
|
||||
)
|
||||
component_cls = entry
|
||||
sig = inspect.signature(component_cls.__init__)
|
||||
has_var_kwargs = any(
|
||||
p.kind == inspect.Parameter.VAR_KEYWORD for p in sig.parameters.values()
|
||||
@@ -114,18 +132,6 @@ class BaseFactory(ABC, Generic[T]):
|
||||
kwargs = {k: v for k, v in kwargs.items() if k in valid}
|
||||
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
|
||||
def get_component_class(cls, name: str) -> Type[T]:
|
||||
"""Get the registered component class without instantiating it."""
|
||||
|
||||
@@ -1,15 +1,29 @@
|
||||
"""Inference module for continuous batching.
|
||||
|
||||
Layers:
|
||||
- core/: Core inference loop (cache, executor, scheduler, task)
|
||||
- api/: HTTP orchestration (ProtocolHandler, server)
|
||||
- protocols/: Response builders (OpenAI, Anthropic)
|
||||
- transport/: SSE transport utilities
|
||||
- engine.py: Facade (InferenceEngine), Value Object (GenerationRequest)
|
||||
- sample.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy, FrequencyPenaltyStrategy)
|
||||
Subpackages:
|
||||
- cache/: KV cache (buffers, strategies, pool)
|
||||
- runtime/: Execution + sampling (executor, CUDA graph, sampling strategies)
|
||||
- task/: Request lifecycle + performance metrics
|
||||
- network/: HTTP protocol handling (server, protocol, OpenAI/Anthropic builders)
|
||||
|
||||
Modules:
|
||||
- scheduler.py: Continuous batching loop
|
||||
- workspace.py: Pre-allocated GPU buffers
|
||||
- engine.py: Facade (InferenceEngine)
|
||||
"""
|
||||
|
||||
from astrai.inference.api import (
|
||||
from astrai.inference.cache import (
|
||||
Allocator,
|
||||
KVCache,
|
||||
KVStorage,
|
||||
PagePool,
|
||||
RadixCache,
|
||||
ReqToTokenPool,
|
||||
TaskCacheManager,
|
||||
page_hash,
|
||||
)
|
||||
from astrai.inference.engine import InferenceEngine
|
||||
from astrai.inference.network import (
|
||||
AnthropicMessage,
|
||||
BaseToolParser,
|
||||
ChatCompletionRequest,
|
||||
@@ -25,30 +39,10 @@ from astrai.inference.api import (
|
||||
get_app,
|
||||
run_server,
|
||||
)
|
||||
from astrai.inference.api.anthropic import AnthropicResponseBuilder
|
||||
from astrai.inference.api.openai import OpenAIResponseBuilder
|
||||
from astrai.inference.core import (
|
||||
STOP,
|
||||
Allocator,
|
||||
CacheView,
|
||||
ContiguousCache,
|
||||
ContiguousCacheView,
|
||||
Executor,
|
||||
InferenceScheduler,
|
||||
KVCache,
|
||||
PageCache,
|
||||
PageCacheView,
|
||||
PagePool,
|
||||
PrefixCache,
|
||||
Storage,
|
||||
Task,
|
||||
TaskManager,
|
||||
TaskStatus,
|
||||
TaskTable,
|
||||
page_hash,
|
||||
)
|
||||
from astrai.inference.engine import GenerationRequest, InferenceEngine
|
||||
from astrai.inference.sample import (
|
||||
from astrai.inference.network.anthropic import AnthropicResponseBuilder
|
||||
from astrai.inference.network.openai import OpenAIResponseBuilder
|
||||
from astrai.inference.runtime.executor import Executor
|
||||
from astrai.inference.runtime.sample import (
|
||||
BaseSamplingStrategy,
|
||||
FrequencyPenaltyStrategy,
|
||||
SamplingPipeline,
|
||||
@@ -57,10 +51,11 @@ from astrai.inference.sample import (
|
||||
TopPStrategy,
|
||||
sample,
|
||||
)
|
||||
from astrai.inference.scheduler import InferenceScheduler
|
||||
from astrai.inference.task import STOP, Task, TaskManager, TaskStatus
|
||||
|
||||
__all__ = [
|
||||
"InferenceEngine",
|
||||
"GenerationRequest",
|
||||
"InferenceScheduler",
|
||||
"Executor",
|
||||
"STOP",
|
||||
@@ -68,16 +63,12 @@ __all__ = [
|
||||
"TaskManager",
|
||||
"TaskStatus",
|
||||
"Allocator",
|
||||
"CacheView",
|
||||
"KVCache",
|
||||
"ContiguousCache",
|
||||
"ContiguousCacheView",
|
||||
"PageCache",
|
||||
"PageCacheView",
|
||||
"KVStorage",
|
||||
"PagePool",
|
||||
"PrefixCache",
|
||||
"Storage",
|
||||
"TaskTable",
|
||||
"RadixCache",
|
||||
"ReqToTokenPool",
|
||||
"TaskCacheManager",
|
||||
"page_hash",
|
||||
"sample",
|
||||
"BaseSamplingStrategy",
|
||||
|
||||
Vendored
+27
@@ -0,0 +1,27 @@
|
||||
"""KV cache subsystem: buffers, strategies, pool management."""
|
||||
|
||||
from astrai.inference.cache.buffer import KVCache, KVStorage, ReqToTokenPool
|
||||
from astrai.inference.cache.pool import PagePool, TaskCacheManager, page_hash
|
||||
from astrai.inference.cache.strategy import (
|
||||
AllocationStrategy,
|
||||
Allocator,
|
||||
ContiguousStrategy,
|
||||
PagedStrategy,
|
||||
RadixCache,
|
||||
TaskCacheState,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"KVCache",
|
||||
"KVStorage",
|
||||
"ReqToTokenPool",
|
||||
"Allocator",
|
||||
"RadixCache",
|
||||
"TaskCacheState",
|
||||
"AllocationStrategy",
|
||||
"ContiguousStrategy",
|
||||
"PagedStrategy",
|
||||
"PagePool",
|
||||
"TaskCacheManager",
|
||||
"page_hash",
|
||||
]
|
||||
Vendored
+104
@@ -0,0 +1,104 @@
|
||||
"""Physical KV cache buffers.
|
||||
|
||||
Layer 1 — ``KVStorage``: flat token-level K/V GPU buffers [n_layers, size, n_kv_heads, head_dim]
|
||||
Layer 2 — ``ReqToTokenPool``: index table [req_idx, pos] → physical token slot
|
||||
Layer 3 — ``KVCache``: pure dataclass passed to the model for direct buffer access
|
||||
|
||||
These classes have no knowledge of tasks, allocation policies, or scheduling.
|
||||
They are the "dumb" physical storage layer.
|
||||
"""
|
||||
|
||||
import threading
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
class ReqToTokenPool:
|
||||
"""Maps [req_idx, pos] → physical token slot in KV storage.
|
||||
|
||||
Each row is one request; each column is a sequence position. The value
|
||||
at [req_idx, pos] is the flat index into the KV storage buffers.
|
||||
"""
|
||||
|
||||
def __init__(self, size: int, max_context_len: int, device: torch.device):
|
||||
self.size = size
|
||||
self.max_context_len = max_context_len
|
||||
self.req_to_token = torch.zeros(
|
||||
(size, max_context_len), dtype=torch.int32, device=device
|
||||
)
|
||||
self.free_slots = list(range(size))
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def alloc(self, num_reqs: int) -> Optional[List[int]]:
|
||||
with self._lock:
|
||||
if num_reqs > len(self.free_slots):
|
||||
return None
|
||||
slots = self.free_slots[:num_reqs]
|
||||
self.free_slots = self.free_slots[num_reqs:]
|
||||
return slots
|
||||
|
||||
def free(self, req_indices: List[int]):
|
||||
with self._lock:
|
||||
self.free_slots.extend(req_indices)
|
||||
|
||||
def write(self, indices, values):
|
||||
self.req_to_token[indices] = values
|
||||
|
||||
|
||||
class KVStorage:
|
||||
"""Token-level KV cache storage.
|
||||
|
||||
Buffers: ``[n_layers, size, n_kv_heads, head_dim]``. Each token occupies
|
||||
one slot indexed by ``ReqToTokenPool``.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size: int,
|
||||
n_layers: int,
|
||||
n_kv_heads: int,
|
||||
head_dim: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
):
|
||||
self.size = size
|
||||
self.k_buffer = torch.empty(
|
||||
(n_layers, size, n_kv_heads, head_dim), device=device, dtype=dtype
|
||||
)
|
||||
self.v_buffer = torch.empty(
|
||||
(n_layers, size, n_kv_heads, head_dim), device=device, dtype=dtype
|
||||
)
|
||||
|
||||
def get_key_buffer(self, layer_id: int) -> Tensor:
|
||||
return self.k_buffer[layer_id]
|
||||
|
||||
def get_value_buffer(self, layer_id: int) -> Tensor:
|
||||
return self.v_buffer[layer_id]
|
||||
|
||||
def set_kv_buffer(self, layer_id: int, loc: Tensor, k: Tensor, v: Tensor) -> None:
|
||||
self.k_buffer[layer_id, loc] = k
|
||||
self.v_buffer[layer_id, loc] = v
|
||||
|
||||
|
||||
@dataclass
|
||||
class KVCache:
|
||||
"""Pure data struct passed to model for KV cache I/O.
|
||||
|
||||
The attention layer does raw buffer indexing — no methods, no abstraction.
|
||||
"""
|
||||
|
||||
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
|
||||
kv_indptr: Optional[Tensor] = None
|
||||
qo_indptr: Optional[Tensor] = None
|
||||
decode_o_part: Optional[Tensor] = None
|
||||
decode_ml_part: Optional[Tensor] = None
|
||||
decode_out: Optional[Tensor] = None
|
||||
Vendored
+364
@@ -0,0 +1,364 @@
|
||||
"""KV cache orchestration: PagePool + TaskCacheManager.
|
||||
|
||||
PagePool owns the physical buffers (``KVStorage`` + ``ReqToTokenPool``)
|
||||
and wires them to an allocation strategy. It assembles the ``KVCache``
|
||||
dataclass passed to the model forward.
|
||||
|
||||
TaskCacheManager owns the ``task_id`` → ``TaskCacheState`` mapping and
|
||||
delegates physical slot allocation to the strategy, and KV bind to the pool.
|
||||
|
||||
See ``cache_buffer.py`` for the raw buffer primitives and ``cache_strategy.py``
|
||||
for the allocation policies.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.inference.cache.buffer import KVCache, KVStorage, ReqToTokenPool
|
||||
from astrai.inference.cache.strategy import (
|
||||
AllocationStrategy,
|
||||
Allocator,
|
||||
ContiguousStrategy,
|
||||
PagedStrategy,
|
||||
RadixCache,
|
||||
TaskCacheState,
|
||||
)
|
||||
from astrai.inference.workspace import InferenceWorkspace
|
||||
|
||||
# Re-export everything so existing ``from astrai.inference.cache import ...``
|
||||
# continues to work unchanged after the file split.
|
||||
__all__ = [
|
||||
"KVCache",
|
||||
"KVStorage",
|
||||
"ReqToTokenPool",
|
||||
"Allocator",
|
||||
"RadixCache",
|
||||
"AllocationStrategy",
|
||||
"ContiguousStrategy",
|
||||
"PagedStrategy",
|
||||
"PagePool",
|
||||
"TaskCacheManager",
|
||||
"TaskCacheState",
|
||||
"page_hash",
|
||||
]
|
||||
|
||||
# ---- helpers ----
|
||||
|
||||
|
||||
def page_hash(
|
||||
token_ids: List[int], page_idx: int, page_size: int, parent_hash: int = 0
|
||||
) -> int:
|
||||
start = page_idx * page_size
|
||||
end = min(start + page_size, len(token_ids))
|
||||
h = parent_hash
|
||||
for i in range(start, end):
|
||||
h = (h * 31 + token_ids[i]) & 0xFFFFFFFFFFFFFFFF
|
||||
return h
|
||||
|
||||
|
||||
def _is_steady_increment(
|
||||
prev_sig: Optional[tuple],
|
||||
prev_vals: Optional[List[int]],
|
||||
cur_sig: tuple,
|
||||
cur_vals: List[int],
|
||||
) -> bool:
|
||||
return (
|
||||
prev_sig is not None
|
||||
and prev_vals is not None
|
||||
and prev_sig == cur_sig
|
||||
and len(prev_vals) == len(cur_vals)
|
||||
and all(c == p + 1 for c, p in zip(cur_vals, prev_vals))
|
||||
)
|
||||
|
||||
|
||||
# ---- task-scoped bind state ----
|
||||
@dataclass
|
||||
class _BindState:
|
||||
"""Cached bind metadata for steady-state decode increment detection."""
|
||||
|
||||
sig: tuple
|
||||
seq_lens: List[int]
|
||||
|
||||
|
||||
# ---- pool + manager ----
|
||||
|
||||
|
||||
class PagePool:
|
||||
"""Physical KV cache: buffers + req-to-token table + allocation strategy + bind.
|
||||
|
||||
Does not know about tasks — task lifecycle is managed by
|
||||
:class:`TaskCacheManager`, which holds a reference to this pool.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
n_layers: int,
|
||||
n_kv_heads: int,
|
||||
head_dim: int,
|
||||
max_batch_size: int,
|
||||
max_seq_len: int,
|
||||
device: torch.device,
|
||||
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.device = device
|
||||
self.dtype = dtype
|
||||
self.n_layers = n_layers
|
||||
self.n_kv_heads = n_kv_heads
|
||||
self.head_dim = head_dim
|
||||
|
||||
self.contiguous = n_tokens is None
|
||||
self.n_tokens = max_batch_size * max_seq_len if self.contiguous else n_tokens
|
||||
if self.n_tokens > torch.iinfo(torch.int32).max:
|
||||
raise ValueError("KV cache token count exceeds the int32 slot index limit")
|
||||
|
||||
self._storage = KVStorage(
|
||||
self.n_tokens, n_layers, n_kv_heads, head_dim, device, dtype
|
||||
)
|
||||
self._req_pool = ReqToTokenPool(max_batch_size, max_seq_len, device)
|
||||
|
||||
if self.contiguous:
|
||||
for i in range(max_batch_size):
|
||||
self._req_pool.req_to_token[i] = torch.arange(
|
||||
i * max_seq_len,
|
||||
(i + 1) * max_seq_len,
|
||||
dtype=torch.int32,
|
||||
device=device,
|
||||
)
|
||||
self._strategy: AllocationStrategy = ContiguousStrategy()
|
||||
else:
|
||||
n_pages = self.n_tokens // page_size
|
||||
alloc = Allocator(n_pages)
|
||||
prefix = RadixCache(page_size) if page_size > 1 else None
|
||||
if prefix is not None:
|
||||
alloc.on_evict = prefix.evict
|
||||
self._strategy = PagedStrategy(
|
||||
alloc, prefix, page_size, self._req_pool, device
|
||||
)
|
||||
|
||||
@property
|
||||
def strategy(self) -> AllocationStrategy:
|
||||
return self._strategy
|
||||
|
||||
@property
|
||||
def req_pool(self) -> ReqToTokenPool:
|
||||
return self._req_pool
|
||||
|
||||
def bind_tasks(
|
||||
self,
|
||||
req_indices: List[int],
|
||||
seq_lens: List[int],
|
||||
workspace: InferenceWorkspace,
|
||||
device: Optional[torch.device] = None,
|
||||
start_pos: Optional[int] = None,
|
||||
incremental: bool = False,
|
||||
) -> KVCache:
|
||||
"""Assemble the ``KVCache`` metadata for a batch of tasks.
|
||||
|
||||
Args:
|
||||
req_indices: request slot indices (from ``ReqToTokenPool``).
|
||||
seq_lens: current sequence length per task.
|
||||
workspace: pre-allocated fixed-shape buffers (CUDA-graph safe).
|
||||
start_pos: if set, produce **prefill** cache (full q_len range).
|
||||
If ``None``, produce **decode** cache (last position).
|
||||
incremental: if ``True``, reuse workspace state from previous step
|
||||
by incrementing counters in-place (decode hot path).
|
||||
|
||||
Returns:
|
||||
``KVCache`` dataclass with the correct output shapes for the
|
||||
attention backend (prefill: ``[B, q_len]``, decode: ``[B, 1]``).
|
||||
"""
|
||||
if device is None:
|
||||
device = workspace.device
|
||||
b = len(req_indices)
|
||||
|
||||
rpi_buf = workspace.req_pool_indices
|
||||
sl_buf = workspace.seq_lens
|
||||
kvp_buf = workspace.kv_indptr
|
||||
inc_buf = workspace.inc
|
||||
ocl_buf = workspace.out_cache_loc
|
||||
|
||||
if incremental:
|
||||
sl_buf[:b] += 1
|
||||
kvp_buf[: b + 1] += inc_buf[: b + 1]
|
||||
else:
|
||||
rpi_buf[:b].copy_(
|
||||
torch.tensor(req_indices, dtype=torch.int32, device=device)
|
||||
)
|
||||
sl_buf[:b].copy_(torch.tensor(seq_lens, dtype=torch.long, device=device))
|
||||
kvp_buf[: b + 1].zero_()
|
||||
kvp_buf[1 : b + 1] = sl_buf[:b].cumsum(0).to(torch.int32)
|
||||
|
||||
req_pool_indices = rpi_buf[:b]
|
||||
seq_lens_t = sl_buf[:b]
|
||||
kv_indptr = kvp_buf[: b + 1]
|
||||
|
||||
if start_pos is not None:
|
||||
# Packed prefill concatenates each request's query tokens.
|
||||
q_lens = [seq_len - start_pos for seq_len in seq_lens]
|
||||
if any(q_len <= 0 for q_len in q_lens):
|
||||
raise ValueError("prefill sequence lengths must exceed start_pos")
|
||||
out_cache_loc = torch.cat(
|
||||
[
|
||||
self._req_pool.req_to_token[
|
||||
req_pool_indices[i], start_pos : seq_lens[i]
|
||||
]
|
||||
for i in range(b)
|
||||
]
|
||||
)
|
||||
workspace.qo_indptr[: b + 1].zero_()
|
||||
workspace.qo_indptr[1 : b + 1].copy_(
|
||||
torch.tensor(q_lens, dtype=torch.int32, device=device).cumsum(0)
|
||||
)
|
||||
qo_indptr = workspace.qo_indptr[: b + 1]
|
||||
decode_o_part = decode_ml_part = decode_out = None
|
||||
else:
|
||||
# ---- decode: out_cache_loc is a single column (last position) ----
|
||||
write_pos = seq_lens_t - 1
|
||||
loc = self._req_pool.req_to_token[req_pool_indices, write_pos].unsqueeze(-1)
|
||||
ocl_buf[:b].copy_(loc)
|
||||
out_cache_loc = ocl_buf[:b].reshape(-1)
|
||||
workspace.qo_indptr[: b + 1].copy_(inc_buf[: b + 1])
|
||||
qo_indptr = workspace.qo_indptr[: b + 1]
|
||||
decode_o_part = getattr(workspace, "decode_o_part", None)
|
||||
decode_ml_part = getattr(workspace, "decode_ml_part", None)
|
||||
decode_out = getattr(workspace, "decode_out", 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),
|
||||
kv_indptr=kv_indptr,
|
||||
qo_indptr=qo_indptr,
|
||||
decode_o_part=decode_o_part,
|
||||
decode_ml_part=decode_ml_part,
|
||||
decode_out=decode_out,
|
||||
)
|
||||
|
||||
|
||||
class TaskCacheManager:
|
||||
"""Task ↔ KV slot lifecycle manager.
|
||||
|
||||
Sole owner of ``task_id → TaskCacheState``. Delegates physical slot
|
||||
allocation to the strategy (via ``pool.strategy``) and KV bind to
|
||||
``pool.bind_tasks()``.
|
||||
|
||||
Usage::
|
||||
|
||||
pool = PagePool(...)
|
||||
mgr = TaskCacheManager(pool)
|
||||
mgr.task_alloc("req_1", [101, 202, 303])
|
||||
...
|
||||
kv = mgr.bind(["req_1"], workspace)
|
||||
"""
|
||||
|
||||
def __init__(self, pool: PagePool):
|
||||
self._pool = pool
|
||||
self._strategy = pool.strategy
|
||||
self._req_pool = pool.req_pool
|
||||
self._max_seq_len = pool.max_seq_len
|
||||
self._states: Dict[str, TaskCacheState] = {}
|
||||
self._bind_state: Optional[_BindState] = None
|
||||
self._bind_was_steady = False
|
||||
|
||||
# -- public task lifecycle --
|
||||
|
||||
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
|
||||
self._bind_state = None
|
||||
req_slots = self._req_pool.alloc(1)
|
||||
if req_slots is None:
|
||||
return False
|
||||
state = TaskCacheState(req_idx=req_slots[0])
|
||||
self._states[task_id] = state
|
||||
if not self._strategy.alloc(state, prompt_ids):
|
||||
self._rollback(state, task_id)
|
||||
return False
|
||||
self._strategy.write_indices(state, prompt_ids)
|
||||
state.length = len(prompt_ids)
|
||||
return True
|
||||
|
||||
def task_free(self, task_id: str):
|
||||
self._bind_state = None
|
||||
state = self._states.pop(task_id, None)
|
||||
if state is None:
|
||||
return
|
||||
self._strategy.free(state)
|
||||
self._req_pool.free([state.req_idx])
|
||||
|
||||
def task_extend(self, task_id: str, pos: int) -> bool:
|
||||
state = self._states.get(task_id)
|
||||
if state is None or pos >= self._max_seq_len:
|
||||
return False
|
||||
if not self._strategy.extend(state, pos):
|
||||
return False
|
||||
state.length = pos + 1
|
||||
return True
|
||||
|
||||
def task_cached(self, task_id: str) -> int:
|
||||
state = self._states.get(task_id)
|
||||
return state.cached if state is not None else 0
|
||||
|
||||
def task_record_hashes(
|
||||
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
|
||||
):
|
||||
state = self._states.get(task_id)
|
||||
if state is not None:
|
||||
self._strategy.record_hashes(state, prompt_ids, start_logical_page)
|
||||
|
||||
@staticmethod
|
||||
def task_cacheable_ids(task_id: str, prompt_ids: List[int], output_ids: List[int]):
|
||||
return list(prompt_ids) + list(output_ids[:-1])
|
||||
|
||||
# -- bind (assemble KVCache for the model forward) --
|
||||
|
||||
def bind(
|
||||
self,
|
||||
task_ids: List[str],
|
||||
workspace: InferenceWorkspace,
|
||||
device: Optional[torch.device] = None,
|
||||
start_pos: Optional[int] = None,
|
||||
) -> KVCache:
|
||||
"""Build ``KVCache`` for an ordered list of task IDs."""
|
||||
states = [self._states[tid] for tid in task_ids]
|
||||
req_indices = [s.req_idx for s in states]
|
||||
seq_lens = [s.length for s in states]
|
||||
sig = tuple(req_indices)
|
||||
|
||||
prev = self._bind_state
|
||||
incremental = (
|
||||
start_pos is None
|
||||
and prev is not None
|
||||
and _is_steady_increment(prev.sig, prev.seq_lens, sig, seq_lens)
|
||||
)
|
||||
self._bind_state = _BindState(sig, list(seq_lens))
|
||||
self._bind_was_steady = incremental
|
||||
|
||||
return self._pool.bind_tasks(
|
||||
req_indices,
|
||||
seq_lens,
|
||||
workspace,
|
||||
device=device,
|
||||
start_pos=start_pos,
|
||||
incremental=incremental,
|
||||
)
|
||||
|
||||
@property
|
||||
def bind_was_steady(self) -> bool:
|
||||
return self._bind_was_steady
|
||||
|
||||
# -- internals --
|
||||
|
||||
def _rollback(self, state: TaskCacheState, task_id: str):
|
||||
self._strategy.free(state)
|
||||
self._req_pool.free([state.req_idx])
|
||||
self._states.pop(task_id, None)
|
||||
Vendored
+320
@@ -0,0 +1,320 @@
|
||||
"""KV cache allocation layer.
|
||||
|
||||
Encapsulates the physical slot allocation policy, isolated from GPU buffers
|
||||
and task lifecycle management.
|
||||
|
||||
- ``TaskCacheState``: data contract between strategy and manager (per-task slot state)
|
||||
- ``Allocator``: bitmask-based page allocator with LRU eviction
|
||||
- ``RadixCache``: page-granular prefix index (exact token match)
|
||||
- ``AllocationStrategy``: ABC for physical slot allocation
|
||||
- ``ContiguousStrategy``: statically partitioned, no dynamic allocation
|
||||
- ``PagedStrategy``: dynamic paged allocation from a shared pool
|
||||
"""
|
||||
|
||||
import threading
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable, Dict, List, Optional, OrderedDict
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.inference.cache.buffer import ReqToTokenPool
|
||||
|
||||
# ---- data contract: per-task slot state ----
|
||||
|
||||
|
||||
@dataclass
|
||||
class TaskCacheState:
|
||||
"""Per-task cache allocation state.
|
||||
|
||||
Co-locates all task-owned cache metadata so the alloc/free/extend
|
||||
lifecycle is atomic. Owned by ``TaskCacheManager``, consumed by
|
||||
every ``AllocationStrategy`` method.
|
||||
"""
|
||||
|
||||
req_idx: int
|
||||
length: int = 0
|
||||
cached: int = 0
|
||||
pages: List[int] = field(default_factory=list)
|
||||
|
||||
|
||||
# ---- allocation primitives ----
|
||||
|
||||
|
||||
class Allocator:
|
||||
"""Bitmask-based page allocator with ref-counting and LRU eviction."""
|
||||
|
||||
def __init__(self, n_pages: int):
|
||||
self._free_mask = (1 << n_pages) - 1
|
||||
self._refs: List[int] = [0] * n_pages
|
||||
self._lru: OrderedDict[int, None] = OrderedDict()
|
||||
self.on_evict: Optional[Callable[[int], None]] = None
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def alloc(self) -> int:
|
||||
with self._lock:
|
||||
if self._free_mask:
|
||||
lsb = self._free_mask & -self._free_mask
|
||||
idx = lsb.bit_length() - 1
|
||||
self._free_mask ^= lsb
|
||||
self._refs[idx] = 1
|
||||
return idx
|
||||
if self._lru:
|
||||
idx, _ = self._lru.popitem(last=False)
|
||||
if self.on_evict:
|
||||
self.on_evict(idx)
|
||||
self._refs[idx] = 1
|
||||
self._free_mask &= ~(1 << idx)
|
||||
return idx
|
||||
return -1
|
||||
|
||||
def free(self, idx: int, keep_cached: bool = False):
|
||||
with self._lock:
|
||||
self._refs[idx] -= 1
|
||||
if self._refs[idx] == 0:
|
||||
if keep_cached:
|
||||
self._lru[idx] = None
|
||||
else:
|
||||
self._free_mask |= 1 << idx
|
||||
|
||||
def inc_ref(self, idx: int):
|
||||
with self._lock:
|
||||
self._refs[idx] += 1
|
||||
self._lru.pop(idx, None)
|
||||
|
||||
def ref_count(self, idx: int) -> int:
|
||||
with self._lock:
|
||||
return self._refs[idx]
|
||||
|
||||
def touch(self, idx: int):
|
||||
with self._lock:
|
||||
if idx in self._lru:
|
||||
self._lru.move_to_end(idx)
|
||||
|
||||
|
||||
class RadixNode:
|
||||
"""A page-aligned edge in the CPU-side prefix radix trie."""
|
||||
|
||||
__slots__ = ("parent", "children", "page_idx", "tokens", "lock_ref")
|
||||
|
||||
def __init__(self, parent=None, tokens=(), page_idx=None):
|
||||
self.parent = parent
|
||||
self.children: Dict[tuple, "RadixNode"] = {}
|
||||
self.page_idx = page_idx
|
||||
self.tokens = tuple(tokens)
|
||||
self.lock_ref = 0
|
||||
|
||||
|
||||
class RadixCache:
|
||||
"""Page-granular radix prefix index with exact token matching."""
|
||||
|
||||
def __init__(self, page_size: int):
|
||||
self._page_size = page_size
|
||||
self._root = RadixNode()
|
||||
self._page_to_node: Dict[int, RadixNode] = {}
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def evict(self, idx: int):
|
||||
with self._lock:
|
||||
node = self._page_to_node.pop(idx, None)
|
||||
if node is None:
|
||||
return
|
||||
node.page_idx = None
|
||||
parent = node.parent
|
||||
if parent is not None:
|
||||
parent.children.pop(node.tokens, None)
|
||||
|
||||
def has_page(self, idx: int) -> bool:
|
||||
with self._lock:
|
||||
return idx in self._page_to_node
|
||||
|
||||
def lookup(self, token_ids: List[int]) -> List[int]:
|
||||
with self._lock:
|
||||
full_pages = len(token_ids) // self._page_size
|
||||
hits: List[int] = []
|
||||
node = self._root
|
||||
for i in range(full_pages):
|
||||
start = i * self._page_size
|
||||
page_tokens = tuple(token_ids[start : start + self._page_size])
|
||||
child = node.children.get(page_tokens)
|
||||
if child is None or child.page_idx is None:
|
||||
break
|
||||
hits.append(child.page_idx)
|
||||
node = child
|
||||
return hits
|
||||
|
||||
def record(self, page_idx: int, token_ids: List[int], logical_page_idx: int):
|
||||
with self._lock:
|
||||
full_pages = len(token_ids) // self._page_size
|
||||
if logical_page_idx >= full_pages:
|
||||
return
|
||||
old = self._page_to_node.pop(page_idx, None)
|
||||
if old is not None and old.parent is not None:
|
||||
old.parent.children.pop(old.tokens, None)
|
||||
|
||||
node = self._root
|
||||
for i in range(logical_page_idx + 1):
|
||||
start = i * self._page_size
|
||||
page_tokens = tuple(token_ids[start : start + self._page_size])
|
||||
child = node.children.get(page_tokens)
|
||||
if child is None:
|
||||
child = RadixNode(node, page_tokens)
|
||||
node.children[page_tokens] = child
|
||||
node = child
|
||||
if node.page_idx is not None and node.page_idx != page_idx:
|
||||
replaced = node.page_idx
|
||||
self._page_to_node.pop(replaced, None)
|
||||
node.page_idx = page_idx
|
||||
self._page_to_node[page_idx] = node
|
||||
|
||||
def release(self, pages: List[int]) -> None:
|
||||
with self._lock:
|
||||
for page_idx in pages:
|
||||
node = self._page_to_node.get(page_idx)
|
||||
if node is not None and node.lock_ref:
|
||||
node.lock_ref -= 1
|
||||
|
||||
|
||||
class AllocationStrategy(ABC):
|
||||
"""Physical slot allocation policy.
|
||||
|
||||
Subclasses implement the actual allocation semantics. This ABC declares
|
||||
the contract; there are no default implementations.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def alloc(self, state: TaskCacheState, prompt_ids: List[int]) -> bool: ...
|
||||
|
||||
@abstractmethod
|
||||
def free(self, state: TaskCacheState) -> None: ...
|
||||
|
||||
@abstractmethod
|
||||
def extend(self, state: TaskCacheState, pos: int) -> bool: ...
|
||||
|
||||
@abstractmethod
|
||||
def write_indices(self, state: TaskCacheState, prompt_ids: List[int]) -> None: ...
|
||||
|
||||
@abstractmethod
|
||||
def record_hashes(
|
||||
self,
|
||||
state: TaskCacheState,
|
||||
prompt_ids: List[int],
|
||||
start: int,
|
||||
) -> None: ...
|
||||
|
||||
|
||||
class ContiguousStrategy(AllocationStrategy):
|
||||
"""Static contiguous allocation: slots are pre-assigned at pool init.
|
||||
|
||||
No dynamic allocation or prefix caching. All operations are no-ops
|
||||
because ``ReqToTokenPool`` is pre-filled with contiguous ranges.
|
||||
"""
|
||||
|
||||
def alloc(self, state: TaskCacheState, prompt_ids: List[int]) -> bool:
|
||||
return True
|
||||
|
||||
def free(self, state: TaskCacheState) -> None:
|
||||
pass
|
||||
|
||||
def extend(self, state: TaskCacheState, pos: int) -> bool:
|
||||
return True
|
||||
|
||||
def write_indices(self, state: TaskCacheState, prompt_ids: List[int]) -> None:
|
||||
pass
|
||||
|
||||
def record_hashes(
|
||||
self,
|
||||
state: TaskCacheState,
|
||||
prompt_ids: List[int],
|
||||
start: int,
|
||||
) -> None:
|
||||
pass
|
||||
|
||||
|
||||
class PagedStrategy(AllocationStrategy):
|
||||
"""Dynamic paged allocation from a shared bitmask pool.
|
||||
|
||||
``page_size`` is a parameter, not a separate strategy: at ``page_size=1``
|
||||
each allocated page *is* one token slot (``page * 1 + 0``), and prefix
|
||||
caching is simply disabled (``prefix=None``). The unified page formula
|
||||
``pages[page_idx] * page_size + offset`` holds for both.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
alloc: Allocator,
|
||||
prefix: Optional[RadixCache],
|
||||
page_size: int,
|
||||
req_pool: ReqToTokenPool,
|
||||
device,
|
||||
):
|
||||
self._alloc = alloc
|
||||
self._prefix = prefix
|
||||
self._page_size = page_size
|
||||
self._req_pool = req_pool
|
||||
self._device = device
|
||||
|
||||
def alloc(self, state: TaskCacheState, prompt_ids: List[int]) -> bool:
|
||||
if self._prefix is not None:
|
||||
hits = self._prefix.lookup(prompt_ids)
|
||||
state.cached = len(hits) * self._page_size
|
||||
for p in hits:
|
||||
self._alloc.inc_ref(p)
|
||||
state.pages = list(hits)
|
||||
|
||||
remaining = len(prompt_ids) - state.cached
|
||||
if remaining <= 0:
|
||||
return True
|
||||
n_new = (remaining + self._page_size - 1) // self._page_size
|
||||
for _ in range(n_new):
|
||||
p = self._alloc.alloc()
|
||||
if p < 0:
|
||||
return False
|
||||
state.pages.append(p)
|
||||
return True
|
||||
|
||||
def free(self, state: TaskCacheState) -> None:
|
||||
if self._prefix is not None:
|
||||
for p in state.pages:
|
||||
keep = self._prefix.has_page(p)
|
||||
self._alloc.free(p, keep_cached=keep)
|
||||
if not keep:
|
||||
self._prefix.evict(p)
|
||||
else:
|
||||
for p in state.pages:
|
||||
self._alloc.free(p)
|
||||
|
||||
def extend(self, state: TaskCacheState, pos: int) -> bool:
|
||||
page_idx = pos // self._page_size
|
||||
if page_idx >= len(state.pages):
|
||||
p = self._alloc.alloc()
|
||||
if p < 0:
|
||||
return False
|
||||
state.pages.append(p)
|
||||
offset = pos % self._page_size
|
||||
self._req_pool.req_to_token[state.req_idx, pos] = (
|
||||
state.pages[page_idx] * self._page_size + offset
|
||||
)
|
||||
return True
|
||||
|
||||
def write_indices(self, state: TaskCacheState, prompt_ids: List[int]) -> None:
|
||||
total = len(prompt_ids)
|
||||
for pos in range(total):
|
||||
page_idx = pos // self._page_size
|
||||
offset = pos % self._page_size
|
||||
if page_idx < len(state.pages):
|
||||
self._req_pool.req_to_token[state.req_idx, pos] = (
|
||||
state.pages[page_idx] * self._page_size + offset
|
||||
)
|
||||
|
||||
def record_hashes(
|
||||
self,
|
||||
state: TaskCacheState,
|
||||
prompt_ids: List[int],
|
||||
start: int,
|
||||
) -> None:
|
||||
if self._prefix is None:
|
||||
return
|
||||
full = len(prompt_ids) // self._page_size
|
||||
for i in range(start, min(full, len(state.pages))):
|
||||
self._prefix.record(state.pages[i], prompt_ids, i)
|
||||
@@ -1,40 +0,0 @@
|
||||
"""Inference core: cache, executor, scheduler, task management."""
|
||||
|
||||
from astrai.inference.core.cache import (
|
||||
Allocator,
|
||||
CacheView,
|
||||
ContiguousCache,
|
||||
ContiguousCacheView,
|
||||
KVCache,
|
||||
PageCache,
|
||||
PageCacheView,
|
||||
PagePool,
|
||||
PrefixCache,
|
||||
Storage,
|
||||
TaskTable,
|
||||
page_hash,
|
||||
)
|
||||
from astrai.inference.core.executor import Executor
|
||||
from astrai.inference.core.scheduler import InferenceScheduler
|
||||
from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
|
||||
|
||||
__all__ = [
|
||||
"Allocator",
|
||||
"CacheView",
|
||||
"KVCache",
|
||||
"ContiguousCache",
|
||||
"ContiguousCacheView",
|
||||
"PageCache",
|
||||
"PageCacheView",
|
||||
"PagePool",
|
||||
"PrefixCache",
|
||||
"Storage",
|
||||
"TaskTable",
|
||||
"page_hash",
|
||||
"Executor",
|
||||
"InferenceScheduler",
|
||||
"STOP",
|
||||
"Task",
|
||||
"TaskManager",
|
||||
"TaskStatus",
|
||||
]
|
||||
@@ -1,525 +0,0 @@
|
||||
import threading
|
||||
from abc import ABC, abstractmethod
|
||||
from collections import OrderedDict
|
||||
from typing import Callable, Dict, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
def page_hash(token_ids: List[int], page_idx: int, page_size: int) -> int:
|
||||
start = page_idx * page_size
|
||||
end = min(start + page_size, len(token_ids))
|
||||
h = 0
|
||||
for i in range(start, end):
|
||||
h = (h * 31 + token_ids[i]) & 0xFFFFFFFFFFFFFFFF
|
||||
return h
|
||||
|
||||
|
||||
class Allocator:
|
||||
"""Bitmask-based page allocator with ref-counting and LRU eviction."""
|
||||
|
||||
def __init__(self, n_pages: int):
|
||||
self._free_mask = (1 << n_pages) - 1
|
||||
self._refs: List[int] = [0] * n_pages
|
||||
self._lru: OrderedDict[int, None] = OrderedDict()
|
||||
self.on_evict: Optional[Callable[[int], None]] = None
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def alloc(self) -> int:
|
||||
with self._lock:
|
||||
if self._free_mask:
|
||||
lsb = self._free_mask & -self._free_mask
|
||||
idx = lsb.bit_length() - 1
|
||||
self._free_mask ^= lsb
|
||||
self._refs[idx] = 1
|
||||
return idx
|
||||
if self._lru:
|
||||
idx, _ = self._lru.popitem(last=False)
|
||||
if self.on_evict:
|
||||
self.on_evict(idx)
|
||||
self._refs[idx] = 1
|
||||
self._free_mask &= ~(1 << idx)
|
||||
return idx
|
||||
return -1
|
||||
|
||||
def free(self, idx: int, keep_cached: bool = False):
|
||||
with self._lock:
|
||||
self._refs[idx] -= 1
|
||||
if self._refs[idx] == 0:
|
||||
if keep_cached:
|
||||
self._lru[idx] = None
|
||||
else:
|
||||
self._free_mask |= 1 << idx
|
||||
|
||||
def inc_ref(self, idx: int):
|
||||
with self._lock:
|
||||
self._refs[idx] += 1
|
||||
self._lru.pop(idx, None)
|
||||
|
||||
def ref_count(self, idx: int) -> int:
|
||||
with self._lock:
|
||||
return self._refs[idx]
|
||||
|
||||
def touch(self, idx: int):
|
||||
with self._lock:
|
||||
if idx in self._lru:
|
||||
self._lru.move_to_end(idx)
|
||||
|
||||
|
||||
class PrefixCache:
|
||||
"""Hash-based prefix matching: maps page hashes to physical page indices."""
|
||||
|
||||
def __init__(self, page_size: int):
|
||||
self._page_size = page_size
|
||||
self._page_to_hash: Dict[int, int] = {}
|
||||
self._hash_to_page: Dict[int, int] = {}
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def evict(self, idx: int):
|
||||
with self._lock:
|
||||
h = self._page_to_hash.pop(idx, None)
|
||||
if h is not None:
|
||||
self._hash_to_page.pop(h, None)
|
||||
|
||||
def has_page(self, idx: int) -> bool:
|
||||
with self._lock:
|
||||
return idx in self._page_to_hash
|
||||
|
||||
def lookup(self, token_ids: List[int]) -> List[int]:
|
||||
with self._lock:
|
||||
full_pages = len(token_ids) // self._page_size
|
||||
hits: List[int] = []
|
||||
for i in range(full_pages):
|
||||
h = page_hash(token_ids, i, self._page_size)
|
||||
p = self._hash_to_page.get(h)
|
||||
if p is None:
|
||||
break
|
||||
hits.append(p)
|
||||
return hits
|
||||
|
||||
def record(self, page_idx: int, token_ids: List[int], logical_page_idx: int):
|
||||
with self._lock:
|
||||
h = page_hash(token_ids, logical_page_idx, self._page_size)
|
||||
old_h = self._page_to_hash.pop(page_idx, None)
|
||||
if old_h is not None:
|
||||
self._hash_to_page.pop(old_h, None)
|
||||
self._page_to_hash[page_idx] = h
|
||||
self._hash_to_page[h] = page_idx
|
||||
|
||||
|
||||
class PagePool:
|
||||
"""Orchestrates allocator (page management) and PrefixCache (content addressing)."""
|
||||
|
||||
def __init__(self, allocator: Allocator, prefix: PrefixCache):
|
||||
self._alloc = allocator
|
||||
self._prefix = prefix
|
||||
self._alloc.on_evict = prefix.evict
|
||||
|
||||
@property
|
||||
def allocator(self) -> Allocator:
|
||||
return self._alloc
|
||||
|
||||
@property
|
||||
def prefix(self) -> PrefixCache:
|
||||
return self._prefix
|
||||
|
||||
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()
|
||||
|
||||
def set(self, task_id: str, page_table: List[int], cached: int):
|
||||
with self._lock:
|
||||
self._pages[task_id] = page_table
|
||||
self._cached[task_id] = cached
|
||||
|
||||
def get(self, task_id: str) -> List[int]:
|
||||
with self._lock:
|
||||
return self._pages.get(task_id, [])
|
||||
|
||||
def get_cached(self, task_id: str) -> int:
|
||||
with self._lock:
|
||||
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:
|
||||
"""KV-cache tensor storage with paged write/gather."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
n_layers: int,
|
||||
n_pages: int,
|
||||
page_size: int,
|
||||
n_kv_heads: int,
|
||||
head_dim: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
):
|
||||
self.page_size = page_size
|
||||
self.k_cache = torch.empty(
|
||||
(n_layers, n_pages, page_size, n_kv_heads, head_dim),
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
self.v_cache = torch.empty(
|
||||
(n_layers, n_pages, page_size, n_kv_heads, head_dim),
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
def write(
|
||||
self,
|
||||
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(
|
||||
self, layer_id: int, page_table: Tensor, total_len: int
|
||||
) -> Tuple[Tensor, Tensor]:
|
||||
safe = page_table.clamp(min=0)
|
||||
k = self.k_cache[layer_id, safe]
|
||||
v = self.v_cache[layer_id, safe]
|
||||
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):
|
||||
"""Abstract view passed to attention layers for KV-cache I/O."""
|
||||
|
||||
@abstractmethod
|
||||
def write(self, layer_id: int, k: Tensor, v: Tensor): ...
|
||||
|
||||
@abstractmethod
|
||||
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]: ...
|
||||
|
||||
|
||||
class KVCache(ABC):
|
||||
"""Abstract KV-cache facade for scheduler/executor."""
|
||||
|
||||
@abstractmethod
|
||||
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool: ...
|
||||
|
||||
@abstractmethod
|
||||
def task_free(self, task_id: str): ...
|
||||
|
||||
@abstractmethod
|
||||
def task_extend(self, task_id: str, pos: int) -> bool: ...
|
||||
|
||||
@abstractmethod
|
||||
def bind_tasks(
|
||||
self,
|
||||
task_ids: List[str],
|
||||
total_len: int,
|
||||
device: torch.device,
|
||||
write_positions: Optional[Tensor] = None,
|
||||
) -> CacheView: ...
|
||||
|
||||
def task_cached(self, task_id: str) -> int:
|
||||
return 0
|
||||
|
||||
def task_record_hashes(
|
||||
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
|
||||
): ...
|
||||
|
||||
|
||||
class PageCacheView(CacheView):
|
||||
"""Bundles Storage + page_table + total_len for attention layers."""
|
||||
|
||||
def __init__(self, storage: Storage, page_table: Tensor, total_len: int = 0):
|
||||
self._storage = storage
|
||||
self._page_table = page_table
|
||||
self._total_len = total_len
|
||||
|
||||
def write(self, layer_id: int, k: Tensor, v: Tensor):
|
||||
start_pos = self._total_len - k.size(1)
|
||||
self._storage.write(layer_id, self._page_table, start_pos, k, v)
|
||||
|
||||
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]:
|
||||
return self._storage.gather(layer_id, self._page_table, self._total_len)
|
||||
|
||||
|
||||
class PageCache(KVCache):
|
||||
"""Paged KV-cache with prefix sharing."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
n_layers: int,
|
||||
n_pages: int,
|
||||
page_size: int,
|
||||
n_kv_heads: int,
|
||||
head_dim: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
):
|
||||
self.page_size = page_size
|
||||
self._pool = PagePool(Allocator(n_pages), PrefixCache(page_size))
|
||||
self._table = TaskTable(page_size)
|
||||
self._storage = Storage(
|
||||
n_layers, n_pages, page_size, n_kv_heads, head_dim, device, dtype
|
||||
)
|
||||
|
||||
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
|
||||
hits = self._pool.lookup(prompt_ids)
|
||||
cached = len(hits) * self.page_size
|
||||
for p in hits:
|
||||
self._pool.inc_ref(p)
|
||||
|
||||
remaining = len(prompt_ids) - cached
|
||||
n_new = (
|
||||
(remaining + self.page_size - 1) // self.page_size if remaining > 0 else 0
|
||||
)
|
||||
new_pages: List[int] = []
|
||||
if n_new > 0:
|
||||
for _ in range(n_new):
|
||||
p = self._pool.alloc()
|
||||
if p < 0:
|
||||
for hp in hits:
|
||||
self._pool.free(hp)
|
||||
for np in new_pages:
|
||||
self._pool.free(np)
|
||||
return False
|
||||
new_pages.append(p)
|
||||
|
||||
self._table.set(task_id, hits + new_pages, cached)
|
||||
return True
|
||||
|
||||
def task_free(self, task_id: str):
|
||||
page_table, _ = self._table.pop(task_id)
|
||||
for idx in page_table:
|
||||
self._pool.free(idx)
|
||||
|
||||
def task_extend(self, task_id: str, pos: int) -> bool:
|
||||
page_table = self._table.get(task_id)
|
||||
needed = (pos + 1 + self.page_size - 1) // self.page_size
|
||||
while len(page_table) < needed:
|
||||
p = self._pool.alloc()
|
||||
if p < 0:
|
||||
return False
|
||||
page_table.append(p)
|
||||
return True
|
||||
|
||||
def task_cached(self, task_id: str) -> int:
|
||||
return self._table.get_cached(task_id)
|
||||
|
||||
def task_record_hashes(
|
||||
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
|
||||
):
|
||||
page_table = self._table.get(task_id)
|
||||
full_pages = len(prompt_ids) // self.page_size
|
||||
for i in range(start_logical_page, full_pages):
|
||||
self._pool.record(page_table[i], prompt_ids, i)
|
||||
|
||||
def bind_tasks(
|
||||
self,
|
||||
task_ids: List[str],
|
||||
total_len: int,
|
||||
device: torch.device,
|
||||
write_positions: Optional[Tensor] = None,
|
||||
) -> PageCacheView:
|
||||
page_table = self._table.table_tensor(task_ids, device)
|
||||
return PageCacheView(self._storage, page_table, total_len)
|
||||
|
||||
|
||||
class ContiguousCacheView(CacheView):
|
||||
"""Contiguous KV-cache view for attention layers."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
cache: "ContiguousCache",
|
||||
batch_indices: Tensor,
|
||||
total_len: int = 0,
|
||||
write_positions: Optional[Tensor] = None,
|
||||
):
|
||||
self._cache = cache
|
||||
self._batch_indices = batch_indices
|
||||
self._total_len = total_len
|
||||
self._write_positions = write_positions
|
||||
|
||||
def write(self, layer_id: int, k: Tensor, v: Tensor):
|
||||
seq_len = k.size(1)
|
||||
indices = self._batch_indices
|
||||
if self._write_positions is not None and seq_len == 1:
|
||||
pos = self._write_positions
|
||||
self._cache.k[layer_id, indices, pos] = k.squeeze(1)
|
||||
self._cache.v[layer_id, indices, pos] = v.squeeze(1)
|
||||
else:
|
||||
start_pos = self._total_len - seq_len
|
||||
self._cache.k[layer_id, indices, start_pos : start_pos + seq_len] = k
|
||||
self._cache.v[layer_id, indices, start_pos : start_pos + seq_len] = v
|
||||
|
||||
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]:
|
||||
max_len = self._total_len
|
||||
indices = self._batch_indices
|
||||
k = self._cache.k[layer_id, indices, :max_len]
|
||||
v = self._cache.v[layer_id, indices, :max_len]
|
||||
return k, v
|
||||
|
||||
|
||||
class ContiguousCache(KVCache):
|
||||
"""Contiguous per-slot KV cache (default implementation)."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
n_layers: int,
|
||||
max_batch_size: int,
|
||||
max_seq_len: int,
|
||||
n_kv_heads: int,
|
||||
head_dim: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
):
|
||||
self.max_seq_len = max_seq_len
|
||||
self.k = torch.zeros(
|
||||
n_layers,
|
||||
max_batch_size,
|
||||
max_seq_len,
|
||||
n_kv_heads,
|
||||
head_dim,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
self.v = torch.zeros(
|
||||
n_layers,
|
||||
max_batch_size,
|
||||
max_seq_len,
|
||||
n_kv_heads,
|
||||
head_dim,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
self._slot_len: Dict[int, int] = {}
|
||||
self._task_slot: Dict[str, int] = {}
|
||||
self._free_slots = list(range(max_batch_size))
|
||||
self._device = device
|
||||
|
||||
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
|
||||
if not self._free_slots:
|
||||
return False
|
||||
slot = self._free_slots.pop(0)
|
||||
self._task_slot[task_id] = slot
|
||||
self._slot_len[slot] = 0
|
||||
return True
|
||||
|
||||
def task_free(self, task_id: str):
|
||||
slot = self._task_slot.pop(task_id, None)
|
||||
if slot is not None:
|
||||
self._slot_len.pop(slot, None)
|
||||
self._free_slots.append(slot)
|
||||
|
||||
def task_extend(self, task_id: str, pos: int) -> bool:
|
||||
return pos < self.max_seq_len
|
||||
|
||||
def task_cached(self, task_id: str) -> int:
|
||||
slot = self._task_slot.get(task_id)
|
||||
if slot is None:
|
||||
return 0
|
||||
return self._slot_len.get(slot, 0)
|
||||
|
||||
def bind_tasks(
|
||||
self,
|
||||
task_ids: List[str],
|
||||
total_len: int,
|
||||
device: torch.device,
|
||||
write_positions: Optional[Tensor] = None,
|
||||
) -> ContiguousCacheView:
|
||||
slots = [self._task_slot[tid] for tid in task_ids]
|
||||
batch_indices = torch.tensor(slots, dtype=torch.long, device=device)
|
||||
for slot in slots:
|
||||
if total_len > self._slot_len.get(slot, 0):
|
||||
self._slot_len[slot] = total_len
|
||||
return ContiguousCacheView(
|
||||
self, batch_indices, total_len, write_positions=write_positions
|
||||
)
|
||||
@@ -1,166 +0,0 @@
|
||||
import logging
|
||||
from typing import List, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.inference.core.cache import KVCache
|
||||
from astrai.inference.core.task import Task
|
||||
from astrai.inference.sample import sample
|
||||
from astrai.model.automodel import AutoModel
|
||||
from astrai.tokenize.tokenizer import AutoTokenizer
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class Executor:
|
||||
"""Model forward passes for prefill and decode phases."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: AutoModel,
|
||||
tokenizer: AutoTokenizer,
|
||||
kv_cache: KVCache,
|
||||
device: Optional[str] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
self.model = model
|
||||
self.tokenizer = tokenizer
|
||||
self.kv_cache = kv_cache
|
||||
self.device = device or next(model.parameters()).device
|
||||
self.dtype = dtype or next(model.parameters()).dtype
|
||||
|
||||
def execute_prefill(self, tasks: List[Task], prompt_len: int, start_pos: int = 0):
|
||||
if start_pos >= prompt_len:
|
||||
return
|
||||
|
||||
tasks = sorted(tasks, key=lambda t: t.task_id)
|
||||
batch_sz = len(tasks)
|
||||
|
||||
input_ids = torch.tensor(
|
||||
[t.prompt_ids[start_pos:prompt_len] for t in tasks],
|
||||
dtype=torch.long,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
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():
|
||||
self.model(
|
||||
input_ids,
|
||||
input_mask=input_mask,
|
||||
position_ids=position_ids,
|
||||
paged_cache=self.kv_cache.bind_tasks(task_ids, prompt_len, self.device),
|
||||
)
|
||||
|
||||
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:
|
||||
return []
|
||||
|
||||
input_ids = torch.tensor(
|
||||
[t.output_ids[-1] if t.output_ids else t.prompt_ids[-1] for t in tasks],
|
||||
dtype=torch.long,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
position_ids = torch.tensor(
|
||||
[t.next_pos for t in tasks], dtype=torch.long, device=self.device
|
||||
)
|
||||
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]
|
||||
|
||||
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_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
|
||||
)
|
||||
|
||||
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
|
||||
|
||||
with torch.inference_mode():
|
||||
outputs = self.model(
|
||||
input_ids.unsqueeze(1),
|
||||
input_mask=input_mask,
|
||||
paged_cache=self.kv_cache.bind_tasks(
|
||||
task_ids,
|
||||
total_len,
|
||||
self.device,
|
||||
write_positions=position_ids,
|
||||
),
|
||||
position_ids=position_ids.unsqueeze(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(
|
||||
logits,
|
||||
temperature=temperatures,
|
||||
top_k=top_ks,
|
||||
top_p=top_ps,
|
||||
frequency_penalty=freq_penalties,
|
||||
input_ids=padded_ids,
|
||||
input_mask=padded_mask,
|
||||
).tolist()
|
||||
@@ -1,311 +0,0 @@
|
||||
import logging
|
||||
import threading
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.inference.core.cache import ContiguousCache, KVCache
|
||||
from astrai.inference.core.executor import Executor
|
||||
from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
|
||||
from astrai.model.automodel import AutoModel
|
||||
from astrai.tokenize.tokenizer import AutoTokenizer
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class InferenceScheduler:
|
||||
"""Continuous batching loop: cleanup -> refill -> prefill -> decode (all groups)."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: AutoModel,
|
||||
tokenizer: AutoTokenizer,
|
||||
max_batch_size: int = 16,
|
||||
max_seq_len: Optional[int] = None,
|
||||
max_prompt_len: int = 2048,
|
||||
device: Optional[str] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
cache: Optional[KVCache] = None,
|
||||
):
|
||||
config = model.config
|
||||
|
||||
if max_seq_len is not None:
|
||||
self.max_seq_len = max_seq_len
|
||||
elif config.max_position_embeddings is not None:
|
||||
self.max_seq_len = config.max_position_embeddings
|
||||
else:
|
||||
raise ValueError(
|
||||
"max_seq_len must be provided either as argument "
|
||||
"or in model config (config.max_position_embeddings)"
|
||||
)
|
||||
self.device = device or next(model.parameters()).device
|
||||
self.dtype = dtype or next(model.parameters()).dtype
|
||||
|
||||
head_dim = config.hidden_size // config.num_attention_heads
|
||||
|
||||
if cache is not None:
|
||||
self._cache = cache
|
||||
else:
|
||||
self._cache = ContiguousCache(
|
||||
config.num_hidden_layers,
|
||||
max_batch_size,
|
||||
self.max_seq_len,
|
||||
config.num_key_value_heads,
|
||||
head_dim,
|
||||
self.device,
|
||||
self.dtype,
|
||||
)
|
||||
|
||||
self._task_mgr = TaskManager(
|
||||
tokenizer=tokenizer,
|
||||
max_batch_size=max_batch_size,
|
||||
max_seq_len=self.max_seq_len,
|
||||
max_prompt_len=max_prompt_len,
|
||||
)
|
||||
|
||||
self._executor = Executor(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
kv_cache=self._cache,
|
||||
device=self.device,
|
||||
dtype=self.dtype,
|
||||
)
|
||||
|
||||
self._stop_event = threading.Event()
|
||||
self._loop_thread: Optional[threading.Thread] = None
|
||||
|
||||
def add_task(self, prompt: str, **kwargs) -> str:
|
||||
return self._task_mgr.add_task(prompt, **kwargs)
|
||||
|
||||
def remove_task(self, task_id: str):
|
||||
for task in self._task_mgr.remove_task(task_id):
|
||||
self._cache.task_free(task.task_id)
|
||||
|
||||
def get_stats(self) -> Dict[str, Any]:
|
||||
return self._task_mgr.get_stats()
|
||||
|
||||
def _run_generation_loop(self):
|
||||
stop_ids = self._task_mgr.tokenizer.stop_ids
|
||||
cache = self._cache
|
||||
try:
|
||||
while not self._stop_event.is_set():
|
||||
finished = self._task_mgr.remove_finished_tasks(stop_ids)
|
||||
for task in finished:
|
||||
cache.task_free(task.task_id)
|
||||
|
||||
active = self._task_mgr.get_active_tasks()
|
||||
available = self._task_mgr.max_batch_size - len(active)
|
||||
if available > 0:
|
||||
candidates = self._task_mgr.pull_candidates(available)
|
||||
failed = []
|
||||
for task in candidates:
|
||||
if cache.task_alloc(task.task_id, task.prompt_ids):
|
||||
self._task_mgr.activate(task)
|
||||
else:
|
||||
failed.append(task)
|
||||
if failed:
|
||||
self._task_mgr.return_to_waiting(failed)
|
||||
|
||||
if not self._task_mgr.has_work():
|
||||
self._task_mgr.wait_for_tasks(timeout=1.0)
|
||||
continue
|
||||
|
||||
to_prefill = [
|
||||
t
|
||||
for t in self._task_mgr.get_active_tasks()
|
||||
if t.output_tokens == 0
|
||||
and cache.task_cached(t.task_id) < len(t.prompt_ids)
|
||||
]
|
||||
if to_prefill:
|
||||
for t in to_prefill:
|
||||
t.input_tokens = len(t.prompt_ids)
|
||||
|
||||
groups: Dict[Tuple[int, int], List[Task]] = {}
|
||||
for t in to_prefill:
|
||||
key = (
|
||||
len(t.prompt_ids),
|
||||
cache.task_cached(t.task_id),
|
||||
)
|
||||
groups.setdefault(key, []).append(t)
|
||||
|
||||
for (prompt_len, start_pos), group in groups.items():
|
||||
self._executor.execute_prefill(group, prompt_len, start_pos)
|
||||
start_logical_page = start_pos // getattr(
|
||||
cache, "page_size", 64
|
||||
)
|
||||
for t in group:
|
||||
cache.task_record_hashes(
|
||||
t.task_id, t.prompt_ids, start_logical_page
|
||||
)
|
||||
|
||||
decode_tasks = self._task_mgr.get_active_tasks()
|
||||
|
||||
valid: List[Task] = []
|
||||
for t in sorted(decode_tasks, key=lambda t: t.task_id):
|
||||
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)
|
||||
|
||||
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
|
||||
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)
|
||||
|
||||
except Exception as e:
|
||||
self._stop_event.set()
|
||||
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
|
||||
for task in self._task_mgr.get_active_tasks():
|
||||
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||
cache.task_free(task.task_id)
|
||||
for task in self._task_mgr.get_waiting_tasks():
|
||||
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||
self._task_mgr.clear_queues()
|
||||
|
||||
def start(self):
|
||||
if self._loop_thread is not None and self._loop_thread.is_alive():
|
||||
return
|
||||
self._stop_event.clear()
|
||||
t = threading.Thread(target=self._run_generation_loop, daemon=True)
|
||||
t.start()
|
||||
self._loop_thread = t
|
||||
|
||||
def stop(self):
|
||||
self._stop_event.set()
|
||||
self._task_mgr.wake()
|
||||
if self._loop_thread is not None:
|
||||
self._loop_thread.join(timeout=2.0)
|
||||
self._loop_thread = None
|
||||
for task in self._task_mgr.get_active_tasks():
|
||||
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||
self._cache.task_free(task.task_id)
|
||||
for task in self._task_mgr.get_waiting_tasks():
|
||||
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||
self._cache.task_free(task.task_id)
|
||||
self._task_mgr.clear_queues()
|
||||
if torch.cuda.is_available():
|
||||
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
|
||||
+72
-177
@@ -8,9 +8,10 @@ from typing import Any, AsyncGenerator, Dict, Generator, List, Optional, Tuple,
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from astrai.inference.core.cache import KVCache
|
||||
from astrai.inference.core.scheduler import InferenceScheduler
|
||||
from astrai.inference.core.task import STOP
|
||||
from astrai.extension import ATTN_BACKEND, AttentionBackend, get_backend
|
||||
from astrai.inference.cache import PagePool
|
||||
from astrai.inference.scheduler import InferenceScheduler
|
||||
from astrai.inference.task import STOP
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
|
||||
@@ -64,44 +65,6 @@ class GenerateResult:
|
||||
return self.results.copy()
|
||||
|
||||
|
||||
class GenerationRequest:
|
||||
"""Request parameters for text generation."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
messages: List[Dict[str, str]],
|
||||
top_k: int = 50,
|
||||
top_p: float = 1.0,
|
||||
temperature: float = 1.0,
|
||||
max_tokens: Optional[int] = None,
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
stream: bool = False,
|
||||
):
|
||||
if not (isinstance(top_k, int) and top_k >= 0):
|
||||
raise ValueError("top_k must be a non-negative integer")
|
||||
if not (0.0 <= top_p <= 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):
|
||||
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.top_k = top_k
|
||||
self.top_p = top_p
|
||||
self.temperature = temperature
|
||||
self.max_tokens = max_tokens
|
||||
self.frequency_penalty = frequency_penalty
|
||||
self.rep_window = rep_window
|
||||
self.stream = stream
|
||||
|
||||
|
||||
class InferenceEngine:
|
||||
"""Unified inference engine backed by continuous-batching scheduler."""
|
||||
|
||||
@@ -111,9 +74,9 @@ class InferenceEngine:
|
||||
tokenizer: AutoTokenizer,
|
||||
max_batch_size: int = 1,
|
||||
max_seq_len: Optional[int] = None,
|
||||
max_prompt_len: int = 2048,
|
||||
page_size: int = 128,
|
||||
cache: Optional[KVCache] = None,
|
||||
cache: Optional[PagePool] = None,
|
||||
enable_cuda_graph: bool = True,
|
||||
backend: Optional[Union[str, ATTN_BACKEND, AttentionBackend, type]] = None,
|
||||
):
|
||||
self.model = model
|
||||
self.tokenizer = tokenizer
|
||||
@@ -122,8 +85,9 @@ class InferenceEngine:
|
||||
tokenizer=self.tokenizer,
|
||||
max_batch_size=max_batch_size,
|
||||
max_seq_len=max_seq_len,
|
||||
max_prompt_len=max_prompt_len,
|
||||
cache=cache,
|
||||
enable_cuda_graph=enable_cuda_graph,
|
||||
backend=backend,
|
||||
)
|
||||
|
||||
self.scheduler.start()
|
||||
@@ -149,28 +113,23 @@ class InferenceEngine:
|
||||
is_batch = isinstance(prompt, list)
|
||||
prompts = prompt if is_batch else [prompt]
|
||||
|
||||
if stream:
|
||||
return self._generate_streaming(
|
||||
prompts,
|
||||
is_batch,
|
||||
max_tokens,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
frequency_penalty,
|
||||
rep_window,
|
||||
)
|
||||
else:
|
||||
return self._generate_non_streaming(
|
||||
prompts,
|
||||
is_batch,
|
||||
max_tokens,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
frequency_penalty,
|
||||
rep_window,
|
||||
)
|
||||
if max_tokens is not None and max_tokens <= 0:
|
||||
if stream:
|
||||
return iter(())
|
||||
results = [""] * len(prompts)
|
||||
return results if is_batch else results[0]
|
||||
|
||||
return self._generate(
|
||||
prompts,
|
||||
is_batch,
|
||||
stream,
|
||||
max_tokens,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
frequency_penalty,
|
||||
rep_window,
|
||||
)
|
||||
|
||||
def generate_async(
|
||||
self,
|
||||
@@ -182,9 +141,10 @@ class InferenceEngine:
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
) -> AsyncGenerator[str, None]:
|
||||
sync_gen = self._generate_streaming(
|
||||
sync_gen = self._generate(
|
||||
[prompt],
|
||||
False,
|
||||
True,
|
||||
max_tokens,
|
||||
temperature,
|
||||
top_p,
|
||||
@@ -196,51 +156,31 @@ class InferenceEngine:
|
||||
async def _agen():
|
||||
loop = asyncio.get_event_loop()
|
||||
while True:
|
||||
token = await loop.run_in_executor(None, self._next_token, sync_gen)
|
||||
if token is None:
|
||||
try:
|
||||
token = await loop.run_in_executor(None, next, sync_gen)
|
||||
except StopIteration:
|
||||
break
|
||||
yield token
|
||||
|
||||
return _agen()
|
||||
|
||||
@staticmethod
|
||||
def _next_token(gen: Generator) -> Optional[str]:
|
||||
try:
|
||||
return next(gen)
|
||||
except StopIteration:
|
||||
return None
|
||||
|
||||
def generate_with_request(
|
||||
self, request: GenerationRequest
|
||||
) -> Union[Generator[str, None, None], str, List[str]]:
|
||||
prompt = self.tokenizer.apply_chat_template(request.messages, tokenize=False)
|
||||
return self.generate(
|
||||
prompt=prompt,
|
||||
stream=request.stream,
|
||||
max_tokens=request.max_tokens,
|
||||
temperature=request.temperature,
|
||||
top_p=request.top_p,
|
||||
top_k=request.top_k,
|
||||
frequency_penalty=request.frequency_penalty,
|
||||
rep_window=request.rep_window,
|
||||
)
|
||||
|
||||
def _submit_tasks(
|
||||
def _generate(
|
||||
self,
|
||||
prompts: List[str],
|
||||
is_batch: bool,
|
||||
stream: bool,
|
||||
max_tokens: Optional[int],
|
||||
temperature: float,
|
||||
top_p: float,
|
||||
top_k: int,
|
||||
frequency_penalty: float,
|
||||
rep_window: int,
|
||||
) -> Tuple[GenerateResult, List[str]]:
|
||||
) -> Union[Generator, str, List[str]]:
|
||||
n = len(prompts)
|
||||
request_backend = get_backend(use_default=False)
|
||||
result = GenerateResult(count=n)
|
||||
task_ids = []
|
||||
for i, p in enumerate(prompts):
|
||||
cb = self._make_callback(result, i)
|
||||
task_id = self.scheduler.add_task(
|
||||
task_ids = [
|
||||
self.scheduler.add_task(
|
||||
prompt=p,
|
||||
max_tokens=max_tokens,
|
||||
temperature=temperature,
|
||||
@@ -248,99 +188,54 @@ class InferenceEngine:
|
||||
top_k=top_k,
|
||||
frequency_penalty=frequency_penalty,
|
||||
rep_window=rep_window,
|
||||
stream_callback=cb,
|
||||
backend=request_backend,
|
||||
stream_callback=lambda token, idx=i: result.append(token, idx),
|
||||
)
|
||||
task_ids.append(task_id)
|
||||
return result, task_ids
|
||||
for i, p in enumerate(prompts)
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def _make_callback(result: GenerateResult, idx: int):
|
||||
def cb(token):
|
||||
result.append(token, idx)
|
||||
if not stream:
|
||||
try:
|
||||
result.wait_completion()
|
||||
except TimeoutError:
|
||||
for tid in task_ids:
|
||||
self.scheduler.remove_task(tid)
|
||||
raise
|
||||
for tid in task_ids:
|
||||
self.scheduler.remove_task(tid)
|
||||
res = result.get_results()
|
||||
return res if is_batch else res[0]
|
||||
|
||||
return cb
|
||||
|
||||
def _generate_streaming(
|
||||
self,
|
||||
prompts: List[str],
|
||||
is_batch: bool,
|
||||
max_tokens: Optional[int],
|
||||
temperature: float,
|
||||
top_p: float,
|
||||
top_k: int,
|
||||
frequency_penalty: float,
|
||||
rep_window: int,
|
||||
) -> Generator:
|
||||
result, task_ids = self._submit_tasks(
|
||||
prompts,
|
||||
max_tokens,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
frequency_penalty,
|
||||
rep_window,
|
||||
)
|
||||
n = len(prompts)
|
||||
remaining = n
|
||||
finished = [False] * n
|
||||
|
||||
def gen():
|
||||
nonlocal remaining
|
||||
try:
|
||||
while remaining > 0:
|
||||
items = result.pop_all()
|
||||
for idx, token in items:
|
||||
if token is STOP:
|
||||
if not finished[idx]:
|
||||
finished[idx] = True
|
||||
remaining -= 1
|
||||
else:
|
||||
yield (idx, token) if is_batch else token
|
||||
if remaining > 0:
|
||||
result.wait(timeout=0.05)
|
||||
finally:
|
||||
for tid in task_ids:
|
||||
self.scheduler.remove_task(tid)
|
||||
while remaining > 0:
|
||||
items = result.pop_all()
|
||||
for idx, token in items:
|
||||
if token is STOP:
|
||||
if not finished[idx]:
|
||||
finished[idx] = True
|
||||
remaining -= 1
|
||||
else:
|
||||
yield (idx, token) if is_batch else token
|
||||
if remaining > 0:
|
||||
result.wait(timeout=0.05)
|
||||
|
||||
return gen()
|
||||
|
||||
def _generate_non_streaming(
|
||||
self,
|
||||
prompts: List[str],
|
||||
is_batch: bool,
|
||||
max_tokens: Optional[int],
|
||||
temperature: float,
|
||||
top_p: float,
|
||||
top_k: int,
|
||||
frequency_penalty: float,
|
||||
rep_window: int,
|
||||
) -> Union[str, List[str]]:
|
||||
result, task_ids = self._submit_tasks(
|
||||
prompts,
|
||||
max_tokens,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
frequency_penalty,
|
||||
rep_window,
|
||||
)
|
||||
|
||||
try:
|
||||
result.wait_completion()
|
||||
except TimeoutError:
|
||||
for tid in task_ids:
|
||||
self.scheduler.remove_task(tid)
|
||||
raise
|
||||
|
||||
for tid in task_ids:
|
||||
self.scheduler.remove_task(tid)
|
||||
|
||||
res = result.get_results()
|
||||
return res if is_batch else res[0]
|
||||
|
||||
def get_stats(self) -> Dict[str, Any]:
|
||||
return self.scheduler.get_stats()
|
||||
|
||||
@property
|
||||
def backend_name(self) -> str:
|
||||
return self.scheduler.backend_name
|
||||
|
||||
@property
|
||||
def cuda_graph_enabled(self) -> bool:
|
||||
return self.scheduler.cuda_graph_enabled
|
||||
|
||||
def shutdown(self):
|
||||
self.scheduler.stop()
|
||||
if torch.cuda.is_available():
|
||||
|
||||
@@ -0,0 +1,223 @@
|
||||
"""Unified per-task perf/stats: timing records, context-manager scopes, aggregate reporting."""
|
||||
|
||||
import time
|
||||
from collections import deque
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Deque, Dict, Generator, List, Literal, Optional
|
||||
|
||||
|
||||
@dataclass
|
||||
class TaskTiming:
|
||||
"""Timestamp snapshots and computed metrics for one generation task.
|
||||
|
||||
Created by :class:`MetricsCollector` at task-registration time;
|
||||
updated via ``record`` / ``mark_finished``.
|
||||
"""
|
||||
|
||||
task_id: str
|
||||
arrival_time: float
|
||||
prefill_start_time: Optional[float] = None
|
||||
first_token_time: Optional[float] = None
|
||||
finish_time: Optional[float] = None
|
||||
input_tokens: int = 0
|
||||
output_tokens: int = 0
|
||||
_decode_steps: int = 0
|
||||
_decode_total_s: float = 0.0
|
||||
|
||||
# derived metrics
|
||||
|
||||
@property
|
||||
def queue_wait_ms(self) -> Optional[float]:
|
||||
if self.prefill_start_time is not None:
|
||||
return (self.prefill_start_time - self.arrival_time) * 1000
|
||||
return None
|
||||
|
||||
@property
|
||||
def ttft_ms(self) -> Optional[float]:
|
||||
if self.first_token_time is not None:
|
||||
return (self.first_token_time - self.arrival_time) * 1000
|
||||
return None
|
||||
|
||||
@property
|
||||
def prefill_tps(self) -> Optional[float]:
|
||||
if self.prefill_start_time is not None and self.first_token_time is not None:
|
||||
d = self.first_token_time - self.prefill_start_time
|
||||
if d > 0 and self.input_tokens > 0:
|
||||
return self.input_tokens / d
|
||||
return None
|
||||
|
||||
@property
|
||||
def decode_tps(self) -> Optional[float]:
|
||||
if self.first_token_time is not None and self.finish_time is not None:
|
||||
d = self.finish_time - self.first_token_time
|
||||
dt = self.output_tokens - 1
|
||||
if dt > 0 and d > 0:
|
||||
return dt / d
|
||||
return None
|
||||
|
||||
@property
|
||||
def decode_avg_ms(self) -> Optional[float]:
|
||||
if self._decode_steps > 0 and self._decode_total_s > 0:
|
||||
return (self._decode_total_s / self._decode_steps) * 1000
|
||||
return None
|
||||
|
||||
@property
|
||||
def e2e_latency_ms(self) -> Optional[float]:
|
||||
if self.finish_time is not None:
|
||||
return (self.finish_time - self.arrival_time) * 1000
|
||||
return None
|
||||
|
||||
@property
|
||||
def total_tps(self) -> Optional[float]:
|
||||
if self.finish_time is not None:
|
||||
total = self.input_tokens + self.output_tokens
|
||||
d = self.finish_time - self.arrival_time
|
||||
if total > 0 and d > 0:
|
||||
return total / d
|
||||
return None
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
return {
|
||||
"task_id": self.task_id,
|
||||
"input_tokens": self.input_tokens,
|
||||
"output_tokens": self.output_tokens,
|
||||
"queue_wait_ms": (
|
||||
round(self.queue_wait_ms, 2) if self.queue_wait_ms is not None else None
|
||||
),
|
||||
"ttft_ms": (round(self.ttft_ms, 2) if self.ttft_ms is not None else None),
|
||||
"prefill_tps": (
|
||||
round(self.prefill_tps, 2) if self.prefill_tps is not None else None
|
||||
),
|
||||
"decode_tps": (
|
||||
round(self.decode_tps, 2) if self.decode_tps is not None else None
|
||||
),
|
||||
"decode_avg_ms": (
|
||||
round(self.decode_avg_ms, 2) if self.decode_avg_ms is not None else None
|
||||
),
|
||||
"total_tps": (
|
||||
round(self.total_tps, 2) if self.total_tps is not None else None
|
||||
),
|
||||
"e2e_latency_ms": (
|
||||
round(self.e2e_latency_ms, 2)
|
||||
if self.e2e_latency_ms is not None
|
||||
else None
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
class MetricsCollector:
|
||||
"""Single-owner perf/stats hub for all generation tasks.
|
||||
|
||||
Usage::
|
||||
|
||||
metrics = MetricsCollector()
|
||||
metrics.register(task_id, arrival_time)
|
||||
|
||||
with metrics.record(task_ids, "prefill"):
|
||||
run_prefill(...)
|
||||
|
||||
metrics.mark_finished(task_id, input_tokens, output_tokens)
|
||||
|
||||
stats = metrics.get_stats()
|
||||
"""
|
||||
|
||||
def __init__(self, max_recent: int = 128):
|
||||
self._timings: Dict[str, TaskTiming] = {}
|
||||
self._completed: Deque[TaskTiming] = deque(maxlen=max_recent)
|
||||
|
||||
self._ttft_ms_sum = 0.0
|
||||
self._ttft_ms_count = 0
|
||||
self._decode_tps_sum = 0.0
|
||||
self._decode_tps_count = 0
|
||||
self._e2e_ms_sum = 0.0
|
||||
self._e2e_ms_count = 0
|
||||
|
||||
def register(self, task_id: str):
|
||||
"""Create a timing record for a newly-created task."""
|
||||
self._timings[task_id] = TaskTiming(task_id=task_id, arrival_time=time.time())
|
||||
|
||||
def mark_finished(self, task_id: str, input_tokens: int, output_tokens: int):
|
||||
"""Close timing for a finished/aborted task and move it to completed."""
|
||||
timing = self._timings.pop(task_id, None)
|
||||
if timing is None:
|
||||
return
|
||||
timing.finish_time = time.time()
|
||||
timing.input_tokens = input_tokens
|
||||
timing.output_tokens = output_tokens
|
||||
self._completed.append(timing)
|
||||
self._accumulate(timing)
|
||||
|
||||
def clear(self):
|
||||
"""Reset all state (e.g. on engine shutdown)."""
|
||||
self._timings.clear()
|
||||
self._completed.clear()
|
||||
self._ttft_ms_sum = 0.0
|
||||
self._ttft_ms_count = 0
|
||||
self._decode_tps_sum = 0.0
|
||||
self._decode_tps_count = 0
|
||||
self._e2e_ms_sum = 0.0
|
||||
self._e2e_ms_count = 0
|
||||
|
||||
# timing scopes
|
||||
|
||||
@contextmanager
|
||||
def record(
|
||||
self, task_ids: List[str], phase: Literal["prefill", "decode"]
|
||||
) -> Generator[None, None, None]:
|
||||
tic = time.time()
|
||||
yield
|
||||
toc = time.time()
|
||||
dt = toc - tic
|
||||
for tid in task_ids:
|
||||
t = self._timings.get(tid)
|
||||
if t is None:
|
||||
continue
|
||||
if phase == "prefill":
|
||||
t.prefill_start_time = tic
|
||||
t.first_token_time = toc
|
||||
elif phase == "decode":
|
||||
t._decode_steps += 1
|
||||
t._decode_total_s += dt
|
||||
|
||||
# access
|
||||
|
||||
def get_timing(self, task_id: str) -> Optional[TaskTiming]:
|
||||
"""Return the timing record for *task_id* (active or completed)."""
|
||||
if task_id in self._timings:
|
||||
return self._timings[task_id]
|
||||
for t in self._completed:
|
||||
if t.task_id == task_id:
|
||||
return t
|
||||
return None
|
||||
|
||||
# aggregate stats
|
||||
|
||||
def get_stats(self) -> Dict[str, Any]:
|
||||
stats: Dict[str, Any] = {}
|
||||
if self._ttft_ms_count > 0:
|
||||
stats["avg_ttft_ms"] = round(self._ttft_ms_sum / self._ttft_ms_count, 2)
|
||||
if self._decode_tps_count > 0:
|
||||
stats["avg_decode_tps"] = round(
|
||||
self._decode_tps_sum / self._decode_tps_count, 2
|
||||
)
|
||||
if self._e2e_ms_count > 0:
|
||||
stats["avg_e2e_latency_ms"] = round(
|
||||
self._e2e_ms_sum / self._e2e_ms_count, 2
|
||||
)
|
||||
if self._completed:
|
||||
stats["recent_tasks"] = [t.to_dict() for t in self._completed]
|
||||
return stats
|
||||
|
||||
# internal
|
||||
|
||||
def _accumulate(self, t: TaskTiming):
|
||||
if t.ttft_ms is not None:
|
||||
self._ttft_ms_sum += t.ttft_ms
|
||||
self._ttft_ms_count += 1
|
||||
if t.decode_tps is not None:
|
||||
self._decode_tps_sum += t.decode_tps
|
||||
self._decode_tps_count += 1
|
||||
if t.e2e_latency_ms is not None:
|
||||
self._e2e_ms_sum += t.e2e_latency_ms
|
||||
self._e2e_ms_count += 1
|
||||
@@ -4,8 +4,7 @@
|
||||
lazy singleton FastAPI instance.
|
||||
"""
|
||||
|
||||
from astrai.inference.api.protocol import GenContext, ProtocolHandler, StopChecker
|
||||
from astrai.inference.api.server import (
|
||||
from astrai.inference.network.app import (
|
||||
AnthropicMessage,
|
||||
ChatCompletionRequest,
|
||||
ChatMessage,
|
||||
@@ -15,7 +14,8 @@ from astrai.inference.api.server import (
|
||||
get_app,
|
||||
run_server,
|
||||
)
|
||||
from astrai.inference.api.tool_parser import (
|
||||
from astrai.inference.network.protocol import GenContext, ProtocolHandler, StopChecker
|
||||
from astrai.inference.network.tool_parser import (
|
||||
BaseToolParser,
|
||||
SimpleJsonToolParser,
|
||||
ToolParserFactory,
|
||||
@@ -6,13 +6,13 @@ from typing import Any, Dict, List, Tuple, Union
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from astrai.inference.api.protocol import (
|
||||
from astrai.inference.engine import InferenceEngine
|
||||
from astrai.inference.network.protocol import (
|
||||
GenContext,
|
||||
ResponseBuilder,
|
||||
StopInfo,
|
||||
sse_event,
|
||||
)
|
||||
from astrai.inference.engine import InferenceEngine
|
||||
|
||||
|
||||
def _extract_text(content: Union[str, List[Dict[str, Any]]]) -> str:
|
||||
@@ -18,10 +18,10 @@ import uvicorn
|
||||
from fastapi import APIRouter, FastAPI, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from astrai.inference.api.anthropic import AnthropicResponseBuilder
|
||||
from astrai.inference.api.openai import OpenAIResponseBuilder
|
||||
from astrai.inference.api.protocol import ProtocolHandler
|
||||
from astrai.inference.engine import InferenceEngine
|
||||
from astrai.inference.network.anthropic import AnthropicResponseBuilder
|
||||
from astrai.inference.network.openai import OpenAIResponseBuilder
|
||||
from astrai.inference.network.protocol import ProtocolHandler
|
||||
from astrai.model import AutoModel
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
@@ -110,6 +110,7 @@ def _create_engine(
|
||||
device: str = "cuda",
|
||||
dtype: torch.dtype = torch.bfloat16,
|
||||
max_batch_size: int = 16,
|
||||
max_seq_len: Optional[int] = None,
|
||||
) -> InferenceEngine:
|
||||
if not param_path.exists():
|
||||
raise FileNotFoundError(f"Parameter directory not found: {param_path}")
|
||||
@@ -123,6 +124,7 @@ def _create_engine(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
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}")
|
||||
return engine
|
||||
@@ -186,6 +188,7 @@ def run_server(
|
||||
device: str = "cuda",
|
||||
dtype: torch.dtype = torch.bfloat16,
|
||||
max_batch_size: int = 16,
|
||||
max_seq_len: Optional[int] = None,
|
||||
):
|
||||
app = get_app()
|
||||
app.state.server_config = {
|
||||
@@ -193,6 +196,7 @@ def run_server(
|
||||
"dtype": dtype,
|
||||
"param_path": param_path,
|
||||
"max_batch_size": max_batch_size,
|
||||
"max_seq_len": max_seq_len,
|
||||
}
|
||||
uvicorn.run(
|
||||
app,
|
||||
@@ -7,14 +7,14 @@ from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from astrai.inference.api.protocol import (
|
||||
from astrai.inference.engine import InferenceEngine
|
||||
from astrai.inference.network.protocol import (
|
||||
GenContext,
|
||||
ResponseBuilder,
|
||||
StopInfo,
|
||||
sse_event,
|
||||
)
|
||||
from astrai.inference.api.tool_parser import BaseToolParser, ToolParserFactory
|
||||
from astrai.inference.engine import InferenceEngine
|
||||
from astrai.inference.network.tool_parser import BaseToolParser, ToolParserFactory
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -181,12 +181,10 @@ class ProtocolHandler:
|
||||
self, agen: AsyncGenerator, ctx: GenContext, stop_sequences: List[str]
|
||||
) -> Dict[str, Any]:
|
||||
checker = StopChecker(stop_sequences)
|
||||
chunks: List[str] = []
|
||||
body = ""
|
||||
matched = None
|
||||
|
||||
async for token in agen:
|
||||
chunks.append(token)
|
||||
body += token
|
||||
|
||||
matched = checker.check(body)
|
||||
@@ -195,6 +193,5 @@ class ProtocolHandler:
|
||||
|
||||
ctx.completion_tokens += 1
|
||||
|
||||
content = "".join(chunks)
|
||||
stop = StopInfo(matched=matched, body=body)
|
||||
return self.builder.format_response(ctx, content, stop)
|
||||
return self.builder.format_response(ctx, body, stop)
|
||||
@@ -22,13 +22,10 @@ class BaseToolParser(ABC):
|
||||
Maintains streaming state internally so that each call to :meth:`feed`
|
||||
can diff against previously emitted content.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
tools : list of dict, optional
|
||||
Tool definitions from the request.
|
||||
tool_choice : str
|
||||
``"auto"`` / ``"required"`` / ``"none"`` or a named tool choice
|
||||
dict.
|
||||
Args:
|
||||
tools (list of dict, optional): Tool definitions from the request.
|
||||
tool_choice (str): ``"auto"`` / ``"required"`` / ``"none"`` or a named
|
||||
tool choice dict.
|
||||
"""
|
||||
|
||||
def __init__(self, tools: Optional[List[Dict]] = None, tool_choice: str = "auto"):
|
||||
@@ -51,14 +48,12 @@ class BaseToolParser(ABC):
|
||||
|
||||
Returns an empty list when nothing new should be emitted.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
body : str
|
||||
The complete accumulated generated text so far.
|
||||
current_token_ids : list of int, optional
|
||||
All token IDs decoded into *body* (cumulative).
|
||||
delta_token_ids : list of int, optional
|
||||
Only the token IDs for this chunk.
|
||||
Args:
|
||||
body (str): The complete accumulated generated text so far.
|
||||
current_token_ids (list of int, optional): All token IDs decoded
|
||||
into *body* (cumulative).
|
||||
delta_token_ids (list of int, optional): Only the token IDs for
|
||||
this chunk.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
@@ -0,0 +1,25 @@
|
||||
"""Execution primitives: forward passes, CUDA graphs, and sampling."""
|
||||
|
||||
from astrai.inference.runtime.executor import Executor
|
||||
from astrai.inference.runtime.graph import CudaGraphContext
|
||||
from astrai.inference.runtime.sample import (
|
||||
BaseSamplingStrategy,
|
||||
FrequencyPenaltyStrategy,
|
||||
SamplingPipeline,
|
||||
TemperatureStrategy,
|
||||
TopKStrategy,
|
||||
TopPStrategy,
|
||||
sample,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"Executor",
|
||||
"CudaGraphContext",
|
||||
"BaseSamplingStrategy",
|
||||
"FrequencyPenaltyStrategy",
|
||||
"SamplingPipeline",
|
||||
"TemperatureStrategy",
|
||||
"TopKStrategy",
|
||||
"TopPStrategy",
|
||||
"sample",
|
||||
]
|
||||
@@ -0,0 +1,421 @@
|
||||
import logging
|
||||
import time
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.extension.backend.attention import (
|
||||
CudaBackend,
|
||||
get_backend,
|
||||
)
|
||||
from astrai.inference.cache import PagePool, TaskCacheManager
|
||||
from astrai.inference.runtime.graph import CudaGraphContext
|
||||
from astrai.inference.runtime.sample import sample
|
||||
from astrai.inference.task import Task
|
||||
from astrai.inference.workspace import InferenceWorkspace
|
||||
from astrai.model.automodel import AutoModel
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def timed(label: str, log: Optional[logging.Logger] = None):
|
||||
"""GPU-precise timer via CUDA events; falls back to perf_counter on CPU."""
|
||||
log = log or logger
|
||||
if not log.isEnabledFor(logging.DEBUG):
|
||||
yield
|
||||
return
|
||||
use_cuda = torch.cuda.is_available()
|
||||
if use_cuda:
|
||||
start = torch.cuda.Event(enable_timing=True)
|
||||
end = torch.cuda.Event(enable_timing=True)
|
||||
start.record()
|
||||
else:
|
||||
tic = time.perf_counter()
|
||||
yield
|
||||
if use_cuda:
|
||||
end.record()
|
||||
torch.cuda.synchronize()
|
||||
elapsed_ms = start.elapsed_time(end)
|
||||
else:
|
||||
elapsed_ms = (time.perf_counter() - tic) * 1000
|
||||
log.debug("%s %.2fms", label, elapsed_ms)
|
||||
|
||||
|
||||
@dataclass
|
||||
class SamplingBatchInfo:
|
||||
"""Per-batch sampling parameters, cached across decode steps.
|
||||
|
||||
Sampling params are constant for a given ordered task set, so they are
|
||||
built once (pinned-memory async H2D) and reused until the task set
|
||||
changes. ``top_ks`` is int32 to match the native consumers.
|
||||
"""
|
||||
|
||||
temperatures: Tensor # float32 [B]
|
||||
top_ks: Tensor # int32 [B]
|
||||
top_ps: Tensor # float32 [B]
|
||||
freq_penalties: Tensor # float32 [B]
|
||||
has_freq: bool # any frequency_penalty != 0 (avoids per-step GPU .any())
|
||||
|
||||
|
||||
@dataclass
|
||||
class DecodeSteadyState:
|
||||
"""Cached decode metadata for the steady-state case.
|
||||
|
||||
When the same ordered task set decodes one token per step, sampling
|
||||
params and task signature are reused; only positions advance by 1.
|
||||
"""
|
||||
|
||||
task_sig: tuple
|
||||
positions: list[int]
|
||||
sampling_info: SamplingBatchInfo
|
||||
|
||||
|
||||
def _build_sampling_batch_info(tasks: List[Task], device) -> SamplingBatchInfo:
|
||||
pin = str(device).startswith("cuda")
|
||||
freq_penalties = torch.tensor(
|
||||
[t.frequency_penalty for t in tasks], dtype=torch.float32, pin_memory=pin
|
||||
).to(device, non_blocking=True)
|
||||
return SamplingBatchInfo(
|
||||
temperatures=torch.tensor(
|
||||
[t.temperature for t in tasks], dtype=torch.float32, pin_memory=pin
|
||||
).to(device, non_blocking=True),
|
||||
top_ks=torch.tensor(
|
||||
[t.top_k for t in tasks], dtype=torch.int32, pin_memory=pin
|
||||
).to(device, non_blocking=True),
|
||||
top_ps=torch.tensor(
|
||||
[t.top_p for t in tasks], dtype=torch.float32, pin_memory=pin
|
||||
).to(device, non_blocking=True),
|
||||
freq_penalties=freq_penalties,
|
||||
has_freq=bool((freq_penalties != 0).any()),
|
||||
)
|
||||
|
||||
|
||||
def _warmup_cuda_graphs(
|
||||
model: AutoModel,
|
||||
pool: PagePool,
|
||||
task_cache: TaskCacheManager,
|
||||
ws: InferenceWorkspace,
|
||||
gctx: CudaGraphContext,
|
||||
max_batch_size: int,
|
||||
prompt_len: int = 1,
|
||||
device: Optional[str] = None,
|
||||
):
|
||||
dev = device or next(model.parameters()).device
|
||||
|
||||
# Prefill warmup: cuBLAS auto-tunes for the actual prompt-length tensor
|
||||
# shapes on first call (F.linear is the dominant cost). This also warms
|
||||
# up the CUDA context (driver init) and compiles the graph-capture trace
|
||||
# that follows. Custom .so kernels do NOT need this — they are pre-built.
|
||||
warmup_len = 64
|
||||
tid = "_warmup_prefill"
|
||||
if task_cache.task_alloc(tid, list(range(warmup_len))):
|
||||
with (
|
||||
torch.inference_mode(),
|
||||
timed("warmup prefill", logger),
|
||||
):
|
||||
kv = task_cache.bind([tid], ws, start_pos=0)
|
||||
ids_in = torch.arange(warmup_len, device=dev)
|
||||
pos_in = ids_in
|
||||
model(
|
||||
ids_in,
|
||||
kv_cache=kv,
|
||||
position_ids=pos_in,
|
||||
fwd="prefill",
|
||||
)
|
||||
task_cache.task_free(tid)
|
||||
|
||||
batch_sizes = [1]
|
||||
n = 2
|
||||
while n <= max_batch_size:
|
||||
batch_sizes.append(n)
|
||||
n *= 2
|
||||
if max_batch_size not in batch_sizes:
|
||||
batch_sizes.append(max_batch_size)
|
||||
|
||||
for b in batch_sizes:
|
||||
task_ids = [f"_warmup_decode_{b}_{i}" for i in range(b)]
|
||||
prompt_tokens = [list(range(prompt_len)) for _ in range(b)]
|
||||
alloc_ok = True
|
||||
for tid, pt in zip(task_ids, prompt_tokens):
|
||||
if not task_cache.task_alloc(tid, pt):
|
||||
alloc_ok = False
|
||||
break
|
||||
if not alloc_ok:
|
||||
for tid in task_ids:
|
||||
task_cache.task_free(tid)
|
||||
continue
|
||||
|
||||
with (
|
||||
torch.inference_mode(),
|
||||
timed(f"warmup decode b={b}", logger),
|
||||
):
|
||||
for step in range(2):
|
||||
seq_pos = step
|
||||
ws.position_ids[:b] = seq_pos
|
||||
for tid in task_ids:
|
||||
task_cache.task_extend(tid, seq_pos)
|
||||
kv = task_cache.bind(task_ids, ws)
|
||||
ids_buf = ws.fill_input_ids([step] * b)
|
||||
gctx.forward(
|
||||
model,
|
||||
key=(b,),
|
||||
input_ids=ids_buf,
|
||||
kv_cache=kv,
|
||||
position_ids=ws.position_ids[:b],
|
||||
fwd="decode",
|
||||
)
|
||||
|
||||
for tid in task_ids:
|
||||
task_cache.task_free(tid)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
|
||||
class Executor:
|
||||
"""Model forward passes for prefill and decode phases."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: AutoModel,
|
||||
kv_cache: PagePool,
|
||||
task_cache: TaskCacheManager,
|
||||
device: Optional[str] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
enable_cuda_graph: bool = True,
|
||||
):
|
||||
self.model = model
|
||||
self.kv_cache = kv_cache
|
||||
self.task_cache = task_cache
|
||||
self.device = device or next(model.parameters()).device
|
||||
self.dtype = dtype or next(model.parameters()).dtype
|
||||
|
||||
# Per-step decode cache for the steady-state case (same ordered
|
||||
# task set decodes one token per step). Sampling params stay
|
||||
# constant; only positions advance.
|
||||
self._decode_cache: Optional[DecodeSteadyState] = None
|
||||
|
||||
# Pre-allocated fixed-shape buffers for the decode hot path
|
||||
# (input_ids, decode mask, KV bind metadata). Eagerly sized at init
|
||||
# so the workspace is CUDA-graph-capture friendly — no allocation
|
||||
# during capture.
|
||||
config = model.config
|
||||
max_q_heads = config.num_attention_heads
|
||||
head_dim = config.hidden_size // config.num_attention_heads
|
||||
backend = get_backend()
|
||||
self._graph_supported = backend.supports_graph() and CudaBackend.supports(
|
||||
head_dim=head_dim
|
||||
)
|
||||
self._workspace = InferenceWorkspace(
|
||||
max_batch_size=kv_cache.max_batch_size,
|
||||
max_seq_len=kv_cache.max_seq_len,
|
||||
max_q_heads=max_q_heads,
|
||||
head_dim=head_dim,
|
||||
device=self.device,
|
||||
dtype=self.dtype,
|
||||
)
|
||||
|
||||
# CUDA-graph capture: one graph per (batch_size,) key.
|
||||
# Enabled at init-time via _warmup_cuda_graphs for CudaBackend
|
||||
# on supported head_dims; left disabled otherwise.
|
||||
self._graph_ctx = CudaGraphContext()
|
||||
if enable_cuda_graph:
|
||||
self._try_enable_cuda_graph()
|
||||
|
||||
def _try_enable_cuda_graph(self):
|
||||
if not self._graph_supported:
|
||||
return
|
||||
|
||||
self._graph_ctx.set_enabled(True)
|
||||
_warmup_cuda_graphs(
|
||||
self.model,
|
||||
self.kv_cache,
|
||||
self.task_cache,
|
||||
self._workspace,
|
||||
self._graph_ctx,
|
||||
max_batch_size=self.kv_cache.max_batch_size,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
@property
|
||||
def cuda_graph_enabled(self) -> bool:
|
||||
return self._graph_ctx.enabled and self._graph_supported
|
||||
|
||||
def _sample_logits(
|
||||
self,
|
||||
logits: Tensor,
|
||||
tasks: List[Task],
|
||||
return_logprobs: bool = False,
|
||||
info: Optional[SamplingBatchInfo] = None,
|
||||
):
|
||||
info = info or _build_sampling_batch_info(tasks, self.device)
|
||||
if info.has_freq:
|
||||
history_lists = [
|
||||
t.prompt_ids[-t.rep_window :] + t.output_ids for t in tasks
|
||||
]
|
||||
history_lens = [len(ids) for ids in history_lists]
|
||||
max_len = max(history_lens, default=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, ids in enumerate(history_lists):
|
||||
length = len(ids)
|
||||
padded_ids[i, :length] = torch.as_tensor(
|
||||
ids, dtype=torch.long, device=self.device
|
||||
)
|
||||
padded_mask[i, :length] = True
|
||||
else:
|
||||
padded_ids = None
|
||||
padded_mask = None
|
||||
|
||||
result = sample(
|
||||
logits,
|
||||
temperature=info.temperatures,
|
||||
top_k=info.top_ks,
|
||||
top_p=info.top_ps,
|
||||
frequency_penalty=info.freq_penalties,
|
||||
input_ids=padded_ids,
|
||||
input_mask=padded_mask,
|
||||
return_logprobs=return_logprobs,
|
||||
)
|
||||
if not return_logprobs:
|
||||
return result.tolist()
|
||||
|
||||
tokens, logprobs = result
|
||||
tokens_list = tokens.tolist()
|
||||
logprobs_list = logprobs.tolist()
|
||||
for task, logprob in zip(tasks, logprobs_list):
|
||||
task.output_logprobs.append(float(logprob))
|
||||
return list(zip(tokens_list, logprobs_list))
|
||||
|
||||
def execute_prefill(
|
||||
self,
|
||||
tasks: List[Task],
|
||||
prompt_len: int,
|
||||
start_pos: int = 0,
|
||||
return_logprobs: bool = False,
|
||||
):
|
||||
if start_pos >= prompt_len:
|
||||
return []
|
||||
|
||||
tasks = sorted(tasks, key=lambda t: t.task_id)
|
||||
batch_sz = len(tasks)
|
||||
|
||||
input_ids = torch.tensor(
|
||||
[token for t in tasks for token in t.prompt_ids[start_pos:prompt_len]],
|
||||
dtype=torch.long,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
task_ids = [t.task_id for t in tasks]
|
||||
position_ids = torch.arange(
|
||||
start_pos, prompt_len, dtype=torch.long, device=self.device
|
||||
).repeat(batch_sz)
|
||||
|
||||
with (
|
||||
torch.inference_mode(),
|
||||
timed(f"execute_prefill b={batch_sz} prompt_len={prompt_len}", logger),
|
||||
):
|
||||
outputs = self.model(
|
||||
input_ids,
|
||||
position_ids=position_ids,
|
||||
kv_cache=self.task_cache.bind(
|
||||
task_ids,
|
||||
self._workspace,
|
||||
start_pos=start_pos,
|
||||
),
|
||||
fwd="prefill",
|
||||
)
|
||||
q_len = prompt_len - start_pos
|
||||
logits = outputs["logits"][
|
||||
torch.arange(1, batch_sz + 1, device=self.device) * q_len - 1
|
||||
]
|
||||
|
||||
return tasks, self._sample_logits(logits, tasks, return_logprobs)
|
||||
|
||||
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:
|
||||
return []
|
||||
|
||||
b = len(tasks)
|
||||
ws = self._workspace
|
||||
|
||||
# ---- pre-replay: update input buffers in-place ----
|
||||
|
||||
input_ids = ws.fill_input_ids(
|
||||
[t.output_ids[-1] if t.output_ids else t.prompt_ids[-1] for t in tasks]
|
||||
)
|
||||
|
||||
task_ids = [t.task_id for t in tasks]
|
||||
cur_positions = [t.next_pos for t in tasks]
|
||||
|
||||
kv_cache = self.task_cache.bind(task_ids, ws)
|
||||
|
||||
task_sig = tuple(task_ids)
|
||||
reuse_decode_state = (
|
||||
self.task_cache.bind_was_steady
|
||||
and self._decode_cache is not None
|
||||
and self._decode_cache.task_sig == task_sig
|
||||
)
|
||||
if reuse_decode_state:
|
||||
info = self._decode_cache.sampling_info
|
||||
ws.position_ids[:b] += 1
|
||||
else:
|
||||
info = _build_sampling_batch_info(tasks, self.device)
|
||||
ws.position_ids[:b].copy_(
|
||||
torch.tensor(cur_positions, dtype=torch.long, device=self.device)
|
||||
)
|
||||
self._decode_cache = DecodeSteadyState(task_sig, cur_positions, info)
|
||||
|
||||
# ---- forward (graph replay or live run + capture) ----
|
||||
|
||||
use_graph = (
|
||||
self._graph_ctx.enabled
|
||||
and self._graph_supported
|
||||
and get_backend().supports_graph()
|
||||
)
|
||||
key = (b,)
|
||||
with (
|
||||
torch.inference_mode(),
|
||||
timed(f"execute_decode forward b={b}", logger),
|
||||
):
|
||||
if use_graph:
|
||||
outputs = self._graph_ctx.forward(
|
||||
self.model,
|
||||
key=key,
|
||||
input_ids=input_ids,
|
||||
kv_cache=kv_cache,
|
||||
position_ids=ws.position_ids[:b],
|
||||
fwd="decode",
|
||||
)
|
||||
else:
|
||||
outputs = self.model(
|
||||
input_ids,
|
||||
kv_cache=kv_cache,
|
||||
position_ids=ws.position_ids[:b],
|
||||
fwd="decode",
|
||||
)
|
||||
logits = outputs["logits"]
|
||||
|
||||
return self._sample_logits(logits, tasks, return_logprobs, info=info)
|
||||
@@ -0,0 +1,103 @@
|
||||
"""CUDA-graph capture for the decode model-forward step.
|
||||
|
||||
Mirrors SGLang's cuda-graph manager: one graph per batch size. The graph
|
||||
pair. The graph captures ``model.forward()`` with workspace-backed inputs
|
||||
(all at fixed addresses). Before each replay the caller updates the input
|
||||
buffer content in-place so the graph sees fresh data at the same tensor
|
||||
addresses.
|
||||
|
||||
Only the model forward is captured — sampling runs outside the graph
|
||||
(via ``torch.multinomial`` which consumes a mutable RNG state).
|
||||
"""
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
class CudaGraphContext:
|
||||
"""CUDA-graph capture/replay for decode steps.
|
||||
|
||||
Parameters:
|
||||
enabled: When ``False``, ``forward()`` always runs the live model
|
||||
forward without capture/replay (graphs are cleared). Toggle at
|
||||
runtime via the ``set_enabled()`` method.
|
||||
|
||||
Usage::
|
||||
|
||||
gctx = CudaGraphContext()
|
||||
with torch.inference_mode():
|
||||
outputs = gctx.forward(
|
||||
model,
|
||||
key=(batch_size,),
|
||||
input_ids=workspace.input_ids[:b].unsqueeze(1),
|
||||
input_mask=input_mask,
|
||||
kv_cache=kv_cache,
|
||||
position_ids=workspace.position_ids[:b].unsqueeze(1),
|
||||
)
|
||||
|
||||
The first call at a given key runs *without* capture (warmup). The
|
||||
second call captures the graph. Subsequent calls replay the captured
|
||||
graph. A ``torch.cuda.synchronize()`` before capture drains in-flight
|
||||
work so the graph trace is clean.
|
||||
"""
|
||||
|
||||
def __init__(self, enabled: bool = False):
|
||||
self._enabled = enabled
|
||||
self._graphs: dict[tuple, torch.cuda.CUDAGraph] = {}
|
||||
self._outputs: dict[tuple, dict[str, Tensor]] = {}
|
||||
self._warmed: set[tuple] = set()
|
||||
|
||||
@property
|
||||
def enabled(self) -> bool:
|
||||
return self._enabled
|
||||
|
||||
def set_enabled(self, flag: bool):
|
||||
"""Enable or disable CUDA-graph capture at runtime.
|
||||
|
||||
Disabling clears all captured graphs (frees GPU memory) and warmup
|
||||
state. Re-enabling after disable starts fresh — graphs are
|
||||
re-captured on the next warmup cycle.
|
||||
"""
|
||||
if flag == self._enabled:
|
||||
return
|
||||
self._enabled = flag
|
||||
if not flag:
|
||||
self._graphs.clear()
|
||||
self._outputs.clear()
|
||||
self._warmed.clear()
|
||||
|
||||
def forward(self, model, *, key, **kwargs) -> dict[str, Tensor]:
|
||||
"""Run ``model(**kwargs)`` via graph replay or live forward.
|
||||
|
||||
Args:
|
||||
model: callable, e.g. ``self.model.forward``.
|
||||
key: ``(batch_size,)`` — the dispatch key (one graph per batch size).
|
||||
**kwargs: arguments forwarded to ``model``. All tensor arguments
|
||||
must reside at stable addresses (workspace buffers).
|
||||
|
||||
Returns:
|
||||
The dict produced by ``model(**kwargs)``, e.g.
|
||||
``{"logits": ..., "h0": ...}``.
|
||||
"""
|
||||
if not self._enabled:
|
||||
self._outputs[key] = model(**kwargs)
|
||||
return self._outputs[key]
|
||||
|
||||
if key in self._graphs:
|
||||
self._graphs[key].replay()
|
||||
elif key in self._warmed:
|
||||
cap_output = model(**kwargs)
|
||||
torch.cuda.synchronize()
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(graph):
|
||||
self._outputs[key] = model(**kwargs)
|
||||
self._graphs[key] = graph
|
||||
self._warmed.discard(key)
|
||||
return cap_output
|
||||
else:
|
||||
self._warmed.add(key)
|
||||
self._outputs[key] = model(**kwargs)
|
||||
return self._outputs[key]
|
||||
|
||||
def has_graph(self, key: tuple) -> bool:
|
||||
return key in self._graphs
|
||||
@@ -266,7 +266,7 @@ class SamplingPipeline(BaseSamplingStrategy):
|
||||
@staticmethod
|
||||
def _is_greedy(temperature: Union[float, Tensor]) -> bool:
|
||||
if isinstance(temperature, Tensor):
|
||||
return temperature.numel() == 1 and temperature.item() == 0
|
||||
return bool((temperature == 0).all())
|
||||
return temperature == 0
|
||||
|
||||
@torch.inference_mode()
|
||||
@@ -305,12 +305,12 @@ class SamplingPipeline(BaseSamplingStrategy):
|
||||
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)
|
||||
if not return_logprobs:
|
||||
return tokens
|
||||
log_probs = torch.log_softmax(transformed.float(), dim=-1)
|
||||
chosen = torch.gather(log_probs, -1, tokens.unsqueeze(-1)).squeeze(-1)
|
||||
return tokens, chosen
|
||||
|
||||
@@ -343,6 +343,10 @@ def sample(
|
||||
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:
|
||||
logits: Raw logits ``[batch, vocab_size]``.
|
||||
frequency_penalty: Penalty per occurrence for repeated tokens
|
||||
@@ -359,14 +363,21 @@ def sample(
|
||||
``True`` — a ``(token_ids, chosen_logprobs)`` tuple where
|
||||
``chosen_logprobs`` has shape ``[batch]``.
|
||||
"""
|
||||
return SamplingPipeline(
|
||||
[
|
||||
TemperatureStrategy(temperature),
|
||||
TopKStrategy(top_k),
|
||||
TopPStrategy(top_p),
|
||||
FrequencyPenaltyStrategy(frequency_penalty),
|
||||
]
|
||||
).sample(
|
||||
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,
|
||||
@@ -0,0 +1,408 @@
|
||||
import logging
|
||||
import threading
|
||||
import uuid
|
||||
from contextlib import nullcontext
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.extension import (
|
||||
ATTN_BACKEND,
|
||||
AttentionBackend,
|
||||
attn_backend,
|
||||
get_backend,
|
||||
)
|
||||
from astrai.inference.cache import PagePool, TaskCacheManager
|
||||
from astrai.inference.metrics import MetricsCollector
|
||||
from astrai.inference.runtime.executor import Executor
|
||||
from astrai.inference.task import STOP, Task, TaskManager, TaskStatus
|
||||
from astrai.model.automodel import AutoModel
|
||||
from astrai.tokenize.tokenizer import AutoTokenizer
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class InferenceScheduler:
|
||||
"""Continuous batching loop: cleanup -> refill -> prefill -> decode (all groups)."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: AutoModel,
|
||||
tokenizer: AutoTokenizer,
|
||||
max_batch_size: int = 16,
|
||||
max_seq_len: Optional[int] = None,
|
||||
device: Optional[str] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
cache: Optional[PagePool] = None,
|
||||
enable_cuda_graph: bool = True,
|
||||
backend: Optional[Union[str, ATTN_BACKEND, AttentionBackend, type]] = None,
|
||||
):
|
||||
config = model.config
|
||||
|
||||
if max_seq_len is not None:
|
||||
self.max_seq_len = max_seq_len
|
||||
elif config.max_position_embeddings is not None:
|
||||
self.max_seq_len = config.max_position_embeddings
|
||||
else:
|
||||
raise ValueError(
|
||||
"max_seq_len must be provided either as argument "
|
||||
"or in model config (config.max_position_embeddings)"
|
||||
)
|
||||
self.device = device or next(model.parameters()).device
|
||||
self.dtype = dtype or next(model.parameters()).dtype
|
||||
|
||||
head_dim = config.hidden_size // config.num_attention_heads
|
||||
|
||||
if cache is not None:
|
||||
self._cache = cache
|
||||
else:
|
||||
self._cache = PagePool(
|
||||
n_layers=config.num_hidden_layers,
|
||||
n_kv_heads=config.num_key_value_heads,
|
||||
head_dim=head_dim,
|
||||
max_batch_size=max_batch_size,
|
||||
max_seq_len=self.max_seq_len,
|
||||
device=self.device,
|
||||
dtype=self.dtype,
|
||||
)
|
||||
|
||||
self._metrics = MetricsCollector()
|
||||
|
||||
self._task_cache = TaskCacheManager(self._cache)
|
||||
|
||||
self._task_mgr = TaskManager(
|
||||
tokenizer=tokenizer,
|
||||
max_batch_size=max_batch_size,
|
||||
max_seq_len=self.max_seq_len,
|
||||
metrics=self._metrics,
|
||||
)
|
||||
|
||||
if backend is None:
|
||||
self._backend = None
|
||||
default_backend = get_backend()
|
||||
self._backend_name = type(default_backend).__name__
|
||||
with attn_backend(default_backend):
|
||||
self._executor = Executor(
|
||||
model=model,
|
||||
kv_cache=self._cache,
|
||||
task_cache=self._task_cache,
|
||||
device=self.device,
|
||||
dtype=self.dtype,
|
||||
enable_cuda_graph=enable_cuda_graph,
|
||||
)
|
||||
else:
|
||||
with attn_backend(backend):
|
||||
self._backend = get_backend()
|
||||
self._backend_name = type(self._backend).__name__
|
||||
self._executor = Executor(
|
||||
model=model,
|
||||
kv_cache=self._cache,
|
||||
task_cache=self._task_cache,
|
||||
device=self.device,
|
||||
dtype=self.dtype,
|
||||
enable_cuda_graph=enable_cuda_graph,
|
||||
)
|
||||
|
||||
self._stop_event = threading.Event()
|
||||
self._loop_thread: Optional[threading.Thread] = None
|
||||
|
||||
def add_task(self, prompt: str, **kwargs) -> str:
|
||||
return self._task_mgr.add_task(prompt, **kwargs)
|
||||
|
||||
def remove_task(self, task_id: str):
|
||||
for task in self._task_mgr.remove_task(task_id):
|
||||
self._task_cache.task_free(task.task_id)
|
||||
|
||||
def get_stats(self) -> Dict[str, Any]:
|
||||
return self._task_mgr.get_stats()
|
||||
|
||||
@property
|
||||
def backend_name(self) -> str:
|
||||
return self._backend_name
|
||||
|
||||
@property
|
||||
def cuda_graph_enabled(self) -> bool:
|
||||
return self._executor.cuda_graph_enabled
|
||||
|
||||
def _backend_context(self):
|
||||
if self._backend is None:
|
||||
return nullcontext()
|
||||
return attn_backend(self._backend)
|
||||
|
||||
@staticmethod
|
||||
def _task_backend_groups(tasks: List[Task]):
|
||||
groups = {}
|
||||
for task in tasks:
|
||||
groups.setdefault(task.backend, (task.backend, []))[1].append(task)
|
||||
return groups.values()
|
||||
|
||||
def _step(
|
||||
self, tasks: List[Task], return_logprobs: bool = False
|
||||
) -> Tuple[List[Task], List[Task]]:
|
||||
"""Advance every active task by one token (prefill + decode).
|
||||
|
||||
Single shared primitive for both the continuous-batching loop and
|
||||
the synchronous ``run_batch`` path, so the two cannot drift.
|
||||
|
||||
Tasks must already be allocated in the KV cache. Tasks without output
|
||||
are prefilled first and sample their first token from the final prompt
|
||||
position. Tasks with output extend the cache by one position and decode
|
||||
from their latest generated token.
|
||||
|
||||
Args:
|
||||
tasks: Active tasks to advance by one token.
|
||||
return_logprobs: Forwarded to ``execute_decode``; per-token
|
||||
logprobs are recorded on each task's ``output_logprobs``.
|
||||
|
||||
Returns:
|
||||
``(decoded, aborted)``: tasks that produced a new token (its ID
|
||||
already appended to ``output_ids``) and tasks that hit the
|
||||
sequence cap and were marked ``ABORTED``.
|
||||
"""
|
||||
to_prefill = [t for t in tasks if not t.prefill_done and t.prompt_ids]
|
||||
prefilled_ids = set()
|
||||
produced: List[Task] = []
|
||||
if to_prefill:
|
||||
for t in to_prefill:
|
||||
t.input_tokens = len(t.prompt_ids)
|
||||
|
||||
groups: Dict[Tuple[int, int, Optional[AttentionBackend]], List[Task]] = {}
|
||||
for t in to_prefill:
|
||||
start_pos = min(
|
||||
self._task_cache.task_cached(t.task_id), len(t.prompt_ids) - 1
|
||||
)
|
||||
groups.setdefault((len(t.prompt_ids), start_pos, t.backend), []).append(
|
||||
t
|
||||
)
|
||||
|
||||
for (prompt_len, start_pos, _), group in groups.items():
|
||||
backend = group[0].backend
|
||||
backend_context = (
|
||||
attn_backend(backend) if backend is not None else nullcontext()
|
||||
)
|
||||
with (
|
||||
backend_context,
|
||||
self._metrics.record([t.task_id for t in group], "prefill"),
|
||||
):
|
||||
prefilled, step_out = self._executor.execute_prefill(
|
||||
group, prompt_len, start_pos, return_logprobs=return_logprobs
|
||||
)
|
||||
|
||||
for t, out in zip(prefilled, step_out):
|
||||
t.output_ids.append(out[0] if return_logprobs else out)
|
||||
t.output_tokens += 1
|
||||
t.mark_prefill_done()
|
||||
prefilled_ids.add(t.task_id)
|
||||
produced.append(t)
|
||||
|
||||
start_logical_page = start_pos // self._cache.page_size
|
||||
for t in group:
|
||||
self._task_cache.task_record_hashes(
|
||||
t.task_id, t.prompt_ids, start_logical_page
|
||||
)
|
||||
|
||||
decoded: List[Task] = []
|
||||
aborted: List[Task] = []
|
||||
for t in tasks:
|
||||
if t.task_id in prefilled_ids:
|
||||
continue
|
||||
if self._task_cache.task_extend(t.task_id, t.next_pos):
|
||||
decoded.append(t)
|
||||
else:
|
||||
t.status = TaskStatus.ABORTED
|
||||
aborted.append(t)
|
||||
|
||||
for backend, group in self._task_backend_groups(decoded):
|
||||
backend_context = (
|
||||
attn_backend(backend) if backend is not None else nullcontext()
|
||||
)
|
||||
with (
|
||||
backend_context,
|
||||
self._metrics.record([t.task_id for t in group], "decode"),
|
||||
):
|
||||
step_out = self._executor.execute_decode(
|
||||
group, return_logprobs=return_logprobs
|
||||
)
|
||||
for t, out in zip(group, step_out):
|
||||
t.output_ids.append(out[0] if return_logprobs else out)
|
||||
t.output_tokens += 1
|
||||
t.advance_kv()
|
||||
produced.append(t)
|
||||
|
||||
return produced, aborted
|
||||
|
||||
def _run_generation_loop(self):
|
||||
stop_ids = self._task_mgr.tokenizer.stop_ids
|
||||
try:
|
||||
with self._backend_context():
|
||||
while not self._stop_event.is_set():
|
||||
finished = self._task_mgr.remove_finished_tasks(stop_ids)
|
||||
for task in finished:
|
||||
if task.status == TaskStatus.FINISHED:
|
||||
self._task_cache.task_record_hashes(
|
||||
task.task_id,
|
||||
self._task_cache.task_cacheable_ids(
|
||||
task.task_id, task.prompt_ids, task.output_ids
|
||||
),
|
||||
)
|
||||
self._task_cache.task_free(task.task_id)
|
||||
|
||||
active = self._task_mgr.get_active_tasks()
|
||||
available = self._task_mgr.max_batch_size - len(active)
|
||||
if available > 0:
|
||||
candidates = self._task_mgr.pull_candidates(available)
|
||||
failed = []
|
||||
for task in candidates:
|
||||
if self._task_cache.task_alloc(
|
||||
task.task_id, task.prompt_ids
|
||||
):
|
||||
self._task_mgr.activate(task)
|
||||
else:
|
||||
failed.append(task)
|
||||
if failed:
|
||||
self._task_mgr.return_to_waiting(failed)
|
||||
|
||||
if not self._task_mgr.has_work():
|
||||
self._task_mgr.wait_for_tasks(timeout=1.0)
|
||||
continue
|
||||
|
||||
active = self._task_mgr.get_active_tasks()
|
||||
|
||||
decoded, aborted = self._step(active)
|
||||
|
||||
for t in aborted:
|
||||
self._task_mgr.invoke_callback(t.task_id, STOP)
|
||||
|
||||
for t in decoded:
|
||||
new_text = t.decode_new_token(self._task_mgr.tokenizer)
|
||||
if new_text:
|
||||
self._task_mgr.invoke_callback(t.task_id, new_text)
|
||||
if t.is_finished(stop_ids):
|
||||
self._task_mgr.invoke_callback(t.task_id, STOP)
|
||||
|
||||
except Exception as e:
|
||||
self._stop_event.set()
|
||||
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
|
||||
for task in self._task_mgr.get_active_tasks():
|
||||
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||
self._task_cache.task_free(task.task_id)
|
||||
for task in self._task_mgr.get_waiting_tasks():
|
||||
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||
self._task_mgr.clear_queues()
|
||||
|
||||
def start(self):
|
||||
if self._loop_thread is not None and self._loop_thread.is_alive():
|
||||
return
|
||||
self._stop_event.clear()
|
||||
t = threading.Thread(target=self._run_generation_loop, daemon=True)
|
||||
t.start()
|
||||
self._loop_thread = t
|
||||
|
||||
def stop(self):
|
||||
self._stop_event.set()
|
||||
self._task_mgr.wake()
|
||||
if self._loop_thread is not None:
|
||||
self._loop_thread.join(timeout=2.0)
|
||||
self._loop_thread = None
|
||||
for task in self._task_mgr.get_active_tasks():
|
||||
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||
self._task_cache.task_free(task.task_id)
|
||||
for task in self._task_mgr.get_waiting_tasks():
|
||||
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||
self._task_cache.task_free(task.task_id)
|
||||
self._task_mgr.clear_queues()
|
||||
if torch.cuda.is_available():
|
||||
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
|
||||
seq_cap = self.max_seq_len
|
||||
request_backend = get_backend(use_default=False)
|
||||
|
||||
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))
|
||||
if t_max <= 0:
|
||||
tasks.append(None)
|
||||
continue
|
||||
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,
|
||||
backend=request_backend,
|
||||
)
|
||||
if not self._task_cache.task_alloc(task.task_id, task.prompt_ids):
|
||||
tasks.append(None)
|
||||
continue
|
||||
task.input_tokens = len(task.prompt_ids)
|
||||
self._metrics.register(task.task_id)
|
||||
tasks.append(task)
|
||||
|
||||
try:
|
||||
live = [t for t in tasks if t is not None]
|
||||
|
||||
with self._backend_context():
|
||||
while live:
|
||||
decoded, _ = self._step(live, return_logprobs=return_logprobs)
|
||||
live = [t for t in decoded if not t.is_finished(stop_ids)]
|
||||
finally:
|
||||
for t in tasks:
|
||||
if t is not None:
|
||||
self._metrics.mark_finished(
|
||||
t.task_id, t.input_tokens, t.output_tokens
|
||||
)
|
||||
self._task_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
|
||||
@@ -1,50 +1,46 @@
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from collections import deque
|
||||
from enum import Enum
|
||||
from typing import Any, Callable, Deque, Dict, List, Optional
|
||||
from typing import TYPE_CHECKING, Any, Callable, Deque, Dict, List, Optional
|
||||
|
||||
from tokenizers.decoders import DecodeStream
|
||||
|
||||
from astrai.inference.metrics import MetricsCollector
|
||||
from astrai.tokenize.tokenizer import AutoTokenizer
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
if TYPE_CHECKING:
|
||||
from astrai.extension import AttentionBackend
|
||||
|
||||
STOP = object()
|
||||
|
||||
|
||||
class StreamDecoder:
|
||||
"""Incremental decoder for byte-level BPE streaming.
|
||||
"""Incremental decoder backed by the tokenizers library's DecodeStream.
|
||||
|
||||
Byte-level BPE may split a single Unicode character (e.g. em-dash,
|
||||
smart quotes) across multiple tokens. Decoding such a token in
|
||||
isolation produces U+FFFD (replacement char). This decoder
|
||||
accumulates token IDs and only emits text once the trailing
|
||||
characters are complete, buffering incomplete multi-byte sequences
|
||||
until the next token arrives.
|
||||
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__ = ("_tokenizer", "_ids", "_emitted")
|
||||
__slots__ = ("_stream", "_tok")
|
||||
|
||||
def __init__(self, tokenizer: AutoTokenizer):
|
||||
self._tokenizer = tokenizer
|
||||
self._ids: List[int] = []
|
||||
self._emitted: str = ""
|
||||
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.
|
||||
"""
|
||||
self._ids.append(token_id)
|
||||
full = self._tokenizer.decode(self._ids, skip_special_tokens=True)
|
||||
if full.endswith("\ufffd"):
|
||||
return ""
|
||||
if len(full) > len(self._emitted):
|
||||
diff = full[len(self._emitted) :]
|
||||
self._emitted = full
|
||||
return diff
|
||||
return ""
|
||||
chunk = self._stream.step(self._tok, token_id)
|
||||
return chunk or ""
|
||||
|
||||
|
||||
class TaskStatus(Enum):
|
||||
@@ -69,6 +65,7 @@ class Task:
|
||||
top_k: int = 50,
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
backend: Optional["AttentionBackend"] = None,
|
||||
):
|
||||
self.task_id = task_id
|
||||
self.prompt_ids = prompt_ids
|
||||
@@ -78,16 +75,25 @@ class Task:
|
||||
self.top_k = top_k
|
||||
self.frequency_penalty = frequency_penalty
|
||||
self.rep_window = rep_window
|
||||
self.backend = backend
|
||||
|
||||
self.status = TaskStatus.PENDING
|
||||
self.output_ids: List[int] = []
|
||||
self.output_logprobs: List[float] = []
|
||||
self.input_tokens: int = 0
|
||||
self.output_tokens: int = 0
|
||||
self.arrival_time = time.time()
|
||||
self.finish_time: Optional[float] = None
|
||||
self._kv_len: int = 0
|
||||
self._decoder: Optional[StreamDecoder] = None
|
||||
|
||||
def mark_prefill_done(self):
|
||||
"""Prompt KV is materialized by prefill; first output sampled but
|
||||
not yet written to KV."""
|
||||
self._kv_len = self.input_tokens
|
||||
|
||||
def advance_kv(self):
|
||||
"""One more position written to KV (after a decode forward)."""
|
||||
self._kv_len += 1
|
||||
|
||||
def decode_new_token(self, tokenizer: AutoTokenizer) -> str:
|
||||
"""Decode the last appended output token, buffering incomplete
|
||||
multi-byte sequences across calls.
|
||||
@@ -98,26 +104,15 @@ class Task:
|
||||
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.
|
||||
|
||||
Called when generation terminates (max_tokens reached, stop
|
||||
sequence, or external removal) to avoid dropping a final
|
||||
incomplete-looking fragment that is actually complete when
|
||||
adjacent to the stop token.
|
||||
"""
|
||||
if self._decoder is None or not self.output_ids:
|
||||
return ""
|
||||
full = tokenizer.decode(self.output_ids, skip_special_tokens=True)
|
||||
if len(full) > len(self._decoder._emitted):
|
||||
diff = full[len(self._decoder._emitted) :]
|
||||
self._decoder._emitted = full
|
||||
return diff
|
||||
return ""
|
||||
|
||||
@property
|
||||
def next_pos(self) -> int:
|
||||
return self.input_tokens + len(self.output_ids)
|
||||
"""KV position where the next decode step will write."""
|
||||
return self._kv_len
|
||||
|
||||
@property
|
||||
def prefill_done(self) -> bool:
|
||||
"""True when all prompt KV entries are materialized."""
|
||||
return self._kv_len >= self.input_tokens > 0
|
||||
|
||||
def is_finished(self, stop_ids: List[int]) -> bool:
|
||||
if self.max_tokens is not None and self.output_tokens >= self.max_tokens:
|
||||
@@ -135,12 +130,11 @@ class TaskManager:
|
||||
tokenizer: AutoTokenizer,
|
||||
max_batch_size: int = 16,
|
||||
max_seq_len: int = 8192,
|
||||
max_prompt_len: int = 512,
|
||||
metrics: Optional["MetricsCollector"] = None,
|
||||
):
|
||||
self.tokenizer = tokenizer
|
||||
self.max_batch_size = max_batch_size
|
||||
self.max_seq_len = max_seq_len
|
||||
self.max_prompt_len = max_prompt_len
|
||||
|
||||
self.waiting_queue: Deque[Task] = deque()
|
||||
self.active_tasks: List[Task] = []
|
||||
@@ -152,6 +146,8 @@ class TaskManager:
|
||||
self._total_tasks = 0
|
||||
self._total_tokens = 0
|
||||
|
||||
self._metrics = metrics
|
||||
|
||||
def add_task(
|
||||
self,
|
||||
prompt: str,
|
||||
@@ -161,17 +157,13 @@ class TaskManager:
|
||||
top_k: int = 50,
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
backend: Optional["AttentionBackend"] = None,
|
||||
stream_callback: Optional[Callable[[str], None]] = None,
|
||||
) -> str:
|
||||
task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}"
|
||||
prompt_ids = self.tokenizer.encode(prompt)
|
||||
if len(prompt_ids) > self.max_prompt_len:
|
||||
prompt_ids = prompt_ids[-self.max_prompt_len :]
|
||||
|
||||
if len(prompt_ids) >= self.max_seq_len:
|
||||
if stream_callback:
|
||||
stream_callback(STOP)
|
||||
return task_id
|
||||
if len(prompt_ids) > self.max_seq_len:
|
||||
prompt_ids = prompt_ids[-self.max_seq_len :]
|
||||
|
||||
if max_tokens is None:
|
||||
max_tokens = self.max_seq_len - len(prompt_ids)
|
||||
@@ -187,6 +179,7 @@ class TaskManager:
|
||||
top_k=top_k,
|
||||
frequency_penalty=frequency_penalty,
|
||||
rep_window=rep_window,
|
||||
backend=backend,
|
||||
)
|
||||
|
||||
with self._lock:
|
||||
@@ -195,6 +188,9 @@ class TaskManager:
|
||||
if stream_callback:
|
||||
self._callbacks[task_id] = stream_callback
|
||||
|
||||
if self._metrics is not None:
|
||||
self._metrics.register(task_id)
|
||||
|
||||
self._task_event.set()
|
||||
return task_id
|
||||
|
||||
@@ -214,26 +210,33 @@ class TaskManager:
|
||||
cb(token)
|
||||
|
||||
def get_stats(self) -> Dict[str, Any]:
|
||||
return {
|
||||
stats: Dict[str, Any] = {
|
||||
"total_tasks": self._total_tasks,
|
||||
"total_tokens": self._total_tokens,
|
||||
"active_tasks": len(self.active_tasks),
|
||||
"waiting_queue": len(self.waiting_queue),
|
||||
}
|
||||
if self._metrics is not None:
|
||||
stats.update(self._metrics.get_stats())
|
||||
return stats
|
||||
|
||||
def remove_finished_tasks(self, stop_ids: List[int]) -> List[Task]:
|
||||
with self._lock:
|
||||
finished = []
|
||||
for task in self.active_tasks:
|
||||
if task.status == TaskStatus.ABORTED:
|
||||
task.finish_time = time.time()
|
||||
finished.append(task)
|
||||
elif task.is_finished(stop_ids):
|
||||
task.status = TaskStatus.FINISHED
|
||||
task.finish_time = time.time()
|
||||
finished.append(task)
|
||||
self._total_tokens += task.output_tokens
|
||||
|
||||
if self._metrics is not None:
|
||||
for task in finished:
|
||||
self._metrics.mark_finished(
|
||||
task.task_id, task.input_tokens, task.output_tokens
|
||||
)
|
||||
|
||||
self.active_tasks = [
|
||||
t
|
||||
for t in self.active_tasks
|
||||
@@ -0,0 +1,151 @@
|
||||
"""Pre-allocated buffers for the inference decode hot path.
|
||||
|
||||
Mirrors FlashInfer / SGLang's global workspace pattern: all per-step tensors
|
||||
are allocated eagerly at init (nothing is lazy), so the decode step
|
||||
reads/writes fixed-address tensors with zero ``torch.empty`` calls during
|
||||
the hot loop — a prerequisite for CUDA-graph capture.
|
||||
"""
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
_MAX_SPLITS = 32
|
||||
|
||||
|
||||
class InferenceWorkspace:
|
||||
"""Reusable fixed-shape per-step buffers for decode.
|
||||
|
||||
Families of buffers, all sized to ``max_batch_size`` / ``max_seq_len``
|
||||
and sliced via views each step:
|
||||
|
||||
- ``decode_mask``: a ``[B, 1, total_len]`` validity mask, the RHS
|
||||
``arange`` pre-computed so only a single ``torch.ge(out=)`` runs per
|
||||
step.
|
||||
- ``input_ids``: per-step token IDs filled from host (pinned, double-
|
||||
buffered so an in-flight async H2D copy never races the next fill).
|
||||
- KV-cache bind metadata (``req_pool_indices``, ``seq_lens``,
|
||||
``kv_indptr``, ``inc``, ``out_cache_loc``), written by
|
||||
``PagePool.bind_tasks`` when the Executor passes this workspace.
|
||||
- ``decode_o_part`` / ``decode_ml_part``: split-KV partial result buffers
|
||||
(mirrors FlashInfer's workspace). One global alloc, reused by every
|
||||
decode step across all layers. Sliced views are passed to the CUDA
|
||||
attention kernel so its internal ``torch.empty`` hot-path alloc goes
|
||||
through a stable address (CUDA-graph capturable).
|
||||
|
||||
No re-allocation while the server's bounds are respected.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
max_batch_size: int,
|
||||
max_seq_len: int,
|
||||
max_q_heads: int,
|
||||
head_dim: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
):
|
||||
self.max_batch_size = max_batch_size
|
||||
self.max_seq_len = max_seq_len
|
||||
self.max_q_heads = max_q_heads
|
||||
self.head_dim = head_dim
|
||||
self.device = device
|
||||
self.dtype = dtype
|
||||
|
||||
# ``position_ids[:, None, None] >= arange`` RHS, reused every step.
|
||||
self.arange = torch.arange(max_seq_len, device=device)
|
||||
# Decode validity mask: [max_batch, 1, max_seq_len] bool.
|
||||
self.input_mask = torch.empty(
|
||||
(max_batch_size, 1, max_seq_len), dtype=torch.bool, device=device
|
||||
)
|
||||
|
||||
# Per-step token IDs. Values come from host Python lists every
|
||||
# step, so the device buffer is pre-allocated (stable address for
|
||||
# CUDA-graph capture) and filled via a host staging buffer. A
|
||||
# double buffer keeps a copy in flight from being overwritten by
|
||||
# the next fill.
|
||||
self.input_ids = torch.empty((max_batch_size,), dtype=torch.long, device=device)
|
||||
self._pin = [
|
||||
torch.empty((max_batch_size,), dtype=torch.long),
|
||||
torch.empty((max_batch_size,), dtype=torch.long),
|
||||
]
|
||||
self._pin_idx = 0
|
||||
|
||||
# KV-cache bind metadata (fixed shape, written by ``PagePool.bind_tasks``
|
||||
# when the Executor passes this workspace). Stable addresses make the
|
||||
# decode forward CUDA-graph capturable.
|
||||
self.req_pool_indices = torch.empty(
|
||||
(max_batch_size,), dtype=torch.int32, device=device
|
||||
)
|
||||
self.seq_lens = torch.empty((max_batch_size,), dtype=torch.long, device=device)
|
||||
self.kv_indptr = torch.empty(
|
||||
(max_batch_size + 1,), dtype=torch.int32, device=device
|
||||
)
|
||||
self.qo_indptr = torch.empty(
|
||||
(max_batch_size + 1,), dtype=torch.int32, device=device
|
||||
)
|
||||
self.inc = torch.arange(max_batch_size + 1, dtype=torch.int32, device=device)
|
||||
self.out_cache_loc = torch.empty(
|
||||
(max_batch_size, 1), dtype=torch.int32, device=device
|
||||
)
|
||||
|
||||
# Per-step position IDs (must be at a fixed address for CUDA-graph capture).
|
||||
self.position_ids = torch.empty(
|
||||
(max_batch_size,), dtype=torch.long, device=device
|
||||
)
|
||||
|
||||
# Split-KV partial-result buffers for decode (persistent, one global
|
||||
# alloc per process — mirrors FlashInfer's workspace pattern).
|
||||
# Shape: [max_batch_size, max_q_heads, _MAX_SPLITS, head_dim] (o_part)
|
||||
# [max_batch_size, max_q_heads, _MAX_SPLITS, 2] (ml_part)
|
||||
self.decode_o_part = torch.empty(
|
||||
(max_batch_size, max_q_heads, _MAX_SPLITS, head_dim),
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
)
|
||||
self.decode_ml_part = torch.empty(
|
||||
(max_batch_size, max_q_heads, _MAX_SPLITS, 2),
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
)
|
||||
|
||||
# Decode output buffer (graph-safe pre-alloc). Shape matches the
|
||||
# decode kernel's output: [batch, q_head, head_dim].
|
||||
self.decode_out = torch.empty(
|
||||
(max_batch_size, max_q_heads, head_dim),
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
)
|
||||
|
||||
def decode_buffers(self, batch: int, q_heads: int):
|
||||
"""Return ``(o_part, ml_part)`` view sliced to live dimensions."""
|
||||
return (
|
||||
self.decode_o_part[:batch, :q_heads],
|
||||
self.decode_ml_part[:batch, :q_heads],
|
||||
)
|
||||
|
||||
def fill_input_ids(self, ids: "list[int]") -> Tensor:
|
||||
"""Write ``ids`` into the device buffer and return ``[B]``.
|
||||
|
||||
Host values are staged through the double buffer and copied into the
|
||||
stable device buffer (``copy_`` without pinning is synchronous, so
|
||||
the alternating buffers guard against an in-flight transfer).
|
||||
"""
|
||||
b = len(ids)
|
||||
pin = self._pin[self._pin_idx]
|
||||
self._pin_idx ^= 1
|
||||
for i, v in enumerate(ids):
|
||||
pin[i] = v
|
||||
self.input_ids[:b].copy_(pin[:b])
|
||||
return self.input_ids[:b]
|
||||
|
||||
def decode_mask(self, position_ids: Tensor, total_len: int) -> Tensor:
|
||||
"""Return the ``[B, 1, total_len]`` validity mask for this step.
|
||||
|
||||
Written into the pre-allocated buffer via ``torch.ge(out=)`` — no
|
||||
new tensor is allocated. ``position_ids`` is the current step's
|
||||
``[B]`` positions; ``total_len`` must not exceed ``max_seq_len``.
|
||||
"""
|
||||
b = position_ids.size(0)
|
||||
out = self.input_mask[:b, :, :total_len]
|
||||
torch.ge(position_ids[:, None, None], self.arange[:total_len], out=out)
|
||||
return out
|
||||
@@ -0,0 +1,27 @@
|
||||
import logging
|
||||
import os
|
||||
|
||||
|
||||
def setup_logging(level: str = "INFO"):
|
||||
"""Attach a StreamHandler to the ``astrai`` logger (idempotent).
|
||||
|
||||
Call once per process at the top of CLI scripts.
|
||||
Set ``ASTR_LOG_LEVEL`` env var to override the default level.
|
||||
|
||||
Level names: ``DEBUG``, ``INFO``, ``WARNING``, ``ERROR``, ``CRITICAL``.
|
||||
``DEBUG`` enables per-step prefill/decode timing logs
|
||||
(:func:`astrai.inference.runtime.executor.timed`).
|
||||
"""
|
||||
logger = logging.getLogger("astrai")
|
||||
if logger.handlers:
|
||||
return
|
||||
level_name = os.environ.get("ASTR_LOG_LEVEL", level).upper()
|
||||
logger.setLevel(getattr(logging, level_name, logging.INFO))
|
||||
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)
|
||||
@@ -9,7 +9,7 @@ from astrai.model.components.lora import (
|
||||
merge_lora,
|
||||
save_lora,
|
||||
)
|
||||
from astrai.model.components.mlp import MLP
|
||||
from astrai.model.components.mlp import MLP, DeepSeekMoE
|
||||
from astrai.model.components.norm import RMSNorm
|
||||
from astrai.model.encoder import EmbeddingEncoder
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
@@ -19,6 +19,7 @@ __all__ = [
|
||||
"Linear",
|
||||
"RMSNorm",
|
||||
"MLP",
|
||||
"DeepSeekMoE",
|
||||
"GQA",
|
||||
"DecoderBlock",
|
||||
# Models
|
||||
|
||||
@@ -40,11 +40,12 @@ def _disable_random_init(enable: bool = True):
|
||||
setattr(nn.init, n, fn)
|
||||
|
||||
|
||||
class AutoModel(BaseFactory["AutoModel"], nn.Module):
|
||||
"""
|
||||
Autoregressive language model base class.
|
||||
Provides model loading/saving, registration, and generation.
|
||||
"""
|
||||
class ModelFactory(BaseFactory[nn.Module]):
|
||||
"""Pure factory for model dispatch, separated from nn.Module state."""
|
||||
|
||||
|
||||
class AutoModel(nn.Module):
|
||||
"""Model base class with loading/saving and generation."""
|
||||
|
||||
def __init__(self, config: BaseModelConfig):
|
||||
super().__init__()
|
||||
@@ -68,7 +69,7 @@ class AutoModel(BaseFactory["AutoModel"], nn.Module):
|
||||
config = ConfigFactory.load(raw)
|
||||
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):
|
||||
model = actual_cls(config)
|
||||
|
||||
@@ -1,12 +1,12 @@
|
||||
from astrai.model.components.attention import GQA, MLA, repeat_kv
|
||||
from astrai.extension.backend.rotary import apply_rotary_emb
|
||||
from astrai.model.components.attention import GQA, MLA
|
||||
from astrai.model.components.decoder_block import DecoderBlock
|
||||
from astrai.model.components.embedding import Embedding
|
||||
from astrai.model.components.linear import Linear
|
||||
from astrai.model.components.mlp import MLP
|
||||
from astrai.model.components.mlp import MLP, DeepSeekMoE
|
||||
from astrai.model.components.norm import RMSNorm
|
||||
from astrai.model.components.rope import (
|
||||
RotaryEmbedding,
|
||||
apply_rotary_emb,
|
||||
get_rotary_emb,
|
||||
)
|
||||
|
||||
@@ -14,6 +14,7 @@ __all__ = [
|
||||
"Linear",
|
||||
"RMSNorm",
|
||||
"MLP",
|
||||
"DeepSeekMoE",
|
||||
"Embedding",
|
||||
"GQA",
|
||||
"MLA",
|
||||
@@ -21,5 +22,4 @@ __all__ = [
|
||||
"RotaryEmbedding",
|
||||
"apply_rotary_emb",
|
||||
"get_rotary_emb",
|
||||
"repeat_kv",
|
||||
]
|
||||
|
||||
@@ -5,22 +5,11 @@ import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.extension.backend import apply_rotary_emb, attention
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.inference.core.cache import CacheView
|
||||
from astrai.inference.cache import KVCache
|
||||
from astrai.model.components.linear import Linear
|
||||
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]):
|
||||
@@ -66,17 +55,16 @@ class GQA(nn.Module):
|
||||
self.gate = Linear(dim, dim)
|
||||
|
||||
def _split_heads(self, x: Tensor, n_heads) -> Tensor:
|
||||
batch_size, seq_len, _ = x.shape
|
||||
x = x.reshape(batch_size, seq_len, n_heads, self.head_dim)
|
||||
return x
|
||||
return x.reshape(*x.shape[:-1], n_heads, self.head_dim)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: Tensor,
|
||||
rotary_emb: Tensor,
|
||||
attn_mask: Tensor = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
kv_cache: Optional[KVCache] = None,
|
||||
is_causal: bool = False,
|
||||
fwd: Optional[str] = None,
|
||||
) -> Tensor:
|
||||
q = self._split_heads(self.q_proj(x), self.n_heads)
|
||||
k = self._split_heads(self.k_proj(x), self.n_kv_heads)
|
||||
@@ -86,19 +74,9 @@ class GQA(nn.Module):
|
||||
if self.use_qk_norm:
|
||||
q, k = self.q_norm(q), self.k_norm(k)
|
||||
|
||||
if paged_cache is not None:
|
||||
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)
|
||||
)
|
||||
sdqa_out = attention(
|
||||
q, k, v, kv_cache, self.layer_id, attn_mask, is_causal, fwd
|
||||
).reshape(*x.shape[:-1], self.dim)
|
||||
|
||||
if self.use_gated_attention:
|
||||
sdqa_out = sdqa_out * F.sigmoid(self.gate(x))
|
||||
@@ -161,19 +139,18 @@ class MLA(nn.Module):
|
||||
x: Tensor,
|
||||
rotary_emb: Tensor,
|
||||
attn_mask: Tensor = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
kv_cache: Optional[KVCache] = None,
|
||||
is_causal: bool = False,
|
||||
fwd: Optional[str] = None,
|
||||
) -> Tensor:
|
||||
bsz, seq_len, _ = x.size()
|
||||
|
||||
q = self.q_proj(x)
|
||||
q = q.view(bsz, seq_len, self.n_heads, self.head_dim)
|
||||
q = q.reshape(*x.shape[:-1], self.n_heads, self.head_dim)
|
||||
|
||||
kv_compressed = self.kv_a_proj(x)
|
||||
kv_compressed = self.kv_norm(kv_compressed)
|
||||
|
||||
kv = self.kv_b_proj(kv_compressed)
|
||||
kv = kv.view(bsz, seq_len, self.n_kv_heads, -1)
|
||||
kv = kv.reshape(*x.shape[:-1], self.n_kv_heads, -1)
|
||||
|
||||
k_nope, k_rope, v = torch.split(
|
||||
kv, [self.qk_nope_head_dim, self.qk_rope_head_dim, self.head_dim], dim=-1
|
||||
@@ -193,18 +170,9 @@ class MLA(nn.Module):
|
||||
q = self.q_norm(q)
|
||||
k = self.k_norm(k)
|
||||
|
||||
if paged_cache is not None:
|
||||
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)
|
||||
attn_out = attention(
|
||||
q, k, v, kv_cache, self.layer_id, attn_mask, is_causal, fwd
|
||||
).reshape(*x.shape[:-1], self.dim)
|
||||
|
||||
if self.use_gated_attention:
|
||||
attn_out = attn_out * F.sigmoid(self.gate(x))
|
||||
|
||||
@@ -1,15 +1,21 @@
|
||||
from dataclasses import asdict
|
||||
from typing import Optional
|
||||
from typing import Optional, TypedDict
|
||||
|
||||
import torch.nn as nn
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.inference.core.cache import CacheView
|
||||
from astrai.inference.cache import KVCache
|
||||
from astrai.model.components.attention import AttnFactory
|
||||
from astrai.model.components.mlp import FFNFactory
|
||||
from astrai.model.components.mlp import FFNFactory, RouterStats
|
||||
from astrai.model.components.norm import RMSNorm
|
||||
|
||||
|
||||
class DecoderOutput(TypedDict):
|
||||
hidden_states: Tensor
|
||||
aux_loss: Optional[Tensor]
|
||||
router_stats: Optional[RouterStats]
|
||||
|
||||
|
||||
class DecoderBlock(nn.Module):
|
||||
def __init__(self, config, layer_id: int):
|
||||
super().__init__()
|
||||
@@ -26,24 +32,45 @@ class DecoderBlock(nn.Module):
|
||||
self.attention = AttnFactory.create(config.attn_type, **cfg, layer_id=layer_id)
|
||||
self.input_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
||||
self.post_attention_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
||||
self.mlp = FFNFactory.create(config.ffn_type, **cfg)
|
||||
ffn_type = self._resolve_ffn_type(config, layer_id)
|
||||
self.mlp = FFNFactory.create(ffn_type, **cfg)
|
||||
|
||||
@staticmethod
|
||||
def _resolve_ffn_type(config, layer_id: int) -> str:
|
||||
if config.ffn_type != "moe":
|
||||
return config.ffn_type
|
||||
mlp_only = config.mlp_only_layers or []
|
||||
if layer_id in mlp_only:
|
||||
return "mlp"
|
||||
if config.decoder_sparse_step > 1:
|
||||
if (layer_id + 1) % config.decoder_sparse_step != 0:
|
||||
return "mlp"
|
||||
return "moe"
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: Tensor,
|
||||
rotary_emb: Tensor,
|
||||
attention_mask: Optional[Tensor] = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
kv_cache: Optional[KVCache] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
fwd: Optional[str] = None,
|
||||
) -> DecoderOutput:
|
||||
attn_output = self.attention(
|
||||
self.input_norm(x),
|
||||
rotary_emb,
|
||||
attention_mask,
|
||||
paged_cache,
|
||||
kv_cache,
|
||||
is_causal,
|
||||
fwd,
|
||||
)
|
||||
x = attn_output + x
|
||||
x = self.mlp(self.post_attention_norm(x)) + x
|
||||
normalized = self.post_attention_norm(x)
|
||||
mlp_output = self.mlp(normalized)
|
||||
x = mlp_output["hidden_states"] + x
|
||||
|
||||
return x
|
||||
return {
|
||||
"hidden_states": x,
|
||||
"aux_loss": mlp_output["aux_loss"],
|
||||
"router_stats": mlp_output.get("router_stats"),
|
||||
}
|
||||
|
||||
@@ -1,11 +1,12 @@
|
||||
import logging
|
||||
from dataclasses import asdict, dataclass
|
||||
from dataclasses import asdict
|
||||
from pathlib import Path
|
||||
from typing import Optional, Set
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from pydantic.dataclasses import dataclass
|
||||
|
||||
from astrai.model.components.linear import Linear
|
||||
from astrai.serialization import (
|
||||
|
||||
+100
-22
@@ -1,3 +1,5 @@
|
||||
from typing import Optional, TypedDict
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
@@ -11,6 +13,28 @@ class FFNFactory(BaseFactory[nn.Module]):
|
||||
pass
|
||||
|
||||
|
||||
class RouterStats(TypedDict):
|
||||
"""Per-layer MoE routing statistics for training diagnostics.
|
||||
|
||||
Both tensors are detached monitoring data produced during forward.
|
||||
"""
|
||||
|
||||
probs: Tensor
|
||||
topk_indices: Tensor
|
||||
|
||||
|
||||
class FFNOutput(TypedDict):
|
||||
hidden_states: Tensor
|
||||
aux_loss: Optional[Tensor]
|
||||
router_stats: Optional[RouterStats]
|
||||
|
||||
|
||||
class RoutedOutput(TypedDict):
|
||||
hidden_states: Tensor
|
||||
aux_loss: Optional[Tensor]
|
||||
router_stats: Optional[RouterStats]
|
||||
|
||||
|
||||
@FFNFactory.register("mlp")
|
||||
class MLP(nn.Module):
|
||||
def __init__(self, dim: int, dim_ffn: int, down_init_std: float = 0.02):
|
||||
@@ -19,10 +43,10 @@ class MLP(nn.Module):
|
||||
self.gate = Linear(dim, dim_ffn)
|
||||
self.down = Linear(dim_ffn, dim, init_std=down_init_std)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
def forward(self, x: Tensor) -> FFNOutput:
|
||||
gated = self.up(x) * F.silu(self.gate(x))
|
||||
out = self.down(gated)
|
||||
return out
|
||||
return {"hidden_states": out, "aux_loss": None, "router_stats": None}
|
||||
|
||||
|
||||
@FFNFactory.register("moe")
|
||||
@@ -36,6 +60,9 @@ class DeepSeekMoE(nn.Module):
|
||||
n_activated_experts: int = 2,
|
||||
topk_method: str = "greedy",
|
||||
n_layers: int = 1,
|
||||
moe_intermediate_size: Optional[int] = None,
|
||||
shared_expert_intermediate_size: Optional[int] = None,
|
||||
norm_topk_prob: bool = True,
|
||||
):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
@@ -43,6 +70,16 @@ class DeepSeekMoE(nn.Module):
|
||||
self.n_shared_experts = n_shared_experts
|
||||
self.n_activated_experts = n_activated_experts
|
||||
self.topk_method = topk_method
|
||||
self.norm_topk_prob = norm_topk_prob
|
||||
|
||||
expert_dim_ffn = (
|
||||
moe_intermediate_size if moe_intermediate_size is not None else dim_ffn
|
||||
)
|
||||
shared_dim_ffn = (
|
||||
shared_expert_intermediate_size
|
||||
if shared_expert_intermediate_size is not None
|
||||
else dim_ffn
|
||||
)
|
||||
|
||||
self.router = Linear(dim, n_routed_experts, bias=False)
|
||||
moe_scale = 1 / max(n_shared_experts, 1) + 1 / n_activated_experts
|
||||
@@ -50,51 +87,92 @@ class DeepSeekMoE(nn.Module):
|
||||
|
||||
self.shared_experts = nn.ModuleList(
|
||||
[
|
||||
MLP(dim, dim_ffn, down_init_std=down_init_std)
|
||||
MLP(dim, shared_dim_ffn, down_init_std=down_init_std)
|
||||
for _ in range(n_shared_experts)
|
||||
]
|
||||
)
|
||||
self.routed_experts = nn.ModuleList(
|
||||
[
|
||||
MLP(dim, dim_ffn, down_init_std=down_init_std)
|
||||
MLP(dim, expert_dim_ffn, down_init_std=down_init_std)
|
||||
for _ in range(n_routed_experts)
|
||||
]
|
||||
)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
bsz, seq_len, dim = x.shape
|
||||
def forward(self, x: Tensor) -> FFNOutput:
|
||||
include_aux_loss = self.training and torch.is_grad_enabled()
|
||||
shape = x.shape
|
||||
dim = shape[-1]
|
||||
x_flat = x.view(-1, dim)
|
||||
|
||||
shared_out = self._shared_forward(x_flat)
|
||||
routed_out = self._routed_forward(x_flat)
|
||||
routed_output = self._routed_forward(x_flat, include_aux_loss)
|
||||
|
||||
out = (shared_out + routed_out).view(bsz, seq_len, dim)
|
||||
return out
|
||||
out = (shared_out + routed_output["hidden_states"]).view(shape)
|
||||
return {
|
||||
"hidden_states": out,
|
||||
"aux_loss": routed_output["aux_loss"],
|
||||
"router_stats": routed_output["router_stats"],
|
||||
}
|
||||
|
||||
def _shared_forward(self, x: Tensor) -> Tensor:
|
||||
if self.n_shared_experts == 0:
|
||||
return torch.zeros_like(x)
|
||||
return sum(e(x) for e in self.shared_experts) / self.n_shared_experts
|
||||
return (
|
||||
sum(e(x)["hidden_states"] for e in self.shared_experts)
|
||||
/ self.n_shared_experts
|
||||
)
|
||||
|
||||
def _routed_forward(self, x: Tensor) -> Tensor:
|
||||
def _routed_forward(self, x: Tensor, include_aux_loss: bool) -> RoutedOutput:
|
||||
N, D = x.shape
|
||||
K = self.n_activated_experts
|
||||
E = self.n_routed_experts
|
||||
|
||||
router_logits = self.router(x)
|
||||
router_probs = torch.softmax(router_logits.float(), dim=-1).to(x.dtype)
|
||||
|
||||
topk_weights, topk_indices = torch.topk(router_probs, K, dim=-1)
|
||||
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
|
||||
topk_weights, topk_indices = torch.topk(router_probs, K, dim=-1, sorted=False)
|
||||
if self.norm_topk_prob:
|
||||
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
|
||||
|
||||
aux_loss = None
|
||||
router_stats = None
|
||||
if include_aux_loss:
|
||||
expert_load = F.one_hot(topk_indices, num_classes=E).float()
|
||||
expert_load = expert_load.mean(dim=(0, 1))
|
||||
router_prob = router_probs.float().mean(dim=0)
|
||||
aux_loss = E * (expert_load * router_prob).sum()
|
||||
router_stats = {
|
||||
"probs": router_probs.detach(),
|
||||
"topk_indices": topk_indices,
|
||||
}
|
||||
|
||||
# Grouped dispatch: sort (token, slot) pairs by expert so each expert
|
||||
# consumes one contiguous slice instead of a per-expert mask scan.
|
||||
flat_experts = topk_indices.reshape(-1)
|
||||
sorted_experts, order = torch.sort(flat_experts)
|
||||
flat_tokens = x.repeat_interleave(K, dim=0)[order]
|
||||
flat_weights = topk_weights.reshape(-1, 1)[order]
|
||||
boundaries = torch.cumsum(
|
||||
torch.bincount(sorted_experts, minlength=E), dim=0
|
||||
).tolist()
|
||||
|
||||
output = torch.zeros(N, D, device=x.device, dtype=x.dtype)
|
||||
for expert_idx in range(self.n_routed_experts):
|
||||
expert_mask = topk_indices == expert_idx
|
||||
token_idx, k_idx = expert_mask.nonzero(as_tuple=True)
|
||||
if token_idx.numel() == 0:
|
||||
start = 0
|
||||
for expert_idx, end in enumerate(boundaries):
|
||||
if end == start:
|
||||
continue
|
||||
expert_input = x[token_idx]
|
||||
expert_output = self.routed_experts[expert_idx](expert_input)
|
||||
weights = topk_weights[token_idx, k_idx].unsqueeze(-1)
|
||||
output.index_add_(0, token_idx, expert_output * weights)
|
||||
expert_output = self.routed_experts[expert_idx](flat_tokens[start:end])[
|
||||
"hidden_states"
|
||||
]
|
||||
output.index_add_(
|
||||
0,
|
||||
order[start:end] // K,
|
||||
expert_output * flat_weights[start:end],
|
||||
)
|
||||
start = end
|
||||
|
||||
return output
|
||||
return {
|
||||
"hidden_states": output,
|
||||
"aux_loss": aux_loss,
|
||||
"router_stats": router_stats,
|
||||
}
|
||||
|
||||
@@ -11,28 +11,23 @@ def get_rotary_emb(
|
||||
base: float = 10000,
|
||||
device: Optional[torch.device] = None,
|
||||
) -> 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)
|
||||
t = torch.arange(0, max_len, dtype=torch.float64, device=device)
|
||||
freqs = torch.outer(t, theta).float()
|
||||
cos = torch.cos(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:
|
||||
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):
|
||||
def __init__(
|
||||
self,
|
||||
@@ -56,16 +51,26 @@ class RotaryEmbedding(nn.Module):
|
||||
self._set_rotary_buffer(self.max_len)
|
||||
|
||||
def _set_rotary_buffer(self, max_len: int):
|
||||
rotary_emb = get_rotary_emb(self.dim, max_len, self.base)
|
||||
freqs_cis = torch.view_as_real(rotary_emb)
|
||||
freqs_cis = get_rotary_emb(self.dim, max_len, self.base)
|
||||
self.register_buffer("freqs_cis", freqs_cis, persistent=False)
|
||||
|
||||
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:
|
||||
position_ids = (
|
||||
torch.arange(x.size(1), device=x.device)
|
||||
.unsqueeze(0)
|
||||
.expand(x.size(0), -1)
|
||||
)
|
||||
position_freq_cis = self.freqs_cis[position_ids].float()
|
||||
return torch.view_as_complex(position_freq_cis)
|
||||
if x.ndim == 2:
|
||||
position_ids = torch.arange(x.size(0), device=x.device)
|
||||
else:
|
||||
position_ids = (
|
||||
torch.arange(x.size(1), device=x.device)
|
||||
.unsqueeze(0)
|
||||
.expand(x.size(0), -1)
|
||||
)
|
||||
return self.freqs_cis[position_ids].float()
|
||||
|
||||
@@ -5,7 +5,7 @@ import torch.nn as nn
|
||||
from torch import Tensor
|
||||
|
||||
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.embedding import Embedding
|
||||
from astrai.model.components.norm import RMSNorm
|
||||
@@ -13,7 +13,7 @@ from astrai.model.components.rope import RotaryEmbedding
|
||||
from astrai.model.transformer import process_attention_mask
|
||||
|
||||
|
||||
@AutoModel.register("embedding")
|
||||
@ModelFactory.register("embedding")
|
||||
class EmbeddingEncoder(AutoModel):
|
||||
def __init__(self, config: EncoderConfig):
|
||||
super().__init__(config)
|
||||
@@ -70,7 +70,7 @@ class EmbeddingEncoder(AutoModel):
|
||||
attn_mask = process_attention_mask(input_mask)
|
||||
|
||||
for layer in self.layers:
|
||||
x = layer(x, rotary_emb, attn_mask)
|
||||
x = layer(x, rotary_emb, attn_mask)["hidden_states"]
|
||||
|
||||
hidden_states = self.norm(x)
|
||||
|
||||
|
||||
@@ -5,8 +5,8 @@ import torch.nn as nn
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.inference.core.cache import CacheView
|
||||
from astrai.model.automodel import AutoModel
|
||||
from astrai.inference.cache import KVCache
|
||||
from astrai.model.automodel import AutoModel, ModelFactory
|
||||
from astrai.model.components.decoder_block import DecoderBlock
|
||||
from astrai.model.components.embedding import Embedding
|
||||
from astrai.model.components.linear import Linear
|
||||
@@ -26,7 +26,7 @@ def process_attention_mask(
|
||||
return input_mask
|
||||
|
||||
|
||||
@AutoModel.register("autoregressive_lm")
|
||||
@ModelFactory.register("autoregressive_lm")
|
||||
class AutoRegressiveLM(AutoModel):
|
||||
"""Autoregressive language model with paged KV cache."""
|
||||
|
||||
@@ -103,20 +103,50 @@ class AutoRegressiveLM(AutoModel):
|
||||
self,
|
||||
input_ids: Tensor,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
kv_cache: Optional[KVCache] = None,
|
||||
position_ids: Optional[Tensor] = None,
|
||||
fwd: Optional[str] = None,
|
||||
) -> Dict[str, Tensor]:
|
||||
assert input_ids.ndim == 2
|
||||
if fwd is None:
|
||||
if input_ids.ndim != 2:
|
||||
raise ValueError("training input_ids must be [batch, seq_len]")
|
||||
if kv_cache is not None:
|
||||
raise ValueError("training forward does not accept a KV cache")
|
||||
elif fwd in ("prefill", "decode"):
|
||||
if input_ids.ndim != 1:
|
||||
raise ValueError("inference input_ids must be packed [tokens]")
|
||||
if kv_cache is None:
|
||||
raise ValueError("inference forward requires a KV cache")
|
||||
else:
|
||||
raise ValueError(f"unsupported forward mode: {fwd}")
|
||||
|
||||
x = self.embed_tokens(input_ids)
|
||||
rotary_emb = self.rotary_embedding(x, position_ids)
|
||||
attn_mask = process_attention_mask(input_mask)
|
||||
use_sdpa_causal_mask = attn_mask is None
|
||||
|
||||
aux_losses = []
|
||||
router_stats_list = []
|
||||
for layer in self.layers:
|
||||
x = layer(x, rotary_emb, attn_mask, paged_cache, use_sdpa_causal_mask)
|
||||
layer_output = layer(
|
||||
x,
|
||||
rotary_emb,
|
||||
attn_mask,
|
||||
kv_cache,
|
||||
use_sdpa_causal_mask,
|
||||
fwd,
|
||||
)
|
||||
x = layer_output["hidden_states"]
|
||||
stats = layer_output.get("router_stats")
|
||||
if stats is not None:
|
||||
aux_losses.append(layer_output["aux_loss"])
|
||||
router_stats_list.append(stats)
|
||||
|
||||
hidden_states = self.norm(x)
|
||||
logits = self.lm_head(hidden_states)
|
||||
|
||||
return {"logits": logits, "hidden_states": hidden_states}
|
||||
output = {"logits": logits, "hidden_states": hidden_states}
|
||||
if aux_losses:
|
||||
output["aux_loss"] = torch.stack(aux_losses).mean()
|
||||
output["router_stats"] = router_stats_list
|
||||
return output
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
"""Optimizer implementations and factory registration."""
|
||||
|
||||
from astrai.optim.composite import (
|
||||
OptimizerFactory,
|
||||
composite_state_dict,
|
||||
composite_step,
|
||||
composite_zero_grad,
|
||||
refresh_param_groups,
|
||||
)
|
||||
from astrai.optim.mano_adamw import Mano, ManoAdamW
|
||||
from astrai.optim.muon_adamw import MuonAdamW
|
||||
from astrai.optim.nora_nadamw import (
|
||||
NAdamW,
|
||||
Nora,
|
||||
NoraNAdamW,
|
||||
OptimizerParameterGroups,
|
||||
nora_direction,
|
||||
nora_lr_scale,
|
||||
partition_optimizer_parameters,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"Mano",
|
||||
"ManoAdamW",
|
||||
"MuonAdamW",
|
||||
"NAdamW",
|
||||
"Nora",
|
||||
"NoraNAdamW",
|
||||
"OptimizerFactory",
|
||||
"OptimizerParameterGroups",
|
||||
"composite_state_dict",
|
||||
"composite_step",
|
||||
"composite_zero_grad",
|
||||
"nora_direction",
|
||||
"nora_lr_scale",
|
||||
"partition_optimizer_parameters",
|
||||
"refresh_param_groups",
|
||||
]
|
||||
@@ -0,0 +1,71 @@
|
||||
"""Shared infrastructure for the optim package.
|
||||
|
||||
This module hosts two things:
|
||||
|
||||
* ``OptimizerFactory`` — the registry for built-in optimizers. Defining it
|
||||
here (rather than in ``__init__.py``) lets each optimizer module import it
|
||||
and register itself with a decorator, avoiding circular imports.
|
||||
* Composite-optimizer helpers — ``step``/``zero_grad``/``state_dict``/
|
||||
``param_groups`` delegation shared by every optimizer that routes different
|
||||
parameter groups through distinct sub-optimizers.
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch.optim import Optimizer
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
|
||||
|
||||
class OptimizerFactory(BaseFactory[Optimizer]):
|
||||
"""Factory for built-in training optimizers."""
|
||||
|
||||
|
||||
def composite_step(
|
||||
sub_optimizers: list[Optimizer],
|
||||
closure=None,
|
||||
) -> torch.Tensor | None:
|
||||
"""Run ``step`` on every sub-optimizer, invoking the closure once.
|
||||
|
||||
The closure (if given) is executed inside ``torch.enable_grad`` exactly
|
||||
once before any sub-optimizer steps, matching the contract of a single
|
||||
``Optimizer.step``. Sub-optimizers receive ``None`` so they do not
|
||||
re-execute it.
|
||||
"""
|
||||
loss = None
|
||||
if closure is not None:
|
||||
with torch.enable_grad():
|
||||
loss = closure()
|
||||
for sub in sub_optimizers:
|
||||
sub.step()
|
||||
return loss
|
||||
|
||||
|
||||
def composite_zero_grad(
|
||||
sub_optimizers: list[Optimizer],
|
||||
set_to_none: bool = True,
|
||||
) -> None:
|
||||
for sub in sub_optimizers:
|
||||
sub.zero_grad(set_to_none=set_to_none)
|
||||
|
||||
|
||||
def composite_state_dict(
|
||||
named_sub_optimizers: dict[str, Optimizer | None],
|
||||
) -> dict[str, Any]:
|
||||
"""Serialize sub-optimizers, preserving ``None`` slots."""
|
||||
return {
|
||||
name: sub.state_dict() if sub is not None else None
|
||||
for name, sub in named_sub_optimizers.items()
|
||||
}
|
||||
|
||||
|
||||
def refresh_param_groups(
|
||||
sub_optimizers: list[Optimizer],
|
||||
) -> list[dict]:
|
||||
"""Concatenate param_groups from every non-None sub-optimizer."""
|
||||
groups: list[dict] = []
|
||||
for sub in sub_optimizers:
|
||||
if sub is not None:
|
||||
groups.extend(sub.param_groups)
|
||||
return groups
|
||||
@@ -0,0 +1,214 @@
|
||||
"""Mano manifold optimizer combined with AdamW.
|
||||
|
||||
Mano projects the momentum onto the tangent space of the Oblique manifold
|
||||
(axis-wise tangent projection) and normalizes it, replacing the expensive
|
||||
Newton-Schulz iteration in Muon with a cheaper manifold normalization.
|
||||
|
||||
Reference: https://arxiv.org/abs/2601.23000
|
||||
"""
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
from torch import nn, optim
|
||||
from torch.optim import Optimizer
|
||||
|
||||
from astrai.optim.composite import (
|
||||
OptimizerFactory,
|
||||
composite_state_dict,
|
||||
composite_step,
|
||||
composite_zero_grad,
|
||||
refresh_param_groups,
|
||||
)
|
||||
from astrai.optim.nora_nadamw import partition_optimizer_parameters
|
||||
|
||||
|
||||
class Mano(Optimizer):
|
||||
"""Manifold Normalized Optimizer for two-dimensional matrices.
|
||||
|
||||
Each step alternates the projection axis (dim 0 / dim 1) to restrike the
|
||||
manifold along both rows and columns. The tangent momentum is computed
|
||||
without normalizing the parameter itself (v2 simplification) and the
|
||||
epsilon is added (not clamped) to the norm denominator.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
params,
|
||||
lr: float = 1e-3,
|
||||
weight_decay: float = 0.1,
|
||||
momentum: float = 0.95,
|
||||
nesterov: bool = True,
|
||||
eps: float = 1e-8,
|
||||
):
|
||||
if lr < 0:
|
||||
raise ValueError(f"Invalid learning rate: {lr}")
|
||||
if weight_decay < 0:
|
||||
raise ValueError(f"Invalid weight decay: {weight_decay}")
|
||||
if not 0 <= momentum <= 1:
|
||||
raise ValueError(f"Invalid momentum: {momentum}")
|
||||
if eps <= 0:
|
||||
raise ValueError(f"Invalid epsilon: {eps}")
|
||||
|
||||
defaults = {
|
||||
"lr": lr,
|
||||
"weight_decay": weight_decay,
|
||||
"momentum": momentum,
|
||||
"nesterov": nesterov,
|
||||
"eps": eps,
|
||||
"steps": 0,
|
||||
}
|
||||
super().__init__(params, defaults)
|
||||
for group in self.param_groups:
|
||||
for param in group["params"]:
|
||||
if param.ndim != 2:
|
||||
raise ValueError(
|
||||
f"Mano only supports 2D matrices, got shape {tuple(param.shape)}"
|
||||
)
|
||||
|
||||
@torch.no_grad()
|
||||
def step(self, closure=None):
|
||||
loss = None
|
||||
if closure is not None:
|
||||
with torch.enable_grad():
|
||||
loss = closure()
|
||||
|
||||
for group in self.param_groups:
|
||||
lr = group["lr"]
|
||||
weight_decay = group["weight_decay"]
|
||||
momentum = group["momentum"]
|
||||
nesterov = group["nesterov"]
|
||||
eps = group["eps"]
|
||||
dim = int(group["steps"] % 2)
|
||||
|
||||
for param in group["params"]:
|
||||
if param.grad is None:
|
||||
continue
|
||||
if param.grad.is_sparse:
|
||||
raise RuntimeError("Mano does not support sparse gradients")
|
||||
|
||||
grad = param.grad
|
||||
state = self.state[param]
|
||||
momentum_buffer = state.get("momentum_buffer")
|
||||
if momentum_buffer is None:
|
||||
momentum_buffer = torch.zeros_like(grad)
|
||||
momentum_buffer.mul_(momentum).add_(grad)
|
||||
update = (
|
||||
grad.add(momentum_buffer, alpha=momentum)
|
||||
if nesterov
|
||||
else momentum_buffer
|
||||
)
|
||||
|
||||
tangent = update - (
|
||||
torch.sum(update * param.data, dim=dim, keepdim=True) * param.data
|
||||
)
|
||||
direction = tangent / (
|
||||
torch.norm(tangent, p=2, dim=dim, keepdim=True) + eps
|
||||
)
|
||||
|
||||
if weight_decay != 0:
|
||||
param.mul_(1 - lr * weight_decay)
|
||||
adjusted_lr = lr * 0.2 * math.sqrt(direction.shape[dim])
|
||||
param.add_(direction, alpha=-adjusted_lr)
|
||||
state["momentum_buffer"] = momentum_buffer
|
||||
|
||||
group["steps"] += 1
|
||||
|
||||
return loss
|
||||
|
||||
|
||||
@OptimizerFactory.register("mano_adamw")
|
||||
class ManoAdamW(Optimizer):
|
||||
"""Mano for internal linear weights and AdamW for remaining parameters."""
|
||||
|
||||
optimizer_name = "mano_adamw"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: nn.Module,
|
||||
lr: float = 3e-4,
|
||||
weight_decay: float = 0.1,
|
||||
momentum: float = 0.95,
|
||||
nesterov: bool = True,
|
||||
):
|
||||
groups = partition_optimizer_parameters(model)
|
||||
all_params = [
|
||||
*groups.nora,
|
||||
*groups.nadamw_decay,
|
||||
*groups.nadamw_no_decay,
|
||||
]
|
||||
if not all_params:
|
||||
raise ValueError(
|
||||
"Cannot build an optimizer for a model with no trainable parameters"
|
||||
)
|
||||
super().__init__(all_params, {})
|
||||
|
||||
self.mano = (
|
||||
Mano(
|
||||
groups.nora,
|
||||
lr=lr,
|
||||
weight_decay=weight_decay,
|
||||
momentum=momentum,
|
||||
nesterov=nesterov,
|
||||
)
|
||||
if groups.nora
|
||||
else None
|
||||
)
|
||||
|
||||
adamw_groups = []
|
||||
if groups.nadamw_decay:
|
||||
adamw_groups.append(
|
||||
{"params": groups.nadamw_decay, "weight_decay": weight_decay}
|
||||
)
|
||||
if groups.nadamw_no_decay:
|
||||
adamw_groups.append({"params": groups.nadamw_no_decay, "weight_decay": 0.0})
|
||||
self.adamw = (
|
||||
optim.AdamW(
|
||||
adamw_groups,
|
||||
lr=lr,
|
||||
betas=(0.9, 0.95),
|
||||
fused=True,
|
||||
)
|
||||
if adamw_groups
|
||||
else None
|
||||
)
|
||||
self.param_groups = refresh_param_groups([self.mano, self.adamw])
|
||||
|
||||
@torch.no_grad()
|
||||
def step(self, closure=None):
|
||||
return composite_step(
|
||||
[opt for opt in (self.mano, self.adamw) if opt is not None],
|
||||
closure,
|
||||
)
|
||||
|
||||
def zero_grad(self, set_to_none: bool = True):
|
||||
composite_zero_grad(
|
||||
[opt for opt in (self.mano, self.adamw) if opt is not None],
|
||||
set_to_none,
|
||||
)
|
||||
|
||||
def state_dict(self) -> dict:
|
||||
return composite_state_dict({"mano": self.mano, "adamw": self.adamw})
|
||||
|
||||
def load_state_dict(self, state_dict: dict):
|
||||
if "muon" in state_dict or "nora" in state_dict:
|
||||
raise ValueError(
|
||||
"Checkpoint uses a different optimizer; select the matching "
|
||||
"--optimizer to resume it"
|
||||
)
|
||||
if "mano" not in state_dict or "adamw" not in state_dict:
|
||||
raise ValueError(
|
||||
"Checkpoint optimizer state is not compatible with mano_adamw"
|
||||
)
|
||||
|
||||
saved_mano = state_dict["mano"]
|
||||
saved_adamw = state_dict["adamw"]
|
||||
if (self.mano is None) != (saved_mano is None):
|
||||
raise ValueError("Checkpoint Mano parameter groups do not match the model")
|
||||
if (self.adamw is None) != (saved_adamw is None):
|
||||
raise ValueError("Checkpoint AdamW parameter groups do not match the model")
|
||||
if self.mano is not None:
|
||||
self.mano.load_state_dict(saved_mano)
|
||||
if self.adamw is not None:
|
||||
self.adamw.load_state_dict(saved_adamw)
|
||||
self.param_groups = refresh_param_groups([self.mano, self.adamw])
|
||||
@@ -0,0 +1,95 @@
|
||||
"""Legacy Muon + AdamW combined optimizer."""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch import Tensor, nn, optim
|
||||
|
||||
from astrai.optim.composite import (
|
||||
OptimizerFactory,
|
||||
composite_state_dict,
|
||||
composite_step,
|
||||
composite_zero_grad,
|
||||
refresh_param_groups,
|
||||
)
|
||||
|
||||
|
||||
@OptimizerFactory.register("muon_adamw")
|
||||
class MuonAdamW(optim.Optimizer):
|
||||
"""Combined Muon (matrix) + AdamW (non-matrix) optimizer."""
|
||||
|
||||
optimizer_name = "muon_adamw"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: nn.Module,
|
||||
lr: float = 3e-4,
|
||||
weight_decay: float = 0.1,
|
||||
momentum: float = 0.95,
|
||||
nesterov: bool = True,
|
||||
ns_steps: int = 5,
|
||||
adjust_lr_fn: str = "match_rms_adamw",
|
||||
):
|
||||
defaults = {
|
||||
"lr": lr,
|
||||
"weight_decay": weight_decay,
|
||||
"momentum": momentum,
|
||||
"nesterov": nesterov,
|
||||
"ns_steps": ns_steps,
|
||||
"adjust_lr_fn": adjust_lr_fn,
|
||||
}
|
||||
params = [param for param in model.parameters() if param.requires_grad]
|
||||
super().__init__(params, defaults)
|
||||
|
||||
matrix_params: list[Tensor] = []
|
||||
other_params: list[Tensor] = []
|
||||
for name, param in model.named_parameters():
|
||||
if not param.requires_grad:
|
||||
continue
|
||||
if (
|
||||
param.dim() >= 2
|
||||
and "norm" not in name
|
||||
and "bias" not in name
|
||||
and "embed" not in name
|
||||
and "lm_head" not in name
|
||||
):
|
||||
matrix_params.append(param)
|
||||
else:
|
||||
other_params.append(param)
|
||||
|
||||
self.muon = optim.Muon(
|
||||
matrix_params,
|
||||
lr=lr,
|
||||
weight_decay=weight_decay,
|
||||
momentum=momentum,
|
||||
nesterov=nesterov,
|
||||
ns_steps=ns_steps,
|
||||
adjust_lr_fn=adjust_lr_fn,
|
||||
)
|
||||
self.adamw = optim.AdamW(
|
||||
[{"params": other_params, "weight_decay": 0.0}],
|
||||
lr=lr,
|
||||
betas=(0.9, 0.95),
|
||||
fused=True,
|
||||
)
|
||||
|
||||
self.param_groups = refresh_param_groups([self.muon, self.adamw])
|
||||
|
||||
@torch.no_grad()
|
||||
def step(self, closure=None):
|
||||
return composite_step([self.muon, self.adamw], closure)
|
||||
|
||||
def zero_grad(self, set_to_none: bool = True):
|
||||
composite_zero_grad([self.muon, self.adamw], set_to_none)
|
||||
|
||||
def state_dict(self) -> dict[str, Any]:
|
||||
return composite_state_dict({"muon": self.muon, "adamw": self.adamw})
|
||||
|
||||
def load_state_dict(self, state_dict: dict[str, Any]):
|
||||
if "muon" not in state_dict or "adamw" not in state_dict:
|
||||
raise ValueError(
|
||||
"Checkpoint optimizer state is not compatible with muon_adamw"
|
||||
)
|
||||
self.muon.load_state_dict(state_dict["muon"])
|
||||
self.adamw.load_state_dict(state_dict["adamw"])
|
||||
self.param_groups = refresh_param_groups([self.muon, self.adamw])
|
||||
@@ -0,0 +1,372 @@
|
||||
"""Nora matrix optimizer combined with Nesterov AdamW."""
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch import Tensor, nn
|
||||
from torch.distributed.tensor import DTensor, Shard
|
||||
from torch.optim import Optimizer
|
||||
|
||||
from astrai.model.components.embedding import Embedding
|
||||
from astrai.model.components.linear import Linear
|
||||
from astrai.model.components.lora import LoRALinear
|
||||
from astrai.model.components.norm import RMSNorm
|
||||
from astrai.optim.composite import (
|
||||
OptimizerFactory,
|
||||
composite_state_dict,
|
||||
composite_step,
|
||||
composite_zero_grad,
|
||||
refresh_param_groups,
|
||||
)
|
||||
|
||||
NORA_EPS = 1e-10
|
||||
|
||||
|
||||
def _row_normalize(tensor: Tensor, eps: float) -> Tensor:
|
||||
return tensor / tensor.norm(dim=-1, keepdim=True).clamp(min=eps)
|
||||
|
||||
|
||||
def nora_direction(update: Tensor, param: Tensor, eps: float = NORA_EPS) -> Tensor:
|
||||
"""Project an update onto each parameter row's tangent space and normalize."""
|
||||
theta_hat = _row_normalize(param.to(torch.float32), eps)
|
||||
update_fp32 = update.to(torch.float32)
|
||||
radial = (update_fp32 * theta_hat).sum(dim=-1, keepdim=True) * theta_hat
|
||||
direction = _row_normalize(update_fp32 - radial, eps)
|
||||
return direction.to(update.dtype)
|
||||
|
||||
|
||||
def nora_lr_scale(lr: float, shape: torch.Size) -> float:
|
||||
"""Scale Nora's LR for tall ``[d_out, d_in]`` linear weights."""
|
||||
return lr * math.sqrt(max(1.0, shape[-2] / shape[-1]))
|
||||
|
||||
|
||||
def _validate_complete_rows(param: Tensor) -> None:
|
||||
if not isinstance(param, DTensor):
|
||||
return
|
||||
last_dim = param.ndim - 1
|
||||
for placement in param.placements:
|
||||
if isinstance(placement, Shard) and placement.dim % param.ndim == last_dim:
|
||||
raise ValueError(
|
||||
"Nora requires complete parameter rows, but this DTensor is sharded "
|
||||
"along its last dimension"
|
||||
)
|
||||
|
||||
|
||||
class Nora(Optimizer):
|
||||
"""Normalized Orthogonal Row Alignment for two-dimensional matrices."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
params,
|
||||
lr: float = 5e-3,
|
||||
weight_decay: float = 0.0,
|
||||
momentum: float = 0.95,
|
||||
beta: float = 0.95,
|
||||
nesterov: bool = True,
|
||||
eps: float = NORA_EPS,
|
||||
):
|
||||
if lr < 0:
|
||||
raise ValueError(f"Invalid learning rate: {lr}")
|
||||
if weight_decay < 0:
|
||||
raise ValueError(f"Invalid weight decay: {weight_decay}")
|
||||
if not 0 <= momentum <= 1:
|
||||
raise ValueError(f"Invalid momentum: {momentum}")
|
||||
if not 0 <= beta < 1:
|
||||
raise ValueError(f"Invalid beta: {beta}")
|
||||
if eps <= 0:
|
||||
raise ValueError(f"Invalid epsilon: {eps}")
|
||||
|
||||
defaults = {
|
||||
"lr": lr,
|
||||
"weight_decay": weight_decay,
|
||||
"momentum": momentum,
|
||||
"beta": beta,
|
||||
"nesterov": nesterov,
|
||||
"eps": eps,
|
||||
}
|
||||
super().__init__(params, defaults)
|
||||
for group in self.param_groups:
|
||||
for param in group["params"]:
|
||||
if param.ndim != 2:
|
||||
raise ValueError(
|
||||
f"Nora only supports 2D matrices, got shape {tuple(param.shape)}"
|
||||
)
|
||||
_validate_complete_rows(param)
|
||||
|
||||
@torch.no_grad()
|
||||
def step(self, closure=None):
|
||||
loss = None
|
||||
if closure is not None:
|
||||
with torch.enable_grad():
|
||||
loss = closure()
|
||||
|
||||
for group in self.param_groups:
|
||||
lr = group["lr"]
|
||||
weight_decay = group["weight_decay"]
|
||||
momentum = group["momentum"]
|
||||
beta = group["beta"]
|
||||
nesterov = group["nesterov"]
|
||||
eps = group["eps"]
|
||||
for param in group["params"]:
|
||||
if param.grad is None:
|
||||
continue
|
||||
if param.grad.is_sparse:
|
||||
raise RuntimeError("Nora does not support sparse gradients")
|
||||
|
||||
grad = param.grad
|
||||
state = self.state[param]
|
||||
momentum_buffer = state.get("momentum_buffer")
|
||||
if momentum_buffer is None:
|
||||
momentum_buffer = torch.zeros_like(grad)
|
||||
momentum_buffer.lerp_(grad, 1 - beta)
|
||||
update = (
|
||||
grad.lerp(momentum_buffer, momentum)
|
||||
if nesterov
|
||||
else momentum_buffer
|
||||
)
|
||||
direction = nora_direction(update, param, eps)
|
||||
|
||||
if weight_decay != 0:
|
||||
param.mul_(1 - lr * weight_decay)
|
||||
param.add_(direction, alpha=-nora_lr_scale(lr, param.shape))
|
||||
state["momentum_buffer"] = momentum_buffer
|
||||
|
||||
return loss
|
||||
|
||||
|
||||
class NAdamW(Optimizer):
|
||||
"""AdamW using the reference Nesterov first-moment update."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
params,
|
||||
lr: float = 3e-4,
|
||||
betas: tuple[float, float] = (0.9, 0.999),
|
||||
eps: float = 1e-8,
|
||||
weight_decay: float = 0.1,
|
||||
):
|
||||
beta1, beta2 = betas
|
||||
if lr < 0:
|
||||
raise ValueError(f"Invalid learning rate: {lr}")
|
||||
if not 0 <= beta1 < 1 or not 0 <= beta2 < 1:
|
||||
raise ValueError(f"Invalid betas: {betas}")
|
||||
if eps <= 0:
|
||||
raise ValueError(f"Invalid epsilon: {eps}")
|
||||
if weight_decay < 0:
|
||||
raise ValueError(f"Invalid weight decay: {weight_decay}")
|
||||
defaults = {
|
||||
"lr": lr,
|
||||
"betas": betas,
|
||||
"eps": eps,
|
||||
"weight_decay": weight_decay,
|
||||
}
|
||||
super().__init__(params, defaults)
|
||||
|
||||
@torch.no_grad()
|
||||
def step(self, closure=None):
|
||||
loss = None
|
||||
if closure is not None:
|
||||
with torch.enable_grad():
|
||||
loss = closure()
|
||||
|
||||
for group in self.param_groups:
|
||||
beta1, beta2 = group["betas"]
|
||||
eps = group["eps"]
|
||||
lr = group["lr"]
|
||||
weight_decay = group["weight_decay"]
|
||||
for param in group["params"]:
|
||||
if param.grad is None:
|
||||
continue
|
||||
if param.grad.is_sparse:
|
||||
raise RuntimeError("NAdamW does not support sparse gradients")
|
||||
|
||||
grad = param.grad
|
||||
state = self.state[param]
|
||||
if not state:
|
||||
state["step"] = 0
|
||||
state["m"] = torch.zeros_like(param)
|
||||
state["v"] = torch.zeros_like(param)
|
||||
|
||||
state["step"] += 1
|
||||
first_moment = state["m"]
|
||||
second_moment = state["v"]
|
||||
first_moment.mul_(beta1).add_(grad, alpha=1 - beta1)
|
||||
second_moment.mul_(beta2).addcmul_(grad, grad, value=1 - beta2)
|
||||
|
||||
bias_correction1 = 1 - beta1 ** state["step"]
|
||||
bias_correction2 = 1 - beta2 ** state["step"]
|
||||
nesterov_moment = (
|
||||
beta1 * first_moment + (1 - beta1) * grad
|
||||
) / bias_correction1
|
||||
corrected_second_moment = second_moment / bias_correction2
|
||||
|
||||
if weight_decay != 0:
|
||||
param.mul_(1 - lr * weight_decay)
|
||||
param.addcdiv_(
|
||||
nesterov_moment,
|
||||
corrected_second_moment.sqrt().add_(eps),
|
||||
value=-lr,
|
||||
)
|
||||
|
||||
return loss
|
||||
|
||||
|
||||
@dataclass
|
||||
class OptimizerParameterGroups:
|
||||
nora: list[Tensor]
|
||||
nadamw_decay: list[Tensor]
|
||||
nadamw_no_decay: list[Tensor]
|
||||
|
||||
|
||||
def partition_optimizer_parameters(model: nn.Module) -> OptimizerParameterGroups:
|
||||
"""Partition trainable parameters by module role and parameter identity."""
|
||||
nora_ids: set[int] = set()
|
||||
no_decay_ids: set[int] = set()
|
||||
|
||||
for module_name, module in model.named_modules():
|
||||
if isinstance(module, LoRALinear):
|
||||
for param in module.parameters(recurse=False):
|
||||
if param.requires_grad:
|
||||
no_decay_ids.add(id(param))
|
||||
continue
|
||||
|
||||
if isinstance(module, (Embedding, RMSNorm)):
|
||||
for param in module.parameters(recurse=False):
|
||||
if param.requires_grad:
|
||||
no_decay_ids.add(id(param))
|
||||
continue
|
||||
|
||||
if not isinstance(module, Linear):
|
||||
continue
|
||||
|
||||
if module.bias is not None and module.bias.requires_grad:
|
||||
no_decay_ids.add(id(module.bias))
|
||||
if not module.weight.requires_grad:
|
||||
continue
|
||||
if module_name.rsplit(".", 1)[-1] == "lm_head":
|
||||
no_decay_ids.add(id(module.weight))
|
||||
elif module.weight.ndim == 2:
|
||||
nora_ids.add(id(module.weight))
|
||||
|
||||
nora: list[Tensor] = []
|
||||
nadamw_decay: list[Tensor] = []
|
||||
nadamw_no_decay: list[Tensor] = []
|
||||
seen: set[int] = set()
|
||||
for param in model.parameters():
|
||||
param_id = id(param)
|
||||
if not param.requires_grad or param_id in seen:
|
||||
continue
|
||||
seen.add(param_id)
|
||||
if param_id in no_decay_ids or param.ndim <= 1:
|
||||
nadamw_no_decay.append(param)
|
||||
elif param_id in nora_ids:
|
||||
nora.append(param)
|
||||
else:
|
||||
nadamw_decay.append(param)
|
||||
|
||||
trainable_ids = {id(param) for param in model.parameters() if param.requires_grad}
|
||||
grouped_ids = {id(param) for param in [*nora, *nadamw_decay, *nadamw_no_decay]}
|
||||
if grouped_ids != trainable_ids:
|
||||
missing = len(trainable_ids - grouped_ids)
|
||||
extra = len(grouped_ids - trainable_ids)
|
||||
raise RuntimeError(
|
||||
f"Optimizer parameter partition is incomplete: missing={missing}, extra={extra}"
|
||||
)
|
||||
|
||||
return OptimizerParameterGroups(nora, nadamw_decay, nadamw_no_decay)
|
||||
|
||||
|
||||
@OptimizerFactory.register("nora_nadamw")
|
||||
class NoraNAdamW(Optimizer):
|
||||
"""Nora for internal linear weights and NAdamW for remaining parameters."""
|
||||
|
||||
optimizer_name = "nora_nadamw"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: nn.Module,
|
||||
lr: float = 3e-4,
|
||||
weight_decay: float = 0.1,
|
||||
nora_lr: float = 5e-3,
|
||||
nora_weight_decay: float = 0.0,
|
||||
nora_beta: float = 0.95,
|
||||
nora_momentum: float = 0.95,
|
||||
):
|
||||
groups = partition_optimizer_parameters(model)
|
||||
all_params = [
|
||||
*groups.nora,
|
||||
*groups.nadamw_decay,
|
||||
*groups.nadamw_no_decay,
|
||||
]
|
||||
if not all_params:
|
||||
raise ValueError(
|
||||
"Cannot build an optimizer for a model with no trainable parameters"
|
||||
)
|
||||
super().__init__(all_params, {})
|
||||
|
||||
self.nora = (
|
||||
Nora(
|
||||
groups.nora,
|
||||
lr=nora_lr,
|
||||
weight_decay=nora_weight_decay,
|
||||
momentum=nora_momentum,
|
||||
beta=nora_beta,
|
||||
)
|
||||
if groups.nora
|
||||
else None
|
||||
)
|
||||
|
||||
nadamw_groups = []
|
||||
if groups.nadamw_decay:
|
||||
nadamw_groups.append(
|
||||
{"params": groups.nadamw_decay, "weight_decay": weight_decay}
|
||||
)
|
||||
if groups.nadamw_no_decay:
|
||||
nadamw_groups.append(
|
||||
{"params": groups.nadamw_no_decay, "weight_decay": 0.0}
|
||||
)
|
||||
self.nadamw = NAdamW(nadamw_groups, lr=lr) if nadamw_groups else None
|
||||
self.param_groups = refresh_param_groups([self.nora, self.nadamw])
|
||||
|
||||
@torch.no_grad()
|
||||
def step(self, closure=None):
|
||||
return composite_step(
|
||||
[opt for opt in (self.nora, self.nadamw) if opt is not None],
|
||||
closure,
|
||||
)
|
||||
|
||||
def zero_grad(self, set_to_none: bool = True):
|
||||
composite_zero_grad(
|
||||
[opt for opt in (self.nora, self.nadamw) if opt is not None],
|
||||
set_to_none,
|
||||
)
|
||||
|
||||
def state_dict(self) -> dict[str, Any]:
|
||||
return composite_state_dict({"nora": self.nora, "nadamw": self.nadamw})
|
||||
|
||||
def load_state_dict(self, state_dict: dict[str, Any]):
|
||||
if "muon" in state_dict or "adamw" in state_dict:
|
||||
raise ValueError(
|
||||
"Checkpoint uses muon_adamw state; select optimizer='muon_adamw' "
|
||||
"to resume it"
|
||||
)
|
||||
if "nora" not in state_dict or "nadamw" not in state_dict:
|
||||
raise ValueError(
|
||||
"Checkpoint optimizer state is not compatible with nora_nadamw"
|
||||
)
|
||||
|
||||
saved_nora = state_dict["nora"]
|
||||
saved_nadamw = state_dict["nadamw"]
|
||||
if (self.nora is None) != (saved_nora is None):
|
||||
raise ValueError("Checkpoint Nora parameter groups do not match the model")
|
||||
if (self.nadamw is None) != (saved_nadamw is None):
|
||||
raise ValueError(
|
||||
"Checkpoint NAdamW parameter groups do not match the model"
|
||||
)
|
||||
if self.nora is not None:
|
||||
self.nora.load_state_dict(saved_nora)
|
||||
if self.nadamw is not None:
|
||||
self.nadamw.load_state_dict(saved_nadamw)
|
||||
self.param_groups = refresh_param_groups([self.nora, self.nadamw])
|
||||
@@ -4,12 +4,12 @@ from astrai.parallel.executor import (
|
||||
BaseExecutor,
|
||||
DDPExecutor,
|
||||
ExecutorFactory,
|
||||
FSDP2Executor,
|
||||
FSDPExecutor,
|
||||
GradientState,
|
||||
NoneExecutor,
|
||||
broadcast_state_dict,
|
||||
create_ref_model,
|
||||
)
|
||||
from astrai.parallel.module import ColumnParallelLinear, RowParallelLinear
|
||||
from astrai.parallel.setup import (
|
||||
get_current_device,
|
||||
get_rank,
|
||||
@@ -26,8 +26,6 @@ __all__ = [
|
||||
"only_on_rank",
|
||||
"setup_parallel",
|
||||
"spawn_parallel_fn",
|
||||
"RowParallelLinear",
|
||||
"ColumnParallelLinear",
|
||||
"ExecutorFactory",
|
||||
"BaseExecutor",
|
||||
"GradientState",
|
||||
@@ -36,5 +34,6 @@ __all__ = [
|
||||
"NoneExecutor",
|
||||
"DDPExecutor",
|
||||
"FSDPExecutor",
|
||||
"FSDP2Executor",
|
||||
"create_ref_model",
|
||||
"broadcast_state_dict",
|
||||
]
|
||||
|
||||
+120
-99
@@ -4,18 +4,15 @@ import contextlib
|
||||
import logging
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
from typing import Any, Callable, Optional, Tuple
|
||||
from typing import Any, Callable, Dict, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn as nn
|
||||
from torch.distributed.fsdp import (
|
||||
FSDPModule,
|
||||
FullStateDictConfig,
|
||||
StateDictType,
|
||||
fully_shard,
|
||||
)
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
from torch.distributed.tensor import DTensor
|
||||
from torch.nn.parallel import DistributedDataParallel as DDP
|
||||
from torch.optim import Optimizer
|
||||
@@ -27,6 +24,82 @@ from astrai.parallel.setup import get_rank, get_world_size
|
||||
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:
|
||||
def __init__(self, grad_accum_steps: int = 1):
|
||||
self.num_steps = max(grad_accum_steps, 1)
|
||||
@@ -95,11 +168,14 @@ class BaseExecutor:
|
||||
optimizer_fn: Optional[Callable[[nn.Module], Optimizer]] = None,
|
||||
scheduler_fn: Optional[Callable[[Optimizer], LRScheduler]] = None,
|
||||
before_wrap: Optional[Callable[[nn.Module], nn.Module]] = None,
|
||||
after_wrap: Optional[Callable[[nn.Module], nn.Module]] = None,
|
||||
) -> Tuple[nn.Module, Optional[Optimizer], Optional[LRScheduler]]:
|
||||
model = model_fn()
|
||||
if before_wrap is not None:
|
||||
model = before_wrap(model)
|
||||
model = self._prepare_model(model)
|
||||
if after_wrap is not None:
|
||||
model = after_wrap(model)
|
||||
optimizer = None
|
||||
scheduler = None
|
||||
if optimizer_fn is not None:
|
||||
@@ -238,88 +314,11 @@ class DDPExecutor(BaseExecutor):
|
||||
|
||||
@ExecutorFactory.register("fsdp")
|
||||
class FSDPExecutor(BaseExecutor):
|
||||
def __init__(
|
||||
self,
|
||||
grad_accum_steps: int = 1,
|
||||
process_group=None,
|
||||
sharding_strategy=None,
|
||||
cpu_offload=None,
|
||||
auto_wrap_policy=None,
|
||||
backward_prefetch=None,
|
||||
mixed_precision=None,
|
||||
ignored_modules=None,
|
||||
param_init_fn=None,
|
||||
sync_module_states: bool = False,
|
||||
forward_prefetch: bool = False,
|
||||
limit_all_gathers: bool = True,
|
||||
ignored_states=None,
|
||||
device_mesh=None,
|
||||
):
|
||||
super().__init__(grad_accum_steps=grad_accum_steps)
|
||||
self._fsdp_kwargs = {
|
||||
k: v
|
||||
for k, v in dict(
|
||||
process_group=process_group,
|
||||
sharding_strategy=sharding_strategy,
|
||||
cpu_offload=cpu_offload,
|
||||
auto_wrap_policy=auto_wrap_policy,
|
||||
backward_prefetch=backward_prefetch,
|
||||
mixed_precision=mixed_precision,
|
||||
ignored_modules=ignored_modules,
|
||||
param_init_fn=param_init_fn,
|
||||
sync_module_states=sync_module_states,
|
||||
forward_prefetch=forward_prefetch,
|
||||
limit_all_gathers=limit_all_gathers,
|
||||
use_orig_params=True,
|
||||
ignored_states=ignored_states,
|
||||
device_mesh=device_mesh,
|
||||
).items()
|
||||
if v is not None
|
||||
}
|
||||
self._original_model: Optional[nn.Module] = None
|
||||
|
||||
def _prepare_model(self, model: nn.Module) -> nn.Module:
|
||||
if not self.use_distributed:
|
||||
logger.warning("FSDP backend selected but world_size=1, model not wrapped")
|
||||
return model
|
||||
self._original_model = model
|
||||
device_id = torch.device("cuda", get_rank())
|
||||
model = FSDP(model, device_id=device_id, **self._fsdp_kwargs)
|
||||
logger.info("Model wrapped with FSDP (world_size=%d)", get_world_size())
|
||||
return model
|
||||
|
||||
def _no_sync(self, model: nn.Module):
|
||||
if isinstance(model, FSDP):
|
||||
return model.no_sync()
|
||||
return contextlib.nullcontext()
|
||||
|
||||
def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float:
|
||||
if isinstance(model, FSDP) and self.use_distributed:
|
||||
total_norm = model.clip_grad_norm_(max_norm)
|
||||
if isinstance(total_norm, torch.Tensor):
|
||||
return total_norm.item()
|
||||
return total_norm
|
||||
return super().clip_grad_norm(model, max_norm)
|
||||
|
||||
def unwrap_model(self, model: nn.Module):
|
||||
if isinstance(model, FSDP) and self.use_distributed:
|
||||
with FSDP.state_dict_type(
|
||||
model,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
|
||||
):
|
||||
return model.state_dict()
|
||||
|
||||
return model.state_dict()
|
||||
|
||||
|
||||
@ExecutorFactory.register("fsdp2")
|
||||
class FSDP2Executor(BaseExecutor):
|
||||
"""FSDP2 executor using `torch.distributed.fsdp.fully_shard` (per-module API).
|
||||
"""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
|
||||
FSDP2's dynamic ``__class__`` assignment fail at the CPython level.
|
||||
``fully_shard``'s dynamic ``__class__`` assignment fail at the CPython level.
|
||||
Original ``Parameter`` objects are preserved (as DTensors) — no
|
||||
``FlatParameter``, no ``use_orig_params=True`` hack.
|
||||
"""
|
||||
@@ -329,7 +328,7 @@ class FSDP2Executor(BaseExecutor):
|
||||
grad_accum_steps: int = 1,
|
||||
mesh: Optional[Any] = None,
|
||||
mp_policy: Optional[Any] = None,
|
||||
reshard_after_forward: bool = True,
|
||||
reshard_after_forward: bool = False,
|
||||
):
|
||||
super().__init__(grad_accum_steps=grad_accum_steps)
|
||||
self._mesh = mesh
|
||||
@@ -338,7 +337,7 @@ class FSDP2Executor(BaseExecutor):
|
||||
|
||||
def _prepare_model(self, model: nn.Module) -> nn.Module:
|
||||
if not self.use_distributed:
|
||||
logger.warning("FSDP2 backend selected but world_size=1, model not wrapped")
|
||||
logger.warning("FSDP backend selected but world_size=1, model not wrapped")
|
||||
return model
|
||||
|
||||
kwargs = dict(
|
||||
@@ -356,7 +355,7 @@ class FSDP2Executor(BaseExecutor):
|
||||
fully_shard(child, **kwargs)
|
||||
|
||||
logger.info(
|
||||
"FSDP2 wrapping applied to %d direct children (root skipped for ABC compat)",
|
||||
"FSDP wrapping applied to %d direct children (root skipped for ABC compat)",
|
||||
len(list(model.children())),
|
||||
)
|
||||
return model
|
||||
@@ -376,32 +375,54 @@ class FSDP2Executor(BaseExecutor):
|
||||
yield
|
||||
|
||||
def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float:
|
||||
if self.use_distributed:
|
||||
total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm)
|
||||
if isinstance(total_norm, torch.Tensor):
|
||||
return total_norm.item()
|
||||
return total_norm
|
||||
return super().clip_grad_norm(model, max_norm)
|
||||
if not self.use_distributed:
|
||||
return super().clip_grad_norm(model, max_norm)
|
||||
|
||||
# FSDP params are DTensors (sharded across ranks).
|
||||
# torch.nn.utils.clip_grad_norm_ computes LOCAL norm per rank,
|
||||
# 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):
|
||||
if not self.use_distributed:
|
||||
return model.state_dict()
|
||||
|
||||
if get_rank() != 0:
|
||||
return None
|
||||
|
||||
# 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 = {
|
||||
k: (v.full_tensor() if isinstance(v, DTensor) else v)
|
||||
for k, v in state_dict.items()
|
||||
}
|
||||
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)
|
||||
@@ -12,7 +12,7 @@ import torch
|
||||
import torch.distributed as dist
|
||||
import torch.multiprocessing as mp
|
||||
|
||||
from astrai.parallel.signal_handler import install_early_signal_handlers
|
||||
from astrai.signal_handler import install_early_signal_handlers
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -416,7 +416,11 @@ class MultiOutputMaskBuilder(BaseMaskBuilder):
|
||||
return None
|
||||
|
||||
result: dict = {}
|
||||
any_output = False
|
||||
required_outputs = {
|
||||
output_key
|
||||
for output_key, spec in sources_spec.items()
|
||||
if spec.get("sections")
|
||||
}
|
||||
|
||||
for output_key, spec in sources_spec.items():
|
||||
sections = spec.get("sections", [])
|
||||
@@ -428,7 +432,6 @@ class MultiOutputMaskBuilder(BaseMaskBuilder):
|
||||
if ids is None:
|
||||
continue
|
||||
result[output_key] = ids
|
||||
any_output = True
|
||||
continue
|
||||
|
||||
list_field = spec.get("list_field", False)
|
||||
@@ -444,7 +447,6 @@ class MultiOutputMaskBuilder(BaseMaskBuilder):
|
||||
result[output_key] = ids
|
||||
if mask is not None:
|
||||
result[mask_key] = mask
|
||||
any_output = True
|
||||
continue
|
||||
|
||||
ids, mask = self.renderer.process_sections(
|
||||
@@ -460,9 +462,7 @@ class MultiOutputMaskBuilder(BaseMaskBuilder):
|
||||
elif "mask_key" in spec:
|
||||
result[mask_key] = mask
|
||||
|
||||
any_output = True
|
||||
|
||||
if not any_output:
|
||||
if not required_outputs or not required_outputs.issubset(result):
|
||||
return None
|
||||
|
||||
result["domain"] = _extract_domain(item, config.output.domain_key)
|
||||
@@ -474,6 +474,11 @@ class MultiOutputMaskBuilder(BaseMaskBuilder):
|
||||
return [None] * len(items)
|
||||
|
||||
results = [{} for _ in items]
|
||||
required_outputs = {
|
||||
output_key
|
||||
for output_key, spec in sources_spec.items()
|
||||
if spec.get("sections")
|
||||
}
|
||||
for output_key, spec in sources_spec.items():
|
||||
sections = spec.get("sections", [])
|
||||
if not sections:
|
||||
@@ -506,7 +511,7 @@ class MultiOutputMaskBuilder(BaseMaskBuilder):
|
||||
|
||||
return [
|
||||
({**result, "domain": _extract_domain(item, config.output.domain_key)})
|
||||
if result
|
||||
if required_outputs and required_outputs.issubset(result)
|
||||
else None
|
||||
for item, result in zip(items, results)
|
||||
]
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""Config-driven JSONL preprocessing pipeline.
|
||||
|
||||
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,
|
||||
dispatched by configuration keys.
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""Storage writer strategies for pipeline output.
|
||||
|
||||
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
|
||||
by ``output.storage_format``.
|
||||
"""
|
||||
@@ -15,7 +15,7 @@ from typing import Dict, List
|
||||
import torch
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.serialization import save_bin, save_h5
|
||||
from astrai.serialization import save_bin
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -54,22 +54,3 @@ class BinWriter(StoreWriter):
|
||||
exc_info=True,
|
||||
)
|
||||
raise
|
||||
|
||||
|
||||
@StoreWriterFactory.register("h5")
|
||||
class H5Writer(StoreWriter):
|
||||
def save(self, output_dir, domain, shard_idx, tensors):
|
||||
chunk_dir = os.path.join(output_dir, domain)
|
||||
file_path = os.path.join(chunk_dir, f"data_{shard_idx:04d}.h5")
|
||||
try:
|
||||
save_h5(chunk_dir, f"data_{shard_idx:04d}", tensors)
|
||||
except Exception:
|
||||
if os.path.exists(file_path):
|
||||
os.remove(file_path)
|
||||
logger.error(
|
||||
"Failed to write shard %s/data_%04d.h5, cleaned up partial output",
|
||||
domain,
|
||||
shard_idx,
|
||||
exc_info=True,
|
||||
)
|
||||
raise
|
||||
|
||||
@@ -20,9 +20,7 @@ from astrai.serialization.checkpoint import (
|
||||
from astrai.serialization.dataset import (
|
||||
load_bin,
|
||||
load_bin_offsets,
|
||||
load_h5,
|
||||
save_bin,
|
||||
save_h5,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
@@ -39,7 +37,5 @@ __all__ = [
|
||||
"save_torch",
|
||||
"load_bin",
|
||||
"load_bin_offsets",
|
||||
"load_h5",
|
||||
"save_bin",
|
||||
"save_h5",
|
||||
]
|
||||
|
||||
@@ -1,55 +1,14 @@
|
||||
"""Dataset storage serialization helpers (HDF5 / memory-mapped binary)."""
|
||||
"""Dataset storage serialization helpers (memory-mapped binary)."""
|
||||
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import h5py
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
def save_h5(file_path: str, file_name: str, tensor_group: Dict[str, List[Tensor]]):
|
||||
os.makedirs(file_path, exist_ok=True)
|
||||
full_file_path = os.path.join(file_path, f"{file_name}.h5")
|
||||
with h5py.File(full_file_path, "w") as f:
|
||||
for key, tensors in tensor_group.items():
|
||||
grp = f.create_group(key)
|
||||
for idx, tensor in enumerate(tensors):
|
||||
arr = tensor.cpu().numpy()
|
||||
grp.create_dataset(f"data_{idx}", data=arr)
|
||||
|
||||
|
||||
def load_h5(file_path: str, share_memory=True) -> Dict[str, List[Tensor]]:
|
||||
tensor_group: Dict[str, List[Tensor]] = {}
|
||||
|
||||
root_path = Path(file_path)
|
||||
if root_path.is_file() and root_path.suffix in (".h5", ".hdf5"):
|
||||
h5_files = [root_path]
|
||||
else:
|
||||
h5_files = list(root_path.rglob("*.h5")) + list(root_path.rglob("*.hdf5"))
|
||||
|
||||
for h5_file in h5_files:
|
||||
with h5py.File(h5_file, "r") as f:
|
||||
for key in f.keys():
|
||||
grp = f[key]
|
||||
dsets = []
|
||||
for dset_name in grp.keys():
|
||||
dset = grp[dset_name]
|
||||
tensor = torch.from_numpy(dset[:])
|
||||
if share_memory:
|
||||
tensor = tensor.share_memory_()
|
||||
dsets.append(tensor)
|
||||
|
||||
if tensor_group.get(key) is None:
|
||||
tensor_group[key] = []
|
||||
tensor_group[key].extend(dsets)
|
||||
|
||||
return tensor_group
|
||||
|
||||
|
||||
def save_bin(
|
||||
file_path: str,
|
||||
tensor_group: Dict[str, List[Tensor]],
|
||||
@@ -65,7 +24,7 @@ def save_bin(
|
||||
offsets, preserving backward compatibility.
|
||||
|
||||
Nested keys (``List[List[Tensor]]`` such as GRPO ``responses``) are
|
||||
not supported in bin format — use H5 for those.
|
||||
not supported in bin format — use JSONL for those.
|
||||
"""
|
||||
os.makedirs(file_path, exist_ok=True)
|
||||
record_keys = set(record_keys or [])
|
||||
@@ -74,7 +33,7 @@ def save_bin(
|
||||
if tensors and isinstance(tensors[0], list):
|
||||
raise ValueError(
|
||||
f"Nested key '{key}' (List[List[Tensor]]) is not supported "
|
||||
f"in bin format. Use H5 or JSONL storage instead."
|
||||
f"in bin format. Use JSONL storage instead."
|
||||
)
|
||||
cat = torch.cat(tensors, dim=0)
|
||||
entry: Dict[str, Any] = {
|
||||
@@ -112,7 +71,7 @@ def load_bin_offsets(file_path: str) -> Dict[str, List[int]]:
|
||||
|
||||
Returns an empty dict when no key has offsets (legacy bin files),
|
||||
in which case record-mode access falls back to per-record segment
|
||||
indexing (H5/JSONL layout).
|
||||
indexing (JSONL layout).
|
||||
"""
|
||||
with open(os.path.join(file_path, "meta.json"), "r") as f:
|
||||
meta = json.load(f)
|
||||
|
||||
@@ -38,12 +38,27 @@ class ChatTemplate:
|
||||
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. By deferring compilation to first access, the
|
||||
default pickle protocol serialises only ``template_str``; each
|
||||
worker rebuilds the cache on first render.
|
||||
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
|
||||
def from_string(
|
||||
cls,
|
||||
|
||||
@@ -20,8 +20,6 @@ Messages = List[Message]
|
||||
class AutoTokenizer:
|
||||
"""Base tokenizer class with automatic loading support"""
|
||||
|
||||
TOKENIZER_CLASSES = {} # Registry for auto-loading
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
path: Optional[Union[str, Path]] = None,
|
||||
@@ -108,17 +106,6 @@ class AutoTokenizer:
|
||||
with open(save_path / "tokenizer_config.json", "w", encoding="utf-8") as f:
|
||||
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(
|
||||
self,
|
||||
tokens: Union[str, List[str]],
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import math
|
||||
from typing import Dict
|
||||
|
||||
import torch
|
||||
@@ -22,6 +23,54 @@ def grad_norm(model: nn.Module, per_param: bool = False) -> float | Dict[str, fl
|
||||
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 reported value is the power ratio in decibels: ``10 * log10(SNR)``.
|
||||
|
||||
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
|
||||
snr = total_signal / (total_noise + self.eps)
|
||||
return 10.0 * math.log10(max(snr, self.eps))
|
||||
|
||||
|
||||
def ctx_get_loss(ctx):
|
||||
return ctx.loss
|
||||
|
||||
@@ -36,3 +85,30 @@ def ctx_get_val_loss(ctx):
|
||||
|
||||
def ctx_get_grad_norm(ctx):
|
||||
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
|
||||
|
||||
|
||||
def ctx_get_moe_aux_loss(ctx):
|
||||
return ctx.strategy._moe_metrics.get("aux_loss")
|
||||
|
||||
|
||||
def ctx_get_router_entropy(ctx):
|
||||
return ctx.strategy._moe_metrics.get("router_entropy")
|
||||
|
||||
|
||||
def ctx_get_dead_expert_fraction(ctx):
|
||||
return ctx.strategy._moe_metrics.get("dead_expert_fraction")
|
||||
|
||||
|
||||
def ctx_get_load_imbalance_mean(ctx):
|
||||
return ctx.strategy._moe_metrics.get("load_imbalance_mean")
|
||||
|
||||
|
||||
def ctx_get_load_imbalance_max(ctx):
|
||||
return ctx.strategy._moe_metrics.get("load_imbalance_max")
|
||||
|
||||
@@ -6,7 +6,7 @@ Provides:
|
||||
- :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`
|
||||
:class:`~astrai.inference.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)``
|
||||
@@ -20,7 +20,7 @@ from typing import Dict, List, Optional, Tuple
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.inference.core.scheduler import InferenceScheduler
|
||||
from astrai.inference.scheduler import InferenceScheduler
|
||||
|
||||
|
||||
@dataclass(kw_only=True)
|
||||
@@ -101,7 +101,7 @@ 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`,
|
||||
:meth:`~astrai.inference.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.
|
||||
|
||||
+204
-40
@@ -1,7 +1,7 @@
|
||||
"""Training strategy implementations with factory pattern."""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Callable, Dict, Union
|
||||
from typing import Callable, Dict, List, Optional, TypedDict, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
@@ -9,18 +9,20 @@ import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.model.components.mlp import RouterStats
|
||||
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
|
||||
class LossOutput(TypedDict):
|
||||
loss: Tensor
|
||||
metrics: Dict[str, float]
|
||||
|
||||
|
||||
class LogprobsOutput(TypedDict):
|
||||
logprobs: Tensor
|
||||
aux_loss: Optional[Tensor]
|
||||
router_stats: Optional[List[RouterStats]]
|
||||
|
||||
|
||||
def move_to_device(batch: Dict[str, Tensor], device: str) -> Dict[str, Tensor]:
|
||||
@@ -34,7 +36,7 @@ def get_logprobs(
|
||||
attn_mask: Tensor,
|
||||
loss_mask: Tensor,
|
||||
reduction: str,
|
||||
) -> Tensor:
|
||||
) -> LogprobsOutput:
|
||||
"""Compute token-wise log probabilities from model outputs.
|
||||
|
||||
Args:
|
||||
@@ -56,10 +58,11 @@ def get_logprobs(
|
||||
shifted_input_ids = input_ids[:, 1:]
|
||||
shifted_loss_mask = loss_mask[:, 1:]
|
||||
|
||||
logits = model(
|
||||
outputs = model(
|
||||
input_ids[:, :-1],
|
||||
attn_mask[:, :, :-1, :-1] if attn_mask.dim() == 4 else attn_mask[:, :-1],
|
||||
)["logits"]
|
||||
)
|
||||
logits = outputs["logits"]
|
||||
log_probs = torch.log_softmax(logits.float(), dim=-1)
|
||||
|
||||
token_logprobs = torch.gather(
|
||||
@@ -67,13 +70,18 @@ def get_logprobs(
|
||||
).squeeze(-1)
|
||||
|
||||
if reduction == "mean":
|
||||
return (token_logprobs * shifted_loss_mask).sum(dim=-1) / shifted_loss_mask.sum(
|
||||
logprobs = (token_logprobs * shifted_loss_mask).sum(
|
||||
dim=-1
|
||||
).clamp(min=1.0)
|
||||
) / shifted_loss_mask.sum(dim=-1).clamp(min=1.0)
|
||||
elif reduction == "sum":
|
||||
return (token_logprobs * shifted_loss_mask).sum(dim=-1)
|
||||
logprobs = (token_logprobs * shifted_loss_mask).sum(dim=-1)
|
||||
else:
|
||||
return token_logprobs * shifted_loss_mask
|
||||
logprobs = token_logprobs * shifted_loss_mask
|
||||
return {
|
||||
"logprobs": logprobs,
|
||||
"aux_loss": outputs.get("aux_loss"),
|
||||
"router_stats": outputs.get("router_stats"),
|
||||
}
|
||||
|
||||
|
||||
def make_doc_boundary_mask(position_ids: Tensor) -> Tensor:
|
||||
@@ -92,6 +100,68 @@ def make_doc_boundary_mask(position_ids: Tensor) -> Tensor:
|
||||
return (same_doc & causal).unsqueeze(1)
|
||||
|
||||
|
||||
def _collect_moe_diagnostics(
|
||||
router_stats_list: List[RouterStats],
|
||||
) -> Dict[str, float]:
|
||||
"""Collect MoE routing diagnostic metrics from per-layer router stats.
|
||||
|
||||
Args:
|
||||
router_stats_list: One :class:`RouterStats` dict per MoE layer with
|
||||
keys ``probs`` (N, E) and ``topk_indices`` (N, K), both detached.
|
||||
|
||||
Returns:
|
||||
Dict with keys: router_entropy, dead_expert_fraction,
|
||||
load_imbalance_mean, load_imbalance_max. Values are averaged
|
||||
across layers.
|
||||
"""
|
||||
layer_entropies: List[Tensor] = []
|
||||
layer_dead_fractions: List[Tensor] = []
|
||||
layer_imbalance_means: List[Tensor] = []
|
||||
layer_imbalance_maxs: List[Tensor] = []
|
||||
|
||||
for stats in router_stats_list:
|
||||
probs = stats["probs"].float()
|
||||
topk_indices = stats["topk_indices"]
|
||||
num_experts = probs.shape[-1]
|
||||
if num_experts == 0:
|
||||
continue
|
||||
probs = probs.reshape(-1, num_experts)
|
||||
if probs.numel() == 0:
|
||||
continue
|
||||
|
||||
# Router entropy
|
||||
entropy = -(probs * torch.log(probs.clamp_min(1e-8))).sum(dim=-1).mean()
|
||||
|
||||
# Load from the actual dispatch: one-hot sum of top-k assignments.
|
||||
expert_counts = F.one_hot(topk_indices, num_experts).sum(dim=(0, 1)).float()
|
||||
ideal_load = expert_counts.mean() # N*K / E
|
||||
load_ratios = expert_counts / max(float(ideal_load), 1.0)
|
||||
imbalance_mean = (load_ratios - 1.0).abs().mean()
|
||||
imbalance_max = load_ratios.max()
|
||||
dead_fraction = (expert_counts == 0).float().mean()
|
||||
|
||||
layer_entropies.append(entropy)
|
||||
layer_dead_fractions.append(dead_fraction)
|
||||
layer_imbalance_means.append(imbalance_mean)
|
||||
layer_imbalance_maxs.append(imbalance_max)
|
||||
|
||||
if not layer_entropies:
|
||||
return {}
|
||||
|
||||
return {
|
||||
"router_entropy": float(torch.stack(layer_entropies).mean().cpu().item()),
|
||||
"dead_expert_fraction": float(
|
||||
torch.stack(layer_dead_fractions).mean().cpu().item()
|
||||
),
|
||||
"load_imbalance_mean": float(
|
||||
torch.stack(layer_imbalance_means).mean().cpu().item()
|
||||
),
|
||||
"load_imbalance_max": float(
|
||||
torch.stack(layer_imbalance_maxs).mean().cpu().item()
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
class BaseStrategy(ABC):
|
||||
"""Abstract base class for training strategies.
|
||||
|
||||
@@ -112,6 +182,8 @@ class BaseStrategy(ABC):
|
||||
self.model = model
|
||||
self.device = device
|
||||
self.executor = kwargs.pop("executor", None)
|
||||
self.moe_aux_loss_coef = kwargs.pop("moe_aux_loss_coef", 0.01)
|
||||
self._moe_metrics: Dict[str, float] = {}
|
||||
self.extra_kwargs = kwargs
|
||||
self._rollout_runner = None
|
||||
|
||||
@@ -127,6 +199,35 @@ class BaseStrategy(ABC):
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
|
||||
return self._normalize_output(self.compute_loss(batch))
|
||||
|
||||
def _loss_output(
|
||||
self,
|
||||
task_loss: Tensor,
|
||||
metrics: Dict[str, Tensor],
|
||||
aux_loss: Optional[Tensor] = None,
|
||||
router_stats: Optional[List[RouterStats]] = None,
|
||||
) -> LossOutput:
|
||||
total_loss = task_loss
|
||||
if aux_loss is not None:
|
||||
weighted_aux_loss = self.moe_aux_loss_coef * aux_loss
|
||||
total_loss = total_loss + weighted_aux_loss
|
||||
metrics["moe_aux_loss"] = aux_loss
|
||||
metrics["moe_aux_loss_weighted"] = weighted_aux_loss
|
||||
self._refresh_moe_diagnostics(aux_loss, router_stats)
|
||||
metrics["loss"] = total_loss
|
||||
return {
|
||||
"loss": total_loss,
|
||||
"metrics": {name: value.detach().item() for name, value in metrics.items()},
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _normalize_output(output: Union[LossOutput, Tensor]) -> LossOutput:
|
||||
if isinstance(output, dict):
|
||||
return output
|
||||
return {"loss": output, "metrics": {"loss": output.detach().item()}}
|
||||
|
||||
def supports_online(self) -> bool:
|
||||
"""Whether this strategy can operate with a rollout runner.
|
||||
|
||||
@@ -158,22 +259,36 @@ class BaseStrategy(ABC):
|
||||
"""
|
||||
pass
|
||||
|
||||
def _refresh_moe_diagnostics(
|
||||
self,
|
||||
aux_loss: Tensor,
|
||||
router_stats: Optional[List[RouterStats]] = None,
|
||||
) -> None:
|
||||
"""Collect MoE routing diagnostics from the latest forward pass.
|
||||
|
||||
Populates ``self._moe_metrics`` with router entropy, dead expert
|
||||
fraction, load imbalance, and aux_loss. Called from
|
||||
:meth:`_loss_output` when an MoE aux loss is present.
|
||||
"""
|
||||
self._moe_metrics = _collect_moe_diagnostics(router_stats or [])
|
||||
self._moe_metrics["aux_loss"] = float(aux_loss.detach().cpu().item())
|
||||
|
||||
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]) -> LossOutput:
|
||||
"""Run offline or online forward depending on runner injection."""
|
||||
if self._rollout_runner is None:
|
||||
return self.compute_loss(batch)
|
||||
return self.compute_loss_output(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)
|
||||
return self.compute_loss_output(train_batch)
|
||||
|
||||
|
||||
class StrategyFactory(BaseFactory["BaseStrategy"]):
|
||||
@@ -200,6 +315,7 @@ class SEQStrategy(BaseStrategy):
|
||||
"""Standard next-token prediction training strategy.
|
||||
|
||||
Computes cross-entropy loss for next token prediction.
|
||||
Optionally adds MoE load balancing auxiliary loss.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
@@ -213,9 +329,13 @@ class SEQStrategy(BaseStrategy):
|
||||
self.label_smoothing = label_smoothing
|
||||
|
||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||
return self.compute_loss_output(batch)["loss"]
|
||||
|
||||
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
|
||||
batch = move_to_device(batch, self.device)
|
||||
input_ids, target_ids = batch["input_ids"], batch["target_ids"]
|
||||
logits = self.model(input_ids=input_ids)["logits"]
|
||||
outputs = self.model(input_ids=input_ids)
|
||||
logits = outputs["logits"]
|
||||
|
||||
loss = F.cross_entropy(
|
||||
input=logits.flatten(0, 1).float(),
|
||||
@@ -223,7 +343,12 @@ class SEQStrategy(BaseStrategy):
|
||||
label_smoothing=self.label_smoothing,
|
||||
)
|
||||
|
||||
return loss
|
||||
return self._loss_output(
|
||||
loss,
|
||||
{"task_loss": loss},
|
||||
outputs.get("aux_loss"),
|
||||
outputs.get("router_stats"),
|
||||
)
|
||||
|
||||
|
||||
@StrategyFactory.register("sft")
|
||||
@@ -231,6 +356,7 @@ class SFTStrategy(BaseStrategy):
|
||||
"""Supervised Fine-tuning strategy with loss masking.
|
||||
|
||||
Applies cross-entropy loss only to tokens where loss_mask is True.
|
||||
Optionally adds MoE load balancing auxiliary loss.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
@@ -244,6 +370,9 @@ class SFTStrategy(BaseStrategy):
|
||||
self.label_smoothing = label_smoothing
|
||||
|
||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||
return self.compute_loss_output(batch)["loss"]
|
||||
|
||||
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
|
||||
batch = move_to_device(batch, self.device)
|
||||
input_ids, target_ids, position_ids, loss_mask = (
|
||||
batch["input_ids"],
|
||||
@@ -255,9 +384,10 @@ class SFTStrategy(BaseStrategy):
|
||||
ignore_index = -100
|
||||
input_mask = make_doc_boundary_mask(position_ids)
|
||||
target_ids = target_ids.masked_fill(~loss_mask, ignore_index)
|
||||
logits = self.model(
|
||||
outputs = self.model(
|
||||
input_ids=input_ids, position_ids=position_ids, input_mask=input_mask
|
||||
)["logits"]
|
||||
)
|
||||
logits = outputs["logits"]
|
||||
|
||||
loss = F.cross_entropy(
|
||||
input=logits.flatten(0, 1).float(),
|
||||
@@ -266,7 +396,12 @@ class SFTStrategy(BaseStrategy):
|
||||
label_smoothing=self.label_smoothing,
|
||||
)
|
||||
|
||||
return loss
|
||||
return self._loss_output(
|
||||
loss,
|
||||
{"task_loss": loss},
|
||||
outputs.get("aux_loss"),
|
||||
outputs.get("router_stats"),
|
||||
)
|
||||
|
||||
|
||||
@StrategyFactory.register("dpo")
|
||||
@@ -292,6 +427,9 @@ class DPOStrategy(BaseStrategy):
|
||||
self.reduction = reduction
|
||||
|
||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||
return self.compute_loss_output(batch)["loss"]
|
||||
|
||||
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
|
||||
batch = move_to_device(batch, self.device)
|
||||
chosen_ids, rejected_ids = batch["chosen"], batch["rejected"]
|
||||
chosen_mask, rejected_mask = batch["chosen_mask"], batch["rejected_mask"]
|
||||
@@ -307,22 +445,25 @@ class DPOStrategy(BaseStrategy):
|
||||
)[None, None, :, :] # [1, 1, S, S]
|
||||
full_mask = key_pad & causal # [B*2, 1, S, S] — composed
|
||||
|
||||
log_pi = get_logprobs(
|
||||
policy_output = get_logprobs(
|
||||
self.model,
|
||||
concat_ids,
|
||||
full_mask,
|
||||
concat_loss_mask,
|
||||
self.reduction,
|
||||
)
|
||||
log_pi = policy_output["logprobs"]
|
||||
aux_loss = policy_output["aux_loss"]
|
||||
|
||||
with torch.no_grad():
|
||||
log_ref = get_logprobs(
|
||||
ref_output = get_logprobs(
|
||||
self.ref_model,
|
||||
concat_ids,
|
||||
full_mask,
|
||||
concat_loss_mask,
|
||||
self.reduction,
|
||||
)
|
||||
log_ref = ref_output["logprobs"]
|
||||
|
||||
log_pi_chosen = log_pi[: chosen_ids.shape[0]]
|
||||
log_pi_rejected = log_pi[chosen_ids.shape[0] :]
|
||||
@@ -335,7 +476,12 @@ class DPOStrategy(BaseStrategy):
|
||||
ratio_diff = pi_log_ratio - ref_log_ratio
|
||||
dpo_loss = -F.logsigmoid(self.beta * ratio_diff).mean()
|
||||
|
||||
return dpo_loss
|
||||
return self._loss_output(
|
||||
dpo_loss,
|
||||
{"dpo_loss": dpo_loss},
|
||||
aux_loss,
|
||||
policy_output.get("router_stats"),
|
||||
)
|
||||
|
||||
def supports_online(self) -> bool:
|
||||
return True
|
||||
@@ -401,9 +547,16 @@ class GRPOStrategy(BaseStrategy):
|
||||
|
||||
def sync_old_model(self):
|
||||
"""Copy current policy weights to old model."""
|
||||
self.old_model.load_state_dict(self.executor.unwrap_model(self.model))
|
||||
state_dict = self.executor.unwrap_model(self.model)
|
||||
if self.executor.use_distributed:
|
||||
state_dict = broadcast_state_dict(state_dict)
|
||||
if state_dict is not None:
|
||||
self.old_model.load_state_dict(state_dict)
|
||||
|
||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||
return self.compute_loss_output(batch)["loss"]
|
||||
|
||||
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
|
||||
batch = move_to_device(batch, self.device)
|
||||
prompts = batch["prompts"]
|
||||
responses = batch["responses"]
|
||||
@@ -444,16 +597,23 @@ class GRPOStrategy(BaseStrategy):
|
||||
# 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(
|
||||
policy_output = get_logprobs(
|
||||
self.model, full_sequences, attn_mask, full_masks, "none"
|
||||
)[:, prompt_len - 1 :]
|
||||
)
|
||||
token_log_probs_policy = policy_output["logprobs"]
|
||||
aux_loss = policy_output["aux_loss"]
|
||||
token_log_probs_policy = token_log_probs_policy[:, prompt_len - 1 :]
|
||||
with torch.no_grad():
|
||||
token_log_probs_old = get_logprobs(
|
||||
old_output = get_logprobs(
|
||||
self.old_model, full_sequences, attn_mask, full_masks, "none"
|
||||
)[:, prompt_len - 1 :]
|
||||
token_log_probs_ref = get_logprobs(
|
||||
)
|
||||
token_log_probs_old = old_output["logprobs"]
|
||||
token_log_probs_old = token_log_probs_old[:, prompt_len - 1 :]
|
||||
ref_output = get_logprobs(
|
||||
self.ref_model, full_sequences, attn_mask, full_masks, "none"
|
||||
)[:, prompt_len - 1 :]
|
||||
)
|
||||
token_log_probs_ref = ref_output["logprobs"]
|
||||
token_log_probs_ref = token_log_probs_ref[:, prompt_len - 1 :]
|
||||
|
||||
# Reshape to [B, G, response_len]
|
||||
token_log_probs_policy = token_log_probs_policy.view(batch_size, group_size, -1)
|
||||
@@ -486,9 +646,13 @@ class GRPOStrategy(BaseStrategy):
|
||||
kl_per_token = r - torch.log(r + eps) - 1.0
|
||||
kl_penalty = self.kl_coef * (kl_per_token * token_masks).sum() / token_count
|
||||
|
||||
total_loss = policy_loss + kl_penalty
|
||||
|
||||
return total_loss
|
||||
task_loss = policy_loss + kl_penalty
|
||||
return self._loss_output(
|
||||
task_loss,
|
||||
{"policy_loss": policy_loss, "kl_loss": kl_penalty},
|
||||
aux_loss,
|
||||
policy_output.get("router_stats"),
|
||||
)
|
||||
|
||||
def supports_online(self) -> bool:
|
||||
return True
|
||||
@@ -510,5 +674,5 @@ class GRPOStrategy(BaseStrategy):
|
||||
# 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._entries["online_grpo"] = GRPOStrategy
|
||||
StrategyFactory._entries["online_dpo"] = DPOStrategy
|
||||
StrategyFactory.register("online_grpo")(GRPOStrategy)
|
||||
StrategyFactory.register("online_dpo")(DPOStrategy)
|
||||
|
||||
@@ -17,9 +17,15 @@ from astrai.parallel import only_on_rank
|
||||
from astrai.parallel.setup import get_current_device
|
||||
from astrai.serialization import Checkpoint
|
||||
from astrai.trainer.metric_util import (
|
||||
ctx_get_dead_expert_fraction,
|
||||
ctx_get_grad_norm,
|
||||
ctx_get_grad_snr,
|
||||
ctx_get_load_imbalance_max,
|
||||
ctx_get_load_imbalance_mean,
|
||||
ctx_get_loss,
|
||||
ctx_get_lr,
|
||||
ctx_get_moe_aux_loss,
|
||||
ctx_get_router_entropy,
|
||||
ctx_get_val_loss,
|
||||
)
|
||||
from astrai.trainer.train_context import TrainContext
|
||||
@@ -235,7 +241,7 @@ class ProgressBarCallback(TrainCallback):
|
||||
class MetricCallback(TrainCallback):
|
||||
def __init__(
|
||||
self,
|
||||
log_dir: str,
|
||||
ckpt_dir: str,
|
||||
save_interval: int,
|
||||
metrics: List[str] = None,
|
||||
val_step: int = 0,
|
||||
@@ -246,8 +252,7 @@ class MetricCallback(TrainCallback):
|
||||
self.val_step = val_step
|
||||
self._next_val_step = 0
|
||||
|
||||
self.log_dir = Path(log_dir) if log_dir else Path.cwd() / "logs"
|
||||
self.log_dir.mkdir(parents=True, exist_ok=True)
|
||||
self.ckpt_dir = Path(ckpt_dir) if ckpt_dir else Path.cwd() / "checkpoint"
|
||||
|
||||
self.log_cache = []
|
||||
|
||||
@@ -256,14 +261,37 @@ class MetricCallback(TrainCallback):
|
||||
"lr": ctx_get_lr,
|
||||
"val_loss": ctx_get_val_loss,
|
||||
"grad_norm": ctx_get_grad_norm,
|
||||
"grad_snr": ctx_get_grad_snr,
|
||||
"moe_aux_loss": ctx_get_moe_aux_loss,
|
||||
"router_entropy": ctx_get_router_entropy,
|
||||
"dead_expert_fraction": ctx_get_dead_expert_fraction,
|
||||
"load_imbalance_mean": ctx_get_load_imbalance_mean,
|
||||
"load_imbalance_max": ctx_get_load_imbalance_max,
|
||||
}
|
||||
|
||||
def _metrics(self, context: TrainContext, names):
|
||||
return {
|
||||
m: self._metric_funcs[m](context)
|
||||
for m in names
|
||||
if self._metric_funcs[m](context) is not None
|
||||
}
|
||||
metrics = dict(context.metrics)
|
||||
for name in names:
|
||||
metric_fn = self._metric_funcs.get(name)
|
||||
if metric_fn is None:
|
||||
continue
|
||||
value = metric_fn(context)
|
||||
if value is not None:
|
||||
metrics[name] = value
|
||||
selected = set(context.metrics) | set(names)
|
||||
selected.discard("*")
|
||||
result = {name: metrics[name] for name in selected if name in metrics}
|
||||
if context.world_size > 1 and dist.is_initialized() and result:
|
||||
metric_names = sorted(result)
|
||||
values = torch.tensor(
|
||||
[result[name] for name in metric_names],
|
||||
dtype=torch.float32,
|
||||
device=get_current_device(),
|
||||
)
|
||||
dist.all_reduce(values, op=dist.ReduceOp.SUM)
|
||||
values /= context.world_size
|
||||
result.update(zip(metric_names, values.tolist()))
|
||||
return result
|
||||
|
||||
@only_on_rank(0)
|
||||
def _append(self, event_type: str, context: TrainContext, **extra):
|
||||
@@ -285,8 +313,8 @@ class MetricCallback(TrainCallback):
|
||||
|
||||
with torch.no_grad():
|
||||
for batch in context.val_dataloader:
|
||||
loss = context.strategy(batch)
|
||||
total_loss += loss.item()
|
||||
loss_output = context.strategy(batch)
|
||||
total_loss += loss_output["loss"].item()
|
||||
num_batches += 1
|
||||
|
||||
if context.world_size > 1 and dist.is_initialized():
|
||||
@@ -306,13 +334,15 @@ class MetricCallback(TrainCallback):
|
||||
|
||||
@only_on_rank(0)
|
||||
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)
|
||||
with open(log_file, "w") as f:
|
||||
for log in self.log_cache:
|
||||
f.write(json.dumps(log) + "\n")
|
||||
|
||||
def on_optimizer_step(self, context):
|
||||
context.grad_snr_tracker.update(context.model)
|
||||
|
||||
if (
|
||||
context.val_dataloader is not None
|
||||
and self.val_step > 0
|
||||
|
||||
+193
-152
@@ -1,3 +1,4 @@
|
||||
import logging
|
||||
import threading
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
@@ -9,15 +10,18 @@ from torch.utils.data import DataLoader, random_split
|
||||
|
||||
from astrai.config.train_config import TrainConfig
|
||||
from astrai.dataset import RDSampler
|
||||
from astrai.inference.core.scheduler import InferenceScheduler
|
||||
from astrai.inference.scheduler import InferenceScheduler
|
||||
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.protocols import OptimizerProtocol, SchedulerProtocol
|
||||
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, create_ref_model
|
||||
from astrai.trainer.strategy import BaseStrategy, StrategyFactory
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -34,7 +38,9 @@ class TrainContext:
|
||||
epoch: int = field(default=0)
|
||||
consumed_samples: int = field(default=0)
|
||||
loss: float = field(default=0.0)
|
||||
metrics: Dict[str, float] = field(default_factory=dict)
|
||||
grad_norm: Optional[float] = field(default=None)
|
||||
grad_snr_tracker: GradSNRTracker = field(default_factory=GradSNRTracker)
|
||||
val_dataloader: Optional[DataLoader] = field(default=None)
|
||||
val_loss: Optional[float] = field(default=None)
|
||||
|
||||
@@ -60,6 +66,15 @@ class TrainContext:
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _PreloadedState:
|
||||
model_config: dict = field(default_factory=dict)
|
||||
state_dict: Optional[dict] = None
|
||||
epoch: int = 0
|
||||
consumed_samples: int = 0
|
||||
checkpoint: Optional[Checkpoint] = None
|
||||
|
||||
|
||||
class TrainContextBuilder:
|
||||
def __init__(
|
||||
self,
|
||||
@@ -75,205 +90,231 @@ class TrainContextBuilder:
|
||||
return self
|
||||
|
||||
def build(self) -> TrainContext:
|
||||
cfg = self.config
|
||||
device = get_current_device()
|
||||
# Resolve persisted state.
|
||||
preloaded_state = self._load_preloaded_state()
|
||||
|
||||
executor = ExecutorFactory.create(
|
||||
# Build the core training components and restore their persisted state.
|
||||
executor = self._create_executor()
|
||||
context = self._create_context(preloaded_state, executor)
|
||||
self._prepare_model(context, executor, preloaded_state)
|
||||
self._restore_optimizer_state(context)
|
||||
|
||||
# Resolve datasets.
|
||||
train_dataset, val_dataset = self._get_datasets()
|
||||
self._create_dataloaders(context, train_dataset, val_dataset)
|
||||
|
||||
# Strategies depend on the prepared model; online rollout depends on both.
|
||||
strategy_kwargs = self._create_strategy(context, executor)
|
||||
self._configure_rollout(context, strategy_kwargs)
|
||||
|
||||
return context
|
||||
|
||||
def _create_executor(self) -> BaseExecutor:
|
||||
cfg = self.config
|
||||
return ExecutorFactory.create(
|
||||
cfg.parallel_mode,
|
||||
grad_accum_steps=cfg.grad_accum_steps,
|
||||
**cfg.executor_kwargs,
|
||||
)
|
||||
|
||||
model_config = {}
|
||||
def _load_preloaded_state(self) -> _PreloadedState:
|
||||
cfg = self.config
|
||||
state = _PreloadedState(
|
||||
epoch=cfg.start_epoch,
|
||||
consumed_samples=cfg.start_samples * get_world_size(),
|
||||
)
|
||||
if self._param_path:
|
||||
config_path = Path(self._param_path) / "config.json"
|
||||
if config_path.exists():
|
||||
model_config = load_json(config_path)
|
||||
|
||||
preloaded_state_dict = None
|
||||
preloaded_epoch = cfg.start_epoch
|
||||
preloaded_consumed = cfg.start_samples * get_world_size()
|
||||
preloaded_checkpoint = None
|
||||
if self._param_path:
|
||||
state.model_config = load_json(config_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
|
||||
state.state_dict = checkpoint.state_dict
|
||||
state.model_config = checkpoint.config or state.model_config
|
||||
if self._resume:
|
||||
preloaded_epoch = checkpoint.epoch or cfg.start_epoch
|
||||
if checkpoint.consumed_samples > 0:
|
||||
per_step = (
|
||||
cfg.batch_per_device
|
||||
* get_world_size()
|
||||
* cfg.grad_accum_steps
|
||||
)
|
||||
preloaded_consumed = (
|
||||
checkpoint.consumed_samples // per_step
|
||||
) * per_step
|
||||
else:
|
||||
preloaded_consumed = cfg.start_samples * get_world_size()
|
||||
preloaded_checkpoint = checkpoint
|
||||
state.epoch = checkpoint.epoch
|
||||
per_step = (
|
||||
cfg.batch_per_device * get_world_size() * cfg.grad_accum_steps
|
||||
)
|
||||
state.consumed_samples = (
|
||||
checkpoint.consumed_samples // per_step * per_step
|
||||
)
|
||||
state.checkpoint = checkpoint
|
||||
if not state.model_config and hasattr(cfg.model_fn(), "config"):
|
||||
state.model_config = cfg.model_fn().config.to_dict()
|
||||
return state
|
||||
|
||||
if not model_config and hasattr(cfg.model_fn(), "config"):
|
||||
model_config = cfg.model_fn().config.to_dict()
|
||||
def _create_context(
|
||||
self, state: _PreloadedState, executor: BaseExecutor
|
||||
) -> TrainContext:
|
||||
return TrainContext(
|
||||
world_size=get_world_size(),
|
||||
rank=get_rank(),
|
||||
config=self.config,
|
||||
model_config=state.model_config,
|
||||
executor=executor,
|
||||
epoch=state.epoch,
|
||||
consumed_samples=state.consumed_samples,
|
||||
checkpoint=state.checkpoint,
|
||||
)
|
||||
|
||||
def _before_wrap(m):
|
||||
m = m.to(device=device)
|
||||
def _prepare_model(
|
||||
self, context: TrainContext, executor: BaseExecutor, state: _PreloadedState
|
||||
) -> None:
|
||||
cfg = self.config
|
||||
device = get_current_device()
|
||||
|
||||
def before_wrap(model):
|
||||
model = model.to(device=device)
|
||||
if cfg.lora is not None:
|
||||
inject_lora(
|
||||
m,
|
||||
model,
|
||||
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
|
||||
if state.state_dict is not None:
|
||||
model.load_state_dict(state.state_dict, strict=False)
|
||||
return model
|
||||
|
||||
context = TrainContext(
|
||||
world_size=get_world_size(),
|
||||
rank=get_rank(),
|
||||
config=cfg,
|
||||
model_config=model_config,
|
||||
executor=executor,
|
||||
epoch=preloaded_epoch,
|
||||
consumed_samples=preloaded_consumed,
|
||||
checkpoint=preloaded_checkpoint,
|
||||
)
|
||||
def after_wrap(model):
|
||||
if cfg.compile_mode is not None:
|
||||
logger.info("torch.compile enabled (mode=%s)", cfg.compile_mode)
|
||||
model = torch.compile(model, mode=cfg.compile_mode)
|
||||
return model
|
||||
|
||||
context.model, context.optimizer, context.scheduler = executor.prepare(
|
||||
cfg.model_fn,
|
||||
cfg.optimizer_fn,
|
||||
cfg.scheduler_fn,
|
||||
before_wrap=_before_wrap,
|
||||
before_wrap=before_wrap,
|
||||
after_wrap=after_wrap,
|
||||
)
|
||||
|
||||
train_dataset = cfg.dataset
|
||||
val_dataset = cfg.val_dataset
|
||||
def _get_datasets(self):
|
||||
cfg = self.config
|
||||
if cfg.val_dataset is not None or cfg.val_split is None:
|
||||
return cfg.dataset, cfg.val_dataset
|
||||
n_val = max(1, int(len(cfg.dataset) * cfg.val_split))
|
||||
generator = torch.Generator().manual_seed(cfg.random_seed)
|
||||
return random_split(
|
||||
cfg.dataset, [len(cfg.dataset) - n_val, n_val], generator=generator
|
||||
)
|
||||
|
||||
if val_dataset is None and cfg.val_split is not None:
|
||||
n_total = len(cfg.dataset)
|
||||
n_val = max(1, int(n_total * cfg.val_split))
|
||||
n_train = n_total - n_val
|
||||
generator = torch.Generator().manual_seed(cfg.random_seed)
|
||||
train_dataset, val_dataset = random_split(
|
||||
cfg.dataset, [n_train, n_val], generator=generator
|
||||
def _create_dataloaders(
|
||||
self, context: TrainContext, train_dataset, val_dataset
|
||||
) -> None:
|
||||
cfg = self.config
|
||||
sampler_offset = context.consumed_samples // context.world_size
|
||||
if self._resume and sampler_offset > 0:
|
||||
samples_per_replica = (
|
||||
len(train_dataset) + context.world_size - 1
|
||||
) // context.world_size
|
||||
if samples_per_replica > 0:
|
||||
context.epoch = sampler_offset // samples_per_replica
|
||||
context.dataloader = self._create_dataloader(
|
||||
train_dataset, context.epoch, sampler_offset
|
||||
)
|
||||
if val_dataset is not None:
|
||||
context.val_dataloader = self._create_dataloader(
|
||||
val_dataset, 0, 0, shuffle=False
|
||||
)
|
||||
|
||||
sampler_offset = context.consumed_samples // context.world_size
|
||||
def _create_dataloader(
|
||||
self, dataset, epoch: int, start_iter: int, shuffle: bool = True
|
||||
):
|
||||
cfg = self.config
|
||||
sampler = RDSampler(
|
||||
data_source=train_dataset,
|
||||
start_epoch=context.epoch,
|
||||
start_iter=sampler_offset,
|
||||
dataset,
|
||||
start_epoch=epoch,
|
||||
start_iter=start_iter,
|
||||
seed=cfg.random_seed,
|
||||
shuffle=shuffle,
|
||||
)
|
||||
context.dataloader = DataLoader(
|
||||
train_dataset,
|
||||
loader_kwargs = dict(
|
||||
dataset=dataset,
|
||||
batch_size=cfg.batch_per_device,
|
||||
sampler=sampler,
|
||||
num_workers=cfg.num_workers,
|
||||
pin_memory=cfg.pin_memory,
|
||||
prefetch_factor=cfg.prefetch_factor,
|
||||
collate_fn=cfg.collate_fn,
|
||||
)
|
||||
|
||||
if val_dataset is not None:
|
||||
val_sampler = RDSampler(
|
||||
data_source=val_dataset,
|
||||
start_epoch=0,
|
||||
start_iter=0,
|
||||
seed=cfg.random_seed,
|
||||
shuffle=False,
|
||||
)
|
||||
context.val_dataloader = DataLoader(
|
||||
val_dataset,
|
||||
batch_size=cfg.batch_per_device,
|
||||
sampler=val_sampler,
|
||||
num_workers=cfg.num_workers,
|
||||
pin_memory=cfg.pin_memory,
|
||||
prefetch_factor=cfg.prefetch_factor,
|
||||
collate_fn=cfg.collate_fn,
|
||||
)
|
||||
|
||||
if context.checkpoint and context.checkpoint.extra:
|
||||
extra = context.checkpoint.extra
|
||||
for name in ("optimizer", "scheduler"):
|
||||
if name in extra:
|
||||
obj = getattr(context, name, None)
|
||||
if obj is not None:
|
||||
obj.load_state_dict(extra[name])
|
||||
|
||||
strategy_kwargs = dict(cfg.extra_kwargs)
|
||||
|
||||
needs_ref = cfg.strategy in (
|
||||
"dpo",
|
||||
"grpo",
|
||||
"online_grpo",
|
||||
"online_dpo",
|
||||
# PyTorch rejects prefetch_factor/persistent_workers when workers=0.
|
||||
if cfg.num_workers > 0:
|
||||
loader_kwargs["persistent_workers"] = cfg.persistent_workers
|
||||
if cfg.prefetch_factor is not None:
|
||||
loader_kwargs["prefetch_factor"] = cfg.prefetch_factor
|
||||
return DataLoader(
|
||||
**loader_kwargs,
|
||||
)
|
||||
needs_old = cfg.strategy in ("grpo", "online_grpo")
|
||||
|
||||
if needs_ref:
|
||||
ref_model = create_ref_model(
|
||||
cfg.model_fn, executor.unwrap_model(context.model)
|
||||
).to(device=device)
|
||||
strategy_kwargs["ref_model"] = ref_model
|
||||
|
||||
old_model = None
|
||||
if needs_old:
|
||||
old_model = create_ref_model(
|
||||
cfg.model_fn, executor.unwrap_model(context.model)
|
||||
).to(device=device)
|
||||
strategy_kwargs["old_model"] = old_model
|
||||
def _restore_optimizer_state(self, context: TrainContext) -> None:
|
||||
if context.checkpoint and context.checkpoint.extra:
|
||||
for name in ("optimizer", "scheduler"):
|
||||
if (
|
||||
name in context.checkpoint.extra
|
||||
and getattr(context, name, None) is not None
|
||||
):
|
||||
getattr(context, name).load_state_dict(
|
||||
context.checkpoint.extra[name]
|
||||
)
|
||||
|
||||
def _create_strategy(self, context: TrainContext, executor: BaseExecutor) -> dict:
|
||||
cfg = self.config
|
||||
kwargs = dict(cfg.extra_kwargs)
|
||||
kwargs.setdefault("moe_aux_loss_coef", cfg.moe_aux_loss_coef)
|
||||
if cfg.strategy in ("dpo", "grpo", "online_grpo", "online_dpo"):
|
||||
kwargs["ref_model"] = create_ref_model(
|
||||
cfg.model_fn,
|
||||
executor=executor,
|
||||
model=context.model,
|
||||
device=get_current_device(),
|
||||
)
|
||||
if cfg.strategy in ("grpo", "online_grpo"):
|
||||
kwargs["old_model"] = create_ref_model(
|
||||
cfg.model_fn,
|
||||
executor=executor,
|
||||
model=context.model,
|
||||
device=get_current_device(),
|
||||
)
|
||||
context.strategy = StrategyFactory.create(
|
||||
cfg.strategy,
|
||||
model=context.model,
|
||||
device=device,
|
||||
device=get_current_device(),
|
||||
executor=executor,
|
||||
**strategy_kwargs,
|
||||
**kwargs,
|
||||
)
|
||||
return 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,
|
||||
max_prompt_len=max_seq_len or 4096,
|
||||
def _configure_rollout(self, context: TrainContext, strategy_kwargs: dict) -> None:
|
||||
cfg = self.config
|
||||
if not cfg.strategy.startswith("online_"):
|
||||
return
|
||||
if not context.strategy.supports_online():
|
||||
raise ValueError(
|
||||
f"Strategy '{cfg.strategy}' does not support online rollout"
|
||||
)
|
||||
|
||||
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(
|
||||
tokenizer = AutoTokenizer.from_pretrained(self._param_path)
|
||||
group_size = strategy_kwargs.get("group_size", 1)
|
||||
scheduler = InferenceScheduler(
|
||||
model=context.model,
|
||||
tokenizer=tokenizer,
|
||||
max_batch_size=group_size * max(1, cfg.batch_per_device),
|
||||
max_seq_len=getattr(context.model.config, "max_position_embeddings", None),
|
||||
)
|
||||
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,
|
||||
)
|
||||
context.strategy.set_rollout_runner(
|
||||
RolloutRunner(
|
||||
generator=generator,
|
||||
reward_model=reward_model,
|
||||
reward_model=cfg.reward_model_fn(),
|
||||
rollout_interval=cfg.rollout_interval,
|
||||
)
|
||||
context.strategy.set_rollout_runner(runner)
|
||||
|
||||
return context
|
||||
)
|
||||
|
||||
@@ -5,7 +5,7 @@ import torch.distributed as dist
|
||||
|
||||
from astrai.config import TrainConfig
|
||||
from astrai.parallel.setup import spawn_parallel_fn
|
||||
from astrai.parallel.signal_handler import (
|
||||
from astrai.signal_handler import (
|
||||
register_signal_handlers,
|
||||
unregister_signal_handlers,
|
||||
)
|
||||
@@ -42,7 +42,7 @@ class Trainer:
|
||||
),
|
||||
CallbackFactory.create(
|
||||
"metric",
|
||||
log_dir=cfg.log_dir,
|
||||
ckpt_dir=cfg.ckpt_dir,
|
||||
save_interval=cfg.ckpt_interval,
|
||||
metrics=cfg.metrics,
|
||||
val_step=cfg.val_step,
|
||||
@@ -82,9 +82,10 @@ class Trainer:
|
||||
break
|
||||
with executor.accumulate(context.model):
|
||||
self._call_callbacks("on_batch_begin", context)
|
||||
loss = context.strategy(batch)
|
||||
context.loss = loss.item()
|
||||
stand_loss = loss / executor.grad_accum_steps
|
||||
loss_output = context.strategy(batch)
|
||||
context.loss = loss_output["loss"].item()
|
||||
context.metrics = loss_output["metrics"]
|
||||
stand_loss = loss_output["loss"] / executor.grad_accum_steps
|
||||
executor.backward(stand_loss)
|
||||
context.consumed_samples += (
|
||||
context.config.batch_per_device * context.world_size
|
||||
|
||||
@@ -0,0 +1,77 @@
|
||||
cmake_minimum_required(VERSION 3.18)
|
||||
project(astrai_kernels LANGUAGES CUDA CXX)
|
||||
|
||||
set(CMAKE_CXX_STANDARD 17)
|
||||
set(CMAKE_CXX_STANDARD_REQUIRED ON)
|
||||
set(CMAKE_CUDA_STANDARD 17)
|
||||
|
||||
find_package(CUDAToolkit REQUIRED)
|
||||
|
||||
if(NOT DEFINED TORCH_HOME)
|
||||
set(TORCH_HOME "$ENV{TORCH_HOME}")
|
||||
endif()
|
||||
if(NOT TORCH_HOME)
|
||||
message(FATAL_ERROR "TORCH_HOME must point at the torch install dir (site-packages/torch)")
|
||||
endif()
|
||||
|
||||
if(NOT DEFINED PYTHON_INCLUDE_DIR)
|
||||
set(PYTHON_INCLUDE_DIR "/usr/include/python${PYTHON_VERSION_MAJOR}.${PYTHON_VERSION_MINOR}")
|
||||
endif()
|
||||
|
||||
if(NOT DEFINED ASTRAI_CUDA_ARCH)
|
||||
if(DEFINED ENV{ASTRAI_CUDA_ARCH})
|
||||
set(ASTRAI_CUDA_ARCH "$ENV{ASTRAI_CUDA_ARCH}")
|
||||
else()
|
||||
set(ASTRAI_CUDA_ARCH 80)
|
||||
endif()
|
||||
endif()
|
||||
|
||||
set(TORCH_LIB_DIR "${TORCH_HOME}/lib")
|
||||
set(CUDA_LIB_DIR "/usr/local/cuda/lib64")
|
||||
|
||||
set(CXX_FLAGS -O3 -funroll-loops)
|
||||
set(NVCC_FLAGS -O3
|
||||
--expt-relaxed-constexpr
|
||||
--use_fast_math
|
||||
"--ptxas-options=-O3,-v"
|
||||
--extra-device-vectorization
|
||||
--threads=16)
|
||||
|
||||
set(TORCH_LIBS
|
||||
"${TORCH_LIB_DIR}/libtorch_python.so"
|
||||
"${TORCH_LIB_DIR}/libtorch_cuda.so"
|
||||
"${TORCH_LIB_DIR}/libc10_cuda.so"
|
||||
"${TORCH_LIB_DIR}/libtorch_cpu.so"
|
||||
"${TORCH_LIB_DIR}/libtorch.so"
|
||||
"${TORCH_LIB_DIR}/libc10.so"
|
||||
CUDA::cudart)
|
||||
|
||||
set(CMAKE_CUDA_ARCHITECTURES "${ASTRAI_CUDA_ARCH}")
|
||||
|
||||
set(KERNELS attn_decode attn_prefill attn_paged_decode attn_paged_prefill rotary_emb fp8_mm)
|
||||
|
||||
foreach(name ${KERNELS})
|
||||
add_library(${name} MODULE "${CMAKE_CURRENT_SOURCE_DIR}/kernels/${name}.cu")
|
||||
|
||||
target_compile_definitions(${name} PRIVATE TORCH_EXTENSION_NAME=${name})
|
||||
|
||||
target_include_directories(${name} PRIVATE
|
||||
"${TORCH_HOME}/include"
|
||||
"${TORCH_HOME}/include/torch/csrc/api/include"
|
||||
"${PYTHON_INCLUDE_DIR}")
|
||||
|
||||
target_link_libraries(${name} PRIVATE ${TORCH_LIBS})
|
||||
if(${name} STREQUAL "fp8_mm")
|
||||
target_link_libraries(${name} PRIVATE CUDA::cublasLt)
|
||||
endif()
|
||||
target_link_options(${name} PRIVATE "-Wl,-rpath,${TORCH_LIB_DIR}")
|
||||
|
||||
target_compile_options(${name} PRIVATE
|
||||
$<$<COMPILE_LANGUAGE:CXX>:${CXX_FLAGS}>
|
||||
$<$<COMPILE_LANGUAGE:CUDA>:${NVCC_FLAGS}>)
|
||||
|
||||
set_target_properties(${name} PROPERTIES
|
||||
PREFIX ""
|
||||
SUFFIX ".${PY_SOABI}.so"
|
||||
LIBRARY_OUTPUT_DIRECTORY "${CMAKE_CURRENT_SOURCE_DIR}/../astrai/extension/lib")
|
||||
endforeach()
|
||||
@@ -1,48 +0,0 @@
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def _arch_flags() -> list[str]:
|
||||
import torch
|
||||
|
||||
if torch.cuda.is_available():
|
||||
cap = torch.cuda.get_device_capability()
|
||||
else:
|
||||
cap = (8, 0)
|
||||
ver = f"{cap[0]}{cap[1]}"
|
||||
flags = [f"-gencode=arch=compute_{ver},code=sm_{ver}"]
|
||||
# tensor-core mma path (mma.sync.m16n8k16.bf16) requires sm_80+; decide the
|
||||
# kernel dispatch at build time via this define rather than at runtime.
|
||||
if cap[0] < 8:
|
||||
flags.append("-DASTRAI_NO_MMA")
|
||||
return flags
|
||||
|
||||
|
||||
_kernels_dir = Path("csrc/kernels")
|
||||
REGISTRY: dict[str, dict] = {}
|
||||
|
||||
CXX_FLAGS = ["-O3", "-funroll-loops"]
|
||||
NVCC_FLAGS = [
|
||||
"-O3",
|
||||
"--expt-relaxed-constexpr",
|
||||
"--use_fast_math",
|
||||
"--ptxas-options=-O3,-v",
|
||||
"--extra-device-vectorization",
|
||||
"--threads=8",
|
||||
]
|
||||
|
||||
|
||||
def register(name: str, sources: list[str] | None = None, **kwargs):
|
||||
if sources is None:
|
||||
sources = [str(_kernels_dir / f"{name}.cu")]
|
||||
REGISTRY[name] = {
|
||||
"sources": sources,
|
||||
"cxx_flags": [*CXX_FLAGS],
|
||||
"nvcc_flags": [*NVCC_FLAGS, *_arch_flags()],
|
||||
"extra_link_args": kwargs.pop("extra_link_args", []),
|
||||
**kwargs,
|
||||
}
|
||||
|
||||
|
||||
register("attn_decode")
|
||||
register("attn_prefill")
|
||||
register("attn_paged_decode")
|
||||
+52
-51
@@ -1,68 +1,69 @@
|
||||
#pragma once
|
||||
|
||||
// Tensor layout for Q/K/V tensors passed to attention kernels.
|
||||
// Internally, kernels always operate on BHLD [batch, n_heads, seq_len, head_dim].
|
||||
// When the caller passes BLHD, dims 1 and 2 are transposed at entry.
|
||||
enum TensorLayout : int {
|
||||
BHLD = 0, // [batch, n_heads, seq_len, head_dim]
|
||||
BLHD = 1, // [batch, seq_len, n_heads, head_dim]
|
||||
};
|
||||
|
||||
|
||||
// Unified attention params covering BOTH addressing modes:
|
||||
// - Contiguous K/V: dense [batch, kv_head, kv_len, head_dim] tensors (k/v).
|
||||
// - Paged (SGLang-style): flat pool [size, kv_head, head_dim] + req_to_token.
|
||||
// Each kernel selects the addressing via a KVSource policy (see
|
||||
// attn_kv_source.cuh); a given call only touches the fields of one mode, so
|
||||
// this is a POD shared by both paths rather than two parallel structs that
|
||||
// drift out of sync.
|
||||
template<typename T, typename AT = float>
|
||||
struct AttentionParams {
|
||||
// Shape
|
||||
int batch;
|
||||
int q_head;
|
||||
int kv_head;
|
||||
int 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;
|
||||
int q_len; // Per-request in contiguous mode; total_q in paged mode.
|
||||
int kv_len; // Contiguous mode; paged mode uses kv_indptr.
|
||||
|
||||
// Attention behavior
|
||||
float scale;
|
||||
|
||||
// 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] (mask_q_stride=0) or 3D [batch, q_len, kv_len]
|
||||
int mask_b_stride; // = kv_len (both 2D and 3D)
|
||||
int mask_q_stride; // 2D: 0 (all q rows share); 3D: kv_len
|
||||
|
||||
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;
|
||||
// -1 = non-causal; >=0 = absolute position of first Q token
|
||||
int causal_offset;
|
||||
float scale;
|
||||
int use_mask;
|
||||
|
||||
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 or 3D)
|
||||
int mask_b_stride;
|
||||
int mask_q_stride;
|
||||
|
||||
const T* __restrict__ q;
|
||||
const T* __restrict__ k_cache;
|
||||
const T* __restrict__ v_cache;
|
||||
// pointers
|
||||
const T* __restrict__ q_ptr;
|
||||
const T* __restrict__ k_ptr;
|
||||
const T* __restrict__ v_ptr;
|
||||
T* __restrict__ o_ptr;
|
||||
const bool* __restrict__ mask;
|
||||
const int64_t* __restrict__ page_table;
|
||||
|
||||
T* __restrict__ o;
|
||||
// strides
|
||||
int q_b_stride;
|
||||
int q_h_stride;
|
||||
int q_l_stride;
|
||||
int q_d_stride;
|
||||
|
||||
int kv_b_stride;
|
||||
int kv_h_stride;
|
||||
int kv_l_stride;
|
||||
int kv_d_stride;
|
||||
|
||||
int mask_b_stride;
|
||||
int mask_h_stride;
|
||||
int mask_l_stride;
|
||||
|
||||
// Paged K/V addressing
|
||||
const int* __restrict__ req_to_token; // [num_reqs, max_context_len]
|
||||
const int* __restrict__ req_pool_indices; // [batch]
|
||||
const int* __restrict__ kv_indptr; // [batch + 1]
|
||||
const int* __restrict__ qo_indptr; // [batch + 1] or nullptr for decode
|
||||
int max_context_len; // req_to_token stride (dim 1)
|
||||
|
||||
// Decode split-KV workspace
|
||||
int num_splits;
|
||||
AT* __restrict__ o_part;
|
||||
AT* __restrict__ ml_part;
|
||||
|
||||
};
|
||||
|
||||
@@ -8,19 +8,43 @@ torch::Tensor attn_decode(
|
||||
c10::optional<torch::Tensor> mask,
|
||||
int64_t causal_offset,
|
||||
double scale,
|
||||
int64_t layout
|
||||
int64_t layout,
|
||||
c10::optional<torch::Tensor> o_part_buf,
|
||||
c10::optional<torch::Tensor> ml_part_buf
|
||||
) {
|
||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||
auto stream = at::cuda::getCurrentCUDAStream();
|
||||
|
||||
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();
|
||||
auto O_view = (layout == BLHD) ? O.transpose(1, 2) : O;
|
||||
p.o_ptr = (bf16*)O_view.data_ptr();
|
||||
|
||||
alloc_split_partials(p);
|
||||
DISPATCH_HEAD_DIM(p.head_dim, dispatch_decode, p);
|
||||
if (o_part_buf.has_value() && ml_part_buf.has_value()
|
||||
&& o_part_buf->defined() && ml_part_buf->defined()) {
|
||||
TORCH_CHECK(o_part_buf->scalar_type() == torch::kFloat32, "o_part_buf must be f32");
|
||||
TORCH_CHECK(ml_part_buf->scalar_type() == torch::kFloat32, "ml_part_buf must be f32");
|
||||
int64_t o_needed = (int64_t)p.batch * p.q_head * MAX_SPLITS * p.head_dim;
|
||||
int64_t ml_needed = (int64_t)p.batch * p.q_head * MAX_SPLITS * 2;
|
||||
TORCH_CHECK(o_part_buf->numel() >= o_needed,
|
||||
"o_part_buf too small: need ", o_needed, " got ", o_part_buf->numel());
|
||||
TORCH_CHECK(ml_part_buf->numel() >= ml_needed,
|
||||
"ml_part_buf too small: need ", ml_needed, " got ", ml_part_buf->numel());
|
||||
TORCH_CHECK(o_part_buf->is_cuda() && ml_part_buf->is_cuda(),
|
||||
"split buffers must be CUDA tensors");
|
||||
TORCH_CHECK(o_part_buf->is_contiguous() && ml_part_buf->is_contiguous(),
|
||||
"split buffers must be contiguous");
|
||||
p.o_part = (float*)o_part_buf->data_ptr();
|
||||
p.ml_part = (float*)ml_part_buf->data_ptr();
|
||||
} else {
|
||||
alloc_split_partials(p);
|
||||
}
|
||||
DISPATCH_HEAD_DIM(p.head_dim, dispatch_decode, p, stream);
|
||||
C10_CUDA_CHECK(cudaGetLastError());
|
||||
return O;
|
||||
}
|
||||
|
||||
@@ -32,6 +56,8 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
py::arg("mask") = py::none(),
|
||||
py::arg("causal_offset") = -1,
|
||||
py::arg("scale") = 0.0,
|
||||
py::arg("layout") = 0,
|
||||
py::arg("layout") = (int64_t)BHLD,
|
||||
py::arg("o_part_buf") = py::none(),
|
||||
py::arg("ml_part_buf") = py::none(),
|
||||
"GQA decode (tensor-core head-packing on sm_80+, scalar fallback)");
|
||||
}
|
||||
|
||||
@@ -2,10 +2,16 @@
|
||||
#include <cuda_bf16.h>
|
||||
#include <float.h>
|
||||
#include "attn_common.h"
|
||||
#include "attn_kv_source.cuh"
|
||||
#include "attn_warp_utils.cuh"
|
||||
constexpr int DC_CHUNK = 64;
|
||||
|
||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||
// Scalar split-KV decode (fallback for sm < 80, no tensor cores), unified
|
||||
// across contiguous and paged (SGLang flat-pool) K/V via the KV template
|
||||
// parameter. For decode the query is the last token, so its valid range
|
||||
// [0, seq_len) IS the causal range; KV::decode_attend_len expresses that
|
||||
// bound per addressing mode (contig clips to causal_offset, paged = seq_len).
|
||||
template <int HEAD_DIM, typename KV, 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;
|
||||
@@ -15,40 +21,46 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
|
||||
int lane = threadIdx.x;
|
||||
int hd_per_thread = p.head_dim / 32;
|
||||
|
||||
const int seq_len = KV::kv_len(p, batch);
|
||||
const KVContext kctx = KV::template make_ctx<HEAD_DIM>(p, batch, kv_head);
|
||||
|
||||
// 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;
|
||||
int q_off = KV::q_decode_base(p, batch, q_head)
|
||||
+ lane * hd_per_thread * p.q_d_stride;
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
q_reg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]);
|
||||
q_reg[i] = __bfloat162float(p.q_ptr[q_off + i * p.q_d_stride]);
|
||||
|
||||
// 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;
|
||||
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[];
|
||||
extern __shared__ __align__(16) bf16 smem[];
|
||||
bf16* k_smem = smem;
|
||||
bf16* v_smem = smem + DC_CHUNK * p.head_dim;
|
||||
|
||||
// Split-KV: each split processes a contiguous subset of chunks
|
||||
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
|
||||
int chunks_total = (seq_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);
|
||||
int this_chunk = min(DC_CHUNK, seq_len - chunk_start);
|
||||
|
||||
// Load K into shared memory (gather from strided global)
|
||||
// Load K and V into shared memory (addressing via KV policy;
|
||||
// paged guards empty slots with zero-fill).
|
||||
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];
|
||||
int kc = chunk_start + s;
|
||||
int token = KV::resolve_token(p, kctx, kc, true);
|
||||
KVAddr a = KV::kv_addr_from_token(p, kctx, token, d_dim);
|
||||
k_smem[i] = a.valid ? *reinterpret_cast<const bf16*>(a.k) : (bf16)0.f;
|
||||
v_smem[i] = a.valid ? *reinterpret_cast<const bf16*>(a.v) : (bf16)0.f;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
@@ -65,20 +77,19 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
|
||||
partial = -FLT_MAX;
|
||||
}
|
||||
if constexpr (IsCausal) {
|
||||
if (kv_idx > p.causal_offset)
|
||||
if (kv_idx >= KV::decode_attend_len(p, batch))
|
||||
partial = -FLT_MAX;
|
||||
}
|
||||
|
||||
float new_m = fmaxf(m, partial);
|
||||
float alpha = expf(m - new_m);
|
||||
float beta = expf(partial - new_m);
|
||||
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);
|
||||
for (int i = 0; i < hd_per_thread; i++) {
|
||||
float vv = __bfloat162float(v_smem[s * p.head_dim + lane * hd_per_thread + i]);
|
||||
acc_reg[i] = fmaf(acc_reg[i], alpha, vv * beta);
|
||||
}
|
||||
m = new_m;
|
||||
}
|
||||
__syncthreads();
|
||||
@@ -98,6 +109,10 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
|
||||
}
|
||||
}
|
||||
|
||||
// Split-combine: merges the per-split partials (o_part/ml_part) into the
|
||||
// final normalised O. KV selects the O addressing (contig batch stride vs
|
||||
// paged row stride).
|
||||
template <typename KV>
|
||||
__global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
|
||||
int bh = blockIdx.x;
|
||||
int d = threadIdx.x;
|
||||
@@ -116,14 +131,14 @@ __global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
|
||||
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);
|
||||
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);
|
||||
int o_off = KV::q_decode_base(p, batch, q_head) + d * p.q_d_stride;
|
||||
p.o_ptr[o_off] = __float2bfloat16(acc * inv);
|
||||
}
|
||||
|
||||
@@ -2,19 +2,22 @@
|
||||
#include <cfloat>
|
||||
#include <cuda_bf16.h>
|
||||
#include "attn_common.h"
|
||||
#include "attn_kv_source.cuh"
|
||||
#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.
|
||||
// Split-K (FlashDecoding) tensor-core decode via GQA head-packing, unified
|
||||
// across contiguous and paged (SGLang flat-pool) K/V via the KV template
|
||||
// parameter. 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.
|
||||
//
|
||||
// KV = ContigKV (dense tensors) or PagedKV (flat pool + req_to_token).
|
||||
// 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>
|
||||
// Traits = KernelTraits<HEAD_DIM, BC=16, WARPS=1, STAGES=2>.
|
||||
template <typename Traits, typename KV, 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;
|
||||
@@ -31,18 +34,21 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
||||
const int G = min(MAX_G, G_total - g_begin);
|
||||
const int q_head0 = kv_head * G_total + g_begin;
|
||||
|
||||
// Per-request seq_len (paged reads kv_indptr; contig uses p.kv_len).
|
||||
const int seq_len = KV::kv_len(p, batch);
|
||||
const KVContext kctx = KV::template make_ctx<Traits::HEAD_DIM>(p, batch, kv_head);
|
||||
|
||||
// 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 q_base = KV::q_decode_base(p, batch, q_head0);
|
||||
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,
|
||||
load_q_mma_frags<Traits::KD>(p.q_ptr + q_base, p.q_h_stride, p.q_d_stride,
|
||||
qra, qrb, va, vb, tid4, Qa);
|
||||
|
||||
float Oacc[Traits::DN8][4];
|
||||
@@ -51,13 +57,12 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
||||
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_total = (seq_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 ----
|
||||
// ---- Load tile lambda: predicated cp.async (addressing via KV policy) ----
|
||||
auto load_tile = [&](int ti, int buf) {
|
||||
int kv0 = ti * Traits::BC;
|
||||
bf16* dK = sK + buf * Traits::BC * Traits::LD;
|
||||
@@ -67,35 +72,27 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
||||
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;
|
||||
bool valid = kc < seq_len;
|
||||
int token = KV::resolve_token(p, kctx, kc, valid);
|
||||
KVAddr a = KV::kv_addr_from_token(p, kctx, token, d);
|
||||
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
|
||||
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_16_pred(&dK[off], a.k, a.valid);
|
||||
cp_async_16_pred(&dV[off], a.v, a.valid);
|
||||
}
|
||||
cp_async_commit();
|
||||
};
|
||||
|
||||
constexpr int BUF_MASK = (Traits::STAGES > 1) ? (Traits::STAGES - 1) : 0;
|
||||
|
||||
// Prologue
|
||||
if (ti_begin < ti_end) {
|
||||
load_tile(ti_begin, 0);
|
||||
}
|
||||
|
||||
for (int ti = ti_begin; ti < ti_end; ti++) {
|
||||
int buf = (ti - ti_begin) & BUF_MASK;
|
||||
|
||||
cp_async_wait_group<0>();
|
||||
__syncwarp();
|
||||
if constexpr (Traits::STAGES > 1) {
|
||||
if (ti + 1 < ti_end)
|
||||
load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK);
|
||||
}
|
||||
// ---- 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 * Traits::BC;
|
||||
int kv0 = (ti_begin + it) * Traits::BC;
|
||||
|
||||
float Sacc[Traits::NC8][4];
|
||||
mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
|
||||
@@ -105,22 +102,45 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
||||
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;
|
||||
// Decode: q_len=1, so qrow0=qrow1=0. Paged treats [0, seq_len) as
|
||||
// the causal range (query is the last token); contig clips to the
|
||||
// causal_offset bound. Dead code eliminated when IsCausal == false.
|
||||
int maxc = IsCausal ? KV::decode_attend_len(p, batch) : seq_len;
|
||||
mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc,
|
||||
0, 0,
|
||||
p.mask_b_stride, 0,
|
||||
batch,
|
||||
p.mask,
|
||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||
p.mask_b_stride, p.mask_h_stride, p.mask_l_stride,
|
||||
batch, q_head0 + gid, q_head0 + gid + 8,
|
||||
p.mask,
|
||||
va, vb,
|
||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||
|
||||
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
|
||||
__syncwarp();
|
||||
};
|
||||
|
||||
if constexpr (Traits::STAGES == 1) {
|
||||
if (ti + 1 < ti_end)
|
||||
load_tile(ti + 1, 0);
|
||||
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++) {
|
||||
if (it + 1 == ntiles)
|
||||
cp_async_wait_group<0>();
|
||||
else
|
||||
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 ----
|
||||
|
||||
+166
-133
@@ -1,195 +1,228 @@
|
||||
#pragma once
|
||||
// Shared attention dispatchers — used by both production .cu and test .cu.
|
||||
// No torch dependency; pure CUDA.
|
||||
//
|
||||
// The paged and contiguous kernels are unified by the KVSource policy
|
||||
// (ContigKV / PagedKV from attn_kv_source.cuh), so each launcher struct
|
||||
// below is templated on KV and the paged dispatch is just the same launcher
|
||||
// instantiated with PagedKV. Only the grid/split math differs, and that is
|
||||
// covered by KV::host_q_len / KV::host_kv_len.
|
||||
|
||||
#include <cuda_runtime.h>
|
||||
#include <algorithm>
|
||||
#include "attn_warp_utils.cuh"
|
||||
#include "attn_kv_source.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.
|
||||
inline int compute_num_splits(int base_blocks, int tiles_total) {
|
||||
int sm_count = 0;
|
||||
cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0);
|
||||
int n = (2 * sm_count + base_blocks - 1) / base_blocks;
|
||||
return std::max(1, std::min(n, std::min(tiles_total, MAX_SPLITS)));
|
||||
// Caps splits so each split processes at least `min_tiles_per_split` tiles,
|
||||
// avoiding excessive loop/prologue overhead when tiles are small.
|
||||
//
|
||||
// Target total grid blocks (`TARGET_BLOCKS`) rather than scaling splits by SM
|
||||
// count. Decode blocks are single-warp (32 threads) and a SM hosts ~11 of
|
||||
// them, so the old `2*sm/base` cap badly undersplit at large batch (B=16 got
|
||||
// 3 splits, optimal ~8). Measured (L20, grid search): bandwidth saturates
|
||||
// near 256-512 total blocks; 512 minimizes worst-case latency across the
|
||||
// B x kv grid; more is pure oversplit overhead.
|
||||
constexpr int DECODE_TARGET_BLOCKS = 512;
|
||||
inline int compute_num_splits(int base_blocks, int tiles_total,
|
||||
int min_tiles_per_split = 1) {
|
||||
int n = (DECODE_TARGET_BLOCKS + 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)));
|
||||
}
|
||||
|
||||
// Dispatch IsCausal × HasMask — eliminates the duplicated 4-way if/else
|
||||
// ladder that appeared in each dispatch_* function. FN must be a function
|
||||
// template <int HEAD_DIM, bool IsCausal, bool HasMask>; HEAD_DIM is forwarded
|
||||
// as the first template argument so callers only spell it once.
|
||||
//
|
||||
// Usage: DISPATCH_CAUSAL_MASK(is_causal, has_mask, launcher<KV>::template launch, HEAD_DIM, p, stream);
|
||||
#define DISPATCH_CAUSAL_MASK(is_causal, has_mask, FN, HEAD_DIM, ...) \
|
||||
do { \
|
||||
if (is_causal) { \
|
||||
if (has_mask) FN<HEAD_DIM, true, true>(__VA_ARGS__); \
|
||||
else FN<HEAD_DIM, true, false>(__VA_ARGS__); \
|
||||
} else { \
|
||||
if (has_mask) FN<HEAD_DIM, false, true>(__VA_ARGS__); \
|
||||
else FN<HEAD_DIM, false, false>(__VA_ARGS__); \
|
||||
} \
|
||||
} while (0)
|
||||
|
||||
// ======================================================================
|
||||
// Prefill
|
||||
// Prefill launchers (KV selects ContigKV or PagedKV addressing)
|
||||
// ======================================================================
|
||||
|
||||
#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);
|
||||
}
|
||||
template <int BC_>
|
||||
struct PrefillKernelConfig {
|
||||
static constexpr int BC = BC_;
|
||||
static constexpr int WARPS = 4;
|
||||
static constexpr int STAGES = 2;
|
||||
};
|
||||
|
||||
// Compile-time configuration map shared by contiguous and paged prefill.
|
||||
// Unsupported head dimensions intentionally have no mapping.
|
||||
template <int HEAD_DIM, bool IsCausal>
|
||||
struct PrefillConfigMap;
|
||||
|
||||
template <> struct PrefillConfigMap<32, false> : PrefillKernelConfig<32> {};
|
||||
template <> struct PrefillConfigMap<32, true> : PrefillKernelConfig<64> {};
|
||||
template <> struct PrefillConfigMap<64, false> : PrefillKernelConfig<32> {};
|
||||
template <> struct PrefillConfigMap<64, true> : PrefillKernelConfig<64> {};
|
||||
template <> struct PrefillConfigMap<128, false> : PrefillKernelConfig<32> {};
|
||||
template <> struct PrefillConfigMap<128, true> : PrefillKernelConfig<32> {};
|
||||
template <> struct PrefillConfigMap<256, false> : PrefillKernelConfig<16> {};
|
||||
template <> struct PrefillConfigMap<256, true> : PrefillKernelConfig<16> {};
|
||||
|
||||
template <typename KV>
|
||||
struct PrefillLauncherMMA {
|
||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||
static void launch(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||
using Config = PrefillConfigMap<HEAD_DIM, IsCausal>;
|
||||
using Traits = KernelTraits<HEAD_DIM, Config::BC, Config::WARPS, Config::STAGES>;
|
||||
constexpr int ROWS = Traits::BR * Config::WARPS;
|
||||
dim3 grid(KV::host_q_blocks(p, ROWS), p.q_head,
|
||||
KV::kPaged ? 1 : p.batch);
|
||||
dim3 block(Traits::NUM_THREADS);
|
||||
attn_prefill_split_q_mma_kernel<Traits, KV, IsCausal, HasMask>
|
||||
<<<grid, block, 0, stream>>>(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 <typename KV>
|
||||
struct PrefillLauncherScalar {
|
||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||
static void launch(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||
constexpr int G = (HEAD_DIM == 32) ? 4 : 8, ROWS = 32, P_BC = 32;
|
||||
dim3 grid(KV::host_q_blocks(p, ROWS), p.q_head,
|
||||
KV::kPaged ? 1 : p.batch);
|
||||
dim3 block(G, ROWS);
|
||||
attn_prefill_split_q_kernel_t<HEAD_DIM, KV, G, ROWS, P_BC, IsCausal, HasMask>
|
||||
<<<grid, block, 0, stream>>>(p);
|
||||
}
|
||||
};
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static inline void dispatch_prefill(AttentionParams<bf16>& p) {
|
||||
static inline void dispatch_prefill(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||
bool is_causal = (p.causal_offset >= 0);
|
||||
bool 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);
|
||||
}
|
||||
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||
PrefillLauncherMMA<ContigKV>::template launch,
|
||||
HEAD_DIM, p, stream);
|
||||
#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);
|
||||
}
|
||||
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||
PrefillLauncherScalar<ContigKV>::template launch,
|
||||
HEAD_DIM, p, stream);
|
||||
#endif
|
||||
}
|
||||
|
||||
// ======================================================================
|
||||
// Decode
|
||||
// ======================================================================
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
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;
|
||||
int tiles_total = (p.kv_len + 32 - 1) / 32;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
|
||||
constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1;
|
||||
using Traits = KernelTraits<HEAD_DIM, 32, 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) {
|
||||
static inline void dispatch_paged_prefill(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||
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);
|
||||
}
|
||||
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||
PrefillLauncherMMA<PagedKV>::template launch,
|
||||
HEAD_DIM, p, stream);
|
||||
#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);
|
||||
}
|
||||
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||
PrefillLauncherScalar<PagedKV>::template launch,
|
||||
HEAD_DIM, p, stream);
|
||||
#endif
|
||||
|
||||
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
|
||||
// ======================================================================
|
||||
// Paged Decode
|
||||
// Decode launchers (KV selects ContigKV or PagedKV addressing)
|
||||
// ======================================================================
|
||||
|
||||
#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;
|
||||
bool page_ok = (p.page_size >= 32);
|
||||
if (G >= 1 && page_ok) {
|
||||
// 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 <typename KV>
|
||||
struct DecodeLauncherMMA {
|
||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||
static void launch(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||
int G = p.q_head / p.kv_head;
|
||||
constexpr int MAX_G = 16;
|
||||
int num_passes = (G + MAX_G - 1) / MAX_G;
|
||||
int tiles_total = (p.kv_len + 32 - 1) / 32;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
|
||||
constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1;
|
||||
using Traits = KernelTraits<HEAD_DIM, 32, 1, STAGES>;
|
||||
constexpr int BC = 16;
|
||||
int kv_len = KV::host_kv_len(p);
|
||||
int tiles_total = (kv_len + BC - 1) / BC;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head * num_passes, 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);
|
||||
} else {
|
||||
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||
size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16);
|
||||
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
|
||||
dim3 block(32, group_size);
|
||||
paged_attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
|
||||
attn_decode_split_kv_mma_kernel<Traits, KV, IsCausal, HasMask>
|
||||
<<<grid, 32, 0, stream>>>(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 <typename KV>
|
||||
struct DecodeLauncherScalar {
|
||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||
static void launch(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||
int kv_len = KV::host_kv_len(p);
|
||||
int chunks_total = (kv_len + DC_CHUNK - 1) / DC_CHUNK;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||
size_t smem = 2 * DC_CHUNK * p.head_dim * sizeof(bf16);
|
||||
int group_size = p.q_head / p.kv_head;
|
||||
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);
|
||||
cudaFuncSetAttribute(
|
||||
attn_decode_split_kv_kernel<HEAD_DIM, KV, IsCausal, HasMask>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
smem);
|
||||
attn_decode_split_kv_kernel<HEAD_DIM, KV, IsCausal, HasMask>
|
||||
<<<grid, block, smem, stream>>>(p);
|
||||
}
|
||||
};
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static inline void dispatch_decode(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||
bool is_causal = (p.causal_offset >= 0);
|
||||
bool has_mask = (p.use_mask && p.mask);
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||
DecodeLauncherMMA<ContigKV>::template launch,
|
||||
HEAD_DIM, p, stream);
|
||||
#else
|
||||
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||
DecodeLauncherScalar<ContigKV>::template launch,
|
||||
HEAD_DIM, p, stream);
|
||||
#endif
|
||||
|
||||
attn_decode_combine_kernel<ContigKV><<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p);
|
||||
}
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static inline void dispatch_paged_decode(PagedAttentionParams<bf16>& p) {
|
||||
static inline void dispatch_paged_decode(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||
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);
|
||||
}
|
||||
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||
DecodeLauncherMMA<PagedKV>::template launch,
|
||||
HEAD_DIM, p, stream);
|
||||
#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);
|
||||
}
|
||||
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||
DecodeLauncherScalar<PagedKV>::template launch,
|
||||
HEAD_DIM, p, stream);
|
||||
#endif
|
||||
|
||||
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
attn_decode_combine_kernel<PagedKV><<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p);
|
||||
}
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
#pragma once
|
||||
#include <float.h>
|
||||
#include <torch/extension.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
#include "attn_common.h"
|
||||
@@ -7,24 +8,28 @@
|
||||
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) \
|
||||
// Usage: DISPATCH_HEAD_DIM(hd, fn, args...)
|
||||
// Expands to: fn<32>(args...); fn<64>(args...); etc.
|
||||
#define DISPATCH_HEAD_DIM(hd, fn, ...) \
|
||||
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; \
|
||||
case 32: fn<32>(__VA_ARGS__); break; \
|
||||
case 64: fn<64>(__VA_ARGS__); break; \
|
||||
case 128: fn<128>(__VA_ARGS__); break; \
|
||||
case 256: fn<256>(__VA_ARGS__); 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({p.batch, p.q_head, MAX_SPLITS, p.head_dim}, fopt);
|
||||
auto ml_part = torch::empty({p.batch, p.q_head, MAX_SPLITS, 2}, fopt);
|
||||
auto 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();
|
||||
}
|
||||
@@ -32,18 +37,21 @@ inline void alloc_split_partials(P& p) {
|
||||
// ---- 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);
|
||||
if (layout == BLHD) 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);
|
||||
p.q_b_stride = (int)q.stride(0);
|
||||
p.q_h_stride = (int)q.stride(1);
|
||||
p.q_l_stride = (int)q.stride(2);
|
||||
p.q_d_stride = (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) {
|
||||
@@ -54,19 +62,27 @@ inline void pack_mask(const c10::optional<torch::Tensor>& mask, P& p) {
|
||||
TORCH_CHECK(m.size(m.dim() - 1) == p.kv_len, "mask kv_len mismatch");
|
||||
if (m.dim() == 2) {
|
||||
p.mask_b_stride = (int)m.stride(0);
|
||||
p.mask_q_stride = 0;
|
||||
p.mask_h_stride = 0;
|
||||
p.mask_l_stride = 0;
|
||||
} else if (m.dim() == 3) {
|
||||
TORCH_CHECK(m.size(1) == p.q_len, "mask q_len mismatch");
|
||||
TORCH_CHECK(m.size(1) == 1 || m.size(1) == p.q_len, "mask q_len mismatch");
|
||||
p.mask_b_stride = (int)m.stride(0);
|
||||
p.mask_q_stride = (int)m.stride(1);
|
||||
p.mask_h_stride = 0;
|
||||
p.mask_l_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_l_stride = (m.size(2) == 1) ? 0 : (int)m.stride(2);
|
||||
} else {
|
||||
TORCH_CHECK(false, "mask must be 2D [batch, kv_len] or 3D [batch, q_len, kv_len]");
|
||||
TORCH_CHECK(false, "mask must be 2D, 3D, or 4D");
|
||||
}
|
||||
p.mask = m.data_ptr<bool>();
|
||||
} else {
|
||||
p.mask = nullptr;
|
||||
p.mask_b_stride = 0;
|
||||
p.mask_q_stride = 0;
|
||||
p.mask_h_stride = 0;
|
||||
p.mask_l_stride = 0;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -90,82 +106,206 @@ inline void attn_pack_params(
|
||||
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);
|
||||
if (layout == BLHD) 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(p.q_head % p.kv_head == 0,
|
||||
"q_head must be divisible by kv_head");
|
||||
TORCH_CHECK(k.size(3) == p.head_dim, "K/V head_dim must match Q");
|
||||
TORCH_CHECK(q.stride(3) == 1 && k.stride(3) == 1 && v.stride(3) == 1,
|
||||
"Q/K/V head_dim must be contiguous");
|
||||
|
||||
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.kv_b_stride = (int)k.stride(0);
|
||||
p.kv_h_stride = (int)k.stride(1);
|
||||
p.kv_l_stride = (int)k.stride(2);
|
||||
p.kv_d_stride = (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.q_ptr = (const T*)q.data_ptr();
|
||||
p.k_ptr = (const T*)k.data_ptr();
|
||||
p.v_ptr = (const T*)v.data_ptr();
|
||||
p.o_ptr = nullptr;
|
||||
p.o_part = nullptr;
|
||||
p.ml_part = nullptr;
|
||||
|
||||
pack_mask(mask, p);
|
||||
}
|
||||
|
||||
// ---- attn_pack_paged_params ----
|
||||
// ---- attn_pack_paged_decode_params ----
|
||||
// SGLang-style: flat KV pool + req_to_token indexing + variable
|
||||
// seq_lens via kv_indptr. Q is [batch, q_head, head_dim] (q_len=1 per req).
|
||||
template<typename T>
|
||||
inline void attn_pack_paged_params(
|
||||
inline void attn_pack_paged_decode_params(
|
||||
torch::Tensor q,
|
||||
torch::Tensor page_table,
|
||||
torch::Tensor k_cache,
|
||||
torch::Tensor v_cache,
|
||||
int64_t page_size,
|
||||
int64_t kv_len,
|
||||
torch::Tensor req_to_token,
|
||||
torch::Tensor req_pool_indices,
|
||||
torch::Tensor kv_indptr,
|
||||
c10::optional<torch::Tensor> mask,
|
||||
int64_t causal_offset,
|
||||
double scale,
|
||||
int64_t layout,
|
||||
PagedAttentionParams<T>& p
|
||||
AttentionParams<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.is_cuda() && k_cache.is_cuda() && v_cache.is_cuda());
|
||||
TORCH_CHECK(req_to_token.is_cuda() && req_pool_indices.is_cuda() && kv_indptr.is_cuda());
|
||||
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16");
|
||||
TORCH_CHECK(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");
|
||||
TORCH_CHECK(req_to_token.dtype() == torch::kInt32, "req_to_token must be int32");
|
||||
TORCH_CHECK(req_pool_indices.dtype() == torch::kInt32,
|
||||
"req_pool_indices must be int32");
|
||||
TORCH_CHECK(kv_indptr.dtype() == torch::kInt32, "kv_indptr must be int32");
|
||||
TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must match");
|
||||
TORCH_CHECK(k_cache.dim() == 3, "k_cache must be 3D [size, kv_head, head_dim]");
|
||||
TORCH_CHECK(q.dim() == 3, "q must be 3D [batch, q_head, head_dim]");
|
||||
|
||||
extract_q_dims_and_strides(q, layout, p);
|
||||
|
||||
p.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)");
|
||||
p.batch = (int)q.size(0);
|
||||
p.q_head = (int)q.size(1);
|
||||
p.head_dim = (int)q.size(2);
|
||||
p.kv_head = (int)k_cache.size(1);
|
||||
TORCH_CHECK(k_cache.size(2) == p.head_dim, "k_cache head_dim mismatch");
|
||||
TORCH_CHECK(q.stride(2) == 1 && k_cache.stride(2) == 1 && v_cache.stride(2) == 1,
|
||||
"Q/K/V head_dim must be contiguous");
|
||||
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);
|
||||
TORCH_CHECK(p.q_head % p.kv_head == 0, "q_head must be divisible by kv_head");
|
||||
|
||||
p.q_l_stride = (int)q.stride(0);
|
||||
p.q_h_stride = (int)q.stride(1);
|
||||
p.q_d_stride = (int)q.stride(2);
|
||||
|
||||
p.k_ptr = (const T*)k_cache.data_ptr();
|
||||
p.v_ptr = (const T*)v_cache.data_ptr();
|
||||
p.q_ptr = (const T*)q.data_ptr();
|
||||
p.req_to_token = req_to_token.data_ptr<int>();
|
||||
p.req_pool_indices = req_pool_indices.data_ptr<int>();
|
||||
p.kv_indptr = kv_indptr.data_ptr<int>();
|
||||
p.qo_indptr = nullptr;
|
||||
p.max_context_len = (int)req_to_token.size(1);
|
||||
|
||||
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;
|
||||
if (p.use_mask) {
|
||||
auto m = mask.value();
|
||||
TORCH_CHECK(m.is_cuda() && m.dtype() == torch::kBool, "mask must be bool CUDA");
|
||||
TORCH_CHECK(m.size(0) == p.batch, "mask batch mismatch");
|
||||
p.mask_b_stride = (int)m.stride(0);
|
||||
p.mask_h_stride = 0;
|
||||
p.mask_l_stride = 0;
|
||||
p.mask = m.data_ptr<bool>();
|
||||
} else {
|
||||
p.mask = nullptr;
|
||||
p.mask_b_stride = 0;
|
||||
p.mask_h_stride = 0;
|
||||
p.mask_l_stride = 0;
|
||||
}
|
||||
|
||||
p.o_ptr = nullptr;
|
||||
p.o_part = nullptr;
|
||||
p.ml_part = nullptr;
|
||||
}
|
||||
|
||||
// ---- attn_pack_paged_prefill_params ----
|
||||
// SGLang-style: flat KV pool + req_to_token + ragged batch via qo_indptr.
|
||||
// Q is [total_q, q_head, head_dim] (flattened across all requests).
|
||||
template<typename T>
|
||||
inline void attn_pack_paged_prefill_params(
|
||||
torch::Tensor q,
|
||||
torch::Tensor k_cache,
|
||||
torch::Tensor v_cache,
|
||||
torch::Tensor req_to_token,
|
||||
torch::Tensor req_pool_indices,
|
||||
torch::Tensor kv_indptr,
|
||||
torch::Tensor qo_indptr,
|
||||
c10::optional<torch::Tensor> mask,
|
||||
int64_t causal_offset,
|
||||
double scale,
|
||||
AttentionParams<T>& p
|
||||
) {
|
||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||
|
||||
TORCH_CHECK(q.is_cuda() && k_cache.is_cuda() && v_cache.is_cuda());
|
||||
TORCH_CHECK(req_to_token.is_cuda() && req_pool_indices.is_cuda());
|
||||
TORCH_CHECK(kv_indptr.is_cuda() && qo_indptr.is_cuda());
|
||||
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16");
|
||||
TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bf16");
|
||||
TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bf16");
|
||||
TORCH_CHECK(req_to_token.dtype() == torch::kInt32, "req_to_token must be int32");
|
||||
TORCH_CHECK(req_pool_indices.dtype() == torch::kInt32,
|
||||
"req_pool_indices must be int32");
|
||||
TORCH_CHECK(kv_indptr.dtype() == torch::kInt32, "kv_indptr must be int32");
|
||||
TORCH_CHECK(qo_indptr.dtype() == torch::kInt32, "qo_indptr must be int32");
|
||||
TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must match");
|
||||
TORCH_CHECK(k_cache.dim() == 3, "k_cache must be 3D [size, kv_head, head_dim]");
|
||||
TORCH_CHECK(q.dim() == 3, "q must be 3D [total_q, q_head, head_dim]");
|
||||
|
||||
p.q_head = (int)q.size(1);
|
||||
p.head_dim = (int)q.size(2);
|
||||
p.q_len = (int)q.size(0);
|
||||
p.kv_head = (int)k_cache.size(1);
|
||||
p.batch = (int)req_pool_indices.size(0);
|
||||
TORCH_CHECK(k_cache.size(2) == p.head_dim, "k_cache head_dim mismatch");
|
||||
TORCH_CHECK(q.stride(2) == 1 && k_cache.stride(2) == 1 && v_cache.stride(2) == 1,
|
||||
"Q/K/V head_dim must be contiguous");
|
||||
TORCH_CHECK(p.head_dim % 16 == 0, "head_dim must be multiple of 16");
|
||||
TORCH_CHECK(p.q_head % p.kv_head == 0, "q_head must be divisible by kv_head");
|
||||
TORCH_CHECK(kv_indptr.size(0) == p.batch + 1, "kv_indptr must be [batch+1]");
|
||||
TORCH_CHECK(qo_indptr.size(0) == p.batch + 1, "qo_indptr must be [batch+1]");
|
||||
|
||||
p.q_l_stride = (int)q.stride(0);
|
||||
p.q_h_stride = (int)q.stride(1);
|
||||
p.q_d_stride = (int)q.stride(2);
|
||||
|
||||
p.k_ptr = (const T*)k_cache.data_ptr();
|
||||
p.v_ptr = (const T*)v_cache.data_ptr();
|
||||
p.q_ptr = (const T*)q.data_ptr();
|
||||
p.req_to_token = req_to_token.data_ptr<int>();
|
||||
p.req_pool_indices = req_pool_indices.data_ptr<int>();
|
||||
p.kv_indptr = kv_indptr.data_ptr<int>();
|
||||
p.qo_indptr = qo_indptr.data_ptr<int>();
|
||||
p.max_context_len = (int)req_to_token.size(1);
|
||||
|
||||
p.causal_offset = (int)causal_offset;
|
||||
p.use_mask = (mask.has_value() && mask.value().defined()) ? 1 : 0;
|
||||
if (p.use_mask) {
|
||||
auto m = mask.value();
|
||||
TORCH_CHECK(m.is_cuda() && m.dtype() == torch::kBool, "mask must be bool CUDA");
|
||||
TORCH_CHECK(m.size(0) == p.batch, "mask batch mismatch");
|
||||
if (m.dim() == 2) {
|
||||
TORCH_CHECK(m.size(1) <= p.max_context_len, "mask kv_len mismatch");
|
||||
p.mask_b_stride = (int)m.stride(0);
|
||||
p.mask_h_stride = 0;
|
||||
p.mask_l_stride = 0;
|
||||
} else if (m.dim() == 4) {
|
||||
TORCH_CHECK(m.size(1) == 1 || m.size(1) == p.q_head, "mask head mismatch");
|
||||
TORCH_CHECK(m.size(2) > 0 && m.size(2) <= p.q_len, "mask q_len mismatch");
|
||||
TORCH_CHECK(m.size(3) <= p.max_context_len, "mask kv_len mismatch");
|
||||
p.mask_b_stride = (int)m.stride(0);
|
||||
p.mask_h_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1);
|
||||
p.mask_l_stride = (m.size(2) == 1) ? 0 : (int)m.stride(2);
|
||||
} else {
|
||||
TORCH_CHECK(false, "mask must be 2D or 4D");
|
||||
}
|
||||
p.mask = m.data_ptr<bool>();
|
||||
} else {
|
||||
p.mask = nullptr;
|
||||
p.mask_b_stride = 0;
|
||||
p.mask_h_stride = 0;
|
||||
p.mask_l_stride = 0;
|
||||
}
|
||||
p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
|
||||
|
||||
p.o_ptr = nullptr;
|
||||
p.o_part = nullptr;
|
||||
p.ml_part = nullptr;
|
||||
|
||||
pack_mask(mask, p);
|
||||
}
|
||||
|
||||
@@ -0,0 +1,224 @@
|
||||
#pragma once
|
||||
#include <cuda_bf16.h>
|
||||
#include "attn_common.h"
|
||||
|
||||
// ============================================================================
|
||||
// KVSource policies — the single dimension along which the paged and
|
||||
// non-paged attention kernels differ. Each kernel is templated on one of
|
||||
// these (ContigKV / PagedKV) and stays fully generic: the policy owns every
|
||||
// place where "where does K/V live" and "what is this request's seq_len"
|
||||
// are answered. All methods are __host__ __device__ so the same policy
|
||||
// serves both the device kernels (addressing, seq_len) and the host-side
|
||||
// launchers (grid / split computation).
|
||||
//
|
||||
// ContigKV: K/V are dense [batch, kv_head, kv_len, head_dim] tensors.
|
||||
// Params fields used: k, v, kv_stride_*, kv_len, q_len,
|
||||
// q_b_stride, causal_offset.
|
||||
// PagedKV: K/V live in a flat pool [size, kv_head, head_dim] indexed via
|
||||
// req_to_token. Params fields used: k_cache, v_cache,
|
||||
// req_to_token, req_pool_indices, kv_indptr, qo_indptr,
|
||||
// max_context_len, q_l_stride.
|
||||
//
|
||||
// Addressing state that is constant across a whole kernel invocation for one
|
||||
// (batch, kv_head) pair is captured once by make_ctx<HEAD_DIM>() and passed
|
||||
// to kv_addr, so the load loops never redo the hoistable base computation
|
||||
// (e.g. the req_pool_indices global read) element-by-element.
|
||||
// ============================================================================
|
||||
|
||||
// Every policy method is static + callable from both host and device code.
|
||||
#define HOST_DEV_FORCEINLINE static __host__ __device__ __forceinline__
|
||||
|
||||
using bf16 = __nv_bfloat16;
|
||||
|
||||
// Hoisted per-(batch, kv_head) addressing context.
|
||||
struct KVContext {
|
||||
int kv_base; // contig: batch*kv_b_stride + kv_head*kv_h_stride
|
||||
int req_idx; // paged: req_pool_indices[batch]
|
||||
int64_t rtt_stride; // paged: max_context_len
|
||||
int64_t pool_stride; // paged: kv_head * HEAD_DIM
|
||||
int64_t head_off; // paged: kv_head * HEAD_DIM
|
||||
};
|
||||
|
||||
// Per-element K/V global addresses for one (kc, d) position of a K/V tile.
|
||||
// The pointers are ALWAYS the computed addresses (never nullptr) — callers
|
||||
// gate on `valid` (cp.async src_size=0, or a guarded scalar deref). `valid`
|
||||
// starts as "within the request's seq_len"; the paged policy further degrades
|
||||
// it when req_to_token maps the position to a negative slot (empty padding).
|
||||
// This matches the original hand-rolled load loops, where the address was
|
||||
// always formed and the predicate decided whether anything was read.
|
||||
struct KVAddr {
|
||||
const void* k;
|
||||
const void* v;
|
||||
bool valid;
|
||||
};
|
||||
|
||||
// ---- Contiguous K/V ----
|
||||
struct ContigKV {
|
||||
static constexpr bool kPaged = false;
|
||||
|
||||
// host-side length hooks (grid + split computation in the launchers)
|
||||
HOST_DEV_FORCEINLINE int host_q_blocks(const AttentionParams<bf16>& p, int rows) {
|
||||
return (p.q_len + rows - 1) / rows;
|
||||
}
|
||||
template <int ROWS>
|
||||
HOST_DEV_FORCEINLINE bool map_q_tile(const AttentionParams<bf16>&,
|
||||
int flat_tile, int grid_batch,
|
||||
int& batch, int& q_tile) {
|
||||
batch = grid_batch;
|
||||
q_tile = flat_tile;
|
||||
return true;
|
||||
}
|
||||
HOST_DEV_FORCEINLINE int host_kv_len(const AttentionParams<bf16>& p) {
|
||||
return p.kv_len;
|
||||
}
|
||||
|
||||
// prefill: element offset of the request's Q rows (kernel adds qrow*q_l_stride)
|
||||
HOST_DEV_FORCEINLINE int q_base(
|
||||
const AttentionParams<bf16>& p, int batch, int q_head) {
|
||||
return batch * p.q_b_stride + q_head * p.q_h_stride;
|
||||
}
|
||||
// decode: same offset (q_len == 1, so there is no row stride component)
|
||||
HOST_DEV_FORCEINLINE int q_decode_base(
|
||||
const AttentionParams<bf16>& p, int batch, int q_head) {
|
||||
return batch * p.q_b_stride + q_head * p.q_h_stride;
|
||||
}
|
||||
|
||||
HOST_DEV_FORCEINLINE int kv_len(const AttentionParams<bf16>& p, int batch) {
|
||||
return p.kv_len;
|
||||
}
|
||||
HOST_DEV_FORCEINLINE int q_len(const AttentionParams<bf16>& p, int batch) {
|
||||
return p.q_len;
|
||||
}
|
||||
HOST_DEV_FORCEINLINE int causal_offset(const AttentionParams<bf16>& p, int batch) {
|
||||
return p.causal_offset;
|
||||
}
|
||||
// decode: exclusive bound of the single query's attend range
|
||||
HOST_DEV_FORCEINLINE int decode_attend_len(const AttentionParams<bf16>& p, int batch) {
|
||||
return (p.kv_len < p.causal_offset + 1) ? p.kv_len : (p.causal_offset + 1);
|
||||
}
|
||||
|
||||
template <int HEAD_DIM>
|
||||
HOST_DEV_FORCEINLINE KVContext make_ctx(
|
||||
const AttentionParams<bf16>& p, int batch, int kv_head) {
|
||||
KVContext c = {};
|
||||
c.kv_base = batch * p.kv_b_stride + kv_head * p.kv_h_stride;
|
||||
return c;
|
||||
}
|
||||
HOST_DEV_FORCEINLINE int resolve_token(
|
||||
const AttentionParams<bf16>& p, const KVContext& c, int kc, bool valid) {
|
||||
return valid ? kc : -1;
|
||||
}
|
||||
HOST_DEV_FORCEINLINE KVAddr kv_addr_from_token(
|
||||
const AttentionParams<bf16>& p, const KVContext& c, int token, int d) {
|
||||
const bool valid = token >= 0;
|
||||
const int safe_token = valid ? token : 0;
|
||||
const int64_t gmem_off = (int64_t)c.kv_base
|
||||
+ (int64_t)safe_token * p.kv_l_stride
|
||||
+ (int64_t)d * p.kv_d_stride;
|
||||
return {&p.k_ptr[gmem_off], &p.v_ptr[gmem_off], valid};
|
||||
}
|
||||
};
|
||||
|
||||
// ---- Paged (SGLang-style flat pool) K/V ----
|
||||
struct PagedKV {
|
||||
static constexpr bool kPaged = true;
|
||||
|
||||
HOST_DEV_FORCEINLINE int host_q_blocks(const AttentionParams<bf16>& p, int rows) {
|
||||
// sum(ceil(q_len[b] / rows)) <= ceil(total_q / rows) + batch - 1.
|
||||
return (p.q_len + rows - 1) / rows + p.batch - 1;
|
||||
}
|
||||
template <int ROWS>
|
||||
HOST_DEV_FORCEINLINE bool map_q_tile(const AttentionParams<bf16>& p,
|
||||
int flat_tile, int,
|
||||
int& batch, int& q_tile) {
|
||||
int tile_base = 0;
|
||||
for (int b = 0; b < p.batch; ++b) {
|
||||
int len = p.qo_indptr[b + 1] - p.qo_indptr[b];
|
||||
int tiles = (len + ROWS - 1) / ROWS;
|
||||
if (flat_tile < tile_base + tiles) {
|
||||
batch = b;
|
||||
q_tile = flat_tile - tile_base;
|
||||
return true;
|
||||
}
|
||||
tile_base += tiles;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
HOST_DEV_FORCEINLINE int host_kv_len(const AttentionParams<bf16>& p) {
|
||||
return p.max_context_len;
|
||||
}
|
||||
|
||||
// prefill: Q rows start at qo_indptr[batch] (ragged batch base)
|
||||
HOST_DEV_FORCEINLINE int q_base(
|
||||
const AttentionParams<bf16>& p, int batch, int q_head) {
|
||||
return p.qo_indptr[batch] * p.q_l_stride + q_head * p.q_h_stride;
|
||||
}
|
||||
// decode: Q is [batch, q_head, head_dim], so batch is the outer row
|
||||
HOST_DEV_FORCEINLINE int q_decode_base(
|
||||
const AttentionParams<bf16>& p, int batch, int q_head) {
|
||||
return batch * p.q_l_stride + q_head * p.q_h_stride;
|
||||
}
|
||||
|
||||
HOST_DEV_FORCEINLINE int kv_len(const AttentionParams<bf16>& p, int batch) {
|
||||
return p.kv_indptr[batch + 1] - p.kv_indptr[batch];
|
||||
}
|
||||
HOST_DEV_FORCEINLINE int q_len(const AttentionParams<bf16>& p, int batch) {
|
||||
return p.qo_indptr[batch + 1] - p.qo_indptr[batch];
|
||||
}
|
||||
HOST_DEV_FORCEINLINE int causal_offset(const AttentionParams<bf16>& p, int batch) {
|
||||
return kv_len(p, batch) - q_len(p, batch);
|
||||
}
|
||||
// decode: the query is the last token, so [0, seq_len) IS its causal range
|
||||
HOST_DEV_FORCEINLINE int decode_attend_len(const AttentionParams<bf16>& p, int batch) {
|
||||
return kv_len(p, batch);
|
||||
}
|
||||
|
||||
template <int HEAD_DIM>
|
||||
HOST_DEV_FORCEINLINE KVContext make_ctx(
|
||||
const AttentionParams<bf16>& p, int batch, int kv_head) {
|
||||
KVContext c = {};
|
||||
c.req_idx = p.req_pool_indices[batch];
|
||||
c.rtt_stride = (int64_t)p.max_context_len;
|
||||
c.pool_stride = (int64_t)p.kv_head * HEAD_DIM;
|
||||
c.head_off = (int64_t)kv_head * HEAD_DIM;
|
||||
return c;
|
||||
}
|
||||
HOST_DEV_FORCEINLINE int resolve_token(
|
||||
const AttentionParams<bf16>& p, const KVContext& c, int kc, bool valid) {
|
||||
return valid ? p.req_to_token[c.req_idx * c.rtt_stride + kc] : -1;
|
||||
}
|
||||
HOST_DEV_FORCEINLINE KVAddr kv_addr_from_token(
|
||||
const AttentionParams<bf16>& p, const KVContext& c, int slot, int d) {
|
||||
const bool valid = slot >= 0;
|
||||
const int safe_slot = valid ? slot : 0;
|
||||
const int64_t gmem_off = (int64_t)safe_slot * c.pool_stride + c.head_off + d;
|
||||
return {&p.k_ptr[gmem_off], &p.v_ptr[gmem_off], valid};
|
||||
}
|
||||
};
|
||||
|
||||
// ---- Q-block mapping ----
|
||||
// Contiguous grids map directly to (batch, q_tile). Paged grids flatten the
|
||||
// ragged Q tiles, so one thread resolves the request and broadcasts it.
|
||||
template <int ROWS, typename KV>
|
||||
__device__ __forceinline__ bool map_q_block(
|
||||
const AttentionParams<bf16>& p, int& batch, int& q_tile) {
|
||||
if constexpr (!KV::kPaged) {
|
||||
batch = blockIdx.z;
|
||||
q_tile = blockIdx.x;
|
||||
return true;
|
||||
} else {
|
||||
__shared__ int mapped_batch;
|
||||
__shared__ int mapped_q_tile;
|
||||
|
||||
if ((threadIdx.x | threadIdx.y) == 0) {
|
||||
mapped_batch = -1;
|
||||
KV::template map_q_tile<ROWS>(
|
||||
p, blockIdx.x, blockIdx.z, mapped_batch, mapped_q_tile);
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
batch = mapped_batch;
|
||||
q_tile = mapped_q_tile;
|
||||
return batch >= 0;
|
||||
}
|
||||
}
|
||||
@@ -3,6 +3,12 @@
|
||||
#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.
|
||||
//
|
||||
@@ -93,22 +99,22 @@ __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.
|
||||
// BypassL1 defaults to .cg (L2 only); false selects .ca (L1 + L2).
|
||||
// src_size=0 means no bytes are read, so an out-of-bounds address is safe.
|
||||
template <bool BypassL1 = true>
|
||||
__device__ __forceinline__ void cp_async_16_pred(bf16* smem_ptr,
|
||||
const void* gmem_ptr,
|
||||
bool pred) {
|
||||
unsigned smem_addr = __cvta_generic_to_shared(smem_ptr);
|
||||
int src_size = pred ? 16 : 0;
|
||||
asm volatile("cp.async.ca.shared.global [%0], [%1], 16, %2;"
|
||||
:: "r"(smem_addr), "l"(gmem_ptr), "r"(src_size));
|
||||
if constexpr (BypassL1) {
|
||||
asm volatile("cp.async.cg.shared.global [%0], [%1], 16, %2;"
|
||||
:: "r"(smem_addr), "l"(gmem_ptr), "r"(src_size));
|
||||
} else {
|
||||
asm volatile("cp.async.ca.shared.global [%0], [%1], 16, %2;"
|
||||
:: "r"(smem_addr), "l"(gmem_ptr), "r"(src_size));
|
||||
}
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void cp_async_commit() {
|
||||
@@ -127,8 +133,8 @@ __device__ __forceinline__ void cp_async_wait_group() {
|
||||
// ---------------------------------------------------------------------------
|
||||
// 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).
|
||||
// stride_row is p.q_h_stride for decode (q_len=1, G heads) or
|
||||
// p.q_l_stride for prefill (multi-q rows).
|
||||
// ---------------------------------------------------------------------------
|
||||
template <int KD>
|
||||
__device__ inline void load_q_mma_frags(
|
||||
@@ -192,9 +198,10 @@ __device__ inline void mma_softmax_tile(
|
||||
int kv0,
|
||||
int maxc0, int maxc1,
|
||||
int qrow0, int qrow1,
|
||||
int mask_b_stride, int mask_q_stride,
|
||||
int mask_batch,
|
||||
int mask_b_stride, int mask_h_stride, int mask_l_stride,
|
||||
int mask_batch, int mask_head0, int mask_head1,
|
||||
const bool* __restrict__ mask,
|
||||
bool valid0, bool valid1,
|
||||
float Sacc[Traits::NC8][4],
|
||||
float Oacc[Traits::DN8][4],
|
||||
float& m0, float& m1,
|
||||
@@ -204,16 +211,16 @@ __device__ inline void mma_softmax_tile(
|
||||
int tid4 = lane & 3;
|
||||
|
||||
float rmax0 = -FLT_MAX, rmax1 = -FLT_MAX;
|
||||
int mask_base0 = mask_batch * mask_b_stride + qrow0 * mask_q_stride;
|
||||
int mask_base1 = mask_batch * mask_b_stride + qrow1 * mask_q_stride;
|
||||
int mask_base0 = mask_batch * mask_b_stride + mask_head0 * mask_h_stride + qrow0 * mask_l_stride;
|
||||
int mask_base1 = mask_batch * mask_b_stride + mask_head1 * mask_h_stride + qrow1 * mask_l_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]);
|
||||
bool b0 = !valid0 || (cc >= maxc0) || (HasMask && !mask[mask_base0 + cc]);
|
||||
bool b1 = !valid0 || (c1 >= maxc0) || (HasMask && !mask[mask_base0 + c1]);
|
||||
bool b2 = !valid1 || (cc >= maxc1) || (HasMask && !mask[mask_base1 + cc]);
|
||||
bool b3 = !valid1 || (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];
|
||||
|
||||
@@ -3,40 +3,79 @@
|
||||
|
||||
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,
|
||||
torch::Tensor req_to_token,
|
||||
torch::Tensor req_pool_indices,
|
||||
torch::Tensor kv_indptr,
|
||||
c10::optional<torch::Tensor> mask,
|
||||
int64_t causal_offset,
|
||||
double scale,
|
||||
int64_t layout
|
||||
c10::optional<torch::Tensor> o_part_buf,
|
||||
c10::optional<torch::Tensor> ml_part_buf,
|
||||
c10::optional<torch::Tensor> out_buf
|
||||
) {
|
||||
PagedAttentionParams<bf16> p;
|
||||
attn_pack_paged_params(q, page_table, k_cache, v_cache,
|
||||
page_size, kv_len, mask, causal_offset, scale, layout, p);
|
||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||
auto stream = at::cuda::getCurrentCUDAStream();
|
||||
|
||||
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();
|
||||
AttentionParams<bf16> p;
|
||||
attn_pack_paged_decode_params(q, k_cache, v_cache,
|
||||
req_to_token, req_pool_indices, kv_indptr,
|
||||
mask, causal_offset, scale, p);
|
||||
|
||||
alloc_split_partials(p);
|
||||
DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_decode, p);
|
||||
torch::Tensor O;
|
||||
if (out_buf.has_value() && out_buf->defined()) {
|
||||
TORCH_CHECK(out_buf->dtype() == q.dtype(), "out_buf dtype must match q");
|
||||
TORCH_CHECK(out_buf->is_cuda() && out_buf->is_contiguous(),
|
||||
"out_buf must be a contiguous CUDA tensor");
|
||||
TORCH_CHECK(out_buf->size(0) >= q.size(0), "out_buf batch too small");
|
||||
TORCH_CHECK(out_buf->size(1) == q.size(1), "out_buf heads must match q");
|
||||
TORCH_CHECK(out_buf->size(2) == q.size(2), "out_buf head_dim must match q");
|
||||
TORCH_CHECK(q.is_contiguous(),
|
||||
"q must be contiguous when out_buf is provided");
|
||||
O = out_buf.value().slice(0, 0, q.size(0));
|
||||
} else {
|
||||
O = torch::empty({q.size(0), q.size(1), q.size(2)}, q.options());
|
||||
}
|
||||
p.o_ptr = (bf16*)O.data_ptr();
|
||||
|
||||
if (o_part_buf.has_value() && ml_part_buf.has_value()
|
||||
&& o_part_buf->defined() && ml_part_buf->defined()) {
|
||||
TORCH_CHECK(o_part_buf->scalar_type() == torch::kFloat32, "o_part_buf must be f32");
|
||||
TORCH_CHECK(ml_part_buf->scalar_type() == torch::kFloat32, "ml_part_buf must be f32");
|
||||
int64_t o_needed = (int64_t)p.batch * p.q_head * MAX_SPLITS * p.head_dim;
|
||||
int64_t ml_needed = (int64_t)p.batch * p.q_head * MAX_SPLITS * 2;
|
||||
TORCH_CHECK(o_part_buf->numel() >= o_needed,
|
||||
"o_part_buf too small: need ", o_needed, " got ", o_part_buf->numel());
|
||||
TORCH_CHECK(ml_part_buf->numel() >= ml_needed,
|
||||
"ml_part_buf too small: need ", ml_needed, " got ", ml_part_buf->numel());
|
||||
TORCH_CHECK(o_part_buf->is_cuda() && ml_part_buf->is_cuda(),
|
||||
"split buffers must be CUDA tensors");
|
||||
TORCH_CHECK(o_part_buf->is_contiguous() && ml_part_buf->is_contiguous(),
|
||||
"split buffers must be contiguous");
|
||||
p.o_part = (float*)o_part_buf->data_ptr();
|
||||
p.ml_part = (float*)ml_part_buf->data_ptr();
|
||||
} else {
|
||||
alloc_split_partials(p);
|
||||
}
|
||||
DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_decode, p, stream);
|
||||
C10_CUDA_CHECK(cudaGetLastError());
|
||||
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("req_to_token"),
|
||||
py::arg("req_pool_indices"),
|
||||
py::arg("kv_indptr"),
|
||||
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.");
|
||||
py::arg("o_part_buf") = py::none(),
|
||||
py::arg("ml_part_buf") = py::none(),
|
||||
py::arg("out_buf") = py::none(),
|
||||
"SGLang-style paged decode: flat KV pool + req_to_token + kv_indptr.");
|
||||
}
|
||||
|
||||
@@ -1,146 +0,0 @@
|
||||
#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;
|
||||
|
||||
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;
|
||||
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 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 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);
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user