Compare commits
166
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a1a1a6bf0f | ||
|
|
7dd184a4e5 | ||
|
|
bf239d194c | ||
|
|
8a353117ea | ||
|
|
04a8e2517a | ||
|
|
c4f7f82725 | ||
|
|
fac9d07542 | ||
|
|
bbb2d95256 | ||
|
|
1c04a0b9fa | ||
|
|
ba8beb81be | ||
|
|
7cfcc6c86a | ||
|
|
f4c44ebf1c | ||
|
|
f86f605f5f | ||
|
|
4c82d5d84b | ||
|
|
76aa4edc9f | ||
|
|
6354dbe8bc | ||
|
|
f45230fb2c | ||
|
|
d4534be8ca | ||
|
|
1d57588d27 | ||
|
|
a92bf79295 | ||
|
|
f7d96455a5 | ||
|
|
8cfe7536ea | ||
|
|
a8b63fa362 | ||
|
|
4d6a244093 | ||
|
|
01eacbde51 | ||
|
|
057c0d33df | ||
|
|
3e57cc8069 | ||
|
|
2eeac02d70 | ||
|
|
5e76fbd1bf | ||
|
|
4dc5e923e0 | ||
|
|
998b443aa3 | ||
|
|
cebdd45d3a | ||
|
|
7da1439c9e | ||
|
|
29e5f571af | ||
|
|
74e694921c | ||
|
|
d5067af064 | ||
|
|
f6db546578 | ||
|
|
31ca357c61 | ||
|
|
34471252ab | ||
|
|
aa08479285 | ||
|
|
4b10d3ca37 | ||
|
|
2bc4d2b8a8 | ||
|
|
4244df2785 | ||
|
|
a29bdfae46 | ||
|
|
10fec8dca1 | ||
|
|
75304d084d | ||
|
|
16a55bb474 | ||
|
|
cb21af38ba | ||
|
|
dcc96de12a | ||
|
|
7d27f3e078 | ||
|
|
84753d3e08 | ||
|
|
53a7149577 | ||
|
|
c79d34eee1 | ||
|
|
398e8a3ea3 | ||
|
|
f252af495c | ||
|
|
00c2c80c8f | ||
|
|
a6c6a54ace | ||
|
|
3d3ea47d37 | ||
|
|
f7f14d0e5f | ||
|
|
7580d80d45 | ||
|
|
cb51a3587b | ||
|
|
1bcd8f53ab | ||
|
|
0d0dc64884 | ||
|
|
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 |
@@ -54,6 +54,9 @@ jobs:
|
|||||||
- name: Build wheel (with CUDA kernels)
|
- name: Build wheel (with CUDA kernels)
|
||||||
run: |
|
run: |
|
||||||
CSRC_KERNELS=true pip wheel . --no-deps --no-build-isolation -w dist/
|
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
|
- uses: actions/upload-artifact@v4
|
||||||
with:
|
with:
|
||||||
|
|||||||
@@ -9,6 +9,7 @@
|
|||||||
!scripts/**/*.py
|
!scripts/**/*.py
|
||||||
!tests/**/*.py
|
!tests/**/*.py
|
||||||
!csrc/**/*.py
|
!csrc/**/*.py
|
||||||
|
!csrc/CMakeLists.txt
|
||||||
|
|
||||||
!csrc/**/*.cu
|
!csrc/**/*.cu
|
||||||
!csrc/**/*.h
|
!csrc/**/*.h
|
||||||
|
|||||||
+13
-10
@@ -20,9 +20,6 @@ Run the following checks **in order** — CI will reject if any fail.
|
|||||||
ruff format .
|
ruff format .
|
||||||
```
|
```
|
||||||
|
|
||||||
> **Note**: `ruff format` may rename parameters (e.g. `mask` → `attn_mask`).
|
|
||||||
> Always review the diff after formatting.
|
|
||||||
|
|
||||||
### 2. Import sorting
|
### 2. Import sorting
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
@@ -42,22 +39,28 @@ ruff format . # re-format after fix
|
|||||||
python -u -m pytest tests/ -v
|
python -u -m pytest tests/ -v
|
||||||
```
|
```
|
||||||
|
|
||||||
> Failed tests may leave orphan tempdirs under `%TEMP%`. Clean them manually if needed.
|
> Failed tests may leave orphan tempdirs under the system temp directory
|
||||||
|
> (`$TMPDIR` on Linux/macOS, `%TEMP%` on Windows). Clean them manually if needed.
|
||||||
|
|
||||||
### 4. (Optional) Full pre-commit check
|
### 4. (Optional) Full pre-commit check script
|
||||||
|
|
||||||
If you have Git Bash available:
|
If you have `bash` available (Git Bash on Windows works too):
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
bash scripts/pre_commit.sh
|
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
|
## 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)
|
- bullet point body (each ~60 chars)
|
||||||
```
|
```
|
||||||
@@ -73,7 +76,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 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 |
|
| `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 |
|
| Tests fail with tempdir left | Test crash | Clean `%TEMP%` manually |
|
||||||
|
|
||||||
## Submitting Changes
|
## Submitting Changes
|
||||||
@@ -93,7 +96,7 @@ fix/feat/chore/docs/refactor/perf/test/style/ci/build/revert : short description
|
|||||||
|
|
||||||
## License
|
## 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).
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
|||||||
+11
-2
@@ -57,8 +57,17 @@ COPY docs/ ./docs/
|
|||||||
COPY pyproject.toml .
|
COPY pyproject.toml .
|
||||||
COPY README.md .
|
COPY README.md .
|
||||||
|
|
||||||
# Create non-root user
|
# Create non-root user matching the host uid/gid (passed via build args).
|
||||||
RUN useradd -m astrai && chown -R astrai:astrai /app
|
# ubuntu:24.04 ships a default 'ubuntu' user/group at uid/gid 1000, so remove
|
||||||
|
# it first to free those ids before creating astrai.
|
||||||
|
ARG USER_UID=1000
|
||||||
|
ARG USER_GID=1000
|
||||||
|
RUN userdel -r ubuntu 2>/dev/null || true \
|
||||||
|
&& groupdel ubuntu 2>/dev/null || true \
|
||||||
|
&& groupadd -g "${USER_GID}" astrai \
|
||||||
|
&& useradd -m -u "${USER_UID}" -g astrai astrai \
|
||||||
|
&& chown -R astrai:astrai /app
|
||||||
|
ENV HOME=/home/astrai
|
||||||
USER astrai
|
USER astrai
|
||||||
|
|
||||||
ENV PYTHONUNBUFFERED=1 \
|
ENV PYTHONUNBUFFERED=1 \
|
||||||
|
|||||||
@@ -1,674 +1,201 @@
|
|||||||
GNU GENERAL PUBLIC LICENSE
|
Apache License
|
||||||
Version 3, 29 June 2007
|
Version 2.0, January 2004
|
||||||
|
http://www.apache.org/licenses/
|
||||||
Copyright (C) 2007 Free Software Foundation, Inc. <https://fsf.org/>
|
|
||||||
Everyone is permitted to copy and distribute verbatim copies
|
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
||||||
of this license document, but changing it is not allowed.
|
|
||||||
|
1. Definitions.
|
||||||
Preamble
|
|
||||||
|
"License" shall mean the terms and conditions for use, reproduction,
|
||||||
The GNU General Public License is a free, copyleft license for
|
and distribution as defined by Sections 1 through 9 of this document.
|
||||||
software and other kinds of works.
|
|
||||||
|
"Licensor" shall mean the copyright owner or entity authorized by
|
||||||
The licenses for most software and other practical works are designed
|
the copyright owner that is granting the License.
|
||||||
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
|
"Legal Entity" shall mean the union of the acting entity and all
|
||||||
share and change all versions of a program--to make sure it remains free
|
other entities that control, are controlled by, or are under common
|
||||||
software for all its users. We, the Free Software Foundation, use the
|
control with that entity. For the purposes of this definition,
|
||||||
GNU General Public License for most of our software; it applies also to
|
"control" means (i) the power, direct or indirect, to cause the
|
||||||
any other work released this way by its authors. You can apply it to
|
direction or management of such entity, whether by contract or
|
||||||
your programs, too.
|
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
||||||
|
outstanding shares, or (iii) beneficial ownership of such entity.
|
||||||
When we speak of free software, we are referring to freedom, not
|
|
||||||
price. Our General Public Licenses are designed to make sure that you
|
"You" (or "Your") shall mean an individual or Legal Entity
|
||||||
have the freedom to distribute copies of free software (and charge for
|
exercising permissions granted by this License.
|
||||||
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
|
"Source" form shall mean the preferred form for making modifications,
|
||||||
free programs, and that you know you can do these things.
|
including but not limited to software source code, documentation
|
||||||
|
source, and configuration files.
|
||||||
To protect your rights, we need to prevent others from denying you
|
|
||||||
these rights or asking you to surrender the rights. Therefore, you have
|
"Object" form shall mean any form resulting from mechanical
|
||||||
certain responsibilities if you distribute copies of the software, or if
|
transformation or translation of a Source form, including but
|
||||||
you modify it: responsibilities to respect the freedom of others.
|
not limited to compiled object code, generated documentation,
|
||||||
|
and conversions to other media types.
|
||||||
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
|
"Work" shall mean the work of authorship, whether in Source or
|
||||||
freedoms that you received. You must make sure that they, too, receive
|
Object form, made available under the License, as indicated by a
|
||||||
or can get the source code. And you must show them these terms so they
|
copyright notice that is included in or attached to the work
|
||||||
know their rights.
|
(an example is provided in the Appendix below).
|
||||||
|
|
||||||
Developers that use the GNU GPL protect your rights with two steps:
|
"Derivative Works" shall mean any work, whether in Source or Object
|
||||||
(1) assert copyright on the software, and (2) offer you this License
|
form, that is based on (or derived from) the Work and for which the
|
||||||
giving you legal permission to copy, distribute and/or modify it.
|
editorial revisions, annotations, elaborations, or other modifications
|
||||||
|
represent, as a whole, an original work of authorship. For the purposes
|
||||||
For the developers' and authors' protection, the GPL clearly explains
|
of this License, Derivative Works shall not include works that remain
|
||||||
that there is no warranty for this free software. For both users' and
|
separable from, or merely link (or bind by name) to the interfaces of,
|
||||||
authors' sake, the GPL requires that modified versions be marked as
|
the Work and Derivative Works thereof.
|
||||||
changed, so that their problems will not be attributed erroneously to
|
|
||||||
authors of previous versions.
|
"Contribution" shall mean any work of authorship, including
|
||||||
|
the original version of the Work and any modifications or additions
|
||||||
Some devices are designed to deny users access to install or run
|
to that Work or Derivative Works thereof, that is intentionally
|
||||||
modified versions of the software inside them, although the manufacturer
|
submitted to Licensor for inclusion in the Work by the copyright owner
|
||||||
can do so. This is fundamentally incompatible with the aim of
|
or by an individual or Legal Entity authorized to submit on behalf of
|
||||||
protecting users' freedom to change the software. The systematic
|
the copyright owner. For the purposes of this definition, "submitted"
|
||||||
pattern of such abuse occurs in the area of products for individuals to
|
means any form of electronic, verbal, or written communication sent
|
||||||
use, which is precisely where it is most unacceptable. Therefore, we
|
to the Licensor or its representatives, including but not limited to
|
||||||
have designed this version of the GPL to prohibit the practice for those
|
communication on electronic mailing lists, source code control systems,
|
||||||
products. If such problems arise substantially in other domains, we
|
and issue tracking systems that are managed by, or on behalf of, the
|
||||||
stand ready to extend this provision to those domains in future versions
|
Licensor for the purpose of discussing and improving the Work, but
|
||||||
of the GPL, as needed to protect the freedom of users.
|
excluding communication that is conspicuously marked or otherwise
|
||||||
|
designated in writing by the copyright owner as "Not a Contribution."
|
||||||
Finally, every program is threatened constantly by software patents.
|
|
||||||
States should not allow patents to restrict development and use of
|
"Contributor" shall mean Licensor and any individual or Legal Entity
|
||||||
software on general-purpose computers, but in those that do, we wish to
|
on behalf of whom a Contribution has been received by Licensor and
|
||||||
avoid the special danger that patents applied to a free program could
|
subsequently incorporated within the Work.
|
||||||
make it effectively proprietary. To prevent this, the GPL assures that
|
|
||||||
patents cannot be used to render the program non-free.
|
2. Grant of Copyright License. Subject to the terms and conditions of
|
||||||
|
this License, each Contributor hereby grants to You a perpetual,
|
||||||
The precise terms and conditions for copying, distribution and
|
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||||
modification follow.
|
copyright license to reproduce, prepare Derivative Works of,
|
||||||
|
publicly display, publicly perform, sublicense, and distribute the
|
||||||
TERMS AND CONDITIONS
|
Work and such Derivative Works in Source or Object form.
|
||||||
|
|
||||||
0. Definitions.
|
3. Grant of Patent License. Subject to the terms and conditions of
|
||||||
|
this License, each Contributor hereby grants to You a perpetual,
|
||||||
"This License" refers to version 3 of the GNU General Public License.
|
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||||
|
(except as stated in this section) patent license to make, have made,
|
||||||
"Copyright" also means copyright-like laws that apply to other kinds of
|
use, offer to sell, sell, import, and otherwise transfer the Work,
|
||||||
works, such as semiconductor masks.
|
where such license applies only to those patent claims licensable
|
||||||
|
by such Contributor that are necessarily infringed by their
|
||||||
"The Program" refers to any copyrightable work licensed under this
|
Contribution(s) alone or by combination of their Contribution(s)
|
||||||
License. Each licensee is addressed as "you". "Licensees" and
|
with the Work to which such Contribution(s) was submitted. If You
|
||||||
"recipients" may be individuals or organizations.
|
institute patent litigation against any entity (including a
|
||||||
|
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
||||||
To "modify" a work means to copy from or adapt all or part of the work
|
or a Contribution incorporated within the Work constitutes direct
|
||||||
in a fashion requiring copyright permission, other than the making of an
|
or contributory patent infringement, then any patent licenses
|
||||||
exact copy. The resulting work is called a "modified version" of the
|
granted to You under this License for that Work shall terminate
|
||||||
earlier work or a work "based on" the earlier work.
|
as of the date such litigation is filed.
|
||||||
|
|
||||||
A "covered work" means either the unmodified Program or a work based
|
4. Redistribution. You may reproduce and distribute copies of the
|
||||||
on the Program.
|
Work or Derivative Works thereof in any medium, with or without
|
||||||
|
modifications, and in Source or Object form, provided that You
|
||||||
To "propagate" a work means to do anything with it that, without
|
meet the following conditions:
|
||||||
permission, would make you directly or secondarily liable for
|
|
||||||
infringement under applicable copyright law, except executing it on a
|
(a) You must give any other recipients of the Work or
|
||||||
computer or modifying a private copy. Propagation includes copying,
|
Derivative Works a copy of this License; and
|
||||||
distribution (with or without modification), making available to the
|
|
||||||
public, and in some countries other activities as well.
|
(b) You must cause any modified files to carry prominent notices
|
||||||
|
stating that You changed the files; and
|
||||||
To "convey" a work means any kind of propagation that enables other
|
|
||||||
parties to make or receive copies. Mere interaction with a user through
|
(c) You must retain, in the Source form of any Derivative Works
|
||||||
a computer network, with no transfer of a copy, is not conveying.
|
that You distribute, all copyright, patent, trademark, and
|
||||||
|
attribution notices from the Source form of the Work,
|
||||||
An interactive user interface displays "Appropriate Legal Notices"
|
excluding those notices that do not pertain to any part of
|
||||||
to the extent that it includes a convenient and prominently visible
|
the Derivative Works; and
|
||||||
feature that (1) displays an appropriate copyright notice, and (2)
|
|
||||||
tells the user that there is no warranty for the work (except to the
|
(d) If the Work includes a "NOTICE" text file as part of its
|
||||||
extent that warranties are provided), that licensees may convey the
|
distribution, then any Derivative Works that You distribute must
|
||||||
work under this License, and how to view a copy of this License. If
|
include a readable copy of the attribution notices contained
|
||||||
the interface presents a list of user commands or options, such as a
|
within such NOTICE file, excluding those notices that do not
|
||||||
menu, a prominent item in the list meets this criterion.
|
pertain to any part of the Derivative Works, in at least one
|
||||||
|
of the following places: within a NOTICE text file distributed
|
||||||
1. Source Code.
|
as part of the Derivative Works; within the Source form or
|
||||||
|
documentation, if provided along with the Derivative Works; or,
|
||||||
The "source code" for a work means the preferred form of the work
|
within a display generated by the Derivative Works, if and
|
||||||
for making modifications to it. "Object code" means any non-source
|
wherever such third-party notices normally appear. The contents
|
||||||
form of a work.
|
of the NOTICE file are for informational purposes only and
|
||||||
|
do not modify the License. You may add Your own attribution
|
||||||
A "Standard Interface" means an interface that either is an official
|
notices within Derivative Works that You distribute, alongside
|
||||||
standard defined by a recognized standards body, or, in the case of
|
or as an addendum to the NOTICE text from the Work, provided
|
||||||
interfaces specified for a particular programming language, one that
|
that such additional attribution notices cannot be construed
|
||||||
is widely used among developers working in that language.
|
as modifying the License.
|
||||||
|
|
||||||
The "System Libraries" of an executable work include anything, other
|
You may add Your own copyright statement to Your modifications and
|
||||||
than the work as a whole, that (a) is included in the normal form of
|
may provide additional or different license terms and conditions
|
||||||
packaging a Major Component, but which is not part of that Major
|
for use, reproduction, or distribution of Your modifications, or
|
||||||
Component, and (b) serves only to enable use of the work with that
|
for any such Derivative Works as a whole, provided Your use,
|
||||||
Major Component, or to implement a Standard Interface for which an
|
reproduction, and distribution of the Work otherwise complies with
|
||||||
implementation is available to the public in source code form. A
|
the conditions stated in this License.
|
||||||
"Major Component", in this context, means a major essential component
|
|
||||||
(kernel, window system, and so on) of the specific operating system
|
5. Submission of Contributions. Unless You explicitly state otherwise,
|
||||||
(if any) on which the executable work runs, or a compiler used to
|
any Contribution intentionally submitted for inclusion in the Work
|
||||||
produce the work, or an object code interpreter used to run it.
|
by You to the Licensor shall be under the terms and conditions of
|
||||||
|
this License, without any additional terms or conditions.
|
||||||
The "Corresponding Source" for a work in object code form means all
|
Notwithstanding the above, nothing herein shall supersede or modify
|
||||||
the source code needed to generate, install, and (for an executable
|
the terms of any separate license agreement you may have executed
|
||||||
work) run the object code and to modify the work, including scripts to
|
with Licensor regarding such Contributions.
|
||||||
control those activities. However, it does not include the work's
|
|
||||||
System Libraries, or general-purpose tools or generally available free
|
6. Trademarks. This License does not grant permission to use the trade
|
||||||
programs which are used unmodified in performing those activities but
|
names, trademarks, service marks, or product names of the Licensor,
|
||||||
which are not part of the work. For example, Corresponding Source
|
except as required for reasonable and customary use in describing the
|
||||||
includes interface definition files associated with source files for
|
origin of the Work and reproducing the content of the NOTICE file.
|
||||||
the work, and the source code for shared libraries and dynamically
|
|
||||||
linked subprograms that the work is specifically designed to require,
|
7. Disclaimer of Warranty. Unless required by applicable law or
|
||||||
such as by intimate data communication or control flow between those
|
agreed to in writing, Licensor provides the Work (and each
|
||||||
subprograms and other parts of the work.
|
Contributor provides its Contributions) on an "AS IS" BASIS,
|
||||||
|
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
||||||
The Corresponding Source need not include anything that users
|
implied, including, without limitation, any warranties or conditions
|
||||||
can regenerate automatically from other parts of the Corresponding
|
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
||||||
Source.
|
PARTICULAR PURPOSE. You are solely responsible for determining the
|
||||||
|
appropriateness of using or redistributing the Work and assume any
|
||||||
The Corresponding Source for a work in source code form is that
|
risks associated with Your exercise of permissions under this License.
|
||||||
same work.
|
|
||||||
|
8. Limitation of Liability. In no event and under no legal theory,
|
||||||
2. Basic Permissions.
|
whether in tort (including negligence), contract, or otherwise,
|
||||||
|
unless required by applicable law (such as deliberate and grossly
|
||||||
All rights granted under this License are granted for the term of
|
negligent acts) or agreed to in writing, shall any Contributor be
|
||||||
copyright on the Program, and are irrevocable provided the stated
|
liable to You for damages, including any direct, indirect, special,
|
||||||
conditions are met. This License explicitly affirms your unlimited
|
incidental, or consequential damages of any character arising as a
|
||||||
permission to run the unmodified Program. The output from running a
|
result of this License or out of the use or inability to use the
|
||||||
covered work is covered by this License only if the output, given its
|
Work (including but not limited to damages for loss of goodwill,
|
||||||
content, constitutes a covered work. This License acknowledges your
|
work stoppage, computer failure or malfunction, or any and all
|
||||||
rights of fair use or other equivalent, as provided by copyright law.
|
other commercial damages or losses), even if such Contributor
|
||||||
|
has been advised of the possibility of such damages.
|
||||||
You may make, run and propagate covered works that you do not
|
|
||||||
convey, without conditions so long as your license otherwise remains
|
9. Accepting Warranty or Additional Liability. While redistributing
|
||||||
in force. You may convey covered works to others for the sole purpose
|
the Work or Derivative Works thereof, You may choose to offer,
|
||||||
of having them make modifications exclusively for you, or provide you
|
and charge a fee for, acceptance of support, warranty, indemnity,
|
||||||
with facilities for running those works, provided that you comply with
|
or other liability obligations and/or rights consistent with this
|
||||||
the terms of this License in conveying all material for which you do
|
License. However, in accepting such obligations, You may act only
|
||||||
not control copyright. Those thus making or running the covered works
|
on Your own behalf and on Your sole responsibility, not on behalf
|
||||||
for you must do so exclusively on your behalf, under your direction
|
of any other Contributor, and only if You agree to indemnify,
|
||||||
and control, on terms that prohibit them from making any copies of
|
defend, and hold each Contributor harmless for any liability
|
||||||
your copyrighted material outside their relationship with you.
|
incurred by, or claims asserted against, such Contributor by reason
|
||||||
|
of your accepting any such warranty or additional liability.
|
||||||
Conveying under any other circumstances is permitted solely under
|
|
||||||
the conditions stated below. Sublicensing is not allowed; section 10
|
END OF TERMS AND CONDITIONS
|
||||||
makes it unnecessary.
|
|
||||||
|
APPENDIX: How to apply the Apache License to your work.
|
||||||
3. Protecting Users' Legal Rights From Anti-Circumvention Law.
|
|
||||||
|
To apply the Apache License to your work, attach the following
|
||||||
No covered work shall be deemed part of an effective technological
|
boilerplate notice, with the fields enclosed by brackets "[]"
|
||||||
measure under any applicable law fulfilling obligations under article
|
replaced with your own identifying information. (Don't include
|
||||||
11 of the WIPO copyright treaty adopted on 20 December 1996, or
|
the brackets!) The text should be enclosed in the appropriate
|
||||||
similar laws prohibiting or restricting circumvention of such
|
comment syntax for the file format. We also recommend that a
|
||||||
measures.
|
file or class name and description of purpose be included on the
|
||||||
|
same "printed page" as the copyright notice for easier
|
||||||
When you convey a covered work, you waive any legal power to forbid
|
identification within third-party archives.
|
||||||
circumvention of technological measures to the extent such circumvention
|
|
||||||
is effected by exercising rights under this License with respect to
|
Copyright [yyyy] [name of copyright owner]
|
||||||
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
|
Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
users, your or third parties' legal rights to forbid circumvention of
|
you may not use this file except in compliance with the License.
|
||||||
technological measures.
|
You may obtain a copy of the License at
|
||||||
|
|
||||||
4. Conveying Verbatim Copies.
|
http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
|
||||||
You may convey verbatim copies of the Program's source code as you
|
Unless required by applicable law or agreed to in writing, software
|
||||||
receive it, in any medium, provided that you conspicuously and
|
distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
appropriately publish on each copy an appropriate copyright notice;
|
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
keep intact all notices stating that this License and any
|
See the License for the specific language governing permissions and
|
||||||
non-permissive terms added in accord with section 7 apply to the code;
|
limitations under the License.
|
||||||
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>.
|
|
||||||
|
|||||||
@@ -8,7 +8,7 @@
|
|||||||
|
|
||||||
<div align="center">
|
<div align="center">
|
||||||
<img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python">
|
<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/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/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">
|
<img src="https://img.shields.io/github/forks/ViperEkura/AstrAI?style=flat&label=Forks&color=76bad9" alt="forks">
|
||||||
@@ -27,7 +27,7 @@
|
|||||||
|
|
||||||
## 📖 Table of Contents
|
## 📖 Table of Contents
|
||||||
|
|
||||||
- [Features](#features)
|
- [Overview](#overview)
|
||||||
- [Getting Started](#getting-started)
|
- [Getting Started](#getting-started)
|
||||||
- [Demo](#demo)
|
- [Demo](#demo)
|
||||||
- [Documentation](#documentation)
|
- [Documentation](#documentation)
|
||||||
@@ -40,15 +40,19 @@
|
|||||||
<a id="english"></a>
|
<a id="english"></a>
|
||||||
## English
|
## English
|
||||||
|
|
||||||
### Features
|
### Overview
|
||||||
|
|
||||||
- 🚀 **High Performance**: Optimized for both training and inference with efficient parallelization.
|
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.
|
||||||
- 🔧 **Flexible**: Support for seq/sft/dpo/grpo training, customizable model architectures.
|
|
||||||
- 💡 **Easy to Use**: Simple API with comprehensive examples and demos.
|
| Area | Capabilities |
|
||||||
- 📦 **Lightweight**: Minimal dependencies, easy to deploy.
|
|---|---|
|
||||||
- 🔬 **Research‑Friendly**: Modular design, easy to experiment with new ideas.
|
| **Models** | Autoregressive language models and embedding models with GQA, MLA, MoE, RoPE, and extensible attention/FFN components |
|
||||||
- 🤗 **HuggingFace-Style API**: AutoModel/AutoTokenizer APIs inspired by HuggingFace for easy model and tokenizer loading.
|
| **Training** | Pre-training (`seq`), supervised fine-tuning (`sft`), DPO, and GRPO with gradient accumulation, checkpointing, DDP, and FSDP |
|
||||||
- 🔌 **Dual API Compatibility**: Supports both OpenAI and Anthropic chat completion APIs out of the box.
|
| **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, ROUGE, and weight-analysis evaluation tools |
|
||||||
|
| **Extensibility** | Factory and registry architecture for models, datasets, training strategies, callbacks, kernels, and protocol components |
|
||||||
|
|
||||||
### Getting Started
|
### Getting Started
|
||||||
|
|
||||||
@@ -56,11 +60,14 @@ End-to-end walkthrough in 5 steps:
|
|||||||
|
|
||||||
**1. Install**
|
**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
|
```bash
|
||||||
git clone https://github.com/ViperEkura/AstrAI.git
|
git clone https://github.com/ViperEkura/AstrAI.git
|
||||||
cd AstrAI
|
cd AstrAI
|
||||||
pip install -e . # pure PyTorch (no CUDA kernels)
|
pip install -e . # kernels auto-build when nvcc + CUDA are detected
|
||||||
# CSRC_KERNELS=true pip install -e . --no-build-isolation # optional: fused CUDA kernels
|
# CSRC_KERNELS=false pip install -e . # skip kernels (pure PyTorch)
|
||||||
|
# CSRC_KERNELS=true pip install -e . --no-build-isolation # force the fused CUDA kernel build
|
||||||
# pip install -e ".[dev]" # dev dependencies (pytest, ruff)
|
# pip install -e ".[dev]" # dev dependencies (pytest, ruff)
|
||||||
```
|
```
|
||||||
|
|
||||||
@@ -132,7 +139,7 @@ Check out the demos in the `scripts/demo/` folder:
|
|||||||
# Download model weights (required before running demos)
|
# Download model weights (required before running demos)
|
||||||
python scripts/demo/download.py # model → params/
|
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
|
python scripts/demo/stream_chat.py
|
||||||
# Type your message after >>, type !exit to quit
|
# Type your message after >>, type !exit to quit
|
||||||
|
|
||||||
@@ -183,8 +190,11 @@ docker run --gpus all -v /path/to/data:/data -it astrai:latest
|
|||||||
# Docker Compose (GPU, default)
|
# Docker Compose (GPU, default)
|
||||||
docker compose up -d
|
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
|
docker compose --profile cpu up -d
|
||||||
|
|
||||||
|
# YAML-driven serving (see serve.yaml; up/run/down/logs/status...)
|
||||||
|
bash scripts/serve.sh up
|
||||||
```
|
```
|
||||||
|
|
||||||
> **Note**: `--gpus all` is required for CUDA support. Without it, `torch.cuda.is_available()` will return `False`.
|
> **Note**: `--gpus all` is required for CUDA support. Without it, `torch.cuda.is_available()` will return `False`.
|
||||||
@@ -230,6 +240,8 @@ See [Inference Guide](docs/guides/inference.md) for SSE streaming format, error
|
|||||||
| [Data Flow](./docs/developer/dataflow.md) | Data pipeline, storage backends & dataset architecture |
|
| [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 |
|
| [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 |
|
| [CUDA Kernels](./docs/developer/cuda_kernels.md) | Custom CUDA attention kernels & benchmarks |
|
||||||
|
| [Docker Serving](./docs/developer/docker-serving.md) | YAML-driven containerized serving (`serve.yaml`, `serve.sh`) |
|
||||||
|
| [Docker Training](./docs/developer/docker-training.md) | YAML-driven containerized training (`train.yaml`, `train.sh`) |
|
||||||
|
|
||||||
### Contributing
|
### Contributing
|
||||||
|
|
||||||
@@ -250,7 +262,7 @@ For major changes, please open an issue first to discuss what you would like to
|
|||||||
|
|
||||||
### License
|
### License
|
||||||
|
|
||||||
This project is licensed under the [GPL-3.0 License](LICENSE).
|
This project is licensed under the [Apache-2.0 License](LICENSE).
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
|||||||
+7
-38
@@ -1,9 +1,6 @@
|
|||||||
__version__ = "1.3.12"
|
__version__ = "1.3.13"
|
||||||
__author__ = "ViperEkura"
|
__author__ = "ViperEkura"
|
||||||
|
|
||||||
import logging
|
|
||||||
import os
|
|
||||||
|
|
||||||
from astrai.config import (
|
from astrai.config import (
|
||||||
AutoRegressiveLMConfig,
|
AutoRegressiveLMConfig,
|
||||||
BaseModelConfig,
|
BaseModelConfig,
|
||||||
@@ -20,15 +17,10 @@ from astrai.dataset import (
|
|||||||
StoreFactory,
|
StoreFactory,
|
||||||
)
|
)
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
from astrai.inference import (
|
from astrai.inference import InferenceEngine, get_app, run_server, sample
|
||||||
GenerationRequest,
|
from astrai.inference.network import ProtocolHandler
|
||||||
InferenceEngine,
|
from astrai.inference.runtime.sample import SamplingPipeline
|
||||||
ProtocolHandler,
|
from astrai.logging import setup_logging
|
||||||
SamplingPipeline,
|
|
||||||
get_app,
|
|
||||||
run_server,
|
|
||||||
sample,
|
|
||||||
)
|
|
||||||
from astrai.model import (
|
from astrai.model import (
|
||||||
AutoModel,
|
AutoModel,
|
||||||
AutoRegressiveLM,
|
AutoRegressiveLM,
|
||||||
@@ -56,30 +48,6 @@ from astrai.trainer import (
|
|||||||
Trainer,
|
Trainer,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def setup_logging(level: str = "INFO"):
|
|
||||||
"""Attach a handler to the ``astrai`` logger (only, not root).
|
|
||||||
|
|
||||||
Call once per process, e.g. at the top of CLI scripts.
|
|
||||||
Set ``ASTR_LOG_LEVEL`` to override the default ``INFO``.
|
|
||||||
"""
|
|
||||||
_logger = logging.getLogger("astrai")
|
|
||||||
if _logger.handlers:
|
|
||||||
return
|
|
||||||
_level = getattr(
|
|
||||||
logging, os.environ.get("ASTR_LOG_LEVEL", level).upper(), logging.INFO
|
|
||||||
)
|
|
||||||
_logger.setLevel(_level)
|
|
||||||
_handler = logging.StreamHandler()
|
|
||||||
_handler.setFormatter(
|
|
||||||
logging.Formatter(
|
|
||||||
"%(asctime)s | %(levelname)-7s | %(name)s | %(message)s",
|
|
||||||
datefmt="%Y-%m-%d %H:%M:%S",
|
|
||||||
)
|
|
||||||
)
|
|
||||||
_logger.addHandler(_handler)
|
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"AutoRegressiveLM",
|
"AutoRegressiveLM",
|
||||||
"AutoRegressiveLMConfig",
|
"AutoRegressiveLMConfig",
|
||||||
@@ -98,7 +66,6 @@ __all__ = [
|
|||||||
"EmbeddingEncoder",
|
"EmbeddingEncoder",
|
||||||
"EncoderConfig",
|
"EncoderConfig",
|
||||||
"ExecutorFactory",
|
"ExecutorFactory",
|
||||||
"GenerationRequest",
|
|
||||||
"InferenceEngine",
|
"InferenceEngine",
|
||||||
"LoRAConfig",
|
"LoRAConfig",
|
||||||
"Pipeline",
|
"Pipeline",
|
||||||
@@ -124,3 +91,5 @@ __all__ = [
|
|||||||
"setup_logging",
|
"setup_logging",
|
||||||
"spawn_parallel_fn",
|
"spawn_parallel_fn",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
setup_logging()
|
||||||
|
|||||||
@@ -63,6 +63,11 @@ class AutoRegressiveLMConfig(BaseModelConfig):
|
|||||||
n_shared_experts (Optional[int]): Number of shared 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.
|
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.
|
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
|
vocab_size: Optional[int] = None
|
||||||
@@ -87,6 +92,12 @@ class AutoRegressiveLMConfig(BaseModelConfig):
|
|||||||
n_shared_experts: Optional[int] = None
|
n_shared_experts: Optional[int] = None
|
||||||
n_activated_experts: Optional[int] = None
|
n_activated_experts: Optional[int] = None
|
||||||
topk_method: Optional[str] = None
|
topk_method: Optional[str] = None
|
||||||
|
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")
|
@field_validator("attn_type")
|
||||||
def _validate_attn_type(cls, v: str) -> str:
|
def _validate_attn_type(cls, v: str) -> str:
|
||||||
@@ -102,6 +113,12 @@ class AutoRegressiveLMConfig(BaseModelConfig):
|
|||||||
raise ValueError(f"ffn_type must be one of {sorted(_FFN_TYPES)}, got {v!r}")
|
raise ValueError(f"ffn_type must be one of {sorted(_FFN_TYPES)}, got {v!r}")
|
||||||
return v
|
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
|
@dataclass
|
||||||
@ConfigFactory.register("embedding")
|
@ConfigFactory.register("embedding")
|
||||||
|
|||||||
@@ -11,10 +11,10 @@ from torch.utils.data import Dataset
|
|||||||
from astrai.config.base import BaseConfig
|
from astrai.config.base import BaseConfig
|
||||||
from astrai.model.components.lora import LoRAConfig
|
from astrai.model.components.lora import LoRAConfig
|
||||||
|
|
||||||
_TRAIN_TYPES = frozenset({"seq", "sft", "dpo", "grpo", "online_grpo", "online_dpo"})
|
TRAIN_TYPES = frozenset({"seq", "sft", "dpo", "grpo", "online_grpo", "online_dpo"})
|
||||||
_PARALLEL_MODES = frozenset({"none", "ddp", "fsdp"})
|
PARALLEL_MODES = frozenset({"none", "ddp", "fsdp"})
|
||||||
_BACKENDS = frozenset({"nccl", "gloo"})
|
BACKENDS = frozenset({"nccl", "gloo"})
|
||||||
_START_METHODS = frozenset({"spawn", "fork", "forkserver"})
|
START_METHODS = frozenset({"spawn", "fork", "forkserver"})
|
||||||
_COMPILE_MODES = frozenset({"default", "reduce-overhead", "max-autotune"})
|
_COMPILE_MODES = frozenset({"default", "reduce-overhead", "max-autotune"})
|
||||||
|
|
||||||
|
|
||||||
@@ -48,6 +48,7 @@ class TrainConfig(BaseConfig):
|
|||||||
random_seed (int): Random seed. Defaults to 3407.
|
random_seed (int): Random seed. Defaults to 3407.
|
||||||
num_workers (int): Number of workers for dataloader. Defaults to 0.
|
num_workers (int): Number of workers for dataloader. Defaults to 0.
|
||||||
prefetch_factor (Optional[int]): Prefetch factor for dataloader. Defaults to None.
|
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.
|
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.
|
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.
|
nprocs (int): Number of processes for distributed training. Defaults to 1.
|
||||||
@@ -61,6 +62,7 @@ class TrainConfig(BaseConfig):
|
|||||||
val_split (Optional[float]): Ratio to split from training dataset for validation, e.g. 0.05. 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.
|
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.
|
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_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_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_k (int): Top-k filtering for online rollout, 0=disable. Defaults to 0.
|
||||||
@@ -68,7 +70,7 @@ class TrainConfig(BaseConfig):
|
|||||||
rollout_max_tokens (int): Maximum generated tokens per response in rollout. Defaults to 1024.
|
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.
|
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 {}.
|
executor_kwargs (Dict[str, Any]): Extra kwargs passed to ExecutorFactory.create(). Defaults to {}.
|
||||||
extra_kwargs (Dict[str, Any]): Other arguments. Defaults to {}.
|
strategy_kwargs (Dict[str, Any]): Extra strategy arguments. Defaults to {}.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
model_fn: Callable[[], nn.Module]
|
model_fn: Callable[[], nn.Module]
|
||||||
@@ -97,6 +99,7 @@ class TrainConfig(BaseConfig):
|
|||||||
random_seed: int = 3407
|
random_seed: int = 3407
|
||||||
num_workers: int = 0
|
num_workers: int = 0
|
||||||
prefetch_factor: Optional[int] = None
|
prefetch_factor: Optional[int] = None
|
||||||
|
persistent_workers: bool = False
|
||||||
pin_memory: bool = False
|
pin_memory: bool = False
|
||||||
collate_fn: Optional[Callable[[List[Any]], Any]] = None
|
collate_fn: Optional[Callable[[List[Any]], Any]] = None
|
||||||
|
|
||||||
@@ -112,6 +115,7 @@ class TrainConfig(BaseConfig):
|
|||||||
val_split: Optional[float] = None
|
val_split: Optional[float] = None
|
||||||
val_step: int = 1000
|
val_step: int = 1000
|
||||||
neftune_alpha: float = 0.0
|
neftune_alpha: float = 0.0
|
||||||
|
moe_aux_loss_coef: float = 0.01
|
||||||
|
|
||||||
rollout_interval: int = 512
|
rollout_interval: int = 512
|
||||||
rollout_temperature: float = 0.7
|
rollout_temperature: float = 0.7
|
||||||
@@ -121,35 +125,35 @@ class TrainConfig(BaseConfig):
|
|||||||
reward_model_fn: Optional[Callable] = None
|
reward_model_fn: Optional[Callable] = None
|
||||||
|
|
||||||
executor_kwargs: Dict[str, Any] = field(default_factory=dict)
|
executor_kwargs: Dict[str, Any] = field(default_factory=dict)
|
||||||
extra_kwargs: Dict[str, Any] = field(default_factory=dict)
|
strategy_kwargs: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
@field_validator("strategy")
|
@field_validator("strategy")
|
||||||
def _validate_strategy(cls, v: str) -> str:
|
def _validate_strategy(cls, v: str) -> str:
|
||||||
if v not in _TRAIN_TYPES:
|
if v not in TRAIN_TYPES:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"strategy must be one of {sorted(_TRAIN_TYPES)}, got {v!r}"
|
f"strategy must be one of {sorted(TRAIN_TYPES)}, got {v!r}"
|
||||||
)
|
)
|
||||||
return v
|
return v
|
||||||
|
|
||||||
@field_validator("parallel_mode")
|
@field_validator("parallel_mode")
|
||||||
def _validate_parallel_mode(cls, v: str) -> str:
|
def _validate_parallel_mode(cls, v: str) -> str:
|
||||||
if v not in _PARALLEL_MODES:
|
if v not in PARALLEL_MODES:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"parallel_mode must be one of {sorted(_PARALLEL_MODES)}, got {v!r}"
|
f"parallel_mode must be one of {sorted(PARALLEL_MODES)}, got {v!r}"
|
||||||
)
|
)
|
||||||
return v
|
return v
|
||||||
|
|
||||||
@field_validator("backend")
|
@field_validator("backend")
|
||||||
def _validate_backend(cls, v: str) -> str:
|
def _validate_backend(cls, v: str) -> str:
|
||||||
if v not in _BACKENDS:
|
if v not in BACKENDS:
|
||||||
raise ValueError(f"backend must be one of {sorted(_BACKENDS)}, got {v!r}")
|
raise ValueError(f"backend must be one of {sorted(BACKENDS)}, got {v!r}")
|
||||||
return v
|
return v
|
||||||
|
|
||||||
@field_validator("start_method")
|
@field_validator("start_method")
|
||||||
def _validate_start_method(cls, v: str) -> str:
|
def _validate_start_method(cls, v: str) -> str:
|
||||||
if v not in _START_METHODS:
|
if v not in START_METHODS:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"start_method must be one of {sorted(_START_METHODS)}, got {v!r}"
|
f"start_method must be one of {sorted(START_METHODS)}, got {v!r}"
|
||||||
)
|
)
|
||||||
return v
|
return v
|
||||||
|
|
||||||
@@ -187,7 +191,9 @@ class TrainConfig(BaseConfig):
|
|||||||
raise ValueError(f"rollout_top_p must be in (0, 1], got {v}")
|
raise ValueError(f"rollout_top_p must be in (0, 1], got {v}")
|
||||||
return v
|
return v
|
||||||
|
|
||||||
@field_validator("rollout_top_k", "num_workers", "neftune_alpha")
|
@field_validator(
|
||||||
|
"rollout_top_k", "num_workers", "neftune_alpha", "moe_aux_loss_coef"
|
||||||
|
)
|
||||||
def _validate_non_negative(cls, v):
|
def _validate_non_negative(cls, v):
|
||||||
if v < 0:
|
if v < 0:
|
||||||
raise ValueError(f"must be non-negative, got {v}")
|
raise ValueError(f"must be non-negative, got {v}")
|
||||||
|
|||||||
+41
-12
@@ -25,20 +25,50 @@ function (pure ``record -> Dict[str, Tensor]``) is forwarded to
|
|||||||
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from functools import partial
|
from functools import partial
|
||||||
|
from pathlib import Path
|
||||||
from typing import Callable, Dict, List, Optional
|
from typing import Callable, Dict, List, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
from torch.utils.data import Dataset
|
from torch.utils.data import Dataset
|
||||||
|
|
||||||
|
from astrai.config.preprocess_config import PipelineConfig
|
||||||
from astrai.dataset.storage import (
|
from astrai.dataset.storage import (
|
||||||
Store,
|
Store,
|
||||||
StoreFactory,
|
StoreFactory,
|
||||||
detect_format,
|
detect_format,
|
||||||
)
|
)
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.preprocessing.transform import TokenizeTransform
|
||||||
from astrai.tokenize import AutoTokenizer
|
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(
|
def dpo_tokenize(
|
||||||
record: dict,
|
record: dict,
|
||||||
@@ -349,16 +379,18 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
|
|||||||
)
|
)
|
||||||
if processor is not None:
|
if processor is not None:
|
||||||
store.load(load_path, processor=processor, **kwargs)
|
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(
|
||||||
|
"JSONL dataset config not found. Expected "
|
||||||
|
"dataset_config.json alongside *.jsonl files, pass "
|
||||||
|
"tokenizer_path= for the built-in messages config, or "
|
||||||
|
"use processor= for lazy on-the-fly tokenisation."
|
||||||
|
)
|
||||||
|
store.load(load_path, transform=transform, **kwargs)
|
||||||
else:
|
else:
|
||||||
load_kwargs = dict(kwargs)
|
store.load(load_path, **kwargs)
|
||||||
if (
|
|
||||||
tokenizer_path is not None
|
|
||||||
and storage_type == "jsonl"
|
|
||||||
and train_type in ("seq", "sft")
|
|
||||||
and "tokenizer_path" not in load_kwargs
|
|
||||||
):
|
|
||||||
load_kwargs["tokenizer_path"] = tokenizer_path
|
|
||||||
store.load(load_path, **load_kwargs)
|
|
||||||
|
|
||||||
return cls.create(train_type, store=store)
|
return cls.create(train_type, store=store)
|
||||||
|
|
||||||
@@ -460,9 +492,6 @@ class DPODataset(BaseDataset):
|
|||||||
|
|
||||||
required_keys = ["chosen", "rejected", "chosen_mask", "rejected_mask"]
|
required_keys = ["chosen", "rejected", "chosen_mask", "rejected_mask"]
|
||||||
|
|
||||||
def make_processor(self, tokenizer, max_len: int):
|
|
||||||
return partial(dpo_processor, tokenizer=tokenizer, max_len=max_len)
|
|
||||||
|
|
||||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||||
return {
|
return {
|
||||||
"chosen": self.store.fetch_record(index, "chosen").to(dtype=torch.long),
|
"chosen": self.store.fetch_record(index, "chosen").to(dtype=torch.long),
|
||||||
|
|||||||
@@ -55,9 +55,7 @@ from typing import Callable, Dict, List, Optional, Tuple, Union
|
|||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from astrai.config.preprocess_config import PipelineConfig
|
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
from astrai.preprocessing.transform import TokenizeTransform
|
|
||||||
from astrai.serialization import (
|
from astrai.serialization import (
|
||||||
load_bin,
|
load_bin,
|
||||||
load_bin_offsets,
|
load_bin_offsets,
|
||||||
@@ -219,7 +217,7 @@ class Store(ABC):
|
|||||||
"""
|
"""
|
||||||
if self._window_size <= 0:
|
if self._window_size <= 0:
|
||||||
raise RuntimeError("sample_window() requires window_size > 0 (stream mode)")
|
raise RuntimeError("sample_window() requires window_size > 0 (stream mode)")
|
||||||
if self._window_size <= 0 or self._length <= self._window_size:
|
if self._length <= self._window_size:
|
||||||
raise IndexError(
|
raise IndexError(
|
||||||
f"Data too short for window: token_count={self._length}, "
|
f"Data too short for window: token_count={self._length}, "
|
||||||
f"window_size={self._window_size}"
|
f"window_size={self._window_size}"
|
||||||
@@ -536,19 +534,8 @@ class JsonlStore(Store, Streamable, Recordable):
|
|||||||
``len(store)`` returns ``num_records``; stream primitives raise.
|
``len(store)`` returns ``num_records``; stream primitives raise.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
CONFIG_NAME = "dataset_config.json"
|
|
||||||
segments_are_records = True
|
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__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
window_size: int = 0,
|
window_size: int = 0,
|
||||||
@@ -569,22 +556,10 @@ class JsonlStore(Store, Streamable, Recordable):
|
|||||||
return
|
return
|
||||||
|
|
||||||
if transform is None:
|
if transform is None:
|
||||||
root = Path(path)
|
raise ValueError(
|
||||||
config_path = root / self.CONFIG_NAME if root.is_dir() else None
|
"JsonlStore eager mode requires transform=. "
|
||||||
if config_path is not None and config_path.exists():
|
"Use DatasetFactory.load() which auto-constructs it."
|
||||||
transform = TokenizeTransform.from_config_file(str(config_path))
|
)
|
||||||
else:
|
|
||||||
tokenizer_path = kwargs.get("tokenizer_path")
|
|
||||||
if not tokenizer_path:
|
|
||||||
raise FileNotFoundError(
|
|
||||||
f"JSONL dataset config not found. Expected "
|
|
||||||
f"{self.CONFIG_NAME} alongside *.jsonl files, pass an "
|
|
||||||
f"explicit transform, pass processor= for lazy "
|
|
||||||
f"on-the-fly tokenisation, or pass tokenizer_path= to "
|
|
||||||
f"use the built-in messages config."
|
|
||||||
)
|
|
||||||
config = PipelineConfig.from_dict(self._DEFAULT_MESSAGES_CONFIG)
|
|
||||||
transform = TokenizeTransform(config, tokenizer_path)
|
|
||||||
|
|
||||||
transformed = transform.apply(records)
|
transformed = transform.apply(records)
|
||||||
self._normalize(transformed)
|
self._normalize(transformed)
|
||||||
|
|||||||
@@ -15,28 +15,34 @@ Each wrapper calls its compiled CUDA kernel directly. Fallback to torch
|
|||||||
SDPA is handled by the attention backend, not the wrapper functions.
|
SDPA is handled by the attention backend, not the wrapper functions.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from astrai.extension.attention_backend import (
|
from astrai.extension.backend import (
|
||||||
ATTN_BACKEND,
|
ATTN_BACKEND,
|
||||||
AttentionBackend,
|
AttentionBackend,
|
||||||
|
AttentionBackendFactory,
|
||||||
CudaBackend,
|
CudaBackend,
|
||||||
|
FlashAttnBackend,
|
||||||
TorchNativeBackend,
|
TorchNativeBackend,
|
||||||
|
apply_rotary_emb,
|
||||||
attention,
|
attention,
|
||||||
attn_backend,
|
attn_backend,
|
||||||
get_backend,
|
get_backend,
|
||||||
)
|
)
|
||||||
from astrai.extension.attention_ops import (
|
from astrai.extension.loader import KERNEL_NAMES, is_available
|
||||||
|
from astrai.extension.ops import (
|
||||||
|
TensorLayout,
|
||||||
attn_decode,
|
attn_decode,
|
||||||
attn_paged_decode,
|
attn_paged_decode,
|
||||||
attn_prefill,
|
attn_prefill,
|
||||||
)
|
)
|
||||||
from astrai.extension.loader import KERNEL_NAMES, is_available
|
|
||||||
from astrai.extension.rotary_backend import apply_rotary_emb
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"ATTN_BACKEND",
|
"ATTN_BACKEND",
|
||||||
"AttentionBackend",
|
"AttentionBackend",
|
||||||
|
"AttentionBackendFactory",
|
||||||
"CudaBackend",
|
"CudaBackend",
|
||||||
"TorchNativeBackend",
|
"TorchNativeBackend",
|
||||||
|
"FlashAttnBackend",
|
||||||
|
"TensorLayout",
|
||||||
"attention",
|
"attention",
|
||||||
"attn_backend",
|
"attn_backend",
|
||||||
"get_backend",
|
"get_backend",
|
||||||
|
|||||||
@@ -1,422 +0,0 @@
|
|||||||
"""Attention backend abstraction with context-manager switching.
|
|
||||||
|
|
||||||
The backend encapsulates KV cache I/O and attention computation. The
|
|
||||||
attention module (GQA/MLA) keeps projections, rotary, QK-norm, gating,
|
|
||||||
and output projection; the backend handles everything from "write K/V
|
|
||||||
to cache" through "SDPA output".
|
|
||||||
|
|
||||||
Usage — mirroring ``torch.nn.attention.sdpa_kernel``:
|
|
||||||
|
|
||||||
from astrai.extension import attn_backend, ATTN_BACKEND
|
|
||||||
|
|
||||||
with attn_backend(ATTN_BACKEND.TORCH_NATIVE):
|
|
||||||
engine.generate("hello")
|
|
||||||
|
|
||||||
# or with an instance:
|
|
||||||
with attn_backend(TorchNativeBackend()):
|
|
||||||
...
|
|
||||||
|
|
||||||
# or the shorthand (instance is itself a context manager):
|
|
||||||
with TorchNativeBackend():
|
|
||||||
...
|
|
||||||
|
|
||||||
Thread-safe via ``contextvars`` — each scheduler thread gets its own
|
|
||||||
active backend. ``get_backend()`` returns the active one, falling back
|
|
||||||
to a process-wide ``TorchNativeBackend`` singleton.
|
|
||||||
|
|
||||||
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
|
|
||||||
(blhd). The backend returns ``[batch, seq_len, n_heads * head_dim]``.
|
|
||||||
"""
|
|
||||||
|
|
||||||
import contextvars
|
|
||||||
import enum
|
|
||||||
from abc import ABC, abstractmethod
|
|
||||||
from contextlib import contextmanager
|
|
||||||
from typing import Optional, Union
|
|
||||||
|
|
||||||
import torch
|
|
||||||
import torch.nn.functional as F
|
|
||||||
from torch import Tensor
|
|
||||||
|
|
||||||
from astrai.extension.attention_ops import attn_paged_decode, attn_prefill
|
|
||||||
from astrai.extension.loader import is_available
|
|
||||||
from astrai.inference.core.cache import KVCache
|
|
||||||
|
|
||||||
_current_backend: contextvars.ContextVar["AttentionBackend"] = contextvars.ContextVar(
|
|
||||||
"attn_backend"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class ATTN_BACKEND(enum.Enum):
|
|
||||||
"""Backend selector enum, mirroring ``torch.nn.attention.SDPBackend``."""
|
|
||||||
|
|
||||||
TORCH_NATIVE = "torch_native"
|
|
||||||
CUDA = "cuda"
|
|
||||||
|
|
||||||
|
|
||||||
def get_backend() -> "AttentionBackend":
|
|
||||||
"""Return the active backend for the current thread/context.
|
|
||||||
|
|
||||||
Falls back to a ``TorchNativeBackend`` singleton when no backend
|
|
||||||
has been activated via ``with``.
|
|
||||||
"""
|
|
||||||
try:
|
|
||||||
return _current_backend.get()
|
|
||||||
except LookupError:
|
|
||||||
return _default_backend
|
|
||||||
|
|
||||||
|
|
||||||
@contextmanager
|
|
||||||
def attn_backend(backend: Union[ATTN_BACKEND, "AttentionBackend", type]):
|
|
||||||
"""Context manager to select an attention backend.
|
|
||||||
|
|
||||||
Mirrors ``torch.nn.attention.sdpa_kernel``. Accepts an
|
|
||||||
``ATTN_BACKEND`` enum value, a backend class, or a backend instance.
|
|
||||||
|
|
||||||
Examples::
|
|
||||||
|
|
||||||
with attn_backend(ATTN_BACKEND.TORCH_NATIVE):
|
|
||||||
...
|
|
||||||
with attn_backend(TorchNativeBackend):
|
|
||||||
...
|
|
||||||
with attn_backend(TorchNativeBackend()):
|
|
||||||
...
|
|
||||||
"""
|
|
||||||
if isinstance(backend, ATTN_BACKEND):
|
|
||||||
instance = _BACKEND_REGISTRY[backend]()
|
|
||||||
elif isinstance(backend, type) and issubclass(backend, AttentionBackend):
|
|
||||||
instance = backend()
|
|
||||||
elif isinstance(backend, AttentionBackend):
|
|
||||||
instance = backend
|
|
||||||
else:
|
|
||||||
raise TypeError(
|
|
||||||
f"expected ATTN_BACKEND, AttentionBackend type, or instance, "
|
|
||||||
f"got {type(backend).__name__}"
|
|
||||||
)
|
|
||||||
token = _current_backend.set(instance)
|
|
||||||
try:
|
|
||||||
yield instance
|
|
||||||
finally:
|
|
||||||
_current_backend.reset(token)
|
|
||||||
|
|
||||||
|
|
||||||
def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
|
|
||||||
"""Expand KV heads to match Q heads for GQA."""
|
|
||||||
bs, slen, n_heads, head_dim = x.shape
|
|
||||||
if n_rep == 1:
|
|
||||||
return x
|
|
||||||
return (
|
|
||||||
x[:, :, :, None, :]
|
|
||||||
.expand(bs, slen, n_heads, n_rep, head_dim)
|
|
||||||
.reshape(bs, slen, n_heads * n_rep, head_dim)
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def attention(
|
|
||||||
q: Tensor,
|
|
||||||
k: Tensor,
|
|
||||||
v: Tensor,
|
|
||||||
kv_cache: Optional[KVCache] = None,
|
|
||||||
layer_id: int = 0,
|
|
||||||
attn_mask: Optional[Tensor] = None,
|
|
||||||
is_causal: bool = False,
|
|
||||||
) -> Tensor:
|
|
||||||
"""Functional attention entry point — mirrors ``F.scaled_dot_product_attention``.
|
|
||||||
|
|
||||||
Delegates to the active backend (set via ``with attn_backend(...)``).
|
|
||||||
Handles KV cache I/O, GQA head expansion, and causal masking so the
|
|
||||||
caller only needs to provide projected q/k/v.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
q: [batch, q_len, n_heads, head_dim] (blhd)
|
|
||||||
k: [batch, q_len, n_kv_heads, head_dim] (blhd)
|
|
||||||
v: [batch, q_len, n_kv_heads, head_dim] (blhd)
|
|
||||||
kv_cache: cache dataclass, or None for training (no cache).
|
|
||||||
layer_id: transformer layer index for buffer access.
|
|
||||||
attn_mask: pre-built attention mask (SDPA-compatible).
|
|
||||||
is_causal: whether to apply causal masking.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
[batch, q_len, n_heads * head_dim]
|
|
||||||
"""
|
|
||||||
backend = get_backend()
|
|
||||||
return backend.forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
|
||||||
|
|
||||||
|
|
||||||
class AttentionBackend(ABC):
|
|
||||||
"""Abstract base for attention computation strategies.
|
|
||||||
|
|
||||||
Subclasses implement ``fwd_decode`` (q_len == 1, with cache) and
|
|
||||||
``fwd_prefill`` (q_len > 1, with or without cache). The public
|
|
||||||
``forward`` method dispatches based on q_len.
|
|
||||||
|
|
||||||
Three equivalent ways to activate a backend::
|
|
||||||
|
|
||||||
with attn_backend(ATTN_BACKEND.TORCH_NATIVE): # enum
|
|
||||||
...
|
|
||||||
with attn_backend(TorchNativeBackend): # class
|
|
||||||
...
|
|
||||||
with TorchNativeBackend(): # instance
|
|
||||||
...
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __enter__(self) -> "AttentionBackend":
|
|
||||||
self._token = _current_backend.set(self)
|
|
||||||
return self
|
|
||||||
|
|
||||||
def __exit__(self, *exc) -> None:
|
|
||||||
_current_backend.reset(self._token)
|
|
||||||
|
|
||||||
def forward(
|
|
||||||
self,
|
|
||||||
q: Tensor,
|
|
||||||
k: Tensor,
|
|
||||||
v: Tensor,
|
|
||||||
kv_cache: Optional[KVCache],
|
|
||||||
layer_id: int,
|
|
||||||
attn_mask: Optional[Tensor] = None,
|
|
||||||
is_causal: bool = False,
|
|
||||||
) -> Tensor:
|
|
||||||
"""Dispatch to decode or extend based on q_len.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
q: [batch, q_len, n_heads, head_dim]
|
|
||||||
k: [batch, q_len, n_kv_heads, head_dim]
|
|
||||||
v: [batch, q_len, n_kv_heads, head_dim]
|
|
||||||
kv_cache: cache dataclass, or None for training (no cache).
|
|
||||||
layer_id: transformer layer index for buffer access.
|
|
||||||
attn_mask: pre-built attention mask compatible with SDPA.
|
|
||||||
is_causal: whether to apply causal masking.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
[batch, q_len, n_heads * head_dim]
|
|
||||||
"""
|
|
||||||
if kv_cache is not None and q.size(1) == 1:
|
|
||||||
return self.fwd_decode(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
|
||||||
return self.fwd_prefill(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def fwd_decode(
|
|
||||||
self,
|
|
||||||
q: Tensor,
|
|
||||||
k: Tensor,
|
|
||||||
v: Tensor,
|
|
||||||
kv_cache: Optional[KVCache],
|
|
||||||
layer_id: int,
|
|
||||||
attn_mask: Optional[Tensor] = None,
|
|
||||||
is_causal: bool = False,
|
|
||||||
) -> Tensor:
|
|
||||||
"""Single-token decode with KV cache."""
|
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def fwd_prefill(
|
|
||||||
self,
|
|
||||||
q: Tensor,
|
|
||||||
k: Tensor,
|
|
||||||
v: Tensor,
|
|
||||||
kv_cache: Optional[KVCache],
|
|
||||||
layer_id: int,
|
|
||||||
attn_mask: Optional[Tensor] = None,
|
|
||||||
is_causal: bool = False,
|
|
||||||
) -> Tensor:
|
|
||||||
"""Multi-token prefill or training forward."""
|
|
||||||
|
|
||||||
|
|
||||||
class TorchNativeBackend(AttentionBackend):
|
|
||||||
"""Reference backend using torch SDPA with indirect KV cache indexing.
|
|
||||||
|
|
||||||
Writes new K/V into the cache buffers, gathers the full sequence K/V
|
|
||||||
via ``req_to_token`` indirect indexing, then calls
|
|
||||||
``F.scaled_dot_product_attention``.
|
|
||||||
|
|
||||||
For training (``kv_cache is None``), skips cache I/O entirely and
|
|
||||||
runs SDPA directly on the projected q/k/v.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def fwd_decode(
|
|
||||||
self,
|
|
||||||
q: Tensor,
|
|
||||||
k: Tensor,
|
|
||||||
v: Tensor,
|
|
||||||
kv_cache: Optional[KVCache],
|
|
||||||
layer_id: int,
|
|
||||||
attn_mask: Optional[Tensor] = None,
|
|
||||||
is_causal: bool = False,
|
|
||||||
) -> Tensor:
|
|
||||||
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
|
||||||
|
|
||||||
def fwd_prefill(
|
|
||||||
self,
|
|
||||||
q: Tensor,
|
|
||||||
k: Tensor,
|
|
||||||
v: Tensor,
|
|
||||||
kv_cache: Optional[KVCache],
|
|
||||||
layer_id: int,
|
|
||||||
attn_mask: Optional[Tensor] = None,
|
|
||||||
is_causal: bool = False,
|
|
||||||
) -> Tensor:
|
|
||||||
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
|
||||||
|
|
||||||
def _forward(
|
|
||||||
self,
|
|
||||||
q: Tensor,
|
|
||||||
k: Tensor,
|
|
||||||
v: Tensor,
|
|
||||||
kv_cache: Optional[KVCache],
|
|
||||||
layer_id: int,
|
|
||||||
attn_mask: Optional[Tensor] = None,
|
|
||||||
is_causal: bool = False,
|
|
||||||
) -> Tensor:
|
|
||||||
if kv_cache is not None:
|
|
||||||
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
|
||||||
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
|
||||||
|
|
||||||
max_len = kv_cache.max_len
|
|
||||||
if kv_cache.page_table is not None:
|
|
||||||
indices = kv_cache.page_table
|
|
||||||
else:
|
|
||||||
indices = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len]
|
|
||||||
if kv_cache.decode_mask is not None:
|
|
||||||
pos_mask = kv_cache.decode_mask
|
|
||||||
else:
|
|
||||||
pos_mask = (
|
|
||||||
torch.arange(max_len, device=q.device)[None, :]
|
|
||||||
< kv_cache.seq_lens[:, None]
|
|
||||||
)
|
|
||||||
indices = torch.where(pos_mask, indices, torch.zeros_like(indices))
|
|
||||||
k = kv_cache.k_buffer[layer_id, indices]
|
|
||||||
v = kv_cache.v_buffer[layer_id, indices]
|
|
||||||
|
|
||||||
n_rep = q.size(2) // k.size(2)
|
|
||||||
if n_rep > 1:
|
|
||||||
k = repeat_kv(k, n_rep)
|
|
||||||
v = repeat_kv(v, n_rep)
|
|
||||||
|
|
||||||
q = q.permute(0, 2, 1, 3)
|
|
||||||
k = k.permute(0, 2, 1, 3)
|
|
||||||
v = v.permute(0, 2, 1, 3)
|
|
||||||
|
|
||||||
out = F.scaled_dot_product_attention(q, k, v, attn_mask, is_causal=is_causal)
|
|
||||||
out = out.permute(0, 2, 1, 3).contiguous().flatten(2)
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
_default_backend = TorchNativeBackend()
|
|
||||||
|
|
||||||
|
|
||||||
class CudaBackend(AttentionBackend):
|
|
||||||
"""CUDA kernel backend with direct KV cache access.
|
|
||||||
|
|
||||||
Decode path: writes K/V to cache, then calls ``attn_paged_decode``
|
|
||||||
with ``page_size=1`` (each token slot is a single-token "page").
|
|
||||||
The ``req_to_token`` table serves directly as the page table.
|
|
||||||
|
|
||||||
Prefill path: writes K/V to cache, gathers full-sequence K/V via
|
|
||||||
indirect indexing (same as TorchNativeBackend), then calls
|
|
||||||
``attn_prefill``.
|
|
||||||
|
|
||||||
Training path (``kv_cache is None``): calls ``attn_prefill`` directly
|
|
||||||
on the projected q/k/v.
|
|
||||||
|
|
||||||
Falls back to ``TorchNativeBackend`` for any path where the
|
|
||||||
corresponding CUDA kernel is not available.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self):
|
|
||||||
self._fallback = TorchNativeBackend()
|
|
||||||
|
|
||||||
def fwd_decode(
|
|
||||||
self,
|
|
||||||
q: Tensor,
|
|
||||||
k: Tensor,
|
|
||||||
v: Tensor,
|
|
||||||
kv_cache: Optional[KVCache],
|
|
||||||
layer_id: int,
|
|
||||||
attn_mask: Optional[Tensor] = None,
|
|
||||||
is_causal: bool = False,
|
|
||||||
) -> Tensor:
|
|
||||||
if kv_cache is None or not is_available("attn_paged_decode"):
|
|
||||||
return self._fallback.fwd_decode(
|
|
||||||
q, k, v, kv_cache, layer_id, attn_mask, is_causal
|
|
||||||
)
|
|
||||||
|
|
||||||
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
|
||||||
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
|
||||||
|
|
||||||
max_len = kv_cache.max_len
|
|
||||||
|
|
||||||
if kv_cache.page_table is not None:
|
|
||||||
page_table = kv_cache.page_table
|
|
||||||
else:
|
|
||||||
page_table = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len]
|
|
||||||
|
|
||||||
k_cache = kv_cache.k_buffer[layer_id].unsqueeze(1)
|
|
||||||
v_cache = kv_cache.v_buffer[layer_id].unsqueeze(1)
|
|
||||||
|
|
||||||
if q.size(0) == 1:
|
|
||||||
mask = None
|
|
||||||
elif kv_cache.decode_mask is not None:
|
|
||||||
mask = kv_cache.decode_mask
|
|
||||||
else:
|
|
||||||
mask = (
|
|
||||||
torch.arange(max_len, device=q.device)[None, :]
|
|
||||||
< kv_cache.seq_lens[:, None]
|
|
||||||
)
|
|
||||||
|
|
||||||
out = attn_paged_decode(
|
|
||||||
q,
|
|
||||||
page_table,
|
|
||||||
k_cache,
|
|
||||||
v_cache,
|
|
||||||
page_size=1,
|
|
||||||
kv_len=max_len,
|
|
||||||
mask=mask,
|
|
||||||
is_causal=is_causal,
|
|
||||||
)
|
|
||||||
|
|
||||||
out = out.flatten(2)
|
|
||||||
return out
|
|
||||||
|
|
||||||
def fwd_prefill(
|
|
||||||
self,
|
|
||||||
q: Tensor,
|
|
||||||
k: Tensor,
|
|
||||||
v: Tensor,
|
|
||||||
kv_cache: Optional[KVCache],
|
|
||||||
layer_id: int,
|
|
||||||
attn_mask: Optional[Tensor] = None,
|
|
||||||
is_causal: bool = False,
|
|
||||||
) -> Tensor:
|
|
||||||
if kv_cache is None:
|
|
||||||
if is_available("attn_prefill"):
|
|
||||||
out = attn_prefill(q, k, v, mask=attn_mask, is_causal=is_causal)
|
|
||||||
return out.flatten(2)
|
|
||||||
return self._fallback.fwd_prefill(
|
|
||||||
q, k, v, kv_cache, layer_id, attn_mask, is_causal
|
|
||||||
)
|
|
||||||
|
|
||||||
if not is_available("attn_prefill"):
|
|
||||||
return self._fallback.fwd_prefill(
|
|
||||||
q, k, v, kv_cache, layer_id, attn_mask, is_causal
|
|
||||||
)
|
|
||||||
|
|
||||||
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
|
||||||
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
|
||||||
|
|
||||||
max_len = kv_cache.max_len
|
|
||||||
indices = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len]
|
|
||||||
pos_mask = (
|
|
||||||
torch.arange(max_len, device=q.device)[None, :] < kv_cache.seq_lens[:, None]
|
|
||||||
)
|
|
||||||
indices = torch.where(pos_mask, indices, torch.zeros_like(indices))
|
|
||||||
k_full = kv_cache.k_buffer[layer_id, indices]
|
|
||||||
v_full = kv_cache.v_buffer[layer_id, indices]
|
|
||||||
|
|
||||||
out = attn_prefill(q, k_full, v_full, mask=attn_mask, is_causal=is_causal)
|
|
||||||
return out.flatten(2)
|
|
||||||
|
|
||||||
|
|
||||||
_BACKEND_REGISTRY: dict[ATTN_BACKEND, type[AttentionBackend]] = {
|
|
||||||
ATTN_BACKEND.TORCH_NATIVE: TorchNativeBackend,
|
|
||||||
ATTN_BACKEND.CUDA: CudaBackend,
|
|
||||||
}
|
|
||||||
@@ -1,117 +0,0 @@
|
|||||||
"""Attention kernel wrapper functions — one entry point per compiled kernel.
|
|
||||||
|
|
||||||
Each wrapper calls its CUDA kernel directly. If the kernel is not
|
|
||||||
available, raises ``RuntimeError``. Fallback to torch SDPA is the
|
|
||||||
responsibility of the attention backend, not this module.
|
|
||||||
|
|
||||||
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
|
|
||||||
(blhd). Scale is always ``1/sqrt(head_dim)``.
|
|
||||||
|
|
||||||
Interface (all functions):
|
|
||||||
is_causal: True = causal mask; False = non-causal
|
|
||||||
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True=keep)
|
|
||||||
"""
|
|
||||||
|
|
||||||
import torch
|
|
||||||
|
|
||||||
from astrai.extension.loader import _available, _modules
|
|
||||||
|
|
||||||
|
|
||||||
def _check_available(name: str):
|
|
||||||
if not _available.get(name):
|
|
||||||
raise RuntimeError(
|
|
||||||
f"CUDA kernel '{name}' is not available. "
|
|
||||||
f"Build with CSRC_KERNELS=true or use a torch-native backend."
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def attn_decode(
|
|
||||||
q: torch.Tensor,
|
|
||||||
k: torch.Tensor,
|
|
||||||
v: torch.Tensor,
|
|
||||||
mask: torch.Tensor | None = None,
|
|
||||||
is_causal: bool = False,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
"""GQA decode attention (q_len == 1).
|
|
||||||
|
|
||||||
Args:
|
|
||||||
q: [batch, 1, n_heads, head_dim] (blhd, bf16)
|
|
||||||
k: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
|
||||||
v: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
|
||||||
mask: 2D [batch, kv_len] or 3D [batch, 1, kv_len] (bool, True=keep)
|
|
||||||
is_causal: apply causal mask
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
[batch, 1, n_heads, head_dim] (blhd, bf16)
|
|
||||||
"""
|
|
||||||
_check_available("attn_decode")
|
|
||||||
causal_offset = (k.size(1) - 1) if is_causal else -1
|
|
||||||
return _modules["attn_decode"].attn_decode(
|
|
||||||
q, k, v, mask=mask, causal_offset=causal_offset, layout=1
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def attn_prefill(
|
|
||||||
q: torch.Tensor,
|
|
||||||
k: torch.Tensor,
|
|
||||||
v: torch.Tensor,
|
|
||||||
mask: torch.Tensor | None = None,
|
|
||||||
is_causal: bool = False,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
"""GQA prefill attention (q_len > 1).
|
|
||||||
|
|
||||||
Args:
|
|
||||||
q: [batch, q_len, n_heads, head_dim] (blhd, bf16)
|
|
||||||
k: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
|
||||||
v: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
|
||||||
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True=keep)
|
|
||||||
is_causal: apply causal mask
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
[batch, q_len, n_heads, head_dim] (blhd, bf16)
|
|
||||||
"""
|
|
||||||
_check_available("attn_prefill")
|
|
||||||
causal_offset = (k.size(1) - q.size(1)) if is_causal else -1
|
|
||||||
return _modules["attn_prefill"].attn_prefill(
|
|
||||||
q, k, v, mask=mask, causal_offset=causal_offset, layout=1
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def attn_paged_decode(
|
|
||||||
q: torch.Tensor,
|
|
||||||
page_table: torch.Tensor,
|
|
||||||
k_cache: torch.Tensor,
|
|
||||||
v_cache: torch.Tensor,
|
|
||||||
page_size: int,
|
|
||||||
kv_len: int,
|
|
||||||
mask: torch.Tensor | None = None,
|
|
||||||
is_causal: bool = False,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
"""Paged GQA decode attention (q_len == 1, direct page-table access).
|
|
||||||
|
|
||||||
Args:
|
|
||||||
q: [batch, 1, n_heads, head_dim] (blhd, bf16)
|
|
||||||
page_table: [batch, max_pages] (int64)
|
|
||||||
k_cache: [n_pages, page_size, n_kv_heads, head_dim] (bf16)
|
|
||||||
v_cache: same as k_cache
|
|
||||||
page_size: tokens per page
|
|
||||||
kv_len: actual sequence length per request
|
|
||||||
mask: 2D [batch, kv_len] or 3D [batch, 1, kv_len] (bool, True=keep)
|
|
||||||
is_causal: apply causal mask
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
[batch, 1, n_heads, head_dim] (blhd, bf16)
|
|
||||||
"""
|
|
||||||
_check_available("attn_paged_decode")
|
|
||||||
causal_offset = (kv_len - 1) if is_causal else -1
|
|
||||||
return _modules["attn_paged_decode"].attn_paged_decode(
|
|
||||||
q,
|
|
||||||
page_table,
|
|
||||||
k_cache,
|
|
||||||
v_cache,
|
|
||||||
page_size,
|
|
||||||
kv_len,
|
|
||||||
mask=mask,
|
|
||||||
causal_offset=causal_offset,
|
|
||||||
layout=1,
|
|
||||||
)
|
|
||||||
@@ -0,0 +1,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,818 @@
|
|||||||
|
"""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. Backend resolution follows a strict precedence:
|
||||||
|
|
||||||
|
1. explicit ``attn_backend(...)`` context (wins over everything),
|
||||||
|
2. the process-wide ``ASTR_BACKEND`` environment override,
|
||||||
|
3. an implicit default picked from the available backends
|
||||||
|
(cuda > flash > torch).
|
||||||
|
|
||||||
|
Capability is polymorphic: every backend declares ``available()``
|
||||||
|
(machine-level) and ``supports_call(...)`` (per-call), so adding a new
|
||||||
|
backend requires no changes to the resolution logic. Training calls
|
||||||
|
(``fwd=None``, no KV cache) resolve through the same priority list: the
|
||||||
|
CUDA cache kernels cannot run without a cache, so they fall back to
|
||||||
|
flash (when it can handle the call — mask-free/causal only) and finally
|
||||||
|
to the reference ``TorchNativeBackend``.
|
||||||
|
|
||||||
|
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
|
||||||
|
(blhd). The backend returns ``[batch, seq_len, n_heads * head_dim]``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import contextvars
|
||||||
|
import enum
|
||||||
|
import functools
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import threading
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from typing import TYPE_CHECKING, Dict, Optional, Tuple, 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
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
_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)
|
||||||
|
)
|
||||||
|
|
||||||
|
# Backends are stateless — one canonical instance per class, created lazily
|
||||||
|
# and reused everywhere (resolution, fallback, context managers).
|
||||||
|
_singletons: Dict[type, "AttentionBackend"] = {}
|
||||||
|
|
||||||
|
|
||||||
|
@functools.lru_cache(maxsize=1)
|
||||||
|
def flash_attn_available() -> bool:
|
||||||
|
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 _instance(backend_cls: type) -> "AttentionBackend":
|
||||||
|
"""Return the canonical singleton instance for a backend class.
|
||||||
|
|
||||||
|
Backends hold no per-instance state, so a single cached instance is
|
||||||
|
safe and avoids per-call allocation on the attention hot path.
|
||||||
|
"""
|
||||||
|
backend = _singletons.get(backend_cls)
|
||||||
|
if backend is None:
|
||||||
|
backend = backend_cls()
|
||||||
|
_singletons[backend_cls] = backend
|
||||||
|
return backend
|
||||||
|
|
||||||
|
|
||||||
|
@functools.lru_cache(maxsize=1)
|
||||||
|
def _priority_backends() -> Tuple["AttentionBackend", ...]:
|
||||||
|
"""Available backends in priority order: cuda -> flash -> torch.
|
||||||
|
|
||||||
|
Computed once (machine availability cannot change at runtime) and
|
||||||
|
cached forever; the tuple always ends with ``TorchNativeBackend``,
|
||||||
|
which is unconditionally available.
|
||||||
|
"""
|
||||||
|
return tuple(
|
||||||
|
_instance(cls)
|
||||||
|
for cls in (CudaBackend, FlashAttnBackend, TorchNativeBackend)
|
||||||
|
if cls.available()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_default_backend() -> "AttentionBackend":
|
||||||
|
"""Pick the highest-priority available backend (cuda -> flash -> torch).
|
||||||
|
|
||||||
|
Resolved lazily on first use and cached via ``_priority_backends``.
|
||||||
|
Per-call capability fallback happens in ``attention()``, so the
|
||||||
|
default is safe for training and fp32 models.
|
||||||
|
"""
|
||||||
|
return _priority_backends()[0]
|
||||||
|
|
||||||
|
|
||||||
|
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 = _resolve_backend(name)
|
||||||
|
except (ValueError, RuntimeError):
|
||||||
|
_env_backend = None
|
||||||
|
logger.warning(
|
||||||
|
"ASTR_BACKEND=%r is not a registered attention backend; "
|
||||||
|
"falling back to default resolution",
|
||||||
|
name,
|
||||||
|
)
|
||||||
|
_env_backend_name = name
|
||||||
|
return _env_backend
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_backend(
|
||||||
|
backend: Optional[Union[str, ATTN_BACKEND, "AttentionBackend", type]] = None,
|
||||||
|
) -> "AttentionBackend":
|
||||||
|
"""Resolve a backend configuration to its canonical instance.
|
||||||
|
|
||||||
|
Accepts a registered name, ``ATTN_BACKEND`` enum value, backend class,
|
||||||
|
or instance. Names/classes resolve to the shared singleton; a caller
|
||||||
|
may still pass its own instance to opt out of sharing.
|
||||||
|
"""
|
||||||
|
if backend is not None:
|
||||||
|
if isinstance(backend, ATTN_BACKEND):
|
||||||
|
return _instance(AttentionBackendFactory.get_component_class(backend.value))
|
||||||
|
if isinstance(backend, str):
|
||||||
|
return _instance(AttentionBackendFactory.get_component_class(backend))
|
||||||
|
if isinstance(backend, type) and issubclass(backend, AttentionBackend):
|
||||||
|
return _instance(backend)
|
||||||
|
if isinstance(backend, AttentionBackend):
|
||||||
|
return backend
|
||||||
|
raise TypeError(
|
||||||
|
f"expected a registered name, ATTN_BACKEND, AttentionBackend type, "
|
||||||
|
f"or instance, got {type(backend).__name__}"
|
||||||
|
)
|
||||||
|
return _resolve_default_backend()
|
||||||
|
|
||||||
|
|
||||||
|
def get_backend(
|
||||||
|
use_default: bool = True,
|
||||||
|
) -> Optional["AttentionBackend"]:
|
||||||
|
"""Resolve the active backend: explicit context > env > default.
|
||||||
|
|
||||||
|
An ``attn_backend(...)`` context is the caller's explicit choice and
|
||||||
|
always wins. ``ASTR_BACKEND`` is a process-wide override consulted
|
||||||
|
only when no context is set. Pass ``use_default=False`` at request
|
||||||
|
submission to retain only an environment override or the caller's
|
||||||
|
:func:`attn_backend` value.
|
||||||
|
"""
|
||||||
|
context_backend = _current_backend.get()
|
||||||
|
if context_backend is not None:
|
||||||
|
return context_backend
|
||||||
|
env_backend = _environment_backend()
|
||||||
|
if env_backend is not None:
|
||||||
|
return env_backend
|
||||||
|
return _resolve_default_backend() if use_default else None
|
||||||
|
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
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 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,
|
||||||
|
backend: Optional[Union[str, ATTN_BACKEND, "AttentionBackend", type]] = None,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Functional attention entry point — mirrors ``F.scaled_dot_product_attention``.
|
||||||
|
|
||||||
|
Delegates to the active backend. ``backend`` (optional) is an explicit
|
||||||
|
escape hatch; when omitted the backend is resolved as
|
||||||
|
explicit context > ``ASTR_BACKEND`` env > default (cuda > flash > torch).
|
||||||
|
Handles KV cache I/O, GQA head expansion, and causal masking so the
|
||||||
|
caller only needs to provide projected q/k/v.
|
||||||
|
|
||||||
|
Training calls (``fwd=None``, ``kv_cache=None``) resolve through the
|
||||||
|
same capability chain — the CUDA cache kernels cannot run without a
|
||||||
|
cache, so they fall back to flash (mask-free/causal calls only) and
|
||||||
|
finally to torch SDPA. An explicitly-selected backend that cannot
|
||||||
|
handle the call raises — an implicit one falls back down the priority
|
||||||
|
list to the first capable backend.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
q: [batch, q_len, n_heads, head_dim] (blhd)
|
||||||
|
k: [batch, q_len, n_kv_heads, head_dim] (blhd)
|
||||||
|
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.
|
||||||
|
fwd: "prefill" / "decode" for inference, None for training.
|
||||||
|
backend: optional explicit backend (name, enum, class, or instance).
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[batch, q_len, n_heads * head_dim]
|
||||||
|
"""
|
||||||
|
if backend is not None:
|
||||||
|
selected = _resolve_backend(backend)
|
||||||
|
explicit = True
|
||||||
|
else:
|
||||||
|
context_backend = _current_backend.get()
|
||||||
|
explicit = context_backend is not None
|
||||||
|
# Resolve through the same chain as inference: explicit context >
|
||||||
|
# ASTR_BACKEND env > default. Training calls (fwd=None, no cache)
|
||||||
|
# land on the CUDA backend and fall back by capability below —
|
||||||
|
# flash when it can handle the call, else torch SDPA.
|
||||||
|
selected = get_backend()
|
||||||
|
assert selected is not None
|
||||||
|
|
||||||
|
if not selected.supports_call(q, kv_cache, attn_mask, is_causal, fwd):
|
||||||
|
if explicit:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Explicitly-set backend {type(selected).__name__} cannot "
|
||||||
|
f"handle this attention call (shape={q.shape}, "
|
||||||
|
f"dtype={q.dtype}, kv_cache={'none' if kv_cache is None else 'present'}, "
|
||||||
|
f"attn_mask={'none' if attn_mask is None else 'present'}). "
|
||||||
|
f"Remove the attn_backend() context or switch to a compatible backend."
|
||||||
|
)
|
||||||
|
selected = next(
|
||||||
|
(
|
||||||
|
candidate
|
||||||
|
for candidate in _priority_backends()
|
||||||
|
if candidate.supports_call(q, kv_cache, attn_mask, is_causal, fwd)
|
||||||
|
),
|
||||||
|
_instance(TorchNativeBackend),
|
||||||
|
)
|
||||||
|
return selected.forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal, fwd)
|
||||||
|
|
||||||
|
|
||||||
|
class AttentionBackend(ABC):
|
||||||
|
"""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.
|
||||||
|
|
||||||
|
Capability contract — every backend declares:
|
||||||
|
|
||||||
|
* ``available()`` — machine-level: can this backend exist here
|
||||||
|
(kernel ``.so`` loaded, flash-attn present, GPU available)?
|
||||||
|
Used once to build the default priority list.
|
||||||
|
* ``supports_call(q, kv_cache, attn_mask, is_causal, fwd)`` — can this
|
||||||
|
backend run this *specific* call (shape/dtype/cache/mask)? Used by
|
||||||
|
``attention()`` for the per-call fallback. Resolution logic never
|
||||||
|
checks concrete backend types, so adding a backend requires no
|
||||||
|
changes outside its own class.
|
||||||
|
|
||||||
|
Three equivalent ways to activate a backend::
|
||||||
|
|
||||||
|
with attn_backend(ATTN_BACKEND.TORCH_NATIVE): # enum
|
||||||
|
...
|
||||||
|
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)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
@abstractmethod
|
||||||
|
def available(cls) -> bool:
|
||||||
|
"""Return True if this backend can run on the current machine.
|
||||||
|
|
||||||
|
Checks static availability only (compiled kernels, optional
|
||||||
|
packages, GPU presence) — not call-specific constraints.
|
||||||
|
"""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def supports_call(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
kv_cache: Optional["KVCache"],
|
||||||
|
attn_mask: Optional[Tensor],
|
||||||
|
is_causal: bool,
|
||||||
|
fwd: Optional[str],
|
||||||
|
) -> bool:
|
||||||
|
"""Return True if this backend can run this specific attention call.
|
||||||
|
|
||||||
|
Called on the canonical singleton instance (or a caller-provided
|
||||||
|
one); must be side-effect free.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
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.
|
||||||
|
"""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def available(cls) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
def supports_call(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
kv_cache: Optional["KVCache"],
|
||||||
|
attn_mask: Optional[Tensor],
|
||||||
|
is_causal: bool,
|
||||||
|
fwd: Optional[str],
|
||||||
|
) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
def fwd_decode(
|
||||||
|
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.
|
||||||
|
"""
|
||||||
|
|
||||||
|
# Head dims supported by the CUDA kernels (single source of truth).
|
||||||
|
HEAD_DIMS = (32, 64, 128, 256)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def available(cls) -> bool:
|
||||||
|
return (
|
||||||
|
torch.cuda.is_available()
|
||||||
|
and is_available("attn_paged_decode")
|
||||||
|
and is_available("attn_paged_prefill")
|
||||||
|
)
|
||||||
|
|
||||||
|
def supports_call(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
kv_cache: Optional["KVCache"],
|
||||||
|
attn_mask: Optional[Tensor],
|
||||||
|
is_causal: bool,
|
||||||
|
fwd: Optional[str],
|
||||||
|
) -> bool:
|
||||||
|
# The CUDA kernels are bf16-only, support head_dim in
|
||||||
|
# HEAD_DIMS, and need a KV cache (decode/prefill); everything
|
||||||
|
# else falls back down the priority list to torch.
|
||||||
|
return (
|
||||||
|
fwd in ("prefill", "decode")
|
||||||
|
and kv_cache is not None
|
||||||
|
and q.ndim == 3
|
||||||
|
and q.dtype == torch.bfloat16
|
||||||
|
and q.size(-1) in self.HEAD_DIMS
|
||||||
|
and is_available(f"attn_paged_{fwd}")
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def supports_graph() -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
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)")
|
||||||
|
|
||||||
|
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,
|
||||||
|
new_k=k,
|
||||||
|
new_v=v,
|
||||||
|
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,
|
||||||
|
kv_cache.q_tile_to_batch,
|
||||||
|
kv_cache.q_tile_to_index,
|
||||||
|
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``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def available(cls) -> bool:
|
||||||
|
return flash_attn_available()
|
||||||
|
|
||||||
|
def supports_call(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
kv_cache: Optional["KVCache"],
|
||||||
|
attn_mask: Optional[Tensor],
|
||||||
|
is_causal: bool,
|
||||||
|
fwd: Optional[str],
|
||||||
|
) -> bool:
|
||||||
|
if not self.available():
|
||||||
|
return False
|
||||||
|
if q.dtype not in (torch.float16, torch.bfloat16):
|
||||||
|
return False
|
||||||
|
if fwd is not None:
|
||||||
|
return q.ndim == 3 and hasattr(_flash_attn, "flash_attn_varlen_func")
|
||||||
|
# Dense (training) path: flash_attn_func cannot apply a custom
|
||||||
|
# mask, so only mask-free calls are supported — ``is_causal`` is
|
||||||
|
# a flag, not a mask. Masked training (SFT/DPO/GRPO) must fall
|
||||||
|
# back to TorchNativeBackend instead of silently ignoring the mask.
|
||||||
|
return attn_mask is None
|
||||||
|
|
||||||
|
def fwd_decode(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
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:
|
||||||
|
raise ValueError(
|
||||||
|
"FlashAttnBackend cannot handle 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,
|
||||||
|
)
|
||||||
|
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
|
||||||
@@ -11,6 +11,7 @@ import torch
|
|||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from astrai.extension.loader import is_available
|
from astrai.extension.loader import is_available
|
||||||
|
from astrai.extension.ops.rotary import rotary_emb as _cuda_rotary
|
||||||
|
|
||||||
_cache = {"available": None}
|
_cache = {"available": None}
|
||||||
|
|
||||||
@@ -26,7 +27,7 @@ def _torch_apply(x: Tensor, freqs_cis: Tensor) -> Tensor:
|
|||||||
dtype = x.dtype
|
dtype = x.dtype
|
||||||
x_ = x.float().reshape(*x.shape[:-1], -1, 2)
|
x_ = x.float().reshape(*x.shape[:-1], -1, 2)
|
||||||
x_complex = torch.view_as_complex(x_)
|
x_complex = torch.view_as_complex(x_)
|
||||||
freqs_cis_complex = torch.complex(cos, sin).unsqueeze(2)
|
freqs_cis_complex = torch.complex(cos, sin).unsqueeze(-2)
|
||||||
x_rotated = x_complex * freqs_cis_complex
|
x_rotated = x_complex * freqs_cis_complex
|
||||||
x_out = torch.view_as_real(x_rotated).flatten(-2)
|
x_out = torch.view_as_real(x_rotated).flatten(-2)
|
||||||
return x_out.to(dtype)
|
return x_out.to(dtype)
|
||||||
@@ -48,7 +49,5 @@ def apply_rotary_emb(x: Tensor, freqs_cis: Tensor) -> Tensor:
|
|||||||
and x.is_cuda
|
and x.is_cuda
|
||||||
and x.dtype == torch.bfloat16
|
and x.dtype == torch.bfloat16
|
||||||
):
|
):
|
||||||
from astrai.extension.rotary_ops import rotary_emb as _cuda_rotary
|
|
||||||
|
|
||||||
return _cuda_rotary(x, freqs_cis)
|
return _cuda_rotary(x, freqs_cis)
|
||||||
return _torch_apply(x, freqs_cis)
|
return _torch_apply(x, freqs_cis)
|
||||||
@@ -0,0 +1,481 @@
|
|||||||
|
"""FP8 training: scaling recipes, per-tensor state, and aten::linear dispatch.
|
||||||
|
|
||||||
|
Layered (see ``ops/fp8.py`` for the CUDA interface adapter):
|
||||||
|
1. ``ops.fp8`` — the only module touching the pybind.
|
||||||
|
2. This module (strategy layer): scaling *recipes* (TE-style delayed scaling
|
||||||
|
or dynamic current-amax scaling), per-tensor scales + amax history, and the
|
||||||
|
``fp8_autocast`` context manager (like ``torch.autocast``).
|
||||||
|
3. aten::linear integration: registers the CUDA + AutogradCUDA impls.
|
||||||
|
|
||||||
|
Usage::
|
||||||
|
|
||||||
|
from astrai.extension.fp8 import fp8_autocast
|
||||||
|
with fp8_autocast(enabled=True, fp8_format="hybrid"):
|
||||||
|
logits = model(input_ids)
|
||||||
|
loss.backward() # fp8 backward runs anywhere; fwd captured state on the node
|
||||||
|
|
||||||
|
Format defaults follow the ecosystem consensus: E4M3 forward / E5M2 backward
|
||||||
|
("hybrid"); every operand's scale is a quantization step derived from its amax
|
||||||
|
history by the active recipe.
|
||||||
|
|
||||||
|
The context mirrors ``torch.autocast`` (``autocast_mode.py``): the active
|
||||||
|
``(enabled, recipe, fp8_format)`` triple is thread-local (a ``contextvars``
|
||||||
|
``ContextVar``, absent outside any region), and the manager is class-based and
|
||||||
|
reentrant with nested ``enabled=False`` disabling dispatch inside it. The module
|
||||||
|
targets *training*: every step quantizes x/w/g fresh (no weight-cast cache — the
|
||||||
|
optimizer bumps the weight version each step, so a torch-style cached_cast would
|
||||||
|
miss anyway), and the per-operand scales come from the delayed/dynamic recipe.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import functools
|
||||||
|
from contextvars import ContextVar, Token
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from enum import Enum
|
||||||
|
from typing import Dict, List, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch.library import Library
|
||||||
|
|
||||||
|
from astrai.extension.ops.fp8 import mm_fp8, quantize
|
||||||
|
|
||||||
|
# Max representable value per FP8 format (E4M3: 448, E5M2: 57344).
|
||||||
|
FP8_MAX = {"e4m3": 448.0, "e5m2": 57344.0}
|
||||||
|
|
||||||
|
|
||||||
|
class FP8Format(str, Enum):
|
||||||
|
"""Per-direction FP8 format. HYBRID = E4M3 forward / E5M2 backward."""
|
||||||
|
|
||||||
|
E4M3 = "e4m3"
|
||||||
|
E5M2 = "e5m2"
|
||||||
|
HYBRID = "hybrid"
|
||||||
|
|
||||||
|
def fwd(self) -> str:
|
||||||
|
return "e4m3" if self is FP8Format.HYBRID else self.value
|
||||||
|
|
||||||
|
def bwd(self) -> str:
|
||||||
|
return "e5m2" if self is FP8Format.HYBRID else self.value
|
||||||
|
|
||||||
|
|
||||||
|
class FP8Recipe:
|
||||||
|
"""Scale-from-amax policy: ``scale = (amax / FP8_MAX[fmt]) / 2^margin``.
|
||||||
|
|
||||||
|
``scale_from_history`` receives the operand's amax tensor (a ring window for
|
||||||
|
delayed scaling, the current amax for dynamic scaling) and returns the
|
||||||
|
quantization step. Subclasses set ``history_len`` / ``margin``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
history_len: int = 16
|
||||||
|
margin: int = 0
|
||||||
|
|
||||||
|
def scale_from_history(self, amax: torch.Tensor, fmt: str) -> torch.Tensor:
|
||||||
|
peak = amax.max()
|
||||||
|
return ((peak / FP8_MAX[fmt]) / (2**self.margin)).clamp_min(1e-12)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class DelayedScaling(FP8Recipe):
|
||||||
|
"""TE-style delayed scaling: max over the amax history window (amax from
|
||||||
|
*previous* steps; the window trades responsiveness against stability)."""
|
||||||
|
|
||||||
|
history_len: int = 16
|
||||||
|
margin: int = 0
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class DynamicScaling(FP8Recipe):
|
||||||
|
"""Current-amax scaling (torchao DYNAMIC): measure, then quantize. No
|
||||||
|
history — the scale is derived from the same-step amax, at an extra pass."""
|
||||||
|
|
||||||
|
history_len: int = 1
|
||||||
|
margin: int = 0
|
||||||
|
|
||||||
|
|
||||||
|
class _ScaleRing:
|
||||||
|
"""One operand's delayed-scaling state: a float32 buffer
|
||||||
|
``[hist[n] | scale | counter]`` (views). ``update`` folds the amax
|
||||||
|
returned by the quantize primitive into ``hist[idx]`` and publishes the
|
||||||
|
next scale from the window; ``idx`` advances host-side each step. The
|
||||||
|
trailing slot is a legacy counter kept for state-buffer compatibility.
|
||||||
|
"""
|
||||||
|
|
||||||
|
__slots__ = ("recipe", "state", "hist", "scale", "idx", "initialized")
|
||||||
|
|
||||||
|
def __init__(self, device: torch.device, recipe: FP8Recipe):
|
||||||
|
self.recipe = recipe
|
||||||
|
n = recipe.history_len
|
||||||
|
self.state = torch.zeros(n + 2, device=device, dtype=torch.float32)
|
||||||
|
self.hist = self.state[:n]
|
||||||
|
self.scale = self.state[n : n + 1]
|
||||||
|
self.idx = 0
|
||||||
|
self.initialized = False
|
||||||
|
|
||||||
|
def advance(self) -> None:
|
||||||
|
"""Rotate to the next history slot after metadata update."""
|
||||||
|
self.idx = (self.idx + 1) % self.hist.numel()
|
||||||
|
|
||||||
|
def seed(self, t: torch.Tensor, fmt: str) -> None:
|
||||||
|
amax = t.abs().amax().to(torch.float32).clamp_min(1e-12)
|
||||||
|
self.hist.fill_(amax)
|
||||||
|
self.scale.copy_(self.recipe.scale_from_history(self.hist, fmt))
|
||||||
|
self.initialized = True
|
||||||
|
|
||||||
|
def update(self, amax: torch.Tensor, fmt: str) -> None:
|
||||||
|
self.hist[self.idx].copy_(amax.reshape(()))
|
||||||
|
self.scale.copy_(self.recipe.scale_from_history(self.hist, fmt))
|
||||||
|
|
||||||
|
|
||||||
|
class FP8TensorMeta:
|
||||||
|
"""Per-weight delayed-scaling state for ``w``, ``x`` and ``g``.
|
||||||
|
|
||||||
|
DynamicScaling never allocates a meta; it measures the current amax inline.
|
||||||
|
"""
|
||||||
|
|
||||||
|
__slots__ = ("w", "x", "g")
|
||||||
|
|
||||||
|
def __init__(self, device: torch.device, recipe: FP8Recipe):
|
||||||
|
self.w = _ScaleRing(device, recipe)
|
||||||
|
self.x = _ScaleRing(device, recipe)
|
||||||
|
self.g = _ScaleRing(device, recipe)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class _ActiveConfig:
|
||||||
|
"""The immutable (enabled, recipe, format) triple of one open region."""
|
||||||
|
|
||||||
|
enabled: bool
|
||||||
|
recipe: FP8Recipe
|
||||||
|
fp8_format: FP8Format
|
||||||
|
|
||||||
|
|
||||||
|
# Thread-local active configuration (torch's autocast TLS analog): set by
|
||||||
|
# fp8_autocast on __enter__, absent outside any region. Autograd engine
|
||||||
|
# threads run backwards with their own empty context — fine, since backward
|
||||||
|
# only reads state captured on ctx at forward time.
|
||||||
|
_active_config: ContextVar[Optional[_ActiveConfig]] = ContextVar(
|
||||||
|
"astrai_fp8_active_config", default=None
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class FP8State:
|
||||||
|
"""Global fp8 training state: per-tensor metas + out-of-region defaults.
|
||||||
|
|
||||||
|
The active ``(enabled, recipe, fp8_format)`` triple is a ``ContextVar`` set
|
||||||
|
by ``fp8_autocast``. The properties below read that active config when a
|
||||||
|
region is open and the global defaults otherwise; the setters (and
|
||||||
|
``fp8_linear_enable``) write the global defaults — the persistent switch
|
||||||
|
applying outside any region. The metas registry is shared across threads
|
||||||
|
(GIL-protected); fp8 backward runs on autograd engine threads and only
|
||||||
|
touches metas captured on ``ctx`` at forward time.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self.default_enabled = False
|
||||||
|
self.default_recipe: FP8Recipe = DelayedScaling()
|
||||||
|
self.default_format: FP8Format = FP8Format.HYBRID
|
||||||
|
self._metas: Dict[tuple, FP8TensorMeta] = {}
|
||||||
|
|
||||||
|
# Active-config views (region config if open, else the defaults).
|
||||||
|
@property
|
||||||
|
def enabled(self) -> bool:
|
||||||
|
cfg = _active_config.get()
|
||||||
|
return cfg.enabled if cfg is not None else self.default_enabled
|
||||||
|
|
||||||
|
@property
|
||||||
|
def recipe(self) -> FP8Recipe:
|
||||||
|
cfg = _active_config.get()
|
||||||
|
return cfg.recipe if cfg is not None else self.default_recipe
|
||||||
|
|
||||||
|
@property
|
||||||
|
def fp8_format(self) -> FP8Format:
|
||||||
|
cfg = _active_config.get()
|
||||||
|
return cfg.fp8_format if cfg is not None else self.default_format
|
||||||
|
|
||||||
|
# Persistent (out-of-region) defaults.
|
||||||
|
@enabled.setter
|
||||||
|
def enabled(self, value: bool) -> None:
|
||||||
|
self.default_enabled = bool(value)
|
||||||
|
|
||||||
|
@recipe.setter
|
||||||
|
def recipe(self, value: FP8Recipe) -> None:
|
||||||
|
self.default_recipe = value
|
||||||
|
|
||||||
|
@fp8_format.setter
|
||||||
|
def fp8_format(self, value: FP8Format) -> None:
|
||||||
|
self.default_format = FP8Format(value)
|
||||||
|
|
||||||
|
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(w.device, self.recipe)
|
||||||
|
self._metas[key] = meta
|
||||||
|
return meta
|
||||||
|
|
||||||
|
def reset(self) -> None:
|
||||||
|
"""Restore construction defaults (switch, recipe, format) and drop all
|
||||||
|
per-weight metas — a full state reset for tests / reconfiguration."""
|
||||||
|
self.default_enabled = False
|
||||||
|
self.default_recipe = DelayedScaling()
|
||||||
|
self.default_format = FP8Format.HYBRID
|
||||||
|
self._metas.clear()
|
||||||
|
|
||||||
|
|
||||||
|
# Process-wide singleton; per-thread/per-region state lives in _active_config.
|
||||||
|
_state = FP8State()
|
||||||
|
|
||||||
|
|
||||||
|
def fp8_state() -> FP8State:
|
||||||
|
return _state
|
||||||
|
|
||||||
|
|
||||||
|
def _active() -> Optional[_ActiveConfig]:
|
||||||
|
"""The active config when fp8 dispatch is on, else ``None`` (fast guard).
|
||||||
|
|
||||||
|
A region config wins (honoring nested ``enabled=False`` regions); with no
|
||||||
|
region open this falls back to the persistent global switch
|
||||||
|
(``fp8_linear_enable``), so that flag still routes aten::linear to fp8.
|
||||||
|
"""
|
||||||
|
cfg = _active_config.get()
|
||||||
|
if cfg is not None:
|
||||||
|
return cfg if cfg.enabled else None
|
||||||
|
if _state.default_enabled:
|
||||||
|
return _ActiveConfig(True, _state.default_recipe, _state.default_format)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _current_config() -> _ActiveConfig:
|
||||||
|
"""Like ``_active()`` but always returns a config (disabled regions and
|
||||||
|
out-of-region direct calls resolve to the global defaults)."""
|
||||||
|
cfg = _active_config.get()
|
||||||
|
if cfg is not None:
|
||||||
|
return cfg
|
||||||
|
return _ActiveConfig(
|
||||||
|
_state.default_enabled, _state.default_recipe, _state.default_format
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class fp8_autocast:
|
||||||
|
"""Autocast-style context: fp8 linear dispatch on this thread.
|
||||||
|
|
||||||
|
Mirrors ``torch.autocast`` — a class-based, reentrant, nestable context
|
||||||
|
over thread-local state::
|
||||||
|
|
||||||
|
with fp8_autocast(enabled=True, fp8_format="hybrid"):
|
||||||
|
logits = model(input_ids) # aten::linear -> fp8 path
|
||||||
|
loss.backward() # fp8 backward; state was captured at forward time
|
||||||
|
|
||||||
|
Nesting follows torch: each ``__enter__`` pushes the new active config, each
|
||||||
|
``__exit__`` restores the previous one, and a nested ``enabled=False`` region
|
||||||
|
simply disables dispatch inside it. The instance doubles as a decorator.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
enabled: bool = True,
|
||||||
|
update_interval: int = 16,
|
||||||
|
recipe: Optional[FP8Recipe] = None,
|
||||||
|
fp8_format: str = "hybrid",
|
||||||
|
margin: int = 0,
|
||||||
|
):
|
||||||
|
if recipe is None:
|
||||||
|
recipe = DelayedScaling(history_len=update_interval, margin=margin)
|
||||||
|
self._config = _ActiveConfig(bool(enabled), recipe, FP8Format(fp8_format))
|
||||||
|
self._tokens: List[Token] = []
|
||||||
|
|
||||||
|
def __enter__(self) -> "fp8_autocast":
|
||||||
|
self._tokens.append(_active_config.set(self._config))
|
||||||
|
return self
|
||||||
|
|
||||||
|
def __exit__(self, exc_type, exc_val, exc_tb) -> bool:
|
||||||
|
token = self._tokens.pop()
|
||||||
|
_active_config.reset(token)
|
||||||
|
return False
|
||||||
|
|
||||||
|
def __call__(self, func):
|
||||||
|
@functools.wraps(func)
|
||||||
|
def decorate(*args, **kwargs):
|
||||||
|
with self:
|
||||||
|
return func(*args, **kwargs)
|
||||||
|
|
||||||
|
return decorate
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Strategy-level forward / backward (called from the aten::linear impl)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def _dynamic_scale(t: torch.Tensor, recipe: FP8Recipe, fmt: str) -> torch.Tensor:
|
||||||
|
amax = t.abs().amax().to(torch.float32).clamp_min(1e-12)
|
||||||
|
return recipe.scale_from_history(amax, fmt)
|
||||||
|
|
||||||
|
|
||||||
|
def _is_fp8(dtype: torch.dtype) -> bool:
|
||||||
|
"""A pre-quantized weight takes the GEMM directly (no re-quantize)."""
|
||||||
|
return dtype in (torch.float8_e4m3fn, torch.float8_e5m2)
|
||||||
|
|
||||||
|
|
||||||
|
def fp8_linear_forward(
|
||||||
|
x: torch.Tensor, w: torch.Tensor, bias=None, cfg: Optional[_ActiveConfig] = None
|
||||||
|
):
|
||||||
|
"""Scaled fp8 linear forward (called from the aten::linear impl).
|
||||||
|
|
||||||
|
Composed from the two stateless primitives: quantize x/w with the active
|
||||||
|
scales, run the pre-quantized GEMM with the bias fused into its epilogue.
|
||||||
|
Delayed scaling folds
|
||||||
|
the returned amax into the history ring and publishes the next scale;
|
||||||
|
dynamic scaling measures the current amax itself. Training quantizes the
|
||||||
|
weight every step (the optimizer bumps its version, so there is no cast
|
||||||
|
cache, matching ``cached_cast``-less behavior).
|
||||||
|
"""
|
||||||
|
state = fp8_state()
|
||||||
|
if cfg is None:
|
||||||
|
cfg = _current_config()
|
||||||
|
fmt = cfg.fp8_format.fwd()
|
||||||
|
if isinstance(cfg.recipe, DynamicScaling):
|
||||||
|
sx = _dynamic_scale(x.reshape(-1, w.size(1)), cfg.recipe, fmt)
|
||||||
|
sw = _dynamic_scale(w, cfg.recipe, fmt)
|
||||||
|
x8, _ = quantize(x, sx.reciprocal(), fmt)
|
||||||
|
w8 = w if _is_fp8(w.dtype) else quantize(w, sw.reciprocal(), fmt)[0]
|
||||||
|
# Bias fuses into the GEMM epilogue (fp32 add before the single bf16
|
||||||
|
# rounding — one rounding fewer than the separate out + bias pass);
|
||||||
|
# None passes through to the kernel's no-bias path.
|
||||||
|
out = mm_fp8(
|
||||||
|
x8.reshape(-1, x8.size(-1)), w8, sx * sw, trans_b=True, bias=bias
|
||||||
|
).reshape(*x.shape[:-1], w.size(0))
|
||||||
|
return out, sx, sw
|
||||||
|
|
||||||
|
meta = state.get_weight_meta(w)
|
||||||
|
if not meta.w.initialized:
|
||||||
|
meta.w.seed(w, fmt)
|
||||||
|
if not meta.x.initialized:
|
||||||
|
meta.x.seed(x, fmt)
|
||||||
|
sx, sw = meta.x.scale.clone(), meta.w.scale.clone()
|
||||||
|
x8, amax_x = quantize(x, sx.reciprocal(), fmt)
|
||||||
|
if _is_fp8(w.dtype):
|
||||||
|
w8, amax_w = w, None
|
||||||
|
else:
|
||||||
|
w8, amax_w = quantize(w, sw.reciprocal(), fmt)
|
||||||
|
out = mm_fp8(
|
||||||
|
x8.reshape(-1, x8.size(-1)), w8, sx * sw, trans_b=True, bias=bias
|
||||||
|
).reshape(*x.shape[:-1], w.size(0))
|
||||||
|
meta.x.update(amax_x, fmt)
|
||||||
|
if amax_w is not None:
|
||||||
|
meta.w.update(amax_w, fmt)
|
||||||
|
meta.x.advance()
|
||||||
|
if amax_w is not None:
|
||||||
|
meta.w.advance()
|
||||||
|
return out, sx, sw
|
||||||
|
|
||||||
|
|
||||||
|
class _LinearFp8(torch.autograd.Function):
|
||||||
|
"""The fp8 linear forward/backward pair (standard Function style).
|
||||||
|
|
||||||
|
The forward runs inside ``fp8_autocast`` and captures the active
|
||||||
|
fmt/recipe/meta on ``ctx``; the backward reads only that captured state, so
|
||||||
|
``loss.backward()`` may run after the context exits. The gradient is
|
||||||
|
quantized once (E5M2 in hybrid) and both dX/dW GEMMs share it; the output
|
||||||
|
masks come from ``needs_input_grad``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def forward(ctx, x, w, bias):
|
||||||
|
cfg = _current_config()
|
||||||
|
out, sx, sw = fp8_linear_forward(x, w, bias, cfg)
|
||||||
|
ctx.save_for_backward(x, w, sx, sw)
|
||||||
|
ctx.fmt_bwd = cfg.fp8_format.bwd()
|
||||||
|
ctx.recipe = cfg.recipe
|
||||||
|
ctx.is_dynamic = isinstance(cfg.recipe, DynamicScaling)
|
||||||
|
ctx.meta = None if ctx.is_dynamic else _state.get_weight_meta(w)
|
||||||
|
return out
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
@torch.autograd.function.once_differentiable
|
||||||
|
def backward(ctx, g):
|
||||||
|
x, w, _sx_fwd, _sw_fwd = ctx.saved_tensors
|
||||||
|
fmt = ctx.fmt_bwd
|
||||||
|
# Flatten leading dims (the forward GEMMs ran on [-1, N] / [-1, K]
|
||||||
|
# views; the kernels only accept 2D operands).
|
||||||
|
g2 = g.reshape(-1, g.size(-1))
|
||||||
|
if ctx.is_dynamic:
|
||||||
|
sg = _dynamic_scale(g2, ctx.recipe, fmt)
|
||||||
|
sw = _dynamic_scale(w, ctx.recipe, fmt)
|
||||||
|
sx = _dynamic_scale(x, ctx.recipe, fmt)
|
||||||
|
else:
|
||||||
|
meta = ctx.meta
|
||||||
|
if not meta.g.initialized:
|
||||||
|
meta.g.seed(g2, fmt)
|
||||||
|
sg = meta.g.scale.clone()
|
||||||
|
sw, sx = _sw_fwd, _sx_fwd
|
||||||
|
# Backward GEMMs route through the NT fast path via transposed
|
||||||
|
# quantize outputs: g8 [m,n] with w8T [k,n] (trans_b=True) gives
|
||||||
|
# grad_x, g8T [n,m] with x8T [k,m] gives grad_w — no NN-swap or TT
|
||||||
|
# crosswise kernel in the training path. g is consumed in both
|
||||||
|
# orientations, so one dual-layout pass feeds both.
|
||||||
|
g8, g8T, amax_g = quantize(g2, sg.reciprocal(), fmt, layout=2)
|
||||||
|
x8T, _ = quantize(x.reshape(-1, x.size(-1)), sx.reciprocal(), fmt, layout=1)
|
||||||
|
if _is_fp8(w.dtype):
|
||||||
|
# Pre-quantized weight has no transposed copy: keep the swap
|
||||||
|
# path for grad_x (grad_w is unaffected).
|
||||||
|
grad_x = mm_fp8(g8, w, sg * sw).reshape(x.shape)
|
||||||
|
else:
|
||||||
|
w8T, _ = quantize(w, sw.reciprocal(), fmt, layout=1)
|
||||||
|
grad_x = mm_fp8(g8, w8T, sg * sw, trans_b=True).reshape(x.shape)
|
||||||
|
grad_w = mm_fp8(g8T, x8T, sg * sx, trans_b=True) # g8.T @ x8
|
||||||
|
# bias-free linears must not pay the column-sum
|
||||||
|
# reduce: g2.sum(0) is another full read of the gradient.
|
||||||
|
grad_b = g2.sum(0).to(torch.bfloat16) if ctx.needs_input_grad[2] else None
|
||||||
|
if not ctx.is_dynamic:
|
||||||
|
meta.g.update(amax_g, fmt)
|
||||||
|
meta.g.advance()
|
||||||
|
return grad_x, grad_w, grad_b
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# aten::linear integration
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def fp8_linear_enable(enabled: bool = True) -> None:
|
||||||
|
"""Toggle fp8 dispatch for aten::linear globally (the out-of-region default;
|
||||||
|
``fp8_autocast`` regions override it thread-locally)."""
|
||||||
|
fp8_state().default_enabled = enabled
|
||||||
|
|
||||||
|
|
||||||
|
def fp8_linear_enabled() -> bool:
|
||||||
|
"""Whether fp8 dispatch is active right now (region config or global)."""
|
||||||
|
return _active() is not None
|
||||||
|
|
||||||
|
|
||||||
|
def _fp8_supported(x: torch.Tensor, w: torch.Tensor) -> bool:
|
||||||
|
"""Shape guard for the fp8 path. Unlike a strict 16-alignment requirement,
|
||||||
|
the kernels handle unaligned M/N via boundary checks (slower but correct) —
|
||||||
|
so no whole-call bf16 fallback for small decode batches. Only the K-dimension
|
||||||
|
contraction must match and the weight must be 2D."""
|
||||||
|
return x.dim() >= 2 and w.dim() == 2 and x.size(-1) == w.size(1)
|
||||||
|
|
||||||
|
|
||||||
|
def _linear_cuda_impl(x: torch.Tensor, w: torch.Tensor, bias=None):
|
||||||
|
if (
|
||||||
|
_active() is not None
|
||||||
|
and x.dtype is torch.bfloat16
|
||||||
|
and w.dtype is torch.bfloat16
|
||||||
|
and _fp8_supported(x, w)
|
||||||
|
):
|
||||||
|
return _LinearFp8.apply(x, w, bias)
|
||||||
|
return torch.ops.aten.linear.default.redispatch(
|
||||||
|
torch._C.DispatchKeySet(torch._C.DispatchKey.CompositeImplicitAutograd),
|
||||||
|
x,
|
||||||
|
w,
|
||||||
|
bias,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
_lib = Library("aten", "IMPL", "CUDA")
|
||||||
|
_lib.impl("linear", _linear_cuda_impl)
|
||||||
|
# Also replace torch's generated linear autograd formula (which would call
|
||||||
|
# aten::linear_backward after the fp8_autocast region exits). The fp8 backward
|
||||||
|
# is owned by _LinearFp8 with state captured at forward time, so loss.backward()
|
||||||
|
# works wherever it is called; the CUDA registration still covers inference_mode.
|
||||||
|
_lib_autograd = Library("aten", "IMPL", "AutogradCUDA")
|
||||||
|
_lib_autograd.impl("linear", _linear_cuda_impl)
|
||||||
+63
-16
@@ -1,36 +1,83 @@
|
|||||||
"""Dynamic discovery and loading of compiled CUDA kernel modules.
|
"""Dynamic discovery and loading of compiled CUDA kernel modules.
|
||||||
|
|
||||||
Each kernel is registered in ``csrc/build.py`` and built into a ``.so`` placed
|
Each kernel is built by the CMake build in ``csrc/CMakeLists.txt`` into a
|
||||||
in this package directory. On import we try to load each one; kernels that
|
``.so`` placed in ``astrai/extension/lib/`` — the module name equals the
|
||||||
failed to build (or are running on a CPU-only machine) are marked unavailable
|
``.so`` name equals the pybind name (e.g. ``attn_decode``, defined via
|
||||||
so the wrapper functions can fall back to ``torch`` SDPA.
|
``TORCH_EXTENSION_NAME``). ``KERNEL_NAMES`` is discovered automatically from
|
||||||
|
the ``.so`` files present, so adding a kernel to the CMake ``KERNELS``
|
||||||
|
registry needs no change here.
|
||||||
|
|
||||||
|
Loading is **lazy and centralized**: module names are discovered eagerly
|
||||||
|
(cheap glob), but each ``.so`` is imported on first use via the single
|
||||||
|
``get_module`` accessor, then cached. The wrapper modules (``ops/*.py``) never
|
||||||
|
touch the internals or keep their own caches — they call ``get_module(name)``
|
||||||
|
(or ``is_available(name)`` when a torch fallback is acceptable). A kernel that
|
||||||
|
failed to build (or is running on a CPU-only machine) is ``None`` in the cache,
|
||||||
|
so ``is_available`` returns ``False`` and ``get_module`` raises a clear error.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import glob
|
||||||
import importlib
|
import importlib
|
||||||
import logging
|
import logging
|
||||||
|
import os
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
KERNEL_NAMES = ["attn_decode", "attn_prefill", "attn_paged_decode", "rotary_emb"]
|
_LIB_DIR = os.path.join(os.path.dirname(__file__), "lib")
|
||||||
|
|
||||||
|
|
||||||
|
def _discover_kernel_names() -> list[str]:
|
||||||
|
"""Return the module names of the compiled kernel ``.so`` files in lib/."""
|
||||||
|
names: list[str] = []
|
||||||
|
for path in glob.glob(os.path.join(_LIB_DIR, "*.so")):
|
||||||
|
# strip the "<soabi>.so" suffix, e.g. attn_decode.cpython-312-...so
|
||||||
|
names.append(os.path.basename(path).split(".", 1)[0])
|
||||||
|
return sorted(names)
|
||||||
|
|
||||||
|
|
||||||
|
KERNEL_NAMES = _discover_kernel_names()
|
||||||
|
|
||||||
_available: dict[str, bool] = {}
|
_available: dict[str, bool] = {}
|
||||||
_modules: dict[str, object] = {}
|
_modules: dict[str, object] = {}
|
||||||
|
|
||||||
for _name in KERNEL_NAMES:
|
|
||||||
try:
|
def _try_load(name: str) -> object:
|
||||||
_mod = importlib.import_module(f".lib.{_name}", package=__package__)
|
"""Import and cache the ``name`` kernel module (lazy, one attempt).
|
||||||
_available[_name] = True
|
|
||||||
_modules[_name] = _mod
|
Returns the module, or ``None`` if it is unavailable. Cached so each
|
||||||
except ImportError:
|
``.so`` is imported at most once per process.
|
||||||
_available[_name] = False
|
"""
|
||||||
_modules[_name] = None
|
if name not in _modules:
|
||||||
|
try:
|
||||||
|
_modules[name] = importlib.import_module(
|
||||||
|
f".lib.{name}", package=__package__
|
||||||
|
)
|
||||||
|
_available[name] = True
|
||||||
|
except ImportError:
|
||||||
|
logger.warning("kernel '%s' failed to import; marking unavailable", name)
|
||||||
|
_modules[name] = None
|
||||||
|
_available[name] = False
|
||||||
|
return _modules[name]
|
||||||
|
|
||||||
|
|
||||||
def is_available(name: str) -> bool:
|
def is_available(name: str) -> bool:
|
||||||
"""Return ``True`` if the compiled kernel ``name`` was loaded."""
|
"""Return ``True`` if the compiled kernel ``name`` could be loaded."""
|
||||||
|
if name not in _available:
|
||||||
|
_try_load(name)
|
||||||
return _available.get(name, False)
|
return _available.get(name, False)
|
||||||
|
|
||||||
|
|
||||||
def get_module(name: str) -> object:
|
def get_module(name: str) -> object:
|
||||||
"""Return the loaded kernel module for ``name``, or ``None`` if unavailable."""
|
"""Return the loaded kernel module for ``name``, importing it on first use.
|
||||||
return _modules.get(name)
|
|
||||||
|
Raises ``RuntimeError`` if the kernel is unavailable (not built, or failed
|
||||||
|
to import) — callers that can tolerate a torch fallback should check
|
||||||
|
``is_available(name)`` first instead.
|
||||||
|
"""
|
||||||
|
mod = _try_load(name)
|
||||||
|
if mod is None:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"CUDA kernel '{name}' is not available. "
|
||||||
|
f"Build with CSRC_KERNELS=true (or use the torch-native fallback)."
|
||||||
|
)
|
||||||
|
return mod
|
||||||
|
|||||||
@@ -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,192 @@
|
|||||||
|
"""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 get_module
|
||||||
|
|
||||||
|
|
||||||
|
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 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)
|
||||||
|
"""
|
||||||
|
mod = get_module("attn_decode")
|
||||||
|
causal_offset = (k.size(1) - 1) if is_causal else -1
|
||||||
|
return mod.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)
|
||||||
|
"""
|
||||||
|
mod = get_module("attn_prefill")
|
||||||
|
causal_offset = (k.size(1) - q.size(1)) if is_causal else -1
|
||||||
|
return mod.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,
|
||||||
|
new_k: Optional[torch.Tensor] = None,
|
||||||
|
new_v: Optional[torch.Tensor] = None,
|
||||||
|
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
|
||||||
|
new_k: current-token K to append, [batch, n_kv_heads, head_dim]
|
||||||
|
new_v: current-token V to append, same shape as new_k
|
||||||
|
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)
|
||||||
|
"""
|
||||||
|
mod = get_module("attn_paged_decode")
|
||||||
|
causal_offset = 0 if is_causal else -1
|
||||||
|
return mod.attn_paged_decode(
|
||||||
|
q,
|
||||||
|
k_cache,
|
||||||
|
v_cache,
|
||||||
|
req_to_token,
|
||||||
|
req_pool_indices,
|
||||||
|
kv_indptr,
|
||||||
|
new_k=new_k,
|
||||||
|
new_v=new_v,
|
||||||
|
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,
|
||||||
|
q_tile_to_batch: torch.Tensor,
|
||||||
|
q_tile_to_index: 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
|
||||||
|
q_tile_to_batch: [num_q_tiles] (int32) — request index per Q tile
|
||||||
|
q_tile_to_index: [num_q_tiles] (int32) — local Q tile index per request
|
||||||
|
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)
|
||||||
|
"""
|
||||||
|
mod = get_module("attn_paged_prefill")
|
||||||
|
causal_offset = 0 if is_causal else -1
|
||||||
|
return mod.attn_paged_prefill(
|
||||||
|
q,
|
||||||
|
k_cache,
|
||||||
|
v_cache,
|
||||||
|
req_to_token,
|
||||||
|
req_pool_indices,
|
||||||
|
kv_indptr,
|
||||||
|
qo_indptr,
|
||||||
|
q_tile_to_batch,
|
||||||
|
q_tile_to_index,
|
||||||
|
mask,
|
||||||
|
causal_offset=causal_offset,
|
||||||
|
)
|
||||||
@@ -0,0 +1,251 @@
|
|||||||
|
"""FP8 CUDA kernel interface adapter (the only module touching the pybind).
|
||||||
|
|
||||||
|
Isolates the ``fp8_ops`` CUDA extension behind stable Python primitives:
|
||||||
|
|
||||||
|
- ``quantize(x, scale, fmt) -> (x8, amax)`` — BF16/FP16/FP32 → FP8 with fused amax
|
||||||
|
- ``mm_fp8(a8, b8, sa, sb) -> out`` — pre-quantized FP8 GEMM (BF16 output)
|
||||||
|
|
||||||
|
Scale semantics: scales are *quantization steps* — the value divided out when
|
||||||
|
quantizing (``x8 = x / scale``). Every primitive computes its own inverse
|
||||||
|
internally; callers never pass ``scale_inv``. ``amax`` values are *returned*,
|
||||||
|
never passed as output arguments. ``fmt`` is ``"e4m3"`` or ``"e5m2"``.
|
||||||
|
|
||||||
|
Policy (scales, amax history, delayed scaling, autocast) lives in ``fp8.py``;
|
||||||
|
this module is stateless.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Optional, Tuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch.library import custom_op
|
||||||
|
|
||||||
|
from astrai.extension.loader import get_module
|
||||||
|
|
||||||
|
# fmt string -> kernel int (0 = E4M3, 1 = E5M2)
|
||||||
|
_FMT_TO_INT = {"e4m3": 0, "e5m2": 1}
|
||||||
|
|
||||||
|
|
||||||
|
def _fmt_int(fmt: str) -> int:
|
||||||
|
try:
|
||||||
|
return _FMT_TO_INT[fmt]
|
||||||
|
except KeyError:
|
||||||
|
raise ValueError(f"unsupported fp8 format {fmt!r} (expected 'e4m3' or 'e5m2')")
|
||||||
|
|
||||||
|
|
||||||
|
def _fmt_name(fmt: int) -> str:
|
||||||
|
if fmt == 0:
|
||||||
|
return "e4m3"
|
||||||
|
if fmt == 1:
|
||||||
|
return "e5m2"
|
||||||
|
raise ValueError(f"unsupported quantization type {fmt!r}")
|
||||||
|
|
||||||
|
|
||||||
|
def _fmt_dtype(fmt: str) -> torch.dtype:
|
||||||
|
return torch.float8_e5m2 if _fmt_int(fmt) else torch.float8_e4m3fn
|
||||||
|
|
||||||
|
|
||||||
|
@custom_op("custom::fp8_quantize", mutates_args=())
|
||||||
|
def fp8_quantize(
|
||||||
|
x: torch.Tensor, scale: torch.Tensor, fmt: int
|
||||||
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
"""Float (bf16/fp16/fp32) -> FP8 quantize with fused amax; ``scale`` is a multiplier."""
|
||||||
|
|
||||||
|
|
||||||
|
@fp8_quantize.register_fake
|
||||||
|
def _fp8_quantize_fake(x, scale, fmt):
|
||||||
|
dtype = torch.float8_e5m2 if fmt == 1 else torch.float8_e4m3fn
|
||||||
|
return (
|
||||||
|
torch.empty(x.shape, device=x.device, dtype=dtype),
|
||||||
|
torch.empty(1, device=x.device, dtype=torch.float32),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
_QUANT_INPUT_DTYPES = (torch.bfloat16, torch.float16, torch.float32)
|
||||||
|
|
||||||
|
|
||||||
|
@custom_op("custom::fp8_quantize_t", mutates_args=())
|
||||||
|
def fp8_quantize_t(
|
||||||
|
x: torch.Tensor, scale: torch.Tensor, fmt: int
|
||||||
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
"""Transposed-output variant of fp8_quantize: returns ``(x8T, amax)``
|
||||||
|
where ``x8T`` is the [cols][rows] row-major transpose of the quantized
|
||||||
|
input (the K-contiguous operand orientation for NT GEMMs)."""
|
||||||
|
|
||||||
|
|
||||||
|
@fp8_quantize_t.register_fake
|
||||||
|
def _fp8_quantize_t_fake(x, scale, fmt):
|
||||||
|
dtype = torch.float8_e5m2 if fmt == 1 else torch.float8_e4m3fn
|
||||||
|
rows, cols = x.shape[-2], x.shape[-1]
|
||||||
|
return (
|
||||||
|
torch.empty((*x.shape[:-2], cols, rows), device=x.device, dtype=dtype),
|
||||||
|
torch.empty(1, device=x.device, dtype=torch.float32),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@fp8_quantize_t.register_kernel("cuda")
|
||||||
|
def _fp8_quantize_t_cuda(x, scale, fmt):
|
||||||
|
if x.dtype not in _QUANT_INPUT_DTYPES:
|
||||||
|
raise TypeError(f"fp8 quantize requires bf16/fp16/fp32 input, got {x.dtype}")
|
||||||
|
return get_module("fp8_ops").quantize(x, scale, int(fmt), 1)
|
||||||
|
|
||||||
|
|
||||||
|
@fp8_quantize_t.register_kernel("cpu")
|
||||||
|
def _fp8_quantize_t_cpu(x, scale, fmt):
|
||||||
|
x8, amax = _fp8_quantize_cpu(x, scale, fmt)
|
||||||
|
return x8.transpose(-2, -1).contiguous(), amax
|
||||||
|
|
||||||
|
|
||||||
|
@custom_op("custom::fp8_quantize_dual", mutates_args=())
|
||||||
|
def fp8_quantize_dual(
|
||||||
|
x: torch.Tensor, scale: torch.Tensor, fmt: int
|
||||||
|
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||||
|
"""Dual-orientation quantize: one read of ``x`` produces both the
|
||||||
|
row-major ``x8`` and its transposed ``x8T`` (plus ``amax``), for tensors
|
||||||
|
consumed by GEMMs on both orientations (backward ``g``)."""
|
||||||
|
|
||||||
|
|
||||||
|
@fp8_quantize_dual.register_fake
|
||||||
|
def _fp8_quantize_dual_fake(x, scale, fmt):
|
||||||
|
dtype = torch.float8_e5m2 if fmt == 1 else torch.float8_e4m3fn
|
||||||
|
rows, cols = x.shape[-2], x.shape[-1]
|
||||||
|
return (
|
||||||
|
torch.empty(x.shape, device=x.device, dtype=dtype),
|
||||||
|
torch.empty((*x.shape[:-2], cols, rows), device=x.device, dtype=dtype),
|
||||||
|
torch.empty(1, device=x.device, dtype=torch.float32),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@fp8_quantize_dual.register_kernel("cuda")
|
||||||
|
def _fp8_quantize_dual_cuda(x, scale, fmt):
|
||||||
|
if x.dtype not in _QUANT_INPUT_DTYPES:
|
||||||
|
raise TypeError(f"fp8 quantize requires bf16/fp16/fp32 input, got {x.dtype}")
|
||||||
|
return get_module("fp8_ops").quantize(x, scale, int(fmt), 2)
|
||||||
|
|
||||||
|
|
||||||
|
@fp8_quantize_dual.register_kernel("cpu")
|
||||||
|
def _fp8_quantize_dual_cpu(x, scale, fmt):
|
||||||
|
x8, amax = _fp8_quantize_cpu(x, scale, fmt)
|
||||||
|
return x8, x8.transpose(-2, -1).contiguous(), amax
|
||||||
|
|
||||||
|
|
||||||
|
@fp8_quantize.register_kernel("cuda")
|
||||||
|
def _fp8_quantize_cuda(x, scale, fmt):
|
||||||
|
if x.dtype not in _QUANT_INPUT_DTYPES:
|
||||||
|
raise TypeError(f"fp8 quantize requires bf16/fp16/fp32 input, got {x.dtype}")
|
||||||
|
return get_module("fp8_ops").quantize(x, scale, int(fmt))
|
||||||
|
|
||||||
|
|
||||||
|
@fp8_quantize.register_kernel("cpu")
|
||||||
|
def _fp8_quantize_cpu(x, scale, fmt):
|
||||||
|
x8 = (x.float() * scale).to(_fmt_dtype(_fmt_name(fmt)))
|
||||||
|
amax = x.abs().amax().float().reshape(1).clamp_min(1e-12)
|
||||||
|
return x8, amax
|
||||||
|
|
||||||
|
|
||||||
|
@custom_op("custom::fp8_gemm", mutates_args=())
|
||||||
|
def fp8_gemm(
|
||||||
|
a: torch.Tensor,
|
||||||
|
b: torch.Tensor,
|
||||||
|
scale: torch.Tensor,
|
||||||
|
trans_a: int = 0,
|
||||||
|
trans_b: int = 0,
|
||||||
|
bias: Optional[torch.Tensor] = None,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""FP8 GEMM: ``a @ b * scale (+ bias)`` with FP32 accumulation.
|
||||||
|
|
||||||
|
2D or 3D (batched) operands; a size-1 batch broadcasts (matmul rules).
|
||||||
|
``bias`` (bf16, length n) fuses into the epilogue in fp32 before the
|
||||||
|
single bf16 rounding. The result is always BF16; FP8 output is a
|
||||||
|
separate quantize operation.
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
@fp8_gemm.register_fake
|
||||||
|
def _fp8_gemm_fake(a, b, scale, trans_a=0, trans_b=0, bias=None):
|
||||||
|
dtype = torch.bfloat16
|
||||||
|
rows = a.size(2) if trans_a else a.size(1)
|
||||||
|
cols = b.size(1) if trans_b else b.size(2)
|
||||||
|
batches = [t.size(0) for t in (a, b) if t.dim() == 3]
|
||||||
|
shape = (max(batches), rows, cols) if batches else (rows, cols)
|
||||||
|
return torch.empty(shape, device=a.device, dtype=dtype)
|
||||||
|
|
||||||
|
|
||||||
|
@fp8_gemm.register_kernel("cuda")
|
||||||
|
def _fp8_gemm_cuda(a, b, scale, trans_a=0, trans_b=0, bias=None):
|
||||||
|
if a.dtype != b.dtype or a.dtype not in (torch.float8_e4m3fn, torch.float8_e5m2):
|
||||||
|
raise TypeError(
|
||||||
|
f"fp8 GEMM requires matching fp8 inputs, got {a.dtype}/{b.dtype}"
|
||||||
|
)
|
||||||
|
return get_module("fp8_ops").mm_fp8(a, b, scale, trans_a, trans_b, bias)
|
||||||
|
|
||||||
|
|
||||||
|
@fp8_gemm.register_kernel("cpu")
|
||||||
|
def _fp8_gemm_cpu(a, b, scale, trans_a=0, trans_b=0, bias=None):
|
||||||
|
aa = a.float().transpose(-2, -1) if trans_a else a.float()
|
||||||
|
bb = b.float().transpose(-2, -1) if trans_b else b.float()
|
||||||
|
acc = aa @ bb * scale
|
||||||
|
if bias is not None and bias.numel() > 0:
|
||||||
|
acc = acc + bias.float()
|
||||||
|
return acc.to(torch.bfloat16)
|
||||||
|
|
||||||
|
|
||||||
|
def quantize(
|
||||||
|
x: torch.Tensor, scale: torch.Tensor, fmt: str = "e4m3", layout: int = 0
|
||||||
|
) -> tuple:
|
||||||
|
"""Float (bf16/fp16/fp32) -> FP8 quantize with fused amax.
|
||||||
|
|
||||||
|
``scale`` is the quantization multiplier (device scalar); ``fmt`` selects
|
||||||
|
E4M3 or E5M2. ``amax`` is a fresh 1-element float32 tensor. ``layout``
|
||||||
|
picks the output orientation: 0 = row-major ``(x8, amax)``; 1 =
|
||||||
|
transposed ``[cols][rows]`` ``(x8T, amax)`` — the K-contiguous operand
|
||||||
|
orientation NT GEMMs want; 2 = both from one read ``(x8, x8T, amax)``
|
||||||
|
(for tensors consumed in both orientations, e.g. backward ``g``).
|
||||||
|
"""
|
||||||
|
# Hot-path bypass of the torch.library dispatch (~5us/call, ~40% of a
|
||||||
|
# 512-wide GEMM): real CUDA tensors of a supported dtype go straight to
|
||||||
|
# the extension. Fake/subclass tensors and non-CUDA inputs keep the
|
||||||
|
# custom_op route so torch.compile / meta / fake-tensor tracing and the
|
||||||
|
# CPU fallback behave exactly as before.
|
||||||
|
if (
|
||||||
|
type(x) is torch.Tensor
|
||||||
|
and x.is_cuda
|
||||||
|
and x.dtype in _QUANT_INPUT_DTYPES
|
||||||
|
and fmt in _FMT_TO_INT
|
||||||
|
):
|
||||||
|
return get_module("fp8_ops").quantize(x, scale, _FMT_TO_INT[fmt], layout)
|
||||||
|
if layout == 0:
|
||||||
|
return fp8_quantize(x, scale, _fmt_int(fmt))
|
||||||
|
if layout == 1:
|
||||||
|
return fp8_quantize_t(x, scale, _fmt_int(fmt))
|
||||||
|
return fp8_quantize_dual(x, scale, _fmt_int(fmt))
|
||||||
|
|
||||||
|
|
||||||
|
def mm_fp8(
|
||||||
|
a: torch.Tensor,
|
||||||
|
b: torch.Tensor,
|
||||||
|
scale: torch.Tensor,
|
||||||
|
trans_a: bool = False,
|
||||||
|
trans_b: bool = False,
|
||||||
|
bias: Optional[torch.Tensor] = None,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Pre-quantized FP8 GEMM: ``a @ b * scale (+ bias)``.
|
||||||
|
|
||||||
|
``a``/``b`` must be FP8 tensors of the same format, 2D or 3D (batched,
|
||||||
|
matmul-style broadcast on the batch dim). Inner-transposed views (e.g.
|
||||||
|
``x.t()``) fold into the layout at zero copy. ``scale`` is their combined
|
||||||
|
dequantization scale. ``bias`` (CUDA bf16 1D of length n) adds inside the
|
||||||
|
kernel epilogue in fp32 — no separate elementwise pass. The result is
|
||||||
|
BF16; FP8 output is a separate quantize operation.
|
||||||
|
"""
|
||||||
|
# Same hot-path bypass as quantize(): the binding's TORCH_CHECKs keep
|
||||||
|
# validation identical on the direct route (bias may be None — the
|
||||||
|
# binding resolves it to the no-bias path).
|
||||||
|
if (
|
||||||
|
type(a) is torch.Tensor
|
||||||
|
and a.is_cuda
|
||||||
|
and a.dtype in (torch.float8_e4m3fn, torch.float8_e5m2)
|
||||||
|
):
|
||||||
|
return get_module("fp8_ops").mm_fp8(
|
||||||
|
a, b, scale, int(trans_a), int(trans_b), bias
|
||||||
|
)
|
||||||
|
return fp8_gemm(a, b, scale, trans_a, trans_b, bias)
|
||||||
@@ -0,0 +1,31 @@
|
|||||||
|
"""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 get_module
|
||||||
|
|
||||||
|
|
||||||
|
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``.
|
||||||
|
"""
|
||||||
|
mod = get_module("rotary_emb")
|
||||||
|
if not x.is_contiguous():
|
||||||
|
x = x.contiguous()
|
||||||
|
if not freqs_cis.is_contiguous():
|
||||||
|
freqs_cis = freqs_cis.contiguous()
|
||||||
|
return mod.rotary_emb(x, freqs_cis)
|
||||||
@@ -1,39 +0,0 @@
|
|||||||
"""Rotary embedding CUDA kernel wrapper.
|
|
||||||
|
|
||||||
Calls the compiled CUDA kernel directly. If the kernel is not available,
|
|
||||||
raises ``RuntimeError``. Fallback to torch complex multiply is the
|
|
||||||
responsibility of ``astrai.extension.rotary_backend.apply_rotary_emb``.
|
|
||||||
|
|
||||||
Layout: x is [batch, seq_len, n_heads, head_dim] (bf16, contiguous).
|
|
||||||
freqs_cis is [batch, seq_len, head_dim/2, 2] (f32, contiguous) — [cos, sin] pairs.
|
|
||||||
"""
|
|
||||||
|
|
||||||
import torch
|
|
||||||
|
|
||||||
from astrai.extension.loader import _available, _modules
|
|
||||||
|
|
||||||
|
|
||||||
def _check_available():
|
|
||||||
if not _available.get("rotary_emb"):
|
|
||||||
raise RuntimeError(
|
|
||||||
"CUDA kernel 'rotary_emb' is not available. "
|
|
||||||
"Build with CSRC_KERNELS=true or use the torch fallback."
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def rotary_emb(x: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor:
|
|
||||||
"""Fused rotary embedding kernel.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
x: [batch, seq_len, n_heads, head_dim] (bf16, contiguous)
|
|
||||||
freqs_cis: [batch, seq_len, head_dim/2, 2] (f32, contiguous) — [cos, sin] pairs
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
[batch, seq_len, n_heads, head_dim] (bf16)
|
|
||||||
"""
|
|
||||||
_check_available()
|
|
||||||
if not x.is_contiguous():
|
|
||||||
x = x.contiguous()
|
|
||||||
if not freqs_cis.is_contiguous():
|
|
||||||
freqs_cis = freqs_cis.contiguous()
|
|
||||||
return _modules["rotary_emb"].rotary_emb(x, freqs_cis)
|
|
||||||
@@ -1,95 +1,33 @@
|
|||||||
"""Inference module for continuous batching.
|
"""Inference module for continuous batching.
|
||||||
|
|
||||||
Layers:
|
Subpackages:
|
||||||
- core/: Core inference loop (cache, executor, scheduler, task)
|
- cache/: KV cache (buffers, strategies, pool)
|
||||||
- api/: HTTP orchestration (ProtocolHandler, server)
|
- runtime/: Execution + sampling (executor, CUDA graph, sampling strategies)
|
||||||
- protocols/: Response builders (OpenAI, Anthropic)
|
- task/: Request lifecycle + performance metrics
|
||||||
- transport/: SSE transport utilities
|
- network/: HTTP protocol handling (server, protocol, OpenAI/Anthropic builders)
|
||||||
- engine.py: Facade (InferenceEngine), Value Object (GenerationRequest)
|
|
||||||
- sample.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy, FrequencyPenaltyStrategy)
|
Modules:
|
||||||
|
- scheduler.py: Continuous batching loop
|
||||||
|
- workspace.py: Pre-allocated GPU buffers
|
||||||
|
- engine.py: Facade (InferenceEngine)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from astrai.inference.api import (
|
from astrai.inference.engine import InferenceEngine
|
||||||
AnthropicMessage,
|
from astrai.inference.network import get_app, run_server
|
||||||
BaseToolParser,
|
from astrai.inference.runtime.executor import Executor
|
||||||
ChatCompletionRequest,
|
from astrai.inference.runtime.sample import sample
|
||||||
ChatMessage,
|
from astrai.inference.scheduler import InferenceScheduler
|
||||||
FunctionDef,
|
from astrai.inference.task import STOP, Task, TaskManager, TaskStatus
|
||||||
GenContext,
|
|
||||||
MessagesRequest,
|
|
||||||
ProtocolHandler,
|
|
||||||
SimpleJsonToolParser,
|
|
||||||
StopChecker,
|
|
||||||
ToolDef,
|
|
||||||
ToolParserFactory,
|
|
||||||
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,
|
|
||||||
Executor,
|
|
||||||
InferenceScheduler,
|
|
||||||
KVCache,
|
|
||||||
KVStorage,
|
|
||||||
PagePool,
|
|
||||||
PrefixCache,
|
|
||||||
ReqToTokenPool,
|
|
||||||
Task,
|
|
||||||
TaskManager,
|
|
||||||
TaskStatus,
|
|
||||||
page_hash,
|
|
||||||
)
|
|
||||||
from astrai.inference.engine import GenerationRequest, InferenceEngine
|
|
||||||
from astrai.inference.sample import (
|
|
||||||
BaseSamplingStrategy,
|
|
||||||
FrequencyPenaltyStrategy,
|
|
||||||
SamplingPipeline,
|
|
||||||
TemperatureStrategy,
|
|
||||||
TopKStrategy,
|
|
||||||
TopPStrategy,
|
|
||||||
sample,
|
|
||||||
)
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"InferenceEngine",
|
"InferenceEngine",
|
||||||
"GenerationRequest",
|
|
||||||
"InferenceScheduler",
|
"InferenceScheduler",
|
||||||
"Executor",
|
"Executor",
|
||||||
"STOP",
|
"STOP",
|
||||||
"Task",
|
"Task",
|
||||||
"TaskManager",
|
"TaskManager",
|
||||||
"TaskStatus",
|
"TaskStatus",
|
||||||
"Allocator",
|
|
||||||
"KVCache",
|
|
||||||
"KVStorage",
|
|
||||||
"PagePool",
|
|
||||||
"PrefixCache",
|
|
||||||
"ReqToTokenPool",
|
|
||||||
"page_hash",
|
|
||||||
"sample",
|
"sample",
|
||||||
"BaseSamplingStrategy",
|
|
||||||
"TemperatureStrategy",
|
|
||||||
"TopKStrategy",
|
|
||||||
"TopPStrategy",
|
|
||||||
"FrequencyPenaltyStrategy",
|
|
||||||
"SamplingPipeline",
|
|
||||||
"ProtocolHandler",
|
|
||||||
"StopChecker",
|
|
||||||
"GenContext",
|
|
||||||
"BaseToolParser",
|
|
||||||
"SimpleJsonToolParser",
|
|
||||||
"ToolParserFactory",
|
|
||||||
"OpenAIResponseBuilder",
|
|
||||||
"AnthropicResponseBuilder",
|
|
||||||
"ChatMessage",
|
|
||||||
"ChatCompletionRequest",
|
|
||||||
"FunctionDef",
|
|
||||||
"ToolDef",
|
|
||||||
"AnthropicMessage",
|
|
||||||
"MessagesRequest",
|
|
||||||
"get_app",
|
"get_app",
|
||||||
"run_server",
|
"run_server",
|
||||||
]
|
]
|
||||||
|
|||||||
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
+96
@@ -0,0 +1,96 @@
|
|||||||
|
"""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
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@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
|
||||||
|
q_tile_to_batch: Optional[Tensor] = None
|
||||||
|
q_tile_to_index: Optional[Tensor] = None
|
||||||
|
decode_o_part: Optional[Tensor] = None
|
||||||
|
decode_ml_part: Optional[Tensor] = None
|
||||||
|
decode_out: Optional[Tensor] = None
|
||||||
Vendored
+382
@@ -0,0 +1,382 @@
|
|||||||
|
"""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 Q_TILE_ROWS, 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]
|
||||||
|
tile_batches = []
|
||||||
|
tile_indices = []
|
||||||
|
for batch, q_len in enumerate(q_lens):
|
||||||
|
n_tiles = (q_len + Q_TILE_ROWS - 1) // Q_TILE_ROWS
|
||||||
|
tile_batches.extend([batch] * n_tiles)
|
||||||
|
tile_indices.extend(range(n_tiles))
|
||||||
|
n_tiles = len(tile_batches)
|
||||||
|
workspace.q_tile_to_batch[:n_tiles].copy_(
|
||||||
|
torch.tensor(tile_batches, dtype=torch.int32, device=device)
|
||||||
|
)
|
||||||
|
workspace.q_tile_to_index[:n_tiles].copy_(
|
||||||
|
torch.tensor(tile_indices, dtype=torch.int32, device=device)
|
||||||
|
)
|
||||||
|
q_tile_to_batch = workspace.q_tile_to_batch[:n_tiles]
|
||||||
|
q_tile_to_index = workspace.q_tile_to_index[:n_tiles]
|
||||||
|
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]
|
||||||
|
q_tile_to_batch = q_tile_to_index = None
|
||||||
|
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,
|
||||||
|
q_tile_to_batch=q_tile_to_batch,
|
||||||
|
q_tile_to_index=q_tile_to_index,
|
||||||
|
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
+318
@@ -0,0 +1,318 @@
|
|||||||
|
"""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
|
||||||
|
|
||||||
|
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,30 +0,0 @@
|
|||||||
"""Inference core: cache, executor, scheduler, task management."""
|
|
||||||
|
|
||||||
from astrai.inference.core.cache import (
|
|
||||||
Allocator,
|
|
||||||
KVCache,
|
|
||||||
KVStorage,
|
|
||||||
PagePool,
|
|
||||||
PrefixCache,
|
|
||||||
ReqToTokenPool,
|
|
||||||
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",
|
|
||||||
"KVCache",
|
|
||||||
"KVStorage",
|
|
||||||
"PagePool",
|
|
||||||
"PrefixCache",
|
|
||||||
"ReqToTokenPool",
|
|
||||||
"page_hash",
|
|
||||||
"Executor",
|
|
||||||
"InferenceScheduler",
|
|
||||||
"STOP",
|
|
||||||
"Task",
|
|
||||||
"TaskManager",
|
|
||||||
"TaskStatus",
|
|
||||||
]
|
|
||||||
@@ -1,501 +0,0 @@
|
|||||||
"""KV cache architecture: three-layer separation (SGLang-inspired).
|
|
||||||
|
|
||||||
Layer 1 — KVStorage: flat token-level K/V buffers [n_layers, size, H, D]
|
|
||||||
Layer 2 — ReqToTokenPool: index table [req_idx, pos] → physical token slot
|
|
||||||
Layer 3 — Allocator: slot/page allocation with ref-counting and LRU
|
|
||||||
|
|
||||||
PagePool orchestrates all three plus PrefixCache (content addressing).
|
|
||||||
KVCache is a pure dataclass passed to the model for direct buffer access.
|
|
||||||
|
|
||||||
Two modes:
|
|
||||||
- contiguous (default): pre-allocated per-request blocks, no dynamic alloc
|
|
||||||
- paged: shared pool with on-demand allocation, prefix caching support
|
|
||||||
"""
|
|
||||||
|
|
||||||
import threading
|
|
||||||
from collections import OrderedDict
|
|
||||||
from dataclasses import dataclass
|
|
||||||
from typing import Callable, Dict, List, Optional
|
|
||||||
|
|
||||||
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 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.long, 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.
|
|
||||||
|
|
||||||
Attributes:
|
|
||||||
k_buffer: [n_layers, size, n_kv_heads, head_dim]
|
|
||||||
v_buffer: [n_layers, size, n_kv_heads, head_dim]
|
|
||||||
req_to_token: [num_reqs, max_ctx_len] — index table
|
|
||||||
req_pool_indices: [batch_size] — row indices into req_to_token
|
|
||||||
seq_lens: [batch_size] — per-request total sequence lengths
|
|
||||||
out_cache_loc: [batch, new_seq_len] or [batch, 1] — write indices
|
|
||||||
max_len: max(seq_lens) as Python int — avoids GPU sync in decode
|
|
||||||
page_table: [batch, max_len] — precomputed gather indices for decode;
|
|
||||||
None for prefill or when not yet computed.
|
|
||||||
decode_mask: [batch, max_len] bool — precomputed position validity
|
|
||||||
mask for decode; None for prefill or single-batch decode.
|
|
||||||
"""
|
|
||||||
|
|
||||||
k_buffer: Tensor
|
|
||||||
v_buffer: Tensor
|
|
||||||
req_to_token: Tensor
|
|
||||||
req_pool_indices: Tensor
|
|
||||||
seq_lens: Tensor
|
|
||||||
out_cache_loc: Tensor
|
|
||||||
max_len: int = 0
|
|
||||||
page_table: Optional[Tensor] = None
|
|
||||||
decode_mask: Optional[Tensor] = None
|
|
||||||
|
|
||||||
|
|
||||||
class PagePool:
|
|
||||||
"""Top-level KV cache manager.
|
|
||||||
|
|
||||||
Combines KVStorage + ReqToTokenPool + Allocator + PrefixCache.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
n_layers: Number of transformer layers.
|
|
||||||
n_kv_heads: Number of KV attention heads.
|
|
||||||
head_dim: Dimension per head.
|
|
||||||
max_batch_size: Maximum concurrent requests.
|
|
||||||
max_seq_len: Maximum sequence length per request.
|
|
||||||
device, dtype: Tensor device and dtype.
|
|
||||||
page_size: Page size for paged mode (1 = token-level).
|
|
||||||
n_tokens: Total token slots for paged mode. None = contiguous mode
|
|
||||||
(pre-allocates max_batch_size * max_seq_len).
|
|
||||||
"""
|
|
||||||
|
|
||||||
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
|
|
||||||
if self.contiguous:
|
|
||||||
self.n_tokens = max_batch_size * max_seq_len
|
|
||||||
else:
|
|
||||||
self.n_tokens = n_tokens
|
|
||||||
|
|
||||||
self._storage = KVStorage(
|
|
||||||
self.n_tokens, n_layers, n_kv_heads, head_dim, device, dtype
|
|
||||||
)
|
|
||||||
self._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, device=device
|
|
||||||
)
|
|
||||||
self._alloc: Optional[Allocator] = None
|
|
||||||
self._prefix: Optional[PrefixCache] = None
|
|
||||||
else:
|
|
||||||
n_pages = self.n_tokens // page_size
|
|
||||||
self._alloc = Allocator(n_pages)
|
|
||||||
self._prefix = PrefixCache(page_size) if page_size > 1 else None
|
|
||||||
if self._prefix is not None:
|
|
||||||
self._alloc.on_evict = self._prefix.evict
|
|
||||||
|
|
||||||
self._task_req: Dict[str, int] = {}
|
|
||||||
self._task_len: Dict[int, int] = {}
|
|
||||||
self._task_cached: Dict[str, int] = {}
|
|
||||||
self._task_slots: Dict[str, List[int]] = {}
|
|
||||||
self._task_pages: Dict[str, List[int]] = {}
|
|
||||||
self._lock = threading.Lock()
|
|
||||||
|
|
||||||
# ---- task lifecycle ----
|
|
||||||
|
|
||||||
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
|
|
||||||
req_slots = self._req_pool.alloc(1)
|
|
||||||
if req_slots is None:
|
|
||||||
return False
|
|
||||||
req_idx = req_slots[0]
|
|
||||||
self._task_req[task_id] = req_idx
|
|
||||||
|
|
||||||
if self.contiguous:
|
|
||||||
self._task_len[req_idx] = len(prompt_ids)
|
|
||||||
self._task_cached[task_id] = 0
|
|
||||||
return True
|
|
||||||
|
|
||||||
n_tokens_needed = len(prompt_ids)
|
|
||||||
cached = 0
|
|
||||||
|
|
||||||
if self._prefix is not None:
|
|
||||||
hits = self._prefix.lookup(prompt_ids)
|
|
||||||
cached = len(hits) * self.page_size
|
|
||||||
for p in hits:
|
|
||||||
self._alloc.inc_ref(p)
|
|
||||||
self._task_pages[task_id] = list(hits)
|
|
||||||
self._task_slots[task_id] = []
|
|
||||||
else:
|
|
||||||
self._task_pages[task_id] = []
|
|
||||||
self._task_slots[task_id] = []
|
|
||||||
|
|
||||||
remaining = n_tokens_needed - cached
|
|
||||||
if remaining > 0:
|
|
||||||
if self.page_size == 1:
|
|
||||||
slots = self._alloc_tokens(remaining)
|
|
||||||
if slots is None:
|
|
||||||
for p in self._task_pages[task_id]:
|
|
||||||
self._alloc.free(p)
|
|
||||||
self._req_pool.free([req_idx])
|
|
||||||
del self._task_req[task_id]
|
|
||||||
return False
|
|
||||||
self._task_slots[task_id] = slots
|
|
||||||
else:
|
|
||||||
n_new_pages = (remaining + self.page_size - 1) // self.page_size
|
|
||||||
new_pages = []
|
|
||||||
for _ in range(n_new_pages):
|
|
||||||
p = self._alloc.alloc()
|
|
||||||
if p < 0:
|
|
||||||
for hp in self._task_pages[task_id]:
|
|
||||||
self._alloc.free(hp)
|
|
||||||
for np_ in new_pages:
|
|
||||||
self._alloc.free(np_)
|
|
||||||
self._req_pool.free([req_idx])
|
|
||||||
del self._task_req[task_id]
|
|
||||||
return False
|
|
||||||
new_pages.append(p)
|
|
||||||
self._task_pages[task_id].extend(new_pages)
|
|
||||||
|
|
||||||
self._write_req_to_token(task_id, prompt_ids, cached)
|
|
||||||
self._task_len[req_idx] = len(prompt_ids)
|
|
||||||
self._task_cached[task_id] = cached
|
|
||||||
return True
|
|
||||||
|
|
||||||
def task_free(self, task_id: str):
|
|
||||||
req_idx = self._task_req.pop(task_id, None)
|
|
||||||
if req_idx is None:
|
|
||||||
return
|
|
||||||
self._task_len.pop(req_idx, None)
|
|
||||||
self._task_cached.pop(task_id, None)
|
|
||||||
|
|
||||||
if not self.contiguous:
|
|
||||||
if self._prefix is not None:
|
|
||||||
for p in self._task_pages.get(task_id, []):
|
|
||||||
keep = self._prefix.has_page(p)
|
|
||||||
self._alloc.free(p, keep_cached=keep)
|
|
||||||
if not keep:
|
|
||||||
self._prefix.evict(p)
|
|
||||||
else:
|
|
||||||
for p in self._task_pages.get(task_id, []):
|
|
||||||
self._alloc.free(p)
|
|
||||||
self._task_pages.pop(task_id, None)
|
|
||||||
self._task_slots.pop(task_id, None)
|
|
||||||
|
|
||||||
self._req_pool.free([req_idx])
|
|
||||||
|
|
||||||
def task_extend(self, task_id: str, pos: int) -> bool:
|
|
||||||
req_idx = self._task_req.get(task_id)
|
|
||||||
if req_idx is None:
|
|
||||||
return False
|
|
||||||
|
|
||||||
if self.contiguous:
|
|
||||||
return pos < self.max_seq_len
|
|
||||||
|
|
||||||
if self.page_size == 1:
|
|
||||||
slots = self._alloc_tokens(1)
|
|
||||||
if slots is None:
|
|
||||||
return False
|
|
||||||
self._task_slots.setdefault(task_id, []).extend(slots)
|
|
||||||
self._req_pool.req_to_token[req_idx, pos] = slots[0]
|
|
||||||
else:
|
|
||||||
page_idx = pos // self.page_size
|
|
||||||
existing = self._task_pages.get(task_id, [])
|
|
||||||
if page_idx >= len(existing):
|
|
||||||
p = self._alloc.alloc()
|
|
||||||
if p < 0:
|
|
||||||
return False
|
|
||||||
existing.append(p)
|
|
||||||
self._task_pages[task_id] = existing
|
|
||||||
page_offset = pos % self.page_size
|
|
||||||
page = existing[page_idx]
|
|
||||||
token_slot = page * self.page_size + page_offset
|
|
||||||
self._req_pool.req_to_token[req_idx, pos] = token_slot
|
|
||||||
|
|
||||||
self._task_len[req_idx] = pos + 1
|
|
||||||
return True
|
|
||||||
|
|
||||||
def task_cached(self, task_id: str) -> int:
|
|
||||||
return self._task_cached.get(task_id, 0)
|
|
||||||
|
|
||||||
def task_record_hashes(
|
|
||||||
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
|
|
||||||
):
|
|
||||||
if self._prefix is None or self.contiguous:
|
|
||||||
return
|
|
||||||
pages = self._task_pages.get(task_id, [])
|
|
||||||
full_pages = len(prompt_ids) // self.page_size
|
|
||||||
for i in range(start_logical_page, min(full_pages, len(pages))):
|
|
||||||
self._prefix.record(pages[i], prompt_ids, i)
|
|
||||||
|
|
||||||
# ---- bind for forward ----
|
|
||||||
|
|
||||||
def bind_tasks(
|
|
||||||
self,
|
|
||||||
task_ids: List[str],
|
|
||||||
seq_lens: List[int],
|
|
||||||
device: torch.device,
|
|
||||||
start_pos: Optional[int] = None,
|
|
||||||
) -> KVCache:
|
|
||||||
req_indices = [self._task_req[tid] for tid in task_ids]
|
|
||||||
req_pool_indices = torch.tensor(req_indices, dtype=torch.long, device=device)
|
|
||||||
seq_lens_t = torch.tensor(seq_lens, dtype=torch.long, device=device)
|
|
||||||
|
|
||||||
if start_pos is not None:
|
|
||||||
seq_len = seq_lens[0]
|
|
||||||
out_cache_loc = self._req_pool.req_to_token[
|
|
||||||
req_pool_indices, start_pos:seq_len
|
|
||||||
]
|
|
||||||
page_table = None
|
|
||||||
decode_mask = None
|
|
||||||
else:
|
|
||||||
write_pos = seq_lens_t - 1
|
|
||||||
out_cache_loc = self._req_pool.req_to_token[
|
|
||||||
req_pool_indices, write_pos
|
|
||||||
].unsqueeze(-1)
|
|
||||||
ml = max(seq_lens)
|
|
||||||
page_table = self._req_pool.req_to_token[req_pool_indices, :ml]
|
|
||||||
if len(task_ids) > 1:
|
|
||||||
decode_mask = (
|
|
||||||
torch.arange(ml, device=device)[None, :] < seq_lens_t[:, None]
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
decode_mask = None
|
|
||||||
|
|
||||||
return KVCache(
|
|
||||||
k_buffer=self._storage.k_buffer,
|
|
||||||
v_buffer=self._storage.v_buffer,
|
|
||||||
req_to_token=self._req_pool.req_to_token,
|
|
||||||
req_pool_indices=req_pool_indices,
|
|
||||||
seq_lens=seq_lens_t,
|
|
||||||
out_cache_loc=out_cache_loc,
|
|
||||||
max_len=max(seq_lens),
|
|
||||||
page_table=page_table,
|
|
||||||
decode_mask=decode_mask,
|
|
||||||
)
|
|
||||||
|
|
||||||
# ---- internals ----
|
|
||||||
|
|
||||||
def _alloc_tokens(self, n: int) -> Optional[List[int]]:
|
|
||||||
if self.page_size != 1:
|
|
||||||
raise RuntimeError("_alloc_tokens is for page_size=1 only")
|
|
||||||
slots = []
|
|
||||||
for _ in range(n):
|
|
||||||
p = self._alloc.alloc()
|
|
||||||
if p < 0:
|
|
||||||
for s in slots:
|
|
||||||
self._alloc.free(s)
|
|
||||||
return None
|
|
||||||
slots.append(p)
|
|
||||||
return slots
|
|
||||||
|
|
||||||
def _write_req_to_token(self, task_id: str, prompt_ids: List[int], cached: int):
|
|
||||||
req_idx = self._task_req[task_id]
|
|
||||||
total = len(prompt_ids)
|
|
||||||
|
|
||||||
if self.contiguous:
|
|
||||||
return
|
|
||||||
|
|
||||||
if self.page_size == 1:
|
|
||||||
slots = self._task_slots.get(task_id, [])
|
|
||||||
all_slots = slots[: total - cached]
|
|
||||||
if all_slots:
|
|
||||||
self._req_pool.req_to_token[req_idx, cached:total] = torch.tensor(
|
|
||||||
all_slots, dtype=torch.long, device=self.device
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
pages = self._task_pages.get(task_id, [])
|
|
||||||
for pos in range(cached, total):
|
|
||||||
page_idx = pos // self.page_size
|
|
||||||
page_offset = pos % self.page_size
|
|
||||||
if page_idx < len(pages):
|
|
||||||
token_slot = pages[page_idx] * self.page_size + page_offset
|
|
||||||
self._req_pool.req_to_token[req_idx, pos] = token_slot
|
|
||||||
@@ -1,174 +0,0 @@
|
|||||||
import logging
|
|
||||||
from typing import List, Optional
|
|
||||||
|
|
||||||
import torch
|
|
||||||
|
|
||||||
from astrai.inference.core.cache import PagePool
|
|
||||||
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: PagePool,
|
|
||||||
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,
|
|
||||||
kv_cache=self.kv_cache.bind_tasks(
|
|
||||||
task_ids, [prompt_len] * batch_sz, self.device, start_pos=start_pos
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
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
|
|
||||||
)
|
|
||||||
|
|
||||||
has_freq = bool((freq_penalties != 0).any())
|
|
||||||
if has_freq:
|
|
||||||
history_lists = []
|
|
||||||
history_lens = []
|
|
||||||
for t in tasks:
|
|
||||||
window = t.rep_window
|
|
||||||
prompt_part = t.prompt_ids[-window:]
|
|
||||||
ids = prompt_part + t.output_ids
|
|
||||||
history_lists.append(ids)
|
|
||||||
history_lens.append(len(ids))
|
|
||||||
|
|
||||||
max_len = max(history_lens) if history_lens else 0
|
|
||||||
padded_ids = torch.zeros(
|
|
||||||
len(tasks), max_len, dtype=torch.long, device=self.device
|
|
||||||
)
|
|
||||||
padded_mask = torch.zeros(
|
|
||||||
len(tasks), max_len, dtype=torch.bool, device=self.device
|
|
||||||
)
|
|
||||||
for i, h in enumerate(history_lists):
|
|
||||||
L = history_lens[i]
|
|
||||||
padded_ids[i, :L] = torch.as_tensor(
|
|
||||||
h, dtype=torch.long, device=self.device
|
|
||||||
)
|
|
||||||
padded_mask[i, :L] = True
|
|
||||||
else:
|
|
||||||
padded_ids = None
|
|
||||||
padded_mask = None
|
|
||||||
|
|
||||||
with torch.inference_mode():
|
|
||||||
outputs = self.model(
|
|
||||||
input_ids.unsqueeze(1),
|
|
||||||
input_mask=input_mask,
|
|
||||||
kv_cache=self.kv_cache.bind_tasks(
|
|
||||||
task_ids,
|
|
||||||
[t.next_pos + 1 for t in tasks],
|
|
||||||
self.device,
|
|
||||||
),
|
|
||||||
position_ids=position_ids.unsqueeze(1),
|
|
||||||
)
|
|
||||||
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 PagePool
|
|
||||||
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,
|
|
||||||
device: Optional[str] = None,
|
|
||||||
dtype: Optional[torch.dtype] = None,
|
|
||||||
cache: Optional[PagePool] = 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._task_mgr = TaskManager(
|
|
||||||
tokenizer=tokenizer,
|
|
||||||
max_batch_size=max_batch_size,
|
|
||||||
max_seq_len=self.max_seq_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
|
|
||||||
|
|
||||||
active = self._task_mgr.get_active_tasks()
|
|
||||||
|
|
||||||
to_prefill = [
|
|
||||||
t
|
|
||||||
for t in active
|
|
||||||
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 = active
|
|
||||||
|
|
||||||
valid: List[Task] = []
|
|
||||||
for t in decode_tasks:
|
|
||||||
if cache.task_extend(t.task_id, t.next_pos):
|
|
||||||
valid.append(t)
|
|
||||||
else:
|
|
||||||
t.status = TaskStatus.ABORTED
|
|
||||||
self._task_mgr.invoke_callback(t.task_id, STOP)
|
|
||||||
|
|
||||||
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
|
|
||||||
+69
-172
@@ -8,9 +8,10 @@ from typing import Any, AsyncGenerator, Dict, Generator, List, Optional, Tuple,
|
|||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
|
||||||
from astrai.inference.core.cache import PagePool
|
from astrai.extension import ATTN_BACKEND, AttentionBackend, get_backend
|
||||||
from astrai.inference.core.scheduler import InferenceScheduler
|
from astrai.inference.cache import PagePool
|
||||||
from astrai.inference.core.task import STOP
|
from astrai.inference.scheduler import InferenceScheduler
|
||||||
|
from astrai.inference.task import STOP
|
||||||
from astrai.tokenize import AutoTokenizer
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
|
||||||
@@ -64,44 +65,6 @@ class GenerateResult:
|
|||||||
return self.results.copy()
|
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:
|
class InferenceEngine:
|
||||||
"""Unified inference engine backed by continuous-batching scheduler."""
|
"""Unified inference engine backed by continuous-batching scheduler."""
|
||||||
|
|
||||||
@@ -112,6 +75,8 @@ class InferenceEngine:
|
|||||||
max_batch_size: int = 1,
|
max_batch_size: int = 1,
|
||||||
max_seq_len: Optional[int] = None,
|
max_seq_len: Optional[int] = None,
|
||||||
cache: Optional[PagePool] = None,
|
cache: Optional[PagePool] = None,
|
||||||
|
enable_cuda_graph: bool = True,
|
||||||
|
backend: Optional[Union[str, ATTN_BACKEND, AttentionBackend, type]] = None,
|
||||||
):
|
):
|
||||||
self.model = model
|
self.model = model
|
||||||
self.tokenizer = tokenizer
|
self.tokenizer = tokenizer
|
||||||
@@ -121,6 +86,8 @@ class InferenceEngine:
|
|||||||
max_batch_size=max_batch_size,
|
max_batch_size=max_batch_size,
|
||||||
max_seq_len=max_seq_len,
|
max_seq_len=max_seq_len,
|
||||||
cache=cache,
|
cache=cache,
|
||||||
|
enable_cuda_graph=enable_cuda_graph,
|
||||||
|
backend=backend,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.scheduler.start()
|
self.scheduler.start()
|
||||||
@@ -146,28 +113,23 @@ class InferenceEngine:
|
|||||||
is_batch = isinstance(prompt, list)
|
is_batch = isinstance(prompt, list)
|
||||||
prompts = prompt if is_batch else [prompt]
|
prompts = prompt if is_batch else [prompt]
|
||||||
|
|
||||||
if stream:
|
if max_tokens is not None and max_tokens <= 0:
|
||||||
return self._generate_streaming(
|
if stream:
|
||||||
prompts,
|
return iter(())
|
||||||
is_batch,
|
results = [""] * len(prompts)
|
||||||
max_tokens,
|
return results if is_batch else results[0]
|
||||||
temperature,
|
|
||||||
top_p,
|
return self._generate(
|
||||||
top_k,
|
prompts,
|
||||||
frequency_penalty,
|
is_batch,
|
||||||
rep_window,
|
stream,
|
||||||
)
|
max_tokens,
|
||||||
else:
|
temperature,
|
||||||
return self._generate_non_streaming(
|
top_p,
|
||||||
prompts,
|
top_k,
|
||||||
is_batch,
|
frequency_penalty,
|
||||||
max_tokens,
|
rep_window,
|
||||||
temperature,
|
)
|
||||||
top_p,
|
|
||||||
top_k,
|
|
||||||
frequency_penalty,
|
|
||||||
rep_window,
|
|
||||||
)
|
|
||||||
|
|
||||||
def generate_async(
|
def generate_async(
|
||||||
self,
|
self,
|
||||||
@@ -179,9 +141,10 @@ class InferenceEngine:
|
|||||||
frequency_penalty: float = 0.0,
|
frequency_penalty: float = 0.0,
|
||||||
rep_window: int = 64,
|
rep_window: int = 64,
|
||||||
) -> AsyncGenerator[str, None]:
|
) -> AsyncGenerator[str, None]:
|
||||||
sync_gen = self._generate_streaming(
|
sync_gen = self._generate(
|
||||||
[prompt],
|
[prompt],
|
||||||
False,
|
False,
|
||||||
|
True,
|
||||||
max_tokens,
|
max_tokens,
|
||||||
temperature,
|
temperature,
|
||||||
top_p,
|
top_p,
|
||||||
@@ -193,51 +156,30 @@ class InferenceEngine:
|
|||||||
async def _agen():
|
async def _agen():
|
||||||
loop = asyncio.get_event_loop()
|
loop = asyncio.get_event_loop()
|
||||||
while True:
|
while True:
|
||||||
token = await loop.run_in_executor(None, self._next_token, sync_gen)
|
token = await loop.run_in_executor(None, next, sync_gen, None)
|
||||||
if token is None:
|
if token is None:
|
||||||
break
|
break
|
||||||
yield token
|
yield token
|
||||||
|
|
||||||
return _agen()
|
return _agen()
|
||||||
|
|
||||||
@staticmethod
|
def _generate(
|
||||||
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(
|
|
||||||
self,
|
self,
|
||||||
prompts: List[str],
|
prompts: List[str],
|
||||||
|
is_batch: bool,
|
||||||
|
stream: bool,
|
||||||
max_tokens: Optional[int],
|
max_tokens: Optional[int],
|
||||||
temperature: float,
|
temperature: float,
|
||||||
top_p: float,
|
top_p: float,
|
||||||
top_k: int,
|
top_k: int,
|
||||||
frequency_penalty: float,
|
frequency_penalty: float,
|
||||||
rep_window: int,
|
rep_window: int,
|
||||||
) -> Tuple[GenerateResult, List[str]]:
|
) -> Union[Generator, str, List[str]]:
|
||||||
n = len(prompts)
|
n = len(prompts)
|
||||||
|
request_backend = get_backend(use_default=False)
|
||||||
result = GenerateResult(count=n)
|
result = GenerateResult(count=n)
|
||||||
task_ids = []
|
task_ids = [
|
||||||
for i, p in enumerate(prompts):
|
self.scheduler.add_task(
|
||||||
cb = self._make_callback(result, i)
|
|
||||||
task_id = self.scheduler.add_task(
|
|
||||||
prompt=p,
|
prompt=p,
|
||||||
max_tokens=max_tokens,
|
max_tokens=max_tokens,
|
||||||
temperature=temperature,
|
temperature=temperature,
|
||||||
@@ -245,99 +187,54 @@ class InferenceEngine:
|
|||||||
top_k=top_k,
|
top_k=top_k,
|
||||||
frequency_penalty=frequency_penalty,
|
frequency_penalty=frequency_penalty,
|
||||||
rep_window=rep_window,
|
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)
|
for i, p in enumerate(prompts)
|
||||||
return result, task_ids
|
]
|
||||||
|
|
||||||
@staticmethod
|
if not stream:
|
||||||
def _make_callback(result: GenerateResult, idx: int):
|
try:
|
||||||
def cb(token):
|
result.wait_completion()
|
||||||
result.append(token, idx)
|
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
|
remaining = n
|
||||||
finished = [False] * n
|
finished = [False] * n
|
||||||
|
|
||||||
def gen():
|
def gen():
|
||||||
nonlocal remaining
|
nonlocal remaining
|
||||||
try:
|
while remaining > 0:
|
||||||
while remaining > 0:
|
items = result.pop_all()
|
||||||
items = result.pop_all()
|
for idx, token in items:
|
||||||
for idx, token in items:
|
if token is STOP:
|
||||||
if token is STOP:
|
if not finished[idx]:
|
||||||
if not finished[idx]:
|
finished[idx] = True
|
||||||
finished[idx] = True
|
remaining -= 1
|
||||||
remaining -= 1
|
else:
|
||||||
else:
|
yield (idx, token) if is_batch else token
|
||||||
yield (idx, token) if is_batch else token
|
if remaining > 0:
|
||||||
if remaining > 0:
|
result.wait(timeout=0.05)
|
||||||
result.wait(timeout=0.05)
|
|
||||||
finally:
|
|
||||||
for tid in task_ids:
|
|
||||||
self.scheduler.remove_task(tid)
|
|
||||||
|
|
||||||
return gen()
|
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]:
|
def get_stats(self) -> Dict[str, Any]:
|
||||||
return self.scheduler.get_stats()
|
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):
|
def shutdown(self):
|
||||||
self.scheduler.stop()
|
self.scheduler.stop()
|
||||||
if torch.cuda.is_available():
|
if torch.cuda.is_available():
|
||||||
|
|||||||
@@ -0,0 +1,201 @@
|
|||||||
|
"""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)
|
||||||
|
|
||||||
|
# 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
|
||||||
|
|
||||||
|
# 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.
|
lazy singleton FastAPI instance.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from astrai.inference.api.protocol import GenContext, ProtocolHandler, StopChecker
|
from astrai.inference.network.app import (
|
||||||
from astrai.inference.api.server import (
|
|
||||||
AnthropicMessage,
|
AnthropicMessage,
|
||||||
ChatCompletionRequest,
|
ChatCompletionRequest,
|
||||||
ChatMessage,
|
ChatMessage,
|
||||||
@@ -15,7 +14,8 @@ from astrai.inference.api.server import (
|
|||||||
get_app,
|
get_app,
|
||||||
run_server,
|
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,
|
BaseToolParser,
|
||||||
SimpleJsonToolParser,
|
SimpleJsonToolParser,
|
||||||
ToolParserFactory,
|
ToolParserFactory,
|
||||||
@@ -6,13 +6,13 @@ from typing import Any, Dict, List, Tuple, Union
|
|||||||
|
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
|
|
||||||
from astrai.inference.api.protocol import (
|
from astrai.inference.engine import InferenceEngine
|
||||||
|
from astrai.inference.network.protocol import (
|
||||||
GenContext,
|
GenContext,
|
||||||
ResponseBuilder,
|
ResponseBuilder,
|
||||||
StopInfo,
|
StopInfo,
|
||||||
sse_event,
|
sse_event,
|
||||||
)
|
)
|
||||||
from astrai.inference.engine import InferenceEngine
|
|
||||||
|
|
||||||
|
|
||||||
def _extract_text(content: Union[str, List[Dict[str, Any]]]) -> str:
|
def _extract_text(content: Union[str, List[Dict[str, Any]]]) -> str:
|
||||||
@@ -18,10 +18,10 @@ import uvicorn
|
|||||||
from fastapi import APIRouter, FastAPI, HTTPException
|
from fastapi import APIRouter, FastAPI, HTTPException
|
||||||
from pydantic import BaseModel, Field
|
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.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.model import AutoModel
|
||||||
from astrai.tokenize import AutoTokenizer
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
@@ -7,14 +7,14 @@ from typing import Any, Dict, List, Optional, Tuple, Union
|
|||||||
|
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
|
|
||||||
from astrai.inference.api.protocol import (
|
from astrai.inference.engine import InferenceEngine
|
||||||
|
from astrai.inference.network.protocol import (
|
||||||
GenContext,
|
GenContext,
|
||||||
ResponseBuilder,
|
ResponseBuilder,
|
||||||
StopInfo,
|
StopInfo,
|
||||||
sse_event,
|
sse_event,
|
||||||
)
|
)
|
||||||
from astrai.inference.api.tool_parser import BaseToolParser, ToolParserFactory
|
from astrai.inference.network.tool_parser import BaseToolParser, ToolParserFactory
|
||||||
from astrai.inference.engine import InferenceEngine
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -181,12 +181,10 @@ class ProtocolHandler:
|
|||||||
self, agen: AsyncGenerator, ctx: GenContext, stop_sequences: List[str]
|
self, agen: AsyncGenerator, ctx: GenContext, stop_sequences: List[str]
|
||||||
) -> Dict[str, Any]:
|
) -> Dict[str, Any]:
|
||||||
checker = StopChecker(stop_sequences)
|
checker = StopChecker(stop_sequences)
|
||||||
chunks: List[str] = []
|
|
||||||
body = ""
|
body = ""
|
||||||
matched = None
|
matched = None
|
||||||
|
|
||||||
async for token in agen:
|
async for token in agen:
|
||||||
chunks.append(token)
|
|
||||||
body += token
|
body += token
|
||||||
|
|
||||||
matched = checker.check(body)
|
matched = checker.check(body)
|
||||||
@@ -195,6 +193,5 @@ class ProtocolHandler:
|
|||||||
|
|
||||||
ctx.completion_tokens += 1
|
ctx.completion_tokens += 1
|
||||||
|
|
||||||
content = "".join(chunks)
|
|
||||||
stop = StopInfo(matched=matched, body=body)
|
stop = StopInfo(matched=matched, body=body)
|
||||||
return self.builder.format_response(ctx, content, stop)
|
return self.builder.format_response(ctx, body, stop)
|
||||||
@@ -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.available() and head_dim in CudaBackend.HEAD_DIMS
|
||||||
|
)
|
||||||
|
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
|
@staticmethod
|
||||||
def _is_greedy(temperature: Union[float, Tensor]) -> bool:
|
def _is_greedy(temperature: Union[float, Tensor]) -> bool:
|
||||||
if isinstance(temperature, Tensor):
|
if isinstance(temperature, Tensor):
|
||||||
return temperature.numel() == 1 and temperature.item() == 0
|
return bool((temperature == 0).all())
|
||||||
return temperature == 0
|
return temperature == 0
|
||||||
|
|
||||||
@torch.inference_mode()
|
@torch.inference_mode()
|
||||||
@@ -305,12 +305,12 @@ class SamplingPipeline(BaseSamplingStrategy):
|
|||||||
return tokens, chosen
|
return tokens, chosen
|
||||||
|
|
||||||
transformed = self.apply(logits, filter_value, input_ids, input_mask)
|
transformed = self.apply(logits, filter_value, input_ids, input_mask)
|
||||||
log_probs = torch.log_softmax(transformed.float(), dim=-1)
|
|
||||||
tokens = torch.multinomial(
|
tokens = torch.multinomial(
|
||||||
torch.softmax(transformed, dim=-1), num_samples=1
|
torch.softmax(transformed, dim=-1), num_samples=1
|
||||||
).squeeze(-1)
|
).squeeze(-1)
|
||||||
if not return_logprobs:
|
if not return_logprobs:
|
||||||
return tokens
|
return tokens
|
||||||
|
log_probs = torch.log_softmax(transformed.float(), dim=-1)
|
||||||
chosen = torch.gather(log_probs, -1, tokens.unsqueeze(-1)).squeeze(-1)
|
chosen = torch.gather(log_probs, -1, tokens.unsqueeze(-1)).squeeze(-1)
|
||||||
return tokens, chosen
|
return tokens, chosen
|
||||||
|
|
||||||
@@ -363,24 +363,6 @@ def sample(
|
|||||||
``True`` — a ``(token_ids, chosen_logprobs)`` tuple where
|
``True`` — a ``(token_ids, chosen_logprobs)`` tuple where
|
||||||
``chosen_logprobs`` has shape ``[batch]``.
|
``chosen_logprobs`` has shape ``[batch]``.
|
||||||
"""
|
"""
|
||||||
greedy = (
|
|
||||||
(
|
|
||||||
isinstance(temperature, Tensor)
|
|
||||||
and temperature.numel() == 1
|
|
||||||
and temperature.item() == 0
|
|
||||||
)
|
|
||||||
if isinstance(temperature, Tensor)
|
|
||||||
else temperature == 0
|
|
||||||
)
|
|
||||||
|
|
||||||
if greedy:
|
|
||||||
tokens = logits.argmax(dim=-1)
|
|
||||||
if not return_logprobs:
|
|
||||||
return tokens
|
|
||||||
log_probs = torch.log_softmax(logits.float(), dim=-1)
|
|
||||||
chosen = torch.gather(log_probs, -1, tokens.unsqueeze(-1)).squeeze(-1)
|
|
||||||
return tokens, chosen
|
|
||||||
|
|
||||||
has_freq = (
|
has_freq = (
|
||||||
(isinstance(frequency_penalty, Tensor) and (frequency_penalty != 0).any())
|
(isinstance(frequency_penalty, Tensor) and (frequency_penalty != 0).any())
|
||||||
if isinstance(frequency_penalty, Tensor)
|
if isinstance(frequency_penalty, Tensor)
|
||||||
@@ -0,0 +1,400 @@
|
|||||||
|
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
|
||||||
|
active_backend = get_backend()
|
||||||
|
else:
|
||||||
|
active_backend = backend
|
||||||
|
with attn_backend(active_backend):
|
||||||
|
if backend is not None:
|
||||||
|
self._backend = get_backend()
|
||||||
|
self._backend_name = type(get_backend()).__name__
|
||||||
|
self._executor = Executor(
|
||||||
|
model=model,
|
||||||
|
kv_cache=self._cache,
|
||||||
|
task_cache=self._task_cache,
|
||||||
|
device=self.device,
|
||||||
|
dtype=self.dtype,
|
||||||
|
enable_cuda_graph=enable_cuda_graph,
|
||||||
|
)
|
||||||
|
|
||||||
|
self._stop_event = threading.Event()
|
||||||
|
self._loop_thread: Optional[threading.Thread] = None
|
||||||
|
|
||||||
|
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)
|
||||||
|
self._abort_and_clear(free_waiting=False)
|
||||||
|
|
||||||
|
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
|
||||||
|
self._abort_and_clear(free_waiting=True)
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
|
def _abort_and_clear(self, free_waiting: bool):
|
||||||
|
"""Invoke STOP callbacks, release cache slots, and clear task queues."""
|
||||||
|
for task in self._task_mgr.get_active_tasks():
|
||||||
|
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||||
|
self._task_cache.task_free(task.task_id)
|
||||||
|
for task in self._task_mgr.get_waiting_tasks():
|
||||||
|
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||||
|
if free_waiting:
|
||||||
|
self._task_cache.task_free(task.task_id)
|
||||||
|
self._task_mgr.clear_queues()
|
||||||
|
|
||||||
|
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,16 +1,17 @@
|
|||||||
import logging
|
|
||||||
import threading
|
import threading
|
||||||
import time
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
from collections import deque
|
from collections import deque
|
||||||
from enum import Enum
|
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 tokenizers.decoders import DecodeStream
|
||||||
|
|
||||||
|
from astrai.inference.metrics import MetricsCollector
|
||||||
from astrai.tokenize.tokenizer import AutoTokenizer
|
from astrai.tokenize.tokenizer import AutoTokenizer
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
if TYPE_CHECKING:
|
||||||
|
from astrai.extension import AttentionBackend
|
||||||
|
|
||||||
STOP = object()
|
STOP = object()
|
||||||
|
|
||||||
@@ -64,6 +65,7 @@ class Task:
|
|||||||
top_k: int = 50,
|
top_k: int = 50,
|
||||||
frequency_penalty: float = 0.0,
|
frequency_penalty: float = 0.0,
|
||||||
rep_window: int = 64,
|
rep_window: int = 64,
|
||||||
|
backend: Optional["AttentionBackend"] = None,
|
||||||
):
|
):
|
||||||
self.task_id = task_id
|
self.task_id = task_id
|
||||||
self.prompt_ids = prompt_ids
|
self.prompt_ids = prompt_ids
|
||||||
@@ -73,16 +75,25 @@ class Task:
|
|||||||
self.top_k = top_k
|
self.top_k = top_k
|
||||||
self.frequency_penalty = frequency_penalty
|
self.frequency_penalty = frequency_penalty
|
||||||
self.rep_window = rep_window
|
self.rep_window = rep_window
|
||||||
|
self.backend = backend
|
||||||
|
|
||||||
self.status = TaskStatus.PENDING
|
self.status = TaskStatus.PENDING
|
||||||
self.output_ids: List[int] = []
|
self.output_ids: List[int] = []
|
||||||
self.output_logprobs: List[float] = []
|
self.output_logprobs: List[float] = []
|
||||||
self.input_tokens: int = 0
|
self.input_tokens: int = 0
|
||||||
self.output_tokens: int = 0
|
self.output_tokens: int = 0
|
||||||
self.arrival_time = time.time()
|
self._kv_len: int = 0
|
||||||
self.finish_time: Optional[float] = None
|
|
||||||
self._decoder: Optional[StreamDecoder] = None
|
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:
|
def decode_new_token(self, tokenizer: AutoTokenizer) -> str:
|
||||||
"""Decode the last appended output token, buffering incomplete
|
"""Decode the last appended output token, buffering incomplete
|
||||||
multi-byte sequences across calls.
|
multi-byte sequences across calls.
|
||||||
@@ -93,19 +104,15 @@ class Task:
|
|||||||
self._decoder = StreamDecoder(tokenizer)
|
self._decoder = StreamDecoder(tokenizer)
|
||||||
return self._decoder.push(self.output_ids[-1])
|
return self._decoder.push(self.output_ids[-1])
|
||||||
|
|
||||||
def flush_remaining(self, tokenizer: AutoTokenizer) -> str:
|
|
||||||
"""Emit any text still buffered in the decoder.
|
|
||||||
|
|
||||||
With the Rust-native DecodeStream, the stream is always in a
|
|
||||||
correct state — any completed text was already emitted by the
|
|
||||||
last ``push``. A trailing incomplete multi-byte sequence has no
|
|
||||||
valid text to emit, so this is a no-op.
|
|
||||||
"""
|
|
||||||
return ""
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def next_pos(self) -> int:
|
def next_pos(self) -> int:
|
||||||
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:
|
def is_finished(self, stop_ids: List[int]) -> bool:
|
||||||
if self.max_tokens is not None and self.output_tokens >= self.max_tokens:
|
if self.max_tokens is not None and self.output_tokens >= self.max_tokens:
|
||||||
@@ -123,6 +130,7 @@ class TaskManager:
|
|||||||
tokenizer: AutoTokenizer,
|
tokenizer: AutoTokenizer,
|
||||||
max_batch_size: int = 16,
|
max_batch_size: int = 16,
|
||||||
max_seq_len: int = 8192,
|
max_seq_len: int = 8192,
|
||||||
|
metrics: Optional["MetricsCollector"] = None,
|
||||||
):
|
):
|
||||||
self.tokenizer = tokenizer
|
self.tokenizer = tokenizer
|
||||||
self.max_batch_size = max_batch_size
|
self.max_batch_size = max_batch_size
|
||||||
@@ -138,6 +146,8 @@ class TaskManager:
|
|||||||
self._total_tasks = 0
|
self._total_tasks = 0
|
||||||
self._total_tokens = 0
|
self._total_tokens = 0
|
||||||
|
|
||||||
|
self._metrics = metrics
|
||||||
|
|
||||||
def add_task(
|
def add_task(
|
||||||
self,
|
self,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
@@ -147,6 +157,7 @@ class TaskManager:
|
|||||||
top_k: int = 50,
|
top_k: int = 50,
|
||||||
frequency_penalty: float = 0.0,
|
frequency_penalty: float = 0.0,
|
||||||
rep_window: int = 64,
|
rep_window: int = 64,
|
||||||
|
backend: Optional["AttentionBackend"] = None,
|
||||||
stream_callback: Optional[Callable[[str], None]] = None,
|
stream_callback: Optional[Callable[[str], None]] = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}"
|
task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}"
|
||||||
@@ -154,11 +165,6 @@ class TaskManager:
|
|||||||
if len(prompt_ids) > self.max_seq_len:
|
if len(prompt_ids) > self.max_seq_len:
|
||||||
prompt_ids = prompt_ids[-self.max_seq_len :]
|
prompt_ids = prompt_ids[-self.max_seq_len :]
|
||||||
|
|
||||||
if len(prompt_ids) > self.max_seq_len:
|
|
||||||
if stream_callback:
|
|
||||||
stream_callback(STOP)
|
|
||||||
return task_id
|
|
||||||
|
|
||||||
if max_tokens is None:
|
if max_tokens is None:
|
||||||
max_tokens = self.max_seq_len - len(prompt_ids)
|
max_tokens = self.max_seq_len - len(prompt_ids)
|
||||||
else:
|
else:
|
||||||
@@ -173,6 +179,7 @@ class TaskManager:
|
|||||||
top_k=top_k,
|
top_k=top_k,
|
||||||
frequency_penalty=frequency_penalty,
|
frequency_penalty=frequency_penalty,
|
||||||
rep_window=rep_window,
|
rep_window=rep_window,
|
||||||
|
backend=backend,
|
||||||
)
|
)
|
||||||
|
|
||||||
with self._lock:
|
with self._lock:
|
||||||
@@ -181,6 +188,9 @@ class TaskManager:
|
|||||||
if stream_callback:
|
if stream_callback:
|
||||||
self._callbacks[task_id] = stream_callback
|
self._callbacks[task_id] = stream_callback
|
||||||
|
|
||||||
|
if self._metrics is not None:
|
||||||
|
self._metrics.register(task_id)
|
||||||
|
|
||||||
self._task_event.set()
|
self._task_event.set()
|
||||||
return task_id
|
return task_id
|
||||||
|
|
||||||
@@ -200,26 +210,33 @@ class TaskManager:
|
|||||||
cb(token)
|
cb(token)
|
||||||
|
|
||||||
def get_stats(self) -> Dict[str, Any]:
|
def get_stats(self) -> Dict[str, Any]:
|
||||||
return {
|
stats: Dict[str, Any] = {
|
||||||
"total_tasks": self._total_tasks,
|
"total_tasks": self._total_tasks,
|
||||||
"total_tokens": self._total_tokens,
|
"total_tokens": self._total_tokens,
|
||||||
"active_tasks": len(self.active_tasks),
|
"active_tasks": len(self.active_tasks),
|
||||||
"waiting_queue": len(self.waiting_queue),
|
"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]:
|
def remove_finished_tasks(self, stop_ids: List[int]) -> List[Task]:
|
||||||
with self._lock:
|
with self._lock:
|
||||||
finished = []
|
finished = []
|
||||||
for task in self.active_tasks:
|
for task in self.active_tasks:
|
||||||
if task.status == TaskStatus.ABORTED:
|
if task.status == TaskStatus.ABORTED:
|
||||||
task.finish_time = time.time()
|
|
||||||
finished.append(task)
|
finished.append(task)
|
||||||
elif task.is_finished(stop_ids):
|
elif task.is_finished(stop_ids):
|
||||||
task.status = TaskStatus.FINISHED
|
task.status = TaskStatus.FINISHED
|
||||||
task.finish_time = time.time()
|
|
||||||
finished.append(task)
|
finished.append(task)
|
||||||
self._total_tokens += task.output_tokens
|
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 = [
|
self.active_tasks = [
|
||||||
t
|
t
|
||||||
for t in self.active_tasks
|
for t in self.active_tasks
|
||||||
@@ -0,0 +1,152 @@
|
|||||||
|
"""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
|
||||||
|
Q_TILE_ROWS = 64
|
||||||
|
|
||||||
|
|
||||||
|
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
|
||||||
|
)
|
||||||
|
max_q_tiles = max_batch_size * ((max_seq_len + Q_TILE_ROWS - 1) // Q_TILE_ROWS)
|
||||||
|
self.q_tile_to_batch = torch.empty(
|
||||||
|
(max_q_tiles,), dtype=torch.int32, device=device
|
||||||
|
)
|
||||||
|
self.q_tile_to_index = torch.empty(
|
||||||
|
(max_q_tiles,), 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 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,37 @@
|
|||||||
|
import logging
|
||||||
|
import os
|
||||||
|
|
||||||
|
from astrai.parallel.setup import get_rank, get_world_size
|
||||||
|
|
||||||
|
|
||||||
|
class _DistributedContextFilter(logging.Filter):
|
||||||
|
def filter(self, record: logging.LogRecord) -> bool:
|
||||||
|
record.rank = str(get_rank())
|
||||||
|
record.world_size = str(get_world_size())
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
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.addFilter(_DistributedContextFilter())
|
||||||
|
handler.setFormatter(
|
||||||
|
logging.Formatter(
|
||||||
|
"%(asctime)s | %(levelname)-8s | rank=%(rank)2s/%(world_size)-2s | %(name)-32s | %(message)s",
|
||||||
|
datefmt="%Y-%m-%d %H:%M:%S",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
logger.addHandler(handler)
|
||||||
@@ -9,7 +9,7 @@ from astrai.model.components.lora import (
|
|||||||
merge_lora,
|
merge_lora,
|
||||||
save_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.components.norm import RMSNorm
|
||||||
from astrai.model.encoder import EmbeddingEncoder
|
from astrai.model.encoder import EmbeddingEncoder
|
||||||
from astrai.model.transformer import AutoRegressiveLM
|
from astrai.model.transformer import AutoRegressiveLM
|
||||||
@@ -19,6 +19,7 @@ __all__ = [
|
|||||||
"Linear",
|
"Linear",
|
||||||
"RMSNorm",
|
"RMSNorm",
|
||||||
"MLP",
|
"MLP",
|
||||||
|
"DeepSeekMoE",
|
||||||
"GQA",
|
"GQA",
|
||||||
"DecoderBlock",
|
"DecoderBlock",
|
||||||
# Models
|
# Models
|
||||||
|
|||||||
@@ -4,13 +4,21 @@ AutoModel base class for model loading and saving.
|
|||||||
|
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Self, Union
|
from typing import Union
|
||||||
|
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
|
||||||
from astrai.config.model_config import BaseModelConfig, ConfigFactory
|
from astrai.config.model_config import BaseModelConfig, ConfigFactory
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
from astrai.serialization import load_model_config, load_model_weights, save_model
|
from astrai.serialization import (
|
||||||
|
HF_MODEL_TYPES,
|
||||||
|
adapt_config,
|
||||||
|
convert_hf_weights,
|
||||||
|
load_model_config,
|
||||||
|
load_model_weights,
|
||||||
|
looks_like_hf_state_dict,
|
||||||
|
save_model,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
@@ -57,7 +65,25 @@ class AutoModel(nn.Module):
|
|||||||
path: Union[str, Path],
|
path: Union[str, Path],
|
||||||
disable_random_init: bool = True,
|
disable_random_init: bool = True,
|
||||||
strict: bool = True,
|
strict: bool = True,
|
||||||
|
weights_format: str = "auto",
|
||||||
) -> nn.Module:
|
) -> nn.Module:
|
||||||
|
"""Load a model directory.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
path: Directory containing ``config.json`` and optionally
|
||||||
|
``model.safetensors``.
|
||||||
|
disable_random_init: Replace parameter initializers with no-ops
|
||||||
|
while building the model.
|
||||||
|
strict: Passed to ``load_state_dict``.
|
||||||
|
weights_format: ``"auto"`` detects HuggingFace checkpoints
|
||||||
|
(LLaMA-style keys and ``model_type``) and converts them;
|
||||||
|
``"astrai"`` skips conversion; ``"hf"`` forces it.
|
||||||
|
"""
|
||||||
|
if weights_format not in ("auto", "astrai", "hf"):
|
||||||
|
raise ValueError(
|
||||||
|
f"weights_format must be one of 'auto', 'astrai', 'hf', "
|
||||||
|
f"got {weights_format!r}"
|
||||||
|
)
|
||||||
|
|
||||||
model_path = Path(path)
|
model_path = Path(path)
|
||||||
|
|
||||||
@@ -66,6 +92,12 @@ class AutoModel(nn.Module):
|
|||||||
raise FileNotFoundError(f"Config file not found: {config_path}")
|
raise FileNotFoundError(f"Config file not found: {config_path}")
|
||||||
|
|
||||||
raw = load_model_config(str(model_path))
|
raw = load_model_config(str(model_path))
|
||||||
|
is_hf_config = weights_format == "hf" or (
|
||||||
|
weights_format == "auto" and raw.get("model_type") in HF_MODEL_TYPES
|
||||||
|
)
|
||||||
|
if is_hf_config:
|
||||||
|
raw = adapt_config(raw)
|
||||||
|
|
||||||
config = ConfigFactory.load(raw)
|
config = ConfigFactory.load(raw)
|
||||||
model_type = config.model_type or "autoregressive_lm"
|
model_type = config.model_type or "autoregressive_lm"
|
||||||
|
|
||||||
@@ -75,8 +107,14 @@ class AutoModel(nn.Module):
|
|||||||
model = actual_cls(config)
|
model = actual_cls(config)
|
||||||
|
|
||||||
weights_path = model_path / "model.safetensors"
|
weights_path = model_path / "model.safetensors"
|
||||||
if weights_path.exists():
|
index_path = model_path / "model.safetensors.index.json"
|
||||||
|
if weights_path.exists() or index_path.exists():
|
||||||
state_dict = load_model_weights(str(model_path))
|
state_dict = load_model_weights(str(model_path))
|
||||||
|
is_hf_weights = is_hf_config or (
|
||||||
|
weights_format == "auto" and looks_like_hf_state_dict(state_dict)
|
||||||
|
)
|
||||||
|
if is_hf_weights:
|
||||||
|
state_dict = convert_hf_weights(state_dict, config)
|
||||||
model.load_state_dict(state_dict, strict=strict)
|
model.load_state_dict(state_dict, strict=strict)
|
||||||
|
|
||||||
return model
|
return model
|
||||||
@@ -90,7 +128,3 @@ class AutoModel(nn.Module):
|
|||||||
state_dict=self.state_dict(),
|
state_dict=self.state_dict(),
|
||||||
save_directory=str(save_directory),
|
save_directory=str(save_directory),
|
||||||
)
|
)
|
||||||
|
|
||||||
def to(self, *args, **kwargs) -> Self:
|
|
||||||
"""Move model to device/dtype."""
|
|
||||||
return super().to(*args, **kwargs)
|
|
||||||
|
|||||||
@@ -1,9 +1,9 @@
|
|||||||
from astrai.extension.rotary_backend import apply_rotary_emb
|
from astrai.extension.backend.rotary import apply_rotary_emb
|
||||||
from astrai.model.components.attention import GQA, MLA
|
from astrai.model.components.attention import GQA, MLA
|
||||||
from astrai.model.components.decoder_block import DecoderBlock
|
from astrai.model.components.decoder_block import DecoderBlock
|
||||||
from astrai.model.components.embedding import Embedding
|
from astrai.model.components.embedding import Embedding
|
||||||
from astrai.model.components.linear import Linear
|
from astrai.model.components.linear import Linear
|
||||||
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.norm import RMSNorm
|
||||||
from astrai.model.components.rope import (
|
from astrai.model.components.rope import (
|
||||||
RotaryEmbedding,
|
RotaryEmbedding,
|
||||||
@@ -14,6 +14,7 @@ __all__ = [
|
|||||||
"Linear",
|
"Linear",
|
||||||
"RMSNorm",
|
"RMSNorm",
|
||||||
"MLP",
|
"MLP",
|
||||||
|
"DeepSeekMoE",
|
||||||
"Embedding",
|
"Embedding",
|
||||||
"GQA",
|
"GQA",
|
||||||
"MLA",
|
"MLA",
|
||||||
|
|||||||
@@ -5,10 +5,9 @@ import torch.nn as nn
|
|||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from astrai.extension import attention
|
from astrai.extension.backend import apply_rotary_emb, attention
|
||||||
from astrai.extension.rotary_backend import apply_rotary_emb
|
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
from astrai.inference.core.cache import KVCache
|
from astrai.inference.cache import KVCache
|
||||||
from astrai.model.components.linear import Linear
|
from astrai.model.components.linear import Linear
|
||||||
from astrai.model.components.norm import RMSNorm
|
from astrai.model.components.norm import RMSNorm
|
||||||
|
|
||||||
@@ -56,9 +55,7 @@ class GQA(nn.Module):
|
|||||||
self.gate = Linear(dim, dim)
|
self.gate = Linear(dim, dim)
|
||||||
|
|
||||||
def _split_heads(self, x: Tensor, n_heads) -> Tensor:
|
def _split_heads(self, x: Tensor, n_heads) -> Tensor:
|
||||||
batch_size, seq_len, _ = x.shape
|
return x.reshape(*x.shape[:-1], n_heads, self.head_dim)
|
||||||
x = x.reshape(batch_size, seq_len, n_heads, self.head_dim)
|
|
||||||
return x
|
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
@@ -67,6 +64,7 @@ class GQA(nn.Module):
|
|||||||
attn_mask: Tensor = None,
|
attn_mask: Tensor = None,
|
||||||
kv_cache: Optional[KVCache] = None,
|
kv_cache: Optional[KVCache] = None,
|
||||||
is_causal: bool = False,
|
is_causal: bool = False,
|
||||||
|
fwd: Optional[str] = None,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
q = self._split_heads(self.q_proj(x), self.n_heads)
|
q = self._split_heads(self.q_proj(x), self.n_heads)
|
||||||
k = self._split_heads(self.k_proj(x), self.n_kv_heads)
|
k = self._split_heads(self.k_proj(x), self.n_kv_heads)
|
||||||
@@ -76,7 +74,9 @@ class GQA(nn.Module):
|
|||||||
if self.use_qk_norm:
|
if self.use_qk_norm:
|
||||||
q, k = self.q_norm(q), self.k_norm(k)
|
q, k = self.q_norm(q), self.k_norm(k)
|
||||||
|
|
||||||
sdqa_out = attention(q, k, v, kv_cache, self.layer_id, attn_mask, is_causal)
|
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:
|
if self.use_gated_attention:
|
||||||
sdqa_out = sdqa_out * F.sigmoid(self.gate(x))
|
sdqa_out = sdqa_out * F.sigmoid(self.gate(x))
|
||||||
@@ -141,17 +141,16 @@ class MLA(nn.Module):
|
|||||||
attn_mask: Tensor = None,
|
attn_mask: Tensor = None,
|
||||||
kv_cache: Optional[KVCache] = None,
|
kv_cache: Optional[KVCache] = None,
|
||||||
is_causal: bool = False,
|
is_causal: bool = False,
|
||||||
|
fwd: Optional[str] = None,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
bsz, seq_len, _ = x.size()
|
|
||||||
|
|
||||||
q = self.q_proj(x)
|
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_a_proj(x)
|
||||||
kv_compressed = self.kv_norm(kv_compressed)
|
kv_compressed = self.kv_norm(kv_compressed)
|
||||||
|
|
||||||
kv = self.kv_b_proj(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(
|
k_nope, k_rope, v = torch.split(
|
||||||
kv, [self.qk_nope_head_dim, self.qk_rope_head_dim, self.head_dim], dim=-1
|
kv, [self.qk_nope_head_dim, self.qk_rope_head_dim, self.head_dim], dim=-1
|
||||||
@@ -171,7 +170,9 @@ class MLA(nn.Module):
|
|||||||
q = self.q_norm(q)
|
q = self.q_norm(q)
|
||||||
k = self.k_norm(k)
|
k = self.k_norm(k)
|
||||||
|
|
||||||
attn_out = attention(q, k, v, kv_cache, self.layer_id, attn_mask, is_causal)
|
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:
|
if self.use_gated_attention:
|
||||||
attn_out = attn_out * F.sigmoid(self.gate(x))
|
attn_out = attn_out * F.sigmoid(self.gate(x))
|
||||||
|
|||||||
@@ -1,15 +1,21 @@
|
|||||||
from dataclasses import asdict
|
from dataclasses import asdict
|
||||||
from typing import Optional
|
from typing import Optional, TypedDict
|
||||||
|
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from astrai.inference.core.cache import KVCache
|
from astrai.inference.cache import KVCache
|
||||||
from astrai.model.components.attention import AttnFactory
|
from astrai.model.components.attention import AttnFactory
|
||||||
from astrai.model.components.mlp import FFNFactory
|
from astrai.model.components.mlp import FFNFactory, RouterStats
|
||||||
from astrai.model.components.norm import RMSNorm
|
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):
|
class DecoderBlock(nn.Module):
|
||||||
def __init__(self, config, layer_id: int):
|
def __init__(self, config, layer_id: int):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
@@ -26,7 +32,20 @@ class DecoderBlock(nn.Module):
|
|||||||
self.attention = AttnFactory.create(config.attn_type, **cfg, layer_id=layer_id)
|
self.attention = AttnFactory.create(config.attn_type, **cfg, layer_id=layer_id)
|
||||||
self.input_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
self.input_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
||||||
self.post_attention_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(
|
def forward(
|
||||||
self,
|
self,
|
||||||
@@ -35,15 +54,23 @@ class DecoderBlock(nn.Module):
|
|||||||
attention_mask: Optional[Tensor] = None,
|
attention_mask: Optional[Tensor] = None,
|
||||||
kv_cache: Optional[KVCache] = None,
|
kv_cache: Optional[KVCache] = None,
|
||||||
is_causal: bool = False,
|
is_causal: bool = False,
|
||||||
) -> Tensor:
|
fwd: Optional[str] = None,
|
||||||
|
) -> DecoderOutput:
|
||||||
attn_output = self.attention(
|
attn_output = self.attention(
|
||||||
self.input_norm(x),
|
self.input_norm(x),
|
||||||
rotary_emb,
|
rotary_emb,
|
||||||
attention_mask,
|
attention_mask,
|
||||||
kv_cache,
|
kv_cache,
|
||||||
is_causal,
|
is_causal,
|
||||||
|
fwd,
|
||||||
)
|
)
|
||||||
x = attn_output + x
|
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,3 +1,5 @@
|
|||||||
|
from typing import Optional, TypedDict
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
@@ -11,6 +13,22 @@ class FFNFactory(BaseFactory[nn.Module]):
|
|||||||
pass
|
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]
|
||||||
|
|
||||||
|
|
||||||
@FFNFactory.register("mlp")
|
@FFNFactory.register("mlp")
|
||||||
class MLP(nn.Module):
|
class MLP(nn.Module):
|
||||||
def __init__(self, dim: int, dim_ffn: int, down_init_std: float = 0.02):
|
def __init__(self, dim: int, dim_ffn: int, down_init_std: float = 0.02):
|
||||||
@@ -19,10 +37,10 @@ class MLP(nn.Module):
|
|||||||
self.gate = Linear(dim, dim_ffn)
|
self.gate = Linear(dim, dim_ffn)
|
||||||
self.down = Linear(dim_ffn, dim, init_std=down_init_std)
|
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))
|
gated = self.up(x) * F.silu(self.gate(x))
|
||||||
out = self.down(gated)
|
out = self.down(gated)
|
||||||
return out
|
return {"hidden_states": out, "aux_loss": None, "router_stats": None}
|
||||||
|
|
||||||
|
|
||||||
@FFNFactory.register("moe")
|
@FFNFactory.register("moe")
|
||||||
@@ -36,6 +54,9 @@ class DeepSeekMoE(nn.Module):
|
|||||||
n_activated_experts: int = 2,
|
n_activated_experts: int = 2,
|
||||||
topk_method: str = "greedy",
|
topk_method: str = "greedy",
|
||||||
n_layers: int = 1,
|
n_layers: int = 1,
|
||||||
|
moe_intermediate_size: Optional[int] = None,
|
||||||
|
shared_expert_intermediate_size: Optional[int] = None,
|
||||||
|
norm_topk_prob: bool = True,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.dim = dim
|
self.dim = dim
|
||||||
@@ -43,6 +64,16 @@ class DeepSeekMoE(nn.Module):
|
|||||||
self.n_shared_experts = n_shared_experts
|
self.n_shared_experts = n_shared_experts
|
||||||
self.n_activated_experts = n_activated_experts
|
self.n_activated_experts = n_activated_experts
|
||||||
self.topk_method = topk_method
|
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)
|
self.router = Linear(dim, n_routed_experts, bias=False)
|
||||||
moe_scale = 1 / max(n_shared_experts, 1) + 1 / n_activated_experts
|
moe_scale = 1 / max(n_shared_experts, 1) + 1 / n_activated_experts
|
||||||
@@ -50,51 +81,92 @@ class DeepSeekMoE(nn.Module):
|
|||||||
|
|
||||||
self.shared_experts = nn.ModuleList(
|
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)
|
for _ in range(n_shared_experts)
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
self.routed_experts = nn.ModuleList(
|
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)
|
for _ in range(n_routed_experts)
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
def forward(self, x: Tensor) -> Tensor:
|
def forward(self, x: Tensor) -> FFNOutput:
|
||||||
bsz, seq_len, dim = x.shape
|
include_aux_loss = self.training and torch.is_grad_enabled()
|
||||||
|
shape = x.shape
|
||||||
|
dim = shape[-1]
|
||||||
x_flat = x.view(-1, dim)
|
x_flat = x.view(-1, dim)
|
||||||
|
|
||||||
shared_out = self._shared_forward(x_flat)
|
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)
|
out = (shared_out + routed_output["hidden_states"]).view(shape)
|
||||||
return out
|
return {
|
||||||
|
"hidden_states": out,
|
||||||
|
"aux_loss": routed_output["aux_loss"],
|
||||||
|
"router_stats": routed_output["router_stats"],
|
||||||
|
}
|
||||||
|
|
||||||
def _shared_forward(self, x: Tensor) -> Tensor:
|
def _shared_forward(self, x: Tensor) -> Tensor:
|
||||||
if self.n_shared_experts == 0:
|
if self.n_shared_experts == 0:
|
||||||
return torch.zeros_like(x)
|
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) -> FFNOutput:
|
||||||
N, D = x.shape
|
N, D = x.shape
|
||||||
K = self.n_activated_experts
|
K = self.n_activated_experts
|
||||||
|
E = self.n_routed_experts
|
||||||
|
|
||||||
router_logits = self.router(x)
|
router_logits = self.router(x)
|
||||||
router_probs = torch.softmax(router_logits.float(), dim=-1).to(x.dtype)
|
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_indices = torch.topk(router_probs, K, dim=-1, sorted=False)
|
||||||
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
|
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)
|
output = torch.zeros(N, D, device=x.device, dtype=x.dtype)
|
||||||
for expert_idx in range(self.n_routed_experts):
|
start = 0
|
||||||
expert_mask = topk_indices == expert_idx
|
for expert_idx, end in enumerate(boundaries):
|
||||||
token_idx, k_idx = expert_mask.nonzero(as_tuple=True)
|
if end == start:
|
||||||
if token_idx.numel() == 0:
|
|
||||||
continue
|
continue
|
||||||
expert_input = x[token_idx]
|
expert_output = self.routed_experts[expert_idx](flat_tokens[start:end])[
|
||||||
expert_output = self.routed_experts[expert_idx](expert_input)
|
"hidden_states"
|
||||||
weights = topk_weights[token_idx, k_idx].unsqueeze(-1)
|
]
|
||||||
output.index_add_(0, token_idx, expert_output * weights)
|
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,
|
||||||
|
}
|
||||||
|
|||||||
@@ -65,9 +65,12 @@ class RotaryEmbedding(nn.Module):
|
|||||||
[batch, seq_len, dim/2, 2] (f32) — [cos, sin] pairs.
|
[batch, seq_len, dim/2, 2] (f32) — [cos, sin] pairs.
|
||||||
"""
|
"""
|
||||||
if position_ids is None:
|
if position_ids is None:
|
||||||
position_ids = (
|
if x.ndim == 2:
|
||||||
torch.arange(x.size(1), device=x.device)
|
position_ids = torch.arange(x.size(0), device=x.device)
|
||||||
.unsqueeze(0)
|
else:
|
||||||
.expand(x.size(0), -1)
|
position_ids = (
|
||||||
)
|
torch.arange(x.size(1), device=x.device)
|
||||||
|
.unsqueeze(0)
|
||||||
|
.expand(x.size(0), -1)
|
||||||
|
)
|
||||||
return self.freqs_cis[position_ids].float()
|
return self.freqs_cis[position_ids].float()
|
||||||
|
|||||||
@@ -70,7 +70,7 @@ class EmbeddingEncoder(AutoModel):
|
|||||||
attn_mask = process_attention_mask(input_mask)
|
attn_mask = process_attention_mask(input_mask)
|
||||||
|
|
||||||
for layer in self.layers:
|
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)
|
hidden_states = self.norm(x)
|
||||||
|
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ import torch.nn as nn
|
|||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
from astrai.inference.core.cache import KVCache
|
from astrai.inference.cache import KVCache
|
||||||
from astrai.model.automodel import AutoModel, ModelFactory
|
from astrai.model.automodel import AutoModel, ModelFactory
|
||||||
from astrai.model.components.decoder_block import DecoderBlock
|
from astrai.model.components.decoder_block import DecoderBlock
|
||||||
from astrai.model.components.embedding import Embedding
|
from astrai.model.components.embedding import Embedding
|
||||||
@@ -105,18 +105,48 @@ class AutoRegressiveLM(AutoModel):
|
|||||||
input_mask: Optional[Tensor] = None,
|
input_mask: Optional[Tensor] = None,
|
||||||
kv_cache: Optional[KVCache] = None,
|
kv_cache: Optional[KVCache] = None,
|
||||||
position_ids: Optional[Tensor] = None,
|
position_ids: Optional[Tensor] = None,
|
||||||
|
fwd: Optional[str] = None,
|
||||||
) -> Dict[str, Tensor]:
|
) -> 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)
|
x = self.embed_tokens(input_ids)
|
||||||
rotary_emb = self.rotary_embedding(x, position_ids)
|
rotary_emb = self.rotary_embedding(x, position_ids)
|
||||||
attn_mask = process_attention_mask(input_mask)
|
attn_mask = process_attention_mask(input_mask)
|
||||||
use_sdpa_causal_mask = attn_mask is None
|
use_sdpa_causal_mask = attn_mask is None
|
||||||
|
|
||||||
|
aux_losses = []
|
||||||
|
router_stats_list = []
|
||||||
for layer in self.layers:
|
for layer in self.layers:
|
||||||
x = layer(x, rotary_emb, attn_mask, kv_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)
|
hidden_states = self.norm(x)
|
||||||
logits = self.lm_head(hidden_states)
|
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
|
||||||
|
|||||||
@@ -30,15 +30,13 @@ def get_current_device():
|
|||||||
def get_world_size() -> int:
|
def get_world_size() -> int:
|
||||||
if dist.is_available() and dist.is_initialized():
|
if dist.is_available() and dist.is_initialized():
|
||||||
return dist.get_world_size()
|
return dist.get_world_size()
|
||||||
else:
|
return int(os.environ.get("WORLD_SIZE", "1"))
|
||||||
return 1
|
|
||||||
|
|
||||||
|
|
||||||
def get_rank() -> int:
|
def get_rank() -> int:
|
||||||
if dist.is_available() and dist.is_initialized():
|
if dist.is_available() and dist.is_initialized():
|
||||||
return dist.get_rank()
|
return dist.get_rank()
|
||||||
else:
|
return int(os.environ.get("RANK", "0"))
|
||||||
return 0
|
|
||||||
|
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
@@ -247,18 +245,15 @@ class LocalStrategy(LaunchStrategy):
|
|||||||
ctx.join()
|
ctx.join()
|
||||||
|
|
||||||
|
|
||||||
def _detect_launcher() -> str:
|
def _is_external_launcher() -> bool:
|
||||||
"""Detect the distributed launcher from environment.
|
"""Whether an external launcher (torchrun/elastic/manual env) started us."""
|
||||||
|
|
||||||
Returns one of: "torchelastic", "torchrun", "external", "local".
|
|
||||||
"""
|
|
||||||
if dist.is_torchelastic_launched():
|
if dist.is_torchelastic_launched():
|
||||||
return "torchelastic"
|
return True
|
||||||
if "LOCAL_WORLD_SIZE" in os.environ:
|
if "LOCAL_WORLD_SIZE" in os.environ:
|
||||||
return "torchrun"
|
return True
|
||||||
if "RANK" in os.environ and "WORLD_SIZE" in os.environ:
|
if "RANK" in os.environ and "WORLD_SIZE" in os.environ:
|
||||||
return "external"
|
return True
|
||||||
return "local"
|
return False
|
||||||
|
|
||||||
|
|
||||||
def spawn_parallel_fn(
|
def spawn_parallel_fn(
|
||||||
@@ -273,8 +268,7 @@ def spawn_parallel_fn(
|
|||||||
):
|
):
|
||||||
if master_port is None:
|
if master_port is None:
|
||||||
master_port = find_free_port()
|
master_port = find_free_port()
|
||||||
launcher = _detect_launcher()
|
if _is_external_launcher():
|
||||||
if launcher in ("torchelastic", "torchrun", "external"):
|
|
||||||
strategy = TorchrunStrategy(
|
strategy = TorchrunStrategy(
|
||||||
world_size, backend, master_addr, master_port, device_type, start_method
|
world_size, backend, master_addr, master_port, device_type, start_method
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -416,7 +416,11 @@ class MultiOutputMaskBuilder(BaseMaskBuilder):
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
result: dict = {}
|
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():
|
for output_key, spec in sources_spec.items():
|
||||||
sections = spec.get("sections", [])
|
sections = spec.get("sections", [])
|
||||||
@@ -428,7 +432,6 @@ class MultiOutputMaskBuilder(BaseMaskBuilder):
|
|||||||
if ids is None:
|
if ids is None:
|
||||||
continue
|
continue
|
||||||
result[output_key] = ids
|
result[output_key] = ids
|
||||||
any_output = True
|
|
||||||
continue
|
continue
|
||||||
|
|
||||||
list_field = spec.get("list_field", False)
|
list_field = spec.get("list_field", False)
|
||||||
@@ -444,7 +447,6 @@ class MultiOutputMaskBuilder(BaseMaskBuilder):
|
|||||||
result[output_key] = ids
|
result[output_key] = ids
|
||||||
if mask is not None:
|
if mask is not None:
|
||||||
result[mask_key] = mask
|
result[mask_key] = mask
|
||||||
any_output = True
|
|
||||||
continue
|
continue
|
||||||
|
|
||||||
ids, mask = self.renderer.process_sections(
|
ids, mask = self.renderer.process_sections(
|
||||||
@@ -460,9 +462,7 @@ class MultiOutputMaskBuilder(BaseMaskBuilder):
|
|||||||
elif "mask_key" in spec:
|
elif "mask_key" in spec:
|
||||||
result[mask_key] = mask
|
result[mask_key] = mask
|
||||||
|
|
||||||
any_output = True
|
if not required_outputs or not required_outputs.issubset(result):
|
||||||
|
|
||||||
if not any_output:
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
result["domain"] = _extract_domain(item, config.output.domain_key)
|
result["domain"] = _extract_domain(item, config.output.domain_key)
|
||||||
@@ -474,6 +474,11 @@ class MultiOutputMaskBuilder(BaseMaskBuilder):
|
|||||||
return [None] * len(items)
|
return [None] * len(items)
|
||||||
|
|
||||||
results = [{} for _ in 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():
|
for output_key, spec in sources_spec.items():
|
||||||
sections = spec.get("sections", [])
|
sections = spec.get("sections", [])
|
||||||
if not sections:
|
if not sections:
|
||||||
@@ -506,7 +511,7 @@ class MultiOutputMaskBuilder(BaseMaskBuilder):
|
|||||||
|
|
||||||
return [
|
return [
|
||||||
({**result, "domain": _extract_domain(item, config.output.domain_key)})
|
({**result, "domain": _extract_domain(item, config.output.domain_key)})
|
||||||
if result
|
if required_outputs and required_outputs.issubset(result)
|
||||||
else None
|
else None
|
||||||
for item, result in zip(items, results)
|
for item, result in zip(items, results)
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -22,9 +22,21 @@ from astrai.serialization.dataset import (
|
|||||||
load_bin_offsets,
|
load_bin_offsets,
|
||||||
save_bin,
|
save_bin,
|
||||||
)
|
)
|
||||||
|
from astrai.serialization.hf_adapter import (
|
||||||
|
HF_MODEL_TYPES,
|
||||||
|
adapt_config,
|
||||||
|
convert_hf_config,
|
||||||
|
convert_hf_weights,
|
||||||
|
looks_like_hf_state_dict,
|
||||||
|
)
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"Checkpoint",
|
"Checkpoint",
|
||||||
|
"HF_MODEL_TYPES",
|
||||||
|
"adapt_config",
|
||||||
|
"convert_hf_config",
|
||||||
|
"convert_hf_weights",
|
||||||
|
"looks_like_hf_state_dict",
|
||||||
"load_json",
|
"load_json",
|
||||||
"load_model_config",
|
"load_model_config",
|
||||||
"load_model_weights",
|
"load_model_weights",
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ import json
|
|||||||
import time
|
import time
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Dict, Optional, Union
|
from typing import Any, Callable, Dict, Optional, Union
|
||||||
|
|
||||||
import safetensors.torch as st
|
import safetensors.torch as st
|
||||||
import torch
|
import torch
|
||||||
@@ -22,39 +22,31 @@ def save_safetensors(state_dict: dict, path: Union[str, Path]):
|
|||||||
st.save_file(state_dict, str(path))
|
st.save_file(state_dict, str(path))
|
||||||
|
|
||||||
|
|
||||||
def load_safetensors(path: Union[str, Path], broadcast: bool = False) -> dict:
|
def _broadcast_load(loader: Callable[[], dict], broadcast: bool) -> dict:
|
||||||
|
"""Load on rank 0 and broadcast the object to all ranks."""
|
||||||
if not broadcast or not dist.is_initialized():
|
if not broadcast or not dist.is_initialized():
|
||||||
return st.load_file(str(path))
|
return loader()
|
||||||
|
|
||||||
rank = get_rank()
|
rank = get_rank()
|
||||||
if rank == 0:
|
if rank == 0:
|
||||||
state_dict = st.load_file(str(path))
|
data = loader()
|
||||||
else:
|
else:
|
||||||
state_dict = {}
|
data = {}
|
||||||
tmp = [state_dict]
|
tmp = [data]
|
||||||
dist.broadcast_object_list(tmp, src=0)
|
dist.broadcast_object_list(tmp, src=0)
|
||||||
return tmp[0]
|
return tmp[0]
|
||||||
|
|
||||||
|
|
||||||
|
def load_safetensors(path: Union[str, Path], broadcast: bool = False) -> dict:
|
||||||
|
return _broadcast_load(lambda: st.load_file(str(path)), broadcast)
|
||||||
|
|
||||||
|
|
||||||
def save_json(data: dict, path: Union[str, Path]):
|
def save_json(data: dict, path: Union[str, Path]):
|
||||||
with open(str(path), "w") as f:
|
with open(str(path), "w") as f:
|
||||||
json.dump(data, f, indent=2)
|
json.dump(data, f, indent=2)
|
||||||
|
|
||||||
|
|
||||||
def load_json(path: Union[str, Path], broadcast: bool = False) -> dict:
|
def load_json(path: Union[str, Path], broadcast: bool = False) -> dict:
|
||||||
if not broadcast or not dist.is_initialized():
|
return _broadcast_load(lambda: json.loads(Path(path).read_text()), broadcast)
|
||||||
with open(str(path), "r") as f:
|
|
||||||
return json.load(f)
|
|
||||||
|
|
||||||
rank = get_rank()
|
|
||||||
if rank == 0:
|
|
||||||
with open(str(path), "r") as f:
|
|
||||||
data = json.load(f)
|
|
||||||
else:
|
|
||||||
data = {}
|
|
||||||
tmp = [data]
|
|
||||||
dist.broadcast_object_list(tmp, src=0)
|
|
||||||
return tmp[0]
|
|
||||||
|
|
||||||
|
|
||||||
def save_torch(obj: Any, path: Union[str, Path]):
|
def save_torch(obj: Any, path: Union[str, Path]):
|
||||||
@@ -99,7 +91,21 @@ def load_model_config(save_directory: str) -> dict:
|
|||||||
|
|
||||||
|
|
||||||
def load_model_weights(save_directory: str) -> dict:
|
def load_model_weights(save_directory: str) -> dict:
|
||||||
return load_state_dict(Path(save_directory) / _WEIGHTS_FILE)
|
save_path = Path(save_directory)
|
||||||
|
weights_file = save_path / _WEIGHTS_FILE
|
||||||
|
if weights_file.exists():
|
||||||
|
return load_state_dict(weights_file)
|
||||||
|
|
||||||
|
index_path = save_path / "model.safetensors.index.json"
|
||||||
|
if index_path.exists():
|
||||||
|
index = load_json(index_path)
|
||||||
|
weight_map = index.get("weight_map", {})
|
||||||
|
state_dict = {}
|
||||||
|
for shard in sorted(set(weight_map.values())):
|
||||||
|
state_dict.update(load_state_dict(save_path / shard))
|
||||||
|
return state_dict
|
||||||
|
|
||||||
|
raise FileNotFoundError(f"No model weights found in {save_directory}")
|
||||||
|
|
||||||
|
|
||||||
def load_state_dict(path: Union[str, Path], broadcast: bool = False) -> dict:
|
def load_state_dict(path: Union[str, Path], broadcast: bool = False) -> dict:
|
||||||
@@ -190,8 +196,10 @@ class Checkpoint:
|
|||||||
if meta_path.exists():
|
if meta_path.exists():
|
||||||
return cls.load(save_dir, broadcast=broadcast)
|
return cls.load(save_dir, broadcast=broadcast)
|
||||||
|
|
||||||
if weights_path.exists():
|
weights_path = save_path / _WEIGHTS_FILE
|
||||||
state_dict = load_state_dict(weights_path, broadcast=broadcast)
|
index_path = save_path / "model.safetensors.index.json"
|
||||||
|
if weights_path.exists() or index_path.exists():
|
||||||
|
state_dict = load_model_weights(save_dir)
|
||||||
config = {}
|
config = {}
|
||||||
config_path = save_path / _CONFIG_FILE
|
config_path = save_path / _CONFIG_FILE
|
||||||
if config_path.exists():
|
if config_path.exists():
|
||||||
|
|||||||
@@ -0,0 +1,271 @@
|
|||||||
|
"""HuggingFace checkpoint adaptation for LLaMA-style decoder models.
|
||||||
|
|
||||||
|
AstrAI stores weights with its own key names (``layers.<i>.input_norm``,
|
||||||
|
``layers.<i>.mlp.gate``), while HuggingFace decoder-only checkpoints use
|
||||||
|
``model.layers.<i>.input_layernorm`` / ``model.layers.<i>.mlp.gate_proj``.
|
||||||
|
This module translates HF configs and state dicts so external checkpoints
|
||||||
|
can be loaded directly.
|
||||||
|
|
||||||
|
Supported families (LLaMA layout, dense and MoE):
|
||||||
|
- dense FFN: llama, mistral, qwen2, gemma, gemma2, phi3
|
||||||
|
- MoE FFN (Mixtral / Qwen2-MoE / DeepSeek-V3 layout): router
|
||||||
|
``mlp.gate``, routed experts ``mlp.experts.<j>``, shared experts
|
||||||
|
``mlp.shared_experts.<j>``
|
||||||
|
|
||||||
|
Not supported:
|
||||||
|
- MLA attention (DeepSeek-V2/V3 ``kv_a_proj_with_mqa``) uses a different
|
||||||
|
KV factorization and cannot be converted numerically.
|
||||||
|
- Attention/MLP bias (``attention_bias`` / ``mlp_bias``) — AstrAI
|
||||||
|
projections are bias-free.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import re
|
||||||
|
from typing import Any, Dict, Mapping
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.config.base import BaseConfig
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
HF_MODEL_TYPES = frozenset(
|
||||||
|
{
|
||||||
|
"llama",
|
||||||
|
"mistral",
|
||||||
|
"mixtral",
|
||||||
|
"qwen2",
|
||||||
|
"qwen2_moe",
|
||||||
|
"gemma",
|
||||||
|
"gemma2",
|
||||||
|
"phi3",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
_EMBED = re.compile(r"^model\.embed_tokens\.weight$")
|
||||||
|
_ATTN = re.compile(r"^model\.layers\.(\d+)\.self_attn\.(q|k|v|o)_proj\.(weight|bias)$")
|
||||||
|
_Q_NORM = re.compile(r"^model\.layers\.(\d+)\.self_attn\.q_norm\.weight$")
|
||||||
|
_K_NORM = re.compile(r"^model\.layers\.(\d+)\.self_attn\.k_norm\.weight$")
|
||||||
|
_INPUT_NORM = re.compile(r"^model\.layers\.(\d+)\.input_layernorm\.weight$")
|
||||||
|
_POST_NORM = re.compile(r"^model\.layers\.(\d+)\.post_attention_layernorm\.weight$")
|
||||||
|
_FINAL_NORM = re.compile(r"^model\.norm\.weight$")
|
||||||
|
_LM_HEAD = re.compile(r"^lm_head\.weight$")
|
||||||
|
_DENSE_MLP = re.compile(
|
||||||
|
r"^model\.layers\.(\d+)\.mlp\.(gate|up|down)_proj\.(weight|bias)$"
|
||||||
|
)
|
||||||
|
_MOE_ROUTER = re.compile(r"^model\.layers\.(\d+)\.mlp\.gate\.weight$")
|
||||||
|
_MOE_EXPERTS = re.compile(
|
||||||
|
r"^model\.layers\.(\d+)\.mlp\.experts\.(\d+)\.(gate|up|down)_proj\.(weight|bias)$"
|
||||||
|
)
|
||||||
|
_MOE_SHARED = re.compile(
|
||||||
|
r"^model\.layers\.(\d+)\.mlp\.shared_expert(?:s)?\.(\d+)\."
|
||||||
|
r"(gate|up|down)_proj\.(weight|bias)$"
|
||||||
|
)
|
||||||
|
|
||||||
|
_ASTR_PREFIXES = ("embed_tokens.", "layers.", "norm.", "lm_head.")
|
||||||
|
|
||||||
|
|
||||||
|
def looks_like_hf_state_dict(state_dict: Mapping[str, Any]) -> bool:
|
||||||
|
"""Return True if *state_dict* uses HuggingFace key names."""
|
||||||
|
return any(
|
||||||
|
key.startswith("model.")
|
||||||
|
or "self_attn." in key
|
||||||
|
or "input_layernorm" in key
|
||||||
|
or "mlp.experts." in key
|
||||||
|
for key in state_dict
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _is_dense_mlp_layer(config: BaseConfig, layer_id: int) -> bool:
|
||||||
|
"""Return whether a layer uses dense MLP instead of routed experts."""
|
||||||
|
if getattr(config, "ffn_type", "mlp") != "moe":
|
||||||
|
return True
|
||||||
|
mlp_only = getattr(config, "mlp_only_layers", None) or []
|
||||||
|
if layer_id in mlp_only:
|
||||||
|
return True
|
||||||
|
step = getattr(config, "decoder_sparse_step", 1) or 1
|
||||||
|
return step > 1 and (layer_id + 1) % step != 0
|
||||||
|
|
||||||
|
|
||||||
|
def adapt_config(raw: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
|
"""Translate *raw* for AstrAI if it looks like an HF model config."""
|
||||||
|
if raw.get("model_type") in HF_MODEL_TYPES:
|
||||||
|
return convert_hf_config(raw)
|
||||||
|
return raw
|
||||||
|
|
||||||
|
|
||||||
|
def convert_hf_config(raw: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
|
"""Convert an HF LLaMA-style config dict to AstrAI field names."""
|
||||||
|
if raw.get("attention_bias") or raw.get("mlp_bias"):
|
||||||
|
raise NotImplementedError(
|
||||||
|
"attention_bias / mlp_bias checkpoints are not supported; "
|
||||||
|
"AstrAI projections are bias-free"
|
||||||
|
)
|
||||||
|
|
||||||
|
cfg: Dict[str, Any] = {}
|
||||||
|
for key in (
|
||||||
|
"vocab_size",
|
||||||
|
"hidden_size",
|
||||||
|
"num_hidden_layers",
|
||||||
|
"intermediate_size",
|
||||||
|
"rms_norm_eps",
|
||||||
|
"tie_word_embeddings",
|
||||||
|
"max_position_embeddings",
|
||||||
|
"rope_theta",
|
||||||
|
"rope_scaling",
|
||||||
|
"num_attention_heads",
|
||||||
|
"num_key_value_heads",
|
||||||
|
"use_qk_norm",
|
||||||
|
"use_gated_attention",
|
||||||
|
"kv_lora_rank",
|
||||||
|
"qk_nope_head_dim",
|
||||||
|
"qk_rope_head_dim",
|
||||||
|
"moe_intermediate_size",
|
||||||
|
"shared_expert_intermediate_size",
|
||||||
|
"topk_method",
|
||||||
|
"norm_topk_prob",
|
||||||
|
"moe_aux_loss_coef",
|
||||||
|
"decoder_sparse_step",
|
||||||
|
"mlp_only_layers",
|
||||||
|
"neftune_alpha",
|
||||||
|
):
|
||||||
|
if key in raw:
|
||||||
|
cfg[key] = raw[key]
|
||||||
|
|
||||||
|
if "qk_norm" in raw and "use_qk_norm" not in cfg:
|
||||||
|
cfg["use_qk_norm"] = raw["qk_norm"]
|
||||||
|
if (
|
||||||
|
raw.get("model_type") in ("gemma", "gemma2")
|
||||||
|
and "use_qk_norm" not in cfg
|
||||||
|
and "qk_norm" not in raw
|
||||||
|
):
|
||||||
|
# Gemma/Gemma2 always apply RMSNorm to Q and K before attention.
|
||||||
|
cfg["use_qk_norm"] = True
|
||||||
|
|
||||||
|
n_heads = raw.get("num_attention_heads")
|
||||||
|
if cfg.get("num_key_value_heads") is None and n_heads is not None:
|
||||||
|
cfg["num_key_value_heads"] = n_heads
|
||||||
|
|
||||||
|
if raw.get("head_dim") is not None and n_heads and raw.get("hidden_size"):
|
||||||
|
expected = raw["hidden_size"] // n_heads
|
||||||
|
if raw["head_dim"] != expected:
|
||||||
|
raise NotImplementedError(
|
||||||
|
f"HF head_dim={raw['head_dim']} differs from the computed "
|
||||||
|
f"head dim {expected}; AstrAI derives head_dim from "
|
||||||
|
"hidden_size / num_attention_heads"
|
||||||
|
)
|
||||||
|
|
||||||
|
if "kv_lora_rank" in raw:
|
||||||
|
cfg["attn_type"] = "mla"
|
||||||
|
|
||||||
|
n_experts = raw.get("num_local_experts") or raw.get("n_routed_experts")
|
||||||
|
if n_experts:
|
||||||
|
cfg["ffn_type"] = "moe"
|
||||||
|
cfg["n_routed_experts"] = n_experts
|
||||||
|
if "num_experts_per_tok" in raw:
|
||||||
|
cfg["n_activated_experts"] = raw["num_experts_per_tok"]
|
||||||
|
if "n_activated_experts" in raw:
|
||||||
|
cfg["n_activated_experts"] = raw["n_activated_experts"]
|
||||||
|
if "n_shared_experts" in raw:
|
||||||
|
cfg["n_shared_experts"] = raw["n_shared_experts"]
|
||||||
|
else:
|
||||||
|
# Mixtral has no shared experts; AstrAI defaults to one.
|
||||||
|
cfg["n_shared_experts"] = 0
|
||||||
|
if cfg.get("moe_intermediate_size") is None and "intermediate_size" in raw:
|
||||||
|
# MoE configs store the per-expert FFN size in intermediate_size.
|
||||||
|
cfg["moe_intermediate_size"] = raw["intermediate_size"]
|
||||||
|
first_k_dense = raw.get("first_k_dense_replace")
|
||||||
|
if isinstance(first_k_dense, int) and first_k_dense > 0:
|
||||||
|
cfg["mlp_only_layers"] = list(range(first_k_dense))
|
||||||
|
cfg["decoder_sparse_step"] = 1
|
||||||
|
|
||||||
|
cfg["model_type"] = "autoregressive_lm"
|
||||||
|
return cfg
|
||||||
|
|
||||||
|
|
||||||
|
def convert_hf_weights(
|
||||||
|
state_dict: Mapping[str, Any],
|
||||||
|
config: BaseConfig,
|
||||||
|
) -> Dict[str, torch.Tensor]:
|
||||||
|
"""Rename HF state dict keys to AstrAI names.
|
||||||
|
|
||||||
|
Keys that are already AstrAI-style pass through unchanged; unmapped
|
||||||
|
HF keys are dropped with a warning. Use with ``strict=True`` to fail
|
||||||
|
loudly when the checkpoint does not match the config.
|
||||||
|
"""
|
||||||
|
if getattr(config, "attn_type", "gqa") == "mla":
|
||||||
|
if any("kv_a_proj_with_mqa" in key for key in state_dict):
|
||||||
|
raise NotImplementedError(
|
||||||
|
"MLA attention (DeepSeek-V2/V3 kv_a_proj_with_mqa) uses a "
|
||||||
|
"different KV factorization and cannot be converted"
|
||||||
|
)
|
||||||
|
|
||||||
|
ffn_type = getattr(config, "ffn_type", "mlp")
|
||||||
|
converted: Dict[str, torch.Tensor] = {}
|
||||||
|
skipped: list[str] = []
|
||||||
|
for key, tensor in state_dict.items():
|
||||||
|
if key.startswith(_ASTR_PREFIXES):
|
||||||
|
converted[key] = tensor
|
||||||
|
continue
|
||||||
|
|
||||||
|
new_key = None
|
||||||
|
if ffn_type == "moe":
|
||||||
|
m = _MOE_ROUTER.match(key)
|
||||||
|
if m:
|
||||||
|
new_key = f"layers.{m.group(1)}.mlp.router.weight"
|
||||||
|
else:
|
||||||
|
m = _MOE_EXPERTS.match(key)
|
||||||
|
if m:
|
||||||
|
new_key = (
|
||||||
|
f"layers.{m.group(1)}.mlp.routed_experts.{m.group(2)}."
|
||||||
|
f"{m.group(3)}.{m.group(4)}"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
m = _MOE_SHARED.match(key)
|
||||||
|
if m:
|
||||||
|
new_key = (
|
||||||
|
f"layers.{m.group(1)}.mlp.shared_experts.{m.group(2)}."
|
||||||
|
f"{m.group(3)}.{m.group(4)}"
|
||||||
|
)
|
||||||
|
if new_key is None:
|
||||||
|
m = _DENSE_MLP.match(key)
|
||||||
|
if m and _is_dense_mlp_layer(config, int(m.group(1))):
|
||||||
|
new_key = f"layers.{m.group(1)}.mlp.{m.group(2)}.{m.group(3)}"
|
||||||
|
else:
|
||||||
|
m = _DENSE_MLP.match(key)
|
||||||
|
if m:
|
||||||
|
new_key = f"layers.{m.group(1)}.mlp.{m.group(2)}.{m.group(3)}"
|
||||||
|
|
||||||
|
if new_key is None:
|
||||||
|
m = _ATTN.match(key)
|
||||||
|
if m:
|
||||||
|
new_key = (
|
||||||
|
f"layers.{m.group(1)}.attention.{m.group(2)}_proj.{m.group(3)}"
|
||||||
|
)
|
||||||
|
elif (m := _Q_NORM.match(key)) is not None:
|
||||||
|
new_key = f"layers.{m.group(1)}.attention.q_norm.weight"
|
||||||
|
elif (m := _K_NORM.match(key)) is not None:
|
||||||
|
new_key = f"layers.{m.group(1)}.attention.k_norm.weight"
|
||||||
|
elif (m := _INPUT_NORM.match(key)) is not None:
|
||||||
|
new_key = f"layers.{m.group(1)}.input_norm.weight"
|
||||||
|
elif (m := _POST_NORM.match(key)) is not None:
|
||||||
|
new_key = f"layers.{m.group(1)}.post_attention_norm.weight"
|
||||||
|
elif (m := _EMBED.match(key)) is not None:
|
||||||
|
new_key = "embed_tokens.weight"
|
||||||
|
elif (m := _FINAL_NORM.match(key)) is not None:
|
||||||
|
new_key = "norm.weight"
|
||||||
|
elif (m := _LM_HEAD.match(key)) is not None:
|
||||||
|
new_key = "lm_head.weight"
|
||||||
|
|
||||||
|
if new_key is None:
|
||||||
|
skipped.append(key)
|
||||||
|
else:
|
||||||
|
converted[new_key] = tensor
|
||||||
|
|
||||||
|
if skipped:
|
||||||
|
logger.warning(
|
||||||
|
"Dropped %d unmapped HuggingFace weight key(s): %s",
|
||||||
|
len(skipped),
|
||||||
|
", ".join(sorted(skipped)[:10]),
|
||||||
|
)
|
||||||
|
return converted
|
||||||
@@ -1,3 +1,4 @@
|
|||||||
|
import math
|
||||||
from typing import Dict
|
from typing import Dict
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -27,6 +28,8 @@ class GradSNRTracker:
|
|||||||
|
|
||||||
SNR = E[g]^2 / Var(g) = E[g]^2 / (E[g^2] - E[g]^2)
|
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.
|
The tracker accumulates per-parameter EMA moments across optimizer steps.
|
||||||
Call ``update`` after backward (before ``optimizer.step``) and read
|
Call ``update`` after backward (before ``optimizer.step``) and read
|
||||||
``snr`` to get the aggregate SNR across all parameters.
|
``snr`` to get the aggregate SNR across all parameters.
|
||||||
@@ -64,7 +67,8 @@ class GradSNRTracker:
|
|||||||
noise = (v - m.pow(2)).clamp(min=0).sum().item()
|
noise = (v - m.pow(2)).clamp(min=0).sum().item()
|
||||||
total_signal += signal
|
total_signal += signal
|
||||||
total_noise += noise
|
total_noise += noise
|
||||||
return total_signal / (total_noise + self.eps)
|
snr = total_signal / (total_noise + self.eps)
|
||||||
|
return 10.0 * math.log10(max(snr, self.eps))
|
||||||
|
|
||||||
|
|
||||||
def ctx_get_loss(ctx):
|
def ctx_get_loss(ctx):
|
||||||
@@ -88,3 +92,7 @@ def ctx_get_grad_snr(ctx):
|
|||||||
if tracker is None:
|
if tracker is None:
|
||||||
return None
|
return None
|
||||||
return tracker.snr
|
return tracker.snr
|
||||||
|
|
||||||
|
|
||||||
|
def ctx_get_moe_metric(ctx, key):
|
||||||
|
return ctx.strategy._moe_metrics.get(key)
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ Provides:
|
|||||||
- :class:`BaseRewardModel` — pluggable reward interface
|
- :class:`BaseRewardModel` — pluggable reward interface
|
||||||
- :class:`RolloutGenerator` — KV-cache-backed generation of grouped
|
- :class:`RolloutGenerator` — KV-cache-backed generation of grouped
|
||||||
responses + decoding (no reward); delegates the generation loop to
|
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
|
so rollout and the production inference server share one code path
|
||||||
- :class:`RolloutRunner` — orchestrates generation + scoring with a
|
- :class:`RolloutRunner` — orchestrates generation + scoring with a
|
||||||
step-driven cache; its ``__call__`` returns ``(RolloutResult, is_fresh)``
|
step-driven cache; its ``__call__`` returns ``(RolloutResult, is_fresh)``
|
||||||
@@ -20,7 +20,7 @@ from typing import Dict, List, Optional, Tuple
|
|||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from astrai.inference.core.scheduler import InferenceScheduler
|
from astrai.inference.scheduler import InferenceScheduler
|
||||||
|
|
||||||
|
|
||||||
@dataclass(kw_only=True)
|
@dataclass(kw_only=True)
|
||||||
@@ -101,7 +101,7 @@ class RolloutGenerator:
|
|||||||
"""Pure generation + decoding for a group of responses per prompt.
|
"""Pure generation + decoding for a group of responses per prompt.
|
||||||
|
|
||||||
Delegates the prefill/decode loop to
|
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
|
which uses a real KV cache (no O(n²) recompute). Has no dependency
|
||||||
on any reward model; can be reused in isolation for offline
|
on any reward model; can be reused in isolation for offline
|
||||||
generation, qualitative sampling, or eval pipelines.
|
generation, qualitative sampling, or eval pipelines.
|
||||||
|
|||||||
@@ -2,7 +2,7 @@
|
|||||||
|
|
||||||
import math
|
import math
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from typing import Any, Dict, List
|
from typing import List
|
||||||
|
|
||||||
from torch.optim.lr_scheduler import LRScheduler
|
from torch.optim.lr_scheduler import LRScheduler
|
||||||
|
|
||||||
@@ -20,12 +20,6 @@ class BaseScheduler(LRScheduler, ABC):
|
|||||||
"""Calculate the current learning rate."""
|
"""Calculate the current learning rate."""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
def state_dict(self) -> Dict[str, Any]:
|
|
||||||
return super().state_dict()
|
|
||||||
|
|
||||||
def load_state_dict(self, state_dict: Dict[str, Any]):
|
|
||||||
super().load_state_dict(state_dict)
|
|
||||||
|
|
||||||
|
|
||||||
class SchedulerFactory(BaseFactory["BaseScheduler"]):
|
class SchedulerFactory(BaseFactory["BaseScheduler"]):
|
||||||
"""Factory class for creating learning rate schedulers.
|
"""Factory class for creating learning rate schedulers.
|
||||||
|
|||||||
+193
-36
@@ -1,7 +1,7 @@
|
|||||||
"""Training strategy implementations with factory pattern."""
|
"""Training strategy implementations with factory pattern."""
|
||||||
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC
|
||||||
from typing import Callable, Dict, Union
|
from typing import Callable, Dict, List, Optional, TypedDict, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
@@ -9,10 +9,22 @@ import torch.nn.functional as F
|
|||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.model.components.mlp import RouterStats
|
||||||
from astrai.parallel.executor import broadcast_state_dict
|
from astrai.parallel.executor import broadcast_state_dict
|
||||||
from astrai.trainer.rollout import RolloutResult
|
from astrai.trainer.rollout import RolloutResult
|
||||||
|
|
||||||
|
|
||||||
|
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]:
|
def move_to_device(batch: Dict[str, Tensor], device: str) -> Dict[str, Tensor]:
|
||||||
"""Move batch tensors to specified device with non-blocking transfer."""
|
"""Move batch tensors to specified device with non-blocking transfer."""
|
||||||
return {key: value.to(device, non_blocking=True) for key, value in batch.items()}
|
return {key: value.to(device, non_blocking=True) for key, value in batch.items()}
|
||||||
@@ -24,7 +36,7 @@ def get_logprobs(
|
|||||||
attn_mask: Tensor,
|
attn_mask: Tensor,
|
||||||
loss_mask: Tensor,
|
loss_mask: Tensor,
|
||||||
reduction: str,
|
reduction: str,
|
||||||
) -> Tensor:
|
) -> LogprobsOutput:
|
||||||
"""Compute token-wise log probabilities from model outputs.
|
"""Compute token-wise log probabilities from model outputs.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -46,10 +58,11 @@ def get_logprobs(
|
|||||||
shifted_input_ids = input_ids[:, 1:]
|
shifted_input_ids = input_ids[:, 1:]
|
||||||
shifted_loss_mask = loss_mask[:, 1:]
|
shifted_loss_mask = loss_mask[:, 1:]
|
||||||
|
|
||||||
logits = model(
|
outputs = model(
|
||||||
input_ids[:, :-1],
|
input_ids[:, :-1],
|
||||||
attn_mask[:, :, :-1, :-1] if attn_mask.dim() == 4 else attn_mask[:, :-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)
|
log_probs = torch.log_softmax(logits.float(), dim=-1)
|
||||||
|
|
||||||
token_logprobs = torch.gather(
|
token_logprobs = torch.gather(
|
||||||
@@ -57,13 +70,18 @@ def get_logprobs(
|
|||||||
).squeeze(-1)
|
).squeeze(-1)
|
||||||
|
|
||||||
if reduction == "mean":
|
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
|
dim=-1
|
||||||
).clamp(min=1.0)
|
) / shifted_loss_mask.sum(dim=-1).clamp(min=1.0)
|
||||||
elif reduction == "sum":
|
elif reduction == "sum":
|
||||||
return (token_logprobs * shifted_loss_mask).sum(dim=-1)
|
logprobs = (token_logprobs * shifted_loss_mask).sum(dim=-1)
|
||||||
else:
|
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:
|
def make_doc_boundary_mask(position_ids: Tensor) -> Tensor:
|
||||||
@@ -82,6 +100,68 @@ def make_doc_boundary_mask(position_ids: Tensor) -> Tensor:
|
|||||||
return (same_doc & causal).unsqueeze(1)
|
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):
|
class BaseStrategy(ABC):
|
||||||
"""Abstract base class for training strategies.
|
"""Abstract base class for training strategies.
|
||||||
|
|
||||||
@@ -102,10 +182,11 @@ class BaseStrategy(ABC):
|
|||||||
self.model = model
|
self.model = model
|
||||||
self.device = device
|
self.device = device
|
||||||
self.executor = kwargs.pop("executor", None)
|
self.executor = kwargs.pop("executor", None)
|
||||||
self.extra_kwargs = kwargs
|
self.moe_aux_loss_coef = kwargs.pop("moe_aux_loss_coef", 0.01)
|
||||||
|
self._moe_metrics: Dict[str, float] = {}
|
||||||
|
self.strategy_kwargs = kwargs
|
||||||
self._rollout_runner = None
|
self._rollout_runner = None
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||||
"""Compute loss for the given batch.
|
"""Compute loss for the given batch.
|
||||||
|
|
||||||
@@ -115,7 +196,36 @@ class BaseStrategy(ABC):
|
|||||||
Returns:
|
Returns:
|
||||||
Computed loss tensor
|
Computed loss tensor
|
||||||
"""
|
"""
|
||||||
raise NotImplementedError
|
return self.compute_loss_output(batch)["loss"]
|
||||||
|
|
||||||
|
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
|
||||||
|
return self._normalize_output(self.compute_loss(batch))
|
||||||
|
|
||||||
|
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:
|
def supports_online(self) -> bool:
|
||||||
"""Whether this strategy can operate with a rollout runner.
|
"""Whether this strategy can operate with a rollout runner.
|
||||||
@@ -148,22 +258,36 @@ class BaseStrategy(ABC):
|
|||||||
"""
|
"""
|
||||||
pass
|
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):
|
def on_optimizer_step(self):
|
||||||
"""Advance online rollout state after a successful optimizer step."""
|
"""Advance online rollout state after a successful optimizer step."""
|
||||||
if self._rollout_runner is not None:
|
if self._rollout_runner is not None:
|
||||||
self._rollout_runner.step()
|
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."""
|
"""Run offline or online forward depending on runner injection."""
|
||||||
if self._rollout_runner is None:
|
if self._rollout_runner is None:
|
||||||
return self.compute_loss(batch)
|
return self.compute_loss_output(batch)
|
||||||
|
|
||||||
result, is_fresh = self._rollout_runner(batch)
|
result, is_fresh = self._rollout_runner(batch)
|
||||||
if is_fresh:
|
if is_fresh:
|
||||||
self._on_rollout_refresh()
|
self._on_rollout_refresh()
|
||||||
|
|
||||||
train_batch = self.prepare_from_rollout(result)
|
train_batch = self.prepare_from_rollout(result)
|
||||||
return self.compute_loss(train_batch)
|
return self.compute_loss_output(train_batch)
|
||||||
|
|
||||||
|
|
||||||
class StrategyFactory(BaseFactory["BaseStrategy"]):
|
class StrategyFactory(BaseFactory["BaseStrategy"]):
|
||||||
@@ -190,6 +314,7 @@ class SEQStrategy(BaseStrategy):
|
|||||||
"""Standard next-token prediction training strategy.
|
"""Standard next-token prediction training strategy.
|
||||||
|
|
||||||
Computes cross-entropy loss for next token prediction.
|
Computes cross-entropy loss for next token prediction.
|
||||||
|
Optionally adds MoE load balancing auxiliary loss.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -202,10 +327,11 @@ class SEQStrategy(BaseStrategy):
|
|||||||
super().__init__(model, device, **kwargs)
|
super().__init__(model, device, **kwargs)
|
||||||
self.label_smoothing = label_smoothing
|
self.label_smoothing = label_smoothing
|
||||||
|
|
||||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
|
||||||
batch = move_to_device(batch, self.device)
|
batch = move_to_device(batch, self.device)
|
||||||
input_ids, target_ids = batch["input_ids"], batch["target_ids"]
|
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(
|
loss = F.cross_entropy(
|
||||||
input=logits.flatten(0, 1).float(),
|
input=logits.flatten(0, 1).float(),
|
||||||
@@ -213,7 +339,12 @@ class SEQStrategy(BaseStrategy):
|
|||||||
label_smoothing=self.label_smoothing,
|
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")
|
@StrategyFactory.register("sft")
|
||||||
@@ -221,6 +352,7 @@ class SFTStrategy(BaseStrategy):
|
|||||||
"""Supervised Fine-tuning strategy with loss masking.
|
"""Supervised Fine-tuning strategy with loss masking.
|
||||||
|
|
||||||
Applies cross-entropy loss only to tokens where loss_mask is True.
|
Applies cross-entropy loss only to tokens where loss_mask is True.
|
||||||
|
Optionally adds MoE load balancing auxiliary loss.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -233,7 +365,7 @@ class SFTStrategy(BaseStrategy):
|
|||||||
super().__init__(model, device, **kwargs)
|
super().__init__(model, device, **kwargs)
|
||||||
self.label_smoothing = label_smoothing
|
self.label_smoothing = label_smoothing
|
||||||
|
|
||||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
|
||||||
batch = move_to_device(batch, self.device)
|
batch = move_to_device(batch, self.device)
|
||||||
input_ids, target_ids, position_ids, loss_mask = (
|
input_ids, target_ids, position_ids, loss_mask = (
|
||||||
batch["input_ids"],
|
batch["input_ids"],
|
||||||
@@ -245,9 +377,10 @@ class SFTStrategy(BaseStrategy):
|
|||||||
ignore_index = -100
|
ignore_index = -100
|
||||||
input_mask = make_doc_boundary_mask(position_ids)
|
input_mask = make_doc_boundary_mask(position_ids)
|
||||||
target_ids = target_ids.masked_fill(~loss_mask, ignore_index)
|
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
|
input_ids=input_ids, position_ids=position_ids, input_mask=input_mask
|
||||||
)["logits"]
|
)
|
||||||
|
logits = outputs["logits"]
|
||||||
|
|
||||||
loss = F.cross_entropy(
|
loss = F.cross_entropy(
|
||||||
input=logits.flatten(0, 1).float(),
|
input=logits.flatten(0, 1).float(),
|
||||||
@@ -256,7 +389,12 @@ class SFTStrategy(BaseStrategy):
|
|||||||
label_smoothing=self.label_smoothing,
|
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")
|
@StrategyFactory.register("dpo")
|
||||||
@@ -281,7 +419,7 @@ class DPOStrategy(BaseStrategy):
|
|||||||
self.beta = beta
|
self.beta = beta
|
||||||
self.reduction = reduction
|
self.reduction = reduction
|
||||||
|
|
||||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
|
||||||
batch = move_to_device(batch, self.device)
|
batch = move_to_device(batch, self.device)
|
||||||
chosen_ids, rejected_ids = batch["chosen"], batch["rejected"]
|
chosen_ids, rejected_ids = batch["chosen"], batch["rejected"]
|
||||||
chosen_mask, rejected_mask = batch["chosen_mask"], batch["rejected_mask"]
|
chosen_mask, rejected_mask = batch["chosen_mask"], batch["rejected_mask"]
|
||||||
@@ -297,22 +435,25 @@ class DPOStrategy(BaseStrategy):
|
|||||||
)[None, None, :, :] # [1, 1, S, S]
|
)[None, None, :, :] # [1, 1, S, S]
|
||||||
full_mask = key_pad & causal # [B*2, 1, S, S] — composed
|
full_mask = key_pad & causal # [B*2, 1, S, S] — composed
|
||||||
|
|
||||||
log_pi = get_logprobs(
|
policy_output = get_logprobs(
|
||||||
self.model,
|
self.model,
|
||||||
concat_ids,
|
concat_ids,
|
||||||
full_mask,
|
full_mask,
|
||||||
concat_loss_mask,
|
concat_loss_mask,
|
||||||
self.reduction,
|
self.reduction,
|
||||||
)
|
)
|
||||||
|
log_pi = policy_output["logprobs"]
|
||||||
|
aux_loss = policy_output["aux_loss"]
|
||||||
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
log_ref = get_logprobs(
|
ref_output = get_logprobs(
|
||||||
self.ref_model,
|
self.ref_model,
|
||||||
concat_ids,
|
concat_ids,
|
||||||
full_mask,
|
full_mask,
|
||||||
concat_loss_mask,
|
concat_loss_mask,
|
||||||
self.reduction,
|
self.reduction,
|
||||||
)
|
)
|
||||||
|
log_ref = ref_output["logprobs"]
|
||||||
|
|
||||||
log_pi_chosen = log_pi[: chosen_ids.shape[0]]
|
log_pi_chosen = log_pi[: chosen_ids.shape[0]]
|
||||||
log_pi_rejected = log_pi[chosen_ids.shape[0] :]
|
log_pi_rejected = log_pi[chosen_ids.shape[0] :]
|
||||||
@@ -325,7 +466,12 @@ class DPOStrategy(BaseStrategy):
|
|||||||
ratio_diff = pi_log_ratio - ref_log_ratio
|
ratio_diff = pi_log_ratio - ref_log_ratio
|
||||||
dpo_loss = -F.logsigmoid(self.beta * ratio_diff).mean()
|
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:
|
def supports_online(self) -> bool:
|
||||||
return True
|
return True
|
||||||
@@ -397,7 +543,7 @@ class GRPOStrategy(BaseStrategy):
|
|||||||
if state_dict is not None:
|
if state_dict is not None:
|
||||||
self.old_model.load_state_dict(state_dict)
|
self.old_model.load_state_dict(state_dict)
|
||||||
|
|
||||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput:
|
||||||
batch = move_to_device(batch, self.device)
|
batch = move_to_device(batch, self.device)
|
||||||
prompts = batch["prompts"]
|
prompts = batch["prompts"]
|
||||||
responses = batch["responses"]
|
responses = batch["responses"]
|
||||||
@@ -438,16 +584,23 @@ class GRPOStrategy(BaseStrategy):
|
|||||||
# get_logprobs returns [B*G, S-1] (S = prompt_len + response_len).
|
# get_logprobs returns [B*G, S-1] (S = prompt_len + response_len).
|
||||||
# Response token logprobs occupy the last ``response_len`` positions
|
# Response token logprobs occupy the last ``response_len`` positions
|
||||||
# (the first response token is predicted from the last prompt token).
|
# (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"
|
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():
|
with torch.no_grad():
|
||||||
token_log_probs_old = get_logprobs(
|
old_output = get_logprobs(
|
||||||
self.old_model, full_sequences, attn_mask, full_masks, "none"
|
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"
|
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]
|
# Reshape to [B, G, response_len]
|
||||||
token_log_probs_policy = token_log_probs_policy.view(batch_size, group_size, -1)
|
token_log_probs_policy = token_log_probs_policy.view(batch_size, group_size, -1)
|
||||||
@@ -480,9 +633,13 @@ class GRPOStrategy(BaseStrategy):
|
|||||||
kl_per_token = r - torch.log(r + eps) - 1.0
|
kl_per_token = r - torch.log(r + eps) - 1.0
|
||||||
kl_penalty = self.kl_coef * (kl_per_token * token_masks).sum() / token_count
|
kl_penalty = self.kl_coef * (kl_per_token * token_masks).sum() / token_count
|
||||||
|
|
||||||
total_loss = policy_loss + kl_penalty
|
task_loss = policy_loss + kl_penalty
|
||||||
|
return self._loss_output(
|
||||||
return total_loss
|
task_loss,
|
||||||
|
{"policy_loss": policy_loss, "kl_loss": kl_penalty},
|
||||||
|
aux_loss,
|
||||||
|
policy_output.get("router_stats"),
|
||||||
|
)
|
||||||
|
|
||||||
def supports_online(self) -> bool:
|
def supports_online(self) -> bool:
|
||||||
return True
|
return True
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ import logging
|
|||||||
import os
|
import os
|
||||||
import sys
|
import sys
|
||||||
import time
|
import time
|
||||||
|
from functools import partial
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import IO, Callable, List, Optional, Protocol, runtime_checkable
|
from typing import IO, Callable, List, Optional, Protocol, runtime_checkable
|
||||||
|
|
||||||
@@ -21,6 +22,7 @@ from astrai.trainer.metric_util import (
|
|||||||
ctx_get_grad_snr,
|
ctx_get_grad_snr,
|
||||||
ctx_get_loss,
|
ctx_get_loss,
|
||||||
ctx_get_lr,
|
ctx_get_lr,
|
||||||
|
ctx_get_moe_metric,
|
||||||
ctx_get_val_loss,
|
ctx_get_val_loss,
|
||||||
)
|
)
|
||||||
from astrai.trainer.train_context import TrainContext
|
from astrai.trainer.train_context import TrainContext
|
||||||
@@ -257,14 +259,40 @@ class MetricCallback(TrainCallback):
|
|||||||
"val_loss": ctx_get_val_loss,
|
"val_loss": ctx_get_val_loss,
|
||||||
"grad_norm": ctx_get_grad_norm,
|
"grad_norm": ctx_get_grad_norm,
|
||||||
"grad_snr": ctx_get_grad_snr,
|
"grad_snr": ctx_get_grad_snr,
|
||||||
|
"moe_aux_loss": partial(ctx_get_moe_metric, key="aux_loss"),
|
||||||
|
"router_entropy": partial(ctx_get_moe_metric, key="router_entropy"),
|
||||||
|
"dead_expert_fraction": partial(
|
||||||
|
ctx_get_moe_metric, key="dead_expert_fraction"
|
||||||
|
),
|
||||||
|
"load_imbalance_mean": partial(
|
||||||
|
ctx_get_moe_metric, key="load_imbalance_mean"
|
||||||
|
),
|
||||||
|
"load_imbalance_max": partial(ctx_get_moe_metric, key="load_imbalance_max"),
|
||||||
}
|
}
|
||||||
|
|
||||||
def _metrics(self, context: TrainContext, names):
|
def _metrics(self, context: TrainContext, names):
|
||||||
return {
|
metrics = dict(context.metrics)
|
||||||
m: self._metric_funcs[m](context)
|
for name in names:
|
||||||
for m in names
|
metric_fn = self._metric_funcs.get(name)
|
||||||
if self._metric_funcs[m](context) is not None
|
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)
|
@only_on_rank(0)
|
||||||
def _append(self, event_type: str, context: TrainContext, **extra):
|
def _append(self, event_type: str, context: TrainContext, **extra):
|
||||||
@@ -286,8 +314,8 @@ class MetricCallback(TrainCallback):
|
|||||||
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
for batch in context.val_dataloader:
|
for batch in context.val_dataloader:
|
||||||
loss = context.strategy(batch)
|
loss_output = context.strategy(batch)
|
||||||
total_loss += loss.item()
|
total_loss += loss_output["loss"].item()
|
||||||
num_batches += 1
|
num_batches += 1
|
||||||
|
|
||||||
if context.world_size > 1 and dist.is_initialized():
|
if context.world_size > 1 and dist.is_initialized():
|
||||||
|
|||||||
+197
-151
@@ -8,14 +8,21 @@ import torch
|
|||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from torch.utils.data import DataLoader, random_split
|
from torch.utils.data import DataLoader, random_split
|
||||||
|
|
||||||
|
from astrai.config.model_config import ConfigFactory
|
||||||
from astrai.config.train_config import TrainConfig
|
from astrai.config.train_config import TrainConfig
|
||||||
from astrai.dataset import RDSampler
|
from astrai.dataset import RDSampler
|
||||||
from astrai.inference.core.scheduler import InferenceScheduler
|
from astrai.inference.scheduler import InferenceScheduler
|
||||||
from astrai.model.components.lora import inject_lora
|
from astrai.model.components.lora import inject_lora
|
||||||
from astrai.parallel.executor import BaseExecutor, ExecutorFactory, create_ref_model
|
from astrai.parallel.executor import BaseExecutor, ExecutorFactory, create_ref_model
|
||||||
from astrai.parallel.setup import get_current_device, get_rank, get_world_size
|
from astrai.parallel.setup import get_current_device, get_rank, get_world_size
|
||||||
from astrai.protocols import OptimizerProtocol, SchedulerProtocol
|
from astrai.protocols import OptimizerProtocol, SchedulerProtocol
|
||||||
from astrai.serialization import Checkpoint, load_json
|
from astrai.serialization import (
|
||||||
|
Checkpoint,
|
||||||
|
adapt_config,
|
||||||
|
convert_hf_weights,
|
||||||
|
load_json,
|
||||||
|
looks_like_hf_state_dict,
|
||||||
|
)
|
||||||
from astrai.tokenize import AutoTokenizer
|
from astrai.tokenize import AutoTokenizer
|
||||||
from astrai.trainer.metric_util import GradSNRTracker
|
from astrai.trainer.metric_util import GradSNRTracker
|
||||||
from astrai.trainer.rollout import RolloutGenerator, RolloutRunner
|
from astrai.trainer.rollout import RolloutGenerator, RolloutRunner
|
||||||
@@ -38,6 +45,7 @@ class TrainContext:
|
|||||||
epoch: int = field(default=0)
|
epoch: int = field(default=0)
|
||||||
consumed_samples: int = field(default=0)
|
consumed_samples: int = field(default=0)
|
||||||
loss: float = field(default=0.0)
|
loss: float = field(default=0.0)
|
||||||
|
metrics: Dict[str, float] = field(default_factory=dict)
|
||||||
grad_norm: Optional[float] = field(default=None)
|
grad_norm: Optional[float] = field(default=None)
|
||||||
grad_snr_tracker: GradSNRTracker = field(default_factory=GradSNRTracker)
|
grad_snr_tracker: GradSNRTracker = field(default_factory=GradSNRTracker)
|
||||||
val_dataloader: Optional[DataLoader] = field(default=None)
|
val_dataloader: Optional[DataLoader] = field(default=None)
|
||||||
@@ -65,6 +73,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:
|
class TrainContextBuilder:
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -80,212 +97,241 @@ class TrainContextBuilder:
|
|||||||
return self
|
return self
|
||||||
|
|
||||||
def build(self) -> TrainContext:
|
def build(self) -> TrainContext:
|
||||||
cfg = self.config
|
# Resolve persisted state.
|
||||||
device = get_current_device()
|
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,
|
cfg.parallel_mode,
|
||||||
grad_accum_steps=cfg.grad_accum_steps,
|
grad_accum_steps=cfg.grad_accum_steps,
|
||||||
**cfg.executor_kwargs,
|
**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:
|
if self._param_path:
|
||||||
config_path = Path(self._param_path) / "config.json"
|
config_path = Path(self._param_path) / "config.json"
|
||||||
if config_path.exists():
|
if config_path.exists():
|
||||||
model_config = load_json(config_path)
|
state.model_config = adapt_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:
|
|
||||||
checkpoint = Checkpoint.load_any(self._param_path)
|
checkpoint = Checkpoint.load_any(self._param_path)
|
||||||
if checkpoint is not None:
|
if checkpoint is not None:
|
||||||
preloaded_state_dict = checkpoint.state_dict
|
|
||||||
if checkpoint.config:
|
if checkpoint.config:
|
||||||
model_config = checkpoint.config
|
checkpoint.config = adapt_config(checkpoint.config)
|
||||||
|
if checkpoint.state_dict and looks_like_hf_state_dict(
|
||||||
|
checkpoint.state_dict
|
||||||
|
):
|
||||||
|
checkpoint.state_dict = convert_hf_weights(
|
||||||
|
checkpoint.state_dict,
|
||||||
|
ConfigFactory.load(checkpoint.config or state.model_config),
|
||||||
|
)
|
||||||
|
state.state_dict = checkpoint.state_dict
|
||||||
|
state.model_config = checkpoint.config or state.model_config
|
||||||
if self._resume:
|
if self._resume:
|
||||||
preloaded_epoch = checkpoint.epoch
|
state.epoch = checkpoint.epoch
|
||||||
per_step = (
|
per_step = (
|
||||||
cfg.batch_per_device * get_world_size() * cfg.grad_accum_steps
|
cfg.batch_per_device * get_world_size() * cfg.grad_accum_steps
|
||||||
)
|
)
|
||||||
preloaded_consumed = (
|
state.consumed_samples = (
|
||||||
checkpoint.consumed_samples // per_step
|
checkpoint.consumed_samples // per_step * per_step
|
||||||
) * per_step
|
)
|
||||||
preloaded_checkpoint = checkpoint
|
state.checkpoint = checkpoint
|
||||||
|
if not state.model_config:
|
||||||
|
model = cfg.model_fn()
|
||||||
|
if hasattr(model, "config"):
|
||||||
|
state.model_config = model.config.to_dict()
|
||||||
|
return state
|
||||||
|
|
||||||
if not model_config and hasattr(cfg.model_fn(), "config"):
|
def _create_context(
|
||||||
model_config = cfg.model_fn().config.to_dict()
|
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):
|
def _prepare_model(
|
||||||
m = m.to(device=device)
|
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:
|
if cfg.lora is not None:
|
||||||
inject_lora(
|
inject_lora(
|
||||||
m,
|
model,
|
||||||
r=cfg.lora.r,
|
r=cfg.lora.r,
|
||||||
alpha=cfg.lora.alpha,
|
alpha=cfg.lora.alpha,
|
||||||
target_modules=set(cfg.lora.target_modules),
|
target_modules=set(cfg.lora.target_modules),
|
||||||
)
|
)
|
||||||
if preloaded_state_dict is not None:
|
if state.state_dict is not None:
|
||||||
m.load_state_dict(preloaded_state_dict, strict=False)
|
model.load_state_dict(state.state_dict, strict=False)
|
||||||
return m
|
return model
|
||||||
|
|
||||||
def _after_wrap(m):
|
def after_wrap(model):
|
||||||
if cfg.compile_mode is not None:
|
if cfg.compile_mode is not None:
|
||||||
logger.info("torch.compile enabled (mode=%s)", cfg.compile_mode)
|
logger.info("torch.compile enabled (mode=%s)", cfg.compile_mode)
|
||||||
m = torch.compile(m, mode=cfg.compile_mode)
|
model = torch.compile(model, mode=cfg.compile_mode)
|
||||||
return m
|
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,
|
|
||||||
)
|
|
||||||
|
|
||||||
context.model, context.optimizer, context.scheduler = executor.prepare(
|
context.model, context.optimizer, context.scheduler = executor.prepare(
|
||||||
cfg.model_fn,
|
cfg.model_fn,
|
||||||
cfg.optimizer_fn,
|
cfg.optimizer_fn,
|
||||||
cfg.scheduler_fn,
|
cfg.scheduler_fn,
|
||||||
before_wrap=_before_wrap,
|
before_wrap=before_wrap,
|
||||||
after_wrap=_after_wrap,
|
after_wrap=after_wrap,
|
||||||
)
|
)
|
||||||
|
|
||||||
train_dataset = cfg.dataset
|
def _get_datasets(self):
|
||||||
val_dataset = cfg.val_dataset
|
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:
|
def _create_dataloaders(
|
||||||
n_total = len(cfg.dataset)
|
self, context: TrainContext, train_dataset, val_dataset
|
||||||
n_val = max(1, int(n_total * cfg.val_split))
|
) -> None:
|
||||||
n_train = n_total - n_val
|
sampler_offset = context.consumed_samples // context.world_size
|
||||||
generator = torch.Generator().manual_seed(cfg.random_seed)
|
if self._resume and sampler_offset > 0:
|
||||||
train_dataset, val_dataset = random_split(
|
samples_per_replica = (
|
||||||
cfg.dataset, [n_train, n_val], generator=generator
|
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
|
||||||
if self._resume and sampler_offset > 0:
|
):
|
||||||
offset = context.world_size - 1
|
cfg = self.config
|
||||||
num_samples_per_replica = (
|
|
||||||
len(train_dataset) + offset
|
|
||||||
) // context.world_size
|
|
||||||
if num_samples_per_replica > 0:
|
|
||||||
context.epoch = sampler_offset // num_samples_per_replica
|
|
||||||
|
|
||||||
sampler = RDSampler(
|
sampler = RDSampler(
|
||||||
data_source=train_dataset,
|
dataset,
|
||||||
start_epoch=context.epoch,
|
start_epoch=epoch,
|
||||||
start_iter=sampler_offset,
|
start_iter=start_iter,
|
||||||
seed=cfg.random_seed,
|
seed=cfg.random_seed,
|
||||||
|
shuffle=shuffle,
|
||||||
)
|
)
|
||||||
context.dataloader = DataLoader(
|
loader_kwargs = dict(
|
||||||
train_dataset,
|
dataset=dataset,
|
||||||
batch_size=cfg.batch_per_device,
|
batch_size=cfg.batch_per_device,
|
||||||
sampler=sampler,
|
sampler=sampler,
|
||||||
num_workers=cfg.num_workers,
|
num_workers=cfg.num_workers,
|
||||||
pin_memory=cfg.pin_memory,
|
pin_memory=cfg.pin_memory,
|
||||||
prefetch_factor=cfg.prefetch_factor,
|
|
||||||
collate_fn=cfg.collate_fn,
|
collate_fn=cfg.collate_fn,
|
||||||
)
|
)
|
||||||
|
# PyTorch rejects prefetch_factor/persistent_workers when workers=0.
|
||||||
if val_dataset is not None:
|
if cfg.num_workers > 0:
|
||||||
val_sampler = RDSampler(
|
loader_kwargs["persistent_workers"] = cfg.persistent_workers
|
||||||
data_source=val_dataset,
|
if cfg.prefetch_factor is not None:
|
||||||
start_epoch=0,
|
loader_kwargs["prefetch_factor"] = cfg.prefetch_factor
|
||||||
start_iter=0,
|
return DataLoader(
|
||||||
seed=cfg.random_seed,
|
**loader_kwargs,
|
||||||
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",
|
|
||||||
)
|
)
|
||||||
needs_old = cfg.strategy in ("grpo", "online_grpo")
|
|
||||||
|
|
||||||
if needs_ref:
|
def _restore_optimizer_state(self, context: TrainContext) -> None:
|
||||||
strategy_kwargs["ref_model"] = create_ref_model(
|
if context.checkpoint and context.checkpoint.extra:
|
||||||
cfg.model_fn, executor=executor, model=context.model, device=device
|
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.strategy_kwargs)
|
||||||
|
kwargs.setdefault("moe_aux_loss_coef", cfg.moe_aux_loss_coef)
|
||||||
|
if cfg.strategy in ("dpo", "grpo", "online_grpo", "online_dpo"):
|
||||||
|
kwargs["ref_model"] = create_ref_model(
|
||||||
|
cfg.model_fn,
|
||||||
|
executor=executor,
|
||||||
|
model=context.model,
|
||||||
|
device=get_current_device(),
|
||||||
)
|
)
|
||||||
|
if cfg.strategy in ("grpo", "online_grpo"):
|
||||||
if needs_old:
|
kwargs["old_model"] = create_ref_model(
|
||||||
strategy_kwargs["old_model"] = create_ref_model(
|
cfg.model_fn,
|
||||||
cfg.model_fn, executor=executor, model=context.model, device=device
|
executor=executor,
|
||||||
|
model=context.model,
|
||||||
|
device=get_current_device(),
|
||||||
)
|
)
|
||||||
|
|
||||||
context.strategy = StrategyFactory.create(
|
context.strategy = StrategyFactory.create(
|
||||||
cfg.strategy,
|
cfg.strategy,
|
||||||
model=context.model,
|
model=context.model,
|
||||||
device=device,
|
device=get_current_device(),
|
||||||
executor=executor,
|
executor=executor,
|
||||||
**strategy_kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
|
return kwargs
|
||||||
|
|
||||||
# Enable online rollout when the train_type is an ``online_*`` variant.
|
def _configure_rollout(self, context: TrainContext, strategy_kwargs: dict) -> None:
|
||||||
is_online = cfg.strategy.startswith("online_")
|
cfg = self.config
|
||||||
if is_online:
|
if not cfg.strategy.startswith("online_"):
|
||||||
if not context.strategy.supports_online():
|
return
|
||||||
raise ValueError(
|
if not context.strategy.supports_online():
|
||||||
f"Strategy '{cfg.strategy}' does not support online rollout"
|
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,
|
|
||||||
)
|
)
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(self._param_path)
|
||||||
generator = RolloutGenerator(
|
group_size = strategy_kwargs.get("group_size", 1)
|
||||||
scheduler=scheduler,
|
scheduler = InferenceScheduler(
|
||||||
tokenizer=tokenizer,
|
model=context.model,
|
||||||
max_tokens=cfg.rollout_max_tokens,
|
tokenizer=tokenizer,
|
||||||
group_size=group_size,
|
max_batch_size=group_size * max(1, cfg.batch_per_device),
|
||||||
temperature=cfg.rollout_temperature,
|
max_seq_len=getattr(context.model.config, "max_position_embeddings", None),
|
||||||
top_k=cfg.rollout_top_k,
|
)
|
||||||
top_p=cfg.rollout_top_p,
|
generator = RolloutGenerator(
|
||||||
)
|
scheduler=scheduler,
|
||||||
runner = RolloutRunner(
|
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,
|
generator=generator,
|
||||||
reward_model=reward_model,
|
reward_model=cfg.reward_model_fn(),
|
||||||
rollout_interval=cfg.rollout_interval,
|
rollout_interval=cfg.rollout_interval,
|
||||||
)
|
)
|
||||||
context.strategy.set_rollout_runner(runner)
|
)
|
||||||
|
|
||||||
return context
|
|
||||||
|
|||||||
@@ -82,9 +82,10 @@ class Trainer:
|
|||||||
break
|
break
|
||||||
with executor.accumulate(context.model):
|
with executor.accumulate(context.model):
|
||||||
self._call_callbacks("on_batch_begin", context)
|
self._call_callbacks("on_batch_begin", context)
|
||||||
loss = context.strategy(batch)
|
loss_output = context.strategy(batch)
|
||||||
context.loss = loss.item()
|
context.loss = loss_output["loss"].item()
|
||||||
stand_loss = loss / executor.grad_accum_steps
|
context.metrics = loss_output["metrics"]
|
||||||
|
stand_loss = loss_output["loss"] / executor.grad_accum_steps
|
||||||
executor.backward(stand_loss)
|
executor.backward(stand_loss)
|
||||||
context.consumed_samples += (
|
context.consumed_samples += (
|
||||||
context.config.batch_per_device * context.world_size
|
context.config.batch_per_device * context.world_size
|
||||||
|
|||||||
@@ -0,0 +1,108 @@
|
|||||||
|
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}")
|
||||||
|
|
||||||
|
# Kernel registry — parallel lists of module names (.so / pybind names,
|
||||||
|
# globally unique across families) and their per-family source paths under
|
||||||
|
# kernels/. `loader.py` auto-discovers the .so files in astrai/extension/lib/,
|
||||||
|
# so this CMake registry is the single place to register a new kernel.
|
||||||
|
#
|
||||||
|
# FP8 MMA instructions require sm_89+. Keep the target out of the build on
|
||||||
|
# older architectures instead of instantiating templates that cannot compile.
|
||||||
|
# The remaining kernels are still useful on sm_80+ (including sm_86).
|
||||||
|
set(KERNEL_NAMES
|
||||||
|
attn_decode
|
||||||
|
attn_prefill
|
||||||
|
attn_paged_decode
|
||||||
|
attn_paged_prefill
|
||||||
|
rotary_emb
|
||||||
|
)
|
||||||
|
set(KERNEL_SRCS
|
||||||
|
attention/decode.cu
|
||||||
|
attention/prefill.cu
|
||||||
|
attention/paged_decode.cu
|
||||||
|
attention/paged_prefill.cu
|
||||||
|
rotary/rotary_emb.cu
|
||||||
|
)
|
||||||
|
|
||||||
|
if(ASTRAI_CUDA_ARCH GREATER_EQUAL 89)
|
||||||
|
list(APPEND KERNEL_NAMES fp8_ops)
|
||||||
|
list(APPEND KERNEL_SRCS fp8/ops.cu)
|
||||||
|
else()
|
||||||
|
message(WARNING
|
||||||
|
"FP8 operator disabled: ASTRAI_CUDA_ARCH=${ASTRAI_CUDA_ARCH} "
|
||||||
|
"requires compute capability 89 or newer")
|
||||||
|
endif()
|
||||||
|
|
||||||
|
list(LENGTH KERNEL_NAMES _kernel_count)
|
||||||
|
math(EXPR _kernel_last "${_kernel_count} - 1")
|
||||||
|
foreach(i RANGE ${_kernel_last})
|
||||||
|
list(GET KERNEL_NAMES ${i} name)
|
||||||
|
list(GET KERNEL_SRCS ${i} src)
|
||||||
|
add_library(${name} MODULE "${CMAKE_CURRENT_SOURCE_DIR}/kernels/${src}")
|
||||||
|
|
||||||
|
target_compile_definitions(${name} PRIVATE TORCH_EXTENSION_NAME=${name})
|
||||||
|
|
||||||
|
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})
|
||||||
|
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
-1
@@ -1,2 +1,2 @@
|
|||||||
# Source directory for CUDA kernels — build-time only.
|
# Source directory for CUDA kernels — build-time only.
|
||||||
# Compiled .so files live in astrAI/_ext/.
|
# Compiled .so files live in astrai/extension/lib/ (see csrc/CMakeLists.txt).
|
||||||
|
|||||||
@@ -1,75 +0,0 @@
|
|||||||
from pathlib import Path
|
|
||||||
|
|
||||||
|
|
||||||
def cuda_toolkit_version() -> tuple[int, int] | None:
|
|
||||||
"""Return ``(major, minor)`` of the nvcc on PATH, or ``None``.
|
|
||||||
|
|
||||||
Used by ``setup.py`` to detect nvcc/torch CUDA version mismatches
|
|
||||||
(e.g. nvcc 13.0 with a cu128 torch wheel) which cause cryptic ABI errors.
|
|
||||||
"""
|
|
||||||
import shutil
|
|
||||||
import subprocess
|
|
||||||
|
|
||||||
nvcc = shutil.which("nvcc")
|
|
||||||
if nvcc is None:
|
|
||||||
return None
|
|
||||||
try:
|
|
||||||
out = subprocess.check_output(
|
|
||||||
[nvcc, "--version"], stderr=subprocess.STDOUT, text=True
|
|
||||||
)
|
|
||||||
for line in out.splitlines():
|
|
||||||
if "release" in line:
|
|
||||||
ver = line.split("release")[1].split(",")[0].strip()
|
|
||||||
major, minor = ver.split(".")
|
|
||||||
return (int(major), int(minor))
|
|
||||||
except Exception:
|
|
||||||
pass
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
def _arch_flags() -> list[str]:
|
|
||||||
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=16",
|
|
||||||
]
|
|
||||||
|
|
||||||
|
|
||||||
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")
|
|
||||||
register("rotary_emb")
|
|
||||||
@@ -0,0 +1,100 @@
|
|||||||
|
#pragma once
|
||||||
|
|
||||||
|
// Pure POD header
|
||||||
|
|
||||||
|
namespace astrai {
|
||||||
|
namespace attention {
|
||||||
|
|
||||||
|
// Tensor layout for Q/K/V tensors passed to attention kernels.
|
||||||
|
// Internally, kernels always operate on BHLD [batch, n_heads, seq_len, head_dim].
|
||||||
|
// When the caller passes BLHD, dims 1 and 2 are transposed at entry.
|
||||||
|
enum TensorLayout : int {
|
||||||
|
BHLD = 0, // [batch, n_heads, seq_len, head_dim]
|
||||||
|
BLHD = 1, // [batch, seq_len, n_heads, head_dim]
|
||||||
|
};
|
||||||
|
|
||||||
|
// Split-KV workspace cap: max decode splits per (batch, q_head).
|
||||||
|
constexpr int MAX_SPLITS = 32;
|
||||||
|
|
||||||
|
// Paged-prefill host Q-tile granularity in q rows: one q_tile_to_index unit
|
||||||
|
// covers this many query rows of one request. Must match Q_TILE_ROWS in
|
||||||
|
// astrai/inference/workspace.py, which builds the device-side tile maps.
|
||||||
|
constexpr int HOST_Q_TILE_ROWS = 64;
|
||||||
|
|
||||||
|
|
||||||
|
// Unified attention params covering BOTH addressing modes:
|
||||||
|
// - Contiguous K/V: dense [batch, kv_head, kv_len, head_dim] tensors (k/v).
|
||||||
|
// - Paged (SGLang-style): flat pool [size, kv_head, head_dim] + req_to_token.
|
||||||
|
// Each kernel selects the addressing via a KVSource policy (see
|
||||||
|
// layout_policies.cuh); a given call only touches the fields of one mode, so
|
||||||
|
// this is a POD shared by both paths rather than two parallel structs that
|
||||||
|
// drift out of sync.
|
||||||
|
//
|
||||||
|
// Pointer/flag members carry default member initializers: the pointers gate
|
||||||
|
// optional paths via null checks (new_k_ptr, mask, o_part, ...), so a stack
|
||||||
|
// `AttentionParams<T> p;` left partially packed must never see garbage
|
||||||
|
// non-null pointers or a garbage use_mask/causal_offset — that class of bug
|
||||||
|
// reads through wild addresses. NSDMI keeps the struct an aggregate (C++17)
|
||||||
|
// and trivially copyable, so `= {}`, memcpy-style packing and by-value kernel
|
||||||
|
// params all behave exactly as before.
|
||||||
|
template<typename T, typename AT = float>
|
||||||
|
struct AttentionParams {
|
||||||
|
// Shape
|
||||||
|
int batch;
|
||||||
|
int q_head;
|
||||||
|
int kv_head;
|
||||||
|
int head_dim;
|
||||||
|
int q_len; // Per-request in contiguous mode; total_q in paged mode.
|
||||||
|
int kv_len; // Contiguous mode; paged mode uses kv_indptr.
|
||||||
|
|
||||||
|
// Attention behavior
|
||||||
|
float scale;
|
||||||
|
// -1 = non-causal; >=0 = absolute position of first Q token
|
||||||
|
int causal_offset = -1;
|
||||||
|
int use_mask = 0;
|
||||||
|
|
||||||
|
// pointers
|
||||||
|
const T* __restrict__ q_ptr = nullptr;
|
||||||
|
const T* __restrict__ k_ptr = nullptr;
|
||||||
|
const T* __restrict__ v_ptr = nullptr;
|
||||||
|
const T* __restrict__ new_k_ptr = nullptr;
|
||||||
|
const T* __restrict__ new_v_ptr = nullptr;
|
||||||
|
T* __restrict__ o_ptr = nullptr;
|
||||||
|
const bool* __restrict__ mask = nullptr;
|
||||||
|
|
||||||
|
// strides
|
||||||
|
int q_b_stride;
|
||||||
|
int q_h_stride;
|
||||||
|
int q_l_stride;
|
||||||
|
int q_d_stride;
|
||||||
|
|
||||||
|
int kv_b_stride;
|
||||||
|
int kv_h_stride;
|
||||||
|
int kv_l_stride;
|
||||||
|
int kv_d_stride;
|
||||||
|
|
||||||
|
int new_kv_b_stride;
|
||||||
|
int new_kv_h_stride;
|
||||||
|
|
||||||
|
int mask_b_stride;
|
||||||
|
int mask_h_stride;
|
||||||
|
int mask_l_stride;
|
||||||
|
|
||||||
|
// Paged K/V addressing
|
||||||
|
const int* __restrict__ req_to_token = nullptr; // [num_reqs, max_context_len]
|
||||||
|
const int* __restrict__ req_pool_indices = nullptr; // [batch]
|
||||||
|
const int* __restrict__ kv_indptr = nullptr; // [batch + 1]
|
||||||
|
const int* __restrict__ qo_indptr = nullptr; // [batch + 1] or nullptr for decode
|
||||||
|
const int* __restrict__ q_tile_to_batch = nullptr; // [num_q_tiles], prefill only
|
||||||
|
const int* __restrict__ q_tile_to_index = nullptr; // [num_q_tiles], prefill only
|
||||||
|
int num_q_tiles;
|
||||||
|
int max_context_len; // req_to_token stride (dim 1)
|
||||||
|
|
||||||
|
// Decode split-KV workspace
|
||||||
|
int num_splits;
|
||||||
|
AT* __restrict__ o_part = nullptr;
|
||||||
|
AT* __restrict__ ml_part = nullptr;
|
||||||
|
};
|
||||||
|
|
||||||
|
} // namespace attention
|
||||||
|
} // namespace astrai
|
||||||
@@ -0,0 +1,65 @@
|
|||||||
|
#include "dispatchers.cuh"
|
||||||
|
#include "entry_utils.cuh"
|
||||||
|
|
||||||
|
using namespace astrai::attention;
|
||||||
|
|
||||||
|
torch::Tensor attn_decode(
|
||||||
|
torch::Tensor q,
|
||||||
|
torch::Tensor k,
|
||||||
|
torch::Tensor v,
|
||||||
|
c10::optional<torch::Tensor> mask,
|
||||||
|
int64_t causal_offset,
|
||||||
|
double scale,
|
||||||
|
int64_t layout,
|
||||||
|
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 == BLHD) ? O.transpose(1, 2) : O;
|
||||||
|
p.o_ptr = (bf16*)O_view.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_decode, p, stream);
|
||||||
|
C10_CUDA_CHECK(cudaGetLastError());
|
||||||
|
return O;
|
||||||
|
}
|
||||||
|
|
||||||
|
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||||
|
m.def("attn_decode", &attn_decode,
|
||||||
|
py::arg("q"),
|
||||||
|
py::arg("k"),
|
||||||
|
py::arg("v"),
|
||||||
|
py::arg("mask") = py::none(),
|
||||||
|
py::arg("causal_offset") = -1,
|
||||||
|
py::arg("scale") = 0.0,
|
||||||
|
py::arg("layout") = (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)");
|
||||||
|
}
|
||||||
+45
-23
@@ -1,11 +1,21 @@
|
|||||||
#pragma once
|
#pragma once
|
||||||
#include <cuda_bf16.h>
|
#include <cuda_bf16.h>
|
||||||
#include <float.h>
|
#include <float.h>
|
||||||
#include "attn_common.h"
|
#include "common.h"
|
||||||
#include "attn_warp_utils.cuh"
|
#include "layout_policies.cuh"
|
||||||
|
#include "../common/reduce.cuh"
|
||||||
|
|
||||||
|
namespace astrai {
|
||||||
|
namespace attention {
|
||||||
|
|
||||||
constexpr int DC_CHUNK = 64;
|
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) {
|
__global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
|
||||||
int batch = blockIdx.x / p.kv_head;
|
int batch = blockIdx.x / p.kv_head;
|
||||||
int kv_head = blockIdx.x % p.kv_head;
|
int kv_head = blockIdx.x % p.kv_head;
|
||||||
@@ -15,40 +25,46 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
|
|||||||
int lane = threadIdx.x;
|
int lane = threadIdx.x;
|
||||||
int hd_per_thread = p.head_dim / 32;
|
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
|
// Q: [batch, q_head, q_len=1, head_dim] — stride-based
|
||||||
float q_reg[8];
|
float q_reg[8];
|
||||||
int q_off = batch * p.q_stride_b + q_head * p.q_stride_h
|
int q_off = KV::q_decode_base(p, batch, q_head)
|
||||||
+ lane * hd_per_thread * p.q_stride_d;
|
+ lane * hd_per_thread * p.q_d_stride;
|
||||||
for (int i = 0; i < hd_per_thread; i++)
|
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 + q_head * p.mask_h_stride;
|
int mask_base = batch * p.mask_b_stride + q_head * p.mask_h_stride;
|
||||||
|
|
||||||
float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f};
|
float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f};
|
||||||
|
|
||||||
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
|
// 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 chunks_per_split = (chunks_total + p.num_splits - 1) / p.num_splits;
|
||||||
int ch_begin = split * chunks_per_split;
|
int ch_begin = split * chunks_per_split;
|
||||||
int ch_end = min(chunks_total, ch_begin + chunks_per_split);
|
int ch_end = min(chunks_total, ch_begin + chunks_per_split);
|
||||||
|
|
||||||
for (int ci = ch_begin; ci < ch_end; ci++) {
|
for (int ci = ch_begin; ci < ch_end; ci++) {
|
||||||
int chunk_start = ci * DC_CHUNK;
|
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;
|
int total = this_chunk * p.head_dim;
|
||||||
for (int i = threadIdx.y * 32 + lane; i < total;
|
for (int i = threadIdx.y * 32 + lane; i < total;
|
||||||
i += blockDim.x * blockDim.y) {
|
i += blockDim.x * blockDim.y) {
|
||||||
int s = i / p.head_dim;
|
int s = i / p.head_dim;
|
||||||
int d_dim = i % p.head_dim;
|
int d_dim = i % p.head_dim;
|
||||||
int kv_idx = chunk_start + s;
|
int kc = chunk_start + s;
|
||||||
int g_off = kv_base + kv_idx * p.kv_stride_l + d_dim * p.kv_stride_d;
|
KVAddr a = KV::template decode_addr<1>(
|
||||||
k_smem[i] = p.k[g_off];
|
p, kctx, batch, kv_head, kc, d_dim, true, true);
|
||||||
|
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();
|
__syncthreads();
|
||||||
|
|
||||||
@@ -65,7 +81,7 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
|
|||||||
partial = -FLT_MAX;
|
partial = -FLT_MAX;
|
||||||
}
|
}
|
||||||
if constexpr (IsCausal) {
|
if constexpr (IsCausal) {
|
||||||
if (kv_idx > p.causal_offset)
|
if (kv_idx >= KV::decode_attend_len(p, batch))
|
||||||
partial = -FLT_MAX;
|
partial = -FLT_MAX;
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -74,11 +90,10 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
|
|||||||
float beta = __expf(partial - new_m);
|
float beta = __expf(partial - new_m);
|
||||||
d = d * alpha + beta;
|
d = d * alpha + beta;
|
||||||
|
|
||||||
int v_off = kv_base + kv_idx * p.kv_stride_l
|
for (int i = 0; i < hd_per_thread; i++) {
|
||||||
+ lane * hd_per_thread * p.kv_stride_d;
|
float vv = __bfloat162float(v_smem[s * p.head_dim + lane * hd_per_thread + i]);
|
||||||
for (int i = 0; i < hd_per_thread; i++)
|
acc_reg[i] = fmaf(acc_reg[i], alpha, vv * beta);
|
||||||
acc_reg[i] = fmaf(acc_reg[i], alpha,
|
}
|
||||||
__bfloat162float(p.v[v_off + i * p.kv_stride_d]) * beta);
|
|
||||||
m = new_m;
|
m = new_m;
|
||||||
}
|
}
|
||||||
__syncthreads();
|
__syncthreads();
|
||||||
@@ -98,6 +113,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) {
|
__global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
|
||||||
int bh = blockIdx.x;
|
int bh = blockIdx.x;
|
||||||
int d = threadIdx.x;
|
int d = threadIdx.x;
|
||||||
@@ -124,6 +143,9 @@ __global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
|
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;
|
int o_off = KV::q_decode_base(p, batch, q_head) + d * p.q_d_stride;
|
||||||
p.o[o_off] = __float2bfloat16(acc * inv);
|
p.o_ptr[o_off] = __float2bfloat16(acc * inv);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
} // namespace attention
|
||||||
|
} // namespace astrai
|
||||||
+47
-28
@@ -1,20 +1,25 @@
|
|||||||
#pragma once
|
#pragma once
|
||||||
#include <cfloat>
|
#include <cfloat>
|
||||||
#include <cuda_bf16.h>
|
#include <cuda_bf16.h>
|
||||||
#include "attn_common.h"
|
#include "common.h"
|
||||||
#include "attn_mma_utils.cuh"
|
#include "layout_policies.cuh"
|
||||||
#include "attn_warp_utils.cuh"
|
#include "mma_utils.cuh"
|
||||||
|
|
||||||
// Split-K (FlashDecoding) tensor-core decode via GQA head-packing.
|
namespace astrai {
|
||||||
// Decode has q_len == 1, so we pack G = q_head/kv_head query heads into the
|
namespace attention {
|
||||||
// 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
|
// IsCausal and HasMask are compile-time bools — no runtime branch in the
|
||||||
// inner compute loop.
|
// inner compute loop.
|
||||||
//
|
//
|
||||||
// Traits = KernelTraits<HEAD_DIM, BC=32, WARPS=1, STAGES=<2 or 1>>.
|
// Traits = KernelTraits<HEAD_DIM, BC=16, WARPS=1, STAGES=2>.
|
||||||
template <typename Traits, bool IsCausal, bool HasMask>
|
template <typename Traits, typename KV, bool IsCausal, bool HasMask>
|
||||||
__global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
__global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
||||||
const int lane = threadIdx.x;
|
const int lane = threadIdx.x;
|
||||||
const int gid = lane >> 2;
|
const int gid = lane >> 2;
|
||||||
@@ -31,18 +36,21 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
|||||||
const int G = min(MAX_G, G_total - g_begin);
|
const int G = min(MAX_G, G_total - g_begin);
|
||||||
const int q_head0 = kv_head * 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)
|
// 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 sK[Traits::STAGES * Traits::BC * Traits::LD];
|
||||||
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
|
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
|
||||||
|
|
||||||
// Load Q directly from global into mma A-operand registers.
|
// 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 = KV::q_decode_base(p, batch, q_head0);
|
||||||
const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h;
|
|
||||||
const int qra = gid;
|
const int qra = gid;
|
||||||
const int qrb = gid + 8;
|
const int qrb = gid + 8;
|
||||||
const bool va = qra < G, vb = qrb < G;
|
const bool va = qra < G, vb = qrb < G;
|
||||||
unsigned Qa[Traits::KD][4];
|
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);
|
qra, qrb, va, vb, tid4, Qa);
|
||||||
|
|
||||||
float Oacc[Traits::DN8][4];
|
float Oacc[Traits::DN8][4];
|
||||||
@@ -51,13 +59,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;
|
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;
|
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 = (seq_len + Traits::BC - 1) / Traits::BC;
|
||||||
const int tiles_total = (p.kv_len + Traits::BC - 1) / Traits::BC;
|
|
||||||
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
|
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
|
||||||
const int ti_begin = split * tiles_per_split;
|
const int ti_begin = split * tiles_per_split;
|
||||||
const int ti_end = min(tiles_total, ti_begin + 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) {
|
auto load_tile = [&](int ti, int buf) {
|
||||||
int kv0 = ti * Traits::BC;
|
int kv0 = ti * Traits::BC;
|
||||||
bf16* dK = sK + buf * Traits::BC * Traits::LD;
|
bf16* dK = sK + buf * Traits::BC * Traits::LD;
|
||||||
@@ -67,13 +74,16 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
|||||||
i += Traits::NUM_THREADS * Traits::VEC) {
|
i += Traits::NUM_THREADS * Traits::VEC) {
|
||||||
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
|
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
|
||||||
int kc = kv0 + r;
|
int kc = kv0 + r;
|
||||||
bool valid = kc < p.kv_len;
|
bool valid = kc < seq_len;
|
||||||
|
// All GQA passes consume new K/V directly. Only the first pass
|
||||||
|
// persists it, so no cross-block synchronization is required.
|
||||||
|
KVAddr a = KV::template decode_addr<Traits::VEC>(
|
||||||
|
p, kctx, batch, kv_head, kc, d, valid, pass == 0);
|
||||||
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
|
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;
|
astrai::cp_async_16(&dK[off], a.k, a.valid);
|
||||||
cp_async_16_pred(&dK[off], &p.k[g_off], valid);
|
astrai::cp_async_16(&dV[off], a.v, a.valid);
|
||||||
cp_async_16_pred(&dV[off], &p.v[g_off], valid);
|
|
||||||
}
|
}
|
||||||
cp_async_commit();
|
astrai::cp_async_commit_group();
|
||||||
};
|
};
|
||||||
|
|
||||||
// ---- Multi-stage cp.async pipeline ----
|
// ---- Multi-stage cp.async pipeline ----
|
||||||
@@ -96,14 +106,17 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
|||||||
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
||||||
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
||||||
|
|
||||||
// Decode: q_len=1, so qrow0=qrow1=0
|
// Decode: q_len=1, so qrow0=qrow1=0. Paged treats [0, seq_len) as
|
||||||
int maxc = IsCausal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
|
// 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,
|
mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc,
|
||||||
0, 0,
|
0, 0,
|
||||||
p.mask_b_stride, 0, 0,
|
p.mask_b_stride, p.mask_h_stride, p.mask_l_stride,
|
||||||
batch, 0,
|
batch, q_head0 + gid, q_head0 + gid + 8,
|
||||||
p.mask,
|
p.mask,
|
||||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
va, vb,
|
||||||
|
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||||
|
|
||||||
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
|
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
|
||||||
};
|
};
|
||||||
@@ -114,7 +127,10 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
|||||||
load_tile(ti_begin + i, i);
|
load_tile(ti_begin + i, i);
|
||||||
|
|
||||||
for (int it = 0; it < ntiles; it++) {
|
for (int it = 0; it < ntiles; it++) {
|
||||||
cp_async_wait_group<STAGES - 1>();
|
if (it + 1 == ntiles)
|
||||||
|
astrai::cp_async_wait_group<0>();
|
||||||
|
else
|
||||||
|
astrai::cp_async_wait_group<STAGES - 1>();
|
||||||
__syncwarp();
|
__syncwarp();
|
||||||
process_tile(it, it & (STAGES - 1));
|
process_tile(it, it & (STAGES - 1));
|
||||||
__syncwarp();
|
__syncwarp();
|
||||||
@@ -125,7 +141,7 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
|||||||
// Fewer tiles than stages: load all, wait for all, process.
|
// Fewer tiles than stages: load all, wait for all, process.
|
||||||
for (int i = 0; i < ntiles; i++)
|
for (int i = 0; i < ntiles; i++)
|
||||||
load_tile(ti_begin + i, i);
|
load_tile(ti_begin + i, i);
|
||||||
cp_async_wait_group<0>();
|
astrai::cp_async_wait_all();
|
||||||
__syncwarp();
|
__syncwarp();
|
||||||
for (int it = 0; it < ntiles; it++)
|
for (int it = 0; it < ntiles; it++)
|
||||||
process_tile(it, it);
|
process_tile(it, it);
|
||||||
@@ -167,3 +183,6 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
} // namespace attention
|
||||||
|
} // namespace astrai
|
||||||
@@ -0,0 +1,247 @@
|
|||||||
|
#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 layout_policies.cuh), so each launcher struct
|
||||||
|
// below is templated on KV and the paged dispatch is just the same launcher
|
||||||
|
// instantiated with PagedKV. Only the grid/split math differs, and that is
|
||||||
|
// covered by KV::host_q_len / KV::host_kv_len.
|
||||||
|
|
||||||
|
#include <cuda_runtime.h>
|
||||||
|
#include <algorithm>
|
||||||
|
#include "layout_policies.cuh"
|
||||||
|
#include "prefill_split_q.cuh"
|
||||||
|
#include "decode_split_kv.cuh"
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
#include "prefill_split_q_mma.cuh"
|
||||||
|
#include "decode_split_kv_mma.cuh"
|
||||||
|
#endif
|
||||||
|
|
||||||
|
namespace astrai {
|
||||||
|
namespace attention {
|
||||||
|
|
||||||
|
// Split-KV: compute number of splits to fill all SMs for small-batch decode.
|
||||||
|
// Caps splits so each split processes at least `min_tiles_per_split` tiles,
|
||||||
|
// avoiding excessive loop/prologue overhead when tiles are small.
|
||||||
|
//
|
||||||
|
// 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 launchers (KV selects ContigKV or PagedKV addressing)
|
||||||
|
// ======================================================================
|
||||||
|
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
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 QSchedule, 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>;
|
||||||
|
// GQA head packing: HB = min(G, WARPS) q-heads of one kv-head group
|
||||||
|
// share each block's K/V stream (~HB× less global K/V traffic).
|
||||||
|
// Each head gets WPH = WARPS/HB 16-row chunks per block, so per-head
|
||||||
|
// rows drop from 64 to BR*WPH while total mma work per K/V byte is
|
||||||
|
// unchanged. G=1 (MHA) reproduces the historical grid exactly.
|
||||||
|
const int G = p.q_head / p.kv_head;
|
||||||
|
const int HB = std::min(G, Config::WARPS);
|
||||||
|
const int WPH = Config::WARPS / HB;
|
||||||
|
constexpr int BR = Traits::BR;
|
||||||
|
dim3 grid(QSchedule::packed_grid_x(p, BR * WPH),
|
||||||
|
p.kv_head * ((G + HB - 1) / HB),
|
||||||
|
QSchedule::host_grid_batch(p));
|
||||||
|
dim3 block(Traits::NUM_THREADS);
|
||||||
|
attn_prefill_split_q_mma_kernel<Traits, QSchedule, KV, IsCausal, HasMask>
|
||||||
|
<<<grid, block, 0, stream>>>(p);
|
||||||
|
}
|
||||||
|
};
|
||||||
|
#endif
|
||||||
|
|
||||||
|
template <typename QSchedule, 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 = 64, P_BC = 32;
|
||||||
|
dim3 grid(QSchedule::host_q_blocks(p, ROWS), p.q_head,
|
||||||
|
QSchedule::host_grid_batch(p));
|
||||||
|
dim3 block(G, ROWS);
|
||||||
|
attn_prefill_split_q_kernel_t<HEAD_DIM, QSchedule, KV, G, ROWS, P_BC,
|
||||||
|
IsCausal, HasMask>
|
||||||
|
<<<grid, block, 0, stream>>>(p);
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
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
|
||||||
|
using Launcher = PrefillLauncherMMA<DenseQSchedule, ContigKV>;
|
||||||
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||||
|
Launcher::template launch,
|
||||||
|
HEAD_DIM, p, stream);
|
||||||
|
#else
|
||||||
|
using Launcher = PrefillLauncherScalar<DenseQSchedule, ContigKV>;
|
||||||
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||||
|
Launcher::template launch,
|
||||||
|
HEAD_DIM, p, stream);
|
||||||
|
#endif
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
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);
|
||||||
|
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
using Launcher = PrefillLauncherMMA<PackedQSchedule, PagedKV>;
|
||||||
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||||
|
Launcher::template launch,
|
||||||
|
HEAD_DIM, p, stream);
|
||||||
|
#else
|
||||||
|
using Launcher = PrefillLauncherScalar<PackedQSchedule, PagedKV>;
|
||||||
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||||
|
Launcher::template launch,
|
||||||
|
HEAD_DIM, p, stream);
|
||||||
|
#endif
|
||||||
|
}
|
||||||
|
|
||||||
|
// ======================================================================
|
||||||
|
// Decode launchers (KV selects ContigKV or PagedKV addressing)
|
||||||
|
// ======================================================================
|
||||||
|
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
// BC=16: halves smem (16KB vs 32KB) → doubles occupancy (6 vs 3 blocks/SM).
|
||||||
|
// For D=256, BC=16 also reduces register pressure (fewer Sacc/PV frags),
|
||||||
|
// enabling STAGES=2 (double-buffer) within the 32KB smem budget — eliminates
|
||||||
|
// the 176-byte spill that STAGES=1+BC=32 suffered.
|
||||||
|
template <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;
|
||||||
|
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);
|
||||||
|
attn_decode_split_kv_mma_kernel<Traits, KV, IsCausal, HasMask>
|
||||||
|
<<<grid, 32, 0, stream>>>(p);
|
||||||
|
}
|
||||||
|
};
|
||||||
|
#endif
|
||||||
|
|
||||||
|
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(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<PagedKV>::template launch,
|
||||||
|
HEAD_DIM, p, stream);
|
||||||
|
#else
|
||||||
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||||
|
DecodeLauncherScalar<PagedKV>::template launch,
|
||||||
|
HEAD_DIM, p, stream);
|
||||||
|
#endif
|
||||||
|
|
||||||
|
attn_decode_combine_kernel<PagedKV><<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p);
|
||||||
|
}
|
||||||
|
|
||||||
|
} // namespace attention
|
||||||
|
} // namespace astrai
|
||||||
@@ -0,0 +1,363 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <float.h>
|
||||||
|
#include <torch/extension.h>
|
||||||
|
#include <c10/cuda/CUDAGuard.h>
|
||||||
|
#include "common.h"
|
||||||
|
|
||||||
|
// Dispatch head_dim: shared macro — avoids C++20 lambda template syntax.
|
||||||
|
// 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>(__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)"); \
|
||||||
|
}
|
||||||
|
|
||||||
|
namespace astrai {
|
||||||
|
namespace attention {
|
||||||
|
|
||||||
|
using bf16 = __nv_bfloat16;
|
||||||
|
|
||||||
|
// The split kernel unconditionally writes every (batch, q_head, split) slot it
|
||||||
|
// owns — including empty split ranges, which store m = -FLT_MAX so the combine
|
||||||
|
// skips them. Allocators are therefore left uninitialized (torch::empty); the
|
||||||
|
// per-call memset (torch::zeros / torch::full) was pure overhead.
|
||||||
|
template<typename P>
|
||||||
|
inline void alloc_split_partials(P& p) {
|
||||||
|
auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA);
|
||||||
|
auto o_part = torch::empty(at::IntArrayRef{p.batch, p.q_head, MAX_SPLITS, p.head_dim}, fopt);
|
||||||
|
auto ml_part = torch::empty(at::IntArrayRef{p.batch, p.q_head, MAX_SPLITS, 2}, fopt);
|
||||||
|
p.o_part = (float*)o_part.data_ptr();
|
||||||
|
p.ml_part = (float*)ml_part.data_ptr();
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- Shared Q-dims + strides extraction ----
|
||||||
|
template <typename P>
|
||||||
|
inline void extract_q_dims_and_strides(torch::Tensor& q, int64_t layout, P& p) {
|
||||||
|
if (layout == 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_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) {
|
||||||
|
auto m = mask.value();
|
||||||
|
TORCH_CHECK(m.is_cuda(), "mask must be on CUDA");
|
||||||
|
TORCH_CHECK(m.dtype() == torch::kBool, "mask must be bool");
|
||||||
|
TORCH_CHECK(m.size(0) == p.batch, "mask batch mismatch");
|
||||||
|
TORCH_CHECK(m.size(m.dim() - 1) == p.kv_len, "mask kv_len mismatch");
|
||||||
|
if (m.dim() == 2) {
|
||||||
|
p.mask_b_stride = (int)m.stride(0);
|
||||||
|
p.mask_h_stride = 0;
|
||||||
|
p.mask_l_stride = 0;
|
||||||
|
} else if (m.dim() == 3) {
|
||||||
|
TORCH_CHECK(m.size(1) == 1 || m.size(1) == p.q_len, "mask q_len mismatch");
|
||||||
|
p.mask_b_stride = (int)m.stride(0);
|
||||||
|
p.mask_h_stride = 0;
|
||||||
|
p.mask_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, 3D, or 4D");
|
||||||
|
}
|
||||||
|
p.mask = m.data_ptr<bool>();
|
||||||
|
} else {
|
||||||
|
p.mask = nullptr;
|
||||||
|
p.mask_b_stride = 0;
|
||||||
|
p.mask_h_stride = 0;
|
||||||
|
p.mask_l_stride = 0;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- attn_pack_params (contiguous KV) ----
|
||||||
|
template<typename T>
|
||||||
|
inline void attn_pack_params(
|
||||||
|
torch::Tensor q,
|
||||||
|
torch::Tensor k,
|
||||||
|
torch::Tensor v,
|
||||||
|
c10::optional<torch::Tensor> mask,
|
||||||
|
int64_t causal_offset,
|
||||||
|
double scale,
|
||||||
|
int64_t layout,
|
||||||
|
AttentionParams<T>& p
|
||||||
|
) {
|
||||||
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||||
|
|
||||||
|
TORCH_CHECK(q.is_cuda() && k.is_cuda() && v.is_cuda());
|
||||||
|
TORCH_CHECK(q.dtype() == torch::kBFloat16);
|
||||||
|
TORCH_CHECK(k.dtype() == torch::kBFloat16);
|
||||||
|
TORCH_CHECK(v.dtype() == torch::kBFloat16);
|
||||||
|
TORCH_CHECK(k.sizes() == v.sizes(), "K and V must have identical shapes");
|
||||||
|
TORCH_CHECK(q.dim() == 4 && k.dim() == 4, "Q/K/V must be 4D");
|
||||||
|
extract_q_dims_and_strides(q, layout, p);
|
||||||
|
|
||||||
|
if (layout == 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_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_ptr = (const T*)q.data_ptr();
|
||||||
|
p.k_ptr = (const T*)k.data_ptr();
|
||||||
|
p.v_ptr = (const T*)v.data_ptr();
|
||||||
|
p.new_k_ptr = nullptr;
|
||||||
|
p.new_v_ptr = nullptr;
|
||||||
|
p.o_ptr = nullptr;
|
||||||
|
p.o_part = nullptr;
|
||||||
|
p.ml_part = nullptr;
|
||||||
|
|
||||||
|
pack_mask(mask, p);
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- 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_decode_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,
|
||||||
|
const c10::optional<torch::Tensor>& new_k,
|
||||||
|
const c10::optional<torch::Tensor>& new_v,
|
||||||
|
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() && 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(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]");
|
||||||
|
|
||||||
|
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(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);
|
||||||
|
|
||||||
|
TORCH_CHECK(new_k.has_value() == new_v.has_value(),
|
||||||
|
"new_k and new_v must be provided together");
|
||||||
|
if (new_k.has_value()) {
|
||||||
|
auto nk = new_k.value();
|
||||||
|
auto nv = new_v.value();
|
||||||
|
TORCH_CHECK(nk.is_cuda() && nv.is_cuda(), "new K/V must be CUDA tensors");
|
||||||
|
TORCH_CHECK(nk.dtype() == torch::kBFloat16 && nv.dtype() == torch::kBFloat16,
|
||||||
|
"new K/V must be bf16");
|
||||||
|
TORCH_CHECK(nk.dim() == 3 && nv.dim() == 3,
|
||||||
|
"new K/V must be 3D [batch, kv_head, head_dim]");
|
||||||
|
TORCH_CHECK(nk.sizes() == nv.sizes(), "new K and V must have identical shapes");
|
||||||
|
TORCH_CHECK(nk.strides() == nv.strides(),
|
||||||
|
"new K and V must have identical strides");
|
||||||
|
TORCH_CHECK(nk.size(0) == p.batch && nk.size(1) == p.kv_head
|
||||||
|
&& nk.size(2) == p.head_dim, "new K/V shape mismatch");
|
||||||
|
TORCH_CHECK(nk.stride(2) == 1 && nv.stride(2) == 1,
|
||||||
|
"new K/V head_dim must be contiguous");
|
||||||
|
p.new_k_ptr = (const T*)nk.data_ptr();
|
||||||
|
p.new_v_ptr = (const T*)nv.data_ptr();
|
||||||
|
p.new_kv_b_stride = (int)nk.stride(0);
|
||||||
|
p.new_kv_h_stride = (int)nk.stride(1);
|
||||||
|
} else {
|
||||||
|
p.new_k_ptr = nullptr;
|
||||||
|
p.new_v_ptr = nullptr;
|
||||||
|
p.new_kv_b_stride = p.new_kv_h_stride = 0;
|
||||||
|
}
|
||||||
|
|
||||||
|
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);
|
||||||
|
|
||||||
|
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,
|
||||||
|
torch::Tensor q_tile_to_batch,
|
||||||
|
torch::Tensor q_tile_to_index,
|
||||||
|
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_tile_to_batch.is_cuda() && q_tile_to_index.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(q_tile_to_batch.dtype() == torch::kInt32,
|
||||||
|
"q_tile_to_batch must be int32");
|
||||||
|
TORCH_CHECK(q_tile_to_index.dtype() == torch::kInt32,
|
||||||
|
"q_tile_to_index 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]");
|
||||||
|
TORCH_CHECK(q_tile_to_batch.dim() == 1 && q_tile_to_index.dim() == 1,
|
||||||
|
"Q tile mappings must be 1D");
|
||||||
|
TORCH_CHECK(q_tile_to_batch.size(0) == q_tile_to_index.size(0),
|
||||||
|
"Q tile mappings must have equal length");
|
||||||
|
|
||||||
|
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.new_k_ptr = nullptr;
|
||||||
|
p.new_v_ptr = nullptr;
|
||||||
|
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.q_tile_to_batch = q_tile_to_batch.data_ptr<int>();
|
||||||
|
p.q_tile_to_index = q_tile_to_index.data_ptr<int>();
|
||||||
|
p.num_q_tiles = (int)q_tile_to_batch.size(0);
|
||||||
|
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;
|
||||||
|
}
|
||||||
|
|
||||||
|
} // namespace attention
|
||||||
|
} // namespace astrai
|
||||||
@@ -0,0 +1,292 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
#include "common.h"
|
||||||
|
|
||||||
|
// ============================================================================
|
||||||
|
// Attention layout policies keep Q scheduling independent from K/V storage.
|
||||||
|
// DenseQSchedule / PackedQSchedule map blocks to Q tiles; ContigKV / PagedKV
|
||||||
|
// resolve logical K/V positions to physical addresses. This lets the shared
|
||||||
|
// kernels compose Q layout and K/V storage without coupling the two concerns.
|
||||||
|
//
|
||||||
|
// 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.
|
||||||
|
// ============================================================================
|
||||||
|
|
||||||
|
#define HOST_FORCEINLINE static __host__ __forceinline__
|
||||||
|
#define DEVICE_FORCEINLINE static __device__ __forceinline__
|
||||||
|
#define HOST_DEV_FORCEINLINE static __host__ __device__ __forceinline__
|
||||||
|
|
||||||
|
namespace astrai {
|
||||||
|
namespace attention {
|
||||||
|
|
||||||
|
using bf16 = __nv_bfloat16;
|
||||||
|
|
||||||
|
// ============================================================================
|
||||||
|
// Q scheduling policies
|
||||||
|
//
|
||||||
|
// Map CUDA blocks to request-local Q tiles independently of K/V storage.
|
||||||
|
// Dense tensors encode the request in blockIdx.z; packed ragged tensors use
|
||||||
|
// a compact precomputed work map indexed by blockIdx.x.
|
||||||
|
// ============================================================================
|
||||||
|
|
||||||
|
struct DenseQSchedule {
|
||||||
|
HOST_FORCEINLINE int host_q_blocks(
|
||||||
|
const AttentionParams<bf16>& p, int rows) {
|
||||||
|
return (p.q_len + rows - 1) / rows;
|
||||||
|
}
|
||||||
|
|
||||||
|
HOST_FORCEINLINE int host_grid_batch(
|
||||||
|
const AttentionParams<bf16>& p) {
|
||||||
|
return p.batch;
|
||||||
|
}
|
||||||
|
|
||||||
|
DEVICE_FORCEINLINE void map_block(
|
||||||
|
const AttentionParams<bf16>&, int& batch, int& q_tile) {
|
||||||
|
batch = blockIdx.z;
|
||||||
|
q_tile = blockIdx.x;
|
||||||
|
}
|
||||||
|
|
||||||
|
// GQA-packed prefill mapping: HB q-heads of one kv-head group share a
|
||||||
|
// block's K/V stream, each head owning `rows` = BR*WPH consecutive q rows
|
||||||
|
// per block. Dense tensors tile q_len directly, one block per range.
|
||||||
|
HOST_FORCEINLINE int packed_grid_x(
|
||||||
|
const AttentionParams<bf16>& p, int rows) {
|
||||||
|
return (p.q_len + rows - 1) / rows;
|
||||||
|
}
|
||||||
|
|
||||||
|
DEVICE_FORCEINLINE void map_packed_block(
|
||||||
|
const AttentionParams<bf16>&, int rows, int& batch, int& row_base) {
|
||||||
|
batch = blockIdx.z;
|
||||||
|
row_base = blockIdx.x * rows;
|
||||||
|
}
|
||||||
|
|
||||||
|
DEVICE_FORCEINLINE int q_len(
|
||||||
|
const AttentionParams<bf16>& p, int) {
|
||||||
|
return p.q_len;
|
||||||
|
}
|
||||||
|
|
||||||
|
DEVICE_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;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
struct PackedQSchedule {
|
||||||
|
HOST_FORCEINLINE int host_q_blocks(
|
||||||
|
const AttentionParams<bf16>& p, int) {
|
||||||
|
return p.num_q_tiles;
|
||||||
|
}
|
||||||
|
|
||||||
|
HOST_FORCEINLINE int host_grid_batch(
|
||||||
|
const AttentionParams<bf16>&) {
|
||||||
|
return 1;
|
||||||
|
}
|
||||||
|
|
||||||
|
DEVICE_FORCEINLINE void map_block(
|
||||||
|
const AttentionParams<bf16>& p, int& batch, int& q_tile) {
|
||||||
|
batch = p.q_tile_to_batch[blockIdx.x];
|
||||||
|
q_tile = p.q_tile_to_index[blockIdx.x];
|
||||||
|
}
|
||||||
|
|
||||||
|
// GQA-packed prefill mapping: the host tile maps are built in
|
||||||
|
// HOST_Q_TILE_ROWS granularity, so each host tile splits into
|
||||||
|
// HOST_Q_TILE_ROWS / rows packed blocks along blockIdx.x.
|
||||||
|
HOST_FORCEINLINE int packed_grid_x(
|
||||||
|
const AttentionParams<bf16>& p, int rows) {
|
||||||
|
return p.num_q_tiles * (HOST_Q_TILE_ROWS / rows);
|
||||||
|
}
|
||||||
|
|
||||||
|
DEVICE_FORCEINLINE void map_packed_block(
|
||||||
|
const AttentionParams<bf16>& p, int rows, int& batch, int& row_base) {
|
||||||
|
const int hb = HOST_Q_TILE_ROWS / rows;
|
||||||
|
const int host_tile = blockIdx.x / hb;
|
||||||
|
batch = p.q_tile_to_batch[host_tile];
|
||||||
|
row_base = p.q_tile_to_index[host_tile] * HOST_Q_TILE_ROWS
|
||||||
|
+ (blockIdx.x - host_tile * hb) * rows;
|
||||||
|
}
|
||||||
|
|
||||||
|
DEVICE_FORCEINLINE int q_len(
|
||||||
|
const AttentionParams<bf16>& p, int batch) {
|
||||||
|
return p.qo_indptr[batch + 1] - p.qo_indptr[batch];
|
||||||
|
}
|
||||||
|
|
||||||
|
DEVICE_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;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
// 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_FORCEINLINE int host_kv_len(const AttentionParams<bf16>& p) {
|
||||||
|
return p.kv_len;
|
||||||
|
}
|
||||||
|
|
||||||
|
// decode: same offset (q_len == 1, so there is no row stride component)
|
||||||
|
DEVICE_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;
|
||||||
|
}
|
||||||
|
|
||||||
|
DEVICE_FORCEINLINE int kv_len(const AttentionParams<bf16>& p, int) {
|
||||||
|
return p.kv_len;
|
||||||
|
}
|
||||||
|
DEVICE_FORCEINLINE int causal_offset(
|
||||||
|
const AttentionParams<bf16>& p, int, int) {
|
||||||
|
return p.causal_offset;
|
||||||
|
}
|
||||||
|
// decode: exclusive bound of the single query's attend range
|
||||||
|
DEVICE_FORCEINLINE int decode_attend_len(const AttentionParams<bf16>& p, int) {
|
||||||
|
return (p.kv_len < p.causal_offset + 1) ? p.kv_len : (p.causal_offset + 1);
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
DEVICE_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;
|
||||||
|
}
|
||||||
|
DEVICE_FORCEINLINE int resolve_token(
|
||||||
|
const AttentionParams<bf16>& p, const KVContext& c, int kc, bool valid) {
|
||||||
|
return valid ? kc : -1;
|
||||||
|
}
|
||||||
|
DEVICE_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};
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int VEC>
|
||||||
|
DEVICE_FORCEINLINE KVAddr decode_addr(
|
||||||
|
const AttentionParams<bf16>& p, const KVContext& c,
|
||||||
|
int, int, int kc, int d, bool valid, bool) {
|
||||||
|
int token = resolve_token(p, c, kc, valid);
|
||||||
|
return kv_addr_from_token(p, c, token, d);
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
// ---- Paged (SGLang-style flat pool) K/V ----
|
||||||
|
struct PagedKV {
|
||||||
|
static constexpr bool kPaged = true;
|
||||||
|
|
||||||
|
HOST_FORCEINLINE int host_kv_len(const AttentionParams<bf16>& p) {
|
||||||
|
return p.max_context_len;
|
||||||
|
}
|
||||||
|
|
||||||
|
// decode: Q is [batch, q_head, head_dim], so batch is the outer row
|
||||||
|
DEVICE_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;
|
||||||
|
}
|
||||||
|
|
||||||
|
DEVICE_FORCEINLINE int kv_len(const AttentionParams<bf16>& p, int batch) {
|
||||||
|
return p.kv_indptr[batch + 1] - p.kv_indptr[batch];
|
||||||
|
}
|
||||||
|
DEVICE_FORCEINLINE int causal_offset(
|
||||||
|
const AttentionParams<bf16>& p, int batch, int q_len) {
|
||||||
|
return kv_len(p, batch) - q_len;
|
||||||
|
}
|
||||||
|
// decode: the query is the last token, so [0, seq_len) IS its causal range
|
||||||
|
DEVICE_FORCEINLINE int decode_attend_len(const AttentionParams<bf16>& p, int batch) {
|
||||||
|
return kv_len(p, batch);
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
DEVICE_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;
|
||||||
|
}
|
||||||
|
DEVICE_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;
|
||||||
|
}
|
||||||
|
DEVICE_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};
|
||||||
|
}
|
||||||
|
|
||||||
|
DEVICE_FORCEINLINE KVAddr new_kv_addr(
|
||||||
|
const AttentionParams<bf16>& p, int batch, int kv_head, int d) {
|
||||||
|
const int64_t off = (int64_t)batch * p.new_kv_b_stride
|
||||||
|
+ (int64_t)kv_head * p.new_kv_h_stride + d;
|
||||||
|
return {&p.new_k_ptr[off], &p.new_v_ptr[off], true};
|
||||||
|
}
|
||||||
|
|
||||||
|
DEVICE_FORCEINLINE void store_new_kv(
|
||||||
|
const AttentionParams<bf16>& p, const KVContext& c,
|
||||||
|
int seq_len, int d, const KVAddr& src) {
|
||||||
|
int slot = resolve_token(p, c, seq_len - 1, true);
|
||||||
|
const int64_t off = (int64_t)slot * c.pool_stride + c.head_off + d;
|
||||||
|
const_cast<bf16*>(p.k_ptr)[off] = *reinterpret_cast<const bf16*>(src.k);
|
||||||
|
const_cast<bf16*>(p.v_ptr)[off] = *reinterpret_cast<const bf16*>(src.v);
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int VEC>
|
||||||
|
DEVICE_FORCEINLINE KVAddr decode_addr(
|
||||||
|
const AttentionParams<bf16>& p, const KVContext& c,
|
||||||
|
int batch, int kv_head, int kc, int d, bool valid, bool persist) {
|
||||||
|
if (p.new_k_ptr && valid && kc == kv_len(p, batch) - 1) {
|
||||||
|
KVAddr src = new_kv_addr(p, batch, kv_head, d);
|
||||||
|
if (persist) {
|
||||||
|
#pragma unroll
|
||||||
|
for (int j = 0; j < VEC; j++) {
|
||||||
|
KVAddr value = new_kv_addr(p, batch, kv_head, d + j);
|
||||||
|
store_new_kv(p, c, kc + 1, d + j, value);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return src;
|
||||||
|
}
|
||||||
|
int token = resolve_token(p, c, kc, valid);
|
||||||
|
return kv_addr_from_token(p, c, token, d);
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
} // namespace attention
|
||||||
|
} // namespace astrai
|
||||||
@@ -3,12 +3,18 @@
|
|||||||
#include <cuda_fp16.h>
|
#include <cuda_fp16.h>
|
||||||
#include <cuda_runtime.h>
|
#include <cuda_runtime.h>
|
||||||
|
|
||||||
|
#include "../common/cp_async.cuh"
|
||||||
|
#include "../common/mma.cuh"
|
||||||
|
|
||||||
// Predicated cp.async (4-operand form) requires CUDA 11.2+.
|
// Predicated cp.async (4-operand form) requires CUDA 11.2+.
|
||||||
// bf16 mma.sync requires sm_80+ (guarded at build time by ASTRAI_NO_MMA).
|
// bf16 mma.sync requires sm_80+ (guarded at build time by ASTRAI_NO_MMA).
|
||||||
#if CUDART_VERSION < 11020
|
#if CUDART_VERSION < 11020
|
||||||
#error "AstrAI CUDA kernels require CUDA 11.2 or later (CUDART_VERSION >= 11020)."
|
#error "AstrAI CUDA kernels require CUDA 11.2 or later (CUDART_VERSION >= 11020)."
|
||||||
#endif
|
#endif
|
||||||
|
|
||||||
|
namespace astrai {
|
||||||
|
namespace attention {
|
||||||
|
|
||||||
// ============================================================================
|
// ============================================================================
|
||||||
// KernelTraits — FlashAttention-v2 style compile-time configuration bundle.
|
// KernelTraits — FlashAttention-v2 style compile-time configuration bundle.
|
||||||
//
|
//
|
||||||
@@ -24,10 +30,10 @@ struct KernelTraits {
|
|||||||
|
|
||||||
static constexpr int BR = 16; // Q rows per warp (mma M=16)
|
static constexpr int BR = 16; // Q rows per warp (mma M=16)
|
||||||
|
|
||||||
// Derived: mma.sync.m16n8k16 tile counts
|
// Derived: mma tile counts from the shared mma_shape (m16n8k16 for bf16)
|
||||||
static constexpr int KD = HEAD_DIM / 16; // Q/K k-slides
|
static constexpr int KD = HEAD_DIM / astrai::mma_shape<bf16>::k; // Q/K k-slides
|
||||||
static constexpr int NC8 = BC / 8; // S n-tiles (N=8)
|
static constexpr int NC8 = BC / 8; // S n-tiles (N=8)
|
||||||
static constexpr int KT2 = BC / 16; // P k-tiles (K=16)
|
static constexpr int KT2 = BC / astrai::mma_shape<bf16>::k; // P k-tiles (K=16)
|
||||||
static constexpr int DN8 = HEAD_DIM / 8; // O n-tiles (N=8)
|
static constexpr int DN8 = HEAD_DIM / 8; // O n-tiles (N=8)
|
||||||
|
|
||||||
static constexpr int LD = HEAD_DIM; // smem leading dim
|
static constexpr int LD = HEAD_DIM; // smem leading dim
|
||||||
@@ -43,16 +49,7 @@ struct KernelTraits {
|
|||||||
|
|
||||||
// ---- PTX wrappers ----
|
// ---- PTX wrappers ----
|
||||||
using bf16 = __nv_bfloat16;
|
using bf16 = __nv_bfloat16;
|
||||||
|
// bf16 mma.sync lives in the shared astrai::mma_sync template (common/mma.cuh).
|
||||||
__device__ __forceinline__ void mma16816(float* d, const unsigned* a,
|
|
||||||
const unsigned* b, const float* c) {
|
|
||||||
asm volatile(
|
|
||||||
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
|
|
||||||
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%10,%11,%12,%13};"
|
|
||||||
: "=f"(d[0]), "=f"(d[1]), "=f"(d[2]), "=f"(d[3])
|
|
||||||
: "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]),
|
|
||||||
"f"(c[0]), "f"(c[1]), "f"(c[2]), "f"(c[3]));
|
|
||||||
}
|
|
||||||
|
|
||||||
// read two adjacent bf16 from smem as one packed .b32 (elem0 low, elem1 high)
|
// read two adjacent bf16 from smem as one packed .b32 (elem0 low, elem1 high)
|
||||||
__device__ __forceinline__ unsigned ld2(const bf16* p) {
|
__device__ __forceinline__ unsigned ld2(const bf16* p) {
|
||||||
@@ -73,68 +70,24 @@ __device__ __forceinline__ unsigned pkb(bf16 a, bf16 b) {
|
|||||||
return *reinterpret_cast<unsigned*>(&v);
|
return *reinterpret_cast<unsigned*>(&v);
|
||||||
}
|
}
|
||||||
|
|
||||||
// ldmatrix: cooperatively load mma fragments from smem (one instruction per
|
// ldmatrix lives in the shared template (common/mma.cuh):
|
||||||
// 16x16 / 16x8 tile) with the exact register layout mma expects.
|
// `astrai::ldmatrix_x2<bf16>` / `<bf16, /*Trans=*/true>` load the K/V
|
||||||
__device__ __forceinline__ void ldmatrix_x4(unsigned* r, const bf16* p) {
|
// fragments with the exact register layout mma expects.
|
||||||
unsigned a = __cvta_generic_to_shared(p);
|
|
||||||
asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];"
|
|
||||||
: "=r"(r[0]), "=r"(r[1]), "=r"(r[2]), "=r"(r[3])
|
|
||||||
: "r"(a));
|
|
||||||
}
|
|
||||||
__device__ __forceinline__ void ldmatrix_x2(unsigned* r, const bf16* p) {
|
|
||||||
unsigned a = __cvta_generic_to_shared(p);
|
|
||||||
asm volatile("ldmatrix.sync.aligned.m8n8.x2.shared.b16 {%0,%1}, [%2];"
|
|
||||||
: "=r"(r[0]), "=r"(r[1])
|
|
||||||
: "r"(a));
|
|
||||||
}
|
|
||||||
__device__ __forceinline__ void ldmatrix_x2_trans(unsigned* r, const bf16* p) {
|
|
||||||
unsigned a = __cvta_generic_to_shared(p);
|
|
||||||
asm volatile("ldmatrix.sync.aligned.m8n8.x2.trans.shared.b16 {%0,%1}, [%2];"
|
|
||||||
: "=r"(r[0]), "=r"(r[1])
|
|
||||||
: "r"(a));
|
|
||||||
}
|
|
||||||
|
|
||||||
// XOR swizzle for shared-memory column at 8-bf16 chunk granularity.
|
// XOR swizzle for shared-memory column at 8-bf16 chunk granularity.
|
||||||
__device__ __forceinline__ int swiz_col(int d, int r, int mask = 7) {
|
__device__ __forceinline__ int swiz_col(int d, int r, int mask = 7) {
|
||||||
return ((d >> 3) ^ (r & mask)) << 3 | (d & 7);
|
return ((d >> 3) ^ (r & mask)) << 3 | (d & 7);
|
||||||
}
|
}
|
||||||
|
|
||||||
// cp.async: copy 16 bytes (8 bf16) from global to shared memory directly.
|
// cp.async primitives live in the shared template (common/cp_async.cuh):
|
||||||
__device__ __forceinline__ void cp_async_16(bf16* smem_ptr, const void* gmem_ptr) {
|
// `astrai::cp_async_16` (predicated), `astrai::cp_async_commit_group`,
|
||||||
unsigned smem_addr = __cvta_generic_to_shared(smem_ptr);
|
// `astrai::cp_async_wait_group<N>` / `_wait_all` stage the K/V tiles.
|
||||||
asm volatile("cp.async.ca.shared.global [%0], [%1], 16;"
|
|
||||||
:: "r"(smem_addr), "l"(gmem_ptr));
|
|
||||||
}
|
|
||||||
|
|
||||||
// Predicated cp.async: copy 16 bytes when `pred`, otherwise zero-fill.
|
|
||||||
// src_size=0 → no bytes read from src, so out-of-bounds src address is safe.
|
|
||||||
__device__ __forceinline__ void cp_async_16_pred(bf16* smem_ptr,
|
|
||||||
const void* gmem_ptr,
|
|
||||||
bool pred) {
|
|
||||||
unsigned smem_addr = __cvta_generic_to_shared(smem_ptr);
|
|
||||||
int src_size = pred ? 16 : 0;
|
|
||||||
asm volatile("cp.async.ca.shared.global [%0], [%1], 16, %2;"
|
|
||||||
:: "r"(smem_addr), "l"(gmem_ptr), "r"(src_size));
|
|
||||||
}
|
|
||||||
|
|
||||||
__device__ __forceinline__ void cp_async_commit() {
|
|
||||||
asm volatile("cp.async.commit_group;");
|
|
||||||
}
|
|
||||||
|
|
||||||
__device__ __forceinline__ void cp_async_wait_all() {
|
|
||||||
asm volatile("cp.async.wait_all;");
|
|
||||||
}
|
|
||||||
|
|
||||||
template <int N>
|
|
||||||
__device__ __forceinline__ void cp_async_wait_group() {
|
|
||||||
asm volatile("cp.async.wait_group %0;" :: "n"(N));
|
|
||||||
}
|
|
||||||
|
|
||||||
// ---------------------------------------------------------------------------
|
// ---------------------------------------------------------------------------
|
||||||
// Q-load: load query rows directly from global memory into mma A-operand
|
// 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.
|
// 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
|
// stride_row is p.q_h_stride for decode (q_len=1, G heads) or
|
||||||
// p.q_stride_l for prefill (multi-q rows).
|
// p.q_l_stride for prefill (multi-q rows).
|
||||||
// ---------------------------------------------------------------------------
|
// ---------------------------------------------------------------------------
|
||||||
template <int KD>
|
template <int KD>
|
||||||
__device__ inline void load_q_mma_frags(
|
__device__ inline void load_q_mma_frags(
|
||||||
@@ -180,9 +133,9 @@ __device__ inline void mma_compute_scores(
|
|||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int kt = 0; kt < Traits::KD; kt++) {
|
for (int kt = 0; kt < Traits::KD; kt++) {
|
||||||
unsigned b[2];
|
unsigned b[2];
|
||||||
ldmatrix_x2(b, &sK[krow_l * Traits::LD
|
astrai::ldmatrix_x2<bf16>(b, &sK[krow_l * Traits::LD
|
||||||
+ swiz_col(kt * 16 + kcol_h, krow_l, Traits::SWIZ_MASK)]);
|
+ swiz_col(kt * 16 + kcol_h, krow_l, Traits::SWIZ_MASK)]);
|
||||||
mma16816(Sacc[n8], Qa[kt], b, Sacc[n8]);
|
astrai::mma_sync<bf16>(Sacc[n8], Qa[kt], b, Sacc[n8]);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -198,9 +151,10 @@ __device__ inline void mma_softmax_tile(
|
|||||||
int kv0,
|
int kv0,
|
||||||
int maxc0, int maxc1,
|
int maxc0, int maxc1,
|
||||||
int qrow0, int qrow1,
|
int qrow0, int qrow1,
|
||||||
int mask_b_stride, int mask_h_stride, int mask_q_stride,
|
int mask_b_stride, int mask_h_stride, int mask_l_stride,
|
||||||
int mask_batch, int mask_head,
|
int mask_batch, int mask_head0, int mask_head1,
|
||||||
const bool* __restrict__ mask,
|
const bool* __restrict__ mask,
|
||||||
|
bool valid0, bool valid1,
|
||||||
float Sacc[Traits::NC8][4],
|
float Sacc[Traits::NC8][4],
|
||||||
float Oacc[Traits::DN8][4],
|
float Oacc[Traits::DN8][4],
|
||||||
float& m0, float& m1,
|
float& m0, float& m1,
|
||||||
@@ -210,16 +164,16 @@ __device__ inline void mma_softmax_tile(
|
|||||||
int tid4 = lane & 3;
|
int tid4 = lane & 3;
|
||||||
|
|
||||||
float rmax0 = -FLT_MAX, rmax1 = -FLT_MAX;
|
float rmax0 = -FLT_MAX, rmax1 = -FLT_MAX;
|
||||||
int mask_base0 = mask_batch * mask_b_stride + mask_head * mask_h_stride + qrow0 * 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_head * mask_h_stride + qrow1 * mask_q_stride;
|
int mask_base1 = mask_batch * mask_b_stride + mask_head1 * mask_h_stride + qrow1 * mask_l_stride;
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int n8 = 0; n8 < Traits::NC8; n8++) {
|
for (int n8 = 0; n8 < Traits::NC8; n8++) {
|
||||||
int cc = kv0 + n8 * 8 + 2 * tid4;
|
int cc = kv0 + n8 * 8 + 2 * tid4;
|
||||||
int c1 = cc + 1;
|
int c1 = cc + 1;
|
||||||
bool b0 = (cc >= maxc0) || (HasMask && !mask[mask_base0 + cc]);
|
bool b0 = !valid0 || (cc >= maxc0) || (HasMask && !mask[mask_base0 + cc]);
|
||||||
bool b1 = (c1 >= maxc0) || (HasMask && !mask[mask_base0 + c1]);
|
bool b1 = !valid0 || (c1 >= maxc0) || (HasMask && !mask[mask_base0 + c1]);
|
||||||
bool b2 = (cc >= maxc1) || (HasMask && !mask[mask_base1 + cc]);
|
bool b2 = !valid1 || (cc >= maxc1) || (HasMask && !mask[mask_base1 + cc]);
|
||||||
bool b3 = (c1 >= maxc1) || (HasMask && !mask[mask_base1 + c1]);
|
bool b3 = !valid1 || (c1 >= maxc1) || (HasMask && !mask[mask_base1 + c1]);
|
||||||
float s0 = b0 ? -FLT_MAX : Sacc[n8][0];
|
float s0 = b0 ? -FLT_MAX : Sacc[n8][0];
|
||||||
float s1 = b1 ? -FLT_MAX : Sacc[n8][1];
|
float s1 = b1 ? -FLT_MAX : Sacc[n8][1];
|
||||||
float s2 = b2 ? -FLT_MAX : Sacc[n8][2];
|
float s2 = b2 ? -FLT_MAX : Sacc[n8][2];
|
||||||
@@ -289,9 +243,12 @@ __device__ inline void mma_pv_accumulate(
|
|||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
|
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
|
||||||
unsigned b[2];
|
unsigned b[2];
|
||||||
ldmatrix_x2_trans(b, &sV[vrow_l * Traits::LD
|
astrai::ldmatrix_x2<bf16, true>(b, &sV[vrow_l * Traits::LD
|
||||||
+ swiz_col(dn8 * 8, vrow_l, Traits::SWIZ_MASK)]);
|
+ swiz_col(dn8 * 8, vrow_l, Traits::SWIZ_MASK)]);
|
||||||
mma16816(Oacc[dn8], Pa, b, Oacc[dn8]);
|
astrai::mma_sync<bf16>(Oacc[dn8], Pa, b, Oacc[dn8]);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
} // namespace attention
|
||||||
|
} // namespace astrai
|
||||||
@@ -0,0 +1,88 @@
|
|||||||
|
#include "dispatchers.cuh"
|
||||||
|
#include "entry_utils.cuh"
|
||||||
|
|
||||||
|
using namespace astrai::attention;
|
||||||
|
|
||||||
|
torch::Tensor attn_paged_decode(
|
||||||
|
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,
|
||||||
|
c10::optional<torch::Tensor> new_k,
|
||||||
|
c10::optional<torch::Tensor> new_v,
|
||||||
|
c10::optional<torch::Tensor> mask,
|
||||||
|
int64_t causal_offset,
|
||||||
|
double scale,
|
||||||
|
c10::optional<torch::Tensor> o_part_buf,
|
||||||
|
c10::optional<torch::Tensor> ml_part_buf,
|
||||||
|
c10::optional<torch::Tensor> out_buf
|
||||||
|
) {
|
||||||
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||||
|
auto stream = at::cuda::getCurrentCUDAStream();
|
||||||
|
|
||||||
|
AttentionParams<bf16> p;
|
||||||
|
attn_pack_paged_decode_params(q, k_cache, v_cache,
|
||||||
|
req_to_token, req_pool_indices, kv_indptr,
|
||||||
|
new_k, new_v,
|
||||||
|
mask, causal_offset, scale, 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("k_cache"),
|
||||||
|
py::arg("v_cache"),
|
||||||
|
py::arg("req_to_token"),
|
||||||
|
py::arg("req_pool_indices"),
|
||||||
|
py::arg("kv_indptr"),
|
||||||
|
py::arg("new_k") = py::none(),
|
||||||
|
py::arg("new_v") = py::none(),
|
||||||
|
py::arg("mask") = py::none(),
|
||||||
|
py::arg("causal_offset") = -1,
|
||||||
|
py::arg("scale") = 0.0,
|
||||||
|
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.");
|
||||||
|
}
|
||||||
@@ -0,0 +1,53 @@
|
|||||||
|
#include "dispatchers.cuh"
|
||||||
|
#include "entry_utils.cuh"
|
||||||
|
|
||||||
|
using namespace astrai::attention;
|
||||||
|
|
||||||
|
torch::Tensor attn_paged_prefill(
|
||||||
|
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,
|
||||||
|
torch::Tensor q_tile_to_batch,
|
||||||
|
torch::Tensor q_tile_to_index,
|
||||||
|
c10::optional<torch::Tensor> mask,
|
||||||
|
int64_t causal_offset,
|
||||||
|
double scale
|
||||||
|
) {
|
||||||
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||||
|
auto stream = at::cuda::getCurrentCUDAStream();
|
||||||
|
|
||||||
|
AttentionParams<bf16> p;
|
||||||
|
attn_pack_paged_prefill_params(q, k_cache, v_cache,
|
||||||
|
req_to_token, req_pool_indices,
|
||||||
|
kv_indptr, qo_indptr,
|
||||||
|
q_tile_to_batch, q_tile_to_index, mask,
|
||||||
|
causal_offset, scale, p);
|
||||||
|
|
||||||
|
auto O = torch::empty({q.size(0), q.size(1), q.size(2)}, q.options());
|
||||||
|
p.o_ptr = (bf16*)O.data_ptr();
|
||||||
|
|
||||||
|
DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_prefill, p, stream);
|
||||||
|
C10_CUDA_CHECK(cudaGetLastError());
|
||||||
|
return O;
|
||||||
|
}
|
||||||
|
|
||||||
|
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||||
|
m.def("attn_paged_prefill", &attn_paged_prefill,
|
||||||
|
py::arg("q"),
|
||||||
|
py::arg("k_cache"),
|
||||||
|
py::arg("v_cache"),
|
||||||
|
py::arg("req_to_token"),
|
||||||
|
py::arg("req_pool_indices"),
|
||||||
|
py::arg("kv_indptr"),
|
||||||
|
py::arg("qo_indptr"),
|
||||||
|
py::arg("q_tile_to_batch"),
|
||||||
|
py::arg("q_tile_to_index"),
|
||||||
|
py::arg("mask") = py::none(),
|
||||||
|
py::arg("causal_offset") = -1,
|
||||||
|
py::arg("scale") = 0.0,
|
||||||
|
"SGLang-style paged prefill: flat KV pool + ragged batch.");
|
||||||
|
}
|
||||||
@@ -1,5 +1,7 @@
|
|||||||
#include "attn_dispatchers.cuh"
|
#include "dispatchers.cuh"
|
||||||
#include "attn_entry_utils.cuh"
|
#include "entry_utils.cuh"
|
||||||
|
|
||||||
|
using namespace astrai::attention;
|
||||||
|
|
||||||
torch::Tensor attn_prefill(
|
torch::Tensor attn_prefill(
|
||||||
torch::Tensor q,
|
torch::Tensor q,
|
||||||
@@ -10,15 +12,19 @@ torch::Tensor attn_prefill(
|
|||||||
double scale,
|
double scale,
|
||||||
int64_t layout
|
int64_t layout
|
||||||
) {
|
) {
|
||||||
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||||
|
auto stream = at::cuda::getCurrentCUDAStream();
|
||||||
|
|
||||||
AttentionParams<bf16> p;
|
AttentionParams<bf16> p;
|
||||||
attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p);
|
attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p);
|
||||||
TORCH_CHECK(p.head_dim % 16 == 0, "head_dim must be multiple of 16");
|
TORCH_CHECK(p.head_dim % 16 == 0, "head_dim must be multiple of 16");
|
||||||
|
|
||||||
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
|
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
|
||||||
auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
|
auto O_view = (layout == BLHD) ? O.transpose(1, 2) : O;
|
||||||
p.o = (bf16*)O_view.data_ptr();
|
p.o_ptr = (bf16*)O_view.data_ptr();
|
||||||
|
|
||||||
DISPATCH_HEAD_DIM(p.head_dim, dispatch_prefill, p);
|
DISPATCH_HEAD_DIM(p.head_dim, dispatch_prefill, p, stream);
|
||||||
|
C10_CUDA_CHECK(cudaGetLastError());
|
||||||
return O;
|
return O;
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -30,6 +36,6 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
|||||||
py::arg("mask") = py::none(),
|
py::arg("mask") = py::none(),
|
||||||
py::arg("causal_offset") = -1,
|
py::arg("causal_offset") = -1,
|
||||||
py::arg("scale") = 0.0,
|
py::arg("scale") = 0.0,
|
||||||
py::arg("layout") = 0,
|
py::arg("layout") = (int64_t)BHLD,
|
||||||
"GQA prefill (tensor-core mma on sm_80+, scalar fallback)");
|
"GQA prefill (tensor-core mma on sm_80+, scalar fallback)");
|
||||||
}
|
}
|
||||||
+42
-34
@@ -1,22 +1,21 @@
|
|||||||
#pragma once
|
#pragma once
|
||||||
#include <cfloat>
|
#include <cfloat>
|
||||||
#include <cuda_bf16.h>
|
#include <cuda_bf16.h>
|
||||||
#include "attn_common.h"
|
#include "common.h"
|
||||||
|
#include "layout_policies.cuh"
|
||||||
|
#include "../common/reduce.cuh"
|
||||||
|
|
||||||
|
namespace astrai {
|
||||||
|
namespace attention {
|
||||||
|
|
||||||
using bf16 = __nv_bfloat16;
|
using bf16 = __nv_bfloat16;
|
||||||
|
|
||||||
// v9: group-split register blocking. G threads cooperate on one query row,
|
// v9: group-split register blocking. G threads cooperate on one query row,
|
||||||
// each owning HEAD_DIM/G dims of qreg[]/acc[]. IsCausal and HasMask are
|
// each owning HEAD_DIM/G dims of qreg[]/acc[]. IsCausal and HasMask are
|
||||||
// compile-time bools — the compiler eliminates dead branches.
|
// compile-time bools — the compiler eliminates dead branches.
|
||||||
// Templated on <HEAD_DIM, G, ROWS, P_BC, IsCausal, HasMask>.
|
// Unified across contiguous and paged (SGLang flat-pool) K/V via KV.
|
||||||
|
// Templated on <HEAD_DIM, KV, G, ROWS, P_BC, IsCausal, HasMask>.
|
||||||
template <int G>
|
// group_reduce_sum<G> lives in common/reduce.cuh (astrai::).
|
||||||
__device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
|
|
||||||
#pragma unroll
|
|
||||||
for (int o = G / 2; o > 0; o >>= 1)
|
|
||||||
v += __shfl_xor_sync(mask, v, o);
|
|
||||||
return v;
|
|
||||||
}
|
|
||||||
|
|
||||||
// load 8 contiguous bf16 from (16-byte aligned) smem as one float4
|
// load 8 contiguous bf16 from (16-byte aligned) smem as one float4
|
||||||
__device__ __forceinline__ void ld8(const bf16* p, float* o) {
|
__device__ __forceinline__ void ld8(const bf16* p, float* o) {
|
||||||
@@ -30,30 +29,37 @@ __device__ __forceinline__ void ld8(const bf16* p, float* o) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
template <int HEAD_DIM, int G, int ROWS, int P_BC, bool IsCausal, bool HasMask>
|
template <int HEAD_DIM, typename QSchedule, typename KV, int G, int ROWS, int P_BC,
|
||||||
|
bool IsCausal, bool HasMask>
|
||||||
__global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
|
__global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
|
||||||
constexpr int DPT = HEAD_DIM / G;
|
constexpr int DPT = HEAD_DIM / G;
|
||||||
|
|
||||||
int q_tile = blockIdx.x;
|
int batch, q_tile;
|
||||||
|
QSchedule::map_block(p, batch, q_tile);
|
||||||
|
|
||||||
int q_head = blockIdx.y;
|
int q_head = blockIdx.y;
|
||||||
int batch = blockIdx.z;
|
|
||||||
int gpos = threadIdx.x; // 0..G-1 (which d-chunk)
|
int gpos = threadIdx.x; // 0..G-1 (which d-chunk)
|
||||||
int row = threadIdx.y; // 0..ROWS-1
|
int row = threadIdx.y; // 0..ROWS-1
|
||||||
int q_row = q_tile * ROWS + row;
|
int q_row = q_tile * ROWS + row;
|
||||||
|
|
||||||
int kv_head = q_head / (p.q_head / p.kv_head);
|
// Per-request dims (from KV policy — paged reads kv_indptr/qo_indptr).
|
||||||
|
const int seq_len = KV::kv_len(p, batch);
|
||||||
|
const int q_len = QSchedule::q_len(p, batch);
|
||||||
|
const int causal_off = KV::causal_offset(p, batch, q_len);
|
||||||
|
const int kv_head = q_head / (p.q_head / p.kv_head);
|
||||||
|
const KVContext kctx = KV::template make_ctx<HEAD_DIM>(p, batch, kv_head);
|
||||||
|
|
||||||
__shared__ __align__(16) bf16 sK[P_BC * HEAD_DIM];
|
__shared__ __align__(16) bf16 sK[P_BC * HEAD_DIM];
|
||||||
__shared__ __align__(16) bf16 sV[P_BC * HEAD_DIM];
|
__shared__ __align__(16) bf16 sV[P_BC * HEAD_DIM];
|
||||||
|
|
||||||
// Q: stride-based load [batch, q_head, q_len, head_dim]
|
// Q: stride-based load [batch, q_head, q_len, head_dim]
|
||||||
|
const int q_base = QSchedule::q_base(p, batch, q_head);
|
||||||
float qreg[DPT];
|
float qreg[DPT];
|
||||||
if (q_row < p.q_len) {
|
if (q_row < q_len) {
|
||||||
int q_off = batch * p.q_stride_b + q_head * p.q_stride_h
|
int q_off = q_base + q_row * p.q_l_stride + gpos * DPT * p.q_d_stride;
|
||||||
+ q_row * p.q_stride_l + gpos * DPT * p.q_stride_d;
|
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int i = 0; i < DPT; i++)
|
for (int i = 0; i < DPT; i++)
|
||||||
qreg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]);
|
qreg[i] = __bfloat162float(p.q_ptr[q_off + i * p.q_d_stride]);
|
||||||
}
|
}
|
||||||
|
|
||||||
float m = -FLT_MAX, l = 0.0f;
|
float m = -FLT_MAX, l = 0.0f;
|
||||||
@@ -62,10 +68,8 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
|
|||||||
for (int i = 0; i < DPT; i++)
|
for (int i = 0; i < DPT; i++)
|
||||||
acc[i] = 0.0f;
|
acc[i] = 0.0f;
|
||||||
|
|
||||||
// KV: stride-based base
|
|
||||||
int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
|
||||||
int mask_batch_base = batch * p.mask_b_stride + q_head * p.mask_h_stride;
|
int mask_batch_base = batch * p.mask_b_stride + q_head * p.mask_h_stride;
|
||||||
int tiles = (p.kv_len + P_BC - 1) / P_BC;
|
int tiles = (seq_len + P_BC - 1) / P_BC;
|
||||||
int tt = G * ROWS;
|
int tt = G * ROWS;
|
||||||
int lid = row * G + gpos;
|
int lid = row * G + gpos;
|
||||||
|
|
||||||
@@ -75,23 +79,25 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
|
|||||||
|
|
||||||
for (int ti = 0; ti < tiles; ti++) {
|
for (int ti = 0; ti < tiles; ti++) {
|
||||||
int kv0 = ti * P_BC;
|
int kv0 = ti * P_BC;
|
||||||
int tlen = min(P_BC, p.kv_len - kv0);
|
int tlen = min(P_BC, seq_len - kv0);
|
||||||
|
|
||||||
// Load K/V into shared memory from strided global
|
// Load K/V into shared memory (addressing via KV policy; paged
|
||||||
|
// guards empty slots with zero-fill).
|
||||||
for (int i = lid; i < tlen * HEAD_DIM; i += tt) {
|
for (int i = lid; i < tlen * HEAD_DIM; i += tt) {
|
||||||
int s = i / HEAD_DIM;
|
int s = i / HEAD_DIM;
|
||||||
int d_dim = i % HEAD_DIM;
|
int d_dim = i % HEAD_DIM;
|
||||||
int kv_idx = kv0 + s;
|
int kc = kv0 + s;
|
||||||
int g_off = kv_base + kv_idx * p.kv_stride_l + d_dim * p.kv_stride_d;
|
int token = KV::resolve_token(p, kctx, kc, true);
|
||||||
sK[i] = p.k[g_off];
|
KVAddr a = KV::kv_addr_from_token(p, kctx, token, d_dim);
|
||||||
sV[i] = p.v[g_off];
|
sK[i] = a.valid ? *reinterpret_cast<const bf16*>(a.k) : (bf16)0.f;
|
||||||
|
sV[i] = a.valid ? *reinterpret_cast<const bf16*>(a.v) : (bf16)0.f;
|
||||||
}
|
}
|
||||||
__syncthreads();
|
__syncthreads();
|
||||||
|
|
||||||
int lim = tlen;
|
int lim = tlen;
|
||||||
if constexpr (IsCausal) {
|
if constexpr (IsCausal) {
|
||||||
if (q_row < p.q_len) {
|
if (q_row < q_len) {
|
||||||
int ep = q_row + p.causal_offset + 1;
|
int ep = causal_off + q_row + 1;
|
||||||
if (kv0 >= ep)
|
if (kv0 >= ep)
|
||||||
lim = 0;
|
lim = 0;
|
||||||
else if (kv0 + tlen > ep)
|
else if (kv0 + tlen > ep)
|
||||||
@@ -99,7 +105,7 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
int mask_row_base = mask_batch_base + q_row * p.mask_q_stride;
|
int mask_row_base = mask_batch_base + q_row * p.mask_l_stride;
|
||||||
for (int s = 0; s < lim; s++) {
|
for (int s = 0; s < lim; s++) {
|
||||||
const bf16* kr = sK + s * HEAD_DIM + gpos * DPT;
|
const bf16* kr = sK + s * HEAD_DIM + gpos * DPT;
|
||||||
float part = 0.0f;
|
float part = 0.0f;
|
||||||
@@ -138,12 +144,14 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
|
|||||||
__syncthreads();
|
__syncthreads();
|
||||||
}
|
}
|
||||||
|
|
||||||
if (q_row < p.q_len) {
|
if (q_row < q_len) {
|
||||||
int o_off = batch * p.q_stride_b + q_head * p.q_stride_h
|
int o_off = q_base + q_row * p.q_l_stride + gpos * DPT * p.q_d_stride;
|
||||||
+ q_row * p.q_stride_l + gpos * DPT * p.q_stride_d;
|
|
||||||
float rl = (l > 1e-20f) ? (1.0f / l) : 0.0f;
|
float rl = (l > 1e-20f) ? (1.0f / l) : 0.0f;
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int i = 0; i < DPT; i++)
|
for (int i = 0; i < DPT; i++)
|
||||||
p.o[o_off + i * p.q_stride_d] = __float2bfloat16(acc[i] * rl);
|
p.o_ptr[o_off + i * p.q_d_stride] = __float2bfloat16(acc[i] * rl);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
} // namespace attention
|
||||||
|
} // namespace astrai
|
||||||
+74
-38
@@ -1,28 +1,61 @@
|
|||||||
#pragma once
|
#pragma once
|
||||||
#include <cfloat>
|
#include <cfloat>
|
||||||
#include <cuda_bf16.h>
|
#include <cuda_bf16.h>
|
||||||
#include "attn_common.h"
|
#include "common.h"
|
||||||
#include "attn_mma_utils.cuh"
|
#include "layout_policies.cuh"
|
||||||
|
#include "mma_utils.cuh"
|
||||||
|
|
||||||
// Tensor-core prefill flash attention (raw mma.sync PTX).
|
namespace astrai {
|
||||||
|
namespace attention {
|
||||||
|
|
||||||
|
// Tensor-core prefill flash attention (raw mma.sync PTX), unified across
|
||||||
|
// contiguous and paged (SGLang flat-pool) K/V via the KV template parameter.
|
||||||
// One warp owns BR=16 query rows. S = Q@K^T and O = P@V run on bf16 tensor
|
// One warp owns BR=16 query rows. S = Q@K^T and O = P@V run on bf16 tensor
|
||||||
// cores via mma.sync.m16n8k16 (f32 accumulate).
|
// cores via mma.sync.m16n8k16 (f32 accumulate).
|
||||||
//
|
//
|
||||||
|
// GQA head packing (FA2/FA3-style): HB = min(G, WARPS) query heads of one
|
||||||
|
// kv-head group share a block's K/V tiles, so each K/V element is read from
|
||||||
|
// global memory once per block instead of once per q head (~HB× less K/V
|
||||||
|
// traffic). WARPS = WPH × HB: warp w handles head slot w/WPH, chunk w%WPH;
|
||||||
|
// all warps of a block cover the same token range, keeping the causal sweep
|
||||||
|
// end block-uniform. G=1 (MHA) degenerates to the unpadded layout.
|
||||||
|
//
|
||||||
|
// KV = ContigKV (dense [batch, kv_head, kv_len, head_dim]) or PagedKV
|
||||||
|
// (flat pool + req_to_token, ragged batches via qo_indptr/kv_indptr).
|
||||||
// IsCausal and HasMask are compile-time bools — the compiler eliminates all
|
// IsCausal and HasMask are compile-time bools — the compiler eliminates all
|
||||||
// dead branches in the inner compute loop (FA2-style).
|
// dead branches in the inner compute loop (FA2-style).
|
||||||
//
|
//
|
||||||
// Traits = KernelTraits<HEAD_DIM, BC, WARPS=4, STAGES=2>.
|
// Traits = KernelTraits<HEAD_DIM, BC, WARPS=4, STAGES=2>.
|
||||||
template <typename Traits, bool IsCausal, bool HasMask>
|
template <typename Traits, typename QSchedule, typename KV, bool IsCausal, bool HasMask>
|
||||||
__global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
|
__global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
|
||||||
const int warp = threadIdx.x / 32;
|
const int warp = threadIdx.x / 32;
|
||||||
const int lane = threadIdx.x % 32;
|
const int lane = threadIdx.x % 32;
|
||||||
const int gid = lane >> 2; // 0..7
|
const int gid = lane >> 2; // 0..7
|
||||||
const int tid4 = lane & 3; // 0..3
|
const int tid4 = lane & 3; // 0..3
|
||||||
|
|
||||||
const int q_head = blockIdx.y;
|
const int G = p.q_head / p.kv_head;
|
||||||
const int batch = blockIdx.z;
|
const int HB = min(G, Traits::WARPS); // q heads packed per block
|
||||||
const int kv_head = q_head / (p.q_head / p.kv_head);
|
const int WPH = Traits::WARPS / HB; // 16-row chunks per head
|
||||||
const int qrow0 = (blockIdx.x * Traits::WARPS + warp) * Traits::BR;
|
const int BPG = (G + HB - 1) / HB; // blocks per GQA group
|
||||||
|
const int chunk = warp % WPH;
|
||||||
|
|
||||||
|
int batch, row_base;
|
||||||
|
QSchedule::map_packed_block(p, Traits::BR * WPH, batch, row_base);
|
||||||
|
const int kv_head = blockIdx.y / BPG;
|
||||||
|
const int slot = blockIdx.y - kv_head * BPG;
|
||||||
|
const int head_idx = slot * HB + warp / WPH;
|
||||||
|
// G % HB tail blocks have idle head slots: clamp to the last head so all
|
||||||
|
// warps do valid work (cp.async + __syncthreads stay block-uniform) and
|
||||||
|
// just skip the O store via `active`.
|
||||||
|
const bool active = head_idx < G;
|
||||||
|
const int q_head = kv_head * G + min(head_idx, G - 1);
|
||||||
|
const int qrow0 = row_base + chunk * Traits::BR;
|
||||||
|
|
||||||
|
// Per-request dims (from KV policy — paged reads kv_indptr/qo_indptr).
|
||||||
|
const int seq_len = KV::kv_len(p, batch);
|
||||||
|
const int q_len = QSchedule::q_len(p, batch);
|
||||||
|
const int causal_off = KV::causal_offset(p, batch, q_len);
|
||||||
|
const KVContext kctx = KV::template make_ctx<Traits::HEAD_DIM>(p, batch, kv_head);
|
||||||
|
|
||||||
// Static shared memory: double-buffered K/V (no sQ — Q goes direct
|
// Static shared memory: double-buffered K/V (no sQ — Q goes direct
|
||||||
// to registers in mma A-operand layout).
|
// to registers in mma A-operand layout).
|
||||||
@@ -30,12 +63,12 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
|
|||||||
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
|
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
|
||||||
|
|
||||||
// Load Q fragments straight from global into mma A-operand layout.
|
// Load Q fragments straight from global into mma A-operand layout.
|
||||||
const int q_base = batch * p.q_stride_b + q_head * p.q_stride_h;
|
const int q_base = QSchedule::q_base(p, batch, q_head);
|
||||||
const int qra = qrow0 + gid;
|
const int qra = qrow0 + gid;
|
||||||
const int qrb = qrow0 + gid + 8;
|
const int qrb = qrow0 + gid + 8;
|
||||||
const bool va = qra < p.q_len, vb = qrb < p.q_len;
|
const bool va = qra < q_len, vb = qrb < q_len;
|
||||||
unsigned Qa[Traits::KD][4];
|
unsigned Qa[Traits::KD][4];
|
||||||
load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_l, p.q_stride_d,
|
load_q_mma_frags<Traits::KD>(p.q_ptr + q_base, p.q_l_stride, p.q_d_stride,
|
||||||
qra, qrb, va, vb, tid4, Qa);
|
qra, qrb, va, vb, tid4, Qa);
|
||||||
|
|
||||||
float Oacc[Traits::DN8][4];
|
float Oacc[Traits::DN8][4];
|
||||||
@@ -44,17 +77,15 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
|
|||||||
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
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;
|
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
||||||
|
|
||||||
// KV: stride-based base
|
const int tiles = (seq_len + Traits::BC - 1) / Traits::BC;
|
||||||
const int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
|
||||||
const int tiles = (p.kv_len + Traits::BC - 1) / Traits::BC;
|
|
||||||
const int qr0 = qrow0 + gid;
|
const int qr0 = qrow0 + gid;
|
||||||
const int qr1 = qrow0 + gid + 8;
|
const int qr1 = qrow0 + gid + 8;
|
||||||
|
|
||||||
// Causal tile-skip bounds (dead code when IsCausal == false)
|
// Causal tile-skip bounds (dead code when IsCausal == false).
|
||||||
const int max_kv = qrow0 + Traits::BR - 1 + p.causal_offset;
|
// max_kv is per-warp (its own 16 rows); block_max_kv is the last row of
|
||||||
const int block_max_kv =
|
// the whole block's range and must be uniform for the shared sweep loop.
|
||||||
blockIdx.x * Traits::WARPS * Traits::BR + Traits::WARPS * Traits::BR - 1
|
const int max_kv = qrow0 + Traits::BR - 1 + causal_off;
|
||||||
+ p.causal_offset;
|
const int block_max_kv = row_base + WPH * Traits::BR - 1 + causal_off;
|
||||||
|
|
||||||
int t_end = tiles - 1;
|
int t_end = tiles - 1;
|
||||||
if constexpr (IsCausal) {
|
if constexpr (IsCausal) {
|
||||||
@@ -62,7 +93,7 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
|
|||||||
if (bt < t_end) t_end = bt;
|
if (bt < t_end) t_end = bt;
|
||||||
}
|
}
|
||||||
|
|
||||||
// ---- Load tile lambda: predicated cp.async ----
|
// ---- Load tile lambda: predicated cp.async (addressing via KV policy) ----
|
||||||
auto load_tile = [&](int ti, int buf) {
|
auto load_tile = [&](int ti, int buf) {
|
||||||
int kv0 = ti * Traits::BC;
|
int kv0 = ti * Traits::BC;
|
||||||
bf16* dK = sK + buf * Traits::BC * Traits::LD;
|
bf16* dK = sK + buf * Traits::BC * Traits::LD;
|
||||||
@@ -72,13 +103,14 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
|
|||||||
i += Traits::NUM_THREADS * Traits::VEC) {
|
i += Traits::NUM_THREADS * Traits::VEC) {
|
||||||
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
|
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
|
||||||
int kc = kv0 + r;
|
int kc = kv0 + r;
|
||||||
bool valid = kc < p.kv_len;
|
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 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;
|
astrai::cp_async_16(&dK[off], a.k, a.valid);
|
||||||
cp_async_16_pred(&dK[off], &p.k[g_off], valid);
|
astrai::cp_async_16(&dV[off], a.v, a.valid);
|
||||||
cp_async_16_pred(&dV[off], &p.v[g_off], valid);
|
|
||||||
}
|
}
|
||||||
cp_async_commit();
|
astrai::cp_async_commit_group();
|
||||||
};
|
};
|
||||||
|
|
||||||
// ---- Prologue: issue first tile load ----
|
// ---- Prologue: issue first tile load ----
|
||||||
@@ -88,7 +120,7 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
|
|||||||
int buf = ti & 1;
|
int buf = ti & 1;
|
||||||
|
|
||||||
// Wait for current tile, then publish cross-warp + guard buffer reuse.
|
// Wait for current tile, then publish cross-warp + guard buffer reuse.
|
||||||
cp_async_wait_group<0>();
|
astrai::cp_async_wait_group<0>();
|
||||||
__syncthreads();
|
__syncthreads();
|
||||||
if (ti < t_end) load_tile(ti + 1, (ti + 1) & 1);
|
if (ti < t_end) load_tile(ti + 1, (ti + 1) & 1);
|
||||||
|
|
||||||
@@ -108,15 +140,16 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
|
|||||||
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
||||||
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
||||||
|
|
||||||
int maxc0 = IsCausal ? min(p.kv_len, qr0 + p.causal_offset + 1)
|
int maxc0 = IsCausal ? min(seq_len, causal_off + qr0 + 1)
|
||||||
: p.kv_len;
|
: seq_len;
|
||||||
int maxc1 = IsCausal ? min(p.kv_len, qr1 + p.causal_offset + 1)
|
int maxc1 = IsCausal ? min(seq_len, causal_off + qr1 + 1)
|
||||||
: p.kv_len;
|
: seq_len;
|
||||||
mma_softmax_tile<Traits, HasMask>(kv0, maxc0, maxc1,
|
mma_softmax_tile<Traits, HasMask>(kv0, maxc0, maxc1,
|
||||||
qr0, qr1,
|
qr0, qr1,
|
||||||
p.mask_b_stride, p.mask_h_stride, p.mask_q_stride,
|
p.mask_b_stride, p.mask_h_stride, p.mask_l_stride,
|
||||||
batch, q_head,
|
batch, q_head, q_head,
|
||||||
p.mask,
|
p.mask,
|
||||||
|
va, vb,
|
||||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||||
|
|
||||||
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
|
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
|
||||||
@@ -126,21 +159,24 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
|
|||||||
// ---- write output: packed bf16x2 stores ----
|
// ---- write output: packed bf16x2 stores ----
|
||||||
float rl0 = (l0 > 1e-20f) ? (1.0f / l0) : 0.0f;
|
float rl0 = (l0 > 1e-20f) ? (1.0f / l0) : 0.0f;
|
||||||
float rl1 = (l1 > 1e-20f) ? (1.0f / l1) : 0.0f;
|
float rl1 = (l1 > 1e-20f) ? (1.0f / l1) : 0.0f;
|
||||||
const int o_base = batch * p.q_stride_b + q_head * p.q_stride_h;
|
const int o_base = QSchedule::q_base(p, batch, q_head);
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
|
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
|
||||||
int d = dn8 * 8 + 2 * tid4;
|
int d = dn8 * 8 + 2 * tid4;
|
||||||
if (qr0 < p.q_len) {
|
if (active && qr0 < q_len) {
|
||||||
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][0] * rl0,
|
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][0] * rl0,
|
||||||
Oacc[dn8][1] * rl0);
|
Oacc[dn8][1] * rl0);
|
||||||
*reinterpret_cast<__nv_bfloat162*>(
|
*reinterpret_cast<__nv_bfloat162*>(
|
||||||
&p.o[o_base + qr0 * p.q_stride_l + d * p.q_stride_d]) = v;
|
&p.o_ptr[o_base + qr0 * p.q_l_stride + d * p.q_d_stride]) = v;
|
||||||
}
|
}
|
||||||
if (qr1 < p.q_len) {
|
if (active && qr1 < q_len) {
|
||||||
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][2] * rl1,
|
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][2] * rl1,
|
||||||
Oacc[dn8][3] * rl1);
|
Oacc[dn8][3] * rl1);
|
||||||
*reinterpret_cast<__nv_bfloat162*>(
|
*reinterpret_cast<__nv_bfloat162*>(
|
||||||
&p.o[o_base + qr1 * p.q_stride_l + d * p.q_stride_d]) = v;
|
&p.o_ptr[o_base + qr1 * p.q_l_stride + d * p.q_d_stride]) = v;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
} // namespace attention
|
||||||
|
} // namespace astrai
|
||||||
@@ -1,71 +0,0 @@
|
|||||||
#pragma once
|
|
||||||
|
|
||||||
|
|
||||||
template<typename T, typename AT = float>
|
|
||||||
struct AttentionParams {
|
|
||||||
int batch;
|
|
||||||
int q_head;
|
|
||||||
int kv_head;
|
|
||||||
int q_len;
|
|
||||||
int kv_len;
|
|
||||||
int head_dim;
|
|
||||||
int use_mask;
|
|
||||||
int causal_offset; // -1 = non-causal; >=0 = absolute position of first Q token
|
|
||||||
int num_splits;
|
|
||||||
float scale;
|
|
||||||
|
|
||||||
// Q strides (element offsets for each dim — layout-agnostic)
|
|
||||||
int q_stride_b, q_stride_h, q_stride_l, q_stride_d;
|
|
||||||
// KV strides (K and V share the same layout — only base pointers differ)
|
|
||||||
int kv_stride_b, kv_stride_h, kv_stride_l, kv_stride_d;
|
|
||||||
|
|
||||||
// Mask: 2D [batch, kv_len], 3D [batch, q_len, kv_len],
|
|
||||||
// or 4D [batch, n_heads, q_len, kv_len] (head dim broadcasts when stride=0)
|
|
||||||
int mask_b_stride; // batch stride
|
|
||||||
int mask_h_stride; // head stride (0 = broadcast across heads)
|
|
||||||
int mask_q_stride; // q stride (0 = all q rows share)
|
|
||||||
|
|
||||||
const T* __restrict__ q;
|
|
||||||
const T* __restrict__ k;
|
|
||||||
const T* __restrict__ v;
|
|
||||||
const bool* __restrict__ mask;
|
|
||||||
|
|
||||||
T* __restrict__ o;
|
|
||||||
AT* __restrict__ o_part;
|
|
||||||
AT* __restrict__ ml_part;
|
|
||||||
};
|
|
||||||
|
|
||||||
template<typename T, typename AT = float>
|
|
||||||
struct PagedAttentionParams {
|
|
||||||
int batch;
|
|
||||||
int q_head;
|
|
||||||
int kv_head;
|
|
||||||
int q_len;
|
|
||||||
int kv_len;
|
|
||||||
int head_dim;
|
|
||||||
int use_mask;
|
|
||||||
int causal_offset;
|
|
||||||
float scale;
|
|
||||||
|
|
||||||
int num_splits;
|
|
||||||
int page_size;
|
|
||||||
int max_pages;
|
|
||||||
|
|
||||||
// Q strides (layout-agnostic)
|
|
||||||
int q_stride_b, q_stride_h, q_stride_l, q_stride_d;
|
|
||||||
|
|
||||||
// Mask strides (2D, 3D, or 4D)
|
|
||||||
int mask_b_stride;
|
|
||||||
int mask_h_stride;
|
|
||||||
int mask_q_stride;
|
|
||||||
|
|
||||||
const T* __restrict__ q;
|
|
||||||
const T* __restrict__ k_cache;
|
|
||||||
const T* __restrict__ v_cache;
|
|
||||||
const bool* __restrict__ mask;
|
|
||||||
const int64_t* __restrict__ page_table;
|
|
||||||
|
|
||||||
T* __restrict__ o;
|
|
||||||
AT* __restrict__ o_part;
|
|
||||||
AT* __restrict__ ml_part;
|
|
||||||
};
|
|
||||||
@@ -1,37 +0,0 @@
|
|||||||
#include "attn_dispatchers.cuh"
|
|
||||||
#include "attn_entry_utils.cuh"
|
|
||||||
|
|
||||||
torch::Tensor attn_decode(
|
|
||||||
torch::Tensor q,
|
|
||||||
torch::Tensor k,
|
|
||||||
torch::Tensor v,
|
|
||||||
c10::optional<torch::Tensor> mask,
|
|
||||||
int64_t causal_offset,
|
|
||||||
double scale,
|
|
||||||
int64_t layout
|
|
||||||
) {
|
|
||||||
AttentionParams<bf16> p;
|
|
||||||
attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p);
|
|
||||||
TORCH_CHECK(p.q_len == 1, "Q seq_len must be 1");
|
|
||||||
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
|
|
||||||
|
|
||||||
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
|
|
||||||
auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
|
|
||||||
p.o = (bf16*)O_view.data_ptr();
|
|
||||||
|
|
||||||
alloc_split_partials(p);
|
|
||||||
DISPATCH_HEAD_DIM(p.head_dim, dispatch_decode, p);
|
|
||||||
return O;
|
|
||||||
}
|
|
||||||
|
|
||||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
|
||||||
m.def("attn_decode", &attn_decode,
|
|
||||||
py::arg("q"),
|
|
||||||
py::arg("k"),
|
|
||||||
py::arg("v"),
|
|
||||||
py::arg("mask") = py::none(),
|
|
||||||
py::arg("causal_offset") = -1,
|
|
||||||
py::arg("scale") = 0.0,
|
|
||||||
py::arg("layout") = 0,
|
|
||||||
"GQA decode (tensor-core head-packing on sm_80+, scalar fallback)");
|
|
||||||
}
|
|
||||||
@@ -1,195 +0,0 @@
|
|||||||
#pragma once
|
|
||||||
// Shared attention dispatchers — used by both production .cu and test .cu.
|
|
||||||
// No torch dependency; pure CUDA.
|
|
||||||
|
|
||||||
#include <cuda_runtime.h>
|
|
||||||
#include <algorithm>
|
|
||||||
#include "attn_warp_utils.cuh"
|
|
||||||
#include "attn_prefill_split_q.cuh"
|
|
||||||
#include "attn_decode_split_kv.cuh"
|
|
||||||
#include "attn_paged_decode_split_kv.cuh"
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
|
||||||
#include "attn_prefill_split_q_mma.cuh"
|
|
||||||
#include "attn_decode_split_kv_mma.cuh"
|
|
||||||
#include "attn_paged_decode_split_kv_mma.cuh"
|
|
||||||
#endif
|
|
||||||
|
|
||||||
// Split-KV: compute number of splits to fill all SMs for small-batch decode.
|
|
||||||
// Caps splits so each split processes at least `min_tiles_per_split` tiles,
|
|
||||||
// avoiding excessive loop/prologue overhead when tiles are small.
|
|
||||||
inline int compute_num_splits(int base_blocks, int tiles_total,
|
|
||||||
int min_tiles_per_split = 1) {
|
|
||||||
int sm_count = 0;
|
|
||||||
cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0);
|
|
||||||
int n = (2 * sm_count + base_blocks - 1) / base_blocks;
|
|
||||||
int max_by_work = tiles_total / min_tiles_per_split;
|
|
||||||
return std::max(1, std::min(n, std::min(max_by_work, MAX_SPLITS)));
|
|
||||||
}
|
|
||||||
|
|
||||||
// ======================================================================
|
|
||||||
// Prefill
|
|
||||||
// ======================================================================
|
|
||||||
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
|
||||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
|
||||||
static inline void launch_prefill_mma(AttentionParams<bf16>& p) {
|
|
||||||
constexpr int WARPS = 4;
|
|
||||||
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
|
|
||||||
using Traits = KernelTraits<HEAD_DIM, BC, WARPS, 2>;
|
|
||||||
dim3 grid((p.q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS), p.q_head, p.batch);
|
|
||||||
dim3 block(Traits::NUM_THREADS);
|
|
||||||
attn_prefill_split_q_mma_kernel<Traits, IsCausal, HasMask><<<grid, block>>>(p);
|
|
||||||
}
|
|
||||||
#endif
|
|
||||||
|
|
||||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
|
||||||
static inline void launch_prefill_scalar(AttentionParams<bf16>& p) {
|
|
||||||
constexpr int G = 8, ROWS = 32, P_BC = 32;
|
|
||||||
dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch);
|
|
||||||
dim3 block(G, ROWS);
|
|
||||||
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC, IsCausal, HasMask><<<grid, block>>>(p);
|
|
||||||
}
|
|
||||||
|
|
||||||
template <int HEAD_DIM>
|
|
||||||
static inline void dispatch_prefill(AttentionParams<bf16>& p) {
|
|
||||||
bool is_causal = (p.causal_offset >= 0);
|
|
||||||
bool has_mask = (p.use_mask && p.mask);
|
|
||||||
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
|
||||||
if (is_causal) {
|
|
||||||
if (has_mask) launch_prefill_mma<HEAD_DIM, true, true>(p);
|
|
||||||
else launch_prefill_mma<HEAD_DIM, true, false>(p);
|
|
||||||
} else {
|
|
||||||
if (has_mask) launch_prefill_mma<HEAD_DIM, false, true>(p);
|
|
||||||
else launch_prefill_mma<HEAD_DIM, false, false>(p);
|
|
||||||
}
|
|
||||||
#else
|
|
||||||
if (is_causal) {
|
|
||||||
if (has_mask) launch_prefill_scalar<HEAD_DIM, true, true>(p);
|
|
||||||
else launch_prefill_scalar<HEAD_DIM, true, false>(p);
|
|
||||||
} else {
|
|
||||||
if (has_mask) launch_prefill_scalar<HEAD_DIM, false, true>(p);
|
|
||||||
else launch_prefill_scalar<HEAD_DIM, false, false>(p);
|
|
||||||
}
|
|
||||||
#endif
|
|
||||||
}
|
|
||||||
|
|
||||||
// ======================================================================
|
|
||||||
// Decode
|
|
||||||
// ======================================================================
|
|
||||||
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
|
||||||
// BC=16: halves smem (16KB vs 32KB) → doubles occupancy (6 vs 3 blocks/SM).
|
|
||||||
// For D=256, BC=16 also reduces register pressure (fewer Sacc/PV frags),
|
|
||||||
// enabling STAGES=2 (double-buffer) within the 32KB smem budget — eliminates
|
|
||||||
// the 176-byte spill that STAGES=1+BC=32 suffered.
|
|
||||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
|
||||||
static inline void launch_decode_mma(AttentionParams<bf16>& p, int group_size) {
|
|
||||||
int G = p.q_head / p.kv_head;
|
|
||||||
constexpr int MAX_G = 16;
|
|
||||||
int num_passes = (G + MAX_G - 1) / MAX_G;
|
|
||||||
constexpr int BC = 16;
|
|
||||||
int tiles_total = (p.kv_len + BC - 1) / BC;
|
|
||||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total, 2);
|
|
||||||
constexpr int STAGES = 2;
|
|
||||||
using Traits = KernelTraits<HEAD_DIM, BC, 1, STAGES>;
|
|
||||||
dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
|
|
||||||
attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask><<<grid, 32>>>(p);
|
|
||||||
}
|
|
||||||
#endif
|
|
||||||
|
|
||||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
|
||||||
static inline void launch_decode_scalar(AttentionParams<bf16>& p, int group_size) {
|
|
||||||
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
|
|
||||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
|
||||||
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
|
|
||||||
int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit
|
|
||||||
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
|
|
||||||
dim3 block(32, g);
|
|
||||||
attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
|
|
||||||
}
|
|
||||||
|
|
||||||
template <int HEAD_DIM>
|
|
||||||
static inline void dispatch_decode(AttentionParams<bf16>& p) {
|
|
||||||
bool is_causal = (p.causal_offset >= 0);
|
|
||||||
bool has_mask = (p.use_mask && p.mask);
|
|
||||||
int group_size = p.q_head / p.kv_head;
|
|
||||||
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
|
||||||
if (is_causal) {
|
|
||||||
if (has_mask) launch_decode_mma<HEAD_DIM, true, true>(p, group_size);
|
|
||||||
else launch_decode_mma<HEAD_DIM, true, false>(p, group_size);
|
|
||||||
} else {
|
|
||||||
if (has_mask) launch_decode_mma<HEAD_DIM, false, true>(p, group_size);
|
|
||||||
else launch_decode_mma<HEAD_DIM, false, false>(p, group_size);
|
|
||||||
}
|
|
||||||
#else
|
|
||||||
if (is_causal) {
|
|
||||||
if (has_mask) launch_decode_scalar<HEAD_DIM, true, true>(p, group_size);
|
|
||||||
else launch_decode_scalar<HEAD_DIM, true, false>(p, group_size);
|
|
||||||
} else {
|
|
||||||
if (has_mask) launch_decode_scalar<HEAD_DIM, false, true>(p, group_size);
|
|
||||||
else launch_decode_scalar<HEAD_DIM, false, false>(p, group_size);
|
|
||||||
}
|
|
||||||
#endif
|
|
||||||
|
|
||||||
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
|
||||||
}
|
|
||||||
|
|
||||||
// ======================================================================
|
|
||||||
// Paged Decode
|
|
||||||
// ======================================================================
|
|
||||||
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
|
||||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
|
||||||
static inline void launch_paged_decode_mma(PagedAttentionParams<bf16>& p, int group_size) {
|
|
||||||
int G = p.q_head / p.kv_head;
|
|
||||||
constexpr int MAX_G = 16;
|
|
||||||
constexpr int BC = 16;
|
|
||||||
int num_passes = (G + MAX_G - 1) / MAX_G;
|
|
||||||
int tiles_total = (p.kv_len + BC - 1) / BC;
|
|
||||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total, 2);
|
|
||||||
constexpr int STAGES = 2;
|
|
||||||
using Traits = KernelTraits<HEAD_DIM, BC, 1, STAGES>;
|
|
||||||
dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
|
|
||||||
paged_attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask> <<<grid, 32>>>(p);
|
|
||||||
}
|
|
||||||
#endif
|
|
||||||
|
|
||||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
|
||||||
static inline void launch_paged_decode_scalar(PagedAttentionParams<bf16>& p, int group_size) {
|
|
||||||
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
|
|
||||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
|
||||||
size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16);
|
|
||||||
int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit
|
|
||||||
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
|
|
||||||
dim3 block(32, g);
|
|
||||||
paged_attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
|
|
||||||
}
|
|
||||||
|
|
||||||
template <int HEAD_DIM>
|
|
||||||
static inline void dispatch_paged_decode(PagedAttentionParams<bf16>& p) {
|
|
||||||
bool is_causal = (p.causal_offset >= 0);
|
|
||||||
bool has_mask = (p.use_mask && p.mask);
|
|
||||||
int group_size = p.q_head / p.kv_head;
|
|
||||||
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
|
||||||
if (is_causal) {
|
|
||||||
if (has_mask) launch_paged_decode_mma<HEAD_DIM, true, true>(p, group_size);
|
|
||||||
else launch_paged_decode_mma<HEAD_DIM, true, false>(p, group_size);
|
|
||||||
} else {
|
|
||||||
if (has_mask) launch_paged_decode_mma<HEAD_DIM, false, true>(p, group_size);
|
|
||||||
else launch_paged_decode_mma<HEAD_DIM, false, false>(p, group_size);
|
|
||||||
}
|
|
||||||
#else
|
|
||||||
if (is_causal) {
|
|
||||||
if (has_mask) launch_paged_decode_scalar<HEAD_DIM, true, true>(p, group_size);
|
|
||||||
else launch_paged_decode_scalar<HEAD_DIM, true, false>(p, group_size);
|
|
||||||
} else {
|
|
||||||
if (has_mask) launch_paged_decode_scalar<HEAD_DIM, false, true>(p, group_size);
|
|
||||||
else launch_paged_decode_scalar<HEAD_DIM, false, false>(p, group_size);
|
|
||||||
}
|
|
||||||
#endif
|
|
||||||
|
|
||||||
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
|
||||||
}
|
|
||||||
@@ -1,187 +0,0 @@
|
|||||||
#pragma once
|
|
||||||
#include <float.h>
|
|
||||||
#include <torch/extension.h>
|
|
||||||
#include <c10/cuda/CUDAGuard.h>
|
|
||||||
#include "attn_common.h"
|
|
||||||
#include "attn_warp_utils.cuh"
|
|
||||||
|
|
||||||
using bf16 = __nv_bfloat16;
|
|
||||||
|
|
||||||
// Dispatch head_dim: shared macro — avoids C++20 lambda template syntax.
|
|
||||||
// Usage: DISPATCH_HEAD_DIM(hd, fn, arg)
|
|
||||||
// Expands to: fn<32>(arg); fn<64>(arg); etc.
|
|
||||||
#define DISPATCH_HEAD_DIM(hd, fn, arg) \
|
|
||||||
switch (hd) { \
|
|
||||||
case 32: fn<32>(arg); break; \
|
|
||||||
case 64: fn<64>(arg); break; \
|
|
||||||
case 128: fn<128>(arg); break; \
|
|
||||||
case 256: fn<256>(arg); break; \
|
|
||||||
default: \
|
|
||||||
TORCH_CHECK(false, "unsupported head_dim ", hd, \
|
|
||||||
" (supported: 32, 64, 128, 256)"); \
|
|
||||||
}
|
|
||||||
|
|
||||||
// The split kernel unconditionally writes every (batch, q_head, split) slot it
|
|
||||||
// owns — including empty split ranges, which store m = -FLT_MAX so the combine
|
|
||||||
// skips them. Allocators are therefore left uninitialized (torch::empty); the
|
|
||||||
// per-call memset (torch::zeros / torch::full) was pure overhead.
|
|
||||||
template<typename P>
|
|
||||||
inline void alloc_split_partials(P& p) {
|
|
||||||
auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA);
|
|
||||||
auto o_part = torch::empty(at::IntArrayRef{p.batch, p.q_head, MAX_SPLITS, p.head_dim}, fopt);
|
|
||||||
auto ml_part = torch::empty(at::IntArrayRef{p.batch, p.q_head, MAX_SPLITS, 2}, fopt);
|
|
||||||
p.o_part = (float*)o_part.data_ptr();
|
|
||||||
p.ml_part = (float*)ml_part.data_ptr();
|
|
||||||
}
|
|
||||||
|
|
||||||
// ---- Shared Q-dims + strides extraction ----
|
|
||||||
template <typename P>
|
|
||||||
inline void extract_q_dims_and_strides(torch::Tensor& q, int64_t layout, P& p) {
|
|
||||||
if (layout == 1) q = q.transpose(1, 2);
|
|
||||||
p.batch = (int)q.size(0);
|
|
||||||
p.q_head = (int)q.size(1);
|
|
||||||
p.q_len = (int)q.size(2);
|
|
||||||
p.head_dim = (int)q.size(3);
|
|
||||||
p.q_stride_b = (int)q.stride(0);
|
|
||||||
p.q_stride_h = (int)q.stride(1);
|
|
||||||
p.q_stride_l = (int)q.stride(2);
|
|
||||||
p.q_stride_d = (int)q.stride(3);
|
|
||||||
}
|
|
||||||
|
|
||||||
// ---- Shared mask packing ----
|
|
||||||
// Accepts 2D [batch, kv_len], 3D [batch, q_len, kv_len],
|
|
||||||
// or 4D [batch, n_heads, q_len, kv_len].
|
|
||||||
// Head/q dimensions with size 1 broadcast (stride set to 0).
|
|
||||||
template <typename P>
|
|
||||||
inline void pack_mask(const c10::optional<torch::Tensor>& mask, P& p) {
|
|
||||||
if (p.use_mask) {
|
|
||||||
auto m = mask.value();
|
|
||||||
TORCH_CHECK(m.is_cuda(), "mask must be on CUDA");
|
|
||||||
TORCH_CHECK(m.dtype() == torch::kBool, "mask must be bool");
|
|
||||||
TORCH_CHECK(m.size(0) == p.batch, "mask batch mismatch");
|
|
||||||
TORCH_CHECK(m.size(m.dim() - 1) == p.kv_len, "mask kv_len mismatch");
|
|
||||||
if (m.dim() == 2) {
|
|
||||||
p.mask_b_stride = (int)m.stride(0);
|
|
||||||
p.mask_h_stride = 0;
|
|
||||||
p.mask_q_stride = 0;
|
|
||||||
} else if (m.dim() == 3) {
|
|
||||||
TORCH_CHECK(m.size(1) == 1 || m.size(1) == p.q_len, "mask q_len mismatch");
|
|
||||||
p.mask_b_stride = (int)m.stride(0);
|
|
||||||
p.mask_h_stride = 0;
|
|
||||||
p.mask_q_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1);
|
|
||||||
} else if (m.dim() == 4) {
|
|
||||||
TORCH_CHECK(m.size(2) == 1 || m.size(2) == p.q_len, "mask q_len mismatch");
|
|
||||||
p.mask_b_stride = (int)m.stride(0);
|
|
||||||
p.mask_h_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1);
|
|
||||||
p.mask_q_stride = (m.size(2) == 1) ? 0 : (int)m.stride(2);
|
|
||||||
} else {
|
|
||||||
TORCH_CHECK(false, "mask must be 2D, 3D, or 4D");
|
|
||||||
}
|
|
||||||
p.mask = m.data_ptr<bool>();
|
|
||||||
} else {
|
|
||||||
p.mask = nullptr;
|
|
||||||
p.mask_b_stride = 0;
|
|
||||||
p.mask_h_stride = 0;
|
|
||||||
p.mask_q_stride = 0;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ---- attn_pack_params (contiguous KV) ----
|
|
||||||
template<typename T>
|
|
||||||
inline void attn_pack_params(
|
|
||||||
torch::Tensor q,
|
|
||||||
torch::Tensor k,
|
|
||||||
torch::Tensor v,
|
|
||||||
c10::optional<torch::Tensor> mask,
|
|
||||||
int64_t causal_offset,
|
|
||||||
double scale,
|
|
||||||
int64_t layout,
|
|
||||||
AttentionParams<T>& p
|
|
||||||
) {
|
|
||||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
|
||||||
|
|
||||||
TORCH_CHECK(q.is_cuda() && k.is_cuda() && v.is_cuda());
|
|
||||||
TORCH_CHECK(q.dtype() == torch::kBFloat16);
|
|
||||||
TORCH_CHECK(k.dtype() == torch::kBFloat16);
|
|
||||||
TORCH_CHECK(v.dtype() == torch::kBFloat16);
|
|
||||||
TORCH_CHECK(k.sizes() == v.sizes(), "K and V must have identical shapes");
|
|
||||||
TORCH_CHECK(q.dim() == 4 && k.dim() == 4, "Q/K/V must be 4D");
|
|
||||||
|
|
||||||
extract_q_dims_and_strides(q, layout, p);
|
|
||||||
|
|
||||||
if (layout == 1) k = k.transpose(1, 2), v = v.transpose(1, 2);
|
|
||||||
|
|
||||||
p.kv_head = (int)k.size(1);
|
|
||||||
p.kv_len = (int)k.size(2);
|
|
||||||
TORCH_CHECK(k.size(3) == p.head_dim, "K/V head_dim must match Q");
|
|
||||||
|
|
||||||
p.kv_stride_b = (int)k.stride(0);
|
|
||||||
p.kv_stride_h = (int)k.stride(1);
|
|
||||||
p.kv_stride_l = (int)k.stride(2);
|
|
||||||
p.kv_stride_d = (int)k.stride(3);
|
|
||||||
|
|
||||||
p.causal_offset = (int)causal_offset;
|
|
||||||
p.use_mask = mask.has_value() ? 1 : 0;
|
|
||||||
p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
|
|
||||||
|
|
||||||
p.q = (const T*)q.data_ptr();
|
|
||||||
p.k = (const T*)k.data_ptr();
|
|
||||||
p.v = (const T*)v.data_ptr();
|
|
||||||
p.o = nullptr;
|
|
||||||
p.o_part = nullptr;
|
|
||||||
p.ml_part = nullptr;
|
|
||||||
|
|
||||||
pack_mask(mask, p);
|
|
||||||
}
|
|
||||||
|
|
||||||
// ---- attn_pack_paged_params ----
|
|
||||||
template<typename T>
|
|
||||||
inline void attn_pack_paged_params(
|
|
||||||
torch::Tensor q,
|
|
||||||
torch::Tensor page_table,
|
|
||||||
torch::Tensor k_cache,
|
|
||||||
torch::Tensor v_cache,
|
|
||||||
int64_t page_size,
|
|
||||||
int64_t kv_len,
|
|
||||||
c10::optional<torch::Tensor> mask,
|
|
||||||
int64_t causal_offset,
|
|
||||||
double scale,
|
|
||||||
int64_t layout,
|
|
||||||
PagedAttentionParams<T>& p
|
|
||||||
) {
|
|
||||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
|
||||||
|
|
||||||
TORCH_CHECK(q.is_cuda() && page_table.is_cuda() && k_cache.is_cuda() && v_cache.is_cuda());
|
|
||||||
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16");
|
|
||||||
TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bf16");
|
|
||||||
TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bf16");
|
|
||||||
TORCH_CHECK(page_table.dtype() == torch::kLong, "page_table must be int64");
|
|
||||||
TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must have identical shapes");
|
|
||||||
|
|
||||||
extract_q_dims_and_strides(q, layout, p);
|
|
||||||
|
|
||||||
p.kv_head = (int)k_cache.size(2);
|
|
||||||
p.kv_len = (int)kv_len;
|
|
||||||
p.page_size = (int)page_size;
|
|
||||||
p.max_pages = (int)page_table.size(1);
|
|
||||||
|
|
||||||
TORCH_CHECK(q.size(2) == 1, "Q seq_len must be 1 (decode)");
|
|
||||||
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
|
|
||||||
TORCH_CHECK(k_cache.size(1) == page_size,
|
|
||||||
"k_cache dim 1 must equal page_size, got ",
|
|
||||||
k_cache.size(1), " vs ", page_size);
|
|
||||||
|
|
||||||
p.causal_offset = (int)causal_offset;
|
|
||||||
p.use_mask = (mask.has_value() && mask.value().defined()) ? 1 : 0;
|
|
||||||
p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
|
|
||||||
|
|
||||||
p.page_table = page_table.data_ptr<int64_t>();
|
|
||||||
p.k_cache = (const T*)k_cache.data_ptr();
|
|
||||||
p.v_cache = (const T*)v_cache.data_ptr();
|
|
||||||
p.q = (const T*)q.data_ptr();
|
|
||||||
p.o = nullptr;
|
|
||||||
p.o_part = nullptr;
|
|
||||||
p.ml_part = nullptr;
|
|
||||||
|
|
||||||
pack_mask(mask, p);
|
|
||||||
}
|
|
||||||
@@ -1,42 +0,0 @@
|
|||||||
#include "attn_dispatchers.cuh"
|
|
||||||
#include "attn_entry_utils.cuh"
|
|
||||||
|
|
||||||
torch::Tensor attn_paged_decode(
|
|
||||||
torch::Tensor q,
|
|
||||||
torch::Tensor page_table,
|
|
||||||
torch::Tensor k_cache,
|
|
||||||
torch::Tensor v_cache,
|
|
||||||
int64_t page_size,
|
|
||||||
int64_t kv_len,
|
|
||||||
c10::optional<torch::Tensor> mask,
|
|
||||||
int64_t causal_offset,
|
|
||||||
double scale,
|
|
||||||
int64_t layout
|
|
||||||
) {
|
|
||||||
PagedAttentionParams<bf16> p;
|
|
||||||
attn_pack_paged_params(q, page_table, k_cache, v_cache,
|
|
||||||
page_size, kv_len, mask, causal_offset, scale, layout, p);
|
|
||||||
|
|
||||||
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
|
|
||||||
auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
|
|
||||||
p.o = (bf16*)O_view.data_ptr();
|
|
||||||
|
|
||||||
alloc_split_partials(p);
|
|
||||||
DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_decode, p);
|
|
||||||
return O;
|
|
||||||
}
|
|
||||||
|
|
||||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
|
||||||
m.def("attn_paged_decode", &attn_paged_decode,
|
|
||||||
py::arg("q"),
|
|
||||||
py::arg("page_table"),
|
|
||||||
py::arg("k_cache"),
|
|
||||||
py::arg("v_cache"),
|
|
||||||
py::arg("page_size"),
|
|
||||||
py::arg("kv_len"),
|
|
||||||
py::arg("mask") = py::none(),
|
|
||||||
py::arg("causal_offset") = -1,
|
|
||||||
py::arg("scale") = 0.0,
|
|
||||||
py::arg("layout") = 0,
|
|
||||||
"Paged GQA decode — split-KV with direct page-table access.");
|
|
||||||
}
|
|
||||||
@@ -1,153 +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 + q_head * p.mask_h_stride;
|
|
||||||
|
|
||||||
for (int ci = ch_begin; ci < ch_end; ci++) {
|
|
||||||
int chunk_start = ci * PDC_CHUNK;
|
|
||||||
int this_chunk = min(PDC_CHUNK, p.kv_len - chunk_start);
|
|
||||||
|
|
||||||
int total = this_chunk * p.head_dim;
|
|
||||||
for (int i = threadIdx.y * 32 + lane; i < total;
|
|
||||||
i += blockDim.x * blockDim.y) {
|
|
||||||
int s = i / p.head_dim;
|
|
||||||
int d_dim = i % p.head_dim;
|
|
||||||
int pos = chunk_start + s;
|
|
||||||
int logical_page = pos / p.page_size;
|
|
||||||
int page_offset = pos % p.page_size;
|
|
||||||
int phys_page = p.page_table[batch * p.max_pages + logical_page];
|
|
||||||
if (phys_page >= 0) {
|
|
||||||
int64_t off = (int64_t)phys_page * p.page_size * p.kv_head * p.head_dim
|
|
||||||
+ (int64_t)page_offset * p.kv_head * p.head_dim
|
|
||||||
+ (int64_t)kv_head * p.head_dim
|
|
||||||
+ d_dim;
|
|
||||||
k_smem[i] = p.k_cache[off];
|
|
||||||
} else {
|
|
||||||
k_smem[i] = __float2bfloat16(0.0f);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
__syncthreads();
|
|
||||||
|
|
||||||
for (int s = 0; s < this_chunk; s++) {
|
|
||||||
float partial = 0.0f;
|
|
||||||
#pragma unroll
|
|
||||||
for (int i = 0; i < hd_per_thread; i++)
|
|
||||||
partial += q_reg[i] * __bfloat162float(
|
|
||||||
k_smem[s * p.head_dim + lane * hd_per_thread + i]);
|
|
||||||
partial = warp_reduce_sum(partial) * p.scale;
|
|
||||||
|
|
||||||
int kv_idx = chunk_start + s;
|
|
||||||
bool masked = false;
|
|
||||||
if constexpr (HasMask) {
|
|
||||||
if (!p.mask[mask_base + kv_idx])
|
|
||||||
masked = true;
|
|
||||||
}
|
|
||||||
if constexpr (IsCausal) {
|
|
||||||
if (kv_idx > p.causal_offset)
|
|
||||||
masked = true;
|
|
||||||
}
|
|
||||||
if (masked)
|
|
||||||
partial = -FLT_MAX;
|
|
||||||
|
|
||||||
float new_m = fmaxf(m, partial);
|
|
||||||
float alpha = __expf(m - new_m);
|
|
||||||
float beta = __expf(partial - new_m);
|
|
||||||
d = d * alpha + beta;
|
|
||||||
|
|
||||||
int pos = chunk_start + s;
|
|
||||||
int logical_page = pos / p.page_size;
|
|
||||||
int page_offset = pos % p.page_size;
|
|
||||||
int phys_page = p.page_table[batch * p.max_pages + logical_page];
|
|
||||||
if (masked) {
|
|
||||||
#pragma unroll
|
|
||||||
for (int i = 0; i < hd_per_thread; i++)
|
|
||||||
acc_reg[i] = fmaf(acc_reg[i], alpha, 0.0f);
|
|
||||||
} else if (phys_page >= 0) {
|
|
||||||
int64_t v_base = (int64_t)phys_page * p.page_size * p.kv_head * p.head_dim
|
|
||||||
+ (int64_t)page_offset * p.kv_head * p.head_dim
|
|
||||||
+ (int64_t)kv_head * p.head_dim;
|
|
||||||
#pragma unroll
|
|
||||||
for (int i = 0; i < hd_per_thread; i++)
|
|
||||||
acc_reg[i] = fmaf(acc_reg[i], alpha,
|
|
||||||
__bfloat162float(p.v_cache[v_base + lane * hd_per_thread + i]) * beta);
|
|
||||||
} else {
|
|
||||||
#pragma unroll
|
|
||||||
for (int i = 0; i < hd_per_thread; i++)
|
|
||||||
acc_reg[i] = fmaf(acc_reg[i], alpha, 0.0f);
|
|
||||||
}
|
|
||||||
m = new_m;
|
|
||||||
}
|
|
||||||
__syncthreads();
|
|
||||||
}
|
|
||||||
|
|
||||||
size_t bh = (size_t)batch * p.q_head + q_head;
|
|
||||||
size_t slot = bh * MAX_SPLITS + split;
|
|
||||||
int d0 = lane * hd_per_thread;
|
|
||||||
#pragma unroll
|
|
||||||
for (int i = 0; i < hd_per_thread; i++)
|
|
||||||
p.o_part[slot * p.head_dim + (d0 + i)] = acc_reg[i];
|
|
||||||
if (lane == 0) {
|
|
||||||
p.ml_part[slot * 2] = m;
|
|
||||||
p.ml_part[slot * 2 + 1] = d;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
__global__ void paged_attn_decode_combine_kernel(PagedAttentionParams<bf16> p) {
|
|
||||||
int bh = blockIdx.x;
|
|
||||||
int d = threadIdx.x;
|
|
||||||
if (d >= p.head_dim) return;
|
|
||||||
|
|
||||||
int batch = bh / p.q_head;
|
|
||||||
int q_head = bh % p.q_head;
|
|
||||||
|
|
||||||
size_t split_base = (size_t)bh * MAX_SPLITS;
|
|
||||||
const float* mlp = p.ml_part + split_base * 2;
|
|
||||||
const float* op = p.o_part + split_base * p.head_dim;
|
|
||||||
|
|
||||||
float m = -FLT_MAX, l = 0.0f, acc = 0.0f;
|
|
||||||
for (int s = 0; s < p.num_splits; s++) {
|
|
||||||
float mi = mlp[s * 2];
|
|
||||||
if (mi <= -FLT_MAX) continue;
|
|
||||||
float li = mlp[s * 2 + 1];
|
|
||||||
float nm = fmaxf(m, mi);
|
|
||||||
float corr = __expf(m - nm);
|
|
||||||
float e = __expf(mi - nm);
|
|
||||||
acc = fmaf(acc, corr, op[s * p.head_dim + d] * e);
|
|
||||||
l = fmaf(l, corr, li * e);
|
|
||||||
m = nm;
|
|
||||||
}
|
|
||||||
|
|
||||||
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
|
|
||||||
int o_off = batch * p.q_stride_b + q_head * p.q_stride_h + d * p.q_stride_d;
|
|
||||||
p.o[o_off] = __float2bfloat16(acc * inv);
|
|
||||||
}
|
|
||||||
@@ -1,182 +0,0 @@
|
|||||||
#pragma once
|
|
||||||
#include <cfloat>
|
|
||||||
#include <cuda_bf16.h>
|
|
||||||
#include "attn_common.h"
|
|
||||||
#include "attn_mma_utils.cuh"
|
|
||||||
#include "attn_warp_utils.cuh"
|
|
||||||
|
|
||||||
// Paged split-KV tensor-core decode via GQA head-packing.
|
|
||||||
// Reads K/V directly from the page pool through a page table — one tile
|
|
||||||
// (BC=32) fits within a single page (page_size >= 32), so the page-table
|
|
||||||
// lookup happens once per tile for cp.async.
|
|
||||||
//
|
|
||||||
// IsCausal and HasMask are compile-time bools.
|
|
||||||
template <typename Traits, bool IsCausal, bool HasMask>
|
|
||||||
__global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16> p) {
|
|
||||||
const int lane = threadIdx.x;
|
|
||||||
const int gid = lane >> 2;
|
|
||||||
const int tid4 = lane & 3;
|
|
||||||
|
|
||||||
const int pass = blockIdx.x / p.kv_head;
|
|
||||||
const int kv_head = blockIdx.x % p.kv_head;
|
|
||||||
const int batch = blockIdx.y;
|
|
||||||
const int split = blockIdx.z;
|
|
||||||
|
|
||||||
constexpr int MAX_G = 16;
|
|
||||||
const int G_total = p.q_head / p.kv_head;
|
|
||||||
const int g_begin = pass * MAX_G;
|
|
||||||
const int G = min(MAX_G, G_total - g_begin);
|
|
||||||
const int q_head0 = kv_head * G_total + g_begin;
|
|
||||||
|
|
||||||
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
|
|
||||||
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
|
|
||||||
|
|
||||||
#pragma unroll
|
|
||||||
for (int i = lane; i < Traits::STAGES * Traits::BC * Traits::LD; i += 32) {
|
|
||||||
sK[i] = __float2bfloat16(0.0f);
|
|
||||||
sV[i] = __float2bfloat16(0.0f);
|
|
||||||
}
|
|
||||||
__syncwarp();
|
|
||||||
|
|
||||||
const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h;
|
|
||||||
const int qra = gid;
|
|
||||||
const int qrb = gid + 8;
|
|
||||||
const bool va = qra < G, vb = qrb < G;
|
|
||||||
unsigned Qa[Traits::KD][4];
|
|
||||||
load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
|
|
||||||
qra, qrb, va, vb, tid4, Qa);
|
|
||||||
|
|
||||||
float Oacc[Traits::DN8][4];
|
|
||||||
#pragma unroll
|
|
||||||
for (int j = 0; j < Traits::DN8; j++)
|
|
||||||
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
|
||||||
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
|
||||||
|
|
||||||
const int tiles_total = (p.kv_len + Traits::BC - 1) / Traits::BC;
|
|
||||||
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
|
|
||||||
const int ti_begin = split * tiles_per_split;
|
|
||||||
const int ti_end = min(tiles_total, ti_begin + tiles_per_split);
|
|
||||||
|
|
||||||
const int64_t page_stride = (int64_t)p.page_size * p.kv_head * Traits::HEAD_DIM;
|
|
||||||
const int64_t pos_stride = (int64_t)p.kv_head * Traits::HEAD_DIM;
|
|
||||||
const int64_t head_off = (int64_t)kv_head * Traits::HEAD_DIM;
|
|
||||||
|
|
||||||
// ---- Load tile lambda: paged addressing ----
|
|
||||||
// Unified per-element page-table lookup. When page_size >= BC, all
|
|
||||||
// elements in a tile share the same page, so the lookup is redundant
|
|
||||||
// but harmless (L1-cached). This avoids a branch on page_size.
|
|
||||||
auto load_tile = [&](int ti, int buf) {
|
|
||||||
int kv0 = ti * Traits::BC;
|
|
||||||
bf16* dK = sK + buf * Traits::BC * Traits::LD;
|
|
||||||
bf16* dV = sV + buf * Traits::BC * Traits::LD;
|
|
||||||
#pragma unroll
|
|
||||||
for (int i = lane * Traits::VEC; i < Traits::TOTAL;
|
|
||||||
i += Traits::NUM_THREADS * Traits::VEC) {
|
|
||||||
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
|
|
||||||
int kc = kv0 + r;
|
|
||||||
bool valid = (kc < p.kv_len);
|
|
||||||
if constexpr (HasMask) {
|
|
||||||
valid = valid && p.mask[batch * p.mask_b_stride + kc];
|
|
||||||
}
|
|
||||||
int phys_page = valid ? p.page_table[batch * p.max_pages + kc] : 0;
|
|
||||||
valid = valid && (phys_page >= 0);
|
|
||||||
int page_off = kc % p.page_size;
|
|
||||||
int64_t gmem_base = (int64_t)phys_page * page_stride
|
|
||||||
+ (int64_t)page_off * pos_stride
|
|
||||||
+ head_off;
|
|
||||||
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
|
|
||||||
cp_async_16_pred(&dK[off], &p.k_cache[gmem_base + d], valid);
|
|
||||||
cp_async_16_pred(&dV[off], &p.v_cache[gmem_base + d], valid);
|
|
||||||
}
|
|
||||||
cp_async_commit();
|
|
||||||
};
|
|
||||||
|
|
||||||
// ---- Multi-stage cp.async pipeline ----
|
|
||||||
// Prologue loads STAGES tiles; each loop iteration waits only for the
|
|
||||||
// oldest outstanding group (wait_group<STAGES-1>) so the STAGES-1 newer
|
|
||||||
// tile loads stay in flight and overlap with the current tile's compute.
|
|
||||||
constexpr int STAGES = Traits::STAGES;
|
|
||||||
const int ntiles = ti_end - ti_begin;
|
|
||||||
|
|
||||||
auto process_tile = [&](int it, int buf) {
|
|
||||||
const bf16* bK = sK + buf * Traits::BC * Traits::LD;
|
|
||||||
const bf16* bV = sV + buf * Traits::BC * Traits::LD;
|
|
||||||
int kv0 = (ti_begin + it) * Traits::BC;
|
|
||||||
|
|
||||||
float Sacc[Traits::NC8][4];
|
|
||||||
mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
|
|
||||||
|
|
||||||
#pragma unroll
|
|
||||||
for (int n8 = 0; n8 < Traits::NC8; n8++)
|
|
||||||
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
|
||||||
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
|
||||||
|
|
||||||
int maxc = IsCausal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
|
|
||||||
mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc,
|
|
||||||
0, 0,
|
|
||||||
p.mask_b_stride, 0, 0,
|
|
||||||
batch, 0,
|
|
||||||
p.mask,
|
|
||||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
|
||||||
|
|
||||||
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
|
|
||||||
};
|
|
||||||
|
|
||||||
if (ntiles >= STAGES) {
|
|
||||||
#pragma unroll
|
|
||||||
for (int i = 0; i < STAGES; i++)
|
|
||||||
load_tile(ti_begin + i, i);
|
|
||||||
|
|
||||||
for (int it = 0; it < ntiles; it++) {
|
|
||||||
cp_async_wait_group<STAGES - 1>();
|
|
||||||
__syncwarp();
|
|
||||||
process_tile(it, it & (STAGES - 1));
|
|
||||||
__syncwarp();
|
|
||||||
if (it + STAGES < ntiles)
|
|
||||||
load_tile(ti_begin + it + STAGES, (it + STAGES) & (STAGES - 1));
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
// Fewer tiles than stages: load all, wait for all, process.
|
|
||||||
for (int i = 0; i < ntiles; i++)
|
|
||||||
load_tile(ti_begin + i, i);
|
|
||||||
cp_async_wait_group<0>();
|
|
||||||
__syncwarp();
|
|
||||||
for (int it = 0; it < ntiles; it++)
|
|
||||||
process_tile(it, it);
|
|
||||||
}
|
|
||||||
|
|
||||||
auto split_slot = [&](int h) -> size_t {
|
|
||||||
size_t bh = (size_t)batch * p.q_head + h;
|
|
||||||
return bh * MAX_SPLITS + split;
|
|
||||||
};
|
|
||||||
#pragma unroll
|
|
||||||
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
|
|
||||||
int d = dn8 * 8 + 2 * tid4;
|
|
||||||
int r0 = gid, r1 = gid + 8;
|
|
||||||
if (r0 < G) {
|
|
||||||
int h = q_head0 + r0;
|
|
||||||
float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
|
|
||||||
op[d] = Oacc[dn8][0];
|
|
||||||
op[d + 1] = Oacc[dn8][1];
|
|
||||||
}
|
|
||||||
if (r1 < G) {
|
|
||||||
int h = q_head0 + r1;
|
|
||||||
float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
|
|
||||||
op[d] = Oacc[dn8][2];
|
|
||||||
op[d + 1] = Oacc[dn8][3];
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if (tid4 == 0) {
|
|
||||||
int r0 = gid, r1 = gid + 8;
|
|
||||||
if (r0 < G) {
|
|
||||||
int h = q_head0 + r0;
|
|
||||||
float* mp = p.ml_part + split_slot(h) * 2;
|
|
||||||
mp[0] = m0; mp[1] = l0;
|
|
||||||
}
|
|
||||||
if (r1 < G) {
|
|
||||||
int h = q_head0 + r1;
|
|
||||||
float* mp = p.ml_part + split_slot(h) * 2;
|
|
||||||
mp[0] = m1; mp[1] = l1;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,13 +0,0 @@
|
|||||||
#pragma once
|
|
||||||
#include <cuda_bf16.h>
|
|
||||||
|
|
||||||
using bf16 = __nv_bfloat16;
|
|
||||||
|
|
||||||
static constexpr int MAX_SPLITS = 32;
|
|
||||||
|
|
||||||
__device__ inline float warp_reduce_sum(float val) {
|
|
||||||
#pragma unroll
|
|
||||||
for (int offset = 16; offset > 0; offset >>= 1)
|
|
||||||
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
|
|
||||||
return val;
|
|
||||||
}
|
|
||||||
@@ -0,0 +1,82 @@
|
|||||||
|
// Shared cp.async primitives — pure CUDA, no torch.
|
||||||
|
//
|
||||||
|
// One header for the async-copy pipeline used by both the attention kernels
|
||||||
|
// (predicated 16-byte K/V tile staging) and the fp8 GEMM (predicated operand
|
||||||
|
// staging + the fixed-depth wait_group). The emitter is split from its
|
||||||
|
// policies: cp_async_16_raw owns the single PTX site, and each wrapper states
|
||||||
|
// one destination contract (generic pointer vs loop-carried shared offset)
|
||||||
|
// and one predication contract (unconditional vs zero-fill-when-false), so
|
||||||
|
// call sites never pass a dead `true` predicate or re-convert a carried
|
||||||
|
// offset. PTX requires wait_group's operand to be an immediate, hence the
|
||||||
|
// template form below.
|
||||||
|
|
||||||
|
#pragma once
|
||||||
|
|
||||||
|
#include <cuda_runtime.h>
|
||||||
|
|
||||||
|
namespace astrai {
|
||||||
|
|
||||||
|
// Raw emitter: read src_size bytes (<= 16) from gmem into the shared
|
||||||
|
// offset. src_size = 0 reads nothing, so a predicated-off call zero-fills
|
||||||
|
// its destination without touching the (possibly out-of-range) source.
|
||||||
|
// BypassL1 selects .cg (L2 only, default) vs .ca (L1 + L2).
|
||||||
|
template <bool BypassL1 = true>
|
||||||
|
__device__ __forceinline__ void cp_async_16_raw(unsigned smem_addr,
|
||||||
|
const void* gmem_ptr,
|
||||||
|
int 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));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Unconditional 16-byte copy to a generic shared pointer.
|
||||||
|
// `T` is the smem element type; only the destination pointer's type matters.
|
||||||
|
template <typename T, bool BypassL1 = true>
|
||||||
|
__device__ __forceinline__ void cp_async_16(T* smem_ptr,
|
||||||
|
const void* gmem_ptr) {
|
||||||
|
cp_async_16_raw<BypassL1>(__cvta_generic_to_shared(smem_ptr), gmem_ptr,
|
||||||
|
16);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Predicated: full copy when `pred`, zero-fill otherwise.
|
||||||
|
template <typename T, bool BypassL1 = true>
|
||||||
|
__device__ __forceinline__ void cp_async_16(T* smem_ptr, const void* gmem_ptr,
|
||||||
|
bool pred) {
|
||||||
|
cp_async_16_raw<BypassL1>(__cvta_generic_to_shared(smem_ptr), gmem_ptr,
|
||||||
|
pred ? 16 : 0);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Predicated raw-offset form: the destination is an already-converted
|
||||||
|
// shared-memory offset (e.g. a loop-carried swizzled stage address), so
|
||||||
|
// steady-state prefetch sites issue one LDGSTS straight from the register.
|
||||||
|
template <bool BypassL1 = true>
|
||||||
|
__device__ __forceinline__ void cp_async_16(unsigned smem_addr,
|
||||||
|
const void* gmem_ptr, bool pred) {
|
||||||
|
cp_async_16_raw<BypassL1>(smem_addr, gmem_ptr, pred ? 16 : 0);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Commit all outstanding cp.async ops of this thread as one group.
|
||||||
|
__device__ __forceinline__ void cp_async_commit_group() {
|
||||||
|
asm volatile("cp.async.commit_group;");
|
||||||
|
}
|
||||||
|
|
||||||
|
// Wait for every committed group (pipeline drain).
|
||||||
|
__device__ __forceinline__ void cp_async_wait_all() {
|
||||||
|
asm volatile("cp.async.wait_all;");
|
||||||
|
}
|
||||||
|
|
||||||
|
// Wait until at most KeepGroups committed groups are still in flight.
|
||||||
|
// PTX requires an immediate operand; keep it as a template argument so the
|
||||||
|
// stage policy stays compile-time configurable.
|
||||||
|
template <int KeepGroups>
|
||||||
|
__device__ __forceinline__ void cp_async_wait_group() {
|
||||||
|
static_assert(KeepGroups >= 0 && KeepGroups <= 7,
|
||||||
|
"cp.async.wait_group supports immediates in [0, 7]");
|
||||||
|
asm volatile("cp.async.wait_group %0;" :: "n"(KeepGroups));
|
||||||
|
}
|
||||||
|
|
||||||
|
} // namespace astrai
|
||||||
@@ -0,0 +1,23 @@
|
|||||||
|
// Pure-CUDA device helpers shared across kernel families (no torch).
|
||||||
|
//
|
||||||
|
// Family-local headers under kernels/<family>/ own their POD params and
|
||||||
|
// strategy traits; anything cross-cutting (compute-capability checks, device
|
||||||
|
// constants) lives here.
|
||||||
|
|
||||||
|
#pragma once
|
||||||
|
|
||||||
|
namespace astrai {
|
||||||
|
|
||||||
|
// Compute-capability comparison: is the device at least (major, minor)?
|
||||||
|
inline bool sm_at_least(int device_major, int device_minor, int major,
|
||||||
|
int minor) {
|
||||||
|
return device_major > major ||
|
||||||
|
(device_major == major && device_minor >= minor);
|
||||||
|
}
|
||||||
|
|
||||||
|
// FP8 tensor-core MMA (`mma.sync.aligned.m16n8k32` with fp8 inputs) exists on
|
||||||
|
// Ada (sm_89) and Hopper (sm_90+); sm_80 has no fp8 instructions.
|
||||||
|
inline constexpr int kMinSmForFp8Major = 8;
|
||||||
|
inline constexpr int kMinSmForFp8Minor = 9;
|
||||||
|
|
||||||
|
} // namespace astrai
|
||||||
@@ -0,0 +1,165 @@
|
|||||||
|
// Shared mma.sync wrappers — pure CUDA, no torch.
|
||||||
|
//
|
||||||
|
// One template for every tensor-core MMA used by the kernel families. The
|
||||||
|
// instruction shape follows from the input element type:
|
||||||
|
// __nv_bfloat16 -> mma.sync.aligned.m16n8k16 (sm_80+), A = 4x b32, B = 2x b32
|
||||||
|
// __nv_fp8_e4m3/e5m2 -> mma.sync.aligned.m16n8k32 (sm_89+), A = 4x b32, B = 2x b32
|
||||||
|
// All variants accumulate into fp32: d = a*b + c, with the PTX mnemonic and
|
||||||
|
// the K dimension differing per type. `d` may alias `c` (in-place accumulate,
|
||||||
|
// as the FP8 GEMM does).
|
||||||
|
|
||||||
|
#pragma once
|
||||||
|
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
#include <cuda_fp8.h>
|
||||||
|
#include <cuda_runtime.h>
|
||||||
|
#include <type_traits>
|
||||||
|
|
||||||
|
|
||||||
|
#define DEVICE_FORCEINLINE static __device__ __forceinline__
|
||||||
|
|
||||||
|
namespace astrai {
|
||||||
|
|
||||||
|
// Compute capability of the current compilation pass: 0 in the host pass,
|
||||||
|
// the numeric CC (e.g. 890) in device passes where __CUDA_ARCH__ is defined.
|
||||||
|
// Defined() cannot appear in expressions, so this macro lets mma_sync use
|
||||||
|
// the arch in a static_assert instead of per-branch #if guards.
|
||||||
|
#ifndef __CUDA_ARCH__
|
||||||
|
#define ASTRAI_DEVICE_ARCH 0
|
||||||
|
#else
|
||||||
|
#define ASTRAI_DEVICE_ARCH __CUDA_ARCH__
|
||||||
|
#endif
|
||||||
|
|
||||||
|
// Compile-time shape of the MMA instruction for an input element type.
|
||||||
|
// `min_arch` is the numeric compute capability the instruction requires —
|
||||||
|
// the single place that encodes the hardware floor for each type.
|
||||||
|
template <typename InT>
|
||||||
|
struct mma_shape {
|
||||||
|
static constexpr int k = 16; // m16n8k16
|
||||||
|
static constexpr int a_regs = 4; // A fragment: 4x b32
|
||||||
|
static constexpr int b_regs = 2; // B fragment: 2x b32
|
||||||
|
static constexpr int min_arch = 800; // bf16 mma.sync, sm_80+
|
||||||
|
};
|
||||||
|
|
||||||
|
template <>
|
||||||
|
struct mma_shape<__nv_fp8_e4m3> {
|
||||||
|
static constexpr int k = 32; // m16n8k32
|
||||||
|
static constexpr int a_regs = 4;
|
||||||
|
static constexpr int b_regs = 2;
|
||||||
|
static constexpr int min_arch = 890; // fp8 mma.sync, sm_89+ (Ada/Hopper)
|
||||||
|
};
|
||||||
|
|
||||||
|
template <>
|
||||||
|
struct mma_shape<__nv_fp8_e5m2> {
|
||||||
|
static constexpr int k = 32;
|
||||||
|
static constexpr int a_regs = 4;
|
||||||
|
static constexpr int b_regs = 2;
|
||||||
|
static constexpr int min_arch = 890;
|
||||||
|
};
|
||||||
|
|
||||||
|
// d[4] = a[4] x b[2] + c[4], row-major A, col-major B, fp32 accumulator.
|
||||||
|
// The PTX mnemonic is selected from InT. Building for a compute capability
|
||||||
|
// below `mma_shape<InT>::min_arch` is a **compile error** — the instruction
|
||||||
|
// does not exist there, and a silent no-op would produce wrong results.
|
||||||
|
template <typename InT>
|
||||||
|
DEVICE_FORCEINLINE void mma_sync(float d[4], const unsigned a[4],
|
||||||
|
const unsigned b[2],
|
||||||
|
const float c[4]) {
|
||||||
|
static_assert(ASTRAI_DEVICE_ARCH == 0 ||
|
||||||
|
ASTRAI_DEVICE_ARCH >= mma_shape<InT>::min_arch,
|
||||||
|
"mma_sync: this MMA shape requires a newer compute "
|
||||||
|
"capability than the build target");
|
||||||
|
if constexpr (std::is_same_v<InT, __nv_bfloat16>) {
|
||||||
|
asm volatile(
|
||||||
|
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
|
||||||
|
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%10,%11,%12,%13};"
|
||||||
|
: "=f"(d[0]), "=f"(d[1]), "=f"(d[2]), "=f"(d[3])
|
||||||
|
: "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]),
|
||||||
|
"f"(c[0]), "f"(c[1]), "f"(c[2]), "f"(c[3]));
|
||||||
|
} else if constexpr (std::is_same_v<InT, __nv_fp8_e5m2>) {
|
||||||
|
asm volatile(
|
||||||
|
"mma.sync.aligned.m16n8k32.row.col.f32.e5m2.e5m2.f32 "
|
||||||
|
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%10,%11,%12,%13};"
|
||||||
|
: "=f"(d[0]), "=f"(d[1]), "=f"(d[2]), "=f"(d[3])
|
||||||
|
: "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]),
|
||||||
|
"f"(c[0]), "f"(c[1]), "f"(c[2]), "f"(c[3]));
|
||||||
|
} else {
|
||||||
|
asm volatile(
|
||||||
|
"mma.sync.aligned.m16n8k32.row.col.f32.e4m3.e4m3.f32 "
|
||||||
|
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%10,%11,%12,%13};"
|
||||||
|
: "=f"(d[0]), "=f"(d[1]), "=f"(d[2]), "=f"(d[3])
|
||||||
|
: "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]),
|
||||||
|
"f"(c[0]), "f"(c[1]), "f"(c[2]), "f"(c[3]));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#undef ASTRAI_DEVICE_ARCH
|
||||||
|
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
// ldmatrix — cooperatively load 8x8 b16 matrices from smem into registers.
|
||||||
|
//
|
||||||
|
// The instruction is identical for every 16-bit-storage element type: bf16
|
||||||
|
// maps 1:1 onto b16 slots; fp8 is stored packed two-per-slot (see
|
||||||
|
// fp8/gemm.cuh), so one b16 slot holds two fp8 values. `T` is the element
|
||||||
|
// type and only serves as a semantic tag.
|
||||||
|
//
|
||||||
|
// x2 (single address): matrix0 = p (8 rows), matrix1 = p + 8*16 bytes
|
||||||
|
// x4: four matrices at p, +128, +256, +384 bytes
|
||||||
|
// Trans: transpose variant (V fragments of attention)
|
||||||
|
//
|
||||||
|
// ldmatrix takes a *single* smem address per thread, but the addresses of
|
||||||
|
// the 32 lanes are *not* all the same: lane i supplies the start address of
|
||||||
|
// matrix-row i (modulo 8) for matrix (i/8) — lanes 0-7 feed matrix 0's rows,
|
||||||
|
// lanes 8-15 matrix 1's rows (x2/x4), lanes 16-23 / 24-31 matrix 2 / 3's rows
|
||||||
|
// (x4 only; their addresses are ignored by x2). Each matrix is 8 rows x 16
|
||||||
|
// bytes, and consecutive matrices of one instruction are contiguous at
|
||||||
|
// 128-byte strides. fp8 fragment layouts in fp8/gemm.cuh are arranged around
|
||||||
|
// this constraint.
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
template <typename T, bool Trans = false>
|
||||||
|
DEVICE_FORCEINLINE void ldmatrix_x2(unsigned r[2], const T* p) {
|
||||||
|
const unsigned a = __cvta_generic_to_shared(p);
|
||||||
|
if constexpr (Trans) {
|
||||||
|
asm volatile(
|
||||||
|
"ldmatrix.sync.aligned.m8n8.x2.trans.shared.b16 {%0,%1}, [%2];"
|
||||||
|
: "=r"(r[0]), "=r"(r[1])
|
||||||
|
: "r"(a));
|
||||||
|
} else {
|
||||||
|
asm volatile(
|
||||||
|
"ldmatrix.sync.aligned.m8n8.x2.shared.b16 {%0,%1}, [%2];"
|
||||||
|
: "=r"(r[0]), "=r"(r[1])
|
||||||
|
: "r"(a));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Four matrices at p, p+128, p+256, p+384 bytes (16-byte row stride).
|
||||||
|
template <typename T>
|
||||||
|
DEVICE_FORCEINLINE void ldmatrix_x4(unsigned r[4], const T* p) {
|
||||||
|
const unsigned a = __cvta_generic_to_shared(p);
|
||||||
|
asm volatile(
|
||||||
|
"ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];"
|
||||||
|
: "=r"(r[0]), "=r"(r[1]), "=r"(r[2]), "=r"(r[3])
|
||||||
|
: "r"(a));
|
||||||
|
}
|
||||||
|
|
||||||
|
// Per-lane-address variants: the caller supplies a raw shared-memory address
|
||||||
|
// per lane instead of one common pointer. Use when the fragment tiles are
|
||||||
|
// XOR-swizzled per 16B chunk so each lane must compute its own row and chunk
|
||||||
|
// address (see fp8/gemm.cuh's frag_addr + lane selectors for the m16n8k32
|
||||||
|
// operand layouts).
|
||||||
|
DEVICE_FORCEINLINE void ldmatrix_x2_lane(unsigned r[2],
|
||||||
|
unsigned addr) {
|
||||||
|
asm volatile("ldmatrix.sync.aligned.m8n8.x2.shared.b16 {%0,%1}, [%2];"
|
||||||
|
: "=r"(r[0]), "=r"(r[1])
|
||||||
|
: "r"(addr));
|
||||||
|
}
|
||||||
|
|
||||||
|
DEVICE_FORCEINLINE void ldmatrix_x4_lane(unsigned r[4],
|
||||||
|
unsigned addr) {
|
||||||
|
asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];"
|
||||||
|
: "=r"(r[0]), "=r"(r[1]), "=r"(r[2]), "=r"(r[3])
|
||||||
|
: "r"(addr));
|
||||||
|
}
|
||||||
|
|
||||||
|
} // namespace astrai
|
||||||
@@ -0,0 +1,48 @@
|
|||||||
|
// Shared warp/block reduction + atomic helpers — pure CUDA, no torch.
|
||||||
|
//
|
||||||
|
// Extracted from the attention and fp8 families so both share one
|
||||||
|
// implementation: warp_reduce_sum (decode scalar kernel), warp_reduce_max +
|
||||||
|
// atomic_max_float (fp8 quantize amax), group_reduce_sum<G> (prefill scalar
|
||||||
|
// kernel).
|
||||||
|
|
||||||
|
#pragma once
|
||||||
|
|
||||||
|
namespace astrai {
|
||||||
|
|
||||||
|
// Full-warp butterfly sum reduction (32 lanes).
|
||||||
|
__device__ __forceinline__ float warp_reduce_sum(float val) {
|
||||||
|
#pragma unroll
|
||||||
|
for (int offset = 16; offset > 0; offset >>= 1)
|
||||||
|
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
|
||||||
|
return val;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Full-warp butterfly max reduction (32 lanes).
|
||||||
|
__device__ __forceinline__ float warp_reduce_max(float value) {
|
||||||
|
#pragma unroll
|
||||||
|
for (int offset = 16; offset > 0; offset >>= 1)
|
||||||
|
value = fmaxf(value, __shfl_xor_sync(0xffffffffu, value, offset));
|
||||||
|
return value;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Sub-warp group reduction over G consecutive lanes (G a power of two).
|
||||||
|
// `mask` is the full participating-lane mask of the group (see the
|
||||||
|
// prefill scalar kernel's gmask computation).
|
||||||
|
template <int G>
|
||||||
|
__device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
|
||||||
|
#pragma unroll
|
||||||
|
for (int o = G / 2; o > 0; o >>= 1)
|
||||||
|
v += __shfl_xor_sync(mask, v, o);
|
||||||
|
return v;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Unsigned-bit-pattern atomicMax for non-negative floats; a null
|
||||||
|
// destination disables the update (kernels with optional amax slots).
|
||||||
|
__device__ __forceinline__ void atomic_max_float(float* destination,
|
||||||
|
float value) {
|
||||||
|
if (destination)
|
||||||
|
atomicMax(reinterpret_cast<unsigned*>(destination),
|
||||||
|
__float_as_uint(value));
|
||||||
|
}
|
||||||
|
|
||||||
|
} // namespace astrai
|
||||||
@@ -0,0 +1,109 @@
|
|||||||
|
#pragma once
|
||||||
|
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
#include <cuda_fp8.h>
|
||||||
|
#include <cuda_runtime.h>
|
||||||
|
#include <cstdint>
|
||||||
|
|
||||||
|
// Pure POD/traits header — no .cuh/CUDA-kernel includes; raw __nv_* type
|
||||||
|
// spellings only.
|
||||||
|
|
||||||
|
namespace astrai {
|
||||||
|
namespace fp8 {
|
||||||
|
|
||||||
|
// Compile-time FP8 format: E4M3 (forward, max 448) or E5M2 (gradients,
|
||||||
|
// max 57344).
|
||||||
|
enum class FP8Format : int {
|
||||||
|
E4M3 = 0,
|
||||||
|
E5M2 = 1,
|
||||||
|
};
|
||||||
|
|
||||||
|
// Operand storage tags (CUTLASS-style) relative to the canonical matrices
|
||||||
|
// A [M][K] / B [K][N]: A RowMajor = [M][K] (default), A ColMajor = [K][M],
|
||||||
|
// B RowMajor = [K][N], B ColMajor = [N][K] (the nn.Linear weight). Selection
|
||||||
|
// is by type at compile time (see gemm.cuh's stage loads).
|
||||||
|
struct RowMajor {};
|
||||||
|
struct ColMajor {};
|
||||||
|
|
||||||
|
// Compile-time tile configuration, mirroring KernelTraits in the attention
|
||||||
|
// kernels: CTA tile, warp tile (WarpM x WarpN — e.g. 64x32 on the 128x128
|
||||||
|
// CTA, 32x32 on the 64x64 small CTA) and cp.async pipeline depth.
|
||||||
|
template <FP8Format Fmt, int BlockM, int BlockN, int K, int Stages,
|
||||||
|
int WarpM = 64, int WarpN = 32>
|
||||||
|
struct Fp8GemmTraits {
|
||||||
|
static constexpr FP8Format kFormat = Fmt;
|
||||||
|
static constexpr int kBlockM = BlockM;
|
||||||
|
static constexpr int kBlockN = BlockN;
|
||||||
|
static constexpr int kK = K;
|
||||||
|
static constexpr int kStages = Stages;
|
||||||
|
static constexpr int kWarpM = WarpM;
|
||||||
|
static constexpr int kWarpN = WarpN;
|
||||||
|
static constexpr bool kIsE5M2 = (Fmt == FP8Format::E5M2);
|
||||||
|
static constexpr __nv_fp8_interpretation_t kNvFormat =
|
||||||
|
kIsE5M2 ? __NV_E5M2 : __NV_E4M3;
|
||||||
|
static constexpr float kFp8Max = kIsE5M2 ? 57344.0f : 448.0f;
|
||||||
|
|
||||||
|
// Derived geometry: warp tiles tile the CTA. The smem budget is
|
||||||
|
// layout-aware, so it lives in Fp8GemmSmem (gemm.cuh).
|
||||||
|
static constexpr int kWarpsM = BlockM / WarpM;
|
||||||
|
static constexpr int kWarpsN = BlockN / WarpN;
|
||||||
|
static constexpr int kCtaThreads = kWarpsM * kWarpsN * 32;
|
||||||
|
static_assert(kWarpsM * WarpM == BlockM && kWarpsN * WarpN == BlockN,
|
||||||
|
"warp tiles must exactly tile the CTA");
|
||||||
|
static_assert(WarpM % 16 == 0 && WarpN % 8 == 0,
|
||||||
|
"warp tile must be a multiple of the m16n8 MMA shape");
|
||||||
|
};
|
||||||
|
|
||||||
|
// Quantize-kernel parameter POD: float input -> FP8 with fused amax.
|
||||||
|
struct FP8QuantizeParams {
|
||||||
|
const void* __restrict__ input_ptr = nullptr;
|
||||||
|
void* __restrict__ output_ptr = nullptr;
|
||||||
|
void* __restrict__ output_transposed_ptr = nullptr; // [cols][rows]
|
||||||
|
// Output layout: 0 = row-major only, 1 = transposed only, 2 = both from
|
||||||
|
// a single read. Modes 1/2 produce K-contiguous operands so crosswise
|
||||||
|
// consumers (backward grad_x / grad_w) route through the NT fast path.
|
||||||
|
int out_layout = 0;
|
||||||
|
|
||||||
|
const float* __restrict__ scale = nullptr; // device multiplier
|
||||||
|
float* __restrict__ amax = nullptr; // raw-domain max out
|
||||||
|
|
||||||
|
// Element count (elementwise kernel); the tiled kernel views the same
|
||||||
|
// buffer as [rows][cols] row-major.
|
||||||
|
int total = 0;
|
||||||
|
int rows = 0;
|
||||||
|
int cols = 0;
|
||||||
|
};
|
||||||
|
|
||||||
|
// Unified GEMM parameter POD, mirroring AttentionParams: one struct flows
|
||||||
|
// through the kernels; each kernel touches only the fields it needs.
|
||||||
|
struct FP8Params {
|
||||||
|
// FP8 operands + output; scales are quantization steps (device
|
||||||
|
// scalars). Optional bf16 bias fuses into the epilogue (fp32 add before
|
||||||
|
// the single bf16 rounding); null disables.
|
||||||
|
const void* __restrict__ a_ptr = nullptr;
|
||||||
|
const void* __restrict__ b_ptr = nullptr;
|
||||||
|
const void* __restrict__ bias_ptr = nullptr;
|
||||||
|
void* __restrict__ out_ptr = nullptr;
|
||||||
|
|
||||||
|
const float* __restrict__ scale = nullptr;
|
||||||
|
// NN-swap mode (canonicalize_gemm): the kernel computes the transposed
|
||||||
|
// problem and the epilogue scatters D[row][col] to out[col * p.m + row]
|
||||||
|
// in the caller's [M][N] buffer. Zero in the plain orientation.
|
||||||
|
int out_transposed = 0;
|
||||||
|
int m, n, k; // int covers LLM shapes; kernels promote to int64
|
||||||
|
|
||||||
|
// Batched (bmm) geometry: grid.z steps these element strides (0
|
||||||
|
// broadcasts the operand across batches).
|
||||||
|
int batch = 1;
|
||||||
|
int64_t a_batch_stride = 0;
|
||||||
|
int64_t b_batch_stride = 0;
|
||||||
|
int64_t out_batch_stride = 0;
|
||||||
|
|
||||||
|
// Physical leading dims (row strides) of A and B; the binding packs
|
||||||
|
// them so the kernel reads each buffer naturally or transposed per the
|
||||||
|
// LayoutA/LayoutB tags.
|
||||||
|
int a_ld, b_ld;
|
||||||
|
};
|
||||||
|
|
||||||
|
} // namespace fp8
|
||||||
|
} // namespace astrai
|
||||||
@@ -0,0 +1,271 @@
|
|||||||
|
#pragma once
|
||||||
|
// FP8 GEMM umbrella: the kernel orchestrator and the host-side launch
|
||||||
|
// planning. Device layers live in gemm/ (policy / load / scheduler /
|
||||||
|
// mainloop / epilogue) — pure CUDA, no torch; launchers are plain functions
|
||||||
|
// shared by the torch binding and the C tests. Layout tags and the NN swap
|
||||||
|
// semantics are documented in common.h and the design notes
|
||||||
|
// (docs/developer/cuda_kernels.md).
|
||||||
|
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
#include <cuda_fp8.h>
|
||||||
|
#include <cuda_runtime.h>
|
||||||
|
#include <type_traits>
|
||||||
|
|
||||||
|
#include "../common/cp_async.cuh"
|
||||||
|
#include "common.h"
|
||||||
|
#include "gemm/epilogue.cuh"
|
||||||
|
#include "gemm/load.cuh"
|
||||||
|
#include "gemm/mainloop.cuh"
|
||||||
|
#include "gemm/policy.cuh"
|
||||||
|
#include "gemm/scheduler.cuh"
|
||||||
|
|
||||||
|
namespace astrai {
|
||||||
|
namespace fp8 {
|
||||||
|
|
||||||
|
template <typename Policy>
|
||||||
|
__global__ void __launch_bounds__(Policy::kCtaThreads, Policy::kMinCtas)
|
||||||
|
fp8_gemm_kernel(FP8Params p) {
|
||||||
|
using Traits = typename Policy::Traits;
|
||||||
|
using Mainloop = Fp8CollectiveMainloop<Policy>;
|
||||||
|
using Epilogue = Fp8CollectiveEpilogue<Policy>;
|
||||||
|
// Stages live in dynamic shared memory so deep pipelines (> 48KB
|
||||||
|
// static limit) opt in via cudaFuncSetAttribute in the launcher.
|
||||||
|
extern __shared__ __align__(16) char fp8_gemm_smem[];
|
||||||
|
|
||||||
|
// Batch slice (grid.z): broadcast operands carry a 0 stride, so the
|
||||||
|
// same pointer serves every batch.
|
||||||
|
using T8 = typename Mainloop::T8;
|
||||||
|
const T8* a = reinterpret_cast<const T8*>(p.a_ptr) +
|
||||||
|
(int64_t)blockIdx.z * p.a_batch_stride;
|
||||||
|
const T8* b = reinterpret_cast<const T8*>(p.b_ptr) +
|
||||||
|
(int64_t)blockIdx.z * p.b_batch_stride;
|
||||||
|
auto* out_bf16 = reinterpret_cast<__nv_bfloat16*>(p.out_ptr) +
|
||||||
|
(int64_t)blockIdx.z * p.out_batch_stride;
|
||||||
|
|
||||||
|
static_assert(Mainloop::kBlockM * Mainloop::kBlockN * 2 <=
|
||||||
|
Mainloop::kARing * Mainloop::kBlockM * Mainloop::kK +
|
||||||
|
Mainloop::kBRing * Mainloop::kBlockN * Mainloop::kK,
|
||||||
|
"output tile must fit the reclaimed operand smem");
|
||||||
|
const int2 bn = Fp8GemmTileScheduler<Policy::kGroupRaster>::tile(blockIdx, gridDim);
|
||||||
|
Mainloop mainloop(fp8_gemm_smem, a, b, p.m, p.n, p.k, p.a_ld, p.b_ld,
|
||||||
|
threadIdx.x, bn);
|
||||||
|
float acc[Mainloop::kNt][Mainloop::kMt][4] = {}; // [nt][mt][acc]
|
||||||
|
mainloop.prologue();
|
||||||
|
mainloop.accumulate(acc);
|
||||||
|
// Drain the pipeline before the epilogue reclaims the operand rings.
|
||||||
|
astrai::cp_async_wait_all();
|
||||||
|
Epilogue(fp8_gemm_smem, p, bn.x, bn.y, threadIdx.x).run(acc, out_bf16);
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
// Launchers — pure CUDA (no torch), usable from the binding and pure C tests.
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
// SM count of the current device (cached per device; benign init race —
|
||||||
|
// every writer stores the same value).
|
||||||
|
inline int device_sm_count() {
|
||||||
|
static int cached[64] = {};
|
||||||
|
int dev = 0;
|
||||||
|
cudaGetDevice(&dev);
|
||||||
|
const bool cacheable = dev >= 0 && dev < 64;
|
||||||
|
int sms = cacheable ? cached[dev] : 0;
|
||||||
|
if (!sms) {
|
||||||
|
cudaDeviceGetAttribute(&sms, cudaDevAttrMultiProcessorCount, dev);
|
||||||
|
sms = sms > 0 ? sms : 1;
|
||||||
|
if (cacheable) cached[dev] = sms;
|
||||||
|
}
|
||||||
|
return sms;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Launch one kernel instantiation with its shared-memory budget: budgets
|
||||||
|
// beyond the 48KB static limit opt in once per instantiation via
|
||||||
|
// cudaFuncSetAttribute. Templated on the kernel *value* (auto NTTP) so
|
||||||
|
// every instantiation owns its own armed flag — same-signature kernels
|
||||||
|
// must not share it. A failed opt-in arms nothing, so the launch below
|
||||||
|
// fails loudly through the caller's error checks.
|
||||||
|
template <auto Kernel, typename... Args>
|
||||||
|
void launch_with_smem(int smem_bytes, dim3 grid, dim3 block,
|
||||||
|
cudaStream_t stream, Args... args) {
|
||||||
|
if (smem_bytes > 48 * 1024) {
|
||||||
|
static bool armed = false; // per instantiation
|
||||||
|
if (!armed) {
|
||||||
|
const cudaError_t err = cudaFuncSetAttribute(
|
||||||
|
Kernel, cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||||
|
smem_bytes);
|
||||||
|
armed = (err == cudaSuccess);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Kernel<<<grid, block, smem_bytes, stream>>>(args...);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Padding-driven small-CTA rule: m or n <= 64 wastes half a 128-row CTA's
|
||||||
|
// MMA work, and a non-128-divisible shape drags its edge tiles through the
|
||||||
|
// predicated generic path — when 64 divides both dims, the 64x64 CTA tiles
|
||||||
|
// exactly and wins that band.
|
||||||
|
inline bool small_cta_padding(int64_t m, int64_t n) {
|
||||||
|
if (m <= 64 || n <= 64) return true;
|
||||||
|
const bool big_div = (m % 128 == 0) && (n % 128 == 0);
|
||||||
|
const bool small_div = (m % 64 == 0) && (n % 64 == 0);
|
||||||
|
return !big_div && small_div;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Launch configuration — a pure function of the problem (unit-testable
|
||||||
|
// without a GPU). Raster order is not a plan field: every canonical layout
|
||||||
|
// runs grouped raster; the plain-raster knob stays available through
|
||||||
|
// launch_plan's GroupRaster parameter for experiments.
|
||||||
|
struct Fp8GemmPlan {
|
||||||
|
enum class Cta { kSmall64, kNarrow128x64, kBig128 };
|
||||||
|
Cta cta;
|
||||||
|
bool small_s3; // kSmall64 only: cp.async pipeline depth (2 vs 3 stages)
|
||||||
|
};
|
||||||
|
|
||||||
|
// crosswise_ops counts the operands taking the direct crosswise load
|
||||||
|
// (A ColMajor / B RowMajor storage): 0 = dual-congruous NT, 1 = TN and the
|
||||||
|
// NN swap, 2 = TT. The layout shifts the crossovers (measured tables in
|
||||||
|
// the design notes): the small CTA hides the crosswise LDG+PRMT latency
|
||||||
|
// far better, while the big CTA's operand reuse buys back load bandwidth
|
||||||
|
// the crosswise path does not traffic in.
|
||||||
|
inline Fp8GemmPlan plan_gemm(const FP8Params& p, int crosswise_ops = 0) {
|
||||||
|
const int64_t sm = device_sm_count();
|
||||||
|
const int64_t tiles_128 =
|
||||||
|
(int64_t)p.batch * ((p.m + 127) / 128) * ((p.n + 127) / 128);
|
||||||
|
const auto small = [&](bool s3) {
|
||||||
|
return Fp8GemmPlan{Fp8GemmPlan::Cta::kSmall64, s3};
|
||||||
|
};
|
||||||
|
const auto big = [] {
|
||||||
|
return Fp8GemmPlan{Fp8GemmPlan::Cta::kBig128, false};
|
||||||
|
};
|
||||||
|
const auto narrow = [] {
|
||||||
|
return Fp8GemmPlan{Fp8GemmPlan::Cta::kNarrow128x64, false};
|
||||||
|
};
|
||||||
|
// Padding rules first: predication waste beats any wave-fill effect.
|
||||||
|
if (small_cta_padding(p.m, p.n)) return small(crosswise_ops > 0);
|
||||||
|
if (crosswise_ops > 0) {
|
||||||
|
// Crosswise ladder (L20 measured): the small s3 CTA holds ~3/4 of
|
||||||
|
// the big CTA's per-SM throughput but tiles 4x finer, so it owns
|
||||||
|
// the whole sub-wave band and past it; the big CTA takes over once
|
||||||
|
// its grid fills ~1.5 waves.
|
||||||
|
if (tiles_128 >= sm * 3 / 2) return big();
|
||||||
|
return small(true);
|
||||||
|
}
|
||||||
|
if (tiles_128 >= sm) {
|
||||||
|
// Wave band: pick by the wave-quantization cost ceil(tiles/sm) *
|
||||||
|
// T_tile. The narrow tile carries half the big tile's MMA work at
|
||||||
|
// ~94% of its per-SM efficiency (T_narrow ~= 0.53 * T_big,
|
||||||
|
// integer-scaled by 100 below) — reproduces every measured
|
||||||
|
// crossover.
|
||||||
|
const int64_t tiles_narrow =
|
||||||
|
(int64_t)p.batch * ((p.m + 127) / 128) * ((p.n + 63) / 64);
|
||||||
|
const auto waves = [sm](int64_t tiles) { return (tiles + sm - 1) / sm; };
|
||||||
|
if (waves(tiles_narrow) * 53 < waves(tiles_128) * 100) return narrow();
|
||||||
|
return big();
|
||||||
|
}
|
||||||
|
// Sub-wave band: the narrow CTA fills the wave with N-tiles at full
|
||||||
|
// warp depth once its grid passes ~3/8 of a wave; below that the plain
|
||||||
|
// 64x64 CTA's extra parallelism wins, and past ~5/8 of a wave of
|
||||||
|
// 128x128 tiles the big CTA's operand reuse wins instead.
|
||||||
|
if (tiles_128 >= sm * 5 / 8) return big();
|
||||||
|
const int64_t tiles_narrow =
|
||||||
|
(int64_t)p.batch * ((p.m + 127) / 128) * ((p.n + 63) / 64);
|
||||||
|
if (tiles_narrow >= sm * 3 / 8) return narrow();
|
||||||
|
// Full-ring small CTAs: the 24KB s2 variant keeps 4 CTAs/SM while the
|
||||||
|
// whole grid stays resident; past that the 32KB s3 variant's deeper
|
||||||
|
// pipeline wins on multi-wave grids.
|
||||||
|
const int64_t tiles_64 =
|
||||||
|
(int64_t)p.batch * ((p.m + 63) / 64) * ((p.n + 63) / 64);
|
||||||
|
return small(tiles_64 > sm * 3);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Grid + launch for one concrete Policy — the only place a GEMM kernel goes
|
||||||
|
// to the wire.
|
||||||
|
template <typename Policy>
|
||||||
|
void launch_policy(const FP8Params& p, cudaStream_t stream) {
|
||||||
|
using Traits = typename Policy::Traits;
|
||||||
|
dim3 grid((p.n + Traits::kBlockN - 1) / Traits::kBlockN,
|
||||||
|
(p.m + Traits::kBlockM - 1) / Traits::kBlockM, p.batch);
|
||||||
|
launch_with_smem<fp8_gemm_kernel<Policy>>(
|
||||||
|
Policy::kSmemBytes, grid, dim3(Traits::kCtaThreads), stream, p);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Plan -> Policy: the production-tuned configs. Big CTA: 128x128 of 8 warps
|
||||||
|
// x 64x32, kK=64, 2-stage full ring, fast loop only for dual-congruous
|
||||||
|
// layouts. Narrow: 128x64. Small CTA: 64x64 of 4 warps x 32x32, kK=64,
|
||||||
|
// kFastLoop always on.
|
||||||
|
template <FP8Format Fmt, typename LayoutA, typename LayoutB, int GroupRaster>
|
||||||
|
void launch_plan(const FP8Params& p, const Fp8GemmPlan& plan,
|
||||||
|
cudaStream_t stream) {
|
||||||
|
constexpr bool kBigFast = !std::is_same_v<LayoutA, ColMajor> &&
|
||||||
|
!std::is_same_v<LayoutB, RowMajor>;
|
||||||
|
switch (plan.cta) {
|
||||||
|
case Fp8GemmPlan::Cta::kBig128: {
|
||||||
|
using Policy =
|
||||||
|
Fp8GemmPolicy<Fmt, 128, 128, LayoutA, LayoutB, 64, 32, 64, 2,
|
||||||
|
GroupRaster, false, kBigFast>;
|
||||||
|
launch_policy<Policy>(p, stream);
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
case Fp8GemmPlan::Cta::kNarrow128x64: {
|
||||||
|
using Policy =
|
||||||
|
Fp8GemmPolicy<Fmt, 128, 64, LayoutA, LayoutB, 32, 32, 64, 2,
|
||||||
|
GroupRaster, false, true>;
|
||||||
|
launch_policy<Policy>(p, stream);
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
case Fp8GemmPlan::Cta::kSmall64: {
|
||||||
|
if (plan.small_s3) {
|
||||||
|
using Policy = Fp8GemmPolicy<Fmt, 64, 64, LayoutA, LayoutB, 32, 32,
|
||||||
|
64, 3, GroupRaster, false, true>;
|
||||||
|
launch_policy<Policy>(p, stream);
|
||||||
|
} else {
|
||||||
|
using Policy = Fp8GemmPolicy<Fmt, 64, 64, LayoutA, LayoutB, 32, 32,
|
||||||
|
64, 2, GroupRaster, false, true>;
|
||||||
|
launch_policy<Policy>(p, stream);
|
||||||
|
}
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Pure problem rewrite: the dual-N-contiguous problem (trans_a/trans_b both
|
||||||
|
// false) has no dedicated instantiation — it runs as its transpose
|
||||||
|
// E[N][M] = B^T @ A^T (CUTLASS-sm90's is_swapAB) over swapped operands,
|
||||||
|
// with p.out_transposed making the epilogue scatter into the caller's
|
||||||
|
// [M][N] row-major buffer. The rewritten trans flags become the layout tags
|
||||||
|
// the launcher instantiates; the NN path pays a scalar-store scatter, which
|
||||||
|
// its rare usage makes the right trade.
|
||||||
|
inline void canonicalize_gemm(FP8Params& p, bool& trans_a, bool& trans_b) {
|
||||||
|
if (!trans_a && !trans_b) {
|
||||||
|
FP8Params s = p; // E = B^T * A^T: swap roles, M <-> N
|
||||||
|
s.m = p.n;
|
||||||
|
s.n = p.m;
|
||||||
|
s.a_ptr = p.b_ptr;
|
||||||
|
s.b_ptr = p.a_ptr;
|
||||||
|
s.a_ld = p.b_ld;
|
||||||
|
s.b_ld = p.a_ld;
|
||||||
|
s.a_batch_stride = p.b_batch_stride;
|
||||||
|
s.b_batch_stride = p.a_batch_stride;
|
||||||
|
s.out_transposed = 1;
|
||||||
|
p = s;
|
||||||
|
trans_a = trans_b = true;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Entry point: canonicalize the problem, plan the launch, wire the layout
|
||||||
|
// tags through.
|
||||||
|
template <FP8Format Fmt>
|
||||||
|
void gemm(FP8Params p, cudaStream_t stream, bool trans_a, bool trans_b) {
|
||||||
|
canonicalize_gemm(p, trans_a, trans_b);
|
||||||
|
// Crosswise operand count for the plan: transposed-A storage (ColMajor)
|
||||||
|
// and plain-B storage (RowMajor) both take the direct crosswise load.
|
||||||
|
const int crosswise = (trans_a ? 1 : 0) + (trans_b ? 0 : 1);
|
||||||
|
const Fp8GemmPlan plan = plan_gemm(p, crosswise);
|
||||||
|
if (trans_a && trans_b)
|
||||||
|
launch_plan<Fmt, ColMajor, ColMajor, 8>(p, plan, stream);
|
||||||
|
else if (trans_b)
|
||||||
|
launch_plan<Fmt, RowMajor, ColMajor, 8>(p, plan, stream);
|
||||||
|
else
|
||||||
|
launch_plan<Fmt, ColMajor, RowMajor, 8>(p, plan, stream);
|
||||||
|
}
|
||||||
|
|
||||||
|
} // namespace fp8
|
||||||
|
} // namespace astrai
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user